diff --git a/src/cli.rs b/src/cli.rs index 14c2d5b..cd68470 100644 --- a/src/cli.rs +++ b/src/cli.rs @@ -16,7 +16,7 @@ use crate::evaluation::compare_measurements; use crate::fitness::FitnessMode; use crate::patch::{Patch, PatchValidation}; use crate::probe::{Probe, built_in, compile_probes, load_probe_file, measure_probes}; -use crate::search::{SearchCheckpoint, SearchConfig, SearchEvent, run_search}; +use crate::search::{SearchCheckpoint, SearchConfig, SearchEvent, default_search_layers, run_search}; #[derive(Parser)] #[command( @@ -161,7 +161,11 @@ struct SearchArgs { report: Option, #[arg(long, default_value_t = 200)] iters: usize, - #[arg(long, value_delimiter = ',')] + #[arg( + long, + value_delimiter = ',', + help = "Comma-separated layer indices to search. Default: an even spread across the model's full depth (derived from its layer count)." + )] layers: Option>, #[arg(long, value_delimiter = ',', default_value = "gate_proj,up_proj")] projections: Vec, @@ -518,7 +522,9 @@ fn search(args: SearchArgs, json: bool) -> Result<()> { .map(|projection| Projection::from_str(projection)) .collect::>>()?; let config = SearchConfig { - search_layers: args.layers.unwrap_or_else(|| vec![1, 2, 3, 4, 34]), + search_layers: args + .layers + .unwrap_or_else(|| default_search_layers(backend.architecture().layer_count())), search_projections: projections, max_iters: args.iters, control_penalty: args.penalty, diff --git a/src/probe.rs b/src/probe.rs index 41822e5..ce4d91e 100644 --- a/src/probe.rs +++ b/src/probe.rs @@ -92,9 +92,31 @@ pub fn compile_probes( .collect() } +/// Measure probe logit gaps. Non-finite gaps are a hard error — used by `eval` +/// and search baselines so reports never contain silent NaN transitions. pub fn measure_probes( backend: &mut B, probes: &[CompiledProbe], +) -> Result> { + let measurements = measure_probes_allowing_non_finite(backend, probes)?; + for measurement in &measurements { + if !measurement.gap.is_finite() { + return Err(Error::MeasurementMismatch(format!( + "probe {} produced non-finite gap {}", + measurement.name, measurement.gap + ))); + } + } + Ok(measurements) +} + +/// Like [`measure_probes`], but returns non-finite gaps as data. +/// +/// Search candidate evaluation uses this so a NaN from one flip can be rejected +/// and reverted without aborting the whole search. Callers must handle NaN. +pub fn measure_probes_allowing_non_finite( + backend: &mut B, + probes: &[CompiledProbe], ) -> Result> { if probes.is_empty() { return Err(Error::EmptyProbeSet); @@ -107,12 +129,6 @@ pub fn measure_probes( compiled.correct_id, compiled.wrong_id, )?; - if !gap.is_finite() { - return Err(Error::MeasurementMismatch(format!( - "probe {} produced non-finite gap {gap}", - compiled.probe.name - ))); - } Ok(ProbeMeasurement { name: compiled.probe.name.clone(), category: compiled.probe.category.clone(), diff --git a/src/search.rs b/src/search.rs index 73cc689..57ff809 100644 --- a/src/search.rs +++ b/src/search.rs @@ -11,7 +11,9 @@ use crate::backend::MiyagiBackend; use crate::error::{Error, Result}; use crate::fitness::{FitnessMode, compute_fitness}; use crate::patch::{Patch, PatchFlip}; -use crate::probe::{CompiledProbe, ProbeMeasurement, measure_probes}; +use crate::probe::{ + CompiledProbe, ProbeMeasurement, measure_probes, measure_probes_allowing_non_finite, +}; const CHECKPOINT_VERSION: u32 = 1; @@ -33,7 +35,10 @@ pub struct SearchConfig { impl Default for SearchConfig { fn default() -> Self { Self { - search_layers: vec![1, 2, 3, 4, 34], + // Empty means "resolve from the live backend in run_search via + // default_search_layers(layer_count)". Do not bake a fixed layer + // count here — that reintroduces 8B-era under-search on larger models. + search_layers: Vec::new(), search_projections: vec![Projection::Gate, Projection::Up], max_iters: 200, control_penalty: 2.0, @@ -47,6 +52,46 @@ impl Default for SearchConfig { } } +/// Upper bound on how many layers the computed default samples. +const MAX_DEFAULT_LAYERS: usize = 12; + +/// Evenly spaced layer indices spanning the model's full depth, used when the +/// caller does not pass `--layers` (or leaves `SearchConfig.search_layers` empty). +/// +/// A fixed absolute list (the old `[1, 2, 3, 4, 34]`) silently under-searches +/// any model deeper than the 8B it was tuned for — e.g. on a 64-layer model it +/// never samples past layer 34, capping results with no warning. +/// +/// - **Large models** (`layer_count > MAX_DEFAULT_LAYERS`): up to +/// `MAX_DEFAULT_LAYERS` points across `1..=layer_count-1` (skips layer 0). +/// - **Small models** (`layer_count ≤ MAX_DEFAULT_LAYERS`): every layer index +/// in `0..layer_count` (includes layer 0). +pub fn default_search_layers(layer_count: usize) -> Vec { + if layer_count == 0 { + return Vec::new(); + } + if layer_count <= MAX_DEFAULT_LAYERS { + return (0..layer_count).collect(); + } + let last = layer_count - 1; + let span = last - 1; // spread across [1, last] + let steps = MAX_DEFAULT_LAYERS - 1; + let mut layers: Vec = (0..MAX_DEFAULT_LAYERS) + .map(|i| 1 + (span * i + steps / 2) / steps) // integer round + .collect(); + layers.dedup(); + layers +} + +/// Fill empty `search_layers` from the live backend architecture. +pub fn resolve_search_layers(layer_count: usize, search_layers: &[usize]) -> Vec { + if search_layers.is_empty() { + default_search_layers(layer_count) + } else { + search_layers.to_vec() + } +} + #[derive(Clone, Debug, Deserialize, Serialize)] pub struct SearchCheckpoint { pub version: u32, @@ -165,12 +210,19 @@ where B: MiyagiBackend, F: FnMut(&SearchEvent), { + let mut config = config; + config.search_layers = resolve_search_layers( + backend.architecture().layer_count(), + &config.search_layers, + ); validate_config(backend, target_probes, control_probes, &config)?; let candidates = build_candidates(backend, &config)?; if candidates.is_empty() { return Err(Error::InvalidSearch("candidate pool is empty".to_owned())); } + // Baselines must be finite (strict measure_probes). Candidate flips may + // produce NaN — those use measure_probes_allowing_non_finite below. let fresh_target_baseline = measure_probes(backend, target_probes)?; let fresh_control_baseline = measure_probes(backend, control_probes)?; let architecture_signature = backend.architecture().signature().to_owned(); @@ -405,11 +457,12 @@ fn evaluate_candidate( .iter() .map(|index| target_probes[*index].clone()) .collect::>(); - let screen_measurements = measure_probes(backend, &screen_probes)?; + // Candidate path: allow non-finite gaps (reject this flip, keep searching). + let screen_measurements = measure_probes_allowing_non_finite(backend, &screen_probes)?; if !screen_measurements.iter().any(|measurement| { - baseline_by_name - .get(&measurement.name) - .is_some_and(|baseline| measurement.gap > *baseline) + baseline_by_name.get(&measurement.name).is_some_and(|baseline| { + measurement.gap.is_finite() && measurement.gap > *baseline + }) }) { return Ok(CandidateEvaluation::ScreenedOut); } @@ -421,7 +474,7 @@ fn evaluate_candidate( .collect::>(); for (index, probe) in target_probes.iter().enumerate() { if let std::collections::btree_map::Entry::Vacant(entry) = measured_by_index.entry(index) { - let measurement = measure_probes(backend, std::slice::from_ref(probe))? + let measurement = measure_probes_allowing_non_finite(backend, std::slice::from_ref(probe))? .into_iter() .next() .expect("one probe returns one measurement"); @@ -435,7 +488,15 @@ fn evaluate_candidate( .expect("every target probe was measured") }) .collect::>(); - let control = measure_probes(backend, control_probes)?; + let control = measure_probes_allowing_non_finite(backend, control_probes)?; + // Non-finite candidate gap → reject this flip only (NaN fails fitness >). + if target + .iter() + .chain(control.iter()) + .any(|measurement| !measurement.gap.is_finite()) + { + return Ok(CandidateEvaluation::Measured { fitness: f32::NAN }); + } let fitness = compute_fitness( config.fitness_mode, &target, @@ -727,6 +788,35 @@ mod tests { assert_eq!(resumed.next_u64(), expected); } + #[test] + fn default_layers_span_full_depth_on_large_models() { + // 64-layer model must reach deep layers, unlike the old fixed list. + let layers = default_search_layers(64); + assert_eq!(layers.len(), MAX_DEFAULT_LAYERS); + assert_eq!(*layers.first().unwrap(), 1); + assert_eq!(*layers.last().unwrap(), 63); + assert!(layers.iter().all(|&l| l < 64)); + assert!(layers.windows(2).all(|w| w[0] < w[1]), "strictly increasing"); + // The old default's deepest layer was 34; the new one goes well past it. + assert!(layers.iter().any(|&l| l > 34)); + } + + #[test] + fn default_layers_cover_every_layer_on_small_models() { + // Small models intentionally include layer 0 (see default_search_layers docs). + assert_eq!(default_search_layers(4), vec![0, 1, 2, 3]); + assert_eq!(default_search_layers(0), Vec::::new()); + } + + #[test] + fn empty_search_layers_resolve_from_backend_depth() { + let resolved = resolve_search_layers(64, &[]); + assert_eq!(resolved, default_search_layers(64)); + assert!(resolved.iter().any(|&l| l > 34)); + let explicit = resolve_search_layers(64, &[1, 2, 3]); + assert_eq!(explicit, vec![1, 2, 3]); + } + #[test] fn screen_uses_worst_gaps_then_name() { fn measurement(name: &str, gap: f32) -> ProbeMeasurement { diff --git a/tests/search_engine.rs b/tests/search_engine.rs index 21e6917..0a0f162 100644 --- a/tests/search_engine.rs +++ b/tests/search_engine.rs @@ -152,3 +152,162 @@ fn cancellation_is_reported_without_losing_current_state() { assert!(matches!(result, Err(miyagi::Error::SearchCancelled))); assert!(backend.flipped.is_empty()); } + +/// Baseline NaN is a broken model/probe — search must hard-fail before iterating. +#[test] +fn non_finite_baseline_aborts_search() { + let mut backend = AlwaysNanBackend::new(); + let result = run_search( + &mut backend, + &[compiled("target", "target")], + &[compiled("control", "control")], + SearchConfig { + search_layers: vec![0], + search_projections: vec![Projection::Gate], + max_iters: 3, + ..SearchConfig::default() + }, + None, + None, + None, + |_| {}, + ); + assert!( + matches!(result, Err(miyagi::Error::MeasurementMismatch(_))), + "expected baseline MeasurementMismatch, got {result:?}" + ); +} + +/// Candidate NaN after a flip is rejectable noise — search continues and reverts. +#[test] +fn non_finite_candidate_is_rejected_not_fatal() { + let mut backend = NanAfterFlipBackend::new(); + let result = run_search( + &mut backend, + &[compiled("target", "target")], + &[compiled("control", "control")], + SearchConfig { + search_layers: vec![0], + search_projections: vec![Projection::Gate], + max_iters: 3, + screen_probe_count: 1, + patch_name: "nan-cand".to_owned(), + base_model: "fake".to_owned(), + ..SearchConfig::default() + }, + None, + None, + None, + |_| {}, + ) + .expect("NaN candidates must not abort the whole search"); + assert!( + result.patch.flips.is_empty(), + "NaN fitness cannot pass strict improvement" + ); + assert!( + backend.flipped.is_empty(), + "rejected NaN candidates must be reverted" + ); + assert_eq!(result.completed_iterations, 3); +} + +/// Returns NaN for every probe gap (broken baseline). +struct AlwaysNanBackend { + architecture: ArchitectureMap, +} + +impl AlwaysNanBackend { + fn new() -> Self { + Self { + architecture: FakeBackend::new().architecture, + } + } +} + +impl MiyagiBackend for AlwaysNanBackend { + fn architecture(&self) -> &ArchitectureMap { + &self.architecture + } + + fn model_label(&self) -> &str { + "always-nan" + } + + fn tokenize(&self, _text: &str) -> Result> { + Ok(vec![1]) + } + + fn row_scales(&mut self, _layer: usize, _projection: Projection) -> Result> { + Ok(vec![1.0, 1.0, 1.0]) + } + + fn flip_row(&mut self, _layer: usize, _projection: Projection, _row: usize) -> Result<()> { + Ok(()) + } + + fn logit_gap(&mut self, _prompt: &[i32], _correct: i32, _wrong: i32) -> Result { + Ok(f32::NAN) + } + + fn generate(&mut self, _prompt: &str, _config: &GenerateConfig) -> Result { + Ok(String::new()) + } +} + +/// Finite baseline; any flipped row yields NaN (candidate-only defect). +struct NanAfterFlipBackend { + architecture: ArchitectureMap, + flipped: BTreeSet<(usize, Projection, usize)>, +} + +impl NanAfterFlipBackend { + fn new() -> Self { + Self { + architecture: FakeBackend::new().architecture, + flipped: BTreeSet::new(), + } + } +} + +impl MiyagiBackend for NanAfterFlipBackend { + fn architecture(&self) -> &ArchitectureMap { + &self.architecture + } + + fn model_label(&self) -> &str { + "nan-after-flip" + } + + fn tokenize(&self, text: &str) -> Result> { + Ok(vec![if text == "target" { 1 } else { 0 }]) + } + + fn row_scales(&mut self, _layer: usize, _projection: Projection) -> Result> { + Ok(vec![1.0, 1.0, 1.0]) + } + + fn flip_row(&mut self, layer: usize, projection: Projection, row: usize) -> Result<()> { + let key = (layer, projection, row); + if !self.flipped.insert(key) { + self.flipped.remove(&key); + } + Ok(()) + } + + fn logit_gap(&mut self, prompt: &[i32], _correct: i32, _wrong: i32) -> Result { + if !self.flipped.is_empty() { + return Ok(f32::NAN); + } + // Finite baseline: target wrong, control right. + if prompt.first() == Some(&1) { + Ok(-1.0) + } else { + Ok(1.0) + } + } + + fn generate(&mut self, _prompt: &str, _config: &GenerateConfig) -> Result { + Ok(String::new()) + } +}