← Back to index

YAQA-UMA Root-Cause Fix — Float64 Curvature Construction

2026-09-06

Not to be confused with: a second, unrelated BF16/F32 bug found 2026-09-07, in a completely different part of the pipeline — this document is about the curvature construction (the Hessian used during correction) needing enough precision to stay positive-semidefinite. The later bug is about the scales/biases metadata (saved after correction is already done) being written in too much precision, which cost real inference speed. Different stage, different symptom, different fix. See RCA_scales_f32_bug.html for that one.

Author: Hakim Ghelab
Organization: VegaLaboratories LTD
Project: YAQA-UMA — Apple MLX / Unified Memory Architecture port
Status: Root-cause fix specification — implemented and shipped. The real 363-tensor production run (2026-09-07) used this fix; every tensor previously processed under the buggy bf16 curvature was auto-invalidated (CURVATURE_VERSION) and recomputed correctly under it. Not just a specification anymore — see CHANGELOG.md’s 2026-09-06 entry for the applied-fix record.


1. Confirmed Root Cause

The production YAQA curvature collector currently constructs and accumulates the two Sketch-B curvature factors in FP32:

For the real failing tensor:

language_model.model.layers.14.linear_attn.in_proj_qkv
shape = (10240, 5120)

the FP32 curvature matrices developed impossible negative eigenvalues:

FP32 H_I:
min_eig  = -0.871
neg_eigs = 2225 / 5120
eff_rank ≈ 1.6

FP32 H_O:
min_eig  = -2.326
neg_eigs = 4685 / 10240
eff_rank ≈ 1.9

But H_I and H_O are sums of Gram matrices and are therefore PSD in exact arithmetic.

The corrected CPU-float64 experiment reduced the negative spectrum to numerical zero:

FLOAT64 H_I:
min_eig ≈ -1.63e-12
eff_rank ≈ 1.58

FLOAT64 H_O:
min_eig ≈ -9.18e-12
eff_rank ≈ 1.90

Therefore:

The low effective rank is real. The large negative eigenvalues are not. They are an FP32 numerical-precision artifact.

The root cause is not YAQA, not the calibration stratification, and not missing Hadamard rotation.

The root cause is:

extremely ill-conditioned / near-low-rank curvature
        +
FP32 curvature construction and accumulation
        ↓
fake negative curvature
        ↓
unstable LDL / LDLQ error propagation
        ↓
catastrophic YAQA correction on sensitive tensors

2. Required Production Fix

Do NOT change the model forward/backward precision

The model may continue to run in BF16.

The required change is specifically in the YAQA curvature construction path.

The production collector must:

  1. take the captured activation x and output-gradient grad_output;
  2. move/cast them to CPU float64 before any curvature matmul;
  3. construct the per-example curvature contribution in float64;
  4. accumulate H_I and H_O in float64;
  5. keep the raw curvature matrices in float64 for:
  6. apply damping/regularization in float64;
  7. only after the curvature has been stabilized may the regularized factors be converted to FP32 for the existing Metal LDLQ path if required.

3. Correct Float64 Curvature Construction

The old pattern must not do this:

Gb = gb.T @ xb
Gb64 = Gb.astype(mx.float64)

That is too late.

The cast must happen before the Gram construction.

Preferred implementation

Confirmed directly, 2026-09-06 (crashed without this): every float64 array creation/allocation — mx.zeros, mx.eye, casts, temporaries — must occur inside with mx.stream(mx.cpu):, or MLX raises ValueError: float64 is not supported on the GPU, even for an op whose output dtype is float64 but whose call site is not wrapped. Also confirmed directly: never combine a GPU-resident FP32/BF16 array and a CPU-resident float64 array in the same mx.eval() call — this raises the same error even though each array evaluates fine on its own. Keep GPU loss evaluation (mx.eval(loss_value)) and CPU curvature evaluation (mx.eval(H_in64, H_out64)) in two separate calls, always.

with mx.stream(mx.cpu):
    x64 = x[b].astype(mx.float64)
    g64 = grad_output[b].astype(mx.float64)

    G64 = g64.T @ x64

    H_in64  += G64.T @ G64
    H_out64 += G64 @ G64.T

    mx.eval(H_in64, H_out64)  # NEVER mx.eval(loss_value, H_in64, H_out64)

Initialize (also inside the CPU stream — see note above):

with mx.stream(mx.cpu):
    H_in64 = mx.zeros((in_features, in_features), dtype=mx.float64)
    H_out64 = mx.zeros((out_features, out_features), dtype=mx.float64)

4. Memory-Safer Equivalent Form

For a tensor such as:

W = [10240, 5120]

materializing:

G = [10240, 5120]

in float64 is expensive.

Use the mathematically equivalent token-space formulation:

with mx.stream(mx.cpu):
    x64 = x[b].astype(mx.float64)
    g64 = grad_output[b].astype(mx.float64)

    Kg = g64 @ g64.T
    Kx = x64 @ x64.T

    H_in64  += x64.T @ (Kg @ x64)
    H_out64 += g64.T @ (Kx @ g64)

    mx.eval(H_in64, H_out64)

This is algebraically equivalent to:

G = g64.T @ x64
H_in64  += G.T @ G
H_out64 += G @ G.T

but avoids materializing the full out × in matrix G.

For 128-token calibration windows, the intermediate Kg and Kx are only:

128 × 128

5. Post-Collection Processing

After all 24 stratified calibration windows have been accumulated:

H_in64 = 0.5 * (H_in64 + H_in64.T)
H_out64 = 0.5 * (H_out64 + H_out64.T)

Keep these as the trusted raw curvature reference.

Then normalize, still in float64:

mean_diag_in = mx.mean(mx.diag(H_in64))
mean_diag_out = mx.mean(mx.diag(H_out64))

H_in_norm64 = H_in64 / mean_diag_in
H_out_norm64 = H_out64 / mean_diag_out

Apply damping in float64 (mx.eye(..., dtype=mx.float64) must also stay inside the CPU stream — see the note in section 3):

with mx.stream(mx.cpu):
    H_in_reg64 = (
        H_in_norm64
        + sigma_I * mx.eye(H_in_norm64.shape[0], dtype=mx.float64)
    )

    H_out_reg64 = (
        H_out_norm64
        + sigma_O * mx.eye(H_out_norm64.shape[0], dtype=mx.float64)
    )

6. Metal / LDLQ Boundary

If the existing regularized_block_ldl() / LDLQ_2hess implementation must run on Metal in FP32, convert only the already-regularized matrices or resulting factors:

H_in_reg32 = H_in_reg64.astype(mx.float32)
H_out_reg32 = H_out_reg64.astype(mx.float32)

Then run the existing block-LDL / LDLQ implementation.

Do not discard the raw float64 curvature matrices.

The raw float64 matrices remain the authoritative scoring reference.


7. Candidate Scoring Must Use Trusted Curvature

For every candidate produced by adaptive damping, score it against the undamped float64 curvature, not the corrupted FP32 matrices and not the damped matrices.

For:

E = W - W_hat

use:

J = tr(E H_I Eᵀ H_O)

with trusted float64 H_I and H_O.

The safety gate must never use the old indefinite FP32 curvature.

The old observations such as:

weighted = -4038x naive

are not valid physical error scores because the curvature matrices were indefinite.


8. Adaptive Damping Stays

Keep the per-tensor adaptive damping mechanism:

default sigma
    ↓
safety gate
    ↓
if failure:
    search stronger sigma_I / sigma_O
    ↓
accept first/best safe candidate

But damping must operate on a numerically trustworthy curvature estimate.

Damping is a stabilizer. It is not a substitute for fixing corrupted FP32 curvature.


9. Fixed Grid vs Adaptive Grid

The real failing tensor produced:

sigma_O = 1.0

adaptive grid:
frob ≈ 300.6x naive

fixed original-W grid:
frob ≈ 20.7x naive

while the accepted candidate produced:

sigma_O = 10.0

adaptive grid:
frob ≈ 1.016x naive

fixed original-W grid:
frob ≈ 1.017x naive

Therefore:

The first priority is the float64 curvature fix.


10. Calibration Size Finding

In YAQA-UMA:

N=4 = 4 windows × 6 domains = 24 real calibration windows
N=8 = 8 windows × 6 domains = 48 real calibration windows

The low effective rank remained stable:

24 windows:
H_I eff_rank ≈ 1.6
H_O eff_rank ≈ 1.9

48 windows:
H_I eff_rank ≈ 1.62
H_O eff_rank ≈ 2.34

Therefore the extreme low-rank structure is not simply a calibration-size artifact.

Do not change the 24-window stratified production calibration merely to hide this numerical problem.


11. Required Regression Test Before Production

Before touching the live production cache, run the corrected collector on:

  1. the known catastrophic tensor:

    language_model.model.layers.14.linear_attn.in_proj_qkv
  2. one previously marginal tensor;

  3. one normal tensor that previously passed immediately.

For each tensor compare:

OLD FP32 curvature
NEW stable float64 curvature
CPU float64 reference

Report:

min eigenvalue
max eigenvalue
PSD violation magnitude
effective rank
selected sigma_I
selected sigma_O
weighted YAQA objective
naive weighted objective
Frobenius ratio
accepted bit/group configuration
teacher → candidate logit KL
teacher → naive logit KL
runtime
peak memory

12. Cache Invalidation

The current resume system only checks whether the tensor exists and whether bits/group size match.

Add a curvature/cache schema identifier, for example:

{
  "yaqa_curvature_version": 2,
  "curvature_accumulation": "cpu_float64",
  "curvature_formula": "sketch_b_token_factorized",
  "curvature_dtype": "float64"
}

Resume logic must reject old entries generated using the FP32 curvature path.


13. Scope of Reprocessing

Once the new collector is validated:

Regenerate every YAQA tensor whose correction was produced from the old FP32 curvature, not only the 14 catastrophic tensors (real, fresh manifest audit 2026-09-06 – 12 was a stale earlier estimate).

The 14 failures are merely the cases where numerical corruption became obvious.

Other tensors may still have:

while still passing the old safety gate.

The calibration corpus, quantization plan, sidecars, MLX packing format, MTP preservation, vision preservation, and model architecture do not need to be rebuilt merely because the curvature collector changes.


14. Production Architecture After the Fix

BF16 model
   │
   ├── forward/backward on MLX / Metal
   │
   ├── capture x + grad_output
   │
   ▼
CPU FLOAT64 curvature construction
   │
   ├── x/g cast BEFORE curvature matmul
   ├── H_I float64
   ├── H_O float64
   ├── symmetrize
   ├── normalize
   └── preserve raw undamped H for scoring
   │
   ▼
FLOAT64 adaptive regularization
   │
   ├── sigma_I
   └── sigma_O
   │
   ▼
stable SPD factors
   │
   ▼
FP32 / Metal block-LDL + LDLQ_2hess
(if required by current implementation)
   │
   ▼
standard MLX affine quantization
Q2 / Q3 / Q4 / Q5 / Q6 / Q8
   │
   ▼
safety gate scored against trusted float64 curvature
   │
   ▼
pack exact accepted codes/scales/biases
   │
   ▼
standard MLX model artifact

15. Explicit Non-Fixes

Do not treat any of the following as the root-cause repair:

MLX_ENABLE_TF32=0

That does not change the Hessian accumulation dtype.

Do not:


16. Acceptance Criteria

The fix is production-ready when:


Final Decision

Fix the curvature collector, not YAQA.

The corrected design is:

BF16 model execution + CPU float64 YAQA curvature construction + float64 normalization/damping/scoring + Metal FP32 LDLQ after stabilization + standard MLX affine packing.

This repairs the confirmed numerical root cause while preserving YAQA, Apple MLX compatibility, unified-memory operation, mixed-bit support, standard model loading, and the rest of the YAQA-UMA architecture.


Hakim Ghelab
VegaLaboratories LTD
YAQA-UMA


© 2026 Hakim Ghelab, VegaLaboratories LTD. All rights reserved.