| |
| """Rewrite fruit_pick's observation.state from its native 15-dim layout |
| (eef_xyz[3] + eef_quat_xyzw[4] + joints[7] + gripper_width[1]) down to the |
| SAME 8-dim convention every other dataset in this repo uses (eef_pos[3] + |
| axis-angle rotation[3] + [gripper_width/2, -gripper_width/2]) -- per user |
| request, this training run must match place_cube_in_bowl/lift_new exactly. |
| |
| Formula copied verbatim from |
| fastwam_train/lift2lerobot/convert_lift_hdf5_to_lerobot_v21.py::quat2axisangle |
| (itself from fastwam/experiments/libero/libero_utils.py / robosuite), so the |
| state convention is bit-identical to how place_cube_in_bowl and lift_new were |
| built -- NOT a from-scratch reimplementation. |
| |
| Usage (run once, before staging/training): |
| python fastwam_train/scripts/fix_fruit_pick_state_schema.py |
| |
| Idempotent: skips if observation.state is already 8-dim. Backs up the |
| original parquet (.orig_15dim_state) and updates meta/info.json's declared |
| feature shape to match. |
| """ |
| import json |
| import math |
| import shutil |
| from pathlib import Path |
|
|
| import numpy as np |
| import pandas as pd |
|
|
| REPO = Path("/shared_work/amin-a/realworld-wam") |
| DATASET_DIR = REPO / "data/fruit_pick" |
| DATA_PARQUET = DATASET_DIR / "data/chunk-000/file-000.parquet" |
| INFO_JSON = DATASET_DIR / "meta/info.json" |
|
|
|
|
| def quat2axisangle(quat: np.ndarray) -> np.ndarray: |
| """(x, y, z, w) vec4 -> (ax, ay, az) axis-angle exponential coordinates.""" |
| quat = quat.copy() |
| if quat[3] > 1.0: |
| quat[3] = 1.0 |
| elif quat[3] < -1.0: |
| quat[3] = -1.0 |
| den = np.sqrt(1.0 - quat[3] * quat[3]) |
| if math.isclose(den, 0.0): |
| return np.zeros(3) |
| return (quat[:3] * 2.0 * math.acos(quat[3])) / den |
|
|
|
|
| def main(): |
| df = pd.read_parquet(DATA_PARQUET) |
| sample_state = np.asarray(df["observation.state"].iloc[0]) |
| if sample_state.shape[-1] == 8: |
| print(f"already 8-dim ({DATA_PARQUET}); nothing to do.") |
| return |
| assert sample_state.shape[-1] == 15, f"unexpected state dim {sample_state.shape[-1]}" |
|
|
| backup = DATA_PARQUET.with_suffix(".parquet.orig_15dim_state") |
| if not backup.exists(): |
| shutil.copy2(DATA_PARQUET, backup) |
| print(f"backed up original -> {backup}") |
|
|
| new_states = [] |
| for raw in df["observation.state"]: |
| raw = np.asarray(raw, dtype=np.float64) |
| eef_pos = raw[0:3] |
| eef_quat_xyzw = raw[3:7] |
| gripper_width = raw[14] |
| ee_state = np.concatenate([eef_pos, quat2axisangle(eef_quat_xyzw)]) |
| gripper_state = np.array([gripper_width / 2.0, -gripper_width / 2.0]) |
| new_states.append(np.concatenate([ee_state, gripper_state]).astype(np.float32)) |
|
|
| df["observation.state"] = new_states |
| df.to_parquet(DATA_PARQUET) |
| print(f"rewrote {DATA_PARQUET.name}: observation.state 15-dim -> 8-dim ({len(df)} rows)") |
|
|
| info = json.loads(INFO_JSON.read_text()) |
| info["features"]["observation.state"] = { |
| "dtype": "float32", |
| "shape": [8], |
| "names": ["eef_x", "eef_y", "eef_z", "axis_x", "axis_y", "axis_z", |
| "gripper_pos", "gripper_neg"], |
| } |
| backup_info = INFO_JSON.with_suffix(".json.orig_15dim_state") |
| if not backup_info.exists(): |
| shutil.copy2(INFO_JSON, backup_info) |
| INFO_JSON.write_text(json.dumps(info, indent=4)) |
| print(f"updated {INFO_JSON} feature shape to [8]") |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|