#![cfg(feature = "burn")]
use gradcheck::{
adapters::burn::Burn, assert_detects_wrong_gradient, gradcheck, shape_sweep, Config,
};
fn cfg() -> Config {
Config::f32_defaults()
}
fn ramp(n: usize, seed: f64) -> Vec<f64> {
(0..n)
.map(|i| ((i as f64) * 0.37 + seed).sin() * 1.7 + 0.15)
.collect()
}
fn matmul_cfg() -> Config {
Config::f32_defaults().with_step(1e-1)
}
#[test]
fn negative_control_detects_a_wrong_gradient() {
assert_detects_wrong_gradient::<Burn<2>>();
}
#[test]
fn elementwise_gradients_are_correct() {
let d = ramp(6, 0.3);
let s = [2usize, 3];
let c = cfg();
gradcheck::<Burn<2>, _>("tanh", &d, &s, |t| t.tanh(), &c).assert_pass();
gradcheck::<Burn<2>, _>("exp", &d, &s, |t| t.exp(), &c).assert_pass();
gradcheck::<Burn<2>, _>("sin", &d, &s, |t| t.sin(), &c).assert_pass();
gradcheck::<Burn<2>, _>("erf", &d, &s, |t| t.erf(), &c).assert_pass();
gradcheck::<Burn<2>, _>("square", &d, &s, |t| t.clone() * t, &c).assert_pass();
gradcheck::<Burn<2>, _>("log_positive", &d, &s, |t| (t.abs() + 0.5).log(), &c).assert_pass();
gradcheck::<Burn<2>, _>("sqrt_positive", &d, &s, |t| (t.abs() + 0.5).sqrt(), &c).assert_pass();
gradcheck::<Burn<2>, _>("chained", &d, &s, |t| t.tanh().exp(), &c).assert_pass();
}
#[test]
fn reduction_gradients_are_correct() {
let d = ramp(6, 0.3);
let s = [2usize, 3];
let c = cfg();
gradcheck::<Burn<2>, _>("sum_dim0", &d, &s, |t| t.sum_dim(0), &c).assert_pass();
gradcheck::<Burn<2>, _>("sum_dim1", &d, &s, |t| t.sum_dim(1), &c).assert_pass();
gradcheck::<Burn<2>, _>("mean_dim1", &d, &s, |t| t.mean_dim(1), &c).assert_pass();
gradcheck::<Burn<2>, _>("cumsum", &d, &s, |t| t.cumsum(1), &c).assert_pass();
}
#[test]
fn broadcast_and_view_gradients_are_correct() {
let d = ramp(6, 0.3);
let s = [2usize, 3];
let c = cfg();
gradcheck::<Burn<2>, _>(
"broadcast_by_colsum",
&d,
&s,
|t| {
let col = t.clone().sum_dim(1);
t * col
},
&c,
)
.assert_pass();
gradcheck::<Burn<2>, _>(
"transpose_roundtrip",
&d,
&s,
|t| t.transpose().transpose(),
&c,
)
.assert_pass();
gradcheck::<Burn<2>, _>("slice", &d, &s, |t| t.slice([0..2, 0..2]), &c).assert_no_mismatch();
}
#[test]
fn transpose_matmul_gradient_across_shapes() {
let c = matmul_cfg();
let shapes: Vec<Vec<usize>> = vec![
vec![2, 3],
vec![3, 2],
vec![2, 2],
vec![4, 4],
vec![4, 3],
vec![6, 3],
vec![8, 2],
vec![2, 5],
vec![5, 5],
vec![3, 8],
];
let sweep = shape_sweep::<Burn<2>, _, _>(
"self_transpose_matmul",
&shapes,
|n| ramp(n, 0.3),
|t| t.clone().transpose().matmul(t),
&c,
);
for r in &sweep.reports {
println!("{r}");
}
sweep.assert_all_pass();
}
#[test]
fn matmul_self_transpose_gradient_across_shapes() {
let c = matmul_cfg();
let shapes: Vec<Vec<usize>> = vec![
vec![2, 3],
vec![3, 2],
vec![4, 4],
vec![2, 8],
vec![8, 2],
vec![5, 7],
];
shape_sweep::<Burn<2>, _, _>(
"matmul_self_transpose",
&shapes,
|n| ramp(n, 0.3),
|t| t.clone().matmul(t.transpose()),
&c,
)
.assert_all_pass();
}