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.