#!/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()