X-WAM-depth / README.md
rooty2020's picture
Add files using upload-large-folder tool
4fbae0d verified
|
Raw History Blame Contribute Delete
5.41 kB
metadata
license: apache-2.0
base_model: sharinka0715/X-WAM-checkpoints
pipeline_tag: robotics
tags:
  - robotics
  - vla
  - world-model
  - diffusion
  - manipulation
  - droid
  - x-wam
  - depth
  - joint-position-control

X-WAM-depth

X-WAM (Wan2.2-TI2V-5B world action model) fine-tuned on DROID with 8-D raw joint-position actions, three camera views and the depth branch enabled. Depth supervision comes from Depth Anything 3 inverse depth. The layout matches the official sharinka0715/X-WAM-checkpoints release, so the checkpoint drops into the upstream X-WAM code.

This is the depth counterpart of rooty2020/X-WAM-DROID, which was trained without a depth branch.

Model details

Base checkpoint X-WAM pretrained/ (40k steps, cross-embodiment), depth branch included
Backbone Wan2.2-TI2V-5B DiT, UMT5-XXL text encoder, Wan2.2 VAE (stride 4×16×16)
Depth branch yes (use_depth: true, num_extra_layers: 10 → extra_blocks / extra_heads)
Weights bf16, full runner state dict (DiT 6.67 B incl. depth branch + frozen T5 + VAE). X-WAM trains without EMA.
Training step 5,200 (of a 20,000-step schedule)
Views exterior_1_left, exterior_2_left, wrist_left at 192×320
Horizon 9 frames (frame skip 4 → 3.75 fps video), 4 actions per frame step → 32 actions at 15 Hz

Action space

DROID raw joint positions [joint_0 … joint_6, gripper] written into X-WAM's 14-D action slots (and 16-D proprio slots). See action_mapping.json.

  • action slots in the 14-D vector: [0, 1, 2, 3, 4, 5, 7, 6] (the gripper goes to slot 6)
  • normalization: y = clip(2·(x − q01)/(q99 − q01) − 1, −1, 1), then the gripper channel is negated. After normalization +1 = open, −1 = closed (the X-WAM convention; raw DROID is 0 = open, 1 = closed)
  • q01 / q99 are in both config.yaml and action_mapping.json

To decode predictions, undo these steps in reverse order.

Depth

  • Target: inverse depth, as in the X-WAM paper, with near = bright.
  • Source: Depth Anything 3 (DA3-LARGE-1.1). Each camera stream was processed in 48-frame chunks. Each chunk was one multi-view DA3 scene, with an 8-frame overlap and median-ratio scale chaining between chunks, which keeps flicker low.
  • Normalization: per (episode, camera), a robust 0.5 / 99.5-percentile affine map to [0, 1], stored as uint8. At train time each view is then min/max-normalized over the window to [−1, 1] (normalize_depths_per_view: true, since DA3 depth is relative).
  • Output: the model's depth predictions are therefore relative inverse depth per view and window, not metric depth.

Training

  • Data: DROID (OXE LeRobot cache), episodes with all three views and a depth cache, with augmentation on (crop 0.95, brightness/contrast/saturation 0.2). The depth caches were generated while training ran, so coverage grew over the run: about 2.6k episodes for the first ~1.2k steps, 19k and then 39k episodes up to ~3k steps, and all 50,441 usable DROID episodes (11.8 M windows) from ~3k steps to step 5,200.
  • Objective: X-WAM flow matching on video, depth, action and proprio (every loss weight 1.0), uniform timestep distribution with shift 5, joint distribution with a 50% clean-action ratio, text dropout 0.1.
  • Optimization: AdamW, LR 1e-5, weight decay 0.01, 200 warmup steps then cosine over 20,000 steps, grad clip 1.0, batch 56 (4 × GH200 × 14), FSDP with bf16-mixed precision.

Files

config.yaml                                          training config (+ action_num, normalization stats)
action_mapping.json                                  DROID 8-D <-> X-WAM 14-D mapping, normalization, depth notes
checkpoints/last.ckpt/checkpoint/mp_rank_00_model_states.pt   {"module": state_dict, "global_step", "epoch"}

config.yaml is the run's own config with one addition: action_num: 4. The training dataset sets this value at runtime, and the runner needs it to build the model.

Usage

hf download rooty2020/X-WAM-depth --local-dir checkpoints/droid_depth

Then point the upstream X-WAM scripts at it like any official checkpoint. For example, evaluation/policy_server.py loads checkpoints/last.ckpt/checkpoint/mp_rank_00_model_states.pt with a strict load_state_dict(ckpt["module"]). The key set is identical to the official pretrained release (1,555 tensors, depth branch included). Build the runner from this config.yaml so that the action/proprio normalization and action_num match.

Provenance

Converted from the Lightning FSDP sharded checkpoint epoch=0-step=5200.ckpt with xwam-droid/scripts/export_hf.py. The DiT tensors, depth branch included, are cast from fp32 to bf16 and keyed at runner level (model.*). The frozen text_encoder.* / vae.* tensors are copied from the official pretrained release, since they are never trained.

License

Apache 2.0, following X-WAM and Wan2.2. Training data comes from DROID, and its terms apply to the data. Depth labels were produced with Depth Anything 3, and its license applies to that model.