#![cfg(feature = "burn")]
use burn_tensor::module::avg_pool1d;
use burn_tensor::{Tensor, TensorData};
const SENTINEL: f32 = 1234.5;
fn device() -> burn_tensor::Device {
burn_tensor::Device::default().autodiff()
}
fn poison() -> f32 {
let big = vec![SENTINEL; 8192];
let t = Tensor::<2>::from_data(TensorData::new(big, [64, 128]), &device());
let s: Vec<f32> = (t.clone() * t).sum().to_data().to_vec().unwrap();
s[0]
}
fn pool_grad(c: usize) -> Vec<f32> {
let len = 6usize;
let n = c * len;
let d: Vec<f32> = (0..n).map(|i| (i as f32) * 0.1 + 0.5).collect();
let x = Tensor::<3>::from_data(TensorData::new(d, [1, c, len]), &device()).require_grad();
let g = avg_pool1d(x.clone(), 3, 2, 1, true, false).sum().backward();
x.grad(&g).unwrap().to_data().to_vec::<f32>().unwrap()
}
#[test]
fn published_burn_pooling_does_not_leak_uninitialised_memory() {
let expected = [
1.0 / 3.0,
2.0 / 3.0,
1.0 / 3.0,
2.0 / 3.0,
1.0 / 3.0,
1.0 / 3.0,
];
let mut leaked = vec![];
let mut wrong = vec![];
for c in [1usize, 2, 3, 4] {
let _ = poison();
let g = pool_grad(c);
let has_sentinel = g.iter().any(|v| (v - SENTINEL).abs() < 1e-3);
let matches_truth = g
.iter()
.enumerate()
.all(|(i, v)| (v - expected[i % 6]).abs() < 1e-4);
println!(
"BOUNDARY c={c} sentinel_leaked={has_sentinel} correct={matches_truth} grad={:?}",
&g[..g.len().min(6)]
);
if has_sentinel {
leaked.push(c);
}
if !matches_truth {
wrong.push(c);
}
}
println!("BOUNDARY summary leaked_at={leaked:?} wrong_at={wrong:?}");
assert!(
leaked.is_empty(),
"uninitialised memory leaked into the gradient at channel counts {leaked:?}"
);
assert!(
wrong.is_empty(),
"gradient incorrect at channel counts {wrong:?}"
);
}