pusht-flashwam / training_code /PACKAGE_README.md
SleepMastger's picture
add model card, conditioning, and training-time processing
2622c40 verified
|
Raw History Blame Contribute Delete
8.63 kB
# pusht_train β€” FastWAM / FlashWAM from scratch on real-robot Push-T
Two from-scratch training runs on `SleepMastger/pusht-manipulation`, 4 GPUs
each on `gpu-h200-103`. Port of `dish_utensil_train` with the dataset, configs
and staging swapped; the recipe is deliberately identical so pusht results are
comparable to the dish_utensil and fruit runs.
## Dataset
100 human teleop demonstrations of the real-hardware Push-T task on a Franka
Panda: push an orange T-block until it aligns with a pink T outline, using a
marker pen clamped in the gripper as a single-point pusher.
| | |
|---|---|
| source | `SleepMastger/pusht-manipulation` (HF, public, 6.3 GB, 100 HDF5) |
| raw snapshot | `raw/pusht_manipulation/<session>/episode_<n>.hdf5` |
| converter shim | `raw/pusht_layout/data/<session>` β†’ symlinks to the 4 session dirs |
| converted | `datasets/pusht_lerobot_v21` |
| episodes / frames | 100 / 32,131 |
| rate | **10 Hz** (`--fps 10`; note the converter's default is 20) |
| sessions | `0826_1637` (50), `0827_1114` (2), `0827_1118` (33), `0827_1144` (15) |
| episode length | 100–767 frames, median 292 |
Sessions sort chronologically by name, so convert order is
`0826_1637/episode_2` … `0827_1144/episode_14`.
### Task string β€” deploy-critical
```
push the T block to the target outline
```
No trailing period. The T5 cache filename is the sha256 of the wrapped prompt:
```
71cb088954da46da4e3cb6c6f73ac812690e9bccf886c7858377c61918b4e3e8.t5_len128.wan22ti2v5b.pt
```
Deployment must byte-match this string or the dataloader/policy sees a
different embedding. Both sbatch scripts assert this file exists before
training starts.
## Preprocessing
Standard pipeline, `lift2lerobot/convert_lift_hdf5_to_lerobot_v21.py`:
```bash
python convert_lift_hdf5_to_lerobot_v21.py \
--raw-dir .../raw/pusht_layout \
--out-root .../datasets/pusht_lerobot_v21 \
--fps 10 \
--task "push the T block to the target outline" \
--repo-id george/pusht
```
- `observation.state` (8) = `eef_pos`(3) + `quat2axisangle(eef_quat)`(3) +
`[width/2, -width/2]`(2)
- `action` (7) = raw `[dx,dy,dz,drx,dry,drz,gripper]` with the gripper remapped
`g β†’ (1-g)/2` (robosuite `{-1 open, +1 close}` β†’ LIBERO/RLDS `{1 open, 0 close}`)
- images: both cameras 256Γ—256 uint8 β†’ per-frame JPEG q95 β†’ AV1 video
### The degenerate channels (read this before touching normalization)
The dataset README warns that a min/max normalizer will divide by zero on this
data. **It does not here**, but the reason is worth knowing.
The teleop rig commanded translation only and the pen stayed clamped for the
whole task, so 6 of the 15 numeric channels are constant. Measured across all
100 episodes / 32,131 frames before conversion:
| field | dim | range | |
|---|---|---|---|
| action | dx, dy, dz | 3.5e-02 … 3.6e-02 | live |
| action | drx, dry, drz | **0.0** | constant |
| action | gripper | **0.0** (raw +1 β†’ 0.0) | constant |
| state | x, y, z | 1.7e-01 … 4.4e-01 | live |
| state | rx, ry, rz | 8.9e-02 … 1.6e-01 | live |
| state | gripL, gripR | **3.2e-05** | constant |
`SingleFieldLinearNormalizer` (`normalizer.py:96-118`) guards this:
```python
input_range = input_max - input_min
ignore_dim = input_range < self.range_tol # range_tol = 1e-4
input_range[ignore_dim] = self.output_max - self.output_min
```
All six constant channels fall under `range_tol=1e-4` β€” the tightest of them,
the gripper state dims at 3.2e-05, clears it by about 3x β€” so each gets
`scale = 1.0` and `offset = -min` instead of a division by ~0. Nothing NaNs.
Note what that offset actually does: an ignored dim normalizes to `x - min`,
which is identically 0 only when the raw channel is *exactly* constant. So the
four action dims (raw range exactly 0.0) do come out at 0, while the two
gripper state dims come out spanning `[0, 3.2e-05]`. Both are finite and
bounded by the raw range, which is the guarantee that matters β€” negligible
beside the Β±1 live channels. `preflight.py` asserts this measured through the
real dataloader rather than reasoned from the source.
Consequence: keeping the full 7-dim action / 8-dim proprio costs nothing and
keeps the architecture byte-identical to the dish_utensil and fruit runs, so
the results stay comparable. Both sbatch scripts re-assert the classification at launch (they fail if the
set of sub-`range_tol` dims is not exactly `action[3,4,5,6]` +
`state[6,7]`, or if any of them drifts to within 2x of `range_tol`).
If you ever lower `range_tol` below 3.2e-05, the gripper state dims go inf.
`delta_action_dim_mask` in the data config is *not* a "convert to delta"
switch β€” the raw actions are already per-step deltas in metres. The processor
only uses the mask to zero padded action steps
(`fastwam_processor.py:300-308`).
### Episode `0827_1144/episode_13`
The one edited episode in the set: 30 lead-in frames were trimmed upstream
(`attrs["trimmed_lead_in_frames"] = 30`), which brought its start pose from
292 mm off the median down to 26 mm. That is still ~3.7x the worst of the
other 99. **Kept** β€” 26 mm is small in absolute terms and dropping it would
cost 1% of the data. Drop it if you later need a strictly homogeneous
initial-state distribution.
## Training
| | |
|---|---|
| variants | `pusht_flashwam_scratch` (M1 FusedKV/RopeFixed), `pusht_fastwam_scratch` (full 30/30 MoT, 6.0B) |
| init | **from scratch** (`resume: null`) |
| GPUs | 4 per job, `gpu-h200-103` |
| batch | 8 per GPU Γ— 4 Γ— accum 1 = **global 32** |
| schedule | cosine, lr 1e-4, wd 1e-2, 30 epochs |
| steps | 1,005 / epoch β†’ **30,150** total |
| checkpoints | every 5 epochs β†’ 6 per run |
| ports | 29566 (flashwam), 29567 (fastwam) |
| output | `runs/pusht_{flash,fast}wam_scratch/<timestamp>/` |
| wandb | `huaweiwam` / `fastwam-realrobot` |
```bash
bash precompute_pusht_text_embeds.sh # once, CPU, ~minutes
sbatch train_pusht_flashwam_scratch_n103.sbatch
sbatch train_pusht_fastwam_scratch_n103.sbatch
```
Submitted **independently β€” no `--dependency`**. Slurm decides whether they
overlap.
## GPU / throughput notes
Global batch 32 was kept rather than raised: the dish runs measured 84–100%
GPU utilisation at this batch size, i.e. already compute-bound, so a larger
batch buys little and would break comparability with the other real-robot
runs. The throughput work is elsewhere:
- **Node-local staging.** `stage_pusht_local.sh` copies the dataset to `/tmp`
(node NVMe) rather than reading video off beegfs. An atomic `mkdir` lock
means that when both jobs land on 103 only the first copies and the second
reuses it β€” one copy per node regardless of job count. `/tmp`, not
`/dev/shm`: the latter is wiped by the 853 job and would eat the host-RAM
budget earlyoom watches.
- **glibc malloc tuning** (in `run_one.sh`): `MALLOC_MMAP_THRESHOLD_=64MB`,
`MALLOC_ARENA_MAX=2`, `MALLOC_TRIM_THRESHOLD_=128MB`. The 64 MB threshold
keeps 16 MB decode buffers pooled (avoiding the mmap churn a low threshold
causes) while returning everything larger. The old blanket 1 GB threshold
let each worker hoard ~10 GB of freed buffers, ~1 TB per 4-GPU job, and two
concurrent 4-GPU jobs then overflowed node 103's 2 TB and got SIGTERMed by
`earlyoom` (dish jobs 1676/1677, 2026-08-13).
- **12 dataloader workers** Γ— 4 ranks = 48 of the job's 64 CPUs.
### Known hazards on node 103
- `earlyoom` SIGTERMs the biggest-RSS process when MemAvailable drops below
~10% (~200 GB). Two concurrent 4-GPU jobs are near that line; the malloc
tuning is what keeps them under it. A kill shows up as Slurm
`FAILED/NonZeroExitCode`, not as a manager kill.
- 103's gres accounting has **double-booked GPUs** (same failure mode as
node 102). Treat "free GPUs" on 103 as unreliable while other users have
jobs pending.
**Failure policy:** killed by the cluster manager β†’ do **not** resubmit.
Own-error failure β†’ fix the root cause, resubmit **once**.
## Files
```
configs/data/pusht_2cam.yaml dataset + normalization
configs/task/pusht_flashwam_scratch.yaml FlashWAM recipe
configs/task/pusht_fastwam_scratch.yaml FastWAM recipe
configs/model/lift_flashwam_m1_fusedkv_ropefixed.yaml
configs/model/lift_fastwam_full.yaml
run_one.sh accelerate launcher + malloc tuning
stage_pusht_local.sh /tmp staging with shared lock
precompute_pusht_text_embeds.sh T5 embedding (CPU)
train_pusht_flashwam_scratch_n103.sbatch
train_pusht_fastwam_scratch_n103.sbatch
text_embeds_cache/ 71cb0889….pt
```