Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
23 changes: 21 additions & 2 deletions crates/larql-inference/src/forward/memit.rs
Original file line number Diff line number Diff line change
Expand Up @@ -427,8 +427,27 @@
}
}

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 {

Check warning on line 437 in crates/larql-inference/src/forward/memit.rs

View workflow job for this annotation

GitHub Actions / cargo-mutants (informational)

Missed mutant

replace > with >= in memit_solve_layer

Check warning on line 437 in crates/larql-inference/src/forward/memit.rs

View workflow job for this annotation

GitHub Actions / cargo-mutants (informational)

Missed mutant

replace > with < in memit_solve_layer

Check warning on line 437 in crates/larql-inference/src/forward/memit.rs

View workflow job for this annotation

GitHub Actions / cargo-mutants (informational)

Missed mutant

replace > with == in memit_solve_layer
eprintln!("MEMIT: Cholesky needed ridge escalation {ridge} → {current_ridge}");
}
break l;
}
Err(e) if current_ridge < ridge * 1000.0 => {

Check warning on line 442 in crates/larql-inference/src/forward/memit.rs

View workflow job for this annotation

GitHub Actions / cargo-mutants (informational)

Missed mutant

replace * with / in memit_solve_layer

Check warning on line 442 in crates/larql-inference/src/forward/memit.rs

View workflow job for this annotation

GitHub Actions / cargo-mutants (informational)

Missed mutant

replace * with + in memit_solve_layer

Check warning on line 442 in crates/larql-inference/src/forward/memit.rs

View workflow job for this annotation

GitHub Actions / cargo-mutants (informational)

Missed mutant

replace < with <= in memit_solve_layer

Check warning on line 442 in crates/larql-inference/src/forward/memit.rs

View workflow job for this annotation

GitHub Actions / cargo-mutants (informational)

Missed mutant

replace < with > in memit_solve_layer

Check warning on line 442 in crates/larql-inference/src/forward/memit.rs

View workflow job for this annotation

GitHub Actions / cargo-mutants (informational)

Missed mutant

replace < with == in memit_solve_layer

Check warning on line 442 in crates/larql-inference/src/forward/memit.rs

View workflow job for this annotation

GitHub Actions / cargo-mutants (informational)

Missed mutant

replace match guard current_ridge < ridge * 1000.0 with false in memit_solve_layer

Check warning on line 442 in crates/larql-inference/src/forward/memit.rs

View workflow job for this annotation

GitHub Actions / cargo-mutants (informational)

Missed mutant

replace match guard current_ridge < ridge * 1000.0 with true in memit_solve_layer
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),
Expand Down
28 changes: 18 additions & 10 deletions crates/larql-inference/src/forward/trace.rs
Original file line number Diff line number Diff line change
Expand Up @@ -298,22 +298,29 @@ 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<f32> suffices at our scales
// (we'll round to f32 when writing to disk anyway).
let mut ktk = Array2::<f32>::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::<f64>::zeros((ffn_dim, ffn_dim));
let mut total_samples: usize = 0;

// Re-process the first capture so we don't double-count it.
// `K^T K` for a (seq, ffn_dim) matrix: each row's outer product
// 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;
Expand All @@ -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;
Expand All @@ -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.
Expand Down
Loading