File size: 10,444 Bytes
33766fa 1c5b1f4 33766fa 1c5b1f4 33766fa 1c5b1f4 33766fa | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 | ---
license: mit
library_name: jax
tags:
- robotics
- world-models
- jax
- flax
- isaac-sim
- manipulation
- simdist
---
# UR10e Linear Gripper β Jig / Bottom Enclosure β SimDist World Model
Latent world model pretrained in simulation with **Simulation Distillation (SimDist)**
([arXiv:2603.15759](https://arxiv.org/abs/2603.15759), RSS 2026; code
[CLeARoboticsLab/simdist](https://github.com/CLeARoboticsLab/simdist), MIT).
**Status: trained.** 24,797,184 simulated steps, 90,000 updates. Weights, optimizer state,
configs and logs are in this repo.
## What it is
A planning-oriented latent world model that predicts future latent states, rewards and values from raw
(non-privileged) observations, so that a sampling-based planner can rank candidate action sequences.
```
z_t = E(o_t) latent encoder, newest observation only
h_t = C(o_{t-H:t-1}, a_{t-H:t-1}) history encoder β proprio + actions only, no images
αΊ_{t+1:t+T} = f(z_t, a_{t:t+T-1}, h_t) causal transformer, whole horizon in one pass
rΜ_{t:t+T-1} = R(αΊ, a) sequence-to-sequence transformer head
vΜ_{t+1:t+T} = V(αΊ) sequence-to-sequence transformer head
Γ’_{t:t+H} = Ο(z_t, h_t) base policy, action chunks to warm-start planning
```
Keeping images out of the history encoder is the paper's *Minimal History Representation*; it is what
makes planning affordable. Predicting the horizon in a single forward pass, rather than unrolling, is
what makes sampling thousands of candidate trajectories tractable.
## Architecture
Paper Table II. Embedding dimension 64; all transformer MLPs hidden 256; dynamics 3 layers / 4 heads;
reward 1 / 1; value 1 / 1; base policy 4 layers / 8 heads. Horizons `H = T = 5`.
The encoder passes each of three camera views through an ImageNet-pretrained ResNet-18 (shared trunk by
default, 11.18 M parameters against 33.53 M for separate trunks) to 3Γ512, concatenates with the 20-d
proprioceptive observation, and projects to the 64-d latent. The torchvisionβFlax NNX weight conversion
is exact: max absolute deviation **1.9e-06**.
## Training objective
Four terms, weighted 1 / 1 / 1 / 4:
- **latent dynamics** β MSE against `stop_grad(E(o_{t+i+1}))`
- **reward** β MSE
- **value** β MSE against the expert critic's V
- **behaviour cloning** β MSE, masked by the *cumulative* expert flag so it stops contributing the
moment an environment leaves the expert
There is deliberately **no pixel-reconstruction loss**. The paper's ablation shows adding one drops
manipulation success from 0.90 to 0.32 β reconstruction pressures the latent to encode randomised
texture and lighting that are irrelevant to the task.
## Read this before using the numbers
**No success rate is reported here, and none of these numbers is comparable to the
paper's Table I.** Table I reports *task success*, which requires closed-loop rollouts
under the MPPI planner. The planner exists in the port (`simdist/control/mppi.py`, task
agnostic) but the manipulation closed-loop harness does not β only the locomotion one
(`scripts/simulate_go2.py`). Everything below is a **training-side proxy**: latent
dynamics error, head regression quality, and a value-ranking AUC.
## Results at step 90,000
| metric | value | reference |
| --- | --- | --- |
| `test/latent_dynamics` | **0.0278** | paper reports 0.076 / 0.019 |
| `train/latent_dynamics` | 0.0343 | |
| `eval/latent_rollout_error` h1βh5 | 0.0164, 0.0228, 0.0283, 0.0344, 0.0385 | mean 0.0281, no knee |
| `eval/value_pearson_r` | **0.9510** | data critic ceiling **+0.953** |
| `eval/value_r2` | 0.9040 | |
| `eval/reward_pearson_r` | 0.9300 | |
| `eval/reward_r2` | 0.8641 | |
| `eval/bc_action_mse` | 0.1188 | |
| `eval/value_auc_success` | 0.7344 | see caveat below |
| `eval/latent_variance` encoder / predicted | 0.3695 / 0.3482 | matched β no collapse, no runaway |
The value head is the informative one. It regresses the data-generating critic's output,
so that critic's own correlation (**r = +0.953** against bootstrapped return-to-go) is a
ceiling it cannot meaningfully exceed. At 0.951 it has essentially reached it.
**The AUC is optimistic.** `eval.episodes.holdout` is 0, so the train/test split is at the
chunk level and no episode is fully unseen β the AUC has a step-level leak.
`eval/value_auc_holdout_episodewise: 0` records that regime. It is also computed over only
32 episodes.
### Stability across training
Early rows were transcribed from live reads during the run; auto-park pulled only the log
tail, so the shipped `logs/train_metrics.txt` starts at step 77500.
| step | `latent_dynamics` | encoder variance | lr | value AUC |
| --- | --- | --- | --- | --- |
| 2500 | 0.424 | 0.309 | 5.0e-5 | 0.766 |
| 5000 | 0.080 | 0.243 | 1.0e-4 | 0.684 |
| 7500 | 0.084 | 0.245 | 1.5e-4 | 0.731 |
| 10000 | 0.076 | 0.336 | 2.0e-4 (peak) | 0.773 |
| 90000 | 0.034 | 0.369 | 1.0e-4 | 0.734 |
Encoder variance staying flat through and past peak LR is the load-bearing observation β
see the divergence section.
## The divergence, and the deviation it forced
**The first attempt at this run diverged and had to be thrown away.** `latent_dynamics`
went 9.37 β 8.62 β 3.0e20 over steps 2500/5000/7500 and settled near 1e28, while reward,
value and action losses stayed bounded at 0.5β2.7 the whole time. `latent_variance/encoder`
tracked it exactly: 0.97 β 2.88 β 3.1e13 β 1.6e25.
The cause is a degenerate direction in the latent objective. The dynamics target is
`stop_gradient(encode_latent(...))`, but it is produced by **the same online encoder** β an
encoder that grows also grows its own target. `stop_gradient` bounds the gradient path, not
the magnitude. Nothing in the loss penalises βzβ, and the encoder's `latent_mlp` ends in a
free linear map, so βzβ is free to run away. The optimizer was bare `optax.adam`, with
`global_norm` computed for logging and never applied, so nothing bounded the rate either.
Two fixes were probed at peak LR 2e-4 held for 4000 steps (warmup compressed to 1000):
| configuration | `latent_dynamics` | encoder variance | value AUC |
| --- | --- | --- | --- |
| gradient clipping only | 5.4 β **9.4e6** | 0.99 β **1.1e4** | 0.80 β 0.57 |
| clipping + latent LayerNorm | 0.63 β **0.063** | pinned **0.12β0.45** | ~0.70β0.80 |
Clipping alone only bounds how fast the run travels the degenerate direction. The LayerNorm
removes it. **This model therefore uses a LayerNorm on the encoder latent, which the paper
does not describe** β an explicit deviation, exposed as `model.encoder.latent_norm`
(default `false`, the paper's architecture) rather than hardcoded. Whether this is an
undocumented detail of the reference implementation or a mis-port has not been checked
against the reference code.
## Deviations from the paper
- **Latent LayerNorm** β added, as above. The substantive one.
- **Gradient clipping** at global norm 1.0. The paper logs `grad_norm` and does not apply it.
- **90,000 updates**, roughly 1.06 epochs, against the paper's ~194.5 k / 2 epochs. A budget
decision: per Table I, data scale dominates and epochs are secondary, so the full dataset
with fewer passes is the better trade under a cap.
- **H = T = 5**, not the paper's 25. Deliberate and documented in the port: this policy runs
at 10 Hz, so 5 steps is 0.5 s of history and prediction β the timescale an insertion
evolves on. 25 would be 2.5 s.
- No pixel-reconstruction loss and no image decoder, matching the paper (Table I: adding
reconstruction as an objective is 0.90 β 0.32).
## Training
```
python scripts/train_model.py \
model=manipulator_world_model system=omnireset_ur10e \
data.dataset_name=simdist_merged data.num_train_workers=32 data.num_test_workers=4 \
training.batch_size=256 training.max_steps=90000 training.warmup_steps=10000 \
training.decay_steps=90000 training.eval_interval=2500 \
training.grad_clip_norm=1.0 model.encoder.latent_norm=true \
checkpoint.enabled=True checkpoint.max_to_keep=5 \
run_name=simdist_ur10e_24m_v2
```
Adam, cosine 2e-4 β 1e-4 with 10 k warmup. Loss weights 1/1/1/4 (latent dynamics / reward /
value / behaviour cloning, the last masked by the cumulative expert flag). One RTX 5090,
2.26 updates/s, 11.9 h, GPU util 97 %, `dataloader_wait_frac` 0.015.
Contents: `checkpoint/` is an orbax checkpoint including optimizer state (resumable);
`config/` holds the model, system and training configs; `logs/` and `metrics.json` hold the
eval history that survived.
## Data
24,797,184 steps, 155,370 episodes, 843 shards merged from 116 independent generation runs
(37.2 h on one RTX 5090, 174 rows/s). Three RGB views at 160Γ120 (front, side, wrist), JPEG
q90, plus 20-d proprioception and 7-d relative Cartesian OSC actions. `expert_prob` 0.5,
sub-optimal actions drawn from a 37-checkpoint ladder, visual and physics domain
randomisation on.
**The dataset is not published.** It lived on rented storage that was released at the end of
the campaign; only the model and its logs were retained.
### Known data defects
Both are recorded because they affect anyone reproducing this, and neither caused the
divergence above (reward and value losses stayed bounded throughout):
- **Value outliers.** 0.116 % of rows have value < β20, p0.01 = **β94.97**, against mean 8.87
and std 4.00 β a β28Ο target after scaling. This matches an MDP-mismatch signature seen
earlier in the project (V^e β98 vs +8.95).
- **Action outliers.** `ee_delta_*` are expected to be order ~1, but per-dimension absolute
maxima reach **3788** while p99.9 is only 12β37, inflating `actions.std` to 7.2β13.7.
A further note for reproducers: `episode_ids` from the generator is a **per-environment**
counter that restarts at 0 in every process, so merging N generation runs puts N distinct
episodes under `(env 0, episode 0)`. They must be namespaced per source run before
processing, or the episode stitcher will refuse the merged set.
## Intended use
Sim-only evaluation: held-out prediction losses, latent rollout error against horizon, reward and value
calibration, and MPPI planning **in simulation**. Real-world deployment and dynamics finetuning
(SimDist stages 4a/4b) are out of scope for this release and have not been validated.
## Licence
MIT, following upstream SimDist.
|