Skip to main content

rlx_cpu/
kernels.rs

1// RLX — versatile ML compiler + runtime.
2// Copyright (C) 2026 Eugene Hauptmann, Nataliya Kosmyna.
3//
4// This program is free software: you can redistribute it and/or modify
5// it under the terms of the GNU General Public License as published by
6// the Free Software Foundation, version 3.
7//
8// This program is distributed in the hope that it will be useful,
9// but WITHOUT ANY WARRANTY; without even the implied warranty of
10// MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the
11// GNU General Public License for more details.
12//
13// You should have received a copy of the GNU General Public License
14// along with this program. If not, see <https://www.gnu.org/licenses/>.
15
16//! SIMD kernels for fused operations.
17//!
18//! These are the production kernels extracted from burnembed's ndarray_fused.rs.
19//! Each kernel processes data in-place or into a pre-allocated output buffer
20//! (from the arena). No allocation.
21
22use crate::pool;
23
24// ── NEON vectorized exp ─────────────────────────────────────────────────
25
26/// NEON vectorized exp(x) for 4 floats. Range reduction + 6th-order Taylor.
27/// Max relative error: ~2e-7 across [-87, 88].
28#[cfg(target_arch = "aarch64")]
29#[inline(always)]
30#[allow(unsafe_op_in_unsafe_fn)]
31pub unsafe fn neon_exp4(x: std::arch::aarch64::float32x4_t) -> std::arch::aarch64::float32x4_t {
32    use std::arch::aarch64::*;
33    let x = vmaxq_f32(x, vdupq_n_f32(-87.3));
34    let x = vminq_f32(x, vdupq_n_f32(88.7));
35    let inv_ln2 = vdupq_n_f32(std::f32::consts::LOG2_E);
36    let ln2_hi = vdupq_n_f32(0.693_145_75);
37    let ln2_lo = vdupq_n_f32(1.428_606_8e-6);
38    let n = vrndnq_f32(vmulq_f32(x, inv_ln2));
39    let r = vfmsq_f32(vfmsq_f32(x, n, ln2_hi), n, ln2_lo);
40    let c1 = vdupq_n_f32(1.0);
41    let mut p = vdupq_n_f32(0.001_388_888_9);
42    p = vfmaq_f32(vdupq_n_f32(0.008_333_334), p, r);
43    p = vfmaq_f32(vdupq_n_f32(0.041_666_668), p, r);
44    p = vfmaq_f32(vdupq_n_f32(0.166_666_67), p, r);
45    p = vfmaq_f32(vdupq_n_f32(0.5), p, r);
46    p = vfmaq_f32(c1, p, r);
47    p = vfmaq_f32(c1, p, r);
48    let ni = vcvtq_s32_f32(n);
49    vreinterpretq_f32_s32(vaddq_s32(vreinterpretq_s32_f32(p), vshlq_n_s32(ni, 23)))
50}
51
52/// AVX2+FMA vectorised exp(x) for 8 floats. Same range reduction +
53/// 6th-order Taylor polynomial as `neon_exp4`. Max relative error
54/// stays in the ~2e-7 range. Runtime-dispatch via `is_x86_feature_detected`.
55#[cfg(target_arch = "x86_64")]
56#[target_feature(enable = "avx2", enable = "fma")]
57#[allow(unsafe_op_in_unsafe_fn)]
58pub unsafe fn avx2_exp8(x: std::arch::x86_64::__m256) -> std::arch::x86_64::__m256 {
59    use std::arch::x86_64::*;
60    let x = _mm256_max_ps(x, _mm256_set1_ps(-87.3));
61    let x = _mm256_min_ps(x, _mm256_set1_ps(88.7));
62    let inv_ln2 = _mm256_set1_ps(1.442695040888963);
63    let ln2_hi = _mm256_set1_ps(0.693145751953125);
64    let ln2_lo = _mm256_set1_ps(1.428606765330187e-6);
65    // n = round(x / ln2)  (round-to-nearest-even)
66    let n = _mm256_round_ps::<{ _MM_FROUND_TO_NEAREST_INT | _MM_FROUND_NO_EXC }>(_mm256_mul_ps(
67        x, inv_ln2,
68    ));
69    // r = x − n·ln2_hi − n·ln2_lo
70    let r = _mm256_fnmadd_ps(n, ln2_lo, _mm256_fnmadd_ps(n, ln2_hi, x));
71    let c1 = _mm256_set1_ps(1.0);
72    let mut p = _mm256_set1_ps(0.001388888888888889);
73    p = _mm256_fmadd_ps(p, r, _mm256_set1_ps(0.008333333333333333));
74    p = _mm256_fmadd_ps(p, r, _mm256_set1_ps(0.041666666666666664));
75    p = _mm256_fmadd_ps(p, r, _mm256_set1_ps(0.16666666666666666));
76    p = _mm256_fmadd_ps(p, r, _mm256_set1_ps(0.5));
77    p = _mm256_fmadd_ps(p, r, c1);
78    p = _mm256_fmadd_ps(p, r, c1);
79    // 2^n via integer-bias trick on the f32 exponent field.
80    let ni = _mm256_cvtps_epi32(n);
81    let shifted = _mm256_slli_epi32::<23>(ni);
82    _mm256_castsi256_ps(_mm256_add_epi32(_mm256_castps_si256(p), shifted))
83}
84
85// ── Fused bias + GELU ───────────────────────────────────────────────────
86
87/// Fused bias addition + GELU activation on a [m, n] buffer.
88/// Uses Abramowitz & Stegun erf approximation with NEON exp.
89#[cfg(target_arch = "aarch64")]
90pub fn bias_gelu(data: &mut [f32], bias: &[f32], m: usize, n: usize) {
91    use std::arch::aarch64::*;
92    let chunks = n / 4;
93    unsafe {
94        let half = vdupq_n_f32(0.5);
95        let one = vdupq_n_f32(1.0);
96        let inv_sqrt2 = vdupq_n_f32(std::f32::consts::FRAC_1_SQRT_2);
97        let p = vdupq_n_f32(0.3275911);
98        let a1 = vdupq_n_f32(0.254_829_6);
99        let a2 = vdupq_n_f32(-0.284_496_72);
100        let a3 = vdupq_n_f32(1.421_413_8);
101        let a4 = vdupq_n_f32(-1.453_152_1);
102        let a5 = vdupq_n_f32(1.061_405_4);
103        let neg_one = vdupq_n_f32(-1.0);
104        let zero = vdupq_n_f32(0.0);
105
106        for row in 0..m {
107            let base = row * n;
108            for c in 0..chunks {
109                let off = base + c * 4;
110                let ptr = data.as_mut_ptr().add(off);
111                let x = vaddq_f32(vld1q_f32(ptr), vld1q_f32(bias.as_ptr().add(c * 4)));
112                let erf_arg = vmulq_f32(x, inv_sqrt2);
113                let xa = vabsq_f32(erf_arg);
114                let sign = vbslq_f32(vcgeq_f32(erf_arg, zero), one, neg_one);
115                let denom = vfmaq_f32(one, p, xa);
116                let t = vdivq_f32(one, denom);
117                let mut y = a5;
118                y = vfmaq_f32(a4, y, t);
119                y = vfmaq_f32(a3, y, t);
120                y = vfmaq_f32(a2, y, t);
121                y = vfmaq_f32(a1, y, t);
122                y = vmulq_f32(y, t);
123                let exp_val = neon_exp4(vnegq_f32(vmulq_f32(xa, xa)));
124                let erf_val = vmulq_f32(sign, vfmsq_f32(one, y, exp_val));
125                vst1q_f32(ptr, vmulq_f32(x, vmulq_f32(half, vaddq_f32(one, erf_val))));
126            }
127            for i in (chunks * 4)..n {
128                let x = data[base + i] + bias[i];
129                data[base + i] = scalar_gelu(x);
130            }
131        }
132    }
133}
134
135#[cfg(all(
136    target_arch = "x86_64",
137    target_feature = "avx2",
138    target_feature = "fma"
139))]
140pub fn bias_gelu(data: &mut [f32], bias: &[f32], m: usize, n: usize) {
141    use std::arch::x86_64::*;
142    let chunks = n / 8;
143    unsafe {
144        let half = _mm256_set1_ps(0.5);
145        let one = _mm256_set1_ps(1.0);
146        let inv_sqrt2 = _mm256_set1_ps(std::f32::consts::FRAC_1_SQRT_2);
147        let p = _mm256_set1_ps(0.3275911);
148        let a1 = _mm256_set1_ps(0.254829592);
149        let a2 = _mm256_set1_ps(-0.284496736);
150        let a3 = _mm256_set1_ps(1.421413741);
151        let a4 = _mm256_set1_ps(-1.453152027);
152        let a5 = _mm256_set1_ps(1.061405429);
153        let neg_one = _mm256_set1_ps(-1.0);
154        let zero = _mm256_set1_ps(0.0);
155        // Sign bit mask for fabs via AND with 0x7fffffff.
156        let abs_mask = _mm256_castsi256_ps(_mm256_set1_epi32(0x7fff_ffff));
157
158        for row in 0..m {
159            let base = row * n;
160            for c in 0..chunks {
161                let off = base + c * 8;
162                let ptr = data.as_mut_ptr().add(off);
163                let x = _mm256_add_ps(
164                    _mm256_loadu_ps(ptr),
165                    _mm256_loadu_ps(bias.as_ptr().add(c * 8)),
166                );
167                let erf_arg = _mm256_mul_ps(x, inv_sqrt2);
168                let xa = _mm256_and_ps(erf_arg, abs_mask);
169                // sign = (erf_arg >= 0) ? 1 : -1
170                let ge0 = _mm256_cmp_ps::<_CMP_GE_OQ>(erf_arg, zero);
171                let sign = _mm256_blendv_ps(neg_one, one, ge0);
172                let denom = _mm256_fmadd_ps(p, xa, one);
173                let t = _mm256_div_ps(one, denom);
174                let mut y = a5;
175                y = _mm256_fmadd_ps(y, t, a4);
176                y = _mm256_fmadd_ps(y, t, a3);
177                y = _mm256_fmadd_ps(y, t, a2);
178                y = _mm256_fmadd_ps(y, t, a1);
179                y = _mm256_mul_ps(y, t);
180                let exp_val = avx2_exp8(_mm256_sub_ps(zero, _mm256_mul_ps(xa, xa)));
181                // erf = sign * (1 - y*exp(-xa^2))
182                let erf_val = _mm256_mul_ps(sign, _mm256_fnmadd_ps(y, exp_val, one));
183                _mm256_storeu_ps(
184                    ptr,
185                    _mm256_mul_ps(x, _mm256_mul_ps(half, _mm256_add_ps(one, erf_val))),
186                );
187            }
188            for i in (chunks * 8)..n {
189                let x = data[base + i] + bias[i];
190                data[base + i] = scalar_gelu(x);
191            }
192        }
193    }
194}
195
196#[cfg(not(any(
197    target_arch = "aarch64",
198    all(
199        target_arch = "x86_64",
200        target_feature = "avx2",
201        target_feature = "fma"
202    )
203)))]
204pub fn bias_gelu(data: &mut [f32], bias: &[f32], m: usize, n: usize) {
205    for row in 0..m {
206        let base = row * n;
207        for i in 0..n {
208            let x = data[base + i] + bias[i];
209            data[base + i] = scalar_gelu(x);
210        }
211    }
212}
213
214/// Parallel bias + GELU across thread pool.
215pub fn par_bias_gelu(data: &mut [f32], bias: &[f32], m: usize, n: usize) {
216    let cfg = crate::config::RuntimeConfig::global();
217    if m * n < cfg.par_threshold || m < cfg.min_rows_per_thread {
218        bias_gelu(data, bias, m, n);
219        return;
220    }
221    let data_ptr = data.as_mut_ptr() as usize;
222    let bias_ptr = bias.as_ptr() as usize;
223    pool::par_for(m, cfg.min_rows_per_thread, &|off, cnt| unsafe {
224        let d = std::slice::from_raw_parts_mut((data_ptr as *mut f32).add(off * n), cnt * n);
225        let b = std::slice::from_raw_parts(bias_ptr as *const f32, n);
226        bias_gelu(d, b, cnt, n);
227    });
228}
229
230// ── Fused SiLU ──────────────────────────────────────────────────────────
231
232/// SiLU (Swish) in-place: x / (1 + exp(-x))
233#[cfg(target_arch = "aarch64")]
234pub fn silu_inplace(data: &mut [f32]) {
235    use std::arch::aarch64::*;
236    let chunks = data.len() / 4;
237    unsafe {
238        let one = vdupq_n_f32(1.0);
239        for c in 0..chunks {
240            let ptr = data.as_mut_ptr().add(c * 4);
241            let x = vld1q_f32(ptr);
242            let exp_neg = neon_exp4(vnegq_f32(x));
243            let sigmoid = vdivq_f32(one, vaddq_f32(one, exp_neg));
244            vst1q_f32(ptr, vmulq_f32(x, sigmoid));
245        }
246    }
247    for i in (chunks * 4)..data.len() {
248        let x = data[i];
249        data[i] = x / (1.0 + (-x).exp());
250    }
251}
252
253/// SiLU via AVX2+FMA. Caller must have checked `is_x86_feature_detected`.
254#[cfg(target_arch = "x86_64")]
255#[target_feature(enable = "avx2", enable = "fma")]
256#[allow(unsafe_op_in_unsafe_fn)]
257unsafe fn silu_inplace_avx2(data: &mut [f32]) {
258    use std::arch::x86_64::*;
259    let chunks = data.len() / 8;
260    let one = _mm256_set1_ps(1.0);
261    let zero = _mm256_set1_ps(0.0);
262    for c in 0..chunks {
263        let off = c * 8;
264        let ptr = data.as_mut_ptr().add(off);
265        let x = _mm256_loadu_ps(ptr);
266        // silu(x) = x / (1 + exp(-x))
267        let neg_x = _mm256_sub_ps(zero, x);
268        let denom = _mm256_add_ps(one, avx2_exp8(neg_x));
269        _mm256_storeu_ps(ptr, _mm256_div_ps(x, denom));
270    }
271    for i in (chunks * 8)..data.len() {
272        let x = data[i];
273        data[i] = x / (1.0 + (-x).exp());
274    }
275}
276
277#[cfg(target_arch = "x86_64")]
278pub fn silu_inplace(data: &mut [f32]) {
279    if std::arch::is_x86_feature_detected!("avx2") && std::arch::is_x86_feature_detected!("fma") {
280        unsafe { silu_inplace_avx2(data) };
281        return;
282    }
283    for v in data.iter_mut() {
284        let x = *v;
285        *v = x / (1.0 + (-x).exp());
286    }
287}
288
289#[cfg(not(any(target_arch = "aarch64", target_arch = "x86_64")))]
290pub fn silu_inplace(data: &mut [f32]) {
291    for v in data.iter_mut() {
292        let x = *v;
293        *v = x / (1.0 + (-x).exp());
294    }
295}
296
297// ── LayerNorm (2-pass) ──────────────────────────────────────────────────
298
299/// Single-row LayerNorm: out = (x - mean) * inv_std * gamma + beta.
300/// 2-pass: compute mean+variance (E\[x²\]-E\[x\]²), then normalize.
301#[cfg(target_arch = "aarch64")]
302pub fn layer_norm_row(
303    input: &[f32],
304    gamma: &[f32],
305    beta: &[f32],
306    output: &mut [f32],
307    h: usize,
308    eps: f32,
309) {
310    use std::arch::aarch64::*;
311    let inv_hf = 1.0 / h as f32;
312    let chunks = h / 4;
313    unsafe {
314        let mut vsum = vdupq_n_f32(0.0);
315        let mut vsumsq = vdupq_n_f32(0.0);
316        for c in 0..chunks {
317            let x = vld1q_f32(input.as_ptr().add(c * 4));
318            vsum = vaddq_f32(vsum, x);
319            vsumsq = vfmaq_f32(vsumsq, x, x);
320        }
321        let mut sum = vaddvq_f32(vsum);
322        let mut sumsq = vaddvq_f32(vsumsq);
323        for i in (chunks * 4)..h {
324            sum += input[i];
325            sumsq += input[i] * input[i];
326        }
327        let mean = sum * inv_hf;
328        let var = (sumsq * inv_hf - mean * mean).max(0.0);
329        let inv = 1.0 / (var + eps).sqrt();
330        let vmean = vdupq_n_f32(mean);
331        let vinv = vdupq_n_f32(inv);
332        for c in 0..chunks {
333            let off = c * 4;
334            let x = vld1q_f32(input.as_ptr().add(off));
335            let norm = vmulq_f32(vsubq_f32(x, vmean), vinv);
336            vst1q_f32(
337                output.as_mut_ptr().add(off),
338                vfmaq_f32(
339                    vld1q_f32(beta.as_ptr().add(off)),
340                    norm,
341                    vld1q_f32(gamma.as_ptr().add(off)),
342                ),
343            );
344        }
345        for i in (chunks * 4)..h {
346            output[i] = (input[i] - mean) * inv * gamma[i] + beta[i];
347        }
348    }
349}
350
351#[cfg(all(
352    target_arch = "x86_64",
353    target_feature = "avx2",
354    target_feature = "fma"
355))]
356pub fn layer_norm_row(
357    input: &[f32],
358    gamma: &[f32],
359    beta: &[f32],
360    output: &mut [f32],
361    h: usize,
362    eps: f32,
363) {
364    use std::arch::x86_64::*;
365    let inv_hf = 1.0 / h as f32;
366    let chunks = h / 8;
367    unsafe {
368        let mut vsum = _mm256_setzero_ps();
369        let mut vsumsq = _mm256_setzero_ps();
370        for c in 0..chunks {
371            let x = _mm256_loadu_ps(input.as_ptr().add(c * 8));
372            vsum = _mm256_add_ps(vsum, x);
373            vsumsq = _mm256_fmadd_ps(x, x, vsumsq);
374        }
375        // Horizontal reduce: 8 lanes → 1.
376        let hsum = {
377            let lo = _mm256_castps256_ps128(vsum);
378            let hi = _mm256_extractf128_ps::<1>(vsum);
379            let s4 = _mm_add_ps(lo, hi);
380            let s2 = _mm_add_ps(s4, _mm_movehl_ps(s4, s4));
381            let s1 = _mm_add_ss(s2, _mm_shuffle_ps::<0x55>(s2, s2));
382            _mm_cvtss_f32(s1)
383        };
384        let hsumsq = {
385            let lo = _mm256_castps256_ps128(vsumsq);
386            let hi = _mm256_extractf128_ps::<1>(vsumsq);
387            let s4 = _mm_add_ps(lo, hi);
388            let s2 = _mm_add_ps(s4, _mm_movehl_ps(s4, s4));
389            let s1 = _mm_add_ss(s2, _mm_shuffle_ps::<0x55>(s2, s2));
390            _mm_cvtss_f32(s1)
391        };
392        let mut sum = hsum;
393        let mut sumsq = hsumsq;
394        for i in (chunks * 8)..h {
395            sum += input[i];
396            sumsq += input[i] * input[i];
397        }
398        let mean = sum * inv_hf;
399        let var = (sumsq * inv_hf - mean * mean).max(0.0);
400        let inv = 1.0 / (var + eps).sqrt();
401        let vmean = _mm256_set1_ps(mean);
402        let vinv = _mm256_set1_ps(inv);
403        for c in 0..chunks {
404            let off = c * 8;
405            let x = _mm256_loadu_ps(input.as_ptr().add(off));
406            let norm = _mm256_mul_ps(_mm256_sub_ps(x, vmean), vinv);
407            let g = _mm256_loadu_ps(gamma.as_ptr().add(off));
408            let b = _mm256_loadu_ps(beta.as_ptr().add(off));
409            _mm256_storeu_ps(output.as_mut_ptr().add(off), _mm256_fmadd_ps(norm, g, b));
410        }
411        for i in (chunks * 8)..h {
412            output[i] = (input[i] - mean) * inv * gamma[i] + beta[i];
413        }
414    }
415}
416
417#[cfg(not(any(
418    target_arch = "aarch64",
419    all(
420        target_arch = "x86_64",
421        target_feature = "avx2",
422        target_feature = "fma"
423    )
424)))]
425pub fn layer_norm_row(
426    input: &[f32],
427    gamma: &[f32],
428    beta: &[f32],
429    output: &mut [f32],
430    h: usize,
431    eps: f32,
432) {
433    let inv_hf = 1.0 / h as f32;
434    let mut sum = 0f32;
435    let mut sumsq = 0f32;
436    for i in 0..h {
437        sum += input[i];
438        sumsq += input[i] * input[i];
439    }
440    let mean = sum * inv_hf;
441    let var = (sumsq * inv_hf - mean * mean).max(0.0);
442    let inv = 1.0 / (var + eps).sqrt();
443    for i in 0..h {
444        output[i] = (input[i] - mean) * inv * gamma[i] + beta[i];
445    }
446}
447
448/// Inference BatchNorm with frozen running statistics (PyTorch `BatchNorm*d` eval).
449///
450/// `x` is row-major with feature dimension `channels` on the last axis
451/// (`[B, C]`, `[B, P, C]`, …). `gamma`, `beta`, `mean`, `var` are length `C`.
452pub fn batch_norm_inference(
453    x: &[f32],
454    gamma: &[f32],
455    beta: &[f32],
456    mean: &[f32],
457    var: &[f32],
458    out: &mut [f32],
459    channels: usize,
460    eps: f32,
461) {
462    let n = x.len() / channels.max(1);
463    for i in 0..n {
464        for c in 0..channels {
465            let idx = i * channels + c;
466            let inv = 1.0 / (var[c] + eps).sqrt();
467            let xhat = (x[idx] - mean[c]) * inv;
468            out[idx] = gamma[c] * xhat + beta[c];
469        }
470    }
471}
472
473/// `d_x` for [`batch_norm_inference`] (mean/var treated as constants).
474pub fn batch_norm_inference_backward_input(
475    x: &[f32],
476    gamma: &[f32],
477    _mean: &[f32],
478    var: &[f32],
479    dy: &[f32],
480    dx: &mut [f32],
481    channels: usize,
482    eps: f32,
483) {
484    let n = x.len() / channels.max(1);
485    for i in 0..n {
486        for c in 0..channels {
487            let idx = i * channels + c;
488            let inv = 1.0 / (var[c] + eps).sqrt();
489            dx[idx] = dy[idx] * gamma[c] * inv;
490        }
491    }
492}
493
494/// `d_gamma` for [`batch_norm_inference`].
495pub fn batch_norm_inference_backward_gamma(
496    x: &[f32],
497    mean: &[f32],
498    var: &[f32],
499    dy: &[f32],
500    dgamma: &mut [f32],
501    channels: usize,
502    eps: f32,
503) {
504    dgamma.fill(0.0);
505    let n = x.len() / channels.max(1);
506    for i in 0..n {
507        for c in 0..channels {
508            let idx = i * channels + c;
509            let inv = 1.0 / (var[c] + eps).sqrt();
510            let xhat = (x[idx] - mean[c]) * inv;
511            dgamma[c] += dy[idx] * xhat;
512        }
513    }
514}
515
516/// `d_beta` for [`batch_norm_inference`].
517pub fn batch_norm_inference_backward_beta(dy: &[f32], dbeta: &mut [f32], channels: usize) {
518    dbeta.fill(0.0);
519    let n = dy.len() / channels.max(1);
520    for i in 0..n {
521        for c in 0..channels {
522            dbeta[c] += dy[i * channels + c];
523        }
524    }
525}
526
527/// Fused residual + bias + LayerNorm on [n, h] buffers.
528/// Computes: output\[row\] = LN(a\[row\] + b\[row\] + bias, gamma, beta)
529pub fn residual_bias_layer_norm(
530    a: &[f32],
531    b: &[f32],
532    bias: &[f32],
533    gamma: &[f32],
534    beta: &[f32],
535    output: &mut [f32],
536    n: usize,
537    h: usize,
538    eps: f32,
539) {
540    // Temporary per-row buffer for a+b+bias (stack allocated for small h)
541    let mut tmp = vec![0f32; h];
542    for row in 0..n {
543        let base = row * h;
544        for i in 0..h {
545            tmp[i] = a[base + i] + b[base + i] + bias[i];
546        }
547        layer_norm_row(&tmp, gamma, beta, &mut output[base..base + h], h, eps);
548    }
549}
550
551/// Fused residual + bias + RMSNorm on [n, h] buffers.
552/// Computes: output[row] = RmsNorm(a[row] + b[row] + bias, gamma, beta)
553pub fn residual_bias_rms_norm(
554    a: &[f32],
555    b: &[f32],
556    bias: &[f32],
557    gamma: &[f32],
558    beta: &[f32],
559    output: &mut [f32],
560    n: usize,
561    h: usize,
562    eps: f32,
563) {
564    let inv_h = 1.0 / h as f32;
565    for row in 0..n {
566        let base = row * h;
567        let mut sumsq = 0f32;
568        for i in 0..h {
569            let v = a[base + i] + b[base + i] + bias[i];
570            sumsq += v * v;
571        }
572        let inv_rms = (sumsq * inv_h + eps).sqrt().recip();
573        for i in 0..h {
574            let v = a[base + i] + b[base + i] + bias[i];
575            output[base + i] = v * inv_rms * gamma[i] + beta[i];
576        }
577    }
578}
579
580/// Parallel residual + bias + LayerNorm.
581pub fn par_residual_bias_ln(
582    a: &[f32],
583    b: &[f32],
584    bias: &[f32],
585    gamma: &[f32],
586    beta: &[f32],
587    output: &mut [f32],
588    n: usize,
589    h: usize,
590    eps: f32,
591) {
592    let cfg = crate::config::RuntimeConfig::global();
593    if n * h < cfg.par_threshold || n < cfg.min_rows_per_thread {
594        residual_bias_layer_norm(a, b, bias, gamma, beta, output, n, h, eps);
595        return;
596    }
597    let a_ptr = a.as_ptr() as usize;
598    let b_ptr = b.as_ptr() as usize;
599    let o_ptr = output.as_mut_ptr() as usize;
600    let bias_ptr = bias.as_ptr() as usize;
601    let gamma_ptr = gamma.as_ptr() as usize;
602    let beta_ptr = beta.as_ptr() as usize;
603    pool::par_for(n, cfg.min_rows_per_thread, &|off, cnt| unsafe {
604        let a_s = std::slice::from_raw_parts((a_ptr as *const f32).add(off * h), cnt * h);
605        let b_s = std::slice::from_raw_parts((b_ptr as *const f32).add(off * h), cnt * h);
606        let o_s = std::slice::from_raw_parts_mut((o_ptr as *mut f32).add(off * h), cnt * h);
607        let bi = std::slice::from_raw_parts(bias_ptr as *const f32, h);
608        let g = std::slice::from_raw_parts(gamma_ptr as *const f32, h);
609        let be = std::slice::from_raw_parts(beta_ptr as *const f32, h);
610        residual_bias_layer_norm(a_s, b_s, bi, g, be, o_s, cnt, h, eps);
611    });
612}
613
614// ── Softmax (NEON / runtime AVX2 / parallel rows) ───────────────────────
615
616/// Row-parallel wrapper: each row is independent so Rayon can split the outer
617/// loop when there is enough work to amortize hand-off.
618#[inline]
619fn par_softmax_rows<F: Fn(&mut [f32], usize, usize) + Sync>(
620    data: &mut [f32],
621    rows: usize,
622    cols: usize,
623    kernel: &F,
624) {
625    if rows >= 4 && pool::should_parallelize(rows * cols) {
626        let base = data.as_mut_ptr() as usize;
627        pool::par_for(rows, 1, &|off, cnt| {
628            for r in off..off + cnt {
629                let row = unsafe {
630                    std::slice::from_raw_parts_mut((base as *mut f32).add(r * cols), cols)
631                };
632                kernel(row, 1, cols);
633            }
634        });
635    } else {
636        kernel(data, rows, cols);
637    }
638}
639
640/// NEON-vectorized softmax: 3-pass (max, exp+sum, normalize).
641#[cfg(target_arch = "aarch64")]
642fn softmax_rows_neon(data: &mut [f32], rows: usize, cols: usize) {
643    use std::arch::aarch64::*;
644    let chunks = cols / 4;
645    unsafe {
646        for row in 0..rows {
647            let base = row * cols;
648            let ptr = data.as_mut_ptr().add(base);
649
650            // Pass 1: find row max
651            let mut vmax = vdupq_n_f32(f32::NEG_INFINITY);
652            for c in 0..chunks {
653                vmax = vmaxq_f32(vmax, vld1q_f32(ptr.add(c * 4)));
654            }
655            let mut max_val = vmaxvq_f32(vmax);
656            for i in (chunks * 4)..cols {
657                max_val = max_val.max(*ptr.add(i));
658            }
659
660            // Pass 2: exp(x - max) and accumulate sum
661            let vmx = vdupq_n_f32(max_val);
662            let mut vsum = vdupq_n_f32(0.0);
663            for c in 0..chunks {
664                let off = c * 4;
665                let e = neon_exp4(vsubq_f32(vld1q_f32(ptr.add(off)), vmx));
666                vst1q_f32(ptr.add(off), e);
667                vsum = vaddq_f32(vsum, e);
668            }
669            let mut sum = vaddvq_f32(vsum);
670            for i in (chunks * 4)..cols {
671                let e = (*ptr.add(i) - max_val).exp();
672                *ptr.add(i) = e;
673                sum += e;
674            }
675
676            // Pass 3: normalize
677            let vinv = vdupq_n_f32(1.0 / sum);
678            for c in 0..chunks {
679                let off = c * 4;
680                vst1q_f32(ptr.add(off), vmulq_f32(vld1q_f32(ptr.add(off)), vinv));
681            }
682            let inv = 1.0 / sum;
683            for i in (chunks * 4)..cols {
684                *ptr.add(i) *= inv;
685            }
686        }
687    }
688}
689
690#[cfg(target_arch = "aarch64")]
691pub fn neon_softmax(data: &mut [f32], rows: usize, cols: usize) {
692    par_softmax_rows(data, rows, cols, &softmax_rows_neon);
693}
694
695#[cfg(target_arch = "x86_64")]
696#[target_feature(enable = "avx2", enable = "fma")]
697#[allow(unsafe_op_in_unsafe_fn)]
698unsafe fn softmax_rows_avx2(data: &mut [f32], rows: usize, cols: usize) {
699    use std::arch::x86_64::*;
700    let chunks = cols / 8;
701    for r in 0..rows {
702        let row = data.as_mut_ptr().add(r * cols);
703        // 1) Vector max for stability.
704        let mut vmax = _mm256_set1_ps(f32::NEG_INFINITY);
705        for c in 0..chunks {
706            vmax = _mm256_max_ps(vmax, _mm256_loadu_ps(row.add(c * 8)));
707        }
708        let mut max_v = {
709            let lo = _mm256_castps256_ps128(vmax);
710            let hi = _mm256_extractf128_ps::<1>(vmax);
711            let s4 = _mm_max_ps(lo, hi);
712            let s2 = _mm_max_ps(s4, _mm_movehl_ps(s4, s4));
713            let s1 = _mm_max_ss(s2, _mm_shuffle_ps::<0x55>(s2, s2));
714            _mm_cvtss_f32(s1)
715        };
716        for i in (chunks * 8)..cols {
717            let v = *row.add(i);
718            if v > max_v {
719                max_v = v;
720            }
721        }
722        // 2) exp(x − max) and sum.
723        let vmax = _mm256_set1_ps(max_v);
724        let mut vsum = _mm256_setzero_ps();
725        for c in 0..chunks {
726            let off = c * 8;
727            let e = avx2_exp8(_mm256_sub_ps(_mm256_loadu_ps(row.add(off)), vmax));
728            _mm256_storeu_ps(row.add(off), e);
729            vsum = _mm256_add_ps(vsum, e);
730        }
731        let mut sum_v = {
732            let lo = _mm256_castps256_ps128(vsum);
733            let hi = _mm256_extractf128_ps::<1>(vsum);
734            let s4 = _mm_add_ps(lo, hi);
735            let s2 = _mm_add_ps(s4, _mm_movehl_ps(s4, s4));
736            let s1 = _mm_add_ss(s2, _mm_shuffle_ps::<0x55>(s2, s2));
737            _mm_cvtss_f32(s1)
738        };
739        for i in (chunks * 8)..cols {
740            let v = (*row.add(i) - max_v).exp();
741            *row.add(i) = v;
742            sum_v += v;
743        }
744        // 3) Normalize.
745        let vinv = _mm256_set1_ps(1.0 / sum_v);
746        for c in 0..chunks {
747            let off = c * 8;
748            _mm256_storeu_ps(
749                row.add(off),
750                _mm256_mul_ps(_mm256_loadu_ps(row.add(off)), vinv),
751            );
752        }
753        let inv_sum = 1.0 / sum_v;
754        for i in (chunks * 8)..cols {
755            *row.add(i) *= inv_sum;
756        }
757    }
758}
759
760#[cfg(target_arch = "x86_64")]
761pub fn neon_softmax(data: &mut [f32], rows: usize, cols: usize) {
762    let avx2 =
763        std::arch::is_x86_feature_detected!("avx2") && std::arch::is_x86_feature_detected!("fma");
764    if avx2 {
765        par_softmax_rows(data, rows, cols, &|d, r, c| unsafe {
766            softmax_rows_avx2(d, r, c);
767        });
768    } else {
769        par_softmax_rows(data, rows, cols, &crate::naive::softmax);
770    }
771}
772
773#[cfg(not(any(target_arch = "aarch64", target_arch = "x86_64")))]
774pub fn neon_softmax(data: &mut [f32], rows: usize, cols: usize) {
775    par_softmax_rows(data, rows, cols, &crate::naive::softmax);
776}
777
778// ── GELU in-place (no bias) ────────────────────────────────────────────
779
780/// NEON GELU activation in-place (without bias addition).
781#[cfg(target_arch = "aarch64")]
782pub fn gelu_inplace(data: &mut [f32]) {
783    use std::arch::aarch64::*;
784    let len = data.len();
785    let chunks = len / 4;
786    unsafe {
787        let half = vdupq_n_f32(0.5);
788        let one = vdupq_n_f32(1.0);
789        let inv_sqrt2 = vdupq_n_f32(std::f32::consts::FRAC_1_SQRT_2);
790        let p = vdupq_n_f32(0.3275911);
791        let a1 = vdupq_n_f32(0.254_829_6);
792        let a2 = vdupq_n_f32(-0.284_496_72);
793        let a3 = vdupq_n_f32(1.421_413_8);
794        let a4 = vdupq_n_f32(-1.453_152_1);
795        let a5 = vdupq_n_f32(1.061_405_4);
796        let neg_one = vdupq_n_f32(-1.0);
797        let zero = vdupq_n_f32(0.0);
798
799        for c in 0..chunks {
800            let ptr = data.as_mut_ptr().add(c * 4);
801            let x = vld1q_f32(ptr);
802            let erf_arg = vmulq_f32(x, inv_sqrt2);
803            let xa = vabsq_f32(erf_arg);
804            let sign = vbslq_f32(vcgeq_f32(erf_arg, zero), one, neg_one);
805            let denom = vfmaq_f32(one, p, xa);
806            let t = vdivq_f32(one, denom);
807            let mut y = a5;
808            y = vfmaq_f32(a4, y, t);
809            y = vfmaq_f32(a3, y, t);
810            y = vfmaq_f32(a2, y, t);
811            y = vfmaq_f32(a1, y, t);
812            y = vmulq_f32(y, t);
813            let exp_val = neon_exp4(vnegq_f32(vmulq_f32(xa, xa)));
814            let erf_val = vmulq_f32(sign, vfmsq_f32(one, y, exp_val));
815            vst1q_f32(ptr, vmulq_f32(x, vmulq_f32(half, vaddq_f32(one, erf_val))));
816        }
817        for i in (chunks * 4)..len {
818            data[i] = scalar_gelu(data[i]);
819        }
820    }
821}
822
823/// Erf-GELU via AVX2+FMA. Caller must have checked feature bits.
824#[cfg(target_arch = "x86_64")]
825#[target_feature(enable = "avx2", enable = "fma")]
826#[allow(unsafe_op_in_unsafe_fn)]
827unsafe fn gelu_inplace_avx2(data: &mut [f32]) {
828    use std::arch::x86_64::*;
829    let chunks = data.len() / 8;
830    let half = _mm256_set1_ps(0.5);
831    let one = _mm256_set1_ps(1.0);
832    let inv_sqrt2 = _mm256_set1_ps(std::f32::consts::FRAC_1_SQRT_2);
833    let p = _mm256_set1_ps(0.3275911);
834    let a1 = _mm256_set1_ps(0.254829592);
835    let a2 = _mm256_set1_ps(-0.284496736);
836    let a3 = _mm256_set1_ps(1.421413741);
837    let a4 = _mm256_set1_ps(-1.453152027);
838    let a5 = _mm256_set1_ps(1.061405429);
839    let neg_one = _mm256_set1_ps(-1.0);
840    let zero = _mm256_set1_ps(0.0);
841    let abs_mask = _mm256_castsi256_ps(_mm256_set1_epi32(0x7fff_ffff));
842    for c in 0..chunks {
843        let off = c * 8;
844        let ptr = data.as_mut_ptr().add(off);
845        let x = _mm256_loadu_ps(ptr);
846        let erf_arg = _mm256_mul_ps(x, inv_sqrt2);
847        let xa = _mm256_and_ps(erf_arg, abs_mask);
848        let ge0 = _mm256_cmp_ps::<_CMP_GE_OQ>(erf_arg, zero);
849        let sign = _mm256_blendv_ps(neg_one, one, ge0);
850        let denom = _mm256_fmadd_ps(p, xa, one);
851        let t = _mm256_div_ps(one, denom);
852        let mut y = a5;
853        y = _mm256_fmadd_ps(y, t, a4);
854        y = _mm256_fmadd_ps(y, t, a3);
855        y = _mm256_fmadd_ps(y, t, a2);
856        y = _mm256_fmadd_ps(y, t, a1);
857        y = _mm256_mul_ps(y, t);
858        let exp_val = avx2_exp8(_mm256_sub_ps(zero, _mm256_mul_ps(xa, xa)));
859        let erf_val = _mm256_mul_ps(sign, _mm256_fnmadd_ps(y, exp_val, one));
860        _mm256_storeu_ps(
861            ptr,
862            _mm256_mul_ps(x, _mm256_mul_ps(half, _mm256_add_ps(one, erf_val))),
863        );
864    }
865    for i in (chunks * 8)..data.len() {
866        data[i] = scalar_gelu(data[i]);
867    }
868}
869
870#[cfg(target_arch = "x86_64")]
871pub fn gelu_inplace(data: &mut [f32]) {
872    if std::arch::is_x86_feature_detected!("avx2") && std::arch::is_x86_feature_detected!("fma") {
873        unsafe { gelu_inplace_avx2(data) };
874        return;
875    }
876    for v in data.iter_mut() {
877        *v = scalar_gelu(*v);
878    }
879}
880
881#[cfg(not(any(target_arch = "aarch64", target_arch = "x86_64")))]
882pub fn gelu_inplace(data: &mut [f32]) {
883    for v in data.iter_mut() {
884        *v = scalar_gelu(*v);
885    }
886}
887
888/// Parallel GELU in-place (splits work across thread pool).
889///
890/// Activation kernels are O(n) with very low per-element cost
891/// (~10 NEON cycles on aarch64). Pool dispatch overhead — even
892/// with the parked design — is in the multi-µs range under
893/// container scheduling, which dwarfs the actual compute for any
894/// reasonable activation size. Threshold here is 1 Mi elements:
895/// only crossed by very large activation tensors (e.g. an
896/// H=4096, FFN=14336, S=1024 LLM up-projection at ~14M
897/// elements). Single-thread NEON is the clear win below that.
898const ACTIVATION_PAR_MIN: usize = 1 << 20;
899
900/// Tanh-approximation GELU (matches PyTorch/candle `Tensor::gelu`):
901///   y = 0.5 x (1 + tanh(√(2/π) · (x + 0.044715 x³)))
902///
903/// Scalar-only for now; the erf-based `gelu_inplace` above is SIMD.
904/// Routed from `Activation::GeluApprox` so models that need
905/// numerical parity with PyTorch's default GELU (e.g. DINOv2,
906/// many ViTs) get the right formula. Use `Activation::Gelu` for the
907/// erf form (also PyTorch-default in some newer builds).
908#[inline]
909pub fn scalar_gelu_approx(x: f32) -> f32 {
910    const C: f32 = 0.797_884_6; // √(2/π)
911    const A: f32 = 0.044_715;
912    0.5 * x * (1.0 + (C * (x + A * x * x * x)).tanh())
913}
914
915pub fn gelu_approx_inplace(data: &mut [f32]) {
916    for v in data.iter_mut() {
917        *v = scalar_gelu_approx(*v);
918    }
919}
920
921pub fn par_gelu_approx_inplace(data: &mut [f32]) {
922    let len = data.len();
923    if len < ACTIVATION_PAR_MIN {
924        gelu_approx_inplace(data);
925        return;
926    }
927    let cfg = crate::config::RuntimeConfig::global();
928    let chunk = 512;
929    let rows = len / chunk;
930    if rows < 2 {
931        gelu_approx_inplace(data);
932        return;
933    }
934    let data_ptr = data.as_mut_ptr() as usize;
935    pool::par_for(rows, cfg.min_rows_per_thread, &|off, cnt| unsafe {
936        let start = off * chunk;
937        let end = if off + cnt >= rows {
938            len
939        } else {
940            (off + cnt) * chunk
941        };
942        let s = std::slice::from_raw_parts_mut((data_ptr as *mut f32).add(start), end - start);
943        gelu_approx_inplace(s);
944    });
945    let done = rows * chunk;
946    if done < len {
947        gelu_approx_inplace(&mut data[done..]);
948    }
949}
950
951pub fn gelu_approx_out(src: &[f32], dst: &mut [f32]) {
952    debug_assert_eq!(src.len(), dst.len());
953    for (s, d) in src.iter().zip(dst.iter_mut()) {
954        *d = scalar_gelu_approx(*s);
955    }
956}
957
958pub fn par_gelu_approx_out(src: &[f32], dst: &mut [f32]) {
959    debug_assert_eq!(src.len(), dst.len());
960    let len = src.len();
961    if len < ACTIVATION_PAR_MIN {
962        gelu_approx_out(src, dst);
963        return;
964    }
965    let cfg = crate::config::RuntimeConfig::global();
966    let chunk = 512;
967    let rows = len / chunk;
968    if rows < 2 {
969        gelu_approx_out(src, dst);
970        return;
971    }
972    let src_ptr = src.as_ptr() as usize;
973    let dst_ptr = dst.as_mut_ptr() as usize;
974    pool::par_for(rows, cfg.min_rows_per_thread, &|off, cnt| unsafe {
975        let start = off * chunk;
976        let end = if off + cnt >= rows {
977            len
978        } else {
979            (off + cnt) * chunk
980        };
981        let n = end - start;
982        let s = std::slice::from_raw_parts((src_ptr as *const f32).add(start), n);
983        let d = std::slice::from_raw_parts_mut((dst_ptr as *mut f32).add(start), n);
984        gelu_approx_out(s, d);
985    });
986    let done = rows * chunk;
987    if done < len {
988        gelu_approx_out(&src[done..], &mut dst[done..]);
989    }
990}
991
992pub fn par_gelu_inplace(data: &mut [f32]) {
993    let len = data.len();
994    if len < ACTIVATION_PAR_MIN {
995        gelu_inplace(data);
996        return;
997    }
998    let cfg = crate::config::RuntimeConfig::global();
999    let chunk = 512;
1000    let rows = len / chunk;
1001    if rows < 2 {
1002        gelu_inplace(data);
1003        return;
1004    }
1005    let data_ptr = data.as_mut_ptr() as usize;
1006    pool::par_for(rows, cfg.min_rows_per_thread, &|off, cnt| unsafe {
1007        let start = off * chunk;
1008        let end = if off + cnt >= rows {
1009            len
1010        } else {
1011            (off + cnt) * chunk
1012        };
1013        let s = std::slice::from_raw_parts_mut((data_ptr as *mut f32).add(start), end - start);
1014        gelu_inplace(s);
1015    });
1016    let done = rows * chunk;
1017    if done < len {
1018        gelu_inplace(&mut data[done..]);
1019    }
1020}
1021
1022/// Parallel SiLU in-place. Same threshold reasoning as `par_gelu_inplace`.
1023pub fn par_silu_inplace(data: &mut [f32]) {
1024    let len = data.len();
1025    if len < ACTIVATION_PAR_MIN {
1026        silu_inplace(data);
1027        return;
1028    }
1029    let cfg = crate::config::RuntimeConfig::global();
1030    let chunk = 512;
1031    let rows = len / chunk;
1032    if rows < 2 {
1033        silu_inplace(data);
1034        return;
1035    }
1036    let data_ptr = data.as_mut_ptr() as usize;
1037    pool::par_for(rows, cfg.min_rows_per_thread, &|off, cnt| unsafe {
1038        let start = off * chunk;
1039        let end = if off + cnt >= rows {
1040            len
1041        } else {
1042            (off + cnt) * chunk
1043        };
1044        let s = std::slice::from_raw_parts_mut((data_ptr as *mut f32).add(start), end - start);
1045        silu_inplace(s);
1046    });
1047    let done = rows * chunk;
1048    if done < len {
1049        silu_inplace(&mut data[done..]);
1050    }
1051}
1052
1053// ── Small-m NEON matmul ─────────────────────────────────────────────────
1054
1055/// NEON matmul for tiny m (1-8 rows). Avoids BLAS call overhead.
1056/// C = A @ B where A=\[m,k\], B=\[k,n\], C=\[m,n\], all row-major.
1057/// For m≤8 with small k×n (under ~16K elements), this beats cblas_sgemm
1058/// by avoiding AMX setup cost and function call overhead.
1059#[cfg(target_arch = "aarch64")]
1060pub fn neon_sgemm_small(a: &[f32], b: &[f32], c: &mut [f32], m: usize, k: usize, n: usize) {
1061    use std::arch::aarch64::*;
1062    let n4 = n / 4;
1063    unsafe {
1064        for j4 in 0..n4 {
1065            let j = j4 * 4;
1066            // m accumulators (one per output row, 4-wide)
1067            let mut acc = [vdupq_n_f32(0.0); 8];
1068            for kk in 0..k {
1069                let bv = vld1q_f32(b.as_ptr().add(kk * n + j));
1070                for i in 0..m {
1071                    let av = vdupq_n_f32(*a.as_ptr().add(i * k + kk));
1072                    acc[i] = vfmaq_f32(acc[i], av, bv);
1073                }
1074            }
1075            for i in 0..m {
1076                vst1q_f32(c.as_mut_ptr().add(i * n + j), acc[i]);
1077            }
1078        }
1079        // Remainder columns
1080        for j in (n4 * 4)..n {
1081            for i in 0..m {
1082                let mut sum = 0f32;
1083                for kk in 0..k {
1084                    sum += a[i * k + kk] * b[kk * n + j];
1085                }
1086                c[i * n + j] = sum;
1087            }
1088        }
1089    }
1090}
1091
1092#[cfg(not(target_arch = "aarch64"))]
1093pub fn neon_sgemm_small(a: &[f32], b: &[f32], c: &mut [f32], m: usize, k: usize, n: usize) {
1094    crate::naive::matmul(a, b, c, m, k, n);
1095}
1096
1097/// NEON sgemm_bias for tiny m: C = A @ B + bias.
1098#[cfg(target_arch = "aarch64")]
1099pub fn neon_sgemm_bias_small(
1100    a: &[f32],
1101    b: &[f32],
1102    bias: &[f32],
1103    c: &mut [f32],
1104    m: usize,
1105    k: usize,
1106    n: usize,
1107) {
1108    neon_sgemm_small(a, b, c, m, k, n);
1109    crate::blas::bias_add(c, bias, m, n);
1110}
1111
1112#[cfg(not(target_arch = "aarch64"))]
1113pub fn neon_sgemm_bias_small(
1114    a: &[f32],
1115    b: &[f32],
1116    bias: &[f32],
1117    c: &mut [f32],
1118    m: usize,
1119    k: usize,
1120    n: usize,
1121) {
1122    crate::naive::matmul(a, b, c, m, k, n);
1123    crate::naive::bias_add(c, bias, m, n);
1124}
1125
1126// ── Scalar fallbacks ────────────────────────────────────────────────────
1127
1128fn scalar_gelu(x: f32) -> f32 {
1129    x * 0.5 * (1.0 + scalar_erf(x * std::f32::consts::FRAC_1_SQRT_2))
1130}
1131
1132fn scalar_erf(x: f32) -> f32 {
1133    let sign = if x >= 0.0 { 1.0f32 } else { -1.0 };
1134    let xa = x.abs();
1135    let t = 1.0 / (1.0 + 0.3275911 * xa);
1136    let y = t
1137        * (0.254_829_6
1138            + t * (-0.284_496_72 + t * (1.421_413_8 + t * (-1.453_152_1 + t * 1.061_405_4))));
1139    sign * (1.0 - y * (-xa * xa).exp())
1140}
1141
1142/// NCHW LayerNorm2d (candle / SAM semantics): normalize across channels at
1143/// each spatial position. `gamma`/`beta` are per-channel `[C]`.
1144pub fn layer_norm2d_nchw(
1145    input: &[f32],
1146    gamma: &[f32],
1147    beta: &[f32],
1148    output: &mut [f32],
1149    batch: usize,
1150    channels: usize,
1151    h: usize,
1152    w: usize,
1153    eps: f32,
1154) {
1155    let spatial = h * w;
1156    for b in 0..batch {
1157        for i in 0..spatial {
1158            let mut mean = 0.0f32;
1159            for c in 0..channels {
1160                mean += input[((b * channels + c) * spatial) + i];
1161            }
1162            mean /= channels as f32;
1163            let mut var = 0.0f32;
1164            for c in 0..channels {
1165                let d = input[((b * channels + c) * spatial) + i] - mean;
1166                var += d * d;
1167            }
1168            var /= channels as f32;
1169            let inv = 1.0 / (var + eps).sqrt();
1170            for c in 0..channels {
1171                let v = (input[((b * channels + c) * spatial) + i] - mean) * inv;
1172                output[((b * channels + c) * spatial) + i] = v * gamma[c] + beta[c];
1173            }
1174        }
1175    }
1176}
1177
1178/// NCHW transposed convolution (PyTorch `ConvTranspose2d`, no bias).
1179/// Weight layout `[C_in, C_out/groups, kH, kW]`.
1180pub fn conv_transpose2d_nchw(
1181    input: &[f32],
1182    weight: &[f32],
1183    output: &mut [f32],
1184    n: usize,
1185    c_in: usize,
1186    h: usize,
1187    w: usize,
1188    c_out: usize,
1189    h_out: usize,
1190    w_out: usize,
1191    kh: usize,
1192    kw: usize,
1193    sh: usize,
1194    sw: usize,
1195    ph: usize,
1196    pw: usize,
1197    dh: usize,
1198    dw: usize,
1199    groups: usize,
1200) {
1201    output.fill(0.0);
1202    let c_in_per_g = c_in / groups;
1203    let c_out_per_g = c_out / groups;
1204    for ni in 0..n {
1205        for ic in 0..c_in {
1206            let g = ic / c_in_per_g;
1207            let _ic_off = ic % c_in_per_g;
1208            for iy in 0..h {
1209                for ix in 0..w {
1210                    let v = input[((ni * c_in + ic) * h + iy) * w + ix];
1211                    if v == 0.0 {
1212                        continue;
1213                    }
1214                    for ky in 0..kh {
1215                        let oy = iy * sh + ky * dh;
1216                        if oy < ph || oy >= h_out + ph {
1217                            continue;
1218                        }
1219                        let oy = oy - ph;
1220                        if oy >= h_out {
1221                            continue;
1222                        }
1223                        for kx in 0..kw {
1224                            let ox = ix * sw + kx * dw;
1225                            if ox < pw || ox >= w_out + pw {
1226                                continue;
1227                            }
1228                            let ox = ox - pw;
1229                            if ox >= w_out {
1230                                continue;
1231                            }
1232                            for oc_off in 0..c_out_per_g {
1233                                let oc = g * c_out_per_g + oc_off;
1234                                let w_idx = ((ic * c_out_per_g + oc_off) * kh + ky) * kw + kx;
1235                                let wt = weight[w_idx];
1236                                output[((ni * c_out + oc) * h_out + oy) * w_out + ox] += v * wt;
1237                            }
1238                        }
1239                    }
1240                }
1241            }
1242        }
1243    }
1244}
1245
1246/// NCDHW transposed convolution (the depth-axis analogue of
1247/// [`conv_transpose2d_nchw`]). Scatter form: each input voxel adds its
1248/// weighted kernel into the output. Weight `[C_in, C_out/g, kD, kH, kW]`.
1249/// `output_padding` is assumed folded into `d_out/h_out/w_out` by the caller.
1250#[allow(clippy::too_many_arguments)]
1251pub fn conv_transpose3d_ncdhw(
1252    input: &[f32],
1253    weight: &[f32],
1254    output: &mut [f32],
1255    n: usize,
1256    c_in: usize,
1257    d: usize,
1258    h: usize,
1259    w: usize,
1260    c_out: usize,
1261    d_out: usize,
1262    h_out: usize,
1263    w_out: usize,
1264    kd: usize,
1265    kh: usize,
1266    kw: usize,
1267    sd: usize,
1268    sh: usize,
1269    sw: usize,
1270    pd: usize,
1271    ph: usize,
1272    pw: usize,
1273    dd: usize,
1274    dh: usize,
1275    dw: usize,
1276    groups: usize,
1277) {
1278    output.fill(0.0);
1279    let c_in_per_g = c_in / groups;
1280    let c_out_per_g = c_out / groups;
1281    for ni in 0..n {
1282        for ic in 0..c_in {
1283            let g = ic / c_in_per_g;
1284            for id in 0..d {
1285                for iy in 0..h {
1286                    for ix in 0..w {
1287                        let v = input[(((ni * c_in + ic) * d + id) * h + iy) * w + ix];
1288                        if v == 0.0 {
1289                            continue;
1290                        }
1291                        for kz in 0..kd {
1292                            let oz = id * sd + kz * dd;
1293                            if oz < pd || oz >= d_out + pd {
1294                                continue;
1295                            }
1296                            let oz = oz - pd;
1297                            if oz >= d_out {
1298                                continue;
1299                            }
1300                            for ky in 0..kh {
1301                                let oy = iy * sh + ky * dh;
1302                                if oy < ph || oy >= h_out + ph {
1303                                    continue;
1304                                }
1305                                let oy = oy - ph;
1306                                if oy >= h_out {
1307                                    continue;
1308                                }
1309                                for kx in 0..kw {
1310                                    let ox = ix * sw + kx * dw;
1311                                    if ox < pw || ox >= w_out + pw {
1312                                        continue;
1313                                    }
1314                                    let ox = ox - pw;
1315                                    if ox >= w_out {
1316                                        continue;
1317                                    }
1318                                    for oc_off in 0..c_out_per_g {
1319                                        let oc = g * c_out_per_g + oc_off;
1320                                        let w_idx = (((ic * c_out_per_g + oc_off) * kd + kz) * kh
1321                                            + ky)
1322                                            * kw
1323                                            + kx;
1324                                        let wt = weight[w_idx];
1325                                        output[(((ni * c_out + oc) * d_out + oz) * h_out + oy)
1326                                            * w_out
1327                                            + ox] += v * wt;
1328                                    }
1329                                }
1330                            }
1331                        }
1332                    }
1333                }
1334            }
1335        }
1336    }
1337}
1338
1339/// NCHW group normalization: normalizes each `(C/G)×H×W` group.
1340pub fn group_norm_nchw(
1341    input: &[f32],
1342    gamma: &[f32],
1343    beta: &[f32],
1344    output: &mut [f32],
1345    batch: usize,
1346    channels: usize,
1347    h: usize,
1348    w: usize,
1349    num_groups: usize,
1350    eps: f32,
1351) {
1352    let cpg = channels / num_groups;
1353    let spatial = h * w;
1354    let n = (cpg * spatial) as f32;
1355    for b in 0..batch {
1356        for g in 0..num_groups {
1357            let c0 = g * cpg;
1358            let mut mean = 0.0f32;
1359            for c in 0..cpg {
1360                let plane = &input
1361                    [((b * channels + c0 + c) * spatial)..((b * channels + c0 + c + 1) * spatial)];
1362                mean += plane.iter().sum::<f32>();
1363            }
1364            mean /= n;
1365            let mut var = 0.0f32;
1366            for c in 0..cpg {
1367                let plane = &input
1368                    [((b * channels + c0 + c) * spatial)..((b * channels + c0 + c + 1) * spatial)];
1369                for &v in plane {
1370                    let d = v - mean;
1371                    var += d * d;
1372                }
1373            }
1374            var /= n;
1375            let inv = 1.0 / (var + eps).sqrt();
1376            for c in 0..cpg {
1377                let gi = c0 + c;
1378                let gamm = gamma[gi];
1379                let bet = beta[gi];
1380                let src =
1381                    &input[((b * channels + gi) * spatial)..((b * channels + gi + 1) * spatial)];
1382                let dst = &mut output
1383                    [((b * channels + gi) * spatial)..((b * channels + gi + 1) * spatial)];
1384                for (d, &s) in dst.iter_mut().zip(src) {
1385                    *d = (s - mean) * inv * gamm + bet;
1386                }
1387            }
1388        }
1389    }
1390}
1391
1392/// Nearest-neighbor 2× upsample on planar NCHW.
1393pub fn resize_nearest_2x_nchw(
1394    input: &[f32],
1395    output: &mut [f32],
1396    channels: usize,
1397    h: usize,
1398    w: usize,
1399) {
1400    let h2 = h * 2;
1401    let w2 = w * 2;
1402    for c in 0..channels {
1403        let plane = &input[c * h * w..(c + 1) * h * w];
1404        let dst = &mut output[c * h2 * w2..(c + 1) * h2 * w2];
1405        for y in 0..h {
1406            for x in 0..w {
1407                let v = plane[y * w + x];
1408                for dy in 0..2 {
1409                    for dx in 0..2 {
1410                        dst[(y * 2 + dy) * w2 + (x * 2 + dx)] = v;
1411                    }
1412                }
1413            }
1414        }
1415    }
1416}
1417
1418#[cfg(test)]
1419mod tests {
1420    use super::*;
1421
1422    #[test]
1423    fn gelu_correctness() {
1424        let x = 1.5f32;
1425        let g = scalar_gelu(x);
1426        // Reference: gelu(1.5) ≈ 1.3990
1427        assert!((g - 1.3990).abs() < 0.01, "gelu(1.5) = {g}");
1428    }
1429
1430    #[test]
1431    fn bias_gelu_works() {
1432        let mut data = vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0];
1433        let bias = vec![0.1, 0.2, 0.3, 0.4];
1434        bias_gelu(&mut data, &bias, 2, 4);
1435        // After bias+gelu, values should be > 0 (all inputs positive)
1436        for &v in &data {
1437            assert!(v > 0.0, "bias_gelu produced {v}");
1438        }
1439    }
1440
1441    #[test]
1442    fn batch_norm_inference_roundtrip() {
1443        let c = 4usize;
1444        let x: Vec<f32> = vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0];
1445        let gamma = vec![1.0; c];
1446        let beta = vec![0.0; c];
1447        let mean = vec![2.5, 2.5, 2.5, 2.5];
1448        let var = vec![1.0; c];
1449        let mut y = vec![0.0; 8];
1450        batch_norm_inference(&x, &gamma, &beta, &mean, &var, &mut y, c, 1e-5);
1451        let mut dx = vec![0.0; 8];
1452        let dy = vec![1.0; 8];
1453        let mut dgamma = vec![0.0; c];
1454        let mut dbeta = vec![0.0; c];
1455        batch_norm_inference_backward_input(&x, &gamma, &mean, &var, &dy, &mut dx, c, 1e-5);
1456        batch_norm_inference_backward_gamma(&x, &mean, &var, &dy, &mut dgamma, c, 1e-5);
1457        batch_norm_inference_backward_beta(&dy, &mut dbeta, c);
1458        assert!(y.iter().all(|v| v.is_finite()));
1459        assert!(dx.iter().all(|v| v.is_finite()));
1460        assert!(dgamma.iter().any(|&v| v.abs() > 1e-6));
1461        assert_eq!(dbeta, vec![2.0, 2.0, 2.0, 2.0]);
1462    }
1463
1464    #[test]
1465    fn layer_norm_unit_test() {
1466        let input = vec![1.0, 2.0, 3.0, 4.0];
1467        let gamma = vec![1.0; 4];
1468        let beta = vec![0.0; 4];
1469        let mut output = vec![0.0; 4];
1470        layer_norm_row(&input, &gamma, &beta, &mut output, 4, 1e-5);
1471        // Mean=2.5, std≈1.118. output ≈ [-1.342, -0.447, 0.447, 1.342]
1472        assert!((output[0] - -1.342).abs() < 0.01);
1473        assert!((output[3] - 1.342).abs() < 0.01);
1474        // Sum should be ~0 (normalized)
1475        let sum: f32 = output.iter().sum();
1476        assert!(sum.abs() < 0.01, "LN sum should be ~0, got {sum}");
1477    }
1478
1479    #[test]
1480    fn par_bias_gelu_matches_sequential() {
1481        let n = 100;
1482        let m = 64;
1483        let mut data_par = vec![0.5f32; n * m];
1484        let mut data_seq = data_par.clone();
1485        let bias = vec![0.1f32; m];
1486
1487        bias_gelu(&mut data_seq, &bias, n, m);
1488        par_bias_gelu(&mut data_par, &bias, n, m);
1489
1490        let max_diff: f32 = data_par
1491            .iter()
1492            .zip(data_seq.iter())
1493            .map(|(a, b)| (a - b).abs())
1494            .fold(0f32, f32::max);
1495        assert!(max_diff < 1e-6, "par vs seq diff: {max_diff}");
1496    }
1497}