use crate::common::diagnostics::NamErrorCode;
use crate::math::common::{AlignedVec, SimdMath};
#[derive(Clone)]
#[repr(align(64))]
pub struct BatchNorm1D {
pub num_channels: usize,
pub scale: AlignedVec<f32>,
pub offset: AlignedVec<f32>,
}
impl BatchNorm1D {
pub fn from_params(
num_channels: usize,
gamma: &[f32],
beta: &[f32],
running_mean: &[f32],
running_var: &[f32],
eps: f32,
) -> Result<Self, NamErrorCode> {
assert_eq!(gamma.len(), num_channels);
assert_eq!(beta.len(), num_channels);
assert_eq!(running_mean.len(), num_channels);
assert_eq!(running_var.len(), num_channels);
let mut scale = AlignedVec::new(num_channels, 0.0f32)?;
let mut offset = AlignedVec::new(num_channels, 0.0f32)?;
for c in 0..num_channels {
let var = running_var[c];
assert!(var >= 0.0, "running_var[{c}] = {var} is negative");
let inv_std = 1.0 / (var + eps).sqrt();
scale[c] = gamma[c] * inv_std;
offset[c] = beta[c] - running_mean[c] * scale[c];
}
Ok(Self {
num_channels,
scale,
offset,
})
}
pub fn from_fused(
num_channels: usize,
scale: &[f32],
offset: &[f32],
) -> Result<Self, NamErrorCode> {
assert_eq!(scale.len(), num_channels);
assert_eq!(offset.len(), num_channels);
let mut s = AlignedVec::new(num_channels, 0.0f32)?;
let mut o = AlignedVec::new(num_channels, 0.0f32)?;
s.copy_from_slice(scale);
o.copy_from_slice(offset);
Ok(Self {
num_channels,
scale: s,
offset: o,
})
}
#[inline(always)]
pub unsafe fn process(&self, data: &mut [f32], num_frames: usize) {
debug_assert_eq!(data.len(), num_frames * self.num_channels);
unsafe { crate::math::common::dispatch_simd!(self, process_simd, data, num_frames) };
}
#[inline(always)]
pub unsafe fn process_simd<M: SimdMath>(&self, data: &mut [f32], num_frames: usize) {
debug_assert_eq!(data.len(), num_frames * self.num_channels);
unsafe {
M::batch_norm_process(
data,
&self.scale,
&self.offset,
self.num_channels,
num_frames,
);
}
}
#[inline(always)]
pub unsafe fn process_scalar(&self, data: &mut [f32], num_frames: usize) {
unsafe {
process_scalar_ref(
data,
&self.scale,
&self.offset,
self.num_channels,
num_frames,
);
}
}
}
#[inline(always)]
unsafe fn process_scalar_ref(
data: &mut [f32],
scale: &[f32],
offset: &[f32],
n_ch: usize,
num_frames: usize,
) {
for ch in 0..n_ch {
let s = unsafe { *scale.get_unchecked(ch) };
let o = unsafe { *offset.get_unchecked(ch) };
for f in 0..num_frames {
let idx = ch + f * n_ch;
unsafe {
*data.get_unchecked_mut(idx) = (*data.get_unchecked(idx)).mul_add(s, o);
}
}
}
}
#[cfg(test)]
#[path = "batch_norm_test.rs"]
mod tests;