File size: 7,290 Bytes
5a2e445
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1d6d985
5a2e445
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1d6d985
 
 
 
5a2e445
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1d6d985
5a2e445
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1d6d985
5a2e445
 
 
 
 
 
 
 
 
 
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
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
#!/usr/bin/env python
"""Build per-frame canonical EE-pose sidecars (offline, once).

For each frame stores the absolute end-effector pose implied by BOTH the
observation state and the action target, in the robot base frame:
  cols: episode_index, frame_index,
        s_px s_py s_pz s_qx s_qy s_qz s_qw   (state EE pose)
        a_px a_py a_pz a_qx a_qy a_qz a_qw   (action-target EE pose)
        grip                                  (raw gripper channel, for q1-q99 later)

Canonical chunk deltas are assembled at train time from these poses (anchor,
frequency, rotation convention are all decided there — this file is convention-free).

Source kinds:
  so101  : FK(joint_deg) for state and action (both are joint vectors).
  ee     : state/action already EE; reconstruct absolute pose. Bridge/LIBERO/DROID
           store EE pose in observation.state; action is a delta → target = state ⊕ action.
           (implemented per-source as adapters are verified.)

Usage:
  python scripts/build_ee_sidecar.py --kind so101 --glob '~/tinyvla_data/so101_v3/*'
"""

from __future__ import annotations

import argparse
from pathlib import Path

import numpy as np
import pyarrow as pa
import pyarrow.parquet as pq

OUT_DIR = Path.home() / "tinyvla_data" / "ee_sidecar"


def _read_cols(name, root):
    """State/action/ep/frame columns WITHOUT video decode (raw parquet)."""
    from lerobot.datasets.lerobot_dataset import LeRobotDataset

    ds = LeRobotDataset(name, root=root)
    hf = ds.reader.hf_dataset.with_format("numpy")
    return (
        np.asarray(hf["observation.state"], dtype=np.float64),
        np.asarray(hf["action"], dtype=np.float64),
        np.asarray(hf["episode_index"]).astype(int),
        np.asarray(hf["frame_index"]).astype(int),
    )


def build_so101(name, root, fk):
    from scipy.spatial.transform import Rotation

    state, action, ep, fr = _read_cols(name, root)
    n = len(state)
    cols = ["episode_index", "frame_index",
            "s_px", "s_py", "s_pz", "s_qx", "s_qy", "s_qz", "s_qw",
            "a_px", "a_py", "a_pz", "a_qx", "a_qy", "a_qz", "a_qw", "grip"]
    buf = {k: np.empty(n, dtype=np.float64) for k in cols}
    buf["episode_index"] = ep.astype(np.float64)
    buf["frame_index"] = fr.astype(np.float64)
    for i in range(n):
        Ts = fk.ee_pose(state[i])
        Ta = fk.ee_pose(action[i])
        sp, sq = Ts[:3, 3], Rotation.from_matrix(Ts[:3, :3]).as_quat()
        ap, aq = Ta[:3, 3], Rotation.from_matrix(Ta[:3, :3]).as_quat()
        for j, k in enumerate(("s_px", "s_py", "s_pz")): buf[k][i] = sp[j]
        for j, k in enumerate(("s_qx", "s_qy", "s_qz", "s_qw")): buf[k][i] = sq[j]
        for j, k in enumerate(("a_px", "a_py", "a_pz")): buf[k][i] = ap[j]
        for j, k in enumerate(("a_qx", "a_qy", "a_qz", "a_qw")): buf[k][i] = aq[j]
        buf["grip"][i] = action[i, 5]
    return pa.table(buf)


def build_ee_from_state(name, root, rot_format="auto"):
    """EE-native sources (Bridge, RT-1): state carries absolute EE pose.

    Bridge state = [x y z roll pitch yaw pad gripper]; RT-1 state =
    [x y z rx ry rz rw gripper] (quat). We store the STATE pose as canonical
    absolute pose (both s_* and a_* set to state pose; targets reconstructed at
    train time as state ⊕ action-delta if needed, but state-derived deltas are
    the uniform choice per plan → a_* == next-frame not needed here).
    """
    from lerobot.datasets.lerobot_dataset import LeRobotDataset
    from scipy.spatial.transform import Rotation

    ds = LeRobotDataset(name, root=root)
    snames = ds.meta.features["observation.state"].get("names", {})
    flat = snames.get("motors", snames) if isinstance(snames, dict) else snames
    state, _, ep, fr = _read_cols(name, root)
    n = len(state)
    # RoboCasa/PandaOmron: state = base_pos(3) base_quat(4) ee_pos(3) ee_quat(4) grip(2)
    # -> the EE block is not at the front, slice it out first
    if state.shape[1] == 16:
        state = np.concatenate([state[:, 7:14], state[:, 14:15]], axis=1)  # pos,quat,grip
    # detect quaternion: named "rw", or 8-dim state whose dims 3:7 are unit-norm
    is_quat = ("rw" in flat) or (
        state.shape[1] >= 8 and abs(np.linalg.norm(state[:200, 3:7], axis=1).mean() - 1.0) < 1e-2
    )
    cols = ["episode_index", "frame_index",
            "s_px", "s_py", "s_pz", "s_qx", "s_qy", "s_qz", "s_qw",
            "a_px", "a_py", "a_pz", "a_qx", "a_qy", "a_qz", "a_qw", "grip"]
    buf = {k: np.empty(n, dtype=np.float64) for k in cols}
    buf["episode_index"] = ep.astype(np.float64)
    buf["frame_index"] = fr.astype(np.float64)
    pos = state[:, :3]
    if rot_format == "rotvec":  # [x y z rvx rvy rvz grip grip] — LIBERO/robosuite
        quat = Rotation.from_rotvec(state[:, 3:6]).as_quat()
        grip = state[:, 6]
    elif is_quat:  # [x y z rx ry rz rw gripper]
        quat = state[:, 3:7]
        grip = state[:, 7]
    else:  # [x y z roll pitch yaw (pad) (gripper)]
        quat = Rotation.from_euler("xyz", state[:, 3:6]).as_quat()
        grip = state[:, 7] if state.shape[1] > 7 else np.zeros(n)  # free-body (UAV): no gripper
    for j, k in enumerate(("s_px", "s_py", "s_pz")): buf[k] = pos[:, j]
    for j, k in enumerate(("s_qx", "s_qy", "s_qz", "s_qw")): buf[k] = quat[:, j]
    for j, k in enumerate(("a_px", "a_py", "a_pz")): buf[k] = pos[:, j]
    for j, k in enumerate(("a_qx", "a_qy", "a_qz", "a_qw")): buf[k] = quat[:, j]
    buf["grip"] = grip
    return pa.table(buf)


def main():
    parser = argparse.ArgumentParser()
    parser.add_argument("--kind", choices=["so101", "ee"], required=True)
    parser.add_argument("--rot-format", choices=["auto", "rotvec"], default="auto")
    parser.add_argument("--glob", required=True)
    parser.add_argument("--audit", default=str(Path.home() / "tinyvla_data/so101_fk_audit.json"),
                        help="skip datasets not marked KEEP")
    args = parser.parse_args()

    import glob as _glob
    import json

    OUT_DIR.mkdir(parents=True, exist_ok=True)
    keep = None
    audit_path = Path(args.audit).expanduser()
    if args.kind == "so101" and audit_path.exists():
        audit = json.loads(audit_path.read_text())
        keep = {k for k, v in audit.items() if v.get("verdict") == "KEEP"}

    fk = None
    if args.kind == "so101":
        from tinyvla.data.kinematics_so101 import SO101FK
        fk = SO101FK()

    roots = sorted(Path(p) for p in _glob.glob(str(Path(args.glob).expanduser())))
    done = 0
    for r in roots:
        if not (r / "meta" / "info.json").exists():
            continue
        if keep is not None and r.name not in keep:
            print(f"skip {r.name} (not KEEP)")
            continue
        out = OUT_DIR / f"{r.name}.parquet"
        if out.exists():
            done += 1
            continue
        try:
            if args.kind == "so101":
                tbl = build_so101(r.name, r, fk)
            else:
                tbl = build_ee_from_state(r.name, r, args.rot_format)
            pq.write_table(tbl, out)
            done += 1
            print(f"[{done}] {r.name}: {tbl.num_rows} frames -> {out}")
        except Exception as e:
            print(f"FAIL {r.name}: {type(e).__name__}: {str(e)[:150]}")
    print(f"done {done}")


if __name__ == "__main__":
    main()