Add README.md (NVFP4 deflated ASP)
Browse files
README.md
CHANGED
|
@@ -4,6 +4,8 @@ library_name: pytorch
|
|
| 4 |
tags:
|
| 5 |
- robotics
|
| 6 |
- quantization
|
|
|
|
|
|
|
| 7 |
- w4a4
|
| 8 |
- svdquant
|
| 9 |
- world-action-model
|
|
@@ -13,82 +15,150 @@ datasets:
|
|
| 13 |
- armanakbari4/ur3-3task-lerobot
|
| 14 |
---
|
| 15 |
|
| 16 |
-
# FastWAM UR3 3-task —
|
| 17 |
|
| 18 |
-
|
| 19 |
-
|
| 20 |
-
checkpoint that
|
| 21 |
|
| 22 |
-
| | |
|
| 23 |
-
|---|---|
|
| 24 |
-
|
|
| 25 |
-
|
|
| 26 |
-
|
|
| 27 |
-
|
|
| 28 |
-
|
|
| 29 |
-
|
|
| 30 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 31 |
|
| 32 |
## Files
|
| 33 |
|
| 34 |
| file | size | |
|
| 35 |
|---|---|---|
|
| 36 |
-
| `
|
| 37 |
-
| `
|
|
|
|
|
|
|
|
|
|
| 38 |
| `ur3_prompt_embeddings.pt` | 3.0 MiB | the three task instructions, pre-encoded, so the 11 GB umT5-XXL text encoder is not needed at inference |
|
| 39 |
|
| 40 |
-
|
|
|
|
|
|
|
|
|
|
| 41 |
|
| 42 |
-
|
| 43 |
-
|
| 44 |
|
| 45 |
```python
|
| 46 |
-
from
|
| 47 |
-
|
| 48 |
-
|
| 49 |
-
|
|
|
|
|
|
|
| 50 |
```
|
| 51 |
|
| 52 |
-
`
|
| 53 |
-
|
| 54 |
-
|
| 55 |
-
wrong. Run `python adapters/fastwam/ur3_infer.py --self-test` first on any new machine.
|
| 56 |
|
| 57 |
-
|
| 58 |
-
|
| 59 |
-
|
|
|
|
|
|
|
|
|
|
| 60 |
|
| 61 |
-
## Verification
|
| 62 |
|
| 63 |
-
|
| 64 |
-
on fitted data cannot detect the failure it exists to detect.
|
| 65 |
|
| 66 |
-
|
|
|
|
|
|
|
| 67 |
|---|---|
|
| 68 |
-
|
|
| 69 |
-
|
|
| 70 |
-
|
|
| 71 |
-
|
|
| 72 |
-
|
|
| 73 |
-
|
|
| 74 |
|
| 75 |
-
|
| 76 |
-
|
|
|
|
| 77 |
|
| 78 |
-
**
|
| 79 |
-
UR3; this checkpoint has not. Everything above is open-loop agreement with recorded trajectories,
|
| 80 |
-
which is necessary but not sufficient.
|
| 81 |
|
| 82 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 83 |
|
| 84 |
-
|
| 85 |
-
representable in int8 and the int32 accumulation rounds nothing. The 4 bits therefore buy DRAM
|
| 86 |
-
traffic and memory rather than arithmetic.
|
| 87 |
|
| 88 |
-
|
| 89 |
-
|
| 90 |
-
|
| 91 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 92 |
|
| 93 |
## Tasks
|
| 94 |
|
|
|
|
| 4 |
tags:
|
| 5 |
- robotics
|
| 6 |
- quantization
|
| 7 |
+
- nvfp4
|
| 8 |
+
- fp4
|
| 9 |
- w4a4
|
| 10 |
- svdquant
|
| 11 |
- world-action-model
|
|
|
|
| 15 |
- armanakbari4/ur3-3task-lerobot
|
| 16 |
---
|
| 17 |
|
| 18 |
+
# FastWAM UR3 3-task — 4-bit quantisations
|
| 19 |
|
| 20 |
+
Two post-training 4-bit quantisations of the UR3 3-task FastWAM fine-tune (step 7000). Both store
|
| 21 |
+
the weight **packed at 4 bits** and contract it on 4-bit or integer tensor cores. Neither is a
|
| 22 |
+
fake-quant checkpoint that keeps 4-bit values inside 16-bit tensors.
|
| 23 |
|
| 24 |
+
| | `ur3_step7000_asp_nvfp4.pt` | `ur3_step7000_svdquant_w4a4.pt` |
|
| 25 |
+
|---|---|---|
|
| 26 |
+
| method | **ĀFQ / ASP, deflated form** | SVDQuant (Li et al., ICLR 2025) |
|
| 27 |
+
| numeric format | **NVFP4** — E2M1 elements, block 16, E4M3 block scales | INT4, group 64 |
|
| 28 |
+
| weights / activations | W4A4, both at block 16 | W4A4, both at group 64 |
|
| 29 |
+
| tensor cores | **FP4** (`torch._scaled_mm_v2`, recipe `BlockWise1x16`) | INT8 (exact for 4-bit codes) |
|
| 30 |
+
| size | 4.6117 BPW, **3.38 GiB** | 4.5798 BPW, 3.36 GiB |
|
| 31 |
+
| quantised-vs-bf16 action NRMSE | **0.0006** | 0.0010 |
|
| 32 |
+
| target GPU | **RTX 5090 / Blackwell** (sm_100, sm_103, sm_120) | L40S / Ada, Hopper (sm_89 tuned) |
|
| 33 |
+
|
| 34 |
+
The NVFP4 checkpoint is the one to run on an RTX 5090.
|
| 35 |
+
|
| 36 |
+
Both quantise the same 600 Linears — the two experts' 2 × 30 blocks × {`self_attn` q/k/v/o,
|
| 37 |
+
`cross_attn` q/k/v/o, `ffn.0`, `ffn.2`}, **5.914 B parameters**. Patch/text/time embeddings, heads,
|
| 38 |
+
the action encoder, norms, modulation and the proprio encoder stay in bf16. Both are
|
| 39 |
+
**self-contained**: the quantised Linears *and* every unquantised `mot` tensor *and* the proprio
|
| 40 |
+
encoder are inside the file, so the 11.2 GiB bf16 checkpoint is not needed at inference.
|
| 41 |
+
|
| 42 |
+
Both were calibrated on `armanakbari4/ur3-3task-lerobot` with **10 episodes per task × 3 tasks,
|
| 43 |
+
every frame, seed 42 — 9 908 observations**.
|
| 44 |
+
|
| 45 |
+
## What ASP is, and what the deflated form means
|
| 46 |
+
|
| 47 |
+
Action-Subspace Protection keeps a rank-32 subspace of each action-expert layer out of the 4-bit
|
| 48 |
+
grid. The subspace is not chosen by activation magnitude: it is the top eigenspace of the **action
|
| 49 |
+
metric** `G = E_o[JᵀJ]`, `J = ∂action/∂x`, differentiated through all ten denoising steps, so it is
|
| 50 |
+
the set of directions the *emitted action* is most sensitive to rather than the ones that happen to
|
| 51 |
+
be large.
|
| 52 |
+
|
| 53 |
+
Per layer, with `s` the smoothing vector, `H` a block Hadamard, `V` the rank-32 basis and
|
| 54 |
+
`W̃ = W·diag(s)` the smoothed weight, the deflated contract is
|
| 55 |
+
|
| 56 |
+
```
|
| 57 |
+
x̃ = (x / s) H
|
| 58 |
+
y = (x̃ V)(W̃V)ᵀ + NVFP4GEMM( (I − VVᵀ) x̃ , W̃(I − VVᵀ) ) + bias
|
| 59 |
+
```
|
| 60 |
+
|
| 61 |
+
**Deflated** means the protected subspace is subtracted from the 4-bit path on *both* sides: the
|
| 62 |
+
low-rank branch carries the exact `W̃V` and the FP4 weight holds `W̃(I − VVᵀ)`. The cheaper
|
| 63 |
+
shared-weight variant, which quantises one weight and stores `V` alone, is a different arm and a
|
| 64 |
+
measurably worse one at these group sizes. The rank-32 branch stays in bf16 by design — it carries
|
| 65 |
+
the directions the action depends on, which is the point of protecting them.
|
| 66 |
+
|
| 67 |
+
The per-layer smoothing `(α, β)` is not a fixed 0.5. It comes from a 39-candidate grid search
|
| 68 |
+
(`α ∈ {0, 0.05…0.95}`, `β ∈ {0, 1−α}`, per SVDQuant's protocol) scored on **real calibration
|
| 69 |
+
activations** with the objective being this layer's own deflated-ASP NVFP4 output MSE — the same
|
| 70 |
+
objective the deployed arm minimises, at the same granularity it is applied.
|
| 71 |
|
| 72 |
## Files
|
| 73 |
|
| 74 |
| file | size | |
|
| 75 |
|---|---|---|
|
| 76 |
+
| `ur3_step7000_asp_nvfp4.pt` | 3.38 GiB | **NVFP4 deflated-ASP checkpoint.** Weights E2M1 packed, E4M3 block scales pre-swizzled to `SWIZZLE_32_4_4` |
|
| 77 |
+
| `ur3_step7000_svdquant_w4a4.pt` | 3.36 GiB | INT4 SVDQuant checkpoint |
|
| 78 |
+
| `nvfp4.py` | 8 KB | the NVFP4 quantiser and `_scaled_mm_v2` wrapper — quantise, dequantise, scale swizzle, GEMM |
|
| 79 |
+
| `asp_nvfp4_runtime.py` | 9.6 KB | `ASPNVFP4Linear`, `install_asp_nvfp4`, `load_quantized_asp_model`. Self-contained; imports only torch and `nvfp4.py` |
|
| 80 |
+
| `ur3_3task_10k_dataset_stats.json` | 170 KB | proprio z-scoring and action denormalisation — **the checkpoint cannot be run correctly without it** |
|
| 81 |
| `ur3_prompt_embeddings.pt` | 3.0 MiB | the three task instructions, pre-encoded, so the 11 GB umT5-XXL text encoder is not needed at inference |
|
| 82 |
|
| 83 |
+
Besides these you need the Wan2.2 VAE (`Wan2.2_VAE.pth`, 2.7 GiB) to encode the camera image to the
|
| 84 |
+
video latent. You do not need the bf16 checkpoint, the text encoder, or the Wan2.2 DiT weights.
|
| 85 |
+
|
| 86 |
+
## Running the NVFP4 checkpoint
|
| 87 |
|
| 88 |
+
Requires a GPU with **FP4 tensor cores** — RTX 5090 (sm_120), B200/B300 (sm_100/sm_103) — and a
|
| 89 |
+
PyTorch with `torch._scaled_mm_v2`. Verified on torch 2.12.0+cu130.
|
| 90 |
|
| 91 |
```python
|
| 92 |
+
from asp_nvfp4_runtime import load_quantized_asp_model
|
| 93 |
+
|
| 94 |
+
model, cfg = load_quantized_asp_model(
|
| 95 |
+
"ur3_step7000_asp_nvfp4.pt",
|
| 96 |
+
build_model, # your own bf16 FastWAM constructor -> (model, cfg)
|
| 97 |
+
)
|
| 98 |
```
|
| 99 |
|
| 100 |
+
`build_model` supplies the module graph only; every tensor comes from the checkpoint. To quantise a
|
| 101 |
+
model you already built, call `install_asp_nvfp4(model, ckpt_path)` instead — it raises rather than
|
| 102 |
+
swapping a subset, because a partial swap is not a defined arm.
|
|
|
|
| 103 |
|
| 104 |
+
Full pipeline, kernels and the deployment guide: **https://github.com/arashakb/QuantWAM**
|
| 105 |
+
(`adapters/fastwam/`, `quantwam/kernels/nvfp4.py`, `adapters/fastwam/README_ur3_5090.md`). The
|
| 106 |
+
`UR3QuantPolicy` wrapper there takes **raw** `HxWx3` uint8 frames and the **raw** 14-d state and
|
| 107 |
+
applies the whole observation contract itself (top → 320×256, wrists → 160×128,
|
| 108 |
+
`[top ; [left|right]]` → 384×320, `*2/255 − 1`; state z-scored and clamped to ±5), so a caller
|
| 109 |
+
cannot get the preprocessing subtly wrong.
|
| 110 |
|
| 111 |
+
## Verification — NVFP4 deflated ASP
|
| 112 |
|
| 113 |
+
Everything below is measured on **held-out** episodes; the calibration episodes are excluded,
|
| 114 |
+
because a check run on fitted data cannot detect the failure it exists to detect.
|
| 115 |
|
| 116 |
+
**It is really 4-bit, not a simulation.** Read off the loaded model:
|
| 117 |
+
|
| 118 |
+
| | |
|
| 119 |
|---|---|
|
| 120 |
+
| `wq` dtype, all 600 layers | `torch.float4_e2m1fn_x2` |
|
| 121 |
+
| weight bytes resident | 2 820 MiB for 5.914 B weights = **4.00 bits/weight** (bf16 would be 16.00) |
|
| 122 |
+
| E4M3 block scales | 352.5 MiB |
|
| 123 |
+
| any dequantised weight copy anywhere | **none** |
|
| 124 |
+
| `torch._scaled_mm_v2` calls in one full inference | **3 300** = 300 video prefill × 1 + 300 action × 10 steps |
|
| 125 |
+
| recipes those calls used | `BlockWise1x16` only — i.e. NVFP4 |
|
| 126 |
|
| 127 |
+
**It computes the intended contract.** Recomputing each layer independently from the stored
|
| 128 |
+
transforms and comparing against the runtime: worst relative error **1.7e-3** across sampled action
|
| 129 |
+
and video layers, which is the bf16 output cast and not a form mismatch.
|
| 130 |
|
| 131 |
+
**The actions are right.** Twelve held-out observations spanning all three tasks:
|
|
|
|
|
|
|
| 132 |
|
| 133 |
+
| | NRMSE |
|
| 134 |
+
|---|---|
|
| 135 |
+
| **NVFP4 ASP vs bf16, action chunk** | **0.0006** |
|
| 136 |
+
| bf16 vs recorded actions | 0.0068 ← the ceiling |
|
| 137 |
+
| **NVFP4 ASP vs recorded actions** | **0.0066** (corr 0.9997) |
|
| 138 |
+
|
| 139 |
+
The quantised model sits as close to the recorded actions as the bf16 model does — marginally
|
| 140 |
+
closer on these frames, which is noise, not an improvement.
|
| 141 |
+
|
| 142 |
+
**Not yet measured: real-robot success rate.** The bf16 model has been evaluated on the physical
|
| 143 |
+
UR3; neither quantised checkpoint has. Everything above is open-loop agreement with recorded
|
| 144 |
+
trajectories, which is necessary but not sufficient.
|
| 145 |
|
| 146 |
+
## Verification — INT4 SVDQuant
|
|
|
|
|
|
|
| 147 |
|
| 148 |
+
| check | result |
|
| 149 |
+
|---|---|
|
| 150 |
+
| base bf16 model vs recorded actions | NRMSE 0.0055, corr 0.9998 |
|
| 151 |
+
| packed kernels vs the reference SVDQuant formula, per layer | activation scales bit-identical; ≤ 0.10 % of 4-bit codes differ, every one by exactly 1 LSB |
|
| 152 |
+
| weight repacking fidelity at export | 0.0000 % of codes off by 1 LSB |
|
| 153 |
+
| quantised vs bf16, action chunk | NRMSE 0.0010 |
|
| 154 |
+
| quantised vs recorded actions | NRMSE 0.0064 — the same as bf16's own 0.0064 |
|
| 155 |
+
| end-to-end from raw camera frames | NRMSE 0.0061, corr 0.9997 |
|
| 156 |
+
|
| 157 |
+
Its Triton launch configurations are tuned per compute capability and only `sm_89` ships; on other
|
| 158 |
+
architectures the kernel prints `no tuned config for M=… K=… N=…` and falls back to an occupancy
|
| 159 |
+
heuristic — correct, but roughly 2× off the tuned optimum on narrow shapes. Run
|
| 160 |
+
`analysis/iw_gemm_tune.py` once on the target GPU to fix it. This does not apply to the NVFP4
|
| 161 |
+
checkpoint, which calls cuBLAS through `_scaled_mm_v2` rather than a hand-tuned Triton kernel.
|
| 162 |
|
| 163 |
## Tasks
|
| 164 |
|