use super::line::{for_each_unit, LineLayout, UnitKernel, UnitPtr, PANEL};
use super::{
check_reduce_layout_offset_arithmetic, checked_total_len, reduce_uninit_writer, reduce_writer,
ReduceWriter,
};
use crate::erased_common::{check_dtype, validate_uninit_no_overlap};
use crate::*;
#[non_exhaustive]
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub enum NormKind {
Layer,
Rms,
}
#[non_exhaustive]
#[derive(Clone, Copy, Debug, PartialEq)]
pub struct NormSpec {
kind: NormKind,
eps: f64,
weight_stride: Option<isize>,
bias_stride: Option<isize>,
}
impl NormSpec {
#[inline]
pub const fn layer_norm(eps: f64) -> Self {
Self {
kind: NormKind::Layer,
eps,
weight_stride: None,
bias_stride: None,
}
}
#[inline]
pub const fn rms_norm(eps: f64) -> Self {
Self {
kind: NormKind::Rms,
eps,
weight_stride: None,
bias_stride: None,
}
}
#[inline]
pub const fn with_weight(mut self, stride: isize) -> Self {
self.weight_stride = Some(stride);
self
}
#[inline]
pub const fn with_bias(mut self, stride: isize) -> Self {
self.bias_stride = Some(stride);
self
}
#[inline]
pub const fn kind(&self) -> NormKind {
self.kind
}
#[inline]
pub const fn eps(&self) -> f64 {
self.eps
}
#[inline]
pub const fn weight_stride(&self) -> Option<isize> {
self.weight_stride
}
#[inline]
pub const fn bias_stride(&self) -> Option<isize> {
self.bias_stride
}
}
#[derive(Clone, Debug)]
pub struct ErasedNormPlan {
dtype: KernelDType,
spec: NormSpec,
dims: Vec<usize>,
src_strides: Vec<isize>,
dest_strides: Vec<isize>,
layout: LineLayout,
}
impl ErasedNormPlan {
pub fn compile(
dtype: KernelDType,
spec: NormSpec,
dims: &[usize],
src_strides: &[isize],
dest_strides: &[isize],
axis: usize,
) -> Result<Self> {
if !matches!(dtype, KernelDType::F32 | KernelDType::F64) {
return Err(StridedError::UnsupportedDType {
dtype: dtype.label(),
});
}
if !(spec.eps.is_finite() && spec.eps >= 0.0) {
return Err(StridedError::UnsupportedOp {
op: "norm with a negative or non-finite eps",
dtype: dtype.label(),
});
}
if dims.len() != src_strides.len() || dims.len() != dest_strides.len() {
return Err(StridedError::StrideLengthMismatch);
}
if axis >= dims.len() {
return Err(StridedError::InvalidAxis {
axis,
rank: dims.len(),
});
}
checked_total_len(dims)?;
check_reduce_layout_offset_arithmetic(dims, src_strides)?;
check_reduce_layout_offset_arithmetic(dims, dest_strides)?;
for stride in [spec.weight_stride, spec.bias_stride].into_iter().flatten() {
check_reduce_layout_offset_arithmetic(&dims[axis..=axis], &[stride])?;
}
if !crate::layout_check::is_injective_layout(dims, dest_strides) {
return Err(StridedError::NonInjectiveOutputLayout);
}
let dest_outer: Vec<isize> = (0..dims.len())
.filter(|&a| a != axis)
.map(|a| dest_strides[a])
.collect();
let layout = LineLayout::compile(
dims,
src_strides,
&dest_outer,
dest_strides[axis],
axis,
true,
)?;
Ok(Self {
dtype,
spec,
dims: dims.to_vec(),
src_strides: src_strides.to_vec(),
dest_strides: dest_strides.to_vec(),
layout,
})
}
#[inline]
pub fn dtype(&self) -> KernelDType {
self.dtype
}
#[inline]
pub fn spec(&self) -> NormSpec {
self.spec
}
fn check_layouts(
&self,
dest_dims: &[usize],
dest_strides: &[isize],
src: &ErasedRawStridedRef<'_>,
weight: Option<&ErasedRawStridedRef<'_>>,
bias: Option<&ErasedRawStridedRef<'_>>,
) -> Result<()> {
if src.dims() != self.dims.as_slice()
|| src.strides() != self.src_strides.as_slice()
|| dest_dims != self.dims.as_slice()
|| dest_strides != self.dest_strides.as_slice()
{
return Err(StridedError::PlanLayoutMismatch);
}
let n = self.layout.axis_len;
for (param, stride) in [
(weight, self.spec.weight_stride),
(bias, self.spec.bias_stride),
] {
match (param, stride) {
(None, None) => {}
(Some(param), Some(stride)) => {
check_dtype(self.dtype, param.dtype())?;
if param.dims() != [n] || param.strides() != [stride] {
return Err(StridedError::PlanLayoutMismatch);
}
}
_ => return Err(StridedError::PlanLayoutMismatch),
}
}
Ok(())
}
pub fn execute(
&self,
ctx: &ExecContext,
dest: &mut ErasedRawStridedMut<'_>,
src: &ErasedRawStridedRef<'_>,
weight: Option<&ErasedRawStridedRef<'_>>,
bias: Option<&ErasedRawStridedRef<'_>>,
) -> Result<()> {
check_dtype(self.dtype, dest.dtype())?;
check_dtype(self.dtype, src.dtype())?;
self.check_layouts(dest.dims(), dest.strides(), src, weight, bias)?;
match self.dtype {
KernelDType::F32 => {
let mut writer = reduce_writer::<f32>(dest)?;
self.dispatch::<f32, _>(ctx, &mut writer, src, weight, bias)
}
_ => {
let mut writer = reduce_writer::<f64>(dest)?;
self.dispatch::<f64, _>(ctx, &mut writer, src, weight, bias)
}
}
}
pub fn execute_uninit(
&self,
ctx: &ExecContext,
dest: &mut ErasedRawStridedUninitMut<'_>,
src: &ErasedRawStridedPtr<'_>,
weight: Option<&ErasedRawStridedPtr<'_>>,
bias: Option<&ErasedRawStridedPtr<'_>>,
) -> Result<()> {
check_dtype(self.dtype, dest.dtype())?;
check_dtype(self.dtype, src.dtype())?;
validate_uninit_no_overlap(dest, src, 0)?;
if let Some(weight) = weight {
validate_uninit_no_overlap(dest, weight, 1)?;
}
if let Some(bias) = bias {
validate_uninit_no_overlap(dest, bias, 2)?;
}
let src = unsafe { src.try_as_ref_after_no_overlap() }?;
let weight = weight
.map(|weight| unsafe { weight.try_as_ref_after_no_overlap() })
.transpose()?;
let bias = bias
.map(|bias| unsafe { bias.try_as_ref_after_no_overlap() })
.transpose()?;
self.check_layouts(
dest.dims(),
dest.strides(),
&src,
weight.as_ref(),
bias.as_ref(),
)?;
match self.dtype {
KernelDType::F32 => {
let mut writer = reduce_uninit_writer::<f32>(dest)?;
self.dispatch::<f32, _>(ctx, &mut writer, &src, weight.as_ref(), bias.as_ref())
}
_ => {
let mut writer = reduce_uninit_writer::<f64>(dest)?;
self.dispatch::<f64, _>(ctx, &mut writer, &src, weight.as_ref(), bias.as_ref())
}
}
}
fn dispatch<T, W>(
&self,
ctx: &ExecContext,
dest: &mut W,
src: &ErasedRawStridedRef<'_>,
weight: Option<&ErasedRawStridedRef<'_>>,
bias: Option<&ErasedRawStridedRef<'_>>,
) -> Result<()>
where
T: NormScalar,
W: ReduceWriter<T>,
{
let affine = Affine {
weight: param::<T>(weight)?,
bias: param::<T>(bias)?,
};
macro_rules! go {
($layer:literal) => {
match (weight.is_some(), bias.is_some()) {
(false, false) => {
self.run::<T, W, $layer, false, false>(ctx, dest, src, affine)
}
(true, false) => self.run::<T, W, $layer, true, false>(ctx, dest, src, affine),
(false, true) => self.run::<T, W, $layer, false, true>(ctx, dest, src, affine),
(true, true) => self.run::<T, W, $layer, true, true>(ctx, dest, src, affine),
}
};
}
match self.spec.kind {
NormKind::Layer => go!(true),
NormKind::Rms => go!(false),
}
}
fn run<T, W, const LAYER: bool, const WEIGHT: bool, const BIAS: bool>(
&self,
ctx: &ExecContext,
dest: &mut W,
src: &ErasedRawStridedRef<'_>,
affine: Affine<T>,
) -> Result<()>
where
T: NormScalar,
W: ReduceWriter<T>,
{
let layout = &self.layout;
let source = UnitPtr(src.data_as::<T>()?.as_ptr() as *mut T);
let target = UnitPtr(unsafe { dest.ptr() });
let line = Line {
n: layout.axis_len,
n_t: T::from_usize(layout.axis_len),
eps: T::from_f64(self.spec.eps),
ss: layout.src_axis_stride,
ds: layout.dest_axis_stride,
};
let kernel = NormUnit::<T, LAYER, WEIGHT, BIAS> {
source,
target,
line,
affine,
};
unsafe { for_each_unit(ctx, layout, src.offset(), dest.offset(), kernel) }
}
}
#[derive(Clone, Copy)]
struct NormUnit<T, const LAYER: bool, const WEIGHT: bool, const BIAS: bool> {
source: UnitPtr<T>,
target: UnitPtr<T>,
line: Line<T>,
affine: Affine<T>,
}
impl<T: NormScalar, const LAYER: bool, const WEIGHT: bool, const BIAS: bool> UnitKernel
for NormUnit<T, LAYER, WEIGHT, BIAS>
{
#[inline(always)]
unsafe fn unit(self, so: isize, d_o: isize, width: usize) {
let Self {
source,
target,
line,
affine,
} = self;
unsafe {
if width == 1 {
norm_line::<T, LAYER, WEIGHT, BIAS>(
source.get(),
so,
target.get(),
d_o,
line,
affine,
)
} else {
norm_panel::<T, LAYER, WEIGHT, BIAS>(
source.get(),
so,
target.get(),
d_o,
width,
line,
affine,
)
}
}
}
}
#[derive(Clone, Copy)]
struct Param<T> {
ptr: UnitPtr<T>,
offset: isize,
stride: isize,
}
#[derive(Clone, Copy)]
struct Affine<T> {
weight: Option<Param<T>>,
bias: Option<Param<T>>,
}
impl<T> Param<T> {
#[inline(always)]
unsafe fn at(self, k: usize) -> T
where
T: Copy,
{
unsafe {
self.ptr
.get()
.offset(self.offset + k as isize * self.stride)
.read()
}
}
}
fn param<T: NormScalar>(param: Option<&ErasedRawStridedRef<'_>>) -> Result<Option<Param<T>>> {
param
.map(|param| {
Ok(Param {
ptr: UnitPtr(param.data_as::<T>()?.as_ptr() as *mut T),
offset: param.offset(),
stride: param.strides()[0],
})
})
.transpose()
}
#[derive(Clone, Copy)]
struct Line<T> {
n: usize,
n_t: T,
eps: T,
ss: isize,
ds: isize,
}
pub(super) trait NormScalar:
KernelStorageElement
+ MaybeSendSync
+ PartialOrd
+ core::ops::Add<Output = Self>
+ core::ops::Sub<Output = Self>
+ core::ops::Mul<Output = Self>
+ core::ops::Div<Output = Self>
{
const ZERO: Self;
const ONE: Self;
fn from_usize(value: usize) -> Self;
fn from_f64(value: f64) -> Self;
fn sqrt(self) -> Self;
}
macro_rules! impl_norm_scalar {
($($ty:ty),*) => {$(
impl NormScalar for $ty {
const ZERO: Self = 0.0;
const ONE: Self = 1.0;
#[inline(always)]
fn from_usize(value: usize) -> Self {
value as $ty
}
#[inline(always)]
fn from_f64(value: f64) -> Self {
value as $ty
}
#[inline(always)]
fn sqrt(self) -> Self {
<$ty>::sqrt(self)
}
}
)*};
}
impl_norm_scalar!(f32, f64);
const NORM_LANES: usize = 16;
#[inline(always)]
fn lane_sum<T: NormScalar>(values: &[T], map: impl Fn(T) -> T) -> T {
let mut partial = [T::ZERO; NORM_LANES];
let mut chunks = values.chunks_exact(NORM_LANES);
for chunk in chunks.by_ref() {
for (partial, &value) in partial.iter_mut().zip(chunk) {
*partial = *partial + map(value);
}
}
let mut tail = T::ZERO;
for &value in chunks.remainder() {
tail = tail + map(value);
}
partial.into_iter().fold(T::ZERO, |acc, value| acc + value) + tail
}
#[inline(always)]
unsafe fn strided_sum<T: NormScalar>(
src: *const T,
so: isize,
ss: isize,
n: usize,
map: impl Fn(T) -> T,
) -> T {
let mut sum = T::ZERO;
let mut offset = so;
for _ in 0..n {
sum = sum + map(unsafe { src.offset(offset).read() });
offset += ss;
}
sum
}
#[derive(Clone, Copy)]
struct Stats<T> {
shift: T,
mean: T,
inv: T,
}
impl<T: NormScalar> Stats<T> {
#[inline(always)]
fn apply(self, x: T) -> T {
((x - self.shift) - self.mean) * self.inv
}
}
#[inline(always)]
unsafe fn line_stats<T: NormScalar, const LAYER: bool>(
src: *const T,
so: isize,
line: Line<T>,
) -> Stats<T> {
unsafe {
let shift = if LAYER {
src.offset(so).read()
} else {
T::ZERO
};
let (mean, var) = if line.ss == 1 {
let values = core::slice::from_raw_parts(src.offset(so), line.n);
let mean = if LAYER {
lane_sum(values, |x| x - shift) / line.n_t
} else {
T::ZERO
};
let var = lane_sum(values, |x| {
let d = (x - shift) - mean;
d * d
}) / line.n_t;
(mean, var)
} else {
let mean = if LAYER {
strided_sum(src, so, line.ss, line.n, |x| x - shift) / line.n_t
} else {
T::ZERO
};
let var = strided_sum(src, so, line.ss, line.n, |x| {
let d = (x - shift) - mean;
d * d
}) / line.n_t;
(mean, var)
};
Stats {
shift,
mean,
inv: T::ONE / (var + line.eps).sqrt(),
}
}
}
#[inline(always)]
unsafe fn norm_line<T: NormScalar, const LAYER: bool, const WEIGHT: bool, const BIAS: bool>(
src: *const T,
so: isize,
dst: *mut T,
d_o: isize,
line: Line<T>,
affine: Affine<T>,
) {
unsafe {
let stats = line_stats::<T, LAYER>(src, so, line);
let src = src.offset(so);
let dst = dst.offset(d_o);
let (w, ws) = match affine.weight {
Some(p) if WEIGHT => (p.ptr.get().offset(p.offset) as *const T, p.stride),
_ => (src, 0),
};
let (b, bs) = match affine.bias {
Some(p) if BIAS => (p.ptr.get().offset(p.offset) as *const T, p.stride),
_ => (src, 0),
};
let unit = |stride: isize| stride == 1 || stride == 0;
if line.ss == 1 && line.ds == 1 && unit(ws) && unit(bs) {
norm_output::<T, WEIGHT, BIAS>(
src,
1,
dst,
1,
w,
ws.min(1),
b,
bs.min(1),
line.n,
stats,
);
} else {
norm_output::<T, WEIGHT, BIAS>(src, line.ss, dst, line.ds, w, ws, b, bs, line.n, stats);
}
}
}
#[allow(clippy::too_many_arguments)]
#[inline(always)]
unsafe fn norm_output<T: NormScalar, const WEIGHT: bool, const BIAS: bool>(
src: *const T,
ss: isize,
dst: *mut T,
ds: isize,
w: *const T,
ws: isize,
b: *const T,
bs: isize,
n: usize,
stats: Stats<T>,
) {
for k in 0..n as isize {
unsafe {
let mut y = stats.apply(src.offset(k * ss).read());
if WEIGHT {
y = y * w.offset(k * ws).read();
}
if BIAS {
y = y + b.offset(k * bs).read();
}
dst.offset(k * ds).write(y);
}
}
}
#[allow(clippy::too_many_arguments)]
#[inline(always)]
unsafe fn norm_panel<T: NormScalar, const LAYER: bool, const WEIGHT: bool, const BIAS: bool>(
src: *const T,
so: isize,
dst: *mut T,
d_o: isize,
width: usize,
line: Line<T>,
affine: Affine<T>,
) {
debug_assert!(width <= PANEL);
let mut shift = [T::ZERO; PANEL];
let mut mean = [T::ZERO; PANEL];
let mut scale = [T::ZERO; PANEL];
let shift = &mut shift[..width];
let mean = &mut mean[..width];
let scale = &mut scale[..width];
unsafe {
let row =
|k: usize| core::slice::from_raw_parts(src.offset(so + k as isize * line.ss), width);
if LAYER {
shift.copy_from_slice(row(0));
for k in 0..line.n {
for ((mean, &shift), &x) in mean.iter_mut().zip(shift.iter()).zip(row(k)) {
*mean = *mean + (x - shift);
}
}
for mean in mean.iter_mut() {
*mean = *mean / line.n_t;
}
}
for k in 0..line.n {
for (((acc, &shift), &mean), &x) in scale
.iter_mut()
.zip(shift.iter())
.zip(mean.iter())
.zip(row(k))
{
let d = (x - shift) - mean;
*acc = *acc + d * d;
}
}
for scale in scale.iter_mut() {
*scale = T::ONE / (*scale / line.n_t + line.eps).sqrt();
}
for k in 0..line.n {
let w = if WEIGHT {
affine.weight.unwrap_unchecked().at(k)
} else {
T::ONE
};
let b = if BIAS {
affine.bias.unwrap_unchecked().at(k)
} else {
T::ZERO
};
let out = dst.offset(d_o + k as isize * line.ds);
for (lane, (((&x, &shift), &mean), &scale)) in row(k)
.iter()
.zip(shift.iter())
.zip(mean.iter())
.zip(scale.iter())
.enumerate()
{
let mut y = ((x - shift) - mean) * scale;
if WEIGHT {
y = y * w;
}
if BIAS {
y = y + b;
}
out.add(lane).write(y);
}
}
}
}