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:

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:

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 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