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