File size: 8,727 Bytes
9e3dfff
 
 
 
 
 
bb4ec89
 
9e3dfff
 
 
 
 
 
 
 
 
bb4ec89
9e3dfff
bb4ec89
 
 
9e3dfff
bb4ec89
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
9e3dfff
 
 
 
 
bb4ec89
 
 
 
 
9e3dfff
 
bb4ec89
 
 
 
9e3dfff
bb4ec89
 
9e3dfff
 
bb4ec89
 
 
 
 
 
9e3dfff
 
bb4ec89
 
 
9e3dfff
bb4ec89
 
 
 
 
 
9e3dfff
bb4ec89
9e3dfff
bb4ec89
 
9e3dfff
bb4ec89
 
 
9e3dfff
bb4ec89
 
 
 
 
 
9e3dfff
bb4ec89
 
 
9e3dfff
bb4ec89
9e3dfff
bb4ec89
 
 
 
 
 
 
 
 
 
 
 
9e3dfff
bb4ec89
9e3dfff
bb4ec89
 
 
 
 
 
 
 
 
 
 
 
 
 
9e3dfff
 
 
 
 
 
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
---
license: apache-2.0
library_name: pytorch
tags:
  - robotics
  - quantization
  - nvfp4
  - fp4
  - w4a4
  - svdquant
  - world-action-model
  - fastwam
base_model: armanakbari4/fastwam-ur3-3task-10k
datasets:
  - armanakbari4/ur3-3task-lerobot
---

# FastWAM UR3 3-task — 4-bit quantisations

Two post-training 4-bit quantisations of the UR3 3-task FastWAM fine-tune (step 7000). Both store
the weight **packed at 4 bits** and contract it on 4-bit or integer tensor cores. Neither is a
fake-quant checkpoint that keeps 4-bit values inside 16-bit tensors.

| | `ur3_step7000_asp_nvfp4.pt` | `ur3_step7000_svdquant_w4a4.pt` |
|---|---|---|
| method | **ĀFQ / ASP, deflated form** | SVDQuant (Li et al., ICLR 2025) |
| numeric format | **NVFP4** — E2M1 elements, block 16, E4M3 block scales | INT4, group 64 |
| weights / activations | W4A4, both at block 16 | W4A4, both at group 64 |
| tensor cores | **FP4** (`torch._scaled_mm_v2`, recipe `BlockWise1x16`) | INT8 (exact for 4-bit codes) |
| size | 4.6117 BPW, **3.38 GiB** | 4.5798 BPW, 3.36 GiB |
| quantised-vs-bf16 action NRMSE | **0.0006** | 0.0010 |
| target GPU | **RTX 5090 / Blackwell** (sm_100, sm_103, sm_120) | L40S / Ada, Hopper (sm_89 tuned) |

The NVFP4 checkpoint is the one to run on an RTX 5090.

Both quantise the same 600 Linears — the two experts' 2 × 30 blocks × {`self_attn` q/k/v/o,
`cross_attn` q/k/v/o, `ffn.0`, `ffn.2`}, **5.914 B parameters**. Patch/text/time embeddings, heads,
the action encoder, norms, modulation and the proprio encoder stay in bf16. Both are
**self-contained**: the quantised Linears *and* every unquantised `mot` tensor *and* the proprio
encoder are inside the file, so the 11.2 GiB bf16 checkpoint is not needed at inference.

Both were calibrated on `armanakbari4/ur3-3task-lerobot` with **10 episodes per task × 3 tasks,
every frame, seed 42 — 9 908 observations**.

## What ASP is, and what the deflated form means

Action-Subspace Protection keeps a rank-32 subspace of each action-expert layer out of the 4-bit
grid. The subspace is not chosen by activation magnitude: it is the top eigenspace of the **action
metric** `G = E_o[JᵀJ]`, `J = ∂action/∂x`, differentiated through all ten denoising steps, so it is
the set of directions the *emitted action* is most sensitive to rather than the ones that happen to
be large.

Per layer, with `s` the smoothing vector, `H` a block Hadamard, `V` the rank-32 basis and
`W̃ = W·diag(s)` the smoothed weight, the deflated contract is

```
x̃ = (x / s) H
y  = (x̃ V)(W̃V)ᵀ  +  NVFP4GEMM( (I − VVᵀ) x̃ ,  W̃(I − VVᵀ) )  +  bias
```

**Deflated** means the protected subspace is subtracted from the 4-bit path on *both* sides: the
low-rank branch carries the exact `W̃V` and the FP4 weight holds `W̃(I − VVᵀ)`. The cheaper
shared-weight variant, which quantises one weight and stores `V` alone, is a different arm and a
measurably worse one at these group sizes. The rank-32 branch stays in bf16 by design — it carries
the directions the action depends on, which is the point of protecting them.

The per-layer smoothing `(α, β)` is not a fixed 0.5. It comes from a 39-candidate grid search
(`α ∈ {0, 0.05…0.95}`, `β ∈ {0, 1−α}`, per SVDQuant's protocol) scored on **real calibration
activations** with the objective being this layer's own deflated-ASP NVFP4 output MSE — the same
objective the deployed arm minimises, at the same granularity it is applied.

## Files

| file | size | |
|---|---|---|
| `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` |
| `ur3_step7000_svdquant_w4a4.pt` | 3.36 GiB | INT4 SVDQuant checkpoint |
| `nvfp4.py` | 8 KB | the NVFP4 quantiser and `_scaled_mm_v2` wrapper — quantise, dequantise, scale swizzle, GEMM |
| `asp_nvfp4_runtime.py` | 9.6 KB | `ASPNVFP4Linear`, `install_asp_nvfp4`, `load_quantized_asp_model`. Self-contained; imports only torch and `nvfp4.py` |
| `ur3_3task_10k_dataset_stats.json` | 170 KB | proprio z-scoring and action denormalisation — **the checkpoint cannot be run correctly without it** |
| `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 |

Besides these you need the Wan2.2 VAE (`Wan2.2_VAE.pth`, 2.7 GiB) to encode the camera image to the
video latent. You do not need the bf16 checkpoint, the text encoder, or the Wan2.2 DiT weights.

## Running the NVFP4 checkpoint

Requires a GPU with **FP4 tensor cores** — RTX 5090 (sm_120), B200/B300 (sm_100/sm_103) — and a
PyTorch with `torch._scaled_mm_v2`. Verified on torch 2.12.0+cu130.

```python
from asp_nvfp4_runtime import load_quantized_asp_model

model, cfg = load_quantized_asp_model(
    "ur3_step7000_asp_nvfp4.pt",
    build_model,           # your own bf16 FastWAM constructor -> (model, cfg)
)
```

`build_model` supplies the module graph only; every tensor comes from the checkpoint. To quantise a
model you already built, call `install_asp_nvfp4(model, ckpt_path)` instead — it raises rather than
swapping a subset, because a partial swap is not a defined arm.

Full pipeline, kernels and the deployment guide: **https://github.com/arashakb/QuantWAM**
(`adapters/fastwam/`, `quantwam/kernels/nvfp4.py`, `adapters/fastwam/README_ur3_5090.md`). The
`UR3QuantPolicy` wrapper there takes **raw** `HxWx3` uint8 frames and the **raw** 14-d state and
applies the whole observation contract itself (top → 320×256, wrists → 160×128,
`[top ; [left|right]]` → 384×320, `*2/255 − 1`; state z-scored and clamped to ±5), so a caller
cannot get the preprocessing subtly wrong.

## Verification — NVFP4 deflated ASP

Everything below is measured on **held-out** episodes; the calibration episodes are excluded,
because a check run on fitted data cannot detect the failure it exists to detect.

**It is really 4-bit, not a simulation.** Read off the loaded model:

| | |
|---|---|
| `wq` dtype, all 600 layers | `torch.float4_e2m1fn_x2` |
| weight bytes resident | 2 820 MiB for 5.914 B weights = **4.00 bits/weight** (bf16 would be 16.00) |
| E4M3 block scales | 352.5 MiB |
| any dequantised weight copy anywhere | **none** |
| `torch._scaled_mm_v2` calls in one full inference | **3 300** = 300 video prefill × 1 + 300 action × 10 steps |
| recipes those calls used | `BlockWise1x16` only — i.e. NVFP4 |

**It computes the intended contract.** Recomputing each layer independently from the stored
transforms and comparing against the runtime: worst relative error **1.7e-3** across sampled action
and video layers, which is the bf16 output cast and not a form mismatch.

**The actions are right.** Twelve held-out observations spanning all three tasks:

| | NRMSE |
|---|---|
| **NVFP4 ASP vs bf16, action chunk** | **0.0006** |
| bf16 vs recorded actions | 0.0068 ← the ceiling |
| **NVFP4 ASP vs recorded actions** | **0.0066** (corr 0.9997) |

The quantised model sits as close to the recorded actions as the bf16 model does — marginally
closer on these frames, which is noise, not an improvement.

**Not yet measured: real-robot success rate.** The bf16 model has been evaluated on the physical
UR3; neither quantised checkpoint has. Everything above is open-loop agreement with recorded
trajectories, which is necessary but not sufficient.

## Verification — INT4 SVDQuant

| check | result |
|---|---|
| base bf16 model vs recorded actions | NRMSE 0.0055, corr 0.9998 |
| 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 |
| weight repacking fidelity at export | 0.0000 % of codes off by 1 LSB |
| quantised vs bf16, action chunk | NRMSE 0.0010 |
| quantised vs recorded actions | NRMSE 0.0064 — the same as bf16's own 0.0064 |
| end-to-end from raw camera frames | NRMSE 0.0061, corr 0.9997 |

Its Triton launch configurations are tuned per compute capability and only `sm_89` ships; on other
architectures the kernel prints `no tuned config for M=… K=… N=…` and falls back to an occupancy
heuristic — correct, but roughly 2× off the tuned optimum on narrow shapes. Run
`analysis/iw_gemm_tune.py` once on the target GPU to fix it. This does not apply to the NVFP4
checkpoint, which calls cuBLAS through `_scaled_mm_v2` rather than a hand-tuned Triton kernel.

## Tasks

* `blue_basket` — put the medicine then the measuring tape inside the blue basket
* `drawer` — open the drawer, put the white box inside the drawer then close the drawer
* `stacking_cubes` — put the green cube on top of the black cube and put the red cube on top of the green cube