Skip to main content

ErasedNormPlan

Struct ErasedNormPlan 

Source
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 + bias with mean = sum(x) / n and the biased (population) variance var = sum((x - mean)^2) / n, as in PyTorch’s layer_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

Source

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
  • UnsupportedDType for a dtype other than f32 / f64;
  • UnsupportedOp for a negative, infinite or NaN eps;
  • InvalidAxis, StrideLengthMismatch and NonInjectiveOutputLayout for inconsistent layouts;
  • OffsetOverflow when a layout’s offsets are not representable.
Source

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);
Source

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);
Source

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.

Source

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).

Trait Implementations§

Source§

impl Clone for ErasedNormPlan

Source§

fn clone(&self) -> Self

Returns a duplicate of the value. Read more
1.0.0 (const: unstable) · Source§

fn clone_from(&mut self, source: &Self)

Performs copy-assignment from source. Read more
Source§

impl Debug for ErasedNormPlan

Source§

fn fmt(&self, f: &mut Formatter<'_>) -> Result

Formats the value using the given formatter. Read more

Auto Trait Implementations§

Blanket Implementations§

Source§

impl<T> Any for T
where T: 'static + ?Sized,

Source§

fn type_id(&self) -> TypeId

Gets the TypeId of self. Read more
Source§

impl<T> Borrow<T> for T
where T: ?Sized,

Source§

fn borrow(&self) -> &T

Immutably borrows from an owned value. Read more
Source§

impl<T> BorrowMut<T> for T
where T: ?Sized,

Source§

fn borrow_mut(&mut self) -> &mut T

Mutably borrows from an owned value. Read more
Source§

impl<T> CloneToUninit for T
where T: Clone,

Source§

unsafe fn clone_to_uninit(&self, dest: *mut u8)

🔬This is a nightly-only experimental API. (clone_to_uninit)
Performs copy-assignment from self to dest. Read more
Source§

impl<T> From<T> for T

Source§

fn from(t: T) -> T

Returns the argument unchanged.

Source§

impl<T, U> Into<U> for T
where U: From<T>,

Source§

fn into(self) -> U

Calls U::from(self).

That is, this conversion is whatever the implementation of From<T> for U chooses to do.

Source§

impl<T> MaybeSend for T

Source§

impl<T> MaybeSendSync for T

Source§

impl<T> MaybeSync for T

Source§

impl<T> ToOwned for T
where T: Clone,

Source§

type Owned = T

The resulting type after obtaining ownership.
Source§

fn to_owned(&self) -> T

Creates owned data from borrowed data, usually by cloning. Read more
Source§

fn clone_into(&self, target: &mut T)

Uses borrowed data to replace owned data, usually by cloning. Read more
Source§

impl<T, U> TryFrom<U> for T
where U: Into<T>,

Source§

type Error = !

The type returned in the event of a conversion error.
Source§

fn try_from(value: U) -> Result<T, !>

Performs the conversion.
Source§

impl<T, U> TryInto<U> for T
where U: TryFrom<T>,

Source§

type Error = <U as TryFrom<T>>::Error

The type returned in the event of a conversion error.
Source§

fn try_into(self) -> Result<U, <U as TryFrom<T>>::Error>

Performs the conversion.