1use core::marker::PhantomData;
48use core::ops::{Add, AddAssign, Neg, Sub, SubAssign};
49
50use p3_field::{Algebra, PrimeCharacteristicRing};
51
52pub trait ConvolutionElt:
56 Add<Output = Self> + AddAssign + Clone + Neg<Output = Self> + Sub<Output = Self> + SubAssign
57{
58}
59
60impl<T> ConvolutionElt for T where
61 T: Add<Output = T> + AddAssign + Clone + Neg<Output = T> + Sub<Output = T> + SubAssign
62{
63}
64
65pub trait ConvolutionRhs:
69 Add<Output = Self> + Clone + Neg<Output = Self> + Sub<Output = Self>
70{
71}
72
73impl<T> ConvolutionRhs for T where T: Add<Output = T> + Clone + Neg<Output = T> + Sub<Output = T> {}
74
75pub trait Convolve<F, T: ConvolutionElt, U: ConvolutionRhs> {
97 const T_ZERO: T;
102
103 const U_ZERO: U;
108
109 fn halve(val: T) -> T;
114
115 fn read(input: F) -> T;
118
119 fn parity_dot<const N: usize>(lhs: [T; N], rhs: [U; N]) -> T;
126
127 fn reduce(z: T) -> F;
130
131 #[inline(always)]
136 fn apply<const N: usize, C: Fn([T; N], [U; N], &mut [T])>(
137 lhs: [F; N],
138 rhs: [U; N],
139 conv: C,
140 ) -> [F; N] {
141 let lhs = lhs.map(Self::read);
142 let mut output = [Self::T_ZERO; N];
143 conv(lhs, rhs, &mut output);
144 output.map(Self::reduce)
145 }
146
147 #[inline(always)]
148 fn conv3(lhs: [T; 3], rhs: [U; 3], output: &mut [T]) {
149 output[0] = Self::parity_dot(
150 lhs.clone(),
151 [rhs[0].clone(), rhs[2].clone(), rhs[1].clone()],
152 );
153 output[1] = Self::parity_dot(
154 lhs.clone(),
155 [rhs[1].clone(), rhs[0].clone(), rhs[2].clone()],
156 );
157 output[2] = Self::parity_dot(lhs, [rhs[2].clone(), rhs[1].clone(), rhs[0].clone()]);
158 }
159
160 #[inline(always)]
161 fn negacyclic_conv3(lhs: [T; 3], rhs: [U; 3], output: &mut [T]) {
162 output[0] = Self::parity_dot(
163 lhs.clone(),
164 [rhs[0].clone(), -rhs[2].clone(), -rhs[1].clone()],
165 );
166 output[1] = Self::parity_dot(
167 lhs.clone(),
168 [rhs[1].clone(), rhs[0].clone(), -rhs[2].clone()],
169 );
170 output[2] = Self::parity_dot(lhs, [rhs[2].clone(), rhs[1].clone(), rhs[0].clone()]);
171 }
172
173 #[inline(always)]
174 fn conv4(lhs: [T; 4], rhs: [U; 4], output: &mut [T]) {
175 let u_p = [
178 lhs[0].clone() + lhs[2].clone(),
179 lhs[1].clone() + lhs[3].clone(),
180 ];
181 let u_m = [
182 lhs[0].clone() - lhs[2].clone(),
183 lhs[1].clone() - lhs[3].clone(),
184 ];
185 let v_p = [
186 rhs[0].clone() + rhs[2].clone(),
187 rhs[1].clone() + rhs[3].clone(),
188 ];
189 let v_m = [
190 rhs[0].clone() - rhs[2].clone(),
191 rhs[1].clone() - rhs[3].clone(),
192 ];
193
194 output[0] = Self::parity_dot(u_m.clone(), [v_m[0].clone(), -v_m[1].clone()]);
195 output[1] = Self::parity_dot(u_m, [v_m[1].clone(), v_m[0].clone()]);
196 output[2] = Self::parity_dot(u_p.clone(), v_p.clone());
197 output[3] = Self::parity_dot(u_p, [v_p[1].clone(), v_p[0].clone()]);
198
199 output[0] += output[2].clone();
200 output[1] += output[3].clone();
201
202 output[0] = Self::halve(output[0].clone());
203 output[1] = Self::halve(output[1].clone());
204
205 output[2] -= output[0].clone();
206 output[3] -= output[1].clone();
207 }
208
209 #[inline(always)]
210 fn negacyclic_conv4(lhs: [T; 4], rhs: [U; 4], output: &mut [T]) {
211 output[0] = Self::parity_dot(
212 lhs.clone(),
213 [
214 rhs[0].clone(),
215 -rhs[3].clone(),
216 -rhs[2].clone(),
217 -rhs[1].clone(),
218 ],
219 );
220 output[1] = Self::parity_dot(
221 lhs.clone(),
222 [
223 rhs[1].clone(),
224 rhs[0].clone(),
225 -rhs[3].clone(),
226 -rhs[2].clone(),
227 ],
228 );
229 output[2] = Self::parity_dot(
230 lhs.clone(),
231 [
232 rhs[2].clone(),
233 rhs[1].clone(),
234 rhs[0].clone(),
235 -rhs[3].clone(),
236 ],
237 );
238 output[3] = Self::parity_dot(
239 lhs,
240 [
241 rhs[3].clone(),
242 rhs[2].clone(),
243 rhs[1].clone(),
244 rhs[0].clone(),
245 ],
246 );
247 }
248
249 #[inline(always)]
252 fn conv_n_recursive<const N: usize, const HALF_N: usize, C, NC>(
253 lhs: [T; N],
254 rhs: [U; N],
255 output: &mut [T],
256 inner_conv: C,
257 inner_negacyclic_conv: NC,
258 ) where
259 C: Fn([T; HALF_N], [U; HALF_N], &mut [T]),
260 NC: Fn([T; HALF_N], [U; HALF_N], &mut [T]),
261 {
262 debug_assert_eq!(2 * HALF_N, N);
263 let mut lhs_pos = [Self::T_ZERO; HALF_N]; let mut lhs_neg = [Self::T_ZERO; HALF_N]; let mut rhs_pos = [Self::U_ZERO; HALF_N]; let mut rhs_neg = [Self::U_ZERO; HALF_N]; for i in 0..HALF_N {
269 let s = lhs[i].clone();
270 let t = lhs[i + HALF_N].clone();
271
272 lhs_pos[i] = s.clone() + t.clone();
273 lhs_neg[i] = s - t;
274
275 let s = rhs[i].clone();
276 let t = rhs[i + HALF_N].clone();
277
278 rhs_pos[i] = s.clone() + t.clone();
279 rhs_neg[i] = s - t;
280 }
281
282 let (left, right) = output.split_at_mut(HALF_N);
283
284 inner_negacyclic_conv(lhs_neg, rhs_neg, left);
286
287 inner_conv(lhs_pos, rhs_pos, right);
289
290 for i in 0..HALF_N {
291 left[i] += right[i].clone(); left[i] = Self::halve(left[i].clone()); right[i] -= left[i].clone(); }
295 }
296
297 #[inline(always)]
300 fn negacyclic_conv_n_recursive<const N: usize, const HALF_N: usize, NC>(
301 lhs: [T; N],
302 rhs: [U; N],
303 output: &mut [T],
304 inner_negacyclic_conv: NC,
305 ) where
306 NC: Fn([T; HALF_N], [U; HALF_N], &mut [T]),
307 {
308 debug_assert_eq!(2 * HALF_N, N);
309 let mut lhs_even = [Self::T_ZERO; HALF_N];
310 let mut lhs_odd = [Self::T_ZERO; HALF_N];
311 let mut lhs_sum = [Self::T_ZERO; HALF_N];
312 let mut rhs_even = [Self::U_ZERO; HALF_N];
313 let mut rhs_odd = [Self::U_ZERO; HALF_N];
314 let mut rhs_sum = [Self::U_ZERO; HALF_N];
315
316 for i in 0..HALF_N {
317 let s = lhs[2 * i].clone();
318 let t = lhs[2 * i + 1].clone();
319 lhs_sum[i] = s.clone() + t.clone();
320 lhs_even[i] = s;
321 lhs_odd[i] = t;
322
323 let s = rhs[2 * i].clone();
324 let t = rhs[2 * i + 1].clone();
325 rhs_sum[i] = s.clone() + t.clone();
326 rhs_even[i] = s;
327 rhs_odd[i] = t;
328 }
329
330 let mut even_s_conv = [Self::T_ZERO; HALF_N];
331 let (left, right) = output.split_at_mut(HALF_N);
332
333 inner_negacyclic_conv(lhs_even, rhs_even, &mut even_s_conv);
336 inner_negacyclic_conv(lhs_odd, rhs_odd, left);
337 inner_negacyclic_conv(lhs_sum, rhs_sum, right);
338
339 right[0] -= even_s_conv[0].clone() + left[0].clone();
342 even_s_conv[0] -= left[HALF_N - 1].clone();
343
344 for i in 1..HALF_N {
345 right[i] -= even_s_conv[i].clone() + left[i].clone();
346 even_s_conv[i] += left[i - 1].clone();
347 }
348
349 for i in 0..HALF_N {
351 output[2 * i] = even_s_conv[i].clone();
352 output[2 * i + 1] = output[i + HALF_N].clone();
353 }
354 }
355
356 #[inline(always)]
357 fn conv6(lhs: [T; 6], rhs: [U; 6], output: &mut [T]) {
358 Self::conv_n_recursive(lhs, rhs, output, Self::conv3, Self::negacyclic_conv3);
359 }
360
361 #[inline(always)]
362 fn negacyclic_conv6(lhs: [T; 6], rhs: [U; 6], output: &mut [T]) {
363 Self::negacyclic_conv_n_recursive(lhs, rhs, output, Self::negacyclic_conv3);
364 }
365
366 #[inline(always)]
367 fn conv8(lhs: [T; 8], rhs: [U; 8], output: &mut [T]) {
368 Self::conv_n_recursive(lhs, rhs, output, Self::conv4, Self::negacyclic_conv4);
369 }
370
371 #[inline(always)]
372 fn negacyclic_conv8(lhs: [T; 8], rhs: [U; 8], output: &mut [T]) {
373 Self::negacyclic_conv_n_recursive(lhs, rhs, output, Self::negacyclic_conv4);
374 }
375
376 #[inline(always)]
377 fn conv12(lhs: [T; 12], rhs: [U; 12], output: &mut [T]) {
378 Self::conv_n_recursive(lhs, rhs, output, Self::conv6, Self::negacyclic_conv6);
379 }
380
381 #[inline(always)]
382 fn negacyclic_conv12(lhs: [T; 12], rhs: [U; 12], output: &mut [T]) {
383 Self::negacyclic_conv_n_recursive(lhs, rhs, output, Self::negacyclic_conv6);
384 }
385
386 #[inline(always)]
387 fn conv16(lhs: [T; 16], rhs: [U; 16], output: &mut [T]) {
388 Self::conv_n_recursive(lhs, rhs, output, Self::conv8, Self::negacyclic_conv8);
389 }
390
391 #[inline(always)]
392 fn negacyclic_conv16(lhs: [T; 16], rhs: [U; 16], output: &mut [T]) {
393 Self::negacyclic_conv_n_recursive(lhs, rhs, output, Self::negacyclic_conv8);
394 }
395
396 #[inline(always)]
397 fn conv24(lhs: [T; 24], rhs: [U; 24], output: &mut [T]) {
398 Self::conv_n_recursive(lhs, rhs, output, Self::conv12, Self::negacyclic_conv12);
399 }
400
401 #[inline(always)]
402 fn conv32(lhs: [T; 32], rhs: [U; 32], output: &mut [T]) {
403 Self::conv_n_recursive(lhs, rhs, output, Self::conv16, Self::negacyclic_conv16);
404 }
405
406 #[inline(always)]
407 fn negacyclic_conv32(lhs: [T; 32], rhs: [U; 32], output: &mut [T]) {
408 Self::negacyclic_conv_n_recursive(lhs, rhs, output, Self::negacyclic_conv16);
409 }
410
411 #[inline(always)]
412 fn conv64(lhs: [T; 64], rhs: [U; 64], output: &mut [T]) {
413 Self::conv_n_recursive(lhs, rhs, output, Self::conv32, Self::negacyclic_conv32);
414 }
415}
416
417struct FieldConvolve<F, A>(PhantomData<(F, A)>);
422
423impl<F: PrimeCharacteristicRing, A: Algebra<F> + Clone> Convolve<A, A, F> for FieldConvolve<F, A> {
424 const T_ZERO: A = A::ZERO;
425 const U_ZERO: F = F::ZERO;
426
427 #[inline(always)]
428 fn halve(val: A) -> A {
429 val.halve()
430 }
431
432 #[inline(always)]
433 fn read(input: A) -> A {
434 input
435 }
436
437 #[inline(always)]
438 fn parity_dot<const N: usize>(lhs: [A; N], rhs: [F; N]) -> A {
439 A::mixed_dot_product(&lhs, &rhs)
440 }
441
442 #[inline(always)]
443 fn reduce(z: A) -> A {
444 z
445 }
446}
447
448#[inline]
450pub fn mds_circulant_karatsuba_8<F: PrimeCharacteristicRing, A: Algebra<F> + Clone>(
451 state: &mut [A; 8],
452 col: &[F; 8],
453) {
454 let input = state.clone();
455 FieldConvolve::<F, A>::conv8(input, col.clone(), state.as_mut_slice());
456}
457
458#[inline]
460pub fn mds_circulant_karatsuba_12<F: PrimeCharacteristicRing, A: Algebra<F> + Clone>(
461 state: &mut [A; 12],
462 col: &[F; 12],
463) {
464 let input = state.clone();
465 FieldConvolve::<F, A>::conv12(input, col.clone(), state.as_mut_slice());
466}
467
468#[inline]
470pub fn mds_circulant_karatsuba_16<F: PrimeCharacteristicRing, A: Algebra<F> + Clone>(
471 state: &mut [A; 16],
472 col: &[F; 16],
473) {
474 let input = state.clone();
475 FieldConvolve::<F, A>::conv16(input, col.clone(), state.as_mut_slice());
476}
477
478#[inline]
480pub fn mds_circulant_karatsuba_24<F: PrimeCharacteristicRing, A: Algebra<F> + Clone>(
481 state: &mut [A; 24],
482 col: &[F; 24],
483) {
484 let input = state.clone();
485 FieldConvolve::<F, A>::conv24(input, col.clone(), state.as_mut_slice());
486}
487
488#[cfg(test)]
489mod tests {
490 use p3_baby_bear::BabyBear;
491 use p3_field::PrimeCharacteristicRing;
492 use proptest::prelude::*;
493
494 use super::*;
495
496 type F = BabyBear;
497
498 fn arb_f() -> impl Strategy<Value = F> {
499 prop::num::u32::ANY.prop_map(F::from_u32)
500 }
501
502 fn naive_cyclic_conv<const N: usize>(lhs: [F; N], rhs: [F; N]) -> [F; N] {
503 core::array::from_fn(|i| {
505 let mut acc = F::ZERO;
506 for j in 0..N {
507 acc += lhs[j] * rhs[(N + i - j) % N];
508 }
509 acc
510 })
511 }
512
513 fn naive_negacyclic_conv<const N: usize>(lhs: [F; N], rhs: [F; N]) -> [F; N] {
514 let mut out = [F::ZERO; N];
517 for (i, &l) in lhs.iter().enumerate() {
518 for (j, &r) in rhs.iter().enumerate() {
519 let k = i + j;
520 if k < N {
521 out[k] += l * r;
522 } else {
523 out[k - N] -= l * r;
524 }
525 }
526 }
527 out
528 }
529
530 fn check_conv<const N: usize>(
531 lhs: [F; N],
532 rhs: [F; N],
533 conv_fn: fn([F; N], [F; N], &mut [F]),
534 naive_fn: fn([F; N], [F; N]) -> [F; N],
535 ) {
536 let expected = naive_fn(lhs, rhs);
537 let mut output = [F::ZERO; N];
538 conv_fn(lhs, rhs, &mut output);
539 assert_eq!(output, expected, "convolution mismatch");
540 }
541
542 macro_rules! conv_test {
543 ($name:ident, $n:expr, $conv:expr, $naive:expr, $arr:ident) => {
544 proptest! {
545 #[test]
546 fn $name(
547 lhs in prop::array::$arr(arb_f()),
548 rhs in prop::array::$arr(arb_f()),
549 ) {
550 check_conv::<$n>(lhs, rhs, $conv, $naive);
551 }
552 }
553 };
554 }
555
556 conv_test!(
558 conv3_matches_naive,
559 3,
560 FieldConvolve::<F, F>::conv3,
561 naive_cyclic_conv,
562 uniform3
563 );
564 conv_test!(
565 negacyclic_conv3_matches_naive,
566 3,
567 FieldConvolve::<F, F>::negacyclic_conv3,
568 naive_negacyclic_conv,
569 uniform3
570 );
571
572 conv_test!(
574 conv4_matches_naive,
575 4,
576 FieldConvolve::<F, F>::conv4,
577 naive_cyclic_conv,
578 uniform4
579 );
580 conv_test!(
581 negacyclic_conv4_matches_naive,
582 4,
583 FieldConvolve::<F, F>::negacyclic_conv4,
584 naive_negacyclic_conv,
585 uniform4
586 );
587
588 conv_test!(
590 conv6_matches_naive,
591 6,
592 FieldConvolve::<F, F>::conv6,
593 naive_cyclic_conv,
594 uniform6
595 );
596 conv_test!(
597 negacyclic_conv6_matches_naive,
598 6,
599 FieldConvolve::<F, F>::negacyclic_conv6,
600 naive_negacyclic_conv,
601 uniform6
602 );
603
604 conv_test!(
606 conv8_matches_naive,
607 8,
608 FieldConvolve::<F, F>::conv8,
609 naive_cyclic_conv,
610 uniform8
611 );
612 conv_test!(
613 negacyclic_conv8_matches_naive,
614 8,
615 FieldConvolve::<F, F>::negacyclic_conv8,
616 naive_negacyclic_conv,
617 uniform8
618 );
619
620 conv_test!(
622 conv12_matches_naive,
623 12,
624 FieldConvolve::<F, F>::conv12,
625 naive_cyclic_conv,
626 uniform12
627 );
628 conv_test!(
629 negacyclic_conv12_matches_naive,
630 12,
631 FieldConvolve::<F, F>::negacyclic_conv12,
632 naive_negacyclic_conv,
633 uniform12
634 );
635
636 conv_test!(
638 conv16_matches_naive,
639 16,
640 FieldConvolve::<F, F>::conv16,
641 naive_cyclic_conv,
642 uniform16
643 );
644 conv_test!(
645 negacyclic_conv16_matches_naive,
646 16,
647 FieldConvolve::<F, F>::negacyclic_conv16,
648 naive_negacyclic_conv,
649 uniform16
650 );
651
652 conv_test!(
654 conv24_matches_naive,
655 24,
656 FieldConvolve::<F, F>::conv24,
657 naive_cyclic_conv,
658 uniform24
659 );
660
661 conv_test!(
663 conv32_matches_naive,
664 32,
665 FieldConvolve::<F, F>::conv32,
666 naive_cyclic_conv,
667 uniform32
668 );
669 conv_test!(
670 negacyclic_conv32_matches_naive,
671 32,
672 FieldConvolve::<F, F>::negacyclic_conv32,
673 naive_negacyclic_conv,
674 uniform32
675 );
676
677 #[test]
678 fn conv64_matches_naive_fixed() {
679 let lhs: [F; 64] = core::array::from_fn(|i| F::from_u32(i as u32 + 1));
680 let rhs: [F; 64] = core::array::from_fn(|i| F::from_u32(64 - i as u32));
681 check_conv::<64>(lhs, rhs, FieldConvolve::<F, F>::conv64, naive_cyclic_conv);
682 }
683
684 #[test]
685 fn conv64_all_ones() {
686 let ones = [F::ONE; 64];
687 let expected = naive_cyclic_conv(ones, ones);
688 let mut output = [F::ZERO; 64];
689 FieldConvolve::<F, F>::conv64(ones, ones, &mut output);
690 assert_eq!(output, expected);
691 }
692
693 proptest! {
694 #[test]
695 fn karatsuba_16_matches_naive(
696 col in prop::array::uniform16(arb_f()),
697 state in prop::array::uniform16(arb_f()),
698 ) {
699 let expected = naive_cyclic_conv(state, col);
700 let mut actual = state;
701 mds_circulant_karatsuba_16(&mut actual, &col);
702 prop_assert_eq!(actual, expected);
703 }
704
705 #[test]
706 fn karatsuba_24_matches_naive(
707 col in prop::array::uniform24(arb_f()),
708 state in prop::array::uniform24(arb_f()),
709 ) {
710 let expected = naive_cyclic_conv(state, col);
711 let mut actual = state;
712 mds_circulant_karatsuba_24(&mut actual, &col);
713 prop_assert_eq!(actual, expected);
714 }
715 }
716
717 proptest! {
718 #[test]
719 fn conv8_commutative(
720 a in prop::array::uniform8(arb_f()),
721 b in prop::array::uniform8(arb_f()),
722 ) {
723 let mut ab = [F::ZERO; 8];
725 let mut ba = [F::ZERO; 8];
726 FieldConvolve::<F, F>::conv8(a, b, &mut ab);
727 FieldConvolve::<F, F>::conv8(b, a, &mut ba);
728 prop_assert_eq!(ab, ba);
729 }
730
731 #[test]
732 fn conv8_identity(a in prop::array::uniform8(arb_f())) {
733 let mut id = [F::ZERO; 8];
735 id[0] = F::ONE;
736 let mut out = [F::ZERO; 8];
737 FieldConvolve::<F, F>::conv8(a, id, &mut out);
738 prop_assert_eq!(out, a);
739 }
740
741 #[test]
742 fn conv8_zero(a in prop::array::uniform8(arb_f())) {
743 let zeros = [F::ZERO; 8];
745 let mut out = [F::ZERO; 8];
746 FieldConvolve::<F, F>::conv8(a, zeros, &mut out);
747 prop_assert_eq!(out, zeros);
748 }
749 }
750}