← Back to index

GPTQ `lm_head` Safety Gate + Precision Fix — Implementation Plan

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:

  1. Precision fix (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.
  2. Adaptive damping (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


1. Why (brief — full detail in BF16_CURVATURE_BUG_PLAIN_ENGLISH.md)

Two independent, real findings from tonight’s investigation:

  1. Precision — corrected/precise statement (2026-09-06, verified directly): Apple’s 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).
  2. No safety net: unlike YAQA’s tensors (which now have adaptive_safe_correction

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.


1b. Definitive justification: the original, canonical GPTQ reference implementation already does this

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 = False

And 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.


2. Change 1 of 2 — Precision fix (Catcher)

Original (verbatim, lines ~178, 217 of 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))

Proposed replacement

    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))

To revert change 1

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.


3. Change 2 of 2 — Safety gate + naive fallback

Original (verbatim, from 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).")

Proposed replacement

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).")

To revert change 2

Replace the entire block above with the “Original (verbatim…)” block in section 3 of this document, unchanged.


4. Testing plan before trusting this

  1. Run 06_lm_head_gptq.py once, standalone, against a scratch --resume-dir (not the live production resume dir) and confirm:
  2. Re-run the full test_regression.py suite — must still show ALL REGRESSION TESTS PASSED.
  3. Only after both pass: apply to the actual resume dir being used for the real reprocessing run.

5. Regression test to add (matches this project’s one-test-per-real-bug discipline)

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.