#![cfg(feature = "burn")]
use burn_tensor::module::avg_pool1d;
use gradcheck::{adapters::burn::Burn, gradcheck, Config, Verdict};
#[test]
fn v4_still_rejects_the_pooling_defect() {
let cfg = Config::f32_defaults();
let mut rejected = vec![];
let mut certified = vec![];
let mut abstained = vec![];
for c in [1usize, 2, 3, 4] {
let len = 6usize;
let n = c * len;
let data: Vec<f64> = (0..n).map(|i| (i as f64) * 0.1 + 0.5).collect();
let r = gradcheck::<Burn<3>, _>(
&format!("avg_pool1d_c{c}"),
&data,
&[1, c, len],
|t| avg_pool1d(t, 3, 2, 1, true, false),
&cfg,
);
let (checked, total) = r.checked_fraction();
println!(
"FDPC c={c} {:?} checked={checked}/{total} worst_rel={:.3e}",
r.verdict, r.worst_rel_error
);
match r.verdict {
Verdict::Mismatch => rejected.push(c),
Verdict::Pass => certified.push(c),
_ => abstained.push(c),
}
}
println!("FDPC rejected={rejected:?} certified={certified:?} abstained={abstained:?}");
assert!(
abstained.is_empty(),
"no case should abstain at these magnitudes; abstained at {abstained:?}"
);
if cfg!(feature = "burn-cpu") {
assert!(
rejected.contains(&2) && rejected.contains(&4),
"DETECTION REGRESSION: v4 failed to reject the known pooling defect at c=2,4. \
rejected={rejected:?} certified={certified:?}"
);
assert!(
certified.contains(&1) && certified.contains(&3),
"odd channel counts are correct and must certify; certified={certified:?}"
);
} else {
assert_eq!(
rejected,
Vec::<usize>::new(),
"this backend is known-correct here; nothing should be rejected"
);
assert_eq!(certified.len(), 4, "all four channel counts should certify");
}
}