tinyvla / scripts /diag_obs_gap.py
AlexWortega's picture
Upload folder using huggingface_hub
1601a2f verified
Raw
History Blame Contribute Delete
5.61 kB
#!/usr/bin/env python
"""Diagnose the env-obs adapter: compare policy predictions from env-rendered
observations vs dataset observations at the SAME init state.
If pred(dataset obs) is close to GT but pred(env obs) differs, the observation
adapter (image orientation/cameras/state) is the remaining gap.
"""
from __future__ import annotations
import argparse
import numpy as np
import torch
from scipy.spatial.transform import Rotation
@torch.no_grad()
def main():
parser = argparse.ArgumentParser()
parser.add_argument("--checkpoint", default="outputs/libero_ft2/step_15000")
parser.add_argument("--suite", default="libero_spatial")
parser.add_argument("--tasks", type=int, default=3)
args = parser.parse_args()
from lerobot.datasets.lerobot_dataset import LeRobotDataset, LeRobotDatasetMetadata
from lerobot.envs.factory import make_env, make_env_config
from transformers import AutoTokenizer
from tinyvla.data.mixture import CanonicalSource
from tinyvla.modeling_tinyvla import TinyVLAPolicy
policy = TinyVLAPolicy.from_pretrained(args.checkpoint).cuda().eval()
cfg = policy.config
tok = AutoTokenizer.from_pretrained(cfg.lm_model_name)
meta = LeRobotDatasetMetadata("HuggingFaceVLA/libero")
ds = LeRobotDataset(
"HuggingFaceVLA/libero",
delta_timestamps={"action": [t / meta.fps for t in range(cfg.chunk_size)]},
video_backend="torchcodec",
)
src = CanonicalSource(ds, 2, cfg.image_size, cfg.max_state_dim, cfg.max_action_dim)
s_stats = meta.stats["observation.state"]
s_mean = torch.as_tensor(s_stats["mean"]).flatten().float()
s_std = torch.as_tensor(s_stats["std"]).flatten().float().clamp(min=1e-6)
env_cfg = make_env_config("libero", task=args.suite)
task_envs = make_env(env_cfg, n_envs=1)[args.suite]
env_by_task = {}
for tid, env in task_envs.items():
desc = env.get_attr("task_description")[0]
env_by_task[desc.strip().lower()] = (tid, env)
eps_meta = ds.meta.episodes
first_ep_by_task = {}
for ep in range(ds.num_episodes):
start = int(eps_meta["dataset_from_index"][ep])
task = ds[start]["task"].strip().lower()
if task in env_by_task and task not in first_ep_by_task:
first_ep_by_task[task] = ep
def tok_batch(task_text):
t = tok([task_text], padding=True, truncation=True,
max_length=cfg.tokenizer_max_length, return_tensors="pt")
return t["input_ids"].cuda(), t["attention_mask"].bool().cuda()
def env_to_batch(obs, task_text):
imgs = {}
for slot, key in (("cam0", "image"), ("cam1", "image2")):
x = torch.as_tensor(np.asarray(obs["pixels"][key]))[0].flip(0).flip(1)
x = x.permute(2, 0, 1).float() / 255.0
x = torch.nn.functional.interpolate(x[None], size=(cfg.image_size, cfg.image_size),
mode="bilinear", align_corners=False)[0]
imgs[slot] = x
rs = obs["robot_state"]
pos = np.asarray(rs["eef"]["pos"]).flatten()
quat = np.asarray(rs["eef"]["quat"]).flatten()
rotvec = Rotation.from_quat(quat).as_rotvec()
if rotvec[0] < 0:
th = np.linalg.norm(rotvec)
rotvec = rotvec * (th - 2 * np.pi) / th
grip = np.asarray(rs["gripper"]["qpos"]).flatten()
state = torch.tensor(np.concatenate([pos, rotvec, grip]), dtype=torch.float32)
state = (state - s_mean) / s_std
state = torch.nn.functional.pad(state, (0, cfg.max_state_dim - state.shape[-1]))
ids, mask = tok_batch(task_text)
return {
"observation.images.cam0": imgs["cam0"][None].cuda(),
"observation.images.cam1": imgs["cam1"][None].cuda(),
"observation.state": state[None].cuda(),
"observation.language.tokens": ids,
"observation.language.attention_mask": mask,
"embodiment_id": torch.tensor([2], device="cuda"),
}
def ds_to_batch(item):
ids, mask = tok_batch(item.pop("task"))
b = {k: v[None].cuda() for k, v in item.items() if torch.is_tensor(v)}
b["observation.language.tokens"] = ids
b["observation.language.attention_mask"] = mask
return b
for task, ep in list(first_ep_by_task.items())[: args.tasks]:
tid, env = env_by_task[task]
obs, _ = env.reset(seed=0)
start = int(eps_meta["dataset_from_index"][ep])
item = src[start]
gt = item["action"].clone()[None].cuda()
env_b = env_to_batch(obs, task)
ds_b = ds_to_batch(dict(item))
torch.manual_seed(0)
pred_env = policy.predict_action_chunk(env_b)
torch.manual_seed(0)
pred_ds = policy.predict_action_chunk(ds_b)
m = item["action_dim_mask"]
d_env_gt = ((pred_env[0, :, m].cpu() - gt[0, :, m].cpu()) ** 2).mean().sqrt().item()
d_ds_gt = ((pred_ds[0, :, m].cpu() - gt[0, :, m].cpu()) ** 2).mean().sqrt().item()
d_env_ds = ((pred_env[0, :, m] - pred_ds[0, :, m]) ** 2).mean().sqrt().item()
# also compare the raw inputs
img_gap = (env_b["observation.images.cam0"] - ds_b["observation.images.cam0"]).abs().mean().item()
st_gap = (env_b["observation.state"] - ds_b["observation.state"]).abs().max().item()
print(f"task {tid}: RMSE(pred_env, GT)={d_env_gt:.3f} | RMSE(pred_ds, GT)={d_ds_gt:.3f} "
f"| RMSE(pred_env, pred_ds)={d_env_ds:.3f} | img L1 gap={img_gap:.4f} | state max gap={st_gap:.2f}")
if __name__ == "__main__":
main()