From 4bee96cf9730fe4155397881212dd78f826e2958 Mon Sep 17 00:00:00 2001 From: Sxilent Date: Sun, 30 Aug 2026 10:29:20 -0400 Subject: [PATCH] fix(memit): accumulate FFN covariance in f64; escalate Cholesky ridge before failing estimate_ffn_covariance accumulated K^T K in f32. On Gemma 4 26B-A4B (ffn_dim=2112, activations ~1e3) the entries reach ~1e6*N and the 24-bit mantissa drops the low-order contributions, so the Gram matrix stops being numerically PSD and the Cholesky in memit.rs hits a negative pivot ('Cholesky failed'). Converting to f64 after accumulation (cov_f64) cannot recover what was already lost. - trace.rs: accumulate and scale in f64, downcast to f32 once at the boundary. Signature unchanged. - memit.rs: if Cholesky still fails, escalate ridge x10 up to 1000x before returning the error, logging each step. Defensive; the f64 accumulation is the fix. Originally fixed locally 2026-04-21/22; ported onto current main 2026-08-30. --- crates/larql-inference/src/forward/memit.rs | 23 +++++++++++++++-- crates/larql-inference/src/forward/trace.rs | 28 +++++++++++++-------- 2 files changed, 39 insertions(+), 12 deletions(-) diff --git a/crates/larql-inference/src/forward/memit.rs b/crates/larql-inference/src/forward/memit.rs index 4c24fc023..eee7f7ce1 100644 --- a/crates/larql-inference/src/forward/memit.rs +++ b/crates/larql-inference/src/forward/memit.rs @@ -427,8 +427,27 @@ fn memit_solve_layer( } } - let l = larql_compute::cpu::ops::linalg::cholesky(&cov_f64, ridge) - .map_err(|e| format!("MEMIT: Cholesky failed — {e}"))?; + // Adaptive ridge: if f64 accumulation still lands a negative pivot + // (rare — mostly defends against non-Gemma architectures we haven't + // profiled), escalate ridge up to 4 steps before giving up. + let mut current_ridge = ridge; + let l = loop { + match larql_compute::cpu::ops::linalg::cholesky(&cov_f64, current_ridge) { + Ok(l) => { + if current_ridge > ridge { + eprintln!("MEMIT: Cholesky needed ridge escalation {ridge} → {current_ridge}"); + } + break l; + } + Err(e) if current_ridge < ridge * 1000.0 => { + current_ridge *= 10.0; + eprintln!( + "MEMIT: Cholesky pivot failed ({e}); retrying with ridge={current_ridge}" + ); + } + Err(e) => return Err(format!("MEMIT: Cholesky failed after escalation — {e}")), + } + }; // Q = K @ C⁻¹ [N × ffn_dim] // We compute this as: for each fact i, q_i = C⁻¹ @ k_i (column), diff --git a/crates/larql-inference/src/forward/trace.rs b/crates/larql-inference/src/forward/trace.rs index 62ce36756..5c88f1cd1 100644 --- a/crates/larql-inference/src/forward/trace.rs +++ b/crates/larql-inference/src/forward/trace.rs @@ -298,9 +298,16 @@ pub fn estimate_ffn_covariance( let ffn_dim = first.shape()[1]; // Accumulator — K^T K across all sampled token positions. - // Float64 would be safer but Array2 suffices at our scales - // (we'll round to f32 when writing to disk anyway). - let mut ktk = Array2::::zeros((ffn_dim, ffn_dim)); + // + // Accumulate in f64. With ffn_dim in the thousands and activation + // magnitudes ~1e3 (observed on Gemma 4 26B-A4B, ffn_dim = 2112), each + // entry grows to ~1e6 * N; f32 (24-bit mantissa) loses the low-order + // contributions and the accumulated Gram matrix stops being numerically + // PSD, so the downstream Cholesky in memit.rs hits a negative pivot. + // Converting to f64 *after* accumulation (memit.rs `cov_f64`) cannot + // recover precision that was already lost here. We downcast to f32 once, + // at the boundary, when returning. + let mut ktk = Array2::::zeros((ffn_dim, ffn_dim)); let mut total_samples: usize = 0; // Re-process the first capture so we don't double-count it. @@ -308,12 +315,12 @@ pub fn estimate_ffn_covariance( // with itself, summed across rows. for row in first.rows() { for i in 0..ffn_dim { - let vi = row[i]; + let vi = row[i] as f64; if vi == 0.0 { continue; } for j in 0..ffn_dim { - ktk[[i, j]] += vi * row[j]; + ktk[[i, j]] += vi * (row[j] as f64); } } total_samples += 1; @@ -331,12 +338,12 @@ pub fn estimate_ffn_covariance( }; for row in k.rows() { for i in 0..ffn_dim { - let vi = row[i]; + let vi = row[i] as f64; if vi == 0.0 { continue; } for j in 0..ffn_dim { - ktk[[i, j]] += vi * row[j]; + ktk[[i, j]] += vi * (row[j] as f64); } } total_samples += 1; @@ -347,10 +354,11 @@ pub fn estimate_ffn_covariance( return None; } - // C = (K^T K) / N - let scale = 1.0 / total_samples as f32; + // C = (K^T K) / N, then downcast to f32 once at the boundary. + let scale = 1.0_f64 / total_samples as f64; ktk.mapv_inplace(|v| v * scale); - Some((ktk, total_samples)) + let ktk_f32 = ktk.mapv(|v| v as f32); + Some((ktk_f32, total_samples)) } /// Run a forward pass and capture both residuals and sparse activations.