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}