1use crate::error::TruenoError;
16
17pub fn relu(input: &[f32], output: &mut [f32]) -> Result<(), TruenoError> {
29 contract_pre_relu!(input);
30 let n = input.len();
31 if n != output.len() {
32 return Err(TruenoError::InvalidInput(format!(
33 "relu size mismatch: input[{}], output[{}]",
34 n,
35 output.len()
36 )));
37 }
38
39 #[cfg(target_arch = "x86_64")]
40 {
41 if n > 4096 {
48 relu_autovec(input, output);
49 contract_post_elementwise_parity!(output);
50 return Ok(());
51 }
52 if is_x86_feature_detected!("avx512f") {
53 unsafe {
54 relu_avx512(input, output);
55 }
56 contract_post_elementwise_parity!(output);
57 return Ok(());
58 }
59 if is_x86_feature_detected!("avx2") {
60 unsafe {
61 relu_avx2(input, output);
62 }
63 contract_post_elementwise_parity!(output);
64 return Ok(());
65 }
66 }
67
68 relu_autovec(input, output);
69 contract_post_elementwise_parity!(output);
70 Ok(())
71}
72
73#[inline]
78fn relu_autovec(input: &[f32], output: &mut [f32]) {
79 for i in 0..input.len() {
80 output[i] = input[i].max(0.0);
81 }
82}
83
84#[cfg(target_arch = "x86_64")]
86#[target_feature(enable = "avx512f")]
87unsafe fn relu_avx512(input: &[f32], output: &mut [f32]) {
88 use std::arch::x86_64::*;
89 unsafe {
90 let n = input.len();
91 let ip = input.as_ptr();
92 let op = output.as_mut_ptr();
93 let zero = _mm512_setzero_ps();
94 let mut i = 0;
95
96 let data_bytes = n * 4;
97 let op_aligned = (op as usize) % 64 == 0;
98 if data_bytes > NT_STORE_THRESHOLD_BYTES && op_aligned {
99 while i + 64 <= n {
101 _mm_prefetch(ip.add(i + 128).cast::<i8>(), _MM_HINT_T0);
102
103 _mm512_stream_ps(op.add(i), _mm512_max_ps(_mm512_loadu_ps(ip.add(i)), zero));
104 _mm512_stream_ps(
105 op.add(i + 16),
106 _mm512_max_ps(_mm512_loadu_ps(ip.add(i + 16)), zero),
107 );
108 _mm512_stream_ps(
109 op.add(i + 32),
110 _mm512_max_ps(_mm512_loadu_ps(ip.add(i + 32)), zero),
111 );
112 _mm512_stream_ps(
113 op.add(i + 48),
114 _mm512_max_ps(_mm512_loadu_ps(ip.add(i + 48)), zero),
115 );
116 i += 64;
117 }
118 while i + 16 <= n {
119 _mm512_stream_ps(op.add(i), _mm512_max_ps(_mm512_loadu_ps(ip.add(i)), zero));
120 i += 16;
121 }
122 _mm_sfence();
123 } else {
124 while i + 64 <= n {
125 _mm512_storeu_ps(op.add(i), _mm512_max_ps(_mm512_loadu_ps(ip.add(i)), zero));
126 _mm512_storeu_ps(
127 op.add(i + 16),
128 _mm512_max_ps(_mm512_loadu_ps(ip.add(i + 16)), zero),
129 );
130 _mm512_storeu_ps(
131 op.add(i + 32),
132 _mm512_max_ps(_mm512_loadu_ps(ip.add(i + 32)), zero),
133 );
134 _mm512_storeu_ps(
135 op.add(i + 48),
136 _mm512_max_ps(_mm512_loadu_ps(ip.add(i + 48)), zero),
137 );
138 i += 64;
139 }
140 while i + 16 <= n {
141 _mm512_storeu_ps(op.add(i), _mm512_max_ps(_mm512_loadu_ps(ip.add(i)), zero));
142 i += 16;
143 }
144 }
145 for j in i..n {
146 output[j] = input[j].max(0.0);
147 }
148 } }
150
151#[cfg(target_arch = "x86_64")]
155const PREFETCH_DISTANCE: usize = 512;
156
157#[cfg(target_arch = "x86_64")]
163const NT_STORE_THRESHOLD_BYTES: usize = 512 * 1024; #[cfg(target_arch = "x86_64")]
166#[target_feature(enable = "avx2")]
167unsafe fn relu_avx2(input: &[f32], output: &mut [f32]) {
168 use std::arch::x86_64::*;
169
170 let n = input.len();
171 let data_bytes = n * 4;
172
173 let out_aligned = (output.as_ptr() as usize) % 32 == 0;
177 if data_bytes > NT_STORE_THRESHOLD_BYTES && out_aligned {
178 unsafe { relu_avx2_nt(input, output) }
179 return;
180 }
181
182 let chunks = n / 64;
188 let remainder_64 = chunks * 64;
189
190 unsafe {
191 let zero = _mm256_setzero_ps();
192 let inp = input.as_ptr();
193 let out = output.as_mut_ptr();
194
195 for i in 0..chunks {
196 let base = i * 64;
197 let v0 = _mm256_loadu_ps(inp.add(base));
198 let v1 = _mm256_loadu_ps(inp.add(base + 8));
199 let v2 = _mm256_loadu_ps(inp.add(base + 16));
200 let v3 = _mm256_loadu_ps(inp.add(base + 24));
201 let v4 = _mm256_loadu_ps(inp.add(base + 32));
202 let v5 = _mm256_loadu_ps(inp.add(base + 40));
203 let v6 = _mm256_loadu_ps(inp.add(base + 48));
204 let v7 = _mm256_loadu_ps(inp.add(base + 56));
205 _mm256_storeu_ps(out.add(base), _mm256_max_ps(v0, zero));
206 _mm256_storeu_ps(out.add(base + 8), _mm256_max_ps(v1, zero));
207 _mm256_storeu_ps(out.add(base + 16), _mm256_max_ps(v2, zero));
208 _mm256_storeu_ps(out.add(base + 24), _mm256_max_ps(v3, zero));
209 _mm256_storeu_ps(out.add(base + 32), _mm256_max_ps(v4, zero));
210 _mm256_storeu_ps(out.add(base + 40), _mm256_max_ps(v5, zero));
211 _mm256_storeu_ps(out.add(base + 48), _mm256_max_ps(v6, zero));
212 _mm256_storeu_ps(out.add(base + 56), _mm256_max_ps(v7, zero));
213 }
214
215 let mut i = remainder_64;
216 while i + 8 <= n {
217 let v = _mm256_loadu_ps(inp.add(i));
218 _mm256_storeu_ps(out.add(i), _mm256_max_ps(v, zero));
219 i += 8;
220 }
221
222 while i < n {
223 *out.add(i) = (*inp.add(i)).max(0.0);
224 i += 1;
225 }
226 }
227}
228
229#[cfg(target_arch = "x86_64")]
234#[target_feature(enable = "avx2")]
235unsafe fn relu_avx2_nt(input: &[f32], output: &mut [f32]) {
236 use std::arch::x86_64::*;
237
238 let n = input.len();
239 let chunks = n / 32;
240 let remainder_32 = chunks * 32;
241
242 unsafe {
243 let zero = _mm256_setzero_ps();
244
245 for i in 0..chunks {
246 let base = i * 32;
247 _mm_prefetch(
249 input.as_ptr().add(base + PREFETCH_DISTANCE / 4) as *const i8,
250 _MM_HINT_T0,
251 );
252 let v0 = _mm256_loadu_ps(input.as_ptr().add(base));
253 let v1 = _mm256_loadu_ps(input.as_ptr().add(base + 8));
254 let v2 = _mm256_loadu_ps(input.as_ptr().add(base + 16));
255 let v3 = _mm256_loadu_ps(input.as_ptr().add(base + 24));
256 _mm256_stream_ps(output.as_mut_ptr().add(base), _mm256_max_ps(v0, zero));
258 _mm256_stream_ps(output.as_mut_ptr().add(base + 8), _mm256_max_ps(v1, zero));
259 _mm256_stream_ps(output.as_mut_ptr().add(base + 16), _mm256_max_ps(v2, zero));
260 _mm256_stream_ps(output.as_mut_ptr().add(base + 24), _mm256_max_ps(v3, zero));
261 }
262
263 _mm_sfence();
265
266 let mut i = remainder_32;
268 while i + 8 <= n {
269 let v = _mm256_loadu_ps(input.as_ptr().add(i));
270 _mm256_storeu_ps(output.as_mut_ptr().add(i), _mm256_max_ps(v, zero));
271 i += 8;
272 }
273 while i < n {
274 output[i] = input[i].max(0.0);
275 i += 1;
276 }
277 }
278}
279
280pub fn add(a: &[f32], b: &[f32], output: &mut [f32]) -> Result<(), TruenoError> {
292 let n = a.len();
293 if n != b.len() || n != output.len() {
294 return Err(TruenoError::InvalidInput(format!(
295 "add size mismatch: a[{}], b[{}], output[{}]",
296 n,
297 b.len(),
298 output.len()
299 )));
300 }
301 contract_pre_add!(a, b);
302
303 #[cfg(target_arch = "x86_64")]
304 {
305 if n > 4096 {
309 add_autovec(a, b, output);
310 return Ok(());
311 }
312 if is_x86_feature_detected!("avx512f") {
313 unsafe {
314 add_avx512(a, b, output);
315 }
316 return Ok(());
317 }
318 if is_x86_feature_detected!("avx2") {
319 unsafe {
320 add_avx2(a, b, output);
321 }
322 return Ok(());
323 }
324 }
325
326 add_autovec(a, b, output);
327 contract_post_elementwise_parity!(output);
328 Ok(())
329}
330
331#[inline]
333fn add_autovec(a: &[f32], b: &[f32], output: &mut [f32]) {
334 for i in 0..a.len() {
335 output[i] = a[i] + b[i];
336 }
337}
338
339#[cfg(target_arch = "x86_64")]
341#[target_feature(enable = "avx512f")]
342unsafe fn add_avx512(a: &[f32], b: &[f32], output: &mut [f32]) {
343 use std::arch::x86_64::*;
344 unsafe {
345 let n = a.len();
346 let ap = a.as_ptr();
347 let bp = b.as_ptr();
348 let rp = output.as_mut_ptr();
349 let mut i = 0;
350
351 let data_bytes = n * 4;
352 let rp_aligned = (rp as usize) % 64 == 0;
353 if data_bytes > NT_STORE_THRESHOLD_BYTES && rp_aligned {
354 while i + 64 <= n {
356 if i + 128 <= n {
358 _mm_prefetch(ap.add(i + 128).cast::<i8>(), _MM_HINT_T0);
359 _mm_prefetch(bp.add(i + 128).cast::<i8>(), _MM_HINT_T0);
360 }
361
362 _mm512_stream_ps(
363 rp.add(i),
364 _mm512_add_ps(_mm512_loadu_ps(ap.add(i)), _mm512_loadu_ps(bp.add(i))),
365 );
366 _mm512_stream_ps(
367 rp.add(i + 16),
368 _mm512_add_ps(_mm512_loadu_ps(ap.add(i + 16)), _mm512_loadu_ps(bp.add(i + 16))),
369 );
370 _mm512_stream_ps(
371 rp.add(i + 32),
372 _mm512_add_ps(_mm512_loadu_ps(ap.add(i + 32)), _mm512_loadu_ps(bp.add(i + 32))),
373 );
374 _mm512_stream_ps(
375 rp.add(i + 48),
376 _mm512_add_ps(_mm512_loadu_ps(ap.add(i + 48)), _mm512_loadu_ps(bp.add(i + 48))),
377 );
378 i += 64;
379 }
380 while i + 16 <= n {
381 _mm512_stream_ps(
382 rp.add(i),
383 _mm512_add_ps(_mm512_loadu_ps(ap.add(i)), _mm512_loadu_ps(bp.add(i))),
384 );
385 i += 16;
386 }
387 _mm_sfence();
388 } else {
389 while i + 64 <= n {
390 _mm512_storeu_ps(
391 rp.add(i),
392 _mm512_add_ps(_mm512_loadu_ps(ap.add(i)), _mm512_loadu_ps(bp.add(i))),
393 );
394 _mm512_storeu_ps(
395 rp.add(i + 16),
396 _mm512_add_ps(_mm512_loadu_ps(ap.add(i + 16)), _mm512_loadu_ps(bp.add(i + 16))),
397 );
398 _mm512_storeu_ps(
399 rp.add(i + 32),
400 _mm512_add_ps(_mm512_loadu_ps(ap.add(i + 32)), _mm512_loadu_ps(bp.add(i + 32))),
401 );
402 _mm512_storeu_ps(
403 rp.add(i + 48),
404 _mm512_add_ps(_mm512_loadu_ps(ap.add(i + 48)), _mm512_loadu_ps(bp.add(i + 48))),
405 );
406 i += 64;
407 }
408 while i + 16 <= n {
409 _mm512_storeu_ps(
410 rp.add(i),
411 _mm512_add_ps(_mm512_loadu_ps(ap.add(i)), _mm512_loadu_ps(bp.add(i))),
412 );
413 i += 16;
414 }
415 }
416 for j in i..n {
417 output[j] = a[j] + b[j];
418 }
419 } }
421
422#[cfg(target_arch = "x86_64")]
423#[target_feature(enable = "avx2")]
424unsafe fn add_avx2(a: &[f32], b: &[f32], output: &mut [f32]) {
425 use std::arch::x86_64::*;
426
427 let n = a.len();
428 let data_bytes = n * 4;
429
430 let out_aligned = (output.as_ptr() as usize) % 32 == 0;
433 if data_bytes > NT_STORE_THRESHOLD_BYTES && out_aligned {
434 unsafe { add_avx2_nt(a, b, output) }
435 return;
436 }
437
438 let chunks = n / 64;
441 let remainder_64 = chunks * 64;
442
443 unsafe {
444 let ap = a.as_ptr();
445 let bp = b.as_ptr();
446 let op = output.as_mut_ptr();
447
448 for i in 0..chunks {
449 let base = i * 64;
450 let a0 = _mm256_loadu_ps(ap.add(base));
452 let b0 = _mm256_loadu_ps(bp.add(base));
453 let a1 = _mm256_loadu_ps(ap.add(base + 8));
454 let b1 = _mm256_loadu_ps(bp.add(base + 8));
455 let a2 = _mm256_loadu_ps(ap.add(base + 16));
456 let b2 = _mm256_loadu_ps(bp.add(base + 16));
457 let a3 = _mm256_loadu_ps(ap.add(base + 24));
458 let b3 = _mm256_loadu_ps(bp.add(base + 24));
459 let a4 = _mm256_loadu_ps(ap.add(base + 32));
460 let b4 = _mm256_loadu_ps(bp.add(base + 32));
461 let a5 = _mm256_loadu_ps(ap.add(base + 40));
462 let b5 = _mm256_loadu_ps(bp.add(base + 40));
463 let a6 = _mm256_loadu_ps(ap.add(base + 48));
464 let b6 = _mm256_loadu_ps(bp.add(base + 48));
465 let a7 = _mm256_loadu_ps(ap.add(base + 56));
466 let b7 = _mm256_loadu_ps(bp.add(base + 56));
467 _mm256_storeu_ps(op.add(base), _mm256_add_ps(a0, b0));
468 _mm256_storeu_ps(op.add(base + 8), _mm256_add_ps(a1, b1));
469 _mm256_storeu_ps(op.add(base + 16), _mm256_add_ps(a2, b2));
470 _mm256_storeu_ps(op.add(base + 24), _mm256_add_ps(a3, b3));
471 _mm256_storeu_ps(op.add(base + 32), _mm256_add_ps(a4, b4));
472 _mm256_storeu_ps(op.add(base + 40), _mm256_add_ps(a5, b5));
473 _mm256_storeu_ps(op.add(base + 48), _mm256_add_ps(a6, b6));
474 _mm256_storeu_ps(op.add(base + 56), _mm256_add_ps(a7, b7));
475 }
476
477 let mut i = remainder_64;
478 while i + 8 <= n {
479 let av = _mm256_loadu_ps(ap.add(i));
480 let bv = _mm256_loadu_ps(bp.add(i));
481 _mm256_storeu_ps(op.add(i), _mm256_add_ps(av, bv));
482 i += 8;
483 }
484
485 while i < n {
486 *op.add(i) = *ap.add(i) + *bp.add(i);
487 i += 1;
488 }
489 }
490}
491
492#[cfg(target_arch = "x86_64")]
494#[target_feature(enable = "avx2")]
495unsafe fn add_avx2_nt(a: &[f32], b: &[f32], output: &mut [f32]) {
496 use std::arch::x86_64::*;
497
498 let n = a.len();
499 let chunks = n / 32;
500 let remainder_32 = chunks * 32;
501
502 unsafe {
503 for i in 0..chunks {
504 let base = i * 32;
505 _mm_prefetch(a.as_ptr().add(base + PREFETCH_DISTANCE / 4) as *const i8, _MM_HINT_T0);
506 _mm_prefetch(b.as_ptr().add(base + PREFETCH_DISTANCE / 4) as *const i8, _MM_HINT_T0);
507 let a0 = _mm256_loadu_ps(a.as_ptr().add(base));
508 let a1 = _mm256_loadu_ps(a.as_ptr().add(base + 8));
509 let a2 = _mm256_loadu_ps(a.as_ptr().add(base + 16));
510 let a3 = _mm256_loadu_ps(a.as_ptr().add(base + 24));
511 let b0 = _mm256_loadu_ps(b.as_ptr().add(base));
512 let b1 = _mm256_loadu_ps(b.as_ptr().add(base + 8));
513 let b2 = _mm256_loadu_ps(b.as_ptr().add(base + 16));
514 let b3 = _mm256_loadu_ps(b.as_ptr().add(base + 24));
515 _mm256_stream_ps(output.as_mut_ptr().add(base), _mm256_add_ps(a0, b0));
516 _mm256_stream_ps(output.as_mut_ptr().add(base + 8), _mm256_add_ps(a1, b1));
517 _mm256_stream_ps(output.as_mut_ptr().add(base + 16), _mm256_add_ps(a2, b2));
518 _mm256_stream_ps(output.as_mut_ptr().add(base + 24), _mm256_add_ps(a3, b3));
519 }
520
521 _mm_sfence();
522
523 let mut i = remainder_32;
524 while i + 8 <= n {
525 let av = _mm256_loadu_ps(a.as_ptr().add(i));
526 let bv = _mm256_loadu_ps(b.as_ptr().add(i));
527 _mm256_storeu_ps(output.as_mut_ptr().add(i), _mm256_add_ps(av, bv));
528 i += 8;
529 }
530 while i < n {
531 output[i] = a[i] + b[i];
532 i += 1;
533 }
534 }
535}
536
537pub fn mul_scalar(input: &[f32], scalar: f32, output: &mut [f32]) -> Result<(), TruenoError> {
549 debug_assert!(!input.is_empty(), "Contract mul_scalar: input is empty");
551 debug_assert!(scalar.is_finite(), "Contract mul_scalar: scalar is not finite");
552 let n = input.len();
553 if n != output.len() {
554 return Err(TruenoError::InvalidInput(format!(
555 "mul_scalar size mismatch: input[{}], output[{}]",
556 n,
557 output.len()
558 )));
559 }
560
561 #[cfg(target_arch = "x86_64")]
562 {
563 if is_x86_feature_detected!("avx2") {
564 unsafe {
565 mul_scalar_avx2(input, scalar, output);
566 }
567 return Ok(());
568 }
569 }
570
571 for i in 0..n {
572 output[i] = input[i] * scalar;
573 }
574 Ok(())
575}
576
577#[cfg(target_arch = "x86_64")]
578#[target_feature(enable = "avx2")]
579unsafe fn mul_scalar_avx2(input: &[f32], scalar: f32, output: &mut [f32]) {
580 use std::arch::x86_64::*;
581
582 let n = input.len();
583 let chunks = n / 32;
584 let remainder_32 = chunks * 32;
585
586 unsafe {
587 let s = _mm256_set1_ps(scalar);
588
589 for i in 0..chunks {
590 let base = i * 32;
591 let v0 = _mm256_loadu_ps(input.as_ptr().add(base));
592 let v1 = _mm256_loadu_ps(input.as_ptr().add(base + 8));
593 let v2 = _mm256_loadu_ps(input.as_ptr().add(base + 16));
594 let v3 = _mm256_loadu_ps(input.as_ptr().add(base + 24));
595 _mm256_storeu_ps(output.as_mut_ptr().add(base), _mm256_mul_ps(v0, s));
596 _mm256_storeu_ps(output.as_mut_ptr().add(base + 8), _mm256_mul_ps(v1, s));
597 _mm256_storeu_ps(output.as_mut_ptr().add(base + 16), _mm256_mul_ps(v2, s));
598 _mm256_storeu_ps(output.as_mut_ptr().add(base + 24), _mm256_mul_ps(v3, s));
599 }
600
601 let mut i = remainder_32;
602 while i + 8 <= n {
603 let v = _mm256_loadu_ps(input.as_ptr().add(i));
604 _mm256_storeu_ps(output.as_mut_ptr().add(i), _mm256_mul_ps(v, s));
605 i += 8;
606 }
607
608 while i < n {
609 output[i] = input[i] * scalar;
610 i += 1;
611 }
612 }
613}
614
615#[must_use]
625pub fn relu_alloc(input: &[f32]) -> Vec<f32> {
626 let n = input.len();
627 let mut output = vec![0.0f32; n];
628 let _ = relu(input, &mut output);
629 output
630}
631
632#[must_use]
638pub fn add_alloc(a: &[f32], b: &[f32]) -> Vec<f32> {
639 assert_eq!(a.len(), b.len(), "add_alloc: length mismatch");
640 let n = a.len();
641 let mut output = vec![0.0f32; n];
642 let _ = add(a, b, &mut output);
643 output
644}
645
646#[must_use]
648pub fn mul_scalar_alloc(input: &[f32], scalar: f32) -> Vec<f32> {
649 let n = input.len();
650 let mut output = vec![0.0f32; n];
651 let _ = mul_scalar(input, scalar, &mut output);
652 output
653}
654
655pub fn fused_add_relu(a: &[f32], b: &[f32], output: &mut [f32]) -> Result<(), TruenoError> {
673 let n = a.len();
674 if n != b.len() || n != output.len() {
675 return Err(TruenoError::InvalidInput(format!(
676 "fused_add_relu size mismatch: a[{}], b[{}], output[{}]",
677 n,
678 b.len(),
679 output.len()
680 )));
681 }
682 for i in 0..n {
684 output[i] = (a[i] + b[i]).max(0.0);
685 }
686 Ok(())
687}
688
689pub fn fused_mul_add(
699 a: &[f32],
700 b: &[f32],
701 c: &[f32],
702 output: &mut [f32],
703) -> Result<(), TruenoError> {
704 let n = a.len();
705 if n != b.len() || n != c.len() || n != output.len() {
706 return Err(TruenoError::InvalidInput(format!(
707 "fused_mul_add size mismatch: a[{}], b[{}], c[{}], output[{}]",
708 n,
709 b.len(),
710 c.len(),
711 output.len()
712 )));
713 }
714 for i in 0..n {
715 output[i] = a[i].mul_add(b[i], c[i]);
716 }
717 Ok(())
718}
719
720pub fn fused_scale_bias_relu(
731 input: &[f32],
732 scale: f32,
733 bias: f32,
734 output: &mut [f32],
735) -> Result<(), TruenoError> {
736 let n = input.len();
737 if n != output.len() {
738 return Err(TruenoError::InvalidInput(format!(
739 "fused_scale_bias_relu size mismatch: input[{}], output[{}]",
740 n,
741 output.len()
742 )));
743 }
744 for i in 0..n {
745 output[i] = input[i].mul_add(scale, bias).max(0.0);
746 }
747 Ok(())
748}
749
750#[inline]
760pub fn relu_inplace(data: &mut [f32]) {
761 for x in data.iter_mut() {
762 *x = x.max(0.0);
763 }
764}
765
766pub fn add_inplace(a: &mut [f32], b: &[f32]) -> Result<(), TruenoError> {
770 if a.len() != b.len() {
771 return Err(TruenoError::InvalidInput(format!(
772 "add_inplace size mismatch: a[{}], b[{}]",
773 a.len(),
774 b.len()
775 )));
776 }
777 for i in 0..a.len() {
778 a[i] += b[i];
779 }
780 Ok(())
781}
782
783#[inline]
787pub fn scale_inplace(data: &mut [f32], scalar: f32) {
788 for x in data.iter_mut() {
789 *x *= scalar;
790 }
791}
792
793pub fn fused_add_relu_inplace(a: &mut [f32], b: &[f32]) -> Result<(), TruenoError> {
798 if a.len() != b.len() {
799 return Err(TruenoError::InvalidInput(format!(
800 "fused_add_relu_inplace size mismatch: a[{}], b[{}]",
801 a.len(),
802 b.len()
803 )));
804 }
805 for i in 0..a.len() {
806 a[i] = (a[i] + b[i]).max(0.0);
807 }
808 Ok(())
809}
810
811#[cfg(test)]
816mod tests {
817 use super::*;
818
819 #[test]
822 fn test_relu_basic() {
823 let input = [-1.0, 0.0, 1.0, -0.5, 2.0, -3.0, 0.1, -0.1];
824 let expected = [0.0, 0.0, 1.0, 0.0, 2.0, 0.0, 0.1, 0.0];
825 let mut output = vec![0.0f32; 8];
826 relu(&input, &mut output).unwrap();
827 assert_eq!(output, expected);
828 }
829
830 #[test]
831 fn test_relu_large() {
832 let n = 11008; let input: Vec<f32> =
834 (0..n).map(|i| ((i * 17 + 31) % 1000) as f32 / 1000.0 - 0.5).collect();
835 let mut output = vec![0.0f32; n];
836 relu(&input, &mut output).unwrap();
837 for (i, (&inp, &out)) in input.iter().zip(output.iter()).enumerate() {
838 assert_eq!(out, inp.max(0.0), "ReLU mismatch at {i}");
839 }
840 }
841
842 #[test]
843 fn test_relu_avx2_scalar_parity() {
844 for n in [1, 7, 8, 15, 16, 31, 32, 63, 64, 128, 4096] {
845 let input: Vec<f32> =
846 (0..n).map(|i| ((i * 17 + 31) % 1000) as f32 / 500.0 - 1.0).collect();
847 let mut output = vec![0.0f32; n];
848 relu(&input, &mut output).unwrap();
849 for (i, (&inp, &out)) in input.iter().zip(output.iter()).enumerate() {
850 assert_eq!(out, inp.max(0.0), "ReLU parity at [{i}] n={n}");
851 }
852 }
853 }
854
855 #[test]
856 fn test_relu_error_mismatch() {
857 let input = vec![1.0f32; 4];
858 let mut output = vec![0.0f32; 3];
859 assert!(relu(&input, &mut output).is_err());
860 }
861
862 #[test]
865 fn test_add_basic() {
866 let a = [1.0, 2.0, 3.0, 4.0];
867 let b = [10.0, 20.0, 30.0, 40.0];
868 let mut output = vec![0.0f32; 4];
869 add(&a, &b, &mut output).unwrap();
870 assert_eq!(output, vec![11.0, 22.0, 33.0, 44.0]);
871 }
872
873 #[test]
874 fn test_add_large() {
875 let n = 4096;
876 let a: Vec<f32> = (0..n).map(|i| i as f32).collect();
877 let b: Vec<f32> = (0..n).map(|i| (i * 2) as f32).collect();
878 let mut output = vec![0.0f32; n];
879 add(&a, &b, &mut output).unwrap();
880 for i in 0..n {
881 assert_eq!(output[i], a[i] + b[i], "Add mismatch at {i}");
882 }
883 }
884
885 #[test]
886 fn test_add_avx2_scalar_parity() {
887 for n in [1, 7, 8, 15, 16, 31, 32, 63, 64, 128, 4096] {
888 let a: Vec<f32> = (0..n).map(|i| ((i * 17 + 31) % 1000) as f32 / 500.0 - 1.0).collect();
889 let b: Vec<f32> = (0..n).map(|i| ((i * 13 + 7) % 1000) as f32 / 500.0 - 1.0).collect();
890 let mut output = vec![0.0f32; n];
891 add(&a, &b, &mut output).unwrap();
892 for i in 0..n {
893 assert_eq!(output[i], a[i] + b[i], "Add parity at [{i}] n={n}");
894 }
895 }
896 }
897
898 #[test]
899 fn test_add_error_mismatch() {
900 let a = vec![1.0f32; 4];
901 let b = vec![1.0f32; 3];
902 let mut output = vec![0.0f32; 4];
903 assert!(add(&a, &b, &mut output).is_err());
904 }
905
906 #[test]
909 fn test_mul_scalar_basic() {
910 let input = [1.0, 2.0, 3.0, 4.0];
911 let mut output = vec![0.0f32; 4];
912 mul_scalar(&input, 2.5, &mut output).unwrap();
913 assert_eq!(output, vec![2.5, 5.0, 7.5, 10.0]);
914 }
915
916 #[test]
917 fn test_mul_scalar_large() {
918 let n = 4096;
919 let input: Vec<f32> = (0..n).map(|i| i as f32).collect();
920 let mut output = vec![0.0f32; n];
921 mul_scalar(&input, std::f32::consts::PI, &mut output).unwrap();
922 for i in 0..n {
923 assert!(
924 (output[i] - input[i] * std::f32::consts::PI).abs() < 1e-5,
925 "Mul scalar mismatch at {i}"
926 );
927 }
928 }
929
930 #[test]
931 fn test_mul_scalar_avx2_scalar_parity() {
932 for n in [1, 7, 8, 15, 16, 31, 32, 63, 64, 128, 4096] {
933 let input: Vec<f32> =
934 (0..n).map(|i| ((i * 17 + 31) % 1000) as f32 / 500.0 - 1.0).collect();
935 let mut output = vec![0.0f32; n];
936 mul_scalar(&input, std::f32::consts::E, &mut output).unwrap();
937 for i in 0..n {
938 assert!(
939 (output[i] - input[i] * std::f32::consts::E).abs() < 1e-4,
940 "Mul scalar parity at [{i}] n={n}",
941 );
942 }
943 }
944 }
945
946 #[test]
947 fn test_mul_scalar_error_mismatch() {
948 let input = vec![1.0f32; 4];
949 let mut output = vec![0.0f32; 3];
950 assert!(mul_scalar(&input, 1.0, &mut output).is_err());
951 }
952
953 #[test]
956 fn test_fused_add_relu_basic() {
957 let a = vec![-2.0, -1.0, 0.0, 1.0, 2.0, -0.5, 0.5, 3.0];
958 let b = vec![1.0, 0.5, -1.0, -2.0, 0.0, 1.0, -1.0, -4.0];
959 let mut out = vec![0.0f32; 8];
960 fused_add_relu(&a, &b, &mut out).unwrap();
961 let expected: Vec<f32> = a.iter().zip(&b).map(|(a, b)| (a + b).max(0.0)).collect();
962 assert_eq!(out, expected);
963 }
964
965 #[test]
966 fn test_fused_add_relu_large() {
967 let n = 10_000;
968 let a: Vec<f32> = (0..n).map(|i| (i as f32 - 5000.0) / 100.0).collect();
969 let b: Vec<f32> = (0..n).map(|i| (i as f32 * 0.3) - 1500.0).collect();
970 let mut out = vec![0.0f32; n];
971 fused_add_relu(&a, &b, &mut out).unwrap();
972 for i in 0..n {
973 assert_eq!(out[i], (a[i] + b[i]).max(0.0), "mismatch at {i}");
974 }
975 }
976
977 #[test]
978 fn test_fused_mul_add_basic() {
979 let a = vec![1.0, 2.0, 3.0, 4.0];
980 let b = vec![2.0, 3.0, 4.0, 5.0];
981 let c = vec![0.5, 0.5, 0.5, 0.5];
982 let mut out = vec![0.0f32; 4];
983 fused_mul_add(&a, &b, &c, &mut out).unwrap();
984 let expected: Vec<f32> = (0..4).map(|i| a[i].mul_add(b[i], c[i])).collect();
985 assert_eq!(out, expected);
986 }
987
988 #[test]
989 fn test_fused_scale_bias_relu_basic() {
990 let input = vec![-2.0, -1.0, 0.0, 1.0, 2.0];
991 let mut out = vec![0.0f32; 5];
992 fused_scale_bias_relu(&input, 2.0, 1.0, &mut out).unwrap();
993 assert_eq!(out, vec![0.0, 0.0, 1.0, 3.0, 5.0]);
995 }
996}