FastWAM_UR3 / README.md
arashakb's picture
Add README.md
9e3dfff verified
|
Raw History Blame
5.12 kB
metadata
license: apache-2.0
library_name: pytorch
tags:
  - robotics
  - quantization
  - w4a4
  - svdquant
  - world-action-model
  - fastwam
base_model: armanakbari4/fastwam-ur3-3task-10k
datasets:
  - armanakbari4/ur3-3task-lerobot

FastWAM UR3 3-task β€” SVDQuant W4A4 (real packed INT4)

A 4-bit post-training quantisation of the UR3 3-task FastWAM fine-tune (step 7000). The weight is stored packed at 4 bits and contracted on integer tensor cores β€” this is not a fake-quant checkpoint that stores 4-bit values in 16-bit tensors.

base model armanakbari4/fastwam-ur3-3task-10k :: ur3_3task_10k_step7000.pt (11.2 GiB bf16)
method SVDQuant (Li et al., ICLR 2025) β€” SmoothQuant migration + rank-32 FP16 low-rank branch + per-group-64 INT4 residual
precision W4A4, weights per-group-64, activations per-group-64 dynamic
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
kept in bf16 patch/text/time embeddings, heads, action encoder, norms, modulation, proprio encoder
size 4.5798 BPW, 3.36 GiB (down from 11.2 GiB)
calibration armanakbari4/ur3-3task-lerobot β€” 10 episodes per task Γ— 3 tasks, every frame, seed 42 = 9 908 observations

Files

file size
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
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

Running it

Code, kernels and a deployment guide: https://github.com/arashakb/QuantWAM (adapters/fastwam/, quantwam/kernels/w4a4_triton.py, adapters/fastwam/README_ur3_5090.md).

from ur3_infer import UR3QuantPolicy
pol = UR3QuantPolicy()                                     # ~25 s
pol.set_task("drawer_lerobot")                             # or the full instruction string
chunk = pol.predict(top_rgb, left_rgb, right_rgb, state)   # [32, 14], robot units

predict takes raw HxWx3 uint8 RGB 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. Run python adapters/fastwam/ur3_infer.py --self-test first on any new machine.

Besides these files 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 11.2 GiB bf16 checkpoint, the 11 GB text encoder, or the 18.8 GB Wan2.2 DiT weights. Total deployment footprint β‰ˆ 6.1 GiB.

Verification

All 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.

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

Quantisation costs essentially nothing on this checkpoint: the quantised model sits as close to the recorded actions as the bf16 model does.

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

Hardware notes

The kernels execute the INT4 contract on INT8 tensor cores, which is exact β€” every 4-bit code is representable in int8 and the int32 accumulation rounds nothing. The 4 bits therefore buy DRAM traffic and memory rather than arithmetic.

Launch configurations are tuned per compute capability, and only sm_89 (L40S) ships. On other architectures β€” including sm_120 / RTX 5090 β€” 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 the narrow shapes. Run analysis/iw_gemm_tune.py once on the target GPU to fix it.

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