use self::folds::{par_col_dot, par_col_sum, rows_per_block};
use crate::neural_network::Tensor;
use ndarray::{Array1, IxDyn};
use rayon::prelude::*;
tunable_gate! {
pub(crate) GN_ROW_PARALLEL_MIN_ELEMS => gn_row_parallel_min_elems / set_gn_row_parallel_min_elems = 262_144
}
tunable_gate! {
pub(crate) GN_PARAM_GRAD_PARALLEL_MIN_ELEMS => gn_param_grad_parallel_min_elems / set_gn_param_grad_parallel_min_elems = 262_144
}
fn positions_per_block(channels: usize) -> usize {
rows_per_block(channels)
}
struct GroupStats {
shift: Vec<f32>,
s1: Vec<f32>,
s2: Vec<f32>,
}
fn group_stats(x: &[f32], b: usize, layout: &GroupLayout, parallel: bool) -> GroupStats {
let (positions, channels, num_groups) = (layout.positions, layout.channels, layout.num_groups);
let cpg = layout.channels_per_group;
let base = b * positions * channels;
let shift: Vec<f32> = (0..num_groups).map(|gr| x[base + gr * cpg]).collect();
let block = positions_per_block(channels);
let shift_c: Vec<f32> = shift
.iter()
.flat_map(|&k| std::iter::repeat_n(k, cpg))
.collect();
let fold = |p0: usize| -> (Vec<f32>, Vec<f32>) {
let len = block.min(positions - p0);
let mut c1 = vec![0.0f32; channels];
let mut c2 = vec![0.0f32; channels];
for p in p0..p0 + len {
let row = &x[base + p * channels..base + (p + 1) * channels];
for (((a1, a2), &v), &k) in c1
.iter_mut()
.zip(c2.iter_mut())
.zip(row)
.zip(shift_c.iter())
{
let d = v - k;
*a1 += d;
*a2 += d * d;
}
}
let mut b1 = vec![0.0f32; num_groups];
let mut b2 = vec![0.0f32; num_groups];
for gr in 0..num_groups {
for j in gr * cpg..(gr + 1) * cpg {
b1[gr] += c1[j];
b2[gr] += c2[j];
}
}
(b1, b2)
};
let starts: Vec<usize> = (0..positions).step_by(block.max(1)).collect();
let parts: Vec<(Vec<f32>, Vec<f32>)> = if parallel {
starts.par_iter().map(|&p0| fold(p0)).collect()
} else {
starts.iter().map(|&p0| fold(p0)).collect()
};
let mut s1 = vec![0.0f32; num_groups];
let mut s2 = vec![0.0f32; num_groups];
for (b1, b2) in parts {
for gr in 0..num_groups {
s1[gr] += b1[gr];
s2[gr] += b2[gr];
}
}
GroupStats { shift, s1, s2 }
}
fn per_channel_stats(
stats: &GroupStats,
layout: &GroupLayout,
epsilon: f32,
) -> (Vec<f32>, Vec<f32>) {
let cpg = layout.channels_per_group;
let n = (layout.positions * cpg) as f32;
let mut mean_c = vec![0.0f32; layout.channels];
let mut inv_std_c = vec![0.0f32; layout.channels];
for gr in 0..layout.num_groups {
let m = stats.s1[gr] / n;
let var = (stats.s2[gr] / n - m * m).max(0.0);
let inv = 1.0 / (var + epsilon).sqrt();
for j in 0..cpg {
mean_c[gr * cpg + j] = stats.shift[gr] + m;
inv_std_c[gr * cpg + j] = inv;
}
}
(mean_c, inv_std_c)
}
type ForwardSlabs<'a> = (usize, ((&'a mut [f32], &'a mut [f32]), &'a [f32]));
type BackwardSlabs<'a> = (usize, (&'a mut [f32], (&'a [f32], &'a [f32])));
pub(super) struct GroupLayout {
pub batch: usize,
pub positions: usize,
pub channels: usize,
pub num_groups: usize,
pub channels_per_group: usize,
}
impl GroupLayout {
pub(super) fn new(shape: &[usize], num_groups: usize) -> Self {
let channels = shape[shape.len() - 1];
let positions: usize = shape[1..shape.len() - 1].iter().product();
Self {
batch: shape[0],
positions,
channels,
num_groups,
channels_per_group: channels / num_groups.max(1),
}
}
}
pub(super) fn group_norm_forward_core(
input: &Tensor,
num_groups: usize,
gamma: &Tensor,
beta: &Tensor,
epsilon: f32,
) -> (Tensor, Tensor, Tensor) {
let shape = input.shape().to_vec();
let layout = GroupLayout::new(&shape, num_groups);
let total = input.len();
if total == 0 {
return (
Tensor::zeros(IxDyn(&shape)),
Tensor::zeros(IxDyn(&shape)),
Array1::<f32>::zeros(layout.batch * num_groups).into_dyn(),
);
}
let input_std = input.as_standard_layout();
let x = input_std.as_slice().unwrap();
let parallel = total >= gn_row_parallel_min_elems();
let (gamma_s, beta_s) = (gamma.as_slice().unwrap(), beta.as_slice().unwrap());
let mut x_normalized = Tensor::zeros(IxDyn(&shape));
let mut output = Tensor::zeros(IxDyn(&shape));
let mut inv_std = Array1::<f32>::zeros(layout.batch * num_groups);
let item = layout.positions * layout.channels;
let batch_parallel = parallel && layout.batch >= rayon::current_num_threads();
let stats_parallel = parallel && !batch_parallel;
let per_item = |b: usize| {
per_channel_stats(
&group_stats(x, b, &layout, stats_parallel),
&layout,
epsilon,
)
};
let stats: Vec<(Vec<f32>, Vec<f32>)> = if batch_parallel {
(0..layout.batch).into_par_iter().map(per_item).collect()
} else {
(0..layout.batch).map(per_item).collect()
};
for (b, (_, inv_std_c)) in stats.iter().enumerate() {
for gr in 0..num_groups {
inv_std[b * num_groups + gr] = inv_std_c[gr * layout.channels_per_group];
}
}
let sweep = |(b, ((xn, out), src)): ForwardSlabs| {
let (mean_c, inv_std_c) = &stats[b];
for ((x_row, xn_row), out_row) in src
.chunks_exact(layout.channels)
.zip(xn.chunks_exact_mut(layout.channels))
.zip(out.chunks_exact_mut(layout.channels))
{
for c in 0..layout.channels {
let v = (x_row[c] - mean_c[c]) * inv_std_c[c];
xn_row[c] = v;
out_row[c] = v * gamma_s[c] + beta_s[c];
}
}
};
{
let xn = x_normalized.as_slice_mut().unwrap();
let out = output.as_slice_mut().unwrap();
if parallel {
xn.par_chunks_mut(item)
.zip(out.par_chunks_mut(item))
.zip(x.par_chunks(item))
.enumerate()
.for_each(sweep);
} else {
xn.chunks_mut(item)
.zip(out.chunks_mut(item))
.zip(x.chunks(item))
.enumerate()
.for_each(sweep);
}
}
(output, x_normalized, inv_std.into_dyn())
}
pub(super) fn group_norm_backward_core(
grad_output: &Tensor,
x_normalized: &Tensor,
inv_std: &Tensor,
num_groups: usize,
gamma: &Tensor,
) -> (Tensor, Tensor, Tensor) {
let shape = grad_output.shape().to_vec();
let layout = GroupLayout::new(&shape, num_groups);
let total = grad_output.len();
if total == 0 {
return (
Tensor::zeros(IxDyn(&shape)),
Array1::<f32>::zeros(layout.channels).into_dyn(),
Array1::<f32>::zeros(layout.channels).into_dyn(),
);
}
let grad_std = grad_output.as_standard_layout();
let g = grad_std.as_slice().unwrap();
let xn_std = x_normalized.as_standard_layout();
let xn = xn_std.as_slice().unwrap();
let inv_std_s = inv_std.as_slice().unwrap();
let gamma_s = gamma.as_slice().unwrap();
let col_parallel = total >= gn_param_grad_parallel_min_elems();
let grad_beta = par_col_sum(g, layout.channels, col_parallel, 1.0);
let grad_gamma = par_col_dot(g, xn, layout.channels, col_parallel, 1.0);
let (channels, cpg) = (layout.channels, layout.channels_per_group);
let n = (layout.positions * cpg) as f32;
let item = layout.positions * channels;
let mut grad_input = Tensor::zeros(IxDyn(&shape));
let per_item = |(b, (dx, (g_item, xn_item))): BackwardSlabs| {
let mut sum_g = vec![0.0f32; num_groups];
let mut sum_g_xn = vec![0.0f32; num_groups];
for (g_row, xn_row) in g_item
.chunks_exact(channels)
.zip(xn_item.chunks_exact(channels))
{
for gr in 0..num_groups {
let (mut a1, mut a2) = (0.0f32, 0.0f32);
for c in gr * cpg..(gr + 1) * cpg {
let t = g_row[c] * gamma_s[c];
a1 += t;
a2 += t * xn_row[c];
}
sum_g[gr] += a1;
sum_g_xn[gr] += a2;
}
}
let mut is_c = vec![0.0f32; channels];
let mut sg_c = vec![0.0f32; channels];
let mut sgx_c = vec![0.0f32; channels];
for gr in 0..num_groups {
let inv = inv_std_s[b * num_groups + gr];
for j in 0..cpg {
is_c[gr * cpg + j] = inv;
sg_c[gr * cpg + j] = sum_g[gr] / n;
sgx_c[gr * cpg + j] = sum_g_xn[gr] / n;
}
}
for ((dx_row, g_row), xn_row) in dx
.chunks_exact_mut(channels)
.zip(g_item.chunks_exact(channels))
.zip(xn_item.chunks_exact(channels))
{
for c in 0..channels {
dx_row[c] = is_c[c] * (g_row[c] * gamma_s[c] - (sg_c[c] + xn_row[c] * sgx_c[c]));
}
}
};
{
let dx_all = grad_input.as_slice_mut().unwrap();
if total >= gn_row_parallel_min_elems() {
dx_all
.par_chunks_mut(item)
.zip(g.par_chunks(item).zip(xn.par_chunks(item)))
.enumerate()
.for_each(per_item);
} else {
dx_all
.chunks_mut(item)
.zip(g.chunks(item).zip(xn.chunks(item)))
.enumerate()
.for_each(per_item);
}
}
(grad_input, grad_gamma, grad_beta)
}
mod folds;
pub mod batch_normalization;
pub mod group_normalization;
pub mod instance_normalization;
pub mod layer_normalization;
pub use batch_normalization::BatchNormalization;
pub use group_normalization::GroupNormalization;
pub use instance_normalization::InstanceNormalization;
pub use layer_normalization::{LayerNormalization, LayerNormalizationAxis};
macro_rules! normalization_layer_output_shape {
($self:expr) => {
if !$self.input_shape.is_empty() {
format!(
"({})",
$self
.input_shape
.iter()
.map(|x| x.to_string())
.collect::<Vec<_>>()
.join(", ")
)
} else {
String::from("Unknown")
}
};
}
pub(in crate::neural_network::layers::regularization::normalization) use normalization_layer_output_shape;
#[cfg(test)]
mod tests {
use super::*;
use approx::assert_abs_diff_eq;
fn tensor(data: Vec<f32>, shape: &[usize]) -> Tensor {
Tensor::from_shape_vec(IxDyn(shape), data).unwrap()
}
#[test]
fn group_norm_forward_hand_derived() {
let x = tensor(vec![1.0, 2.0, 10.0, 20.0, 3.0, 4.0, 30.0, 40.0], &[1, 2, 4]);
let gamma = tensor(vec![1.0; 4], &[4]);
let beta = tensor(vec![0.0; 4], &[4]);
let (out, _xn, inv_std) = group_norm_forward_core(&x, 2, &gamma, &beta, 1e-5);
let inv0 = 1.0 / (1.25f32 + 1e-5).sqrt();
let inv1 = 1.0 / (125.0f32 + 1e-5).sqrt();
assert_abs_diff_eq!(inv_std.as_slice().unwrap()[0], inv0, epsilon = 1e-6);
assert_abs_diff_eq!(inv_std.as_slice().unwrap()[1], inv1, epsilon = 1e-6);
let got: Vec<f32> = out.iter().copied().collect();
let want = [
-1.5 * inv0,
-0.5 * inv0,
-15.0 * inv1,
-5.0 * inv1,
0.5 * inv0,
1.5 * inv0,
5.0 * inv1,
15.0 * inv1,
];
for (g, w) in got.iter().zip(want) {
assert_abs_diff_eq!(*g, w, epsilon = 1e-5);
}
}
#[test]
fn group_norm_variance_survives_a_large_mean() {
const BASE: f32 = 1.0e6;
let x = tensor(
vec![BASE + 1.0, BASE + 2.0, BASE + 3.0, BASE + 4.0],
&[1, 4, 1],
);
let gamma = tensor(vec![1.0], &[1]);
let beta = tensor(vec![0.0], &[1]);
let (out, _xn, _inv) = group_norm_forward_core(&x, 1, &gamma, &beta, 1e-5);
let inv = 1.0 / (1.25f32).sqrt();
let want = [-1.5 * inv, -0.5 * inv, 0.5 * inv, 1.5 * inv];
for (g, w) in out.iter().zip(want) {
assert_abs_diff_eq!(*g, w, epsilon = 2e-3);
}
}
#[test]
fn group_norm_group_count_extremes() {
let x = tensor(vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0], &[1, 3, 2]);
let gamma = tensor(vec![1.0, 1.0], &[2]);
let beta = tensor(vec![0.0, 0.0], &[2]);
let (one, _, _) = group_norm_forward_core(&x, 1, &gamma, &beta, 0.0);
let mean_all = 3.5f32;
let var_all = (0..6)
.map(|i| (i as f32 + 1.0 - mean_all).powi(2))
.sum::<f32>()
/ 6.0;
assert_abs_diff_eq!(
one.iter().copied().next().unwrap(),
(1.0 - mean_all) / var_all.sqrt(),
epsilon = 1e-5
);
let (two, _, _) = group_norm_forward_core(&x, 2, &gamma, &beta, 0.0);
let var_ch = ((1.0f32 - 3.0).powi(2) + 0.0 + (5.0f32 - 3.0).powi(2)) / 3.0;
assert_abs_diff_eq!(
two.iter().copied().next().unwrap(),
(1.0 - 3.0) / var_ch.sqrt(),
epsilon = 1e-5
);
}
#[test]
fn group_stats_parallel_matches_serial_bitwise() {
let (b, p, c) = (2usize, 300usize, 8usize);
let data: Vec<f32> = (0..b * p * c)
.map(|i| ((i % 23) as f32 - 11.0) * 0.375)
.collect();
let x = tensor(data, &[b, p, c]);
let xs = x.as_slice().unwrap();
let layout = GroupLayout::new(x.shape(), 4);
for item in 0..b {
let serial = group_stats(xs, item, &layout, false);
let par = group_stats(xs, item, &layout, true);
assert_eq!(serial.s1, par.s1, "sums differ across the gate");
assert_eq!(serial.s2, par.s2, "squared sums differ across the gate");
}
}
#[test]
fn group_norm_backward_matches_finite_difference() {
let shape = [1usize, 3, 4];
let base: Vec<f32> = (0..12)
.map(|i| (i as f32 * 0.7).sin() * 2.0 + 0.3)
.collect();
let gamma = tensor(vec![1.3, 0.7, -0.4, 1.1], &[4]);
let beta = tensor(vec![0.2, -0.1, 0.5, 0.0], &[4]);
let groups = 2;
let eps = 1e-5;
let g_vec: Vec<f32> = (0..12).map(|i| ((i * 5 % 7) as f32 - 3.0) * 0.25).collect();
let g = tensor(g_vec.clone(), &shape);
let x = tensor(base.clone(), &shape);
let (_, xn, inv_std) = group_norm_forward_core(&x, groups, &gamma, &beta, eps);
let (grad_in, _, _) = group_norm_backward_core(&g, &xn, &inv_std, groups, &gamma);
let h = 1e-2f32;
for i in 0..12 {
let mut up = base.clone();
let mut dn = base.clone();
up[i] += h;
dn[i] -= h;
let (out_up, _, _) =
group_norm_forward_core(&tensor(up, &shape), groups, &gamma, &beta, eps);
let (out_dn, _, _) =
group_norm_forward_core(&tensor(dn, &shape), groups, &gamma, &beta, eps);
let num: f32 = out_up
.iter()
.zip(out_dn.iter())
.zip(&g_vec)
.map(|((u, d), gv)| gv * (u - d))
.sum::<f32>()
/ (2.0 * h);
assert_abs_diff_eq!(grad_in.as_slice().unwrap()[i], num, epsilon = 5e-3);
}
}
}