1use core::mem::MaybeUninit;
2
3use itertools::izip;
4use p3_field::{Field, PackedField, PackedValue};
5
6pub trait Butterfly<F: Field>: Copy + Send + Sync {
25 fn apply<PF: PackedField<Scalar = F>>(&self, x_1: PF, x_2: PF) -> (PF, PF);
35
36 #[inline]
40 fn apply_in_place<PF: PackedField<Scalar = F>>(&self, x_1: &mut PF, x_2: &mut PF) {
41 (*x_1, *x_2) = self.apply(*x_1, *x_2);
42 }
43
44 #[inline]
52 fn apply_to_rows(&self, row_1: &mut [F], row_2: &mut [F]) {
53 let (shorts_1, suffix_1) = F::Packing::pack_slice_with_suffix_mut(row_1);
54 let (shorts_2, suffix_2) = F::Packing::pack_slice_with_suffix_mut(row_2);
55 debug_assert_eq!(shorts_1.len(), shorts_2.len());
56 debug_assert_eq!(suffix_1.len(), suffix_2.len());
57 for (x_1, x_2) in shorts_1.iter_mut().zip(shorts_2) {
58 self.apply_in_place(x_1, x_2);
59 }
60 for (x_1, x_2) in suffix_1.iter_mut().zip(suffix_2) {
61 self.apply_in_place(x_1, x_2);
62 }
63 }
64
65 #[inline]
76 fn apply_to_rows_oop(
77 &self,
78 src_1: &[F],
79 dst_1: &mut [MaybeUninit<F>],
80 src_2: &[F],
81 dst_2: &mut [MaybeUninit<F>],
82 ) {
83 let (src_shorts_1, src_suffix_1) = F::Packing::pack_slice_with_suffix(src_1);
84 let (src_shorts_2, src_suffix_2) = F::Packing::pack_slice_with_suffix(src_2);
85 let (dst_shorts_1, dst_suffix_1) =
86 F::Packing::pack_maybe_uninit_slice_with_suffix_mut(dst_1);
87 let (dst_shorts_2, dst_suffix_2) =
88 F::Packing::pack_maybe_uninit_slice_with_suffix_mut(dst_2);
89 debug_assert_eq!(src_shorts_1.len(), src_shorts_2.len());
90 debug_assert_eq!(src_suffix_1.len(), src_suffix_2.len());
91 debug_assert_eq!(dst_shorts_1.len(), dst_shorts_2.len());
92 debug_assert_eq!(dst_suffix_1.len(), dst_suffix_2.len());
93 for (s_1, s_2, d_1, d_2) in izip!(src_shorts_1, src_shorts_2, dst_shorts_1, dst_shorts_2) {
94 let (res_1, res_2) = self.apply(*s_1, *s_2);
95 d_1.write(res_1);
96 d_2.write(res_2);
97 }
98 for (s_1, s_2, d_1, d_2) in izip!(src_suffix_1, src_suffix_2, dst_suffix_1, dst_suffix_2) {
99 let (res_1, res_2) = self.apply(*s_1, *s_2);
100 d_1.write(res_1);
101 d_2.write(res_2);
102 }
103 }
104}
105
106#[derive(Copy, Clone)]
117#[repr(transparent)] pub struct DifButterfly<F>(pub F);
119
120impl<F: Field> Butterfly<F> for DifButterfly<F> {
121 #[inline]
122 fn apply<PF: PackedField<Scalar = F>>(&self, x_1: PF, x_2: PF) -> (PF, PF) {
123 (x_1 + x_2, (x_1 - x_2) * self.0)
124 }
125
126 #[inline]
131 fn apply_to_rows(&self, row_1: &mut [F], row_2: &mut [F]) {
132 let (shorts_1, suffix_1) = F::Packing::pack_slice_with_suffix_mut(row_1);
133 let (shorts_2, suffix_2) = F::Packing::pack_slice_with_suffix_mut(row_2);
134 debug_assert_eq!(shorts_1.len(), shorts_2.len());
135 debug_assert_eq!(suffix_1.len(), suffix_2.len());
136 let twiddle_packed = F::Packing::from(self.0);
137 let (c1, rem1) = shorts_1.as_chunks_mut::<4>();
138 let (c2, rem2) = shorts_2.as_chunks_mut::<4>();
139 for (p1, p2) in c1.iter_mut().zip(c2.iter_mut()) {
140 let a1 = p1[0];
141 let b1 = p1[1];
142 let c1_ = p1[2];
143 let d1 = p1[3];
144 let a2 = p2[0];
145 let b2 = p2[1];
146 let c2_ = p2[2];
147 let d2 = p2[3];
148 p1[0] = a1 + a2;
149 p1[1] = b1 + b2;
150 p1[2] = c1_ + c2_;
151 p1[3] = d1 + d2;
152 p2[0] = (a1 - a2) * twiddle_packed;
153 p2[1] = (b1 - b2) * twiddle_packed;
154 p2[2] = (c1_ - c2_) * twiddle_packed;
155 p2[3] = (d1 - d2) * twiddle_packed;
156 }
157 for (x_1, x_2) in rem1.iter_mut().zip(rem2.iter_mut()) {
158 let sum = *x_1 + *x_2;
159 *x_2 = (*x_1 - *x_2) * twiddle_packed;
160 *x_1 = sum;
161 }
162 for (x_1, x_2) in suffix_1.iter_mut().zip(suffix_2.iter_mut()) {
163 self.apply_in_place(x_1, x_2);
164 }
165 }
166}
167
168#[derive(Copy, Clone)]
179#[repr(transparent)] pub struct DifButterflyZeros<F>(pub F);
181
182impl<F: Field> Butterfly<F> for DifButterflyZeros<F> {
183 #[inline]
184 fn apply<PF: PackedField<Scalar = F>>(&self, x_1: PF, x_2: PF) -> (PF, PF) {
185 debug_assert!(x_2.as_slice().iter().all(|x| x.is_zero())); (x_1, x_1 * self.0)
187 }
188
189 #[inline]
190 fn apply_to_rows(&self, row_1: &mut [F], row_2: &mut [F]) {
191 let (shorts_1, suffix_1) = F::Packing::pack_slice_with_suffix(row_1);
192 let (shorts_2, suffix_2) = F::Packing::pack_slice_with_suffix_mut(row_2);
193 debug_assert_eq!(shorts_1.len(), shorts_2.len());
194 debug_assert_eq!(suffix_1.len(), suffix_2.len());
195 for (x_1, x_2) in shorts_1.iter().zip(shorts_2) {
196 debug_assert!(x_2.as_slice().iter().all(|x| x.is_zero())); *x_2 = *x_1 * self.0; }
199 for (x_1, x_2) in suffix_1.iter().zip(suffix_2) {
200 debug_assert!(x_2.is_zero());
201 *x_2 = *x_1 * self.0; }
203 }
204}
205
206#[derive(Copy, Clone)]
217#[repr(transparent)] pub struct DitButterfly<F>(pub F);
219
220impl<F: Field> Butterfly<F> for DitButterfly<F> {
221 #[inline]
222 fn apply<PF: PackedField<Scalar = F>>(&self, x_1: PF, x_2: PF) -> (PF, PF) {
223 let x_2_twiddle = x_2 * self.0;
224 (x_1 + x_2_twiddle, x_1 - x_2_twiddle)
225 }
226
227 #[inline]
234 fn apply_to_rows(&self, row_1: &mut [F], row_2: &mut [F]) {
235 let (shorts_1, suffix_1) = F::Packing::pack_slice_with_suffix_mut(row_1);
236 let (shorts_2, suffix_2) = F::Packing::pack_slice_with_suffix_mut(row_2);
237 debug_assert_eq!(shorts_1.len(), shorts_2.len());
238 debug_assert_eq!(suffix_1.len(), suffix_2.len());
239 let twiddle_packed = F::Packing::from(self.0);
240 let (c1, rem1) = shorts_1.as_chunks_mut::<4>();
241 let (c2, rem2) = shorts_2.as_chunks_mut::<4>();
242 for (p1, p2) in c1.iter_mut().zip(c2.iter_mut()) {
243 let a1 = p1[0];
244 let b1 = p1[1];
245 let c1_ = p1[2];
246 let d1 = p1[3];
247 let a2 = p2[0];
248 let b2 = p2[1];
249 let c2_ = p2[2];
250 let d2 = p2[3];
251 let a2t = a2 * twiddle_packed;
252 let b2t = b2 * twiddle_packed;
253 let c2t = c2_ * twiddle_packed;
254 let d2t = d2 * twiddle_packed;
255 p1[0] = a1 + a2t;
256 p2[0] = a1 - a2t;
257 p1[1] = b1 + b2t;
258 p2[1] = b1 - b2t;
259 p1[2] = c1_ + c2t;
260 p2[2] = c1_ - c2t;
261 p1[3] = d1 + d2t;
262 p2[3] = d1 - d2t;
263 }
264 for (x_1, x_2) in rem1.iter_mut().zip(rem2.iter_mut()) {
265 let x_2_twiddle = *x_2 * twiddle_packed;
266 let new_x1 = *x_1 + x_2_twiddle;
267 *x_2 = *x_1 - x_2_twiddle;
268 *x_1 = new_x1;
269 }
270 for (x_1, x_2) in suffix_1.iter_mut().zip(suffix_2.iter_mut()) {
271 self.apply_in_place(x_1, x_2);
272 }
273 }
274
275 #[inline]
277 fn apply_to_rows_oop(
278 &self,
279 src_1: &[F],
280 dst_1: &mut [MaybeUninit<F>],
281 src_2: &[F],
282 dst_2: &mut [MaybeUninit<F>],
283 ) {
284 let (src_shorts_1, src_suffix_1) = F::Packing::pack_slice_with_suffix(src_1);
285 let (src_shorts_2, src_suffix_2) = F::Packing::pack_slice_with_suffix(src_2);
286 let (dst_shorts_1, dst_suffix_1) =
287 F::Packing::pack_maybe_uninit_slice_with_suffix_mut(dst_1);
288 let (dst_shorts_2, dst_suffix_2) =
289 F::Packing::pack_maybe_uninit_slice_with_suffix_mut(dst_2);
290 debug_assert_eq!(src_shorts_1.len(), src_shorts_2.len());
291 debug_assert_eq!(src_suffix_1.len(), src_suffix_2.len());
292 debug_assert_eq!(dst_shorts_1.len(), dst_shorts_2.len());
293 debug_assert_eq!(dst_suffix_1.len(), dst_suffix_2.len());
294 let twiddle_packed = F::Packing::from(self.0);
295 let n = src_shorts_1.len();
296 let n4 = n - (n & 3);
297 let mut i = 0;
298 while i < n4 {
299 let a1 = src_shorts_1[i];
300 let b1 = src_shorts_1[i + 1];
301 let c1 = src_shorts_1[i + 2];
302 let d1 = src_shorts_1[i + 3];
303 let a2 = src_shorts_2[i];
304 let b2 = src_shorts_2[i + 1];
305 let c2 = src_shorts_2[i + 2];
306 let d2 = src_shorts_2[i + 3];
307 let a2t = a2 * twiddle_packed;
308 let b2t = b2 * twiddle_packed;
309 let c2t = c2 * twiddle_packed;
310 let d2t = d2 * twiddle_packed;
311 dst_shorts_1[i].write(a1 + a2t);
312 dst_shorts_2[i].write(a1 - a2t);
313 dst_shorts_1[i + 1].write(b1 + b2t);
314 dst_shorts_2[i + 1].write(b1 - b2t);
315 dst_shorts_1[i + 2].write(c1 + c2t);
316 dst_shorts_2[i + 2].write(c1 - c2t);
317 dst_shorts_1[i + 3].write(d1 + d2t);
318 dst_shorts_2[i + 3].write(d1 - d2t);
319 i += 4;
320 }
321 while i < n {
322 let s1 = src_shorts_1[i];
323 let s2 = src_shorts_2[i];
324 let x_2_twiddle = s2 * twiddle_packed;
325 dst_shorts_1[i].write(s1 + x_2_twiddle);
326 dst_shorts_2[i].write(s1 - x_2_twiddle);
327 i += 1;
328 }
329 for (s_1, s_2, d_1, d_2) in izip!(src_suffix_1, src_suffix_2, dst_suffix_1, dst_suffix_2) {
330 let (res_1, res_2) = self.apply(*s_1, *s_2);
331 d_1.write(res_1);
332 d_2.write(res_2);
333 }
334 }
335}
336
337#[derive(Copy, Clone)]
356pub struct ScaledDitButterfly<F> {
357 pub twiddle: F,
358 pub scale: F,
359 pub twiddle_times_scale: F,
361}
362
363impl<F: Field> ScaledDitButterfly<F> {
364 #[inline]
366 pub fn new(twiddle: F, scale: F) -> Self {
367 Self {
368 twiddle,
369 scale,
370 twiddle_times_scale: twiddle * scale,
371 }
372 }
373}
374
375impl<F: Field> Butterfly<F> for ScaledDitButterfly<F> {
376 #[inline]
377 fn apply<PF: PackedField<Scalar = F>>(&self, x_1: PF, x_2: PF) -> (PF, PF) {
378 let x_1_scale = x_1 * self.scale;
384 let x_2_twiddle_scale = x_2 * self.twiddle_times_scale;
385 (x_1_scale + x_2_twiddle_scale, x_1_scale - x_2_twiddle_scale)
386 }
387
388 #[inline]
391 fn apply_to_rows(&self, row_1: &mut [F], row_2: &mut [F]) {
392 let (shorts_1, suffix_1) = F::Packing::pack_slice_with_suffix_mut(row_1);
393 let (shorts_2, suffix_2) = F::Packing::pack_slice_with_suffix_mut(row_2);
394 debug_assert_eq!(shorts_1.len(), shorts_2.len());
395 debug_assert_eq!(suffix_1.len(), suffix_2.len());
396 let scale_packed = F::Packing::from(self.scale);
397 let twiddle_times_scale_packed = F::Packing::from(self.twiddle_times_scale);
398 let (c1, rem1) = shorts_1.as_chunks_mut::<4>();
401 let (c2, rem2) = shorts_2.as_chunks_mut::<4>();
402 for (p1, p2) in c1.iter_mut().zip(c2.iter_mut()) {
403 let a1 = p1[0];
404 let b1 = p1[1];
405 let c1_ = p1[2];
406 let d1 = p1[3];
407 let a2 = p2[0];
408 let b2 = p2[1];
409 let c2_ = p2[2];
410 let d2 = p2[3];
411 let a1s = a1 * scale_packed;
412 let b1s = b1 * scale_packed;
413 let c1s = c1_ * scale_packed;
414 let d1s = d1 * scale_packed;
415 let a2t = a2 * twiddle_times_scale_packed;
416 let b2t = b2 * twiddle_times_scale_packed;
417 let c2t = c2_ * twiddle_times_scale_packed;
418 let d2t = d2 * twiddle_times_scale_packed;
419 p1[0] = a1s + a2t;
420 p2[0] = a1s - a2t;
421 p1[1] = b1s + b2t;
422 p2[1] = b1s - b2t;
423 p1[2] = c1s + c2t;
424 p2[2] = c1s - c2t;
425 p1[3] = d1s + d2t;
426 p2[3] = d1s - d2t;
427 }
428 for (x_1, x_2) in rem1.iter_mut().zip(rem2.iter_mut()) {
429 let x_1_scale = *x_1 * scale_packed;
430 let x_2_twiddle_scale = *x_2 * twiddle_times_scale_packed;
431 *x_1 = x_1_scale + x_2_twiddle_scale;
432 *x_2 = x_1_scale - x_2_twiddle_scale;
433 }
434 for (x_1, x_2) in suffix_1.iter_mut().zip(suffix_2.iter_mut()) {
435 self.apply_in_place(x_1, x_2);
436 }
437 }
438}
439
440#[derive(Copy, Clone)]
452pub struct TwiddleFreeButterfly;
453
454impl<F: Field> Butterfly<F> for TwiddleFreeButterfly {
455 #[inline]
456 fn apply<PF: PackedField<Scalar = F>>(&self, x_1: PF, x_2: PF) -> (PF, PF) {
457 (x_1 + x_2, x_1 - x_2)
458 }
459}