Skip to main content

strided_basic/erased/
norm.rs

1//! Dtype-erased fused `layer_norm` / `rms_norm` along one axis.
2
3use super::line::{for_each_unit, LineLayout, UnitKernel, UnitPtr, PANEL};
4use super::{
5    check_reduce_layout_offset_arithmetic, checked_total_len, reduce_uninit_writer, reduce_writer,
6    ReduceWriter,
7};
8use crate::erased_common::{check_dtype, validate_uninit_no_overlap};
9use crate::*;
10
11/// Normalization computed by an [`ErasedNormPlan`].
12///
13/// # Examples
14///
15/// ```
16/// use strided_basic::NormKind;
17/// assert_ne!(NormKind::Layer, NormKind::Rms);
18/// ```
19#[non_exhaustive]
20#[derive(Clone, Copy, Debug, Eq, PartialEq)]
21pub enum NormKind {
22    /// `y = (x - mean) / sqrt(var + eps)` with the biased (population)
23    /// variance `var = mean((x - mean)^2)`.
24    Layer,
25    /// `y = x / sqrt(mean(x^2) + eps)`.
26    Rms,
27}
28
29/// Kind, `eps` and optional affine parameters of an [`ErasedNormPlan`].
30///
31/// The optional weight (scale) and bias (shift) are vectors along the
32/// normalized axis, each with its own element stride; the result is
33/// `y * weight + bias`.
34///
35/// # Examples
36///
37/// ```
38/// use strided_basic::{NormKind, NormSpec};
39/// let spec = NormSpec::layer_norm(1e-5).with_weight(1).with_bias(-1);
40/// assert_eq!(spec.kind(), NormKind::Layer);
41/// assert_eq!(spec.eps(), 1e-5);
42/// assert_eq!(spec.weight_stride(), Some(1));
43/// assert_eq!(spec.bias_stride(), Some(-1));
44/// assert_eq!(NormSpec::rms_norm(0.0).weight_stride(), None);
45/// ```
46#[non_exhaustive]
47#[derive(Clone, Copy, Debug, PartialEq)]
48pub struct NormSpec {
49    kind: NormKind,
50    eps: f64,
51    weight_stride: Option<isize>,
52    bias_stride: Option<isize>,
53}
54
55impl NormSpec {
56    /// Layer normalization with the given `eps` and no affine parameters.
57    ///
58    /// # Examples
59    ///
60    /// ```
61    /// use strided_basic::{NormKind, NormSpec};
62    /// assert_eq!(NormSpec::layer_norm(1e-6).kind(), NormKind::Layer);
63    /// ```
64    #[inline]
65    pub const fn layer_norm(eps: f64) -> Self {
66        Self {
67            kind: NormKind::Layer,
68            eps,
69            weight_stride: None,
70            bias_stride: None,
71        }
72    }
73
74    /// RMS normalization with the given `eps` and no affine parameters.
75    ///
76    /// # Examples
77    ///
78    /// ```
79    /// use strided_basic::{NormKind, NormSpec};
80    /// assert_eq!(NormSpec::rms_norm(1e-6).kind(), NormKind::Rms);
81    /// ```
82    #[inline]
83    pub const fn rms_norm(eps: f64) -> Self {
84        Self {
85            kind: NormKind::Rms,
86            eps,
87            weight_stride: None,
88            bias_stride: None,
89        }
90    }
91
92    /// Multiply by a weight vector with the given element stride.
93    ///
94    /// # Examples
95    ///
96    /// ```
97    /// use strided_basic::NormSpec;
98    /// assert_eq!(NormSpec::rms_norm(0.0).with_weight(2).weight_stride(), Some(2));
99    /// ```
100    #[inline]
101    pub const fn with_weight(mut self, stride: isize) -> Self {
102        self.weight_stride = Some(stride);
103        self
104    }
105
106    /// Add a bias vector with the given element stride.
107    ///
108    /// # Examples
109    ///
110    /// ```
111    /// use strided_basic::NormSpec;
112    /// assert_eq!(NormSpec::rms_norm(0.0).with_bias(1).bias_stride(), Some(1));
113    /// ```
114    #[inline]
115    pub const fn with_bias(mut self, stride: isize) -> Self {
116        self.bias_stride = Some(stride);
117        self
118    }
119
120    /// Normalization kind.
121    ///
122    /// # Examples
123    ///
124    /// ```
125    /// use strided_basic::{NormKind, NormSpec};
126    /// assert_eq!(NormSpec::rms_norm(0.0).kind(), NormKind::Rms);
127    /// ```
128    #[inline]
129    pub const fn kind(&self) -> NormKind {
130        self.kind
131    }
132
133    /// The `eps` added to the variance (or mean square) before the square root.
134    ///
135    /// # Examples
136    ///
137    /// ```
138    /// use strided_basic::NormSpec;
139    /// assert_eq!(NormSpec::rms_norm(0.25).eps(), 0.25);
140    /// ```
141    #[inline]
142    pub const fn eps(&self) -> f64 {
143        self.eps
144    }
145
146    /// Element stride of the weight vector, if any.
147    ///
148    /// # Examples
149    ///
150    /// ```
151    /// use strided_basic::NormSpec;
152    /// assert_eq!(NormSpec::layer_norm(0.0).weight_stride(), None);
153    /// ```
154    #[inline]
155    pub const fn weight_stride(&self) -> Option<isize> {
156        self.weight_stride
157    }
158
159    /// Element stride of the bias vector, if any.
160    ///
161    /// # Examples
162    ///
163    /// ```
164    /// use strided_basic::NormSpec;
165    /// assert_eq!(NormSpec::layer_norm(0.0).bias_stride(), None);
166    /// ```
167    #[inline]
168    pub const fn bias_stride(&self) -> Option<isize> {
169        self.bias_stride
170    }
171}
172
173/// Dtype-erased fused `layer_norm` / `rms_norm` along one axis.
174///
175/// Each line along `axis` is normalized independently in one call:
176///
177/// * [`NormKind::Layer`]: `y = (x - mean) * rsqrt(var + eps) * weight + bias`
178///   with `mean = sum(x) / n` and the biased (population) variance
179///   `var = sum((x - mean)^2) / n`, as in PyTorch's `layer_norm`.
180/// * [`NormKind::Rms`]: `y = x * rsqrt(sum(x^2) / n + eps) * weight + bias`.
181///
182/// `rsqrt(v)` is evaluated as `1 / sqrt(v)` once per line, and the weight and
183/// bias terms are present only when the [`NormSpec`] requests them. The
184/// destination has the source dimensions with independent strides. Supported
185/// dtypes are `f32` and `f64`; every accumulation is in the element dtype.
186///
187/// # Numerics
188///
189/// Layer norm uses the shifted-data two-pass algorithm: every element is first
190/// shifted by the first element of its line, then the mean of the shifted
191/// values and the variance about that mean are computed in two passes. This
192/// avoids the cancellation of the one-pass `E[x^2] - E[x]^2` form and of a
193/// large common offset. A line whose elements are all equal has a variance of
194/// exactly zero and normalizes to exactly `0 * rsqrt(eps) * weight + bias`,
195/// i.e. `bias` (or zero) whenever `eps > 0`. With `eps == 0` such a line
196/// evaluates `0 * inf` and is `NaN`; that is the caller's choice of `eps`, not
197/// an error.
198///
199/// Non-finite inputs follow IEEE evaluation of the formula. A NaN anywhere in
200/// a line makes every output of that line NaN. An infinite element makes a
201/// layer-norm line NaN (its deviation is `inf - inf`), and makes an RMS-norm
202/// line zero at its finite elements and NaN at the infinite ones
203/// (`rsqrt(inf) = 0`). Squares are not rescaled: when the sum of squared
204/// deviations overflows to `inf` (finite deviations above about `1.8e19` for
205/// `f32` or `1.3e154` for `f64`), the line's finite elements normalize to zero.
206///
207/// When the normalized axis has unit source stride, each line is summed with
208/// eight independent partial sums. Otherwise lines that are adjacent in memory
209/// are processed as a block and each line is summed sequentially. The rounding
210/// can therefore differ between layouts, but the result never depends on the
211/// execution context or thread count.
212///
213/// An empty axis or an empty set of lines writes nothing.
214///
215/// # Examples
216///
217/// ```
218/// use strided_basic::{
219///     ErasedNormPlan, ErasedRawStridedMut, ErasedRawStridedRef, ExecContext, KernelDType,
220///     NormSpec,
221/// };
222///
223/// // Two feature-first rows of width 2: (d = 2, len = 2), normalized over d.
224/// let x = [1.0_f64, 3.0, -2.0, 2.0];
225/// let weight = [2.0_f64, 2.0];
226/// let mut y = [0.0_f64; 4];
227/// let spec = NormSpec::layer_norm(0.0).with_weight(1);
228/// let plan =
229///     ErasedNormPlan::compile(KernelDType::F64, spec, &[2, 2], &[1, 2], &[1, 2], 0).unwrap();
230/// let x = ErasedRawStridedRef::from_slice(&x, &[2, 2], &[1, 2], 0).unwrap();
231/// let w = ErasedRawStridedRef::from_slice(&weight, &[2], &[1], 0).unwrap();
232/// let mut out = ErasedRawStridedMut::from_slice_mut(&mut y, &[2, 2], &[1, 2], 0).unwrap();
233/// plan.execute(&ExecContext::serial(), &mut out, &x, Some(&w), None).unwrap();
234/// assert_eq!(y, [-2.0, 2.0, -2.0, 2.0]);
235/// ```
236#[derive(Clone, Debug)]
237pub struct ErasedNormPlan {
238    dtype: KernelDType,
239    spec: NormSpec,
240    dims: Vec<usize>,
241    src_strides: Vec<isize>,
242    dest_strides: Vec<isize>,
243    layout: LineLayout,
244}
245
246impl ErasedNormPlan {
247    /// Validate and store a normalization plan for fixed layouts.
248    ///
249    /// `dims` are the shared source and destination dimensions; `axis` is the
250    /// normalized axis.
251    ///
252    /// # Examples
253    ///
254    /// ```
255    /// use strided_basic::{ErasedNormPlan, KernelDType, NormSpec};
256    /// let spec = NormSpec::rms_norm(1e-6).with_weight(1);
257    /// assert!(ErasedNormPlan::compile(KernelDType::F32, spec, &[8, 3], &[1, 8], &[1, 8], 0).is_ok());
258    /// // Complex and integer input is rejected, as is a negative eps.
259    /// assert!(ErasedNormPlan::compile(KernelDType::C64, spec, &[8], &[1], &[1], 0).is_err());
260    /// let bad = NormSpec::rms_norm(-1.0);
261    /// assert!(ErasedNormPlan::compile(KernelDType::F32, bad, &[8], &[1], &[1], 0).is_err());
262    /// ```
263    ///
264    /// # Errors
265    ///
266    /// * `UnsupportedDType` for a dtype other than `f32` / `f64`;
267    /// * `UnsupportedOp` for a negative, infinite or NaN `eps`;
268    /// * `InvalidAxis`, `StrideLengthMismatch` and `NonInjectiveOutputLayout`
269    ///   for inconsistent layouts;
270    /// * `OffsetOverflow` when a layout's offsets are not representable.
271    pub fn compile(
272        dtype: KernelDType,
273        spec: NormSpec,
274        dims: &[usize],
275        src_strides: &[isize],
276        dest_strides: &[isize],
277        axis: usize,
278    ) -> Result<Self> {
279        if !matches!(dtype, KernelDType::F32 | KernelDType::F64) {
280            return Err(StridedError::UnsupportedDType {
281                dtype: dtype.label(),
282            });
283        }
284        if !(spec.eps.is_finite() && spec.eps >= 0.0) {
285            return Err(StridedError::UnsupportedOp {
286                op: "norm with a negative or non-finite eps",
287                dtype: dtype.label(),
288            });
289        }
290        if dims.len() != src_strides.len() || dims.len() != dest_strides.len() {
291            return Err(StridedError::StrideLengthMismatch);
292        }
293        if axis >= dims.len() {
294            return Err(StridedError::InvalidAxis {
295                axis,
296                rank: dims.len(),
297            });
298        }
299        checked_total_len(dims)?;
300        check_reduce_layout_offset_arithmetic(dims, src_strides)?;
301        check_reduce_layout_offset_arithmetic(dims, dest_strides)?;
302        for stride in [spec.weight_stride, spec.bias_stride].into_iter().flatten() {
303            check_reduce_layout_offset_arithmetic(&dims[axis..=axis], &[stride])?;
304        }
305        if !crate::layout_check::is_injective_layout(dims, dest_strides) {
306            return Err(StridedError::NonInjectiveOutputLayout);
307        }
308        let dest_outer: Vec<isize> = (0..dims.len())
309            .filter(|&a| a != axis)
310            .map(|a| dest_strides[a])
311            .collect();
312        let layout = LineLayout::compile(
313            dims,
314            src_strides,
315            &dest_outer,
316            dest_strides[axis],
317            axis,
318            true,
319        )?;
320        Ok(Self {
321            dtype,
322            spec,
323            dims: dims.to_vec(),
324            src_strides: src_strides.to_vec(),
325            dest_strides: dest_strides.to_vec(),
326            layout,
327        })
328    }
329
330    /// Element dtype.
331    ///
332    /// # Examples
333    ///
334    /// ```
335    /// use strided_basic::{ErasedNormPlan, KernelDType, NormSpec};
336    /// let plan = ErasedNormPlan::compile(
337    ///     KernelDType::F64, NormSpec::layer_norm(0.0), &[4], &[1], &[1], 0,
338    /// )
339    /// .unwrap();
340    /// assert_eq!(plan.dtype(), KernelDType::F64);
341    /// ```
342    #[inline]
343    pub fn dtype(&self) -> KernelDType {
344        self.dtype
345    }
346
347    /// Kind, `eps` and affine parameter layout.
348    ///
349    /// # Examples
350    ///
351    /// ```
352    /// use strided_basic::{ErasedNormPlan, KernelDType, NormSpec};
353    /// let spec = NormSpec::layer_norm(1e-5);
354    /// let plan = ErasedNormPlan::compile(KernelDType::F64, spec, &[4], &[1], &[1], 0).unwrap();
355    /// assert_eq!(plan.spec(), spec);
356    /// ```
357    #[inline]
358    pub fn spec(&self) -> NormSpec {
359        self.spec
360    }
361
362    fn check_layouts(
363        &self,
364        dest_dims: &[usize],
365        dest_strides: &[isize],
366        src: &ErasedRawStridedRef<'_>,
367        weight: Option<&ErasedRawStridedRef<'_>>,
368        bias: Option<&ErasedRawStridedRef<'_>>,
369    ) -> Result<()> {
370        if src.dims() != self.dims.as_slice()
371            || src.strides() != self.src_strides.as_slice()
372            || dest_dims != self.dims.as_slice()
373            || dest_strides != self.dest_strides.as_slice()
374        {
375            return Err(StridedError::PlanLayoutMismatch);
376        }
377        let n = self.layout.axis_len;
378        for (param, stride) in [
379            (weight, self.spec.weight_stride),
380            (bias, self.spec.bias_stride),
381        ] {
382            match (param, stride) {
383                (None, None) => {}
384                (Some(param), Some(stride)) => {
385                    check_dtype(self.dtype, param.dtype())?;
386                    if param.dims() != [n] || param.strides() != [stride] {
387                        return Err(StridedError::PlanLayoutMismatch);
388                    }
389                }
390                _ => return Err(StridedError::PlanLayoutMismatch),
391            }
392        }
393        Ok(())
394    }
395
396    /// Execute into an initialized destination.
397    ///
398    /// `weight` and `bias` must be present exactly when the plan's
399    /// [`NormSpec`] requests them, as rank-one descriptors of the axis length
400    /// with the recorded stride.
401    ///
402    /// # Examples
403    ///
404    /// ```
405    /// use strided_basic::{
406    ///     ErasedNormPlan, ErasedRawStridedMut, ErasedRawStridedRef, ExecContext, KernelDType,
407    ///     NormSpec,
408    /// };
409    /// let x = [3.0_f32, 4.0, 0.0, 0.0];
410    /// let bias = [1.0_f32, 1.0];
411    /// let mut y = [0.0_f32; 4];
412    /// let spec = NormSpec::rms_norm(0.0).with_bias(1);
413    /// let plan =
414    ///     ErasedNormPlan::compile(KernelDType::F32, spec, &[2, 2], &[1, 2], &[1, 2], 0).unwrap();
415    /// let x = ErasedRawStridedRef::from_slice(&x, &[2, 2], &[1, 2], 0).unwrap();
416    /// let b = ErasedRawStridedRef::from_slice(&bias, &[2], &[1], 0).unwrap();
417    /// let mut out = ErasedRawStridedMut::from_slice_mut(&mut y, &[2, 2], &[1, 2], 0).unwrap();
418    /// plan.execute(&ExecContext::serial(), &mut out, &x, None, Some(&b)).unwrap();
419    /// // Row 0: rms = sqrt(12.5); the all-zero row 1 is 0 * inf = NaN with eps = 0.
420    /// assert!((y[0] - (1.0 + 3.0 / 12.5f32.sqrt())).abs() < 1e-6);
421    /// assert!(y[2].is_nan() && y[3].is_nan());
422    /// ```
423    ///
424    /// # Errors
425    ///
426    /// Returns `DTypeMismatch` or `PlanLayoutMismatch` when a descriptor does
427    /// not match the plan, before any destination write.
428    pub fn execute(
429        &self,
430        ctx: &ExecContext,
431        dest: &mut ErasedRawStridedMut<'_>,
432        src: &ErasedRawStridedRef<'_>,
433        weight: Option<&ErasedRawStridedRef<'_>>,
434        bias: Option<&ErasedRawStridedRef<'_>>,
435    ) -> Result<()> {
436        check_dtype(self.dtype, dest.dtype())?;
437        check_dtype(self.dtype, src.dtype())?;
438        self.check_layouts(dest.dims(), dest.strides(), src, weight, bias)?;
439        match self.dtype {
440            KernelDType::F32 => {
441                let mut writer = reduce_writer::<f32>(dest)?;
442                self.dispatch::<f32, _>(ctx, &mut writer, src, weight, bias)
443            }
444            _ => {
445                let mut writer = reduce_writer::<f64>(dest)?;
446                self.dispatch::<f64, _>(ctx, &mut writer, src, weight, bias)
447            }
448        }
449    }
450
451    /// Execute into an uninitialized destination.
452    ///
453    /// On success every reachable destination element is written; validation
454    /// errors, including any overlap between an input and the destination
455    /// allocation, are returned before any write.
456    ///
457    /// # Examples
458    ///
459    /// ```
460    /// use core::mem::MaybeUninit;
461    /// use strided_basic::{
462    ///     ErasedNormPlan, ErasedRawStridedPtr, ErasedRawStridedRef, ErasedRawStridedUninitMut,
463    ///     ExecContext, KernelDType, NormSpec,
464    /// };
465    /// let x = [1.0_f64, 1.0, 1.0];
466    /// let mut y = [MaybeUninit::<f64>::uninit(); 3];
467    /// let plan = ErasedNormPlan::compile(
468    ///     KernelDType::F64, NormSpec::layer_norm(1e-5), &[3], &[1], &[1], 0,
469    /// )
470    /// .unwrap();
471    /// let x = ErasedRawStridedRef::from_slice(&x, &[3], &[1], 0).unwrap();
472    /// let x = ErasedRawStridedPtr::from_ref(&x);
473    /// let mut out = ErasedRawStridedUninitMut::from_uninit_slice(&mut y, &[3], &[1], 0).unwrap();
474    /// plan.execute_uninit(&ExecContext::serial(), &mut out, &x, None, None).unwrap();
475    /// // Zero variance with eps > 0 normalizes to exactly zero.
476    /// assert!(y.iter().all(|v| unsafe { v.assume_init() } == 0.0));
477    /// ```
478    ///
479    /// # Errors
480    ///
481    /// As [`Self::execute`], plus `OverlappingInputOutput` naming input `0`
482    /// (source), `1` (weight) or `2` (bias).
483    pub fn execute_uninit(
484        &self,
485        ctx: &ExecContext,
486        dest: &mut ErasedRawStridedUninitMut<'_>,
487        src: &ErasedRawStridedPtr<'_>,
488        weight: Option<&ErasedRawStridedPtr<'_>>,
489        bias: Option<&ErasedRawStridedPtr<'_>>,
490    ) -> Result<()> {
491        check_dtype(self.dtype, dest.dtype())?;
492        check_dtype(self.dtype, src.dtype())?;
493        validate_uninit_no_overlap(dest, src, 0)?;
494        if let Some(weight) = weight {
495            validate_uninit_no_overlap(dest, weight, 1)?;
496        }
497        if let Some(bias) = bias {
498            validate_uninit_no_overlap(dest, bias, 2)?;
499        }
500        // SAFETY: the owning erased entry rejected all input/output overlap
501        // before each conversion.
502        let src = unsafe { src.try_as_ref_after_no_overlap() }?;
503        let weight = weight
504            .map(|weight| unsafe { weight.try_as_ref_after_no_overlap() })
505            .transpose()?;
506        let bias = bias
507            .map(|bias| unsafe { bias.try_as_ref_after_no_overlap() })
508            .transpose()?;
509        self.check_layouts(
510            dest.dims(),
511            dest.strides(),
512            &src,
513            weight.as_ref(),
514            bias.as_ref(),
515        )?;
516        match self.dtype {
517            KernelDType::F32 => {
518                let mut writer = reduce_uninit_writer::<f32>(dest)?;
519                self.dispatch::<f32, _>(ctx, &mut writer, &src, weight.as_ref(), bias.as_ref())
520            }
521            _ => {
522                let mut writer = reduce_uninit_writer::<f64>(dest)?;
523                self.dispatch::<f64, _>(ctx, &mut writer, &src, weight.as_ref(), bias.as_ref())
524            }
525        }
526    }
527
528    fn dispatch<T, W>(
529        &self,
530        ctx: &ExecContext,
531        dest: &mut W,
532        src: &ErasedRawStridedRef<'_>,
533        weight: Option<&ErasedRawStridedRef<'_>>,
534        bias: Option<&ErasedRawStridedRef<'_>>,
535    ) -> Result<()>
536    where
537        T: NormScalar,
538        W: ReduceWriter<T>,
539    {
540        let affine = Affine {
541            weight: param::<T>(weight)?,
542            bias: param::<T>(bias)?,
543        };
544        // Kind and affine presence are matched once per execution; each loop
545        // is monomorphized for one combination.
546        macro_rules! go {
547            ($layer:literal) => {
548                match (weight.is_some(), bias.is_some()) {
549                    (false, false) => {
550                        self.run::<T, W, $layer, false, false>(ctx, dest, src, affine)
551                    }
552                    (true, false) => self.run::<T, W, $layer, true, false>(ctx, dest, src, affine),
553                    (false, true) => self.run::<T, W, $layer, false, true>(ctx, dest, src, affine),
554                    (true, true) => self.run::<T, W, $layer, true, true>(ctx, dest, src, affine),
555                }
556            };
557        }
558        match self.spec.kind {
559            NormKind::Layer => go!(true),
560            NormKind::Rms => go!(false),
561        }
562    }
563
564    fn run<T, W, const LAYER: bool, const WEIGHT: bool, const BIAS: bool>(
565        &self,
566        ctx: &ExecContext,
567        dest: &mut W,
568        src: &ErasedRawStridedRef<'_>,
569        affine: Affine<T>,
570    ) -> Result<()>
571    where
572        T: NormScalar,
573        W: ReduceWriter<T>,
574    {
575        let layout = &self.layout;
576        let source = UnitPtr(src.data_as::<T>()?.as_ptr() as *mut T);
577        // SAFETY: the validated writer owns the destination allocation.
578        let target = UnitPtr(unsafe { dest.ptr() });
579        let line = Line {
580            n: layout.axis_len,
581            n_t: T::from_usize(layout.axis_len),
582            eps: T::from_f64(self.spec.eps),
583            ss: layout.src_axis_stride,
584            ds: layout.dest_axis_stride,
585        };
586        let kernel = NormUnit::<T, LAYER, WEIGHT, BIAS> {
587            source,
588            target,
589            line,
590            affine,
591        };
592        // INVARIANT: (1) compile checked the signed source, destination,
593        // weight and bias spans and every cursor step/reset; (2) the raw
594        // descriptors validated every reachable offset; (3) execute checked
595        // exact plan-layout equality for all four descriptors. Units cover
596        // disjoint lines, so their destination writes are disjoint.
597        // SAFETY: the three-link invariant above.
598        unsafe { for_each_unit(ctx, layout, src.offset(), dest.offset(), kernel) }
599    }
600}
601
602/// Per-unit normalization kernel of one execution.
603#[derive(Clone, Copy)]
604struct NormUnit<T, const LAYER: bool, const WEIGHT: bool, const BIAS: bool> {
605    source: UnitPtr<T>,
606    target: UnitPtr<T>,
607    line: Line<T>,
608    affine: Affine<T>,
609}
610
611impl<T: NormScalar, const LAYER: bool, const WEIGHT: bool, const BIAS: bool> UnitKernel
612    for NormUnit<T, LAYER, WEIGHT, BIAS>
613{
614    #[inline(always)]
615    unsafe fn unit(self, so: isize, d_o: isize, width: usize) {
616        let Self {
617            source,
618            target,
619            line,
620            affine,
621        } = self;
622        // SAFETY: the caller passes offsets of the validated layout.
623        unsafe {
624            if width == 1 {
625                norm_line::<T, LAYER, WEIGHT, BIAS>(
626                    source.get(),
627                    so,
628                    target.get(),
629                    d_o,
630                    line,
631                    affine,
632                )
633            } else {
634                norm_panel::<T, LAYER, WEIGHT, BIAS>(
635                    source.get(),
636                    so,
637                    target.get(),
638                    d_o,
639                    width,
640                    line,
641                    affine,
642                )
643            }
644        }
645    }
646}
647
648/// Base pointer, offset and stride of an optional affine vector.
649#[derive(Clone, Copy)]
650struct Param<T> {
651    ptr: UnitPtr<T>,
652    offset: isize,
653    stride: isize,
654}
655
656#[derive(Clone, Copy)]
657struct Affine<T> {
658    weight: Option<Param<T>>,
659    bias: Option<Param<T>>,
660}
661
662impl<T> Param<T> {
663    /// Element `k` of the vector.
664    ///
665    /// # Safety
666    ///
667    /// `k` must be below the validated vector length.
668    #[inline(always)]
669    unsafe fn at(self, k: usize) -> T
670    where
671        T: Copy,
672    {
673        // SAFETY: the caller guarantees `k` is in range of the validated vector.
674        unsafe {
675            self.ptr
676                .get()
677                .offset(self.offset + k as isize * self.stride)
678                .read()
679        }
680    }
681}
682
683fn param<T: NormScalar>(param: Option<&ErasedRawStridedRef<'_>>) -> Result<Option<Param<T>>> {
684    param
685        .map(|param| {
686            Ok(Param {
687                ptr: UnitPtr(param.data_as::<T>()?.as_ptr() as *mut T),
688                offset: param.offset(),
689                stride: param.strides()[0],
690            })
691        })
692        .transpose()
693}
694
695/// Per-plan line constants.
696#[derive(Clone, Copy)]
697struct Line<T> {
698    n: usize,
699    n_t: T,
700    eps: T,
701    ss: isize,
702    ds: isize,
703}
704
705/// Floating element type supported by the norm plan.
706pub(super) trait NormScalar:
707    KernelStorageElement
708    + MaybeSendSync
709    + PartialOrd
710    + core::ops::Add<Output = Self>
711    + core::ops::Sub<Output = Self>
712    + core::ops::Mul<Output = Self>
713    + core::ops::Div<Output = Self>
714{
715    const ZERO: Self;
716    const ONE: Self;
717    fn from_usize(value: usize) -> Self;
718    fn from_f64(value: f64) -> Self;
719    fn sqrt(self) -> Self;
720}
721
722macro_rules! impl_norm_scalar {
723    ($($ty:ty),*) => {$(
724        impl NormScalar for $ty {
725            const ZERO: Self = 0.0;
726            const ONE: Self = 1.0;
727            #[inline(always)]
728            fn from_usize(value: usize) -> Self {
729                value as $ty
730            }
731            #[inline(always)]
732            fn from_f64(value: f64) -> Self {
733                value as $ty
734            }
735            #[inline(always)]
736            fn sqrt(self) -> Self {
737                <$ty>::sqrt(self)
738            }
739        }
740    )*};
741}
742impl_norm_scalar!(f32, f64);
743
744/// Independent partial sums of the contiguous line kernels.
745const NORM_LANES: usize = 16;
746
747/// Sum of `map(x)` over a contiguous slice with [`NORM_LANES`] partial sums.
748#[inline(always)]
749fn lane_sum<T: NormScalar>(values: &[T], map: impl Fn(T) -> T) -> T {
750    let mut partial = [T::ZERO; NORM_LANES];
751    let mut chunks = values.chunks_exact(NORM_LANES);
752    for chunk in chunks.by_ref() {
753        for (partial, &value) in partial.iter_mut().zip(chunk) {
754            *partial = *partial + map(value);
755        }
756    }
757    let mut tail = T::ZERO;
758    for &value in chunks.remainder() {
759        tail = tail + map(value);
760    }
761    // A left fold of the partial sums: a pairwise tree here makes the
762    // vectorizer split the accumulators into two-lane groups.
763    partial.into_iter().fold(T::ZERO, |acc, value| acc + value) + tail
764}
765
766/// Sum of `map(x)` over a strided line, sequentially.
767///
768/// # Safety
769///
770/// Every `src.offset(so + k * ss)` for `k < n` must be readable.
771#[inline(always)]
772unsafe fn strided_sum<T: NormScalar>(
773    src: *const T,
774    so: isize,
775    ss: isize,
776    n: usize,
777    map: impl Fn(T) -> T,
778) -> T {
779    let mut sum = T::ZERO;
780    let mut offset = so;
781    for _ in 0..n {
782        // SAFETY: the caller guarantees every visited offset.
783        sum = sum + map(unsafe { src.offset(offset).read() });
784        offset += ss;
785    }
786    sum
787}
788
789/// Statistics of one line: `y = ((x - shift) - mean) * inv`.
790///
791/// Layer norm shifts every element by the line's first element before
792/// summing (the shifted-data two-pass algorithm): `mean` is the mean of the
793/// shifted values and the variance is taken about it. A line of equal
794/// elements therefore has shifted values, mean and variance of exactly zero.
795/// RMS norm uses `shift = mean = 0`.
796#[derive(Clone, Copy)]
797struct Stats<T> {
798    shift: T,
799    mean: T,
800    inv: T,
801}
802
803impl<T: NormScalar> Stats<T> {
804    #[inline(always)]
805    fn apply(self, x: T) -> T {
806        ((x - self.shift) - self.mean) * self.inv
807    }
808}
809
810/// Statistics of one line.
811///
812/// # Safety
813///
814/// Every `src.offset(so + k * line.ss)` for `k < line.n` must be readable, and
815/// `line.n > 0`.
816#[inline(always)]
817unsafe fn line_stats<T: NormScalar, const LAYER: bool>(
818    src: *const T,
819    so: isize,
820    line: Line<T>,
821) -> Stats<T> {
822    // SAFETY: the caller guarantees every visited offset and `n > 0`; a unit
823    // stride line is `n` contiguous elements.
824    unsafe {
825        let shift = if LAYER {
826            src.offset(so).read()
827        } else {
828            T::ZERO
829        };
830        let (mean, var) = if line.ss == 1 {
831            let values = core::slice::from_raw_parts(src.offset(so), line.n);
832            let mean = if LAYER {
833                lane_sum(values, |x| x - shift) / line.n_t
834            } else {
835                T::ZERO
836            };
837            let var = lane_sum(values, |x| {
838                let d = (x - shift) - mean;
839                d * d
840            }) / line.n_t;
841            (mean, var)
842        } else {
843            let mean = if LAYER {
844                strided_sum(src, so, line.ss, line.n, |x| x - shift) / line.n_t
845            } else {
846                T::ZERO
847            };
848            let var = strided_sum(src, so, line.ss, line.n, |x| {
849                let d = (x - shift) - mean;
850                d * d
851            }) / line.n_t;
852            (mean, var)
853        };
854        Stats {
855            shift,
856            mean,
857            inv: T::ONE / (var + line.eps).sqrt(),
858        }
859    }
860}
861
862/// Normalizes one line.
863///
864/// # Safety
865///
866/// Every source offset `so + k * line.ss` and destination offset
867/// `d_o + k * line.ds` for `k < line.n` must be valid, and the affine vectors
868/// must hold `line.n` elements; `line.n > 0`.
869#[inline(always)]
870unsafe fn norm_line<T: NormScalar, const LAYER: bool, const WEIGHT: bool, const BIAS: bool>(
871    src: *const T,
872    so: isize,
873    dst: *mut T,
874    d_o: isize,
875    line: Line<T>,
876    affine: Affine<T>,
877) {
878    // SAFETY: the caller guarantees every offset below.
879    unsafe {
880        let stats = line_stats::<T, LAYER>(src, so, line);
881        let src = src.offset(so);
882        let dst = dst.offset(d_o);
883        let (w, ws) = match affine.weight {
884            Some(p) if WEIGHT => (p.ptr.get().offset(p.offset) as *const T, p.stride),
885            _ => (src, 0),
886        };
887        let (b, bs) = match affine.bias {
888            Some(p) if BIAS => (p.ptr.get().offset(p.offset) as *const T, p.stride),
889            _ => (src, 0),
890        };
891        let unit = |stride: isize| stride == 1 || stride == 0;
892        if line.ss == 1 && line.ds == 1 && unit(ws) && unit(bs) {
893            // Literal unit strides let the compiler emit a contiguous loop.
894            norm_output::<T, WEIGHT, BIAS>(
895                src,
896                1,
897                dst,
898                1,
899                w,
900                ws.min(1),
901                b,
902                bs.min(1),
903                line.n,
904                stats,
905            );
906        } else {
907            norm_output::<T, WEIGHT, BIAS>(src, line.ss, dst, line.ds, w, ws, b, bs, line.n, stats);
908        }
909    }
910}
911
912/// Output pass of one line, indexing every operand from its base.
913///
914/// # Safety
915///
916/// `src + k * ss`, `dst + k * ds`, and (when present) `w + k * ws` and
917/// `b + k * bs` are valid for `k < n`.
918#[allow(clippy::too_many_arguments)]
919#[inline(always)]
920unsafe fn norm_output<T: NormScalar, const WEIGHT: bool, const BIAS: bool>(
921    src: *const T,
922    ss: isize,
923    dst: *mut T,
924    ds: isize,
925    w: *const T,
926    ws: isize,
927    b: *const T,
928    bs: isize,
929    n: usize,
930    stats: Stats<T>,
931) {
932    for k in 0..n as isize {
933        // SAFETY: the caller guarantees every offset.
934        unsafe {
935            let mut y = stats.apply(src.offset(k * ss).read());
936            if WEIGHT {
937                y = y * w.offset(k * ws).read();
938            }
939            if BIAS {
940                y = y + b.offset(k * bs).read();
941            }
942            // Raw write: the destination may be uninitialized.
943            dst.offset(k * ds).write(y);
944        }
945    }
946}
947
948/// Normalizes `width` adjacent lines whose elements are contiguous across the
949/// lines in both source and destination, with the statistics of
950/// [`line_stats`] per line (each line summed sequentially).
951///
952/// # Safety
953///
954/// As [`norm_line`] for each line `j < width` with source base `so + j` and
955/// destination base `d_o + j`; `width <= PANEL`.
956#[allow(clippy::too_many_arguments)]
957#[inline(always)]
958unsafe fn norm_panel<T: NormScalar, const LAYER: bool, const WEIGHT: bool, const BIAS: bool>(
959    src: *const T,
960    so: isize,
961    dst: *mut T,
962    d_o: isize,
963    width: usize,
964    line: Line<T>,
965    affine: Affine<T>,
966) {
967    debug_assert!(width <= PANEL);
968    let mut shift = [T::ZERO; PANEL];
969    let mut mean = [T::ZERO; PANEL];
970    let mut scale = [T::ZERO; PANEL];
971    let shift = &mut shift[..width];
972    let mean = &mut mean[..width];
973    let scale = &mut scale[..width];
974    // SAFETY: the caller guarantees `width` contiguous elements at every axis
975    // position of the source and destination, and `n > 0`.
976    unsafe {
977        let row =
978            |k: usize| core::slice::from_raw_parts(src.offset(so + k as isize * line.ss), width);
979        if LAYER {
980            shift.copy_from_slice(row(0));
981            for k in 0..line.n {
982                for ((mean, &shift), &x) in mean.iter_mut().zip(shift.iter()).zip(row(k)) {
983                    *mean = *mean + (x - shift);
984                }
985            }
986            for mean in mean.iter_mut() {
987                *mean = *mean / line.n_t;
988            }
989        }
990        for k in 0..line.n {
991            for (((acc, &shift), &mean), &x) in scale
992                .iter_mut()
993                .zip(shift.iter())
994                .zip(mean.iter())
995                .zip(row(k))
996            {
997                let d = (x - shift) - mean;
998                *acc = *acc + d * d;
999            }
1000        }
1001        for scale in scale.iter_mut() {
1002            *scale = T::ONE / (*scale / line.n_t + line.eps).sqrt();
1003        }
1004        for k in 0..line.n {
1005            let w = if WEIGHT {
1006                affine.weight.unwrap_unchecked().at(k)
1007            } else {
1008                T::ONE
1009            };
1010            let b = if BIAS {
1011                affine.bias.unwrap_unchecked().at(k)
1012            } else {
1013                T::ZERO
1014            };
1015            // Raw writes: the destination may be uninitialized.
1016            let out = dst.offset(d_o + k as isize * line.ds);
1017            for (lane, (((&x, &shift), &mean), &scale)) in row(k)
1018                .iter()
1019                .zip(shift.iter())
1020                .zip(mean.iter())
1021                .zip(scale.iter())
1022                .enumerate()
1023            {
1024                let mut y = ((x - shift) - mean) * scale;
1025                if WEIGHT {
1026                    y = y * w;
1027                }
1028                if BIAS {
1029                    y = y + b;
1030                }
1031                out.add(lane).write(y);
1032            }
1033        }
1034    }
1035}