avanturist's picture
Add files using upload-large-folder tool
03e20d4 verified
|
Raw
History Blame Contribute Delete
7.5 kB
---
license: mit
base_model: openvla/openvla-7b
tags:
- robotics
- vision-language-action
- openvla
- openvla-oft
- recurrent-memory
- pomdp
- mikasa-robo
- maniskill
pipeline_tag: robotics
---
# mu-VLA (OpenVLA-OFT + recurrent memory), MIKASA-Robo, 64 memory tokens, TBPTT K=8
mu-VLA extends [OpenVLA-OFT](https://github.com/moojink/openvla-oft) with a recurrent memory
module so that the policy can act in partially observable Markov decision processes (POMDPs),
where the information needed for the current action appeared earlier in the episode and is no
longer visible. Vanilla OpenVLA-OFT is a per-step policy and cannot solve such tasks.
Memory is implemented as a set of learnable memory tokens that are injected into the multimodal
prefix of the transformer sequence and carried across environment steps. Gradients flow through
a truncated backpropagation-through-time (TBPTT) window; the memory state is reset at episode
boundaries.
This checkpoint corresponds to experiment 4 in the mu-VLA experiment matrix: multitask training on
5 MIKASA-Robo-VLA environments, 64 memory tokens, TBPTT window K=8.
Sibling checkpoints:
- [mu-vla-openvla-oft-mikasa-robo-5-tasks-m64-k2-tbptt](https://huggingface.co/mu-vla/mu-vla-openvla-oft-mikasa-robo-5-tasks-m64-k2-tbptt), same setup with K=2
- [mu-vla-openvla-oft-libero-4-tasks-m64-k8-tbptt](https://huggingface.co/mu-vla/mu-vla-openvla-oft-libero-4-tasks-m64-k8-tbptt), same memory setup on LIBERO
## Memory configuration
| Setting | Value |
|---|---|
| `num_mem_tokens` | 64 |
| `memory_update` | `tbptt` |
| `tbptt_length` | 8 |
| `attention_mask_mode` | `custom` |
| `ema_alpha` | 0.1 (unused when `memory_update=tbptt`) |
The same values are shipped in `memory_meta.json` so that the evaluation scripts in the code
repository pick them up automatically.
## Training setup
Base model: `openvla/openvla-7b`, fine-tuned with the OpenVLA-OFT recipe (L1 regression action
head, proprioception projector, two input camera views, no FiLM, no diffusion).
| Setting | Value |
|---|---|
| Checkpoint step | 150000 |
| Dataset | `MIKASARoboVLAEpisodicDataset`, `mikasa_five` |
| Training environments | `ShellGamePush-VLA-v0`, `InterceptMedium-VLA-v0`, `RememberColor5-VLA-v0`, `TakeItBack-VLA-v0`, `RememberShapeAndColor3x3-VLA-v0` |
| GPUs | 8 x NVIDIA A100 80GB |
| Batch size | 4 per GPU |
| LoRA rank | 32 |
| Learning rate | 5e-4, cosine schedule, 2000 warmup steps, minimum ratio 0.1 |
| Image augmentation | on |
| Gradient checkpointing | on |
The episodic dataloader feeds `batch_size` independent streams, each replaying whole episodes in
order, and marks every step with `is_first` and `is_last` so that the memory state can be reset on
episode boundaries. Action normalisation statistics (`q01`/`q99`) are computed jointly over all five
environments and stored under the `mikasa_combined` key in `dataset_statistics.json`.
## Evaluation
Protocol: each environment is evaluated independently, 100 episodes, starting seed `4242424242`,
metric `success_once`. Uncertainty is the standard error over 100 Bernoulli trials.
Inference uses receding horizon: instead of executing the whole open-loop chunk of 8 predicted
actions, the model is queried every step and only the first action is executed. This matters for
POMDP tasks where the visual cue is visible only briefly. The cost is 8 times more forward passes.
Only the first five rows below are training environments. Everything else is zero-shot transfer to
environments the model never saw during fine-tuning.
| Environment | mu-VLA K=8 | OpenVLA-OFT baseline |
|---|---|---|
| `ShellGamePush-VLA-v0` (train) | 0.93 ± 0.03 | 0.33 |
| `InterceptMedium-VLA-v0` (train) | 0.44 ± 0.05 | 0.53 |
| `RememberColor5-VLA-v0` (train) | 0.40 ± 0.05 | 0.23 |
| `TakeItBack-VLA-v0` (train) | 0.99 ± 0.01 | 0.98 |
| `RememberShapeAndColor3x3-VLA-v0` (train) | 0.09 ± 0.03 | 0.11 |
| `ShellGameTouch-VLA-v0` | 0.00 ± 0.00 | 0.00 |
| `ShellGamePick-VLA-v0` | 0.01 ± 0.01 | 0.01 |
| `InterceptSlow-VLA-v0` | 0.06 ± 0.02 | 0.06 |
| `InterceptFast-VLA-v0` | 0.19 ± 0.04 | 0.24 |
| `InterceptGrabSlow-VLA-v0` | 0.00 ± 0.00 | 0.00 |
| `InterceptGrabMedium-VLA-v0` | 0.00 ± 0.00 | 0.00 |
| `InterceptGrabFast-VLA-v0` | 0.00 ± 0.00 | 0.00 |
| `RotateLenientPos-VLA-v0` | 0.00 ± 0.00 | 0.08 |
| `RotateLenientPosNeg-VLA-v0` | 0.06 ± 0.02 | 0.07 |
| `RotateStrictPos-VLA-v0` | 0.04 ± 0.02 | 0.03 |
| `RotateStrictPosNeg-VLA-v0` | 0.06 ± 0.02 | 0.05 |
| `RememberColor3-VLA-v0` | 0.38 ± 0.05 | 0.25 |
| `RememberColor9-VLA-v0` | 0.09 ± 0.03 | 0.07 |
| `RememberShape3-VLA-v0` | 0.35 ± 0.05 | 0.17 |
| `RememberShape5-VLA-v0` | 0.20 ± 0.04 | 0.15 |
| `RememberShape9-VLA-v0` | 0.09 ± 0.03 | 0.11 |
| `RememberShapeAndColor3x2-VLA-v0` | 0.11 ± 0.03 | 0.12 |
| `RememberShapeAndColor5x3-VLA-v0` | 0.15 ± 0.04 | 0.08 |
Baseline column: OpenVLA-OFT without memory, trained on the same five environments with the same
episodic dataloader (experiment 2 in the matrix). It is not a perfectly matched control, since it
used the constant-then-decay learning rate schedule rather than cosine.
Note on the TBPTT window: the K=2 sibling checkpoint is substantially stronger on the memory-heavy
environments (for example 0.94 versus 0.40 on `RememberColor5-VLA-v0`). If you are looking for the
best MIKASA-Robo memory policy of the two, start there. This checkpoint is published as the K=8
point of the ablation.
## Files
Loading this checkpoint requires more than the base transformer weights. The extra heads are stored
as separate PyTorch state dicts, matching the layout the training and evaluation scripts expect.
| File | Purpose |
|---|---|
| `model-0000*-of-00004.safetensors`, `model.safetensors.index.json` | merged vision-language backbone |
| `lora_adapter/` | LoRA adapter as saved by PEFT, before merging |
| `action_head--150000_checkpoint.pt` | L1 regression action head |
| `proprio_projector--150000_checkpoint.pt` | proprioception projector |
| `memory_module--150000_checkpoint.pt` | memory module, including `initial_memory` |
| `memory_meta.json` | memory hyperparameters for evaluation auto-detection |
| `dataset_statistics.json` | action normalisation statistics, key `mikasa_combined` |
| `optimizer--150000_checkpoint.pt`, `scheduler--150000_checkpoint.pt` | optimizer and scheduler state, for resuming training |
| `configuration_prismatic.py`, `modeling_prismatic.py`, `processing_prismatic.py` | custom model code, memory token injection and the custom attention mask live here |
## Usage
The checkpoint is not loadable with upstream `transformers`. It needs the mu-VLA code, which
depends on a forked `transformers` that supports the custom attention mask, and on
[MIKASA-Robo](https://github.com/CognitiveAISystems/MIKASA-Robo) for the environments.
```bash
huggingface-cli download mu-vla/mu-vla-openvla-oft-mikasa-robo-5-tasks-m64-k8-tbptt \
--local-dir ./mu-vla-mikasa-m64-k8
```
Then point the evaluation script at the downloaded directory; memory hyperparameters are read from
`memory_meta.json`. See the code repository, `CognitiveAISystems/mu-vla`, for setup and the exact
command line.
## Attribution
mu-VLA is a fork of OpenVLA-OFT by Moo Jin Kim, Chelsea Finn and Percy Liang
([arXiv:2502.19645](https://arxiv.org/abs/2502.19645)), which in turn builds on OpenVLA. The
recurrent memory module, the episodic dataloaders and the MIKASA-Robo integration are the
contribution of this work. Released under the MIT license, following upstream.