use approx::assert_abs_diff_eq;
use ndarray::Array;
use rustyml::neural_network::Tensor;
use rustyml::neural_network::layers::activation::linear::Linear;
use rustyml::neural_network::layers::activation::softmax::Softmax;
use rustyml::neural_network::layers::activation::tanh::Tanh;
use rustyml::neural_network::layers::convolution::PaddingType;
use rustyml::neural_network::layers::convolution::conv_1d::Conv1D;
use rustyml::neural_network::layers::convolution::conv_2d::Conv2D;
use rustyml::neural_network::layers::convolution::conv_3d::Conv3D;
use rustyml::neural_network::layers::convolution::depthwise_conv_2d::DepthwiseConv2D;
use rustyml::neural_network::layers::convolution::separable_conv_2d::SeparableConv2D;
use rustyml::neural_network::layers::dense::Dense;
use rustyml::neural_network::layers::pooling::average_pooling_1d::AveragePooling1D;
use rustyml::neural_network::layers::pooling::average_pooling_2d::AveragePooling2D;
use rustyml::neural_network::layers::pooling::average_pooling_3d::AveragePooling3D;
use rustyml::neural_network::layers::pooling::global_average_pooling_1d::GlobalAveragePooling1D;
use rustyml::neural_network::layers::pooling::global_average_pooling_2d::GlobalAveragePooling2D;
use rustyml::neural_network::layers::pooling::global_average_pooling_3d::GlobalAveragePooling3D;
use rustyml::neural_network::layers::pooling::global_max_pooling_1d::GlobalMaxPooling1D;
use rustyml::neural_network::layers::pooling::global_max_pooling_2d::GlobalMaxPooling2D;
use rustyml::neural_network::layers::pooling::global_max_pooling_3d::GlobalMaxPooling3D;
use rustyml::neural_network::layers::pooling::max_pooling_1d::MaxPooling1D;
use rustyml::neural_network::layers::pooling::max_pooling_2d::MaxPooling2D;
use rustyml::neural_network::layers::pooling::max_pooling_3d::MaxPooling3D;
use rustyml::neural_network::layers::recurrent::gru::GRU;
use rustyml::neural_network::layers::recurrent::lstm::LSTM;
use rustyml::neural_network::layers::recurrent::simple_rnn::SimpleRNN;
use rustyml::neural_network::layers::regularization::normalization::batch_normalization::BatchNormalization;
use rustyml::neural_network::layers::regularization::normalization::group_normalization::GroupNormalization;
use rustyml::neural_network::layers::regularization::normalization::instance_normalization::InstanceNormalization;
use rustyml::neural_network::layers::regularization::normalization::layer_normalization::{
LayerNormalization, LayerNormalizationAxis,
};
use rustyml::neural_network::traits::Layer;
fn check_input_gradient(layer: &mut dyn Layer, x: &Tensor, eps: f32, tol: f32) {
let out = layer.forward(x).unwrap();
let upstream = Tensor::ones(out.raw_dim());
let analytic = layer.backward(&upstream).unwrap();
assert_eq!(
analytic.shape(),
x.shape(),
"input-gradient shape must match input shape"
);
let analytic_flat: Vec<f32> = analytic.iter().cloned().collect();
let mut x_flat: Vec<f32> = x.iter().cloned().collect();
for i in 0..x_flat.len() {
let orig = x_flat[i];
x_flat[i] = orig + eps;
let xp = Tensor::from_shape_vec(x.raw_dim(), x_flat.clone()).unwrap();
let l_plus: f32 = layer.forward(&xp).unwrap().sum();
x_flat[i] = orig - eps;
let xm = Tensor::from_shape_vec(x.raw_dim(), x_flat.clone()).unwrap();
let l_minus: f32 = layer.forward(&xm).unwrap().sum();
x_flat[i] = orig;
let numeric = (l_plus - l_minus) / (2.0 * eps);
assert_abs_diff_eq!(analytic_flat[i], numeric, epsilon = tol);
}
}
#[test]
fn dense_input_gradient_matches_finite_difference() {
let mut dense = Dense::new(3, 2, Linear::new()).unwrap();
let x = Array::from_shape_vec((4, 3), (0..12).map(|v| 0.1 * v as f32 - 0.5).collect())
.unwrap()
.into_dyn();
check_input_gradient(&mut dense, &x, 1e-3, 1e-2);
}
#[test]
fn conv2d_input_gradient_matches_finite_difference() {
let mut conv = Conv2D::new(2, (2, 2), vec![1, 1, 4, 4], (1, 1), Linear::new()).unwrap();
let x = Array::from_shape_vec(
(1, 1, 4, 4),
(0..16).map(|v| 0.1 * v as f32 - 0.7).collect(),
)
.unwrap()
.into_dyn();
check_input_gradient(&mut conv, &x, 1e-3, 1e-2);
}
#[test]
fn conv1d_input_gradient_matches_finite_difference() {
let mut conv = Conv1D::new(2, 2, vec![1, 1, 5], 1, Linear::new()).unwrap();
let x = Array::from_shape_vec((1, 1, 5), (0..5).map(|v| 0.1 * v as f32 - 0.3).collect())
.unwrap()
.into_dyn();
check_input_gradient(&mut conv, &x, 1e-3, 1e-2);
}
#[test]
fn conv3d_input_gradient_matches_finite_difference() {
let mut conv =
Conv3D::new(2, (2, 2, 2), vec![1, 1, 3, 3, 3], (1, 1, 1), Linear::new()).unwrap();
let x = Array::from_shape_vec(
(1, 1, 3, 3, 3),
(0..27).map(|v| 0.05 * v as f32 - 0.4).collect(),
)
.unwrap()
.into_dyn();
check_input_gradient(&mut conv, &x, 1e-3, 1e-2);
}
#[test]
fn separable_conv2d_input_gradient_matches_finite_difference() {
let mut conv =
SeparableConv2D::new(2, (2, 2), vec![1, 2, 4, 4], (1, 1), 1, Linear::new()).unwrap();
let x = Array::from_shape_vec(
(1, 2, 4, 4),
(0..32).map(|v| 0.05 * v as f32 - 0.7).collect(),
)
.unwrap()
.into_dyn();
check_input_gradient(&mut conv, &x, 1e-3, 2e-2);
}
#[test]
fn separable_conv2d_same_padding_input_gradient_matches_finite_difference() {
let mut conv = SeparableConv2D::new(2, (3, 3), vec![1, 2, 4, 4], (1, 1), 1, Linear::new())
.unwrap()
.with_padding(PaddingType::Same);
let x = Array::from_shape_vec(
(1, 2, 4, 4),
(0..32).map(|v| 0.05 * v as f32 - 0.7).collect(),
)
.unwrap()
.into_dyn();
check_input_gradient(&mut conv, &x, 1e-3, 2e-2);
}
#[test]
fn depthwise_conv2d_input_gradient_matches_finite_difference() {
let mut conv =
DepthwiseConv2D::new(2, (2, 2), vec![1, 2, 4, 4], (1, 1), Linear::new()).unwrap();
let x = Array::from_shape_vec(
(1, 2, 4, 4),
(0..32).map(|v| 0.05 * v as f32 - 0.7).collect(),
)
.unwrap()
.into_dyn();
check_input_gradient(&mut conv, &x, 1e-3, 2e-2);
}
#[test]
fn depthwise_conv2d_same_padding_input_gradient_matches_finite_difference() {
let mut conv = DepthwiseConv2D::new(2, (3, 3), vec![1, 2, 4, 4], (1, 1), Linear::new())
.unwrap()
.with_padding(PaddingType::Same);
let x = Array::from_shape_vec(
(1, 2, 4, 4),
(0..32).map(|v| 0.05 * v as f32 - 0.7).collect(),
)
.unwrap()
.into_dyn();
check_input_gradient(&mut conv, &x, 1e-3, 2e-2);
}
#[test]
fn depthwise_conv2d_same_padding_weight_gradient_matches_finite_difference() {
let mut conv = DepthwiseConv2D::new(2, (3, 3), vec![1, 2, 4, 4], (1, 1), Linear::new())
.unwrap()
.with_padding(PaddingType::Same);
let x = Array::from_shape_vec(
(1, 2, 4, 4),
(0..32).map(|v| 0.05 * v as f32 - 0.7).collect(),
)
.unwrap()
.into_dyn();
check_weight_gradient(&mut conv, &x, 1e-3, 2e-2);
}
#[test]
fn simple_rnn_input_gradient_matches_finite_difference() {
let mut rnn = SimpleRNN::new(2, 3, Tanh::new()).unwrap();
let x = Array::from_shape_vec((1, 3, 2), vec![0.3, -0.6, 0.9, -0.2, 0.5, -0.8])
.unwrap()
.into_dyn();
check_input_gradient(&mut rnn, &x, 1e-3, 2e-2);
}
#[test]
fn lstm_input_gradient_matches_finite_difference() {
let mut lstm = LSTM::new(2, 3, Tanh::new()).unwrap();
let x = Array::from_shape_vec((1, 3, 2), vec![0.3, -0.6, 0.9, -0.2, 0.5, -0.8])
.unwrap()
.into_dyn();
check_input_gradient(&mut lstm, &x, 1e-3, 3e-2);
}
#[test]
fn gru_input_gradient_matches_finite_difference() {
let mut gru = GRU::new(2, 3, Tanh::new()).unwrap();
let x = Array::from_shape_vec((1, 3, 2), vec![0.3, -0.6, 0.9, -0.2, 0.5, -0.8])
.unwrap()
.into_dyn();
check_input_gradient(&mut gru, &x, 1e-3, 3e-2);
}
#[test]
fn batch_normalization_input_gradient_matches_finite_difference() {
let mut bn = BatchNormalization::new(vec![4, 3], 0.9, 1e-5).unwrap();
let x = Array::from_shape_vec(
(4, 3),
vec![
0.5, -1.0, 2.0, 1.5, 0.2, -0.7, -1.2, 0.8, 1.1, 0.3, -0.4, 0.9,
],
)
.unwrap()
.into_dyn();
check_input_gradient(&mut bn, &x, 1e-3, 5e-2);
}
#[test]
fn conv1d_same_padding_output_length_is_ceil_of_input() {
let cases = [
(10usize, 3usize, 1usize, 10usize),
(10, 3, 2, 5),
(8, 5, 1, 8),
(7, 3, 2, 4),
];
for (len, kernel, stride, expected) in cases {
let mut conv = Conv1D::new(2, kernel, vec![1, 1, len], stride, Linear::new())
.unwrap()
.with_padding(PaddingType::Same);
let x = Array::ones((1, 1, len)).into_dyn();
let out = conv.forward(&x).unwrap();
assert_eq!(
out.shape(),
&[1, 2, expected],
"Conv1D Same: input_len={}, kernel={}, stride={}",
len,
kernel,
stride
);
}
}
fn check_weight_gradient(layer: &mut dyn Layer, x: &Tensor, eps: f32, tol: f32) {
let out = layer.forward(x).unwrap();
let upstream = Tensor::ones(out.raw_dim());
layer.backward(&upstream).unwrap();
let params: Vec<(Vec<f32>, Vec<f32>)> = layer
.parameters()
.into_iter()
.map(|pg| (pg.value.to_vec(), pg.grad.to_vec()))
.collect();
assert!(!params.is_empty(), "layer exposes no parameters to check");
for (p_idx, (values, grads)) in params.iter().enumerate() {
for i in 0..values.len() {
let orig = values[i];
layer.parameters()[p_idx].value[i] = orig + eps;
let l_plus: f32 = layer.forward(x).unwrap().sum();
layer.parameters()[p_idx].value[i] = orig - eps;
let l_minus: f32 = layer.forward(x).unwrap().sum();
layer.parameters()[p_idx].value[i] = orig;
let numeric = (l_plus - l_minus) / (2.0 * eps);
assert_abs_diff_eq!(grads[i], numeric, epsilon = tol);
}
}
}
#[test]
fn dense_weight_gradient_matches_finite_difference() {
let mut dense = Dense::new(3, 2, Linear::new()).unwrap();
let x = Array::from_shape_vec((4, 3), (0..12).map(|v| 0.1 * v as f32 - 0.5).collect())
.unwrap()
.into_dyn();
check_weight_gradient(&mut dense, &x, 1e-3, 1e-2);
}
#[test]
fn conv1d_weight_gradient_matches_finite_difference() {
let mut conv = Conv1D::new(2, 2, vec![1, 1, 5], 1, Linear::new()).unwrap();
let x = Array::from_shape_vec((1, 1, 5), (0..5).map(|v| 0.1 * v as f32 - 0.3).collect())
.unwrap()
.into_dyn();
check_weight_gradient(&mut conv, &x, 1e-3, 1e-2);
}
#[test]
fn conv2d_weight_gradient_matches_finite_difference() {
let mut conv = Conv2D::new(2, (2, 2), vec![1, 1, 4, 4], (1, 1), Linear::new()).unwrap();
let x = Array::from_shape_vec(
(1, 1, 4, 4),
(0..16).map(|v| 0.1 * v as f32 - 0.7).collect(),
)
.unwrap()
.into_dyn();
check_weight_gradient(&mut conv, &x, 1e-3, 1e-2);
}
#[test]
fn conv3d_weight_gradient_matches_finite_difference() {
let mut conv =
Conv3D::new(2, (2, 2, 2), vec![1, 1, 3, 3, 3], (1, 1, 1), Linear::new()).unwrap();
let x = Array::from_shape_vec(
(1, 1, 3, 3, 3),
(0..27).map(|v| 0.05 * v as f32 - 0.4).collect(),
)
.unwrap()
.into_dyn();
check_weight_gradient(&mut conv, &x, 1e-3, 1e-2);
}
#[test]
fn separable_conv2d_weight_gradient_matches_finite_difference() {
let mut conv =
SeparableConv2D::new(2, (2, 2), vec![1, 2, 4, 4], (1, 1), 1, Linear::new()).unwrap();
let x = Array::from_shape_vec(
(1, 2, 4, 4),
(0..32).map(|v| 0.05 * v as f32 - 0.7).collect(),
)
.unwrap()
.into_dyn();
check_weight_gradient(&mut conv, &x, 1e-3, 2e-2);
}
#[test]
fn separable_conv2d_same_padding_weight_gradient_matches_finite_difference() {
let mut conv = SeparableConv2D::new(2, (3, 3), vec![1, 2, 4, 4], (1, 1), 1, Linear::new())
.unwrap()
.with_padding(PaddingType::Same);
let x = Array::from_shape_vec(
(1, 2, 4, 4),
(0..32).map(|v| 0.05 * v as f32 - 0.7).collect(),
)
.unwrap()
.into_dyn();
check_weight_gradient(&mut conv, &x, 1e-3, 2e-2);
}
#[test]
fn depthwise_conv2d_weight_gradient_matches_finite_difference() {
let mut conv =
DepthwiseConv2D::new(2, (2, 2), vec![1, 2, 4, 4], (1, 1), Linear::new()).unwrap();
let x = Array::from_shape_vec(
(1, 2, 4, 4),
(0..32).map(|v| 0.05 * v as f32 - 0.7).collect(),
)
.unwrap()
.into_dyn();
check_weight_gradient(&mut conv, &x, 1e-3, 2e-2);
}
fn loss_weights(like: &Tensor) -> Tensor {
let n = like.len();
let flat: Vec<f32> = (0..n).map(|k| 1.0 + 0.1 * ((k % 7) as f32 - 3.0)).collect();
Tensor::from_shape_vec(like.raw_dim(), flat).unwrap()
}
fn ramp(shape: &[usize]) -> Tensor {
let n: usize = shape.iter().product();
let data: Vec<f32> = (0..n).map(|v| 0.5 * v as f32 - 0.25 * n as f32).collect();
Array::from_shape_vec(shape.to_vec(), data).unwrap()
}
fn check_input_gradient_weighted(layer: &mut dyn Layer, x: &Tensor, eps: f32, tol: f32) {
let out = layer.forward(x).unwrap();
let w = loss_weights(&out);
let analytic = layer.backward(&w).unwrap();
assert_eq!(
analytic.shape(),
x.shape(),
"input-gradient shape must match input shape"
);
let analytic_flat: Vec<f32> = analytic.iter().cloned().collect();
let mut x_flat: Vec<f32> = x.iter().cloned().collect();
for i in 0..x_flat.len() {
let orig = x_flat[i];
x_flat[i] = orig + eps;
let xp = Tensor::from_shape_vec(x.raw_dim(), x_flat.clone()).unwrap();
let l_plus: f32 = (&layer.forward(&xp).unwrap() * &w).sum();
x_flat[i] = orig - eps;
let xm = Tensor::from_shape_vec(x.raw_dim(), x_flat.clone()).unwrap();
let l_minus: f32 = (&layer.forward(&xm).unwrap() * &w).sum();
x_flat[i] = orig;
let numeric = (l_plus - l_minus) / (2.0 * eps);
assert_abs_diff_eq!(analytic_flat[i], numeric, epsilon = tol);
}
}
#[test]
fn softmax_input_gradient_matches_finite_difference() {
let mut softmax = Softmax::new();
let x = Array::from_shape_vec((2, 3), vec![0.2, -0.5, 1.0, 0.7, 0.1, -0.3])
.unwrap()
.into_dyn();
check_input_gradient_weighted(&mut softmax, &x, 1e-3, 2e-2);
}
#[test]
fn max_pooling_1d_input_gradient_matches_finite_difference() {
let mut pool = MaxPooling1D::new(2, vec![1, 2, 6]).unwrap();
let x = ramp(&[1, 2, 6]);
check_input_gradient_weighted(&mut pool, &x, 1e-3, 1e-2);
}
#[test]
fn max_pooling_2d_input_gradient_matches_finite_difference() {
let mut pool = MaxPooling2D::new((2, 2), vec![1, 2, 4, 4]).unwrap();
let x = ramp(&[1, 2, 4, 4]);
check_input_gradient_weighted(&mut pool, &x, 1e-3, 1e-2);
}
#[test]
fn max_pooling_3d_input_gradient_matches_finite_difference() {
let mut pool = MaxPooling3D::new((2, 2, 2), vec![1, 1, 4, 4, 4]).unwrap();
let x = ramp(&[1, 1, 4, 4, 4]);
check_input_gradient_weighted(&mut pool, &x, 1e-3, 1e-2);
}
#[test]
fn average_pooling_1d_input_gradient_matches_finite_difference() {
let mut pool = AveragePooling1D::new(2, vec![1, 2, 6]).unwrap();
let x = ramp(&[1, 2, 6]);
check_input_gradient_weighted(&mut pool, &x, 1e-3, 1e-2);
}
#[test]
fn average_pooling_2d_input_gradient_matches_finite_difference() {
let mut pool = AveragePooling2D::new((2, 2), vec![1, 2, 4, 4]).unwrap();
let x = ramp(&[1, 2, 4, 4]);
check_input_gradient_weighted(&mut pool, &x, 1e-3, 1e-2);
}
#[test]
fn average_pooling_3d_input_gradient_matches_finite_difference() {
let mut pool = AveragePooling3D::new((2, 2, 2), vec![1, 1, 4, 4, 4]).unwrap();
let x = ramp(&[1, 1, 4, 4, 4]);
check_input_gradient_weighted(&mut pool, &x, 1e-3, 1e-2);
}
#[test]
fn global_max_pooling_1d_input_gradient_matches_finite_difference() {
let mut pool = GlobalMaxPooling1D::new();
let x = ramp(&[1, 2, 5]);
check_input_gradient_weighted(&mut pool, &x, 1e-3, 1e-2);
}
#[test]
fn global_max_pooling_2d_input_gradient_matches_finite_difference() {
let mut pool = GlobalMaxPooling2D::new();
let x = ramp(&[1, 2, 3, 3]);
check_input_gradient_weighted(&mut pool, &x, 1e-3, 1e-2);
}
#[test]
fn global_max_pooling_3d_input_gradient_matches_finite_difference() {
let mut pool = GlobalMaxPooling3D::new();
let x = ramp(&[1, 2, 2, 2, 2]);
check_input_gradient_weighted(&mut pool, &x, 1e-3, 1e-2);
}
#[test]
fn global_average_pooling_1d_input_gradient_matches_finite_difference() {
let mut pool = GlobalAveragePooling1D::new();
let x = ramp(&[1, 2, 5]);
check_input_gradient_weighted(&mut pool, &x, 1e-3, 1e-2);
}
#[test]
fn global_average_pooling_2d_input_gradient_matches_finite_difference() {
let mut pool = GlobalAveragePooling2D::new();
let x = ramp(&[1, 2, 3, 3]);
check_input_gradient_weighted(&mut pool, &x, 1e-3, 1e-2);
}
#[test]
fn global_average_pooling_3d_input_gradient_matches_finite_difference() {
let mut pool = GlobalAveragePooling3D::new();
let x = ramp(&[1, 2, 2, 2, 2]);
check_input_gradient_weighted(&mut pool, &x, 1e-3, 1e-2);
}
#[test]
fn conv1d_same_padding_input_gradient_matches_finite_difference() {
let mut conv = Conv1D::new(2, 3, vec![1, 1, 6], 1, Linear::new())
.unwrap()
.with_padding(PaddingType::Same);
let x = Array::from_shape_vec((1, 1, 6), (0..6).map(|v| 0.1 * v as f32 - 0.3).collect())
.unwrap()
.into_dyn();
check_input_gradient(&mut conv, &x, 1e-3, 1e-2);
}
#[test]
fn conv1d_same_padding_weight_gradient_matches_finite_difference() {
let mut conv = Conv1D::new(2, 3, vec![1, 1, 6], 1, Linear::new())
.unwrap()
.with_padding(PaddingType::Same);
let x = Array::from_shape_vec((1, 1, 6), (0..6).map(|v| 0.1 * v as f32 - 0.3).collect())
.unwrap()
.into_dyn();
check_weight_gradient(&mut conv, &x, 1e-3, 1e-2);
}
#[test]
fn conv2d_same_padding_input_gradient_matches_finite_difference() {
let mut conv = Conv2D::new(2, (3, 3), vec![1, 1, 5, 5], (1, 1), Linear::new())
.unwrap()
.with_padding(PaddingType::Same);
let x = Array::from_shape_vec(
(1, 1, 5, 5),
(0..25).map(|v| 0.05 * v as f32 - 0.6).collect(),
)
.unwrap()
.into_dyn();
check_input_gradient(&mut conv, &x, 1e-3, 1e-2);
}
#[test]
fn conv2d_same_padding_weight_gradient_matches_finite_difference() {
let mut conv = Conv2D::new(2, (3, 3), vec![1, 1, 5, 5], (1, 1), Linear::new())
.unwrap()
.with_padding(PaddingType::Same);
let x = Array::from_shape_vec(
(1, 1, 5, 5),
(0..25).map(|v| 0.05 * v as f32 - 0.6).collect(),
)
.unwrap()
.into_dyn();
check_weight_gradient(&mut conv, &x, 1e-3, 1e-2);
}
#[test]
fn conv3d_same_padding_input_gradient_matches_finite_difference() {
let mut conv = Conv3D::new(2, (3, 3, 3), vec![1, 1, 4, 4, 4], (1, 1, 1), Linear::new())
.unwrap()
.with_padding(PaddingType::Same);
let x = Array::from_shape_vec(
(1, 1, 4, 4, 4),
(0..64).map(|v| 0.03 * v as f32 - 0.9).collect(),
)
.unwrap()
.into_dyn();
check_input_gradient(&mut conv, &x, 1e-3, 2e-2);
}
#[test]
fn conv3d_same_padding_weight_gradient_matches_finite_difference() {
let mut conv = Conv3D::new(2, (3, 3, 3), vec![1, 1, 4, 4, 4], (1, 1, 1), Linear::new())
.unwrap()
.with_padding(PaddingType::Same);
let x = Array::from_shape_vec(
(1, 1, 4, 4, 4),
(0..64).map(|v| 0.03 * v as f32 - 0.9).collect(),
)
.unwrap()
.into_dyn();
check_weight_gradient(&mut conv, &x, 1e-3, 2e-2);
}
#[test]
fn separable_conv2d_same_padding_3x3_dm2_gradients_match_finite_difference() {
let make = || {
SeparableConv2D::new(2, (3, 3), vec![1, 2, 5, 5], (1, 1), 2, Linear::new())
.unwrap()
.with_padding(PaddingType::Same)
};
let x = Array::from_shape_vec(
(1, 2, 5, 5),
(0..50).map(|v| 0.04 * v as f32 - 1.0).collect(),
)
.unwrap()
.into_dyn();
check_input_gradient(&mut make(), &x, 1e-3, 3e-2);
check_weight_gradient(&mut make(), &x, 1e-3, 3e-2);
}
#[test]
fn simple_rnn_weight_gradient_matches_finite_difference() {
let mut rnn = SimpleRNN::new(2, 3, Tanh::new()).unwrap();
let x = Array::from_shape_vec((1, 3, 2), vec![0.3, -0.6, 0.9, -0.2, 0.5, -0.8])
.unwrap()
.into_dyn();
check_weight_gradient(&mut rnn, &x, 1e-3, 3e-2);
}
#[test]
fn lstm_weight_gradient_matches_finite_difference() {
let mut lstm = LSTM::new(2, 3, Tanh::new()).unwrap();
let x = Array::from_shape_vec((1, 3, 2), vec![0.3, -0.6, 0.9, -0.2, 0.5, -0.8])
.unwrap()
.into_dyn();
check_weight_gradient(&mut lstm, &x, 1e-3, 3e-2);
}
#[test]
fn gru_weight_gradient_matches_finite_difference() {
let mut gru = GRU::new(2, 3, Tanh::new()).unwrap();
let x = Array::from_shape_vec((1, 3, 2), vec![0.3, -0.6, 0.9, -0.2, 0.5, -0.8])
.unwrap()
.into_dyn();
check_weight_gradient(&mut gru, &x, 1e-3, 3e-2);
}
fn check_weight_gradient_weighted(layer: &mut dyn Layer, x: &Tensor, eps: f32, tol: f32) {
let out = layer.forward(x).unwrap();
let w = loss_weights(&out);
layer.backward(&w).unwrap();
let params: Vec<(Vec<f32>, Vec<f32>)> = layer
.parameters()
.into_iter()
.map(|pg| (pg.value.to_vec(), pg.grad.to_vec()))
.collect();
assert!(!params.is_empty(), "layer exposes no parameters to check");
for (p_idx, (values, grads)) in params.iter().enumerate() {
for i in 0..values.len() {
let orig = values[i];
layer.parameters()[p_idx].value[i] = orig + eps;
let l_plus: f32 = (&layer.forward(x).unwrap() * &w).sum();
layer.parameters()[p_idx].value[i] = orig - eps;
let l_minus: f32 = (&layer.forward(x).unwrap() * &w).sum();
layer.parameters()[p_idx].value[i] = orig;
let numeric = (l_plus - l_minus) / (2.0 * eps);
assert_abs_diff_eq!(grads[i], numeric, epsilon = tol);
}
}
}
#[test]
fn layer_normalization_default_input_gradient_matches_finite_difference() {
let mut ln = LayerNormalization::new(vec![2, 4], 1e-5).unwrap();
ln.set_training_if_mode_dependent(true);
let x = ramp(&[2, 4]);
check_input_gradient_weighted(&mut ln, &x, 1e-3, 5e-2);
}
#[test]
fn layer_normalization_default_weight_gradient_matches_finite_difference() {
let mut ln = LayerNormalization::new(vec![2, 4], 1e-5).unwrap();
ln.set_training_if_mode_dependent(true);
let x = ramp(&[2, 4]);
check_weight_gradient_weighted(&mut ln, &x, 1e-3, 5e-2);
}
#[test]
fn layer_normalization_custom_axis_input_gradient_matches_finite_difference() {
let mut ln = LayerNormalization::new(vec![3, 4], 1e-5)
.unwrap()
.with_normalized_axis(LayerNormalizationAxis::Custom(0))
.unwrap();
ln.set_training_if_mode_dependent(true);
let x = ramp(&[3, 4]);
check_input_gradient_weighted(&mut ln, &x, 1e-3, 5e-2);
}
#[test]
fn layer_normalization_rank3_default_input_gradient_matches_finite_difference() {
let mut ln = LayerNormalization::new(vec![2, 3, 4], 1e-5).unwrap();
ln.set_training_if_mode_dependent(true);
let x = ramp(&[2, 3, 4]);
check_input_gradient_weighted(&mut ln, &x, 1e-3, 5e-2);
}
#[test]
fn layer_normalization_trailing_custom_weight_gradient_matches_finite_difference() {
let mut ln = LayerNormalization::new(vec![3, 4], 1e-5)
.unwrap()
.with_normalized_axis(LayerNormalizationAxis::Custom(1))
.unwrap();
ln.set_training_if_mode_dependent(true);
let x = ramp(&[3, 4]);
check_weight_gradient_weighted(&mut ln, &x, 1e-3, 5e-2);
}
#[test]
fn layer_normalization_multiple_trailing_input_gradient_matches_finite_difference() {
let mut ln = LayerNormalization::new(vec![2, 3, 4], 1e-5)
.unwrap()
.with_normalized_axis(LayerNormalizationAxis::Multiple(vec![1, 2]))
.unwrap();
ln.set_training_if_mode_dependent(true);
let x = ramp(&[2, 3, 4]);
check_input_gradient_weighted(&mut ln, &x, 1e-3, 5e-2);
}
#[test]
fn layer_normalization_multiple_permuted_input_gradient_matches_finite_difference() {
let mut ln = LayerNormalization::new(vec![2, 3, 4], 1e-5)
.unwrap()
.with_normalized_axis(LayerNormalizationAxis::Multiple(vec![0, 2]))
.unwrap();
ln.set_training_if_mode_dependent(true);
let x = ramp(&[2, 3, 4]);
check_input_gradient_weighted(&mut ln, &x, 1e-3, 5e-2);
}
#[test]
fn group_normalization_input_gradient_matches_finite_difference() {
let mut gn = GroupNormalization::new(vec![1, 4, 4], 2, 1, 1e-5).unwrap();
gn.set_training_if_mode_dependent(true);
let x = ramp(&[1, 4, 4]);
check_input_gradient_weighted(&mut gn, &x, 1e-3, 5e-2);
}
#[test]
fn group_normalization_weight_gradient_matches_finite_difference() {
let mut gn = GroupNormalization::new(vec![1, 4, 4], 2, 1, 1e-5).unwrap();
gn.set_training_if_mode_dependent(true);
let x = ramp(&[1, 4, 4]);
check_weight_gradient_weighted(&mut gn, &x, 1e-3, 5e-2);
}
#[test]
fn group_normalization_channel_axis2_input_gradient_matches_finite_difference() {
let mut gn = GroupNormalization::new(vec![2, 4, 4], 2, 2, 1e-5).unwrap();
gn.set_training_if_mode_dependent(true);
let x = ramp(&[2, 4, 4]);
check_input_gradient_weighted(&mut gn, &x, 1e-3, 5e-2);
}
#[test]
fn instance_normalization_input_gradient_matches_finite_difference() {
let mut inn = InstanceNormalization::new(vec![1, 3, 4], 1, 1e-5).unwrap();
inn.set_training_if_mode_dependent(true);
let x = ramp(&[1, 3, 4]);
check_input_gradient_weighted(&mut inn, &x, 1e-3, 5e-2);
}
#[test]
fn instance_normalization_weight_gradient_matches_finite_difference() {
let mut inn = InstanceNormalization::new(vec![1, 3, 4], 1, 1e-5).unwrap();
inn.set_training_if_mode_dependent(true);
let x = ramp(&[1, 3, 4]);
check_weight_gradient_weighted(&mut inn, &x, 1e-3, 5e-2);
}
#[test]
fn batch_normalization_input_gradient_weighted_matches_finite_difference() {
let mut bn = BatchNormalization::new(vec![4, 3], 0.9, 1e-5).unwrap();
bn.set_training_if_mode_dependent(true);
let x = ramp(&[4, 3]);
check_input_gradient_weighted(&mut bn, &x, 1e-3, 5e-2);
}
#[test]
fn batch_normalization_weight_gradient_matches_finite_difference() {
let mut bn = BatchNormalization::new(vec![4, 3], 0.9, 1e-5).unwrap();
bn.set_training_if_mode_dependent(true);
let x = ramp(&[4, 3]);
check_weight_gradient_weighted(&mut bn, &x, 1e-3, 5e-2);
}
#[test]
fn batch_normalization_spatial_input_gradient_matches_finite_difference() {
let mut bn = BatchNormalization::new(vec![2, 3, 2, 2], 0.9, 1e-5).unwrap();
bn.set_training_if_mode_dependent(true);
let x = ramp(&[2, 3, 2, 2]);
check_input_gradient_weighted(&mut bn, &x, 1e-3, 5e-2);
}
#[test]
fn batch_normalization_spatial_weight_gradient_matches_finite_difference() {
let mut bn = BatchNormalization::new(vec![2, 3, 2, 2], 0.9, 1e-5).unwrap();
bn.set_training_if_mode_dependent(true);
let x = ramp(&[2, 3, 2, 2]);
check_weight_gradient_weighted(&mut bn, &x, 1e-3, 5e-2);
}