#![cfg(feature = "burn")]
use burn_tensor::{Tensor, TensorData};
fn device() -> burn_tensor::Device {
burn_tensor::Device::default().autodiff()
}
fn sum_f32_in_backend(t: Tensor<2>) -> f64 {
let v: Vec<f32> = t.sum().to_data().to_vec().unwrap();
v[0] as f64
}
fn sum_f64_host(t: Tensor<2>) -> f64 {
let v: Vec<f32> = t.to_data().to_vec().unwrap();
v.into_iter().map(|x| x as f64).sum()
}
fn sum_f64_kahan(t: Tensor<2>) -> f64 {
let v: Vec<f32> = t.to_data().to_vec().unwrap();
let (mut s, mut c) = (0.0f64, 0.0f64);
for x in v {
let y = x as f64 - c;
let t = s + y;
c = (t - s) - y;
s = t;
}
s
}
fn build(data: &[f64], shape: [usize; 2]) -> Tensor<2> {
let v: Vec<f32> = data.iter().map(|&x| x as f32).collect();
Tensor::from_data(TensorData::new(v, shape.to_vec()), &device())
}
fn analytic(f: &dyn Fn(Tensor<2>) -> Tensor<2>, data: &[f64], shape: [usize; 2]) -> Vec<f64> {
let x = build(data, shape).require_grad();
let g = f(x.clone()).sum().backward();
x.grad(&g)
.unwrap()
.to_data()
.to_vec::<f32>()
.unwrap()
.into_iter()
.map(|v| v as f64)
.collect()
}
fn numeric(
f: &dyn Fn(Tensor<2>) -> Tensor<2>,
reduce: &dyn Fn(Tensor<2>) -> f64,
data: &[f64],
shape: [usize; 2],
h: f64,
) -> Vec<f64> {
let mut probe = data.to_vec();
let mut out = Vec::with_capacity(data.len());
for i in 0..data.len() {
let orig = probe[i];
probe[i] = orig + h;
let up = reduce(f(build(&probe, shape)));
probe[i] = orig - h;
let down = reduce(f(build(&probe, shape)));
probe[i] = orig;
out.push((up - down) / (2.0 * h));
}
out
}
fn worst_rel(a: &[f64], b: &[f64]) -> f64 {
a.iter()
.zip(b)
.map(|(x, y)| (x - y).abs() / x.abs().max(y.abs()).max(1.0))
.fold(0.0f64, f64::max)
}
#[test]
fn finite_difference_ceiling_by_reduction_precision() {
let n = 24usize;
let shape = [4usize, 6];
let data: Vec<f64> = (0..n)
.map(|i| ((i as f64) * 0.613).sin() * 1.3 + ((i as f64) * 0.271).cos() * 0.7 + 0.11)
.collect();
type Case<'a> = (&'a str, Box<dyn Fn(Tensor<2>) -> Tensor<2>>);
type Reducer<'a> = (&'a str, Box<dyn Fn(Tensor<2>) -> f64>);
let cases: Vec<Case> = vec![
("tanh", Box::new(|t: Tensor<2>| t.tanh())),
("exp", Box::new(|t: Tensor<2>| t.exp())),
("sin", Box::new(|t: Tensor<2>| t.sin())),
("square", Box::new(|t: Tensor<2>| t.clone() * t)),
];
let reducers: Vec<Reducer> = vec![
("f32_in_backend", Box::new(sum_f32_in_backend)),
("f64_host", Box::new(sum_f64_host)),
("f64_kahan", Box::new(sum_f64_kahan)),
];
println!(
"\n{:<10} {:<16} {:>10} {:>14}",
"case", "reduction", "h", "worst_rel"
);
println!("{}", "-".repeat(54));
for (cname, f) in &cases {
let a = analytic(f.as_ref(), &data, shape);
for (rname, red) in &reducers {
for h in [1e-1f64, 1e-2, 1e-3, 1e-4] {
let num = numeric(f.as_ref(), red.as_ref(), &data, shape, h);
println!(
"{cname:<10} {rname:<16} {h:>10.0e} {:>14.3e}",
worst_rel(&a, &num)
);
}
}
println!();
}
}