use super::tests_olmo::{olmo_fixture_path, read_npy_f32_2d};
use super::*;
fn renormalize_rows(mut probs: Array2<f64>) -> Array2<f64> {
for mut row in probs.rows_mut() {
let sum: f64 = row.iter().sum();
assert!(
sum > 0.99 && sum < 1.01,
"fixture row is not near-simplex: {sum}"
);
row /= sum;
}
probs
}
#[test]
fn qwen_real_activation_behavior_fit_selects_identifiable_lambda_y() {
let activation_full = read_npy_f32_2d(&olmo_fixture_path("qwen35_9b_actsL21_pca64_2000.npy"));
let probabilities_full = renormalize_rows(read_npy_f32_2d(&olmo_fixture_path(
"qwen35_9b_behavior_probs64_2000.npy",
)));
assert_eq!(activation_full.dim(), (2000, 64));
assert_eq!(probabilities_full.dim(), (2000, 64));
const GATE_ROWS: usize = 600;
let activation = activation_full
.slice(ndarray::s![0..GATE_ROWS, ..])
.to_owned();
let probabilities = probabilities_full
.slice(ndarray::s![0..GATE_ROWS, ..])
.to_owned();
let mut config = SaeCrosscoderAutoFitConfig::standard(4, 3);
config.max_iter = 30;
config.run_outer_rho_search = false;
let report = run_auto_sae_behavior_fit(SaeBehaviorAutoFitRequest {
activation,
probabilities,
config,
cancel: None,
})
.expect("real Qwen activation/behavior fit must complete");
assert_eq!(report.crosscoder.layers.len(), 2);
for layer in &report.crosscoder.layers {
assert!(
layer.reconstruction_r2.is_finite() && layer.reconstruction_r2 > 0.0,
"{}: shared-chart reconstruction must beat the column-mean baseline, got {}",
layer.label,
layer.reconstruction_r2
);
}
let ident = &report.weight_identifiability;
assert!(
ident.identifiable,
"real data leaves residual variance in BOTH blocks; got {ident:?}"
);
assert!(ident.activation_residual_variance > 0.0);
assert!(ident.behavior_residual_variance > 0.0);
assert!(ident.log_lambda_curvature > 0.0);
assert!(report.behavior_block.log_lambda_y().is_finite());
assert_eq!(
report.kl.infinite_rows, 0,
"no fitted row may decode off-simplex"
);
assert_eq!(report.kl.finite_rows, GATE_ROWS);
let mean_kl = report
.kl
.mean_kl_nats
.expect("all-finite KL implies a mean");
assert!(mean_kl.is_finite() && mean_kl >= 0.0);
assert_eq!(report.isometry.len(), 4);
let wire = report.wire_report().expect("behavior wire report");
assert!(wire.lambda_y > 0.0);
assert_eq!(wire.target_probabilities.len(), GATE_ROWS);
assert_eq!(wire.fitted_probabilities.len(), GATE_ROWS);
serde_json::to_string(&wire).expect("wire report must serialize");
}
#[test]
fn zz2015_tiny_inner_crawl_terminates() {
let activation_full = read_npy_f32_2d(&olmo_fixture_path("qwen35_9b_actsL21_pca64_2000.npy"));
let probabilities_full = renormalize_rows(read_npy_f32_2d(&olmo_fixture_path(
"qwen35_9b_behavior_probs64_2000.npy",
)));
const TINY_ROWS: usize = 48;
let activation = activation_full
.slice(ndarray::s![0..TINY_ROWS, ..])
.to_owned();
let probabilities = probabilities_full
.slice(ndarray::s![0..TINY_ROWS, ..])
.to_owned();
let mut config = SaeCrosscoderAutoFitConfig::standard(4, 3);
config.max_iter = 30;
config.run_outer_rho_search = false;
let report = run_auto_sae_behavior_fit(SaeBehaviorAutoFitRequest {
activation,
probabilities,
config,
cancel: None,
})
.expect("tiny inner-crawl repro: the inner solve must TERMINATE (converge), not refuse");
assert_eq!(report.crosscoder.layers.len(), 2);
assert!(report.behavior_block.log_lambda_y().is_finite());
}