diff --git a/c/tools/quant_ablation.py b/c/tools/quant_ablation.py index 7a870ff..689643a 100644 --- a/c/tools/quant_ablation.py +++ b/c/tools/quant_ablation.py @@ -106,33 +106,106 @@ def rotation(dim, device, seed=417): return q -def quantize_param(w, bits, group, rot=False): +def quantize_param(w, bits, group, rot=False, e8=""): if w.ndim == 3: # fused experts [E, in, out] -> move input last x = w.transpose(1, 2).contiguous() - x = _rot_quant(x, bits, group) if rot else _quant_last_dim(x, bits, group) + x = _rot_quant(x, bits, group, e8) if rot else _grid_or_e8(x, bits, group, e8) return x.transpose(1, 2).contiguous() if rot: - return _rot_quant(w, bits, group) - return _quant_last_dim(w, bits, group) # nn.Linear [out, in] -- input already last + return _rot_quant(w, bits, group, e8) + return _grid_or_e8(w, bits, group, e8) # nn.Linear [out, in] -- input already last -def _rot_quant(x, bits, group): +def _grid_or_e8(x, bits, group, e8): + if e8: + return _quant_e8(x.float(), group, ball=(e8 == "-e8")) + return _quant_last_dim(x, bits, group) + + +def _rot_quant(x, bits, group, e8=""): """W -> Qn(W@Q) @ Q^T along the last (input) dim — see rotation() above.""" q = rotation(x.shape[-1], x.device) - return (_quant_last_dim(x.float() @ q, bits, group) @ q.T).contiguous() + return (_grid_or_e8(x.float() @ q, bits, group, e8) @ q.T).contiguous() -SCHEME_RE = re.compile(r"^int(2|3|4|8)(?:-g(\d+))?(-rot)?(-nohead)?$") +# -------------------------------------------------------------------------------------- +# E8 lattice quantization (#81 follow-up): the -rot schemes above are QuaRot (rotation + +# uniform grid). QuIP#'s 2-bit result needs the second ingredient — an E8 lattice codebook +# instead of the grid. E8 = D8 ∪ (D8 + 1/2), nearest point via Conway–Sloane: round every +# coordinate, and if the sum is odd re-round the worst coordinate the other way; repeat on +# the half-shifted copy and keep the closer of the two. `-e8` clamps points to |p|^2 <= 10, +# the E8P ball QuIP# builds its 2^16 codebook from (2 bits/weight for 8-dim blocks); +# `-e8u` leaves the lattice unbounded — an ideal-codebook upper bound, not a deployable rate. +# Scale: per group, a small MSE search over multiples of the block RMS (absmax is the wrong +# statistic for a lattice — the ball wants energy matched, not the peak). +# -------------------------------------------------------------------------------------- +def _d8_nearest(y): + f = torch.round(y) + d = y - f + odd = (f.sum(-1).long() & 1).bool() + idx = d.abs().argmax(-1, keepdim=True) + step = torch.where(d.gather(-1, idx) >= 0, 1.0, -1.0) + flipped = f.gather(-1, idx) + step + return f.scatter(-1, idx, torch.where(odd[..., None], flipped, f.gather(-1, idx))) + + +def _e8_nearest(y): + a = _d8_nearest(y) + b = _d8_nearest(y - 0.5) + 0.5 + da = ((y - a) ** 2).sum(-1, keepdim=True) + db = ((y - b) ** 2).sum(-1, keepdim=True) + return torch.where(da <= db, a, b) + + +def _e8_ball(y, r2=10.0): + p = _e8_nearest(y) + for _ in range(8): # shrink-and-requantize until inside + n2 = (p ** 2).sum(-1, keepdim=True) + over = n2 > r2 + 1e-6 + if not over.any(): + break + y = torch.where(over, y * torch.sqrt(r2 / torch.clamp(n2, min=r2)) * 0.98, y) + p = torch.where(over, _e8_nearest(y), p) + return p + + +def _quant_e8(x, group, ball): + """Blocks of 8 along the input dim; per-group scale by MSE search over RMS multiples.""" + if x.shape[-1] % 8: + raise SystemExit(f"-e8 needs input dim divisible by 8 (got {x.shape[-1]})") + g = group or x.shape[-1] + if g % 8: + raise SystemExit(f"-e8 group {g} must be a multiple of 8") + shp = x.shape + xg = x.reshape(-1, g) # [G, g] + rms = torch.clamp(xg.pow(2).mean(-1, keepdim=True).sqrt(), min=1e-8) + best_out, best_err = None, None + for k in (0.5, 0.7, 0.9, 1.1, 1.4, 1.8, 2.4): + s = rms * k + yb = (xg / s).reshape(-1, g // 8, 8) + p = _e8_ball(yb) if ball else _e8_nearest(yb) + out = (p.reshape(-1, g) * s) + err = (out - xg).pow(2).sum(-1, keepdim=True) + if best_err is None: + best_out, best_err = out, err + else: + take = err < best_err + best_out = torch.where(take, out, best_out) + best_err = torch.where(take, err, best_err) + return best_out.reshape(shp) + + +SCHEME_RE = re.compile(r"^int(2|3|4|8)(?:-g(\d+))?(-e8u?)?(-rot)?(-nohead)?$") def parse_scheme(name): - """'int4-g128-nohead' -> (bits=4, group=128, skip_head=True). 'fp16' -> None.""" + """'int4-g128-nohead' -> (bits, group, e8, skip_head...). 'fp16' -> None.""" if name == "fp16": return None m = SCHEME_RE.match(name) if not m: - raise SystemExit(f"bad scheme '{name}' (expected fp16 | int{{2,3,4,8}}[-g][-rot][-nohead])") - return int(m.group(1)), int(m.group(2) or 0), bool(m.group(3)), bool(m.group(4)) + raise SystemExit(f"bad scheme '{name}' (expected fp16 | int{{2,3,4,8}}[-g][-e8|-e8u][-rot][-nohead])") + return int(m.group(1)), int(m.group(2) or 0), m.group(3) or "", bool(m.group(4)), bool(m.group(5)) def is_router(name): @@ -152,7 +225,7 @@ def apply_scheme(model, scheme): spec = parse_scheme(scheme) if spec is None: return 0, 0, total - bits, group, rot, skip_head = spec + bits, group, e8, rot, skip_head = spec n = qp = 0 with torch.no_grad(): for name, p in model.named_parameters(): @@ -160,7 +233,7 @@ def apply_scheme(model, scheme): continue if skip_head and is_head_or_embed(name): continue - p.data.copy_(quantize_param(p.data.float(), bits, group, rot).to(p.dtype)) + p.data.copy_(quantize_param(p.data.float(), bits, group, rot, e8).to(p.dtype)) n += 1 qp += p.numel() return n, qp, total