pub struct ErasedNormPlan { /* private fields */ }Expand description
Dtype-erased fused layer_norm / rms_norm along one axis.
Each line along axis is normalized independently in one call:
NormKind::Layer:y = (x - mean) * rsqrt(var + eps) * weight + biaswithmean = sum(x) / nand the biased (population) variancevar = sum((x - mean)^2) / n, as in PyTorch’slayer_norm.NormKind::Rms:y = x * rsqrt(sum(x^2) / n + eps) * weight + bias.
rsqrt(v) is evaluated as 1 / sqrt(v) once per line, and the weight and
bias terms are present only when the NormSpec requests them. The
destination has the source dimensions with independent strides. Supported
dtypes are f32 and f64; every accumulation is in the element dtype.
§Numerics
Layer norm uses the shifted-data two-pass algorithm: every element is first
shifted by the first element of its line, then the mean of the shifted
values and the variance about that mean are computed in two passes. This
avoids the cancellation of the one-pass E[x^2] - E[x]^2 form and of a
large common offset. A line whose elements are all equal has a variance of
exactly zero and normalizes to exactly 0 * rsqrt(eps) * weight + bias,
i.e. bias (or zero) whenever eps > 0. With eps == 0 such a line
evaluates 0 * inf and is NaN; that is the caller’s choice of eps, not
an error.
Non-finite inputs follow IEEE evaluation of the formula. A NaN anywhere in
a line makes every output of that line NaN. An infinite element makes a
layer-norm line NaN (its deviation is inf - inf), and makes an RMS-norm
line zero at its finite elements and NaN at the infinite ones
(rsqrt(inf) = 0). Squares are not rescaled: when the sum of squared
deviations overflows to inf (finite deviations above about 1.8e19 for
f32 or 1.3e154 for f64), the line’s finite elements normalize to zero.
When the normalized axis has unit source stride, each line is summed with eight independent partial sums. Otherwise lines that are adjacent in memory are processed as a block and each line is summed sequentially. The rounding can therefore differ between layouts, but the result never depends on the execution context or thread count.
An empty axis or an empty set of lines writes nothing.
§Examples
use strided_basic::{
ErasedNormPlan, ErasedRawStridedMut, ErasedRawStridedRef, ExecContext, KernelDType,
NormSpec,
};
// Two feature-first rows of width 2: (d = 2, len = 2), normalized over d.
let x = [1.0_f64, 3.0, -2.0, 2.0];
let weight = [2.0_f64, 2.0];
let mut y = [0.0_f64; 4];
let spec = NormSpec::layer_norm(0.0).with_weight(1);
let plan =
ErasedNormPlan::compile(KernelDType::F64, spec, &[2, 2], &[1, 2], &[1, 2], 0).unwrap();
let x = ErasedRawStridedRef::from_slice(&x, &[2, 2], &[1, 2], 0).unwrap();
let w = ErasedRawStridedRef::from_slice(&weight, &[2], &[1], 0).unwrap();
let mut out = ErasedRawStridedMut::from_slice_mut(&mut y, &[2, 2], &[1, 2], 0).unwrap();
plan.execute(&ExecContext::serial(), &mut out, &x, Some(&w), None).unwrap();
assert_eq!(y, [-2.0, 2.0, -2.0, 2.0]);Implementations§
Source§impl ErasedNormPlan
impl ErasedNormPlan
Sourcepub fn compile(
dtype: KernelDType,
spec: NormSpec,
dims: &[usize],
src_strides: &[isize],
dest_strides: &[isize],
axis: usize,
) -> Result<Self>
pub fn compile( dtype: KernelDType, spec: NormSpec, dims: &[usize], src_strides: &[isize], dest_strides: &[isize], axis: usize, ) -> Result<Self>
Validate and store a normalization plan for fixed layouts.
dims are the shared source and destination dimensions; axis is the
normalized axis.
§Examples
use strided_basic::{ErasedNormPlan, KernelDType, NormSpec};
let spec = NormSpec::rms_norm(1e-6).with_weight(1);
assert!(ErasedNormPlan::compile(KernelDType::F32, spec, &[8, 3], &[1, 8], &[1, 8], 0).is_ok());
// Complex and integer input is rejected, as is a negative eps.
assert!(ErasedNormPlan::compile(KernelDType::C64, spec, &[8], &[1], &[1], 0).is_err());
let bad = NormSpec::rms_norm(-1.0);
assert!(ErasedNormPlan::compile(KernelDType::F32, bad, &[8], &[1], &[1], 0).is_err());§Errors
UnsupportedDTypefor a dtype other thanf32/f64;UnsupportedOpfor a negative, infinite or NaNeps;InvalidAxis,StrideLengthMismatchandNonInjectiveOutputLayoutfor inconsistent layouts;OffsetOverflowwhen a layout’s offsets are not representable.
Sourcepub fn dtype(&self) -> KernelDType
pub fn dtype(&self) -> KernelDType
Element dtype.
§Examples
use strided_basic::{ErasedNormPlan, KernelDType, NormSpec};
let plan = ErasedNormPlan::compile(
KernelDType::F64, NormSpec::layer_norm(0.0), &[4], &[1], &[1], 0,
)
.unwrap();
assert_eq!(plan.dtype(), KernelDType::F64);Sourcepub fn spec(&self) -> NormSpec
pub fn spec(&self) -> NormSpec
Kind, eps and affine parameter layout.
§Examples
use strided_basic::{ErasedNormPlan, KernelDType, NormSpec};
let spec = NormSpec::layer_norm(1e-5);
let plan = ErasedNormPlan::compile(KernelDType::F64, spec, &[4], &[1], &[1], 0).unwrap();
assert_eq!(plan.spec(), spec);Sourcepub fn execute(
&self,
ctx: &ExecContext,
dest: &mut ErasedRawStridedMut<'_>,
src: &ErasedRawStridedRef<'_>,
weight: Option<&ErasedRawStridedRef<'_>>,
bias: Option<&ErasedRawStridedRef<'_>>,
) -> Result<()>
pub fn execute( &self, ctx: &ExecContext, dest: &mut ErasedRawStridedMut<'_>, src: &ErasedRawStridedRef<'_>, weight: Option<&ErasedRawStridedRef<'_>>, bias: Option<&ErasedRawStridedRef<'_>>, ) -> Result<()>
Execute into an initialized destination.
weight and bias must be present exactly when the plan’s
NormSpec requests them, as rank-one descriptors of the axis length
with the recorded stride.
§Examples
use strided_basic::{
ErasedNormPlan, ErasedRawStridedMut, ErasedRawStridedRef, ExecContext, KernelDType,
NormSpec,
};
let x = [3.0_f32, 4.0, 0.0, 0.0];
let bias = [1.0_f32, 1.0];
let mut y = [0.0_f32; 4];
let spec = NormSpec::rms_norm(0.0).with_bias(1);
let plan =
ErasedNormPlan::compile(KernelDType::F32, spec, &[2, 2], &[1, 2], &[1, 2], 0).unwrap();
let x = ErasedRawStridedRef::from_slice(&x, &[2, 2], &[1, 2], 0).unwrap();
let b = ErasedRawStridedRef::from_slice(&bias, &[2], &[1], 0).unwrap();
let mut out = ErasedRawStridedMut::from_slice_mut(&mut y, &[2, 2], &[1, 2], 0).unwrap();
plan.execute(&ExecContext::serial(), &mut out, &x, None, Some(&b)).unwrap();
// Row 0: rms = sqrt(12.5); the all-zero row 1 is 0 * inf = NaN with eps = 0.
assert!((y[0] - (1.0 + 3.0 / 12.5f32.sqrt())).abs() < 1e-6);
assert!(y[2].is_nan() && y[3].is_nan());§Errors
Returns DTypeMismatch or PlanLayoutMismatch when a descriptor does
not match the plan, before any destination write.
Sourcepub fn execute_uninit(
&self,
ctx: &ExecContext,
dest: &mut ErasedRawStridedUninitMut<'_>,
src: &ErasedRawStridedPtr<'_>,
weight: Option<&ErasedRawStridedPtr<'_>>,
bias: Option<&ErasedRawStridedPtr<'_>>,
) -> Result<()>
pub fn execute_uninit( &self, ctx: &ExecContext, dest: &mut ErasedRawStridedUninitMut<'_>, src: &ErasedRawStridedPtr<'_>, weight: Option<&ErasedRawStridedPtr<'_>>, bias: Option<&ErasedRawStridedPtr<'_>>, ) -> Result<()>
Execute into an uninitialized destination.
On success every reachable destination element is written; validation errors, including any overlap between an input and the destination allocation, are returned before any write.
§Examples
use core::mem::MaybeUninit;
use strided_basic::{
ErasedNormPlan, ErasedRawStridedPtr, ErasedRawStridedRef, ErasedRawStridedUninitMut,
ExecContext, KernelDType, NormSpec,
};
let x = [1.0_f64, 1.0, 1.0];
let mut y = [MaybeUninit::<f64>::uninit(); 3];
let plan = ErasedNormPlan::compile(
KernelDType::F64, NormSpec::layer_norm(1e-5), &[3], &[1], &[1], 0,
)
.unwrap();
let x = ErasedRawStridedRef::from_slice(&x, &[3], &[1], 0).unwrap();
let x = ErasedRawStridedPtr::from_ref(&x);
let mut out = ErasedRawStridedUninitMut::from_uninit_slice(&mut y, &[3], &[1], 0).unwrap();
plan.execute_uninit(&ExecContext::serial(), &mut out, &x, None, None).unwrap();
// Zero variance with eps > 0 normalizes to exactly zero.
assert!(y.iter().all(|v| unsafe { v.assume_init() } == 0.0));§Errors
As Self::execute, plus OverlappingInputOutput naming input 0
(source), 1 (weight) or 2 (bias).