File size: 5,468 Bytes
23e6b1f
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
#!/usr/bin/env python
"""Evaluate the new wds embodiments against their own trivial baseline.

The ratio metric used for the canonical robots integrates EE deltas, which is
wrong for these packs, so each action space gets the trivial baseline that
actually applies to it:

  nav (recon/sacson/go-stanford/scand/tartandrive): actions are step deltas /
      local waypoints -> integrate the chunk and compare endpoints. Floor =
      predicting zero motion, exactly as for the canonical robots.

  unitree (G1/Z1): actions are ABSOLUTE joint positions, so integrating is
      meaningless. Floor = "hold still", i.e. predict the current joint state for
      the whole chunk. ratio < 1 means the model beats not moving.

Both are computed in the packs' normalized units, so numbers are comparable
across packs but not to millimetres.
"""
from __future__ import annotations

import sys
from pathlib import Path

import numpy as np
import torch
import yaml
from transformers import AutoTokenizer

from tinyvla.data.wds_shards import WdsShardSource
from tinyvla.modeling_tinyvla import TinyVLAPolicy
from tinyvla.modules.embodiment import MORPH_FIELDS

CKPT = sys.argv[1] if len(sys.argv) > 1 else "outputs/tv2_joint/final"
WDS = Path.home() / "tinyvla_data" / "wds"
_SC = {"arm_dof": 0.1, "reach_m": 2, "gripper_width_m": 10, "num_cameras": 1 / 3,
       "control_hz": 1 / 30, "joint_lo_mean": 1 / 3.1416, "joint_hi_mean": 1 / 3.1416,
       "workspace_x": 2, "workspace_y": 2, "workspace_z": 2, "payload_kg": 0.2}

# (label, path, morph_key, embodiment_id, kind)
TARGETS = [
    ("G1 Dex3 ToastedBread", "wds-unitree/G1_Dex3_ToastedBread_Dataset", "g1", 15, "joint"),
    ("G1 Dex1 DualArm", "wds-unitree/G1_Dex1_DiverseManip_DualArm_256x256", "g1", 15, "joint"),
    ("Z1 StackBox", "wds-unitree/Z1_StackBox_Dataset", "z1", 19, "joint"),
    ("recon (outdoor nav)", "wds-recon", "recon", 16, "nav"),
    ("sacson (indoor nav)", "wds-sacson", "sacson", 21, "nav"),
    ("go-stanford (nav)", "wds-go-stanford", "go_stanford", 20, "nav"),
    ("scand (social nav)", "wds-scand", "scand", 18, "nav"),
    ("tartandrive (ATV)", "wds-tartandrive", "tartandrive", 17, "nav"),
]


@torch.no_grad()
def main():
    pol = TinyVLAPolicy.from_pretrained(CKPT).cuda().eval()
    cfg = pol.config
    tok = AutoTokenizer.from_pretrained(cfg.lm_model_name)
    desc = yaml.safe_load(open("configs/morphology/descriptors.yaml"))
    print(f"=== {CKPT} on the new embodiments (ratio to trivial baseline) ===")
    rows = []
    for label, rel, mkey, emb, kind in TARGETS:
        try:
            morph = torch.tensor([desc[mkey].get(f, 0) * _SC.get(f, 1) for f in MORPH_FIELDS],
                                 dtype=torch.float32)
            src = WdsShardSource(WDS / rel, embodiment_id=emb, chunk=cfg.chunk_size,
                                 image_size=cfg.image_size, max_state_dim=cfg.max_state_dim,
                                 max_action_dim=cfg.max_action_dim, morphology=morph)
            n = len(src)
            # held-out tail: the packs are written in episode order, so the last
            # slice is unseen episodes rather than unseen frames of seen ones
            idxs = range(int(n * 0.98), n, max(1, int(n * 0.02) // 120))
            errs, floors = [], []
            for i in idxs:
                it = src[i]
                t = tok([it["task"]], padding=True, truncation=True, max_length=48,
                        return_tensors="pt")
                b = {"observation.images.cam0": it["observation.images.cam0"][None].cuda(),
                     "observation.images.cam1": it["observation.images.cam1"][None].cuda(),
                     "observation.state": it["observation.state"][None].cuda(),
                     "observation.language.tokens": t["input_ids"].cuda(),
                     "observation.language.attention_mask": t["attention_mask"].bool().cuda(),
                     "morphology": morph[None].cuda(),
                     "embodiment_id": torch.tensor([emb]).cuda()}
                with torch.autocast("cuda", torch.bfloat16):
                    pr = pol.predict_action_chunk(b)[0].cpu().float().numpy()
                d = int(it["action_dim_mask"].sum())
                gt = it["action"][:, :d].numpy()
                pd = pr[:, :d]
                if kind == "nav":
                    gp, pp = np.cumsum(gt, 0)[-1], np.cumsum(pd, 0)[-1]
                    errs.append(np.linalg.norm(pp - gp))
                    floors.append(np.linalg.norm(gp))          # predict zero motion
                else:
                    st = it["observation.state"][:d].numpy()   # hold-still baseline
                    errs.append(np.linalg.norm(pd - gt, axis=1).mean())
                    floors.append(np.linalg.norm(st[None] - gt, axis=1).mean())
            err, floor = float(np.mean(errs)), float(np.mean(floors))
            rows.append((label, kind, err / floor))
            print(f"{label:26} [{kind:5}] n={len(errs):>4}  err {err:7.3f}  floor {floor:7.3f}  "
                  f"ratio {err/floor:.2f}", flush=True)
        except Exception as ex:
            print(f"{label:26} FAILED {type(ex).__name__}: {str(ex)[:80]}", flush=True)
    for kind in ("joint", "nav"):
        sel = [r[2] for r in rows if r[1] == kind]
        if sel:
            print(f"mean ratio, {kind:5}: {np.mean(sel):.2f}")
    print("ratio < 1.0 beats the trivial baseline (zero motion / hold still)")


if __name__ == "__main__":
    main()