tools: E8 lattice quantization for the ablation harness — int{2,3,4,8}[-gN][-e8|-e8u][-rot] (#81)
This commit is contained in:
+85
-12
@@ -106,33 +106,106 @@ def rotation(dim, device, seed=417):
|
|||||||
return q
|
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
|
if w.ndim == 3: # fused experts [E, in, out] -> move input last
|
||||||
x = w.transpose(1, 2).contiguous()
|
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()
|
return x.transpose(1, 2).contiguous()
|
||||||
if rot:
|
if rot:
|
||||||
return _rot_quant(w, bits, group)
|
return _rot_quant(w, bits, group, e8)
|
||||||
return _quant_last_dim(w, bits, group) # nn.Linear [out, in] -- input already last
|
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."""
|
"""W -> Qn(W@Q) @ Q^T along the last (input) dim — see rotation() above."""
|
||||||
q = rotation(x.shape[-1], x.device)
|
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):
|
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":
|
if name == "fp16":
|
||||||
return None
|
return None
|
||||||
m = SCHEME_RE.match(name)
|
m = SCHEME_RE.match(name)
|
||||||
if not m:
|
if not m:
|
||||||
raise SystemExit(f"bad scheme '{name}' (expected fp16 | int{{2,3,4,8}}[-g<N>][-rot][-nohead])")
|
raise SystemExit(f"bad scheme '{name}' (expected fp16 | int{{2,3,4,8}}[-g<N>][-e8|-e8u][-rot][-nohead])")
|
||||||
return int(m.group(1)), int(m.group(2) or 0), bool(m.group(3)), bool(m.group(4))
|
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):
|
def is_router(name):
|
||||||
@@ -152,7 +225,7 @@ def apply_scheme(model, scheme):
|
|||||||
spec = parse_scheme(scheme)
|
spec = parse_scheme(scheme)
|
||||||
if spec is None:
|
if spec is None:
|
||||||
return 0, 0, total
|
return 0, 0, total
|
||||||
bits, group, rot, skip_head = spec
|
bits, group, e8, rot, skip_head = spec
|
||||||
n = qp = 0
|
n = qp = 0
|
||||||
with torch.no_grad():
|
with torch.no_grad():
|
||||||
for name, p in model.named_parameters():
|
for name, p in model.named_parameters():
|
||||||
@@ -160,7 +233,7 @@ def apply_scheme(model, scheme):
|
|||||||
continue
|
continue
|
||||||
if skip_head and is_head_or_embed(name):
|
if skip_head and is_head_or_embed(name):
|
||||||
continue
|
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
|
n += 1
|
||||||
qp += p.numel()
|
qp += p.numel()
|
||||||
return n, qp, total
|
return n, qp, total
|
||||||
|
|||||||
Reference in New Issue
Block a user