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.