1use crate::pool;
23
24#[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#[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 let n = _mm256_round_ps::<{ _MM_FROUND_TO_NEAREST_INT | _MM_FROUND_NO_EXC }>(_mm256_mul_ps(
67 x, inv_ln2,
68 ));
69 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 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#[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 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 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 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
214pub 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#[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#[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 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#[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 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
448pub 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
473pub 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
494pub 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
516pub 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
527pub 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 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
551pub 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
580pub 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#[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#[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 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 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 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 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 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 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#[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#[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
888const ACTIVATION_PAR_MIN: usize = 1 << 20;
899
900#[inline]
909pub fn scalar_gelu_approx(x: f32) -> f32 {
910 const C: f32 = 0.797_884_6; 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
1022pub 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#[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 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 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#[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
1126fn 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
1142pub 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
1178pub 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#[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
1339pub 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
1392pub 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 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 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 assert!((output[0] - -1.342).abs() < 0.01);
1473 assert!((output[3] - 1.342).abs() < 0.01);
1474 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}