File size: 4,017 Bytes
fdc6474 | 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 | # `aqlm_converge.py` β activation-aware AQLM convergence for cold experts
## What it does
The shipped checkpoints quantized cold experts with **init-grade** AQLM:
residual k-means in *weight space* (minimize βWβΕ΄βΒ²), which ignores that
some weight columns matter far more than others at inference time. This
tool re-optimizes those same codes/codebooks/scales against the objective
that actually matters β **per-expert output error on real routed traffic**:
```
minimize Ξ£_e β(W_e β Ε΄_e) Β· diag(βh_e)βΒ²_F
```
where `h_e = E[xΒ²]` is a per-expert **diagonal input Hessian** estimated
from calibration activations that were actually routed to expert `e`:
- for **w13** (gate/up): `x` = the layer's hidden-state inputs
- for **w2** (down): `x` = that expert's own `silu(gate(x))Β·up(x)`,
computed with **teacher** (original NVFP4-dequantized) weights
Intuition: input dimensions that carry large activations get their weights
approximated more carefully; near-dead dimensions absorb the quantization
error. This is the same idea as GPTQ/AQLM's calibration objective, with a
diagonal (rather than full) Hessian so every step stays a giant batched
GEMM that a B200 executes in seconds.
### Algorithm (per layer Γ projection, over the layer's cold experts)
1. **Warm start** from the shipped codes/codebook/scales (already a decent
weight-space solution β convergence needs few iterations).
2. Iterate (max 6, early-stop when improvement < 0.2%):
a. **Weighted coordinate-descent re-encode** β for every group of 8
weights pick the codebook entry minimizing the h-weighted error.
Cost per sweep = two GEMMs against the 65,536Γ8 codebook.
b. **Closed-form weighted codebook update** β each entry β the
h-weighted mean of its assigned groups (weighted k-means step);
dead entries respawn on random groups.
c. **Closed-form weighted scale refit** β per-expert per-out-channel
least squares under h.
3. Log the weighted relative error before/after (typical: w13 ~0.09 β
substantially lower on hot input dims; the *unweighted* MSE may barely
move β that's the point).
What this still is NOT: full AQLM beam search (beam=1 CD only) and no
PV-tuning/global finetune. It captures the largest quality lever
(activation-aware objective) at ~1/100th the cost.
## Prerequisites (all already produced on this box)
| input | path | produced by |
|---|---|---|
| calibration activations, ~24k routed tokens/layer | `/data/glm52-acts/acts_layer{N}.pt` | `capture_acts.py` (needs the `VLLM_ACT_CAPTURE_DIR` hook, vLLM patch β₯0002) |
| teacher weights (original per-expert NVFP4) | `/tmp/glm52-hot-dl2` regions + `/data/glm52-old-layerwise` (layers 3,4,5,8,74β77) | `make_hot_manifest2.py` + `range_download.py` |
| warm-start codes + cold-expert ids | `/data/glm52` (live two-tier checkpoint) | β |
## Running
```bash
cd /home/coder/git/glm52 && source .venv/bin/activate
# smoke test (2 layers, idle GPUs):
python tools/aqlm_converge.py --layers 40,76 --gpus 4,5
# full run (75 layers, all 8 GPUs, layer-parallel):
python tools/aqlm_converge.py > /data/glm52-aqlm-conv/run.log 2>&1 &
```
Resumable: finished layers (`/data/glm52-aqlm-conv/layer_N.pt`) are
skipped on rerun. Expect roughly 10β25 min/layer/GPU (dominated by the
fp32 error/dequant temporaries β each worker peaks ~170 GB GPU memory, so
one layer per GPU at a time).
## Output & downstream
`/data/glm52-aqlm-conv/layer_{N}.pt`:
`expert_ids [nC]`, `w13_codes [nC,1,4096,768] i16`, `w13_codebooks
[1,65536,8] f16`, `w13_scales [nC,4096] f16`, `w2c_*` likewise, plus
`*_err_before/after` (h-weighted rel-err).
To ship: rebuild the checkpoints' cold arrays from these parts (slice by
`expert_ids` per target's cold set β the 1M cold set is a superset of the
500k/250k ones), revalidate (coherence + needle + KV/memory), re-upload.
The checkpoint format, kernels, and serving recipe are unchanged β only
the bytes inside `w13_*` / `w2c_*` tensors improve.
|