arashakb commited on
Commit
bb4ec89
·
verified ·
1 Parent(s): f7cf333

Add README.md (NVFP4 deflated ASP)

Browse files
Files changed (1) hide show
  1. README.md +122 -52
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 — SVDQuant W4A4 (real packed INT4)
17
 
18
- A 4-bit post-training quantisation of the UR3 3-task FastWAM fine-tune (step 7000). The weight is
19
- stored **packed at 4 bits** and contracted on integer tensor cores — this is not a fake-quant
20
- checkpoint that stores 4-bit values in 16-bit tensors.
21
 
22
- | | |
23
- |---|---|
24
- | base model | [`armanakbari4/fastwam-ur3-3task-10k`](https://huggingface.co/armanakbari4/fastwam-ur3-3task-10k) :: `ur3_3task_10k_step7000.pt` (11.2 GiB bf16) |
25
- | method | SVDQuant (Li et al., ICLR 2025) — SmoothQuant migration + rank-32 FP16 low-rank branch + per-group-64 INT4 residual |
26
- | precision | **W4A4**, weights per-group-64, activations per-group-64 dynamic |
27
- | quantised | the 600 block Linears of both experts (2 × 30 blocks × {self_attn q/k/v/o, cross_attn q/k/v/o, ffn.0, ffn.2}), **5.914 B params** |
28
- | kept in bf16 | patch/text/time embeddings, heads, action encoder, norms, modulation, proprio encoder |
29
- | size | **4.5798 BPW**, **3.36 GiB** (down from 11.2 GiB) |
30
- | calibration | `armanakbari4/ur3-3task-lerobot` — **10 episodes per task × 3 tasks, every frame, seed 42 = 9 908 observations** |
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
31
 
32
  ## Files
33
 
34
  | file | size | |
35
  |---|---|---|
36
- | `ur3_step7000_svdquant_w4a4.pt` | 3.36 GiB | the checkpoint. **Self-contained**: the packed INT4 Linears *and* every unquantised `mot` tensor *and* the proprio encoder |
37
- | `ur3_3task_10k_dataset_stats.json` | 170 KB | proprio z-scoring and action denormalisation — the checkpoint cannot be run correctly without it |
 
 
 
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
- ## Running it
 
 
 
41
 
42
- Code, kernels and a deployment guide: **https://github.com/arashakb/QuantWAM**
43
- (`adapters/fastwam/`, `quantwam/kernels/w4a4_triton.py`, `adapters/fastwam/README_ur3_5090.md`).
44
 
45
  ```python
46
- from ur3_infer import UR3QuantPolicy
47
- pol = UR3QuantPolicy() # ~25 s
48
- pol.set_task("drawer_lerobot") # or the full instruction string
49
- chunk = pol.predict(top_rgb, left_rgb, right_rgb, state) # [32, 14], robot units
 
 
50
  ```
51
 
52
- `predict` takes **raw** `HxWx3` uint8 RGB frames and the **raw** 14-d state and applies the whole
53
- observation contract itself (top → 320×256, wrists → 160×128, `[top ; [left|right]]` → 384×320,
54
- `*2/255 - 1`; state z-scored and clamped to ±5), so a caller cannot get the preprocessing subtly
55
- wrong. Run `python adapters/fastwam/ur3_infer.py --self-test` first on any new machine.
56
 
57
- Besides these files you need the Wan2.2 VAE (`Wan2.2_VAE.pth`, 2.7 GiB) to encode the camera image
58
- to the video latent. You do **not** need the 11.2 GiB bf16 checkpoint, the 11 GB text encoder, or
59
- the 18.8 GB Wan2.2 DiT weights. Total deployment footprint ≈ **6.1 GiB**.
 
 
 
60
 
61
- ## Verification
62
 
63
- All measured on **held-out** episodes — the calibration episodes are excluded, because a check run
64
- on fitted data cannot detect the failure it exists to detect.
65
 
66
- | check | result |
 
 
67
  |---|---|
68
- | base bf16 model vs recorded actions | NRMSE **0.0055**, corr 0.9998 |
69
- | 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** |
70
- | weight repacking fidelity at export | **0.0000 %** of codes off by 1 LSB |
71
- | quantised vs bf16, action chunk | NRMSE **0.0010** |
72
- | **quantised vs recorded actions** | NRMSE **0.0064** — the same as bf16's own **0.0064** |
73
- | end-to-end from raw camera frames | NRMSE **0.0061**, corr 0.9997 |
74
 
75
- Quantisation costs essentially nothing on this checkpoint: the quantised model sits as close to the
76
- recorded actions as the bf16 model does.
 
77
 
78
- **Not yet measured: real-robot success rate.** The bf16 model has been evaluated on the physical
79
- UR3; this checkpoint has not. Everything above is open-loop agreement with recorded trajectories,
80
- which is necessary but not sufficient.
81
 
82
- ## Hardware notes
 
 
 
 
 
 
 
 
 
 
 
83
 
84
- The kernels execute the INT4 contract on **INT8 tensor cores**, which is exact — every 4-bit code is
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
- Launch configurations are tuned per compute capability, and only `sm_89` (L40S) ships. On other
89
- architectures — including **sm_120 / RTX 5090** — the kernel prints `no tuned config for M=… K=… N=…`
90
- and falls back to an occupancy heuristic: **correct, but roughly 2× off the tuned optimum** on the
91
- narrow shapes. Run `analysis/iw_gemm_tune.py` once on the target GPU to fix it.
 
 
 
 
 
 
 
 
 
 
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