Files
VLA/scripts/draw_openloop_plot.py
T

87 lines
2.7 KiB
Python
Raw Normal View History

2025-09-07 14:59:17 +08:00
import os
import yaml
import torch
2025-09-11 13:18:33 +08:00
import matplotlib.pyplot as plt
2025-09-07 14:59:17 +08:00
from wall_x.model.qwen2_5_based.modeling_qwen2_5_vl_act import Qwen2_5_VLMoEForAction
from wall_x.data.load_lerobot_dataset import load_test_dataset, get_data_configs
model_path = "path/to/model"
action_tokenizer_path = "path/to/action_tokenizer"
save_dir = "path/to/plot"
2025-09-11 13:18:33 +08:00
model = Qwen2_5_VLMoEForAction.from_pretrained(
model_path, action_tokenizer_path=action_tokenizer_path
)
2025-09-07 14:59:17 +08:00
model.eval()
model = model.to("cuda")
model = model.bfloat16()
2025-09-11 13:18:33 +08:00
2025-09-07 14:59:17 +08:00
def load_config(config_path):
"""Load configuration from YAML file."""
with open(config_path, "r") as f:
config = yaml.load(f, Loader=yaml.FullLoader)
2025-09-11 13:18:33 +08:00
2025-09-07 14:59:17 +08:00
config["data"]["model_type"] = config.get("model_type")
2025-09-11 13:18:33 +08:00
2025-09-07 14:59:17 +08:00
return config
2025-09-11 13:18:33 +08:00
2025-09-07 14:59:17 +08:00
# get test dataloader
path = "path/to/config"
config = load_config(path)
dataload_config = get_data_configs(config["data"])
lerobot_config = dataload_config.get("lerobot_config", {})
dataset = load_test_dataset(config, lerobot_config, seed=42)
dataloader = dataset.get_dataloader()
total_frames = len(dataloader)
pred_horizon = 32
action_dim = 14
gt_traj = torch.zeros((total_frames, action_dim))
pred_traj = torch.zeros((total_frames, action_dim))
for idx, batch in enumerate(dataloader):
2025-09-11 13:18:33 +08:00
if idx % pred_horizon == 0 and idx + pred_horizon < total_frames:
2025-09-07 14:59:17 +08:00
batch = batch.to("cuda")
with torch.no_grad():
outputs = model(
**batch,
action_dim=action_dim,
pred_horizon=pred_horizon,
mode="predict",
2025-09-11 13:18:33 +08:00
predict_mode="fast",
2025-09-07 14:59:17 +08:00
)
2025-09-11 13:18:33 +08:00
pred_traj[idx : idx + pred_horizon] = outputs["predict_action"].detach().cpu()
2025-09-08 18:33:52 +08:00
# Denormalize ground truth actions
2025-09-11 13:18:33 +08:00
gt_action_chunk = batch["action_chunk"][:, :, :action_dim]
2025-09-08 18:33:52 +08:00
dof_mask = batch["dof_mask"].to(gt_action_chunk.dtype)
2025-09-11 13:18:33 +08:00
denormalized_gt = model.action_preprocessor.normalizer_action.unnormalize_data(
gt_action_chunk, ["x2_normal"], dof_mask
)
2025-09-08 18:33:52 +08:00
gt_traj[idx : idx + pred_horizon] = denormalized_gt.detach().cpu()
2025-09-11 13:18:33 +08:00
2025-09-07 14:59:17 +08:00
gt_traj_np = gt_traj.numpy()
pred_traj_np = pred_traj.numpy()
timesteps = gt_traj.shape[0]
fig, axs = plt.subplots(action_dim, 1, figsize=(15, 5 * action_dim), sharex=True)
2025-09-11 13:18:33 +08:00
fig.suptitle("Action Comparison for lerobot", fontsize=16)
2025-09-07 14:59:17 +08:00
for i in range(action_dim):
2025-09-11 13:18:33 +08:00
axs[i].plot(range(timesteps), gt_traj_np[:, i], label="Ground Truth")
axs[i].plot(range(timesteps), pred_traj_np[:, i], label="Prediction")
axs[i].set_ylabel(f"Action Dim {i+1}")
2025-09-07 14:59:17 +08:00
axs[i].legend()
axs[i].grid(True)
2025-09-11 13:18:33 +08:00
axs[-1].set_xlabel("Timestep")
2025-09-07 14:59:17 +08:00
plt.tight_layout(rect=[0, 0.03, 1, 0.95])
os.makedirs(save_dir, exist_ok=True)
2025-09-11 13:18:33 +08:00
plt.savefig(os.path.join(save_dir, "lerobot_comparison.png"))
2025-09-07 14:59:17 +08:00
plt.close()