Merge pull request #347 from ZacharyZcR/tools/e8-rate-scaled
tools: rate-scale the E8 lattice ball by bit-width (#81) — int3-e8 is now real int3
This commit is contained in:
@@ -118,7 +118,7 @@ def quantize_param(w, bits, group, rot=False, e8=""):
|
|||||||
|
|
||||||
def _grid_or_e8(x, bits, group, e8):
|
def _grid_or_e8(x, bits, group, e8):
|
||||||
if e8:
|
if e8:
|
||||||
return _quant_e8(x.float(), group, ball=(e8 == "-e8"))
|
return _quant_e8(x.float(), group, bits, ball=(e8 == "-e8"))
|
||||||
return _quant_last_dim(x, bits, group)
|
return _quant_last_dim(x, bits, group)
|
||||||
|
|
||||||
|
|
||||||
@@ -169,8 +169,15 @@ def _e8_ball(y, r2=10.0):
|
|||||||
return p
|
return p
|
||||||
|
|
||||||
|
|
||||||
def _quant_e8(x, group, ball):
|
_E8_R2_REPORTED = set()
|
||||||
"""Blocks of 8 along the input dim; per-group scale by MSE search over RMS multiples."""
|
def _e8_radius(bits):
|
||||||
|
# E8 lattice: points within |p|^2<=r2 grow ~r2^4, so +1 bit (x256 codebook) needs r2 x4.
|
||||||
|
# Anchor: r2=10 is the ~2^16 E8P ball (2 bits over 8 dims). Scale from there.
|
||||||
|
return 10.0 * (4.0 ** (bits - 2))
|
||||||
|
|
||||||
|
def _quant_e8(x, group, bits, ball):
|
||||||
|
"""Blocks of 8 along the input dim; per-group scale by MSE search over RMS multiples.
|
||||||
|
ball=True clamps to the rate-scaled E8 ball for `bits`; ball=False is the unbounded ideal."""
|
||||||
if x.shape[-1] % 8:
|
if x.shape[-1] % 8:
|
||||||
raise SystemExit(f"-e8 needs input dim divisible by 8 (got {x.shape[-1]})")
|
raise SystemExit(f"-e8 needs input dim divisible by 8 (got {x.shape[-1]})")
|
||||||
g = group or x.shape[-1]
|
g = group or x.shape[-1]
|
||||||
@@ -183,7 +190,11 @@ def _quant_e8(x, group, ball):
|
|||||||
for k in (0.5, 0.7, 0.9, 1.1, 1.4, 1.8, 2.4):
|
for k in (0.5, 0.7, 0.9, 1.1, 1.4, 1.8, 2.4):
|
||||||
s = rms * k
|
s = rms * k
|
||||||
yb = (xg / s).reshape(-1, g // 8, 8)
|
yb = (xg / s).reshape(-1, g // 8, 8)
|
||||||
p = _e8_ball(yb) if ball else _e8_nearest(yb)
|
p = _e8_ball(yb, _e8_radius(bits)) if ball else _e8_nearest(yb)
|
||||||
|
if ball and bits not in _E8_R2_REPORTED:
|
||||||
|
_E8_R2_REPORTED.add(bits)
|
||||||
|
import sys as _sys
|
||||||
|
_sys.stderr.write(f"[e8] bits={bits}: ball r2={_e8_radius(bits):.1f}\n")
|
||||||
out = (p.reshape(-1, g) * s)
|
out = (p.reshape(-1, g) * s)
|
||||||
err = (out - xg).pow(2).sum(-1, keepdim=True)
|
err = (out - xg).pow(2).sum(-1, keepdim=True)
|
||||||
if best_err is None:
|
if best_err is None:
|
||||||
|
|||||||
Reference in New Issue
Block a user