christopher-kapic's picture
Upload folder using huggingface_hub
fdc6474 verified
|
Raw
History Blame Contribute Delete
4.02 kB

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

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.