#[cfg(feature = "cuda")]
fn parity_probe_enabled() -> bool {
static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*ON.get_or_init(|| std::env::var("APR_PARITY_PROBE").as_deref() == Ok("1"))
}
#[cfg(feature = "cuda")]
struct LogitsDelta {
argmax_a: usize,
argmax_b: usize,
max_abs: f32,
cosine: f32,
nonfinite_a: usize,
nonfinite_b: usize,
zeros_a: usize,
zeros_b: usize,
}
#[cfg(feature = "cuda")]
fn compare_logits(a: &[f32], b: &[f32]) -> LogitsDelta {
let argmax = |v: &[f32]| {
v.iter()
.enumerate()
.max_by(|(_, x), (_, y)| x.partial_cmp(y).unwrap_or(std::cmp::Ordering::Equal))
.map_or(0, |(i, _)| i)
};
let mut max_abs = 0.0f32;
let (mut dot, mut na, mut nb) = (0.0f64, 0.0f64, 0.0f64);
for (&x, &y) in a.iter().zip(b.iter()) {
max_abs = max_abs.max((x - y).abs());
dot += f64::from(x) * f64::from(y);
na += f64::from(x) * f64::from(x);
nb += f64::from(y) * f64::from(y);
}
let denom = (na.sqrt() * nb.sqrt()).max(f64::MIN_POSITIVE);
LogitsDelta {
argmax_a: argmax(a),
argmax_b: argmax(b),
max_abs,
cosine: (dot / denom) as f32,
nonfinite_a: a.iter().filter(|x| !x.is_finite()).count(),
nonfinite_b: b.iter().filter(|x| !x.is_finite()).count(),
zeros_a: a.iter().filter(|x| **x == 0.0).count(),
zeros_b: b.iter().filter(|x| **x == 0.0).count(),
}
}
#[cfg(feature = "cuda")]
fn report(tag: &str, d: &LogitsDelta, expect: &str) {
eprintln!(
"[CB-006-PARITY] {tag:<9} argmax {:>6} vs {:>6} | cosine {:.6} | max|d| {:.4} \
| nonfinite {}/{} | zeros {}/{} | expect {expect}",
d.argmax_a, d.argmax_b, d.cosine, d.max_abs, d.nonfinite_a, d.nonfinite_b,
d.zeros_a, d.zeros_b
);
}
#[cfg(feature = "cuda")]
impl OwnedQuantizedModelCuda {
pub(crate) fn cb006_parity_probe(&mut self, state: &BatchedDecodeState) {
if !parity_probe_enabled() || state.m != 1 || state.gen_idx != 0 {
return;
}
let (nl, hd, id, vs, eps) = (
state.num_layers,
state.hidden_dim as u32,
state.intermediate_dim as u32,
state.vocab_size as u32,
state.eps,
);
let position = state.positions[0] as u32;
if state.embed_buf.iter().all(|v| *v == 0.0) {
eprintln!(
"[CB-006-PARITY] REFUSING TO REPORT: embed_buf is all zeros, so both paths would \
be fed a degenerate input. The probe is being called before the embedding is \
written. Nothing below would be a measurement of the batched decode."
);
return;
}
eprintln!(
"[CB-006-PARITY] m=1 step 0: token {} at position {position}, \
batched_kv_lengths[0]={:?}",
state.last_tokens[0],
self.executor.batched_kv_lengths().first().copied()
);
self.executor.set_batched_done_mask(&state.done);
let batched = match self
.executor
.forward_batched_to_logits(&state.embed_buf, &state.pos_buf, nl, hd, id, vs, eps)
{
Ok(v) => v,
Err(e) => {
eprintln!("[CB-006-PARITY] batched forward failed: {e}");
return;
},
};
match self
.executor
.forward_batched_to_logits(&state.embed_buf, &state.pos_buf, nl, hd, id, vs, eps)
{
Ok(again) => report("SELF", &compare_logits(&batched, &again), "IDENTICAL"),
Err(e) => eprintln!("[CB-006-PARITY] batched re-run failed: {e}"),
}
let mut sens_in = state.embed_buf.clone();
for v in sens_in.iter_mut() {
*v = -*v;
}
match self
.executor
.forward_batched_to_logits(&sens_in, &state.pos_buf, nl, hd, id, vs, eps)
{
Ok(flipped) => report(
"SENSITIVE",
&compare_logits(&batched, &flipped),
"DIVERGE(else the batched forward ignores its input)",
),
Err(e) => eprintln!("[CB-006-PARITY] sensitivity forward failed: {e}"),
}
let mut oracle = vec![0.0f32; state.vocab_size];
let saved = self.executor.batched_kv_stride;
self.executor.batched_kv_stride = 0;
let ok = self.executor.forward_all_layers_gpu_to_logits(
&state.embed_buf, &mut oracle, position, nl, hd, id, vs, eps,
);
self.executor.batched_kv_stride = saved;
match ok {
Ok(()) => report("ORACLE", &compare_logits(&batched, &oracle), "IDENTICAL"),
Err(e) => eprintln!("[CB-006-PARITY] oracle forward failed: {e}"),
}
let mut perturbed_in = state.embed_buf.clone();
perturbed_in[0] += 10.0;
let mut perturbed = vec![0.0f32; state.vocab_size];
let saved = self.executor.batched_kv_stride;
self.executor.batched_kv_stride = 0;
let ok = self.executor.forward_all_layers_gpu_to_logits(
&perturbed_in, &mut perturbed, position, nl, hd, id, vs, eps,
);
self.executor.batched_kv_stride = saved;
match ok {
Ok(()) => report("PERTURBED", &compare_logits(&batched, &perturbed), "DIVERGE"),
Err(e) => eprintln!("[CB-006-PARITY] perturbed forward failed: {e}"),
}
if let Err(e) = self.executor.init_batched_workspace(
state.hidden_dim,
state.intermediate_dim,
state.m,
) {
eprintln!("[CB-006-PARITY] workspace restore failed: {e}");
}
self.executor.clear_decode_graph();
}
}