--- 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](https://github.com/sharinka0715/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](https://github.com/ByteDance-Seed/Depth-Anything-3) inverse depth. The layout matches the official [`sharinka0715/X-WAM-checkpoints`](https://huggingface.co/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`](https://huggingface.co/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 ```shell 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](https://droid-dataset.github.io/), and its terms apply to the data. Depth labels were produced with Depth Anything 3, and its license applies to that model.