2026-09-06
Author: Hakim Ghelab Organization: VegaLaboratories LTD Date: 2026-09-06 Status: IMPLEMENTED and regression-tested, 2026-09-06, in two real stages:
FP32Catcher) + safety gate
(gptq_safety_gate, test_regression.py
[14], 4 checks). Real, dry-run tested against the
actual model: on lm_head’s real Hessian, the fixed
default damp_ratio=1e-2 did NOT beat naive (frob_gptq=15.72 vs
frob_naive=10.89) — confirmed independently NOT caused by the precision
fix (the original, unmodified Apple Catcher does worse:
16.38). The safety gate correctly caught this and fell back to naive,
real, on the first try.gptq_correct_with_damping +
gptq_adaptive_correction, test_regression.py
[15], 4 checks) — ports the CONCEPT of YAQA’s
adaptive_safe_correction (try default, escalate through
progressively stronger regularization, fall back to naive only if
nothing passes) into GPTQ’s own, structurally different algorithm (full
matrix inverse, no H_O, no block-LDL) rather than reusing
YAQA’s block-LDL code. Real, honest finding from building the regression
test: even a genuinely well-conditioned synthetic Hessian does not
automatically make GPTQ beat naive — GPTQ’s error-propagation needs
REAL, informative curvature structure to exploit; a near-uniform Hessian
gives it nothing useful to compensate with, and it can lose to naive
even in that “nice” case. This matches the real, honest philosophy
already established for YAQA: damping and search widen the chance of
success, they do not guarantee one — the safety gate is what makes both
approaches trustworthy either way.Full suite: ALL REGRESSION TESTS PASSED (both
[14] and [15]).
Real dry run against lm_head with the FULL
adaptive search, completed 2026-09-06:
| damp_ratio | frob_gptq (x naive) |
|---|---|
| 1e-2 (default) | 1.4438x |
| 1e-1 | 1.2319x |
| 1.0 | 1.1267x |
| 10.0 | 1.1017x |
| naive baseline | 1.0000x |
Adaptive damping genuinely, monotonically helped (44% worse -> 10%
worse as damping increased 1000x over the default) but never crossed
below naive — result: method=naive_fallback, all 4 real
trials tried and correctly rejected, real, honest fallback used.
Conclusion: on this specific real tensor, the shortfall is not
purely a damping-strength problem — likely GPTQ’s one-sided,
full-matrix-inverse structure itself, not something a stronger
damp_ratio alone can fully fix. Adaptive damping stays in production
regardless (it is a real, monotonic improvement and would fully rescue a
tensor where the shortfall IS a pure damping problem) — the safety-gate
fallback is doing genuinely necessary work here, not a formality. This
document remains the exact record of what changed, so it can still be
reverted with zero ambiguity if a real run surfaces a problem. This
project is not a git repository at the scripts/yaqa_port/
level yet (the parent project has one; yaqa_port is
intended to become its own standalone repo — see project notes), so
there is no git diff/ git checkout safety net
for this specific file yet — this document is the safety net.
Target file: 06_lm_head_gptq.py
BF16_CURVATURE_BUG_PLAIN_ENGLISH.md)Two independent, real findings from tonight’s investigation:
Catcher.__init__ seeds
self.H = mx.array(0.0), a float32 scalar. Confirmed by
direct test: float32 + bfloat16 -> float32 in MLX (type
promotion, not truncation), so the running SUM across calibration
sequences is already float32 today — this is NOT the same failure mode
as YAQA’s original bug (where the accumulator itself stayed bf16 the
whole way). What remains bf16 is the PER-SEQUENCE Gram formation itself
(xf.T @ xf, computed while xf is still bf16,
before being promoted and added to the float32 running total) — the same
distinction as YAQA’s “accumulation-only” vs “Gram-formation-first”
fixes tested earlier tonight; this is the “Gram-formation-first” (more
rigorous) fix applied to GPTQ. A real, isolated test
(bf16-formation-then-promoted vs fp32-formation throughout) showed FP32
formation reduces the spurious negative-eigenvalue magnitude by ~4
orders of magnitude (bf16 path: min_eig=-0.0676, fp32 path:
min_eig=-0.0000085 — same tensor, same calibration). Less severe than
the catastrophic YAQA tensor (roughly 18-23x smaller relative to this
matrix’s own scale), but the same real mechanism, and
mx.linalg.cholesky already receives a float32 array either
way (it was never actually being handed bfloat16 – the earlier claim
that GPTQ would crash on bf16 input was itself imprecise; GPTQ already
“supports” float32 today via this promotion, the open question was only
ever about Gram-formation precision, not dtype compatibility).adaptive_safe_correction
safety_gate, with automatic fallback
to naive), 06_lm_head_gptq.py computes
frob_gptq vs frob_naive and only
prints them — the GPTQ-corrected result is saved
unconditionally, even if it turns out worse than naive.
lm_head produces every output token; it has the least
protection of any tensor in the pipeline right now.Both are being fixed at once because they are cheap (measured: ~14s vs ~15s for the precision fix — no meaningful cost) and address different failure modes (precision reduces how often a bad result occurs; the safety gate catches whichever bad result gets through regardless of cause).
Explicitly NOT changing: GPTQ’s own algorithm — the
damping formula, the Cholesky/inverse math, the column-by-column
compensation loop. Two-sided (H_O) correction is not being
added; that remains structurally impossible here (246.7GB) and is an
accepted, documented trade-off, not a gap.
Checked directly, 2026-09-06, against the original GPTQ paper
authors’ own implementation (found locally at
/Users/hghelab/HybridOptiQ_FINAL_BENCHMARK/research/quantization_methods_2026/gptq-main/gptq.py,
dated 2023-07-11 — the real reference code Apple’s
mlx_lm.quant.gptq port would have been adapted from, not
assumed, read in full).
The original authors globally disable reduced-precision matmul, at module load time, unconditionally:
torch.backends.cuda.matmul.allow_tf32 = False
torch.backends.cudnn.allow_tf32 = FalseAnd in add_batch() — the exact function
Catcher.__call__ reimplements — the original explicitly
casts to float BEFORE the outer product:
inp = math.sqrt(2 / self.nsamples) * inp.float()
self.H += inp.matmul(inp.t())This is not a new precaution invented for this project. It is
restoring behavior the canonical, original GPTQ reference implementation
has always had, which Apple’s MLX port
(mlx_lm/quant/gptq.py) silently dropped —
Catcher.__call__ computes xf.T @ xf with no
explicit cast at all; the float32 promotion found in section 1 above
happens only by accident, via the mx.array(0.0) scalar
seed, and only protects the running sum, never the per-batch outer
product the original code explicitly protects.
This changes the framing of change 1 (section 2 below) from “a precision improvement worth adding” to “correcting a real, verifiable divergence from the reference algorithm this port is supposed to implement.”
One further, smaller divergence noted but explicitly NOT
being fixed here (out of scope — would be an actual algorithm
change, not a precision fix): the original computes a running AVERAGE
(self.H *= nsamples/(nsamples+tmp), then rescales each new
batch by sqrt(2/nsamples)) rather than Apple’s/this port’s
plain running SUM (self.H = self.H + xf.T @ xf).
Mathematically related, not identical in absolute scale — likely
immaterial since damping is computed relative to H’s own
mean diagonal either way, but flagged here for completeness rather than
silently ignored.
06_lm_head_gptq.py) from mlx_lm.utils import load
from mlx_lm.quant.gptq import Catcher orig_leaves_list = tree_flatten(model.leaf_modules(), is_leaf=nn.Module.is_module)
catcher = Catcher(orig_module)
leaves_list = [(k, catcher if k == LM_HEAD_NAME else l) for k, l in orig_leaves_list]
model.update_modules(tree_unflatten(leaves_list)) from mlx_lm.utils import load
from mlx_lm.quant.gptq import Catcher # kept for reference/revert; not used directly below
class FP32Catcher(nn.Module):
"""Reimplements Apple's mlx_lm.quant.gptq.Catcher (source read
directly; identical logic otherwise) with ONE change: casts x to
float32 before the Hessian outer product. The actual forward call
to lm_head (self.module(x, ...)) still receives the ORIGINAL
(bf16) x unchanged -- this only affects curvature bookkeeping, not
the model's real forward pass. NOTE (corrected 2026-09-06): the
running SUM (self.H) is already float32 in the original Catcher too
-- self.H starts as mx.array(0.0) (float32 scalar), and MLX
promotes float32+bfloat16 to float32 (confirmed directly), so the
accumulator was never the bf16 part. What THIS change fixes is the
PER-SEQUENCE outer product xf.T@xf, which is still computed in bf16
(xf is bf16) before being promoted-and-added to the float32 running
total -- casting xf to float32 BEFORE the outer product, not after,
matches the more rigorous of the two fixes already validated
tonight for YAQA's own curvature (see
BF16_CURVATURE_BUG_PLAIN_ENGLISH.md)."""
def __init__(self, module):
super().__init__()
self.module = module
self.H = mx.array(0.0)
def __call__(self, x, *args, **kwargs):
xf = x.flatten(0, -2).astype(mx.float32)
self.H = self.H + xf.T @ xf
return self.module(x, *args, **kwargs) orig_leaves_list = tree_flatten(model.leaf_modules(), is_leaf=nn.Module.is_module)
catcher = FP32Catcher(orig_module)
leaves_list = [(k, catcher if k == LM_HEAD_NAME else l) for k, l in orig_leaves_list]
model.update_modules(tree_unflatten(leaves_list))Delete the FP32Catcher class definition and change
catcher = FP32Catcher(orig_module) back to
catcher = Catcher(orig_module). The
from mlx_lm.quant.gptq import Catcher import line does not
need to change either way.
with mx.stream(mx.cpu): through the end of
main()) print(f"\nReal GPTQ column-by-column compensation for {LM_HEAD_NAME}...")
with mx.stream(mx.cpu):
damp = 1e-2 * mx.mean(mx.diag(H))
diag_idx = mx.arange(H.shape[0])
H2 = mx.array(H)
H2[diag_idx, diag_idx] += damp
Hchol = mx.linalg.cholesky(H2)
Hchol = mx.linalg.cholesky_inv(Hchol)
Hinv = mx.linalg.cholesky(Hchol, upper=True)
mx.eval(Hinv)
W = orig_module.weight.astype(mx.float32)
mx.eval(W)
@mx.compile
def gptq_error(w, d, scales, biases, bits=bits):
n_bins = 2**bits - 1
q = mx.clip(mx.round((w - biases) / scales), 0.0, n_bins)
q = scales * q + biases
return (w - q) / d
all_scales, all_biases = [], []
for i in range(0, W.shape[-1], group_size):
j = i + group_size
Wl = W[..., i:j]
err = mx.zeros_like(Wl)
_, scales, biases = mx.quantize(Wl, bits=bits, group_size=group_size)
all_scales.append(scales)
all_biases.append(biases)
for k in range(group_size):
k += i
w = W[..., k:k + 1]
d = Hinv[k, k]
e = gptq_error(w, d, scales, biases)
W[..., k:k + j] -= e @ Hinv[k:k + 1, k:k + j]
err[..., k:k + 1] = e
mx.eval(err, W)
W[..., j:] -= err @ Hinv[i:j, j:]
mx.eval(W)
print("Real column-by-column compensation done.")
# Real pack, mx.quantize on the NOW-corrected W -- idempotent here for
# the same reason established in yaqa_core.ldlq2hess_quantize's
# docstring: this is the FIRST and only real quantize call on this
# exact corrected array, not a re-quantize of an already-dequantized
# result, so there is no idempotency risk to guard against.
wq, scales, biases = mx.quantize(W, bits=bits, group_size=group_size)
mx.eval(wq, scales, biases)
finite_frac = float(mx.mean(mx.isfinite(mx.dequantize(wq, scales, biases, bits=bits, group_size=group_size)).astype(mx.float32)))
print(f"Real packed result: {finite_frac*100:.1f}% finite")
assert finite_frac == 1.0, "lm_head's real GPTQ-corrected, packed weight is not fully finite"
# Real comparison against naive (uncorrected) rounding, same spirit as
# every YAQA tensor's Metric A/B, so lm_head's manifest entry is
# directly comparable in kind, even though the underlying method differs.
naive_wq, naive_scales, naive_biases = mx.quantize(orig_module.weight.astype(mx.float32), bits=bits, group_size=group_size)
naive_W = mx.dequantize(naive_wq, naive_scales, naive_biases, bits=bits, group_size=group_size)
corrected_W = mx.dequantize(wq, scales, biases, bits=bits, group_size=group_size)
orig_W = orig_module.weight.astype(mx.float32)
frob_gptq = float(mx.sqrt(mx.sum((orig_W - corrected_W) ** 2)))
frob_naive = float(mx.sqrt(mx.sum((orig_W - naive_W) ** 2)))
print(f"Real Frobenius: gptq-corrected={frob_gptq:.4f} naive={frob_naive:.4f}")
# Write into the SAME resume-cache the real full run already uses --
# pure addition, the other 362 real manifest entries are never touched.
args.resume_dir.mkdir(parents=True, exist_ok=True)
safe_name = LM_HEAD_NAME.replace("/", "_")
mx.save_safetensors(str(args.resume_dir / f"{safe_name}.safetensors"),
{"weight": wq, "scales": scales, "biases": biases})
manifest_path = args.resume_dir / "manifest.json"
manifest = json.loads(manifest_path.read_text(encoding="utf-8")) if manifest_path.is_file() else {}
manifest[LM_HEAD_NAME] = {
"bits": bits, "group_size": group_size, "rotated": False,
"in_features": in_features, "out_features": out_features,
"method": "gptq", # real, explicit tag -- distinguishes this from the "yaqa"-implicit
# entries (no "method" key) written by 05_full_model_quantize.py
"frob_gptq": frob_gptq, "frob_naive": frob_naive,
}
manifest_path.write_text(json.dumps(manifest, indent=2), encoding="utf-8")
print(f"\nWrote real GPTQ-corrected {LM_HEAD_NAME} to {manifest_path} "
f"(1 new entry added, other {len(manifest)-1} real entries untouched).")Only two insertions and one restructure — the damping/inversion/compensation math itself is byte-for-byte unchanged.
print(f"\nReal GPTQ column-by-column compensation for {LM_HEAD_NAME}...")
with mx.stream(mx.cpu):
damp = 1e-2 * mx.mean(mx.diag(H))
diag_idx = mx.arange(H.shape[0])
H2 = mx.array(H)
H2[diag_idx, diag_idx] += damp
Hchol = mx.linalg.cholesky(H2)
Hchol = mx.linalg.cholesky_inv(Hchol)
Hinv = mx.linalg.cholesky(Hchol, upper=True)
mx.eval(Hinv)
# REAL SAFETY GATE, INSERTION 1 (2026-09-06): Hinv itself was never
# checked for finiteness before -- only the FINAL packed weight was,
# which would catch total corruption but let a "quietly bad" Hinv
# propagate silently through the whole compensation loop first. Check
# immediately, before any real compute is spent using it.
hinv_finite = bool(mx.all(mx.isfinite(Hinv)))
if not hinv_finite:
print(" Hinv is non-finite -- GPTQ correction is unusable for this tensor, "
"skipping compensation entirely and falling back to naive.")
W = orig_module.weight.astype(mx.float32)
mx.eval(W)
if hinv_finite:
@mx.compile
def gptq_error(w, d, scales, biases, bits=bits):
n_bins = 2**bits - 1
q = mx.clip(mx.round((w - biases) / scales), 0.0, n_bins)
q = scales * q + biases
return (w - q) / d
all_scales, all_biases = [], []
for i in range(0, W.shape[-1], group_size):
j = i + group_size
Wl = W[..., i:j]
err = mx.zeros_like(Wl)
_, scales, biases = mx.quantize(Wl, bits=bits, group_size=group_size)
all_scales.append(scales)
all_biases.append(biases)
for k in range(group_size):
k += i
w = W[..., k:k + 1]
d = Hinv[k, k]
e = gptq_error(w, d, scales, biases)
W[..., k:k + j] -= e @ Hinv[k:k + 1, k:k + j]
err[..., k:k + 1] = e
mx.eval(err, W)
W[..., j:] -= err @ Hinv[i:j, j:]
mx.eval(W)
print("Real column-by-column compensation done.")
# Real pack, mx.quantize on the NOW-corrected W -- idempotent here for
# the same reason established in yaqa_core.ldlq2hess_quantize's
# docstring: this is the FIRST and only real quantize call on this
# exact corrected array, not a re-quantize of an already-dequantized
# result, so there is no idempotency risk to guard against.
wq, scales, biases = mx.quantize(W, bits=bits, group_size=group_size)
mx.eval(wq, scales, biases)
finite_frac = float(mx.mean(mx.isfinite(mx.dequantize(wq, scales, biases, bits=bits, group_size=group_size)).astype(mx.float32)))
print(f"Real packed result: {finite_frac*100:.1f}% finite")
else:
finite_frac = 0.0
wq = scales = biases = None
# Real comparison against naive (uncorrected) rounding, same spirit as
# every YAQA tensor's Metric A/B, so lm_head's manifest entry is
# directly comparable in kind, even though the underlying method differs.
naive_wq, naive_scales, naive_biases = mx.quantize(orig_module.weight.astype(mx.float32), bits=bits, group_size=group_size)
naive_W = mx.dequantize(naive_wq, naive_scales, naive_biases, bits=bits, group_size=group_size)
orig_W = orig_module.weight.astype(mx.float32)
frob_naive = float(mx.sqrt(mx.sum((orig_W - naive_W) ** 2)))
# REAL SAFETY GATE, INSERTION 2 (2026-09-06): previously the GPTQ result
# was saved UNCONDITIONALLY -- frob_gptq/frob_naive were computed and
# printed but nothing ever acted on them. GPTQ has no H_O / weighted
# metric to trade Frobenius against (unlike YAQA), so beating naive on
# Frobenius is the correct, sufficient bar here. finite_frac==1.0 is
# required too (redundant with the Hinv check above in most real
# failure modes, but cheap and catches anything the Hinv check alone
# would miss).
if finite_frac == 1.0:
corrected_W = mx.dequantize(wq, scales, biases, bits=bits, group_size=group_size)
frob_gptq = float(mx.sqrt(mx.sum((orig_W - corrected_W) ** 2)))
else:
frob_gptq = float("inf")
gptq_safe = (finite_frac == 1.0) and (frob_gptq < frob_naive)
print(f"Real Frobenius: gptq-corrected={frob_gptq:.4f} naive={frob_naive:.4f} SAFE={gptq_safe}")
if gptq_safe:
final_wq, final_scales, final_biases = wq, scales, biases
method = "gptq"
else:
print(" GPTQ correction did not beat naive (or was unusable) -- "
"real, honest fallback to naive quantization for this tensor.")
final_wq, final_scales, final_biases = naive_wq, naive_scales, naive_biases
method = "naive_fallback"
# Write into the SAME resume-cache the real full run already uses --
# pure addition, the other 362 real manifest entries are never touched.
args.resume_dir.mkdir(parents=True, exist_ok=True)
safe_name = LM_HEAD_NAME.replace("/", "_")
mx.save_safetensors(str(args.resume_dir / f"{safe_name}.safetensors"),
{"weight": final_wq, "scales": final_scales, "biases": final_biases})
manifest_path = args.resume_dir / "manifest.json"
manifest = json.loads(manifest_path.read_text(encoding="utf-8")) if manifest_path.is_file() else {}
manifest[LM_HEAD_NAME] = {
"bits": bits, "group_size": group_size, "rotated": False,
"in_features": in_features, "out_features": out_features,
"method": method, # "gptq" or "naive_fallback" -- real, explicit, auditable
"frob_gptq": frob_gptq, "frob_naive": frob_naive,
}
manifest_path.write_text(json.dumps(manifest, indent=2), encoding="utf-8")
print(f"\nWrote real {method}-corrected {LM_HEAD_NAME} to {manifest_path} "
f"(1 new entry added, other {len(manifest)-1} real entries untouched).")Replace the entire block above with the “Original (verbatim…)” block in section 3 of this document, unchanged.
06_lm_head_gptq.py once, standalone, against a
scratch --resume-dir (not the live production resume dir)
and confirm:
SAFE=True or SAFE=False
explicitly (no silent path)."method" field reads
"gptq" or "naive_fallback" — never crashes,
never leaves an ambiguous state.frob_gptq (if computed) is lower under fp32
accumulation than it would have been under the old bf16 Catcher — real,
direct before/after check.test_regression.py suite — must still
show ALL REGRESSION TESTS PASSED.A synthetic test constructing a deliberately non-finite
Hinv (or a correction that’s provably worse than naive on a
small synthetic W/H) and asserting the fallback path is taken and
manifest["method"] == "naive_fallback" — mirrors
test_effective_rank_and_safety_gate’s existing style for
YAQA’s own safety gate.
Hakim Ghelab VegaLaboratories LTD
© 2026 Hakim Ghelab, VegaLaboratories LTD. All rights reserved.