use crate::{Shape, Tape, Tensor};
use crate::Module;
use super::RmsNorm;
#[test]
fn new_allocates_scale_and_epsilon() {
let tape = Tape::new();
let norm = RmsNorm::new(
&tape,
Tensor::filled([2], 1.0_f64),
Tensor::filled([], 1e-5),
);
assert_eq!(tape.len(), 2);
assert_eq!(norm.parameters().count(), 1);
}
#[test]
#[should_panic(expected = "must be rank 1")]
fn new_rejects_non_vector_scale() {
let tape = Tape::new();
RmsNorm::new(
&tape,
Tensor::filled([2, 2], 1.0_f64),
Tensor::filled([], 1e-5),
);
}
#[test]
#[should_panic(expected = "single value")]
fn new_rejects_multi_value_epsilon() {
let tape = Tape::new();
RmsNorm::new(
&tape,
Tensor::filled([2], 1.0_f64),
Tensor::filled([2], 1e-5),
);
}
#[test]
fn express_normalizes_by_the_root_mean_square() {
let tape = Tape::new();
let norm = RmsNorm::new(&tape, Tensor::filled([2], 1.0_f64), Tensor::filled([], 0.0));
let input = tape.leaf(Tensor::new([2, 2], [2.0, 2.0, 3.0, -3.0]));
let output = norm.express(input);
assert_eq!(output.shape(), Shape::new([2, 2]));
let output = output.symbol();
let network = tape.into_network();
let run = network.forward(&network.parameters(), []);
assert_eq!(run.of(output).to_vec(), &[1.0, 1.0, 1.0, -1.0]);
}
#[test]
fn express_applies_the_learned_scale() {
let tape = Tape::new();
let norm = RmsNorm::new(
&tape,
Tensor::new([2], [2.0_f64, 5.0]),
Tensor::filled([], 0.0),
);
let input = tape.leaf(Tensor::new([2, 2], [2.0, 2.0, 3.0, -3.0]));
let output = norm.express(input).symbol();
let network = tape.into_network();
let run = network.forward(&network.parameters(), []);
assert_eq!(run.of(output).to_vec(), &[2.0, 5.0, 2.0, -5.0]);
}
#[test]
fn express_records_tensor_granularity() {
let tape = Tape::new();
let norm = RmsNorm::new(&tape, Tensor::filled([2], 1.0_f64), Tensor::filled([], 0.0));
let input = tape.leaf(Tensor::new([3, 2], vec![1.0; 6]));
let nodes_before = tape.len();
norm.express(input);
assert_eq!(tape.len(), nodes_before + 11);
}
#[test]
#[should_panic(expected = "disagree on features")]
fn express_rejects_mismatched_features() {
let tape = Tape::new();
let norm = RmsNorm::new(&tape, Tensor::filled([2], 1.0_f64), Tensor::filled([], 0.0));
let input = tape.leaf(Tensor::new([2, 3], vec![1.0; 6]));
norm.express(input);
}
#[test]
fn gradients_flow_through_the_root_mean_square() {
let tape = Tape::new();
let norm = RmsNorm::new(&tape, Tensor::filled([2], 1.0_f64), Tensor::filled([], 2.0));
let input = tape.leaf(Tensor::new([1, 2], [2.0, 0.0]));
let output = norm.express(input);
let target = output.narrow(1, 0, 1).sum();
let (target, input) = (target.symbol(), input.symbol());
let network = tape.into_network();
let run = network.forward(&network.parameters(), []);
let gradients = run.backward(target);
assert_eq!(gradients.of(input).to_vec(), &[0.25, 0.0]);
let parameters: Vec<_> = norm.parameters().collect();
assert_eq!(gradients.of(parameters[0]).to_vec(), &[1.0, 0.0]);
}