leepanic's picture
document experiment_cfg + embodiment slot
62164a5 verified
|
Raw
History Blame Contribute Delete
4.75 kB
---
license: apache-2.0
base_model: nvidia/GR00T-N1.7-3B
library_name: gr00t
pipeline_tag: robotics
tags: [robotics, vla, gr00t, isaac-gr00t, decent-vla, cube-stacking, so-101, federated]
---
# GR00T N1.7 CubeStack β€” FedAvg 3-client (global)
Federated global model: FedAvg over 3 colour-pair clients, **50 rounds x 250 local steps**, batch 32. Every round aggregated all 3 clients (`updates=3`).
Trained on Isambard-AI (GH200) with [decent-vla](https://github.com/kevinDuan1/decent-vla).
**Scope: action-head-only** β€” the Cosmos-Reason2 backbone is frozen
(`tune_llm`/`tune_visual` false); the ~1.62B DiT action head plus the embodiment
projector are trained. That is the standard NVIDIA new-embodiment recipe: the SO-101
projector slot ships at random init, so `tune_projector` must stay on.
**Embodiment:** SO-101 (`NEW_EMBODIMENT`), state/action 6-D `single_arm`+`gripper`,
cameras `front`/`wrist`. The `DecentVLA/Cube_*` datasets are LeRobot v3.0, which the
GR00T loader cannot read, so they were converted to v2.1 and the video pre-scaled to
shortest edge 256 (GR00T resizes to exactly that internally; it cut decode ~3.9x).
**Partition:** the SAME colour-pair non-IID split as the pi0.5 and SmolVLA CubeStack
3-client studies β€” c0 harry Green/Orange, c1 zhekai Green/Blue, c2 kevin Orange/Blue,
each client blind to the third colour.
**Caveat worth knowing:** the two kevin prompts are worded differently from the other
four ("Put orange cube on the blue cube." vs "Put the orange cube on top of the blue
cube."). Kept verbatim so this arm trains on exactly the strings the pi0.5 and SmolVLA
arms saw.
Both checkpoints hold the SAME parameter set: **537 tensors, 1.6205 B params, all
`action_head.*`** (verified by reading the tensors) β€” i.e. the frozen-backbone scope.
Checkpoint wrapper differs by arm, because they came off different code paths β€” load with
decent-vla `--resume`, not `from_pretrained`:
* FL `server_round_0050.pt` β€” `{"state": {...}, "round": 49}` (server aggregate; `round` is 0-indexed).
* Centralized `round_0075.pt` β€” `{"state_dict", "include", "extra", "embodiment", "backend"}`
(the LocalRunner end-of-run save, which also records embodiment + backend metadata).
Note there is no `server_round_0075.pt`: periodic saves ran every 10 rounds, and 75 is
not a multiple of 10, so round 75 exists only as this end-of-run file.
## Running these weights (experiment_cfg/)
The checkpoint holds ONLY `action_head.*` (537 tensors, 1.6205 B). Three things GR00T
derives from the training dataset are not in it and not in the public N1.7 base repo, so
they are published here under `experiment_cfg/`:
| file | what it is |
|---|---|
| `embodiment_id.json` | the projector-slot mapping used in this run |
| `processor_config.json` | modality config (state/action index ranges, video keys) + image/state processor settings |
| `statistics.json` | state/action normalization statistics |
**Embodiment slot = 10.** `new_embodiment` is absent from the base repo's
`embodiment_id.json`; the value comes from `EMBODIMENT_TAG_TO_PROJECTOR_INDEX` in
`gr00t/model/gr00t_n1d7/processing_gr00t_n1d7.py`. Verified against the weights, not just
that table: diffing this checkpoint's 14 embodiment-indexed tensors against the base head
puts **100.0% of the drift in slot 10 and exactly 0.0 in all 31 other slots** (see
`experiment_cfg/slot_drift.json`). Slot 0 is untrained β€” falling back to it runs without
error and produces meaningless actions.
Modality (from `processor_config.json`): video keys `front`, `wrist`; state and action
`single_arm` (0-5) + `gripper` (5-6); action horizon 16; `shortest_image_edge` 256.
The training config for this run is included as `experiment_cfg/train_config.yaml`.
### Which statistics for the federated model?
Our GR00T adapter wires normalization **per client partition**, so each FL client trained
against the statistics of its own two repos β€” there is no single "global" normalizer that
training produced. All four are published:
* `statistics.json` β€” **pooled over all 6 repos** (identical to the centralized arm). Use
this one for a single-number evaluation of the global model: it is also what the pi0.5
and SmolVLA CubeStack FL models carry, so it keeps the three policies comparable.
* `statistics_client{0,1,2}.json` β€” what each client actually trained with
(c0 harry Green/Orange, c1 zhekai Green/Blue, c2 kevin Orange/Blue).
They differ substantially β€” `action.gripper.mean` is 4.92 / 8.05 / 19.78 per client vs
10.92 pooled β€” so this choice changes results. Note this is a real difference from the
pi0.5 and SmolVLA CubeStack arms, whose lerobot adapter pooled the statistics at training
time; disclose it in any cross-policy table.