Skip to main content

openvm_benchmarks_fields/
lib.rs

1//! Extension Field Benchmark
2//!
3//! Benchmarks for field arithmetic operations on GPU.
4//! Measures compute throughput by doing many ops per element to amortize memory access.
5//!
6//! Throughput is reported in Gops/s (Giga-operations per second = billion ops/sec).
7
8use std::{ffi::c_void, sync::OnceLock, time::Instant};
9
10use openvm_cuda_common::{
11    copy::MemCopyH2D,
12    d_buffer::DeviceBuffer,
13    stream::{cudaStream_t, GpuDeviceCtx},
14};
15
16pub fn bench_ctx() -> &'static GpuDeviceCtx {
17    static CTX: OnceLock<GpuDeviceCtx> = OnceLock::new();
18    CTX.get_or_init(|| GpuDeviceCtx::for_current_device().expect("benchmark CUDA context"))
19}
20
21// ============================================================================
22// FFI Bindings
23// ============================================================================
24
25// Benchmark kernels
26mod ffi {
27    use super::*;
28
29    #[link(name = "ext_field_bench")]
30    extern "C" {
31        pub(crate) fn init_fp(
32            out: *mut c_void,
33            raw_data: *const u32,
34            n: usize,
35            stream: cudaStream_t,
36        ) -> i32;
37        pub(crate) fn add_fp(
38            out: *mut c_void,
39            a: *const c_void,
40            b: *const c_void,
41            n: usize,
42            reps: i32,
43            stream: cudaStream_t,
44        ) -> i32;
45        pub(crate) fn mul_fp(
46            out: *mut c_void,
47            a: *const c_void,
48            b: *const c_void,
49            n: usize,
50            reps: i32,
51            stream: cudaStream_t,
52        ) -> i32;
53        pub(crate) fn inv_fp(
54            out: *mut c_void,
55            a: *const c_void,
56            n: usize,
57            reps: i32,
58            stream: cudaStream_t,
59        ) -> i32;
60
61        // BabyBear quartic extension (simple implementation)
62        pub fn init_fp4(
63            out: *mut c_void,
64            raw_data: *const u32,
65            n: usize,
66            stream: cudaStream_t,
67        ) -> i32;
68        pub(crate) fn add_fp4(
69            out: *mut c_void,
70            a: *const c_void,
71            b: *const c_void,
72            n: usize,
73            reps: i32,
74            stream: cudaStream_t,
75        ) -> i32;
76        pub(crate) fn mul_fp4(
77            out: *mut c_void,
78            a: *const c_void,
79            b: *const c_void,
80            n: usize,
81            reps: i32,
82            stream: cudaStream_t,
83        ) -> i32;
84        pub(crate) fn inv_fp4(
85            out: *mut c_void,
86            a: *const c_void,
87            n: usize,
88            reps: i32,
89            stream: cudaStream_t,
90        ) -> i32;
91
92        // BabyBear quartic extension (optimized bb31_4_t)
93        pub(crate) fn init_fpext(
94            out: *mut c_void,
95            raw_data: *const u32,
96            n: usize,
97            stream: cudaStream_t,
98        ) -> i32;
99        pub(crate) fn add_fpext(
100            out: *mut c_void,
101            a: *const c_void,
102            b: *const c_void,
103            n: usize,
104            reps: i32,
105            stream: cudaStream_t,
106        ) -> i32;
107        pub(crate) fn mul_fpext(
108            out: *mut c_void,
109            a: *const c_void,
110            b: *const c_void,
111            n: usize,
112            reps: i32,
113            stream: cudaStream_t,
114        ) -> i32;
115        pub(crate) fn inv_fpext(
116            out: *mut c_void,
117            a: *const c_void,
118            n: usize,
119            reps: i32,
120            stream: cudaStream_t,
121        ) -> i32;
122
123        pub fn init_fp5(
124            out: *mut c_void,
125            raw_data: *const u32,
126            n: usize,
127            stream: cudaStream_t,
128        ) -> i32;
129        pub(crate) fn add_fp5(
130            out: *mut c_void,
131            a: *const c_void,
132            b: *const c_void,
133            n: usize,
134            reps: i32,
135            stream: cudaStream_t,
136        ) -> i32;
137        pub(crate) fn mul_fp5(
138            out: *mut c_void,
139            a: *const c_void,
140            b: *const c_void,
141            n: usize,
142            reps: i32,
143            stream: cudaStream_t,
144        ) -> i32;
145        pub(crate) fn inv_fp5(
146            out: *mut c_void,
147            a: *const c_void,
148            n: usize,
149            reps: i32,
150            stream: cudaStream_t,
151        ) -> i32;
152
153        pub fn init_fp6(
154            out: *mut c_void,
155            raw_data: *const u32,
156            n: usize,
157            stream: cudaStream_t,
158        ) -> i32;
159        pub(crate) fn add_fp6(
160            out: *mut c_void,
161            a: *const c_void,
162            b: *const c_void,
163            n: usize,
164            reps: i32,
165            stream: cudaStream_t,
166        ) -> i32;
167        pub(crate) fn mul_fp6(
168            out: *mut c_void,
169            a: *const c_void,
170            b: *const c_void,
171            n: usize,
172            reps: i32,
173            stream: cudaStream_t,
174        ) -> i32;
175        pub(crate) fn inv_fp6(
176            out: *mut c_void,
177            a: *const c_void,
178            n: usize,
179            reps: i32,
180            stream: cudaStream_t,
181        ) -> i32;
182
183        pub fn init_fp2x3(
184            out: *mut c_void,
185            raw_data: *const u32,
186            n: usize,
187            stream: cudaStream_t,
188        ) -> i32;
189        pub(crate) fn add_fp2x3(
190            out: *mut c_void,
191            a: *const c_void,
192            b: *const c_void,
193            n: usize,
194            reps: i32,
195            stream: cudaStream_t,
196        ) -> i32;
197        pub(crate) fn mul_fp2x3(
198            out: *mut c_void,
199            a: *const c_void,
200            b: *const c_void,
201            n: usize,
202            reps: i32,
203            stream: cudaStream_t,
204        ) -> i32;
205        pub(crate) fn inv_fp2x3(
206            out: *mut c_void,
207            a: *const c_void,
208            n: usize,
209            reps: i32,
210            stream: cudaStream_t,
211        ) -> i32;
212
213        pub fn init_fp3x2(
214            out: *mut c_void,
215            raw_data: *const u32,
216            n: usize,
217            stream: cudaStream_t,
218        ) -> i32;
219        pub(crate) fn add_fp3x2(
220            out: *mut c_void,
221            a: *const c_void,
222            b: *const c_void,
223            n: usize,
224            reps: i32,
225            stream: cudaStream_t,
226        ) -> i32;
227        pub(crate) fn mul_fp3x2(
228            out: *mut c_void,
229            a: *const c_void,
230            b: *const c_void,
231            n: usize,
232            reps: i32,
233            stream: cudaStream_t,
234        ) -> i32;
235        pub(crate) fn inv_fp3x2(
236            out: *mut c_void,
237            a: *const c_void,
238            n: usize,
239            reps: i32,
240            stream: cudaStream_t,
241        ) -> i32;
242
243        // KoalaBear base field
244        pub fn init_kb(
245            out: *mut c_void,
246            raw_data: *const u32,
247            n: usize,
248            stream: cudaStream_t,
249        ) -> i32;
250        pub(crate) fn add_kb(
251            out: *mut c_void,
252            a: *const c_void,
253            b: *const c_void,
254            n: usize,
255            reps: i32,
256            stream: cudaStream_t,
257        ) -> i32;
258        pub(crate) fn mul_kb(
259            out: *mut c_void,
260            a: *const c_void,
261            b: *const c_void,
262            n: usize,
263            reps: i32,
264            stream: cudaStream_t,
265        ) -> i32;
266        pub(crate) fn inv_kb(
267            out: *mut c_void,
268            a: *const c_void,
269            n: usize,
270            reps: i32,
271            stream: cudaStream_t,
272        ) -> i32;
273
274        // KoalaBear quintic extension (x^5 + x + 4)
275        pub fn init_kb5(
276            out: *mut c_void,
277            raw_data: *const u32,
278            n: usize,
279            stream: cudaStream_t,
280        ) -> i32;
281        pub(crate) fn add_kb5(
282            out: *mut c_void,
283            a: *const c_void,
284            b: *const c_void,
285            n: usize,
286            reps: i32,
287            stream: cudaStream_t,
288        ) -> i32;
289        pub(crate) fn mul_kb5(
290            out: *mut c_void,
291            a: *const c_void,
292            b: *const c_void,
293            n: usize,
294            reps: i32,
295            stream: cudaStream_t,
296        ) -> i32;
297        pub(crate) fn inv_kb5(
298            out: *mut c_void,
299            a: *const c_void,
300            n: usize,
301            reps: i32,
302            stream: cudaStream_t,
303        ) -> i32;
304
305        // KoalaBear sextic extension (x^6 + x^3 + 1)
306        pub fn init_kb6(
307            out: *mut c_void,
308            raw_data: *const u32,
309            n: usize,
310            stream: cudaStream_t,
311        ) -> i32;
312        pub(crate) fn add_kb6(
313            out: *mut c_void,
314            a: *const c_void,
315            b: *const c_void,
316            n: usize,
317            reps: i32,
318            stream: cudaStream_t,
319        ) -> i32;
320        pub(crate) fn mul_kb6(
321            out: *mut c_void,
322            a: *const c_void,
323            b: *const c_void,
324            n: usize,
325            reps: i32,
326            stream: cudaStream_t,
327        ) -> i32;
328        pub(crate) fn inv_kb6(
329            out: *mut c_void,
330            a: *const c_void,
331            n: usize,
332            reps: i32,
333            stream: cudaStream_t,
334        ) -> i32;
335
336        // KoalaBear 2×3 tower (u²=3, v³=1+u)
337        pub fn init_kb2x3(
338            out: *mut c_void,
339            raw_data: *const u32,
340            n: usize,
341            stream: cudaStream_t,
342        ) -> i32;
343        pub(crate) fn add_kb2x3(
344            out: *mut c_void,
345            a: *const c_void,
346            b: *const c_void,
347            n: usize,
348            reps: i32,
349            stream: cudaStream_t,
350        ) -> i32;
351        pub(crate) fn mul_kb2x3(
352            out: *mut c_void,
353            a: *const c_void,
354            b: *const c_void,
355            n: usize,
356            reps: i32,
357            stream: cudaStream_t,
358        ) -> i32;
359        pub(crate) fn inv_kb2x3(
360            out: *mut c_void,
361            a: *const c_void,
362            n: usize,
363            reps: i32,
364            stream: cudaStream_t,
365        ) -> i32;
366
367        // KoalaBear 3×2 tower (w³=-w-4, z²=3)
368        pub fn init_kb3x2(
369            out: *mut c_void,
370            raw_data: *const u32,
371            n: usize,
372            stream: cudaStream_t,
373        ) -> i32;
374        pub(crate) fn add_kb3x2(
375            out: *mut c_void,
376            a: *const c_void,
377            b: *const c_void,
378            n: usize,
379            reps: i32,
380            stream: cudaStream_t,
381        ) -> i32;
382        pub(crate) fn mul_kb3x2(
383            out: *mut c_void,
384            a: *const c_void,
385            b: *const c_void,
386            n: usize,
387            reps: i32,
388            stream: cudaStream_t,
389        ) -> i32;
390        pub(crate) fn inv_kb3x2(
391            out: *mut c_void,
392            a: *const c_void,
393            n: usize,
394            reps: i32,
395            stream: cudaStream_t,
396        ) -> i32;
397
398        // Goldilocks base field (64-bit prime: 2^64 - 2^32 + 1)
399        pub(crate) fn init_gl(
400            out: *mut c_void,
401            raw_data: *const u64,
402            n: usize,
403            stream: cudaStream_t,
404        ) -> i32;
405        pub(crate) fn add_gl(
406            out: *mut c_void,
407            a: *const c_void,
408            b: *const c_void,
409            n: usize,
410            reps: i32,
411            stream: cudaStream_t,
412        ) -> i32;
413        pub(crate) fn mul_gl(
414            out: *mut c_void,
415            a: *const c_void,
416            b: *const c_void,
417            n: usize,
418            reps: i32,
419            stream: cudaStream_t,
420        ) -> i32;
421        pub(crate) fn inv_gl(
422            out: *mut c_void,
423            a: *const c_void,
424            n: usize,
425            reps: i32,
426            stream: cudaStream_t,
427        ) -> i32;
428
429        // Goldilocks cubic extension (X³ - X - 1)
430        pub(crate) fn init_gl3(
431            out: *mut c_void,
432            raw_data: *const u64,
433            n: usize,
434            stream: cudaStream_t,
435        ) -> i32;
436        pub(crate) fn add_gl3(
437            out: *mut c_void,
438            a: *const c_void,
439            b: *const c_void,
440            n: usize,
441            reps: i32,
442            stream: cudaStream_t,
443        ) -> i32;
444        pub(crate) fn mul_gl3(
445            out: *mut c_void,
446            a: *const c_void,
447            b: *const c_void,
448            n: usize,
449            reps: i32,
450            stream: cudaStream_t,
451        ) -> i32;
452        pub(crate) fn inv_gl3(
453            out: *mut c_void,
454            a: *const c_void,
455            n: usize,
456            reps: i32,
457            stream: cudaStream_t,
458        ) -> i32;
459    }
460
461    // Poseidon2 benchmark kernels
462    #[link(name = "ext_field_bench")]
463    extern "C" {
464        pub(crate) fn init_poseidon2_bb(
465            out: *mut c_void,
466            raw_data: *const u32,
467            n: usize,
468            stream: cudaStream_t,
469        ) -> i32;
470        pub(crate) fn run_poseidon2_bb(
471            states: *mut c_void,
472            n: usize,
473            reps: i32,
474            stream: cudaStream_t,
475        ) -> i32;
476        pub(crate) fn init_poseidon2_kb(
477            out: *mut c_void,
478            raw_data: *const u32,
479            n: usize,
480            stream: cudaStream_t,
481        ) -> i32;
482        pub(crate) fn run_poseidon2_kb(
483            states: *mut c_void,
484            n: usize,
485            reps: i32,
486            stream: cudaStream_t,
487        ) -> i32;
488    }
489}
490
491macro_rules! wrap_init_u32 {
492    ($vis:vis $name:ident) => {
493        /// Launches the benchmark init kernel on the benchmark stream.
494        ///
495        /// # Safety
496        /// The caller must provide valid device pointers for `out` and `raw_data`,
497        /// and `out` must have capacity for `n` field elements of the target type.
498        #[inline]
499        $vis unsafe extern "C" fn $name(out: *mut c_void, raw_data: *const u32, n: usize) -> i32 {
500            ffi::$name(out, raw_data, n, bench_ctx().stream.as_raw())
501        }
502    };
503}
504
505macro_rules! wrap_init_u64 {
506    ($vis:vis $name:ident) => {
507        /// Launches the benchmark init kernel on the benchmark stream.
508        ///
509        /// # Safety
510        /// The caller must provide valid device pointers for `out` and `raw_data`,
511        /// and `out` must have capacity for `n` field elements of the target type.
512        #[inline]
513        $vis unsafe extern "C" fn $name(out: *mut c_void, raw_data: *const u64, n: usize) -> i32 {
514            ffi::$name(out, raw_data, n, bench_ctx().stream.as_raw())
515        }
516    };
517}
518
519macro_rules! wrap_binary {
520    ($name:ident) => {
521        #[inline]
522        unsafe extern "C" fn $name(
523            out: *mut c_void,
524            a: *const c_void,
525            b: *const c_void,
526            n: usize,
527            reps: i32,
528        ) -> i32 {
529            ffi::$name(out, a, b, n, reps, bench_ctx().stream.as_raw())
530        }
531    };
532}
533
534macro_rules! wrap_unary {
535    ($name:ident) => {
536        #[inline]
537        unsafe extern "C" fn $name(out: *mut c_void, a: *const c_void, n: usize, reps: i32) -> i32 {
538            ffi::$name(out, a, n, reps, bench_ctx().stream.as_raw())
539        }
540    };
541}
542
543wrap_init_u32!(pub init_fp);
544wrap_binary!(add_fp);
545wrap_binary!(mul_fp);
546wrap_unary!(inv_fp);
547
548wrap_init_u32!(pub init_fp4);
549wrap_binary!(add_fp4);
550wrap_binary!(mul_fp4);
551wrap_unary!(inv_fp4);
552
553wrap_init_u32!(init_fpext);
554wrap_binary!(add_fpext);
555wrap_binary!(mul_fpext);
556wrap_unary!(inv_fpext);
557
558wrap_init_u32!(pub init_fp5);
559wrap_binary!(add_fp5);
560wrap_binary!(mul_fp5);
561wrap_unary!(inv_fp5);
562
563wrap_init_u32!(pub init_fp6);
564wrap_binary!(add_fp6);
565wrap_binary!(mul_fp6);
566wrap_unary!(inv_fp6);
567
568wrap_init_u32!(pub init_fp2x3);
569wrap_binary!(add_fp2x3);
570wrap_binary!(mul_fp2x3);
571wrap_unary!(inv_fp2x3);
572
573wrap_init_u32!(pub init_fp3x2);
574wrap_binary!(add_fp3x2);
575wrap_binary!(mul_fp3x2);
576wrap_unary!(inv_fp3x2);
577
578wrap_init_u32!(pub init_kb);
579wrap_binary!(add_kb);
580wrap_binary!(mul_kb);
581wrap_unary!(inv_kb);
582
583wrap_init_u32!(pub init_kb5);
584wrap_binary!(add_kb5);
585wrap_binary!(mul_kb5);
586wrap_unary!(inv_kb5);
587
588wrap_init_u32!(pub init_kb6);
589wrap_binary!(add_kb6);
590wrap_binary!(mul_kb6);
591wrap_unary!(inv_kb6);
592
593wrap_init_u32!(pub init_kb2x3);
594wrap_binary!(add_kb2x3);
595wrap_binary!(mul_kb2x3);
596wrap_unary!(inv_kb2x3);
597
598wrap_init_u32!(pub init_kb3x2);
599wrap_binary!(add_kb3x2);
600wrap_binary!(mul_kb3x2);
601wrap_unary!(inv_kb3x2);
602
603wrap_init_u64!(init_gl);
604wrap_binary!(add_gl);
605wrap_binary!(mul_gl);
606wrap_unary!(inv_gl);
607
608wrap_init_u64!(init_gl3);
609wrap_binary!(add_gl3);
610wrap_binary!(mul_gl3);
611wrap_unary!(inv_gl3);
612
613wrap_init_u32!(init_poseidon2_bb);
614
615#[inline]
616unsafe extern "C" fn run_poseidon2_bb(states: *mut c_void, n: usize, reps: i32) -> i32 {
617    ffi::run_poseidon2_bb(states, n, reps, bench_ctx().stream.as_raw())
618}
619
620wrap_init_u32!(init_poseidon2_kb);
621
622#[inline]
623unsafe extern "C" fn run_poseidon2_kb(states: *mut c_void, n: usize, reps: i32) -> i32 {
624    ffi::run_poseidon2_kb(states, n, reps, bench_ctx().stream.as_raw())
625}
626
627/// Check CUDA return code, panic on error
628pub fn cuda_check(code: i32) {
629    assert!(code == 0, "CUDA error: {}", code);
630}
631
632/// Sync and check
633pub fn sync() {
634    bench_ctx().stream.synchronize().expect("sync failed");
635}
636
637// ============================================================================
638// Configuration & Results
639// ============================================================================
640
641pub struct BenchConfig {
642    pub num_elements: usize,
643    pub warmup_iters: usize,
644    pub bench_iters: usize,
645    pub ops_per_element: i32,
646}
647
648impl Default for BenchConfig {
649    fn default() -> Self {
650        Self {
651            num_elements: 1 << 22, // 4M elements
652            warmup_iters: 3,
653            bench_iters: 10,
654            ops_per_element: 100,
655        }
656    }
657}
658
659#[derive(Debug, Clone)]
660pub struct OpResult {
661    pub avg_time_ms: f64,
662    pub throughput_gops: f64,
663}
664
665#[derive(Clone)]
666pub struct FieldBenchResult {
667    pub field_name: String,
668    pub bits: usize, // bits of provable security
669    pub u32s_per_element: usize,
670    pub init: OpResult,
671    pub add: OpResult,
672    pub mul: OpResult,
673    pub inv: OpResult,
674}
675
676/// Print a table of benchmark results
677pub fn print_benchmark_tables(results: &[FieldBenchResult], baseline: &FieldBenchResult) {
678    // Find max field name length for alignment
679    let max_name_len = results
680        .iter()
681        .map(|r| r.field_name.len())
682        .max()
683        .unwrap_or(10);
684
685    // Table 1: Time (ms)
686    println!("### Time (ms)");
687    println!();
688    print!("| {:width$} | bits |", "Field", width = max_name_len);
689    println!("    init |     add |     mul |     inv |");
690    print!("|{:-<width$}--|-----:|", "", width = max_name_len);
691    println!("--------:|--------:|--------:|--------:|");
692    for r in results {
693        print!(
694            "| {:width$} | {:>4} |",
695            r.field_name,
696            r.bits,
697            width = max_name_len
698        );
699        println!(
700            " {:>7.3} | {:>7.3} | {:>7.3} | {:>7.3} |",
701            r.init.avg_time_ms, r.add.avg_time_ms, r.mul.avg_time_ms, r.inv.avg_time_ms
702        );
703    }
704    println!();
705
706    // Table 2: Throughput (Gops/s)
707    println!("### Throughput (Gops/s)");
708    println!();
709    print!("| {:width$} | bits |", "Field", width = max_name_len);
710    println!("    init |     add |     mul |     inv |");
711    print!("|{:-<width$}--|-----:|", "", width = max_name_len);
712    println!("--------:|--------:|--------:|--------:|");
713    for r in results {
714        print!(
715            "| {:width$} | {:>4} |",
716            r.field_name,
717            r.bits,
718            width = max_name_len
719        );
720        println!(
721            " {:>7.1} | {:>7.1} | {:>7.1} | {:>7.1} |",
722            r.init.throughput_gops,
723            r.add.throughput_gops,
724            r.mul.throughput_gops,
725            r.inv.throughput_gops
726        );
727    }
728    println!();
729
730    // Table 3: Relative to baseline (xN)
731    println!("### Relative to {} (xN)", baseline.field_name);
732    println!();
733    print!("| {:width$} | bits |", "Field", width = max_name_len);
734    println!("    init |     add |     mul |     inv |");
735    print!("|{:-<width$}--|-----:|", "", width = max_name_len);
736    println!("--------:|--------:|--------:|--------:|");
737    for r in results {
738        let init_ratio = baseline.init.throughput_gops / r.init.throughput_gops;
739        let add_ratio = baseline.add.throughput_gops / r.add.throughput_gops;
740        let mul_ratio = baseline.mul.throughput_gops / r.mul.throughput_gops;
741        let inv_ratio = baseline.inv.throughput_gops / r.inv.throughput_gops;
742
743        print!(
744            "| {:width$} | {:>4} |",
745            r.field_name,
746            r.bits,
747            width = max_name_len
748        );
749        println!(
750            " {:>7.1} | {:>7.1} | {:>7.1} | {:>7.1} |",
751            init_ratio, add_ratio, mul_ratio, inv_ratio
752        );
753    }
754    println!();
755}
756
757// ============================================================================
758// Benchmark Implementation
759// ============================================================================
760
761pub fn random_u32s(count: usize, seed: u64) -> Vec<u32> {
762    let mut rng = seed;
763    (0..count)
764        .map(|_| {
765            rng = rng
766                .wrapping_mul(6364136223846793005)
767                .wrapping_add(1442695040888963407);
768            (rng >> 32) as u32
769        })
770        .collect()
771}
772
773pub fn random_u64s(count: usize, seed: u64) -> Vec<u64> {
774    let mut rng = seed;
775    (0..count)
776        .map(|_| {
777            rng = rng
778                .wrapping_mul(6364136223846793005)
779                .wrapping_add(1442695040888963407);
780            rng
781        })
782        .collect()
783}
784
785pub fn measure<F: FnMut()>(config: &BenchConfig, total_ops: u64, mut f: F) -> OpResult {
786    // Warmup
787    for _ in 0..config.warmup_iters {
788        f();
789    }
790    sync();
791
792    // Timed
793    let start = Instant::now();
794    for _ in 0..config.bench_iters {
795        f();
796    }
797    sync();
798    let elapsed = start.elapsed();
799
800    let avg_time_ms = elapsed.as_secs_f64() * 1000.0 / config.bench_iters as f64;
801    let throughput_gops =
802        (total_ops * config.bench_iters as u64) as f64 / elapsed.as_secs_f64() / 1e9;
803
804    OpResult {
805        avg_time_ms,
806        throughput_gops,
807    }
808}
809
810pub fn bench_fp(config: &BenchConfig) -> FieldBenchResult {
811    let n = config.num_elements;
812    let reps = config.ops_per_element;
813
814    // Only 3 device buffers: a, b, out (init works in-place)
815    let d_a = random_u32s(n, 12345).to_device_on(bench_ctx()).unwrap();
816    let d_b = random_u32s(n, 67890).to_device_on(bench_ctx()).unwrap();
817    let d_out = DeviceBuffer::<u32>::with_capacity_on(n, bench_ctx());
818
819    // Benchmark init (in-place: raw u32 -> Fp) - also initializes d_a
820    let init = measure(config, n as u64, || {
821        cuda_check(unsafe { init_fp(d_a.as_mut_raw_ptr(), d_a.as_ptr(), n) });
822    });
823    cuda_check(unsafe { init_fp(d_b.as_mut_raw_ptr(), d_b.as_ptr(), n) });
824    sync();
825
826    let ops = n as u64 * reps as u64;
827
828    let add = measure(config, ops, || {
829        cuda_check(unsafe {
830            add_fp(
831                d_out.as_mut_raw_ptr(),
832                d_a.as_raw_ptr(),
833                d_b.as_raw_ptr(),
834                n,
835                reps,
836            )
837        });
838    });
839
840    let mul = measure(config, ops, || {
841        cuda_check(unsafe {
842            mul_fp(
843                d_out.as_mut_raw_ptr(),
844                d_a.as_raw_ptr(),
845                d_b.as_raw_ptr(),
846                n,
847                reps,
848            )
849        });
850    });
851
852    let inv = measure(config, ops, || {
853        cuda_check(unsafe { inv_fp(d_out.as_mut_raw_ptr(), d_a.as_raw_ptr(), n, reps) });
854    });
855
856    FieldBenchResult {
857        field_name: "Fp".into(),
858        bits: 31,
859        u32s_per_element: 1,
860        init,
861        add,
862        mul,
863        inv,
864    }
865}
866
867pub fn bench_fp4(config: &BenchConfig) -> FieldBenchResult {
868    let n = config.num_elements;
869    let reps = config.ops_per_element;
870
871    let d_a = random_u32s(n * 4, 11111).to_device_on(bench_ctx()).unwrap();
872    let d_b = random_u32s(n * 4, 22222).to_device_on(bench_ctx()).unwrap();
873    let d_out = DeviceBuffer::<u32>::with_capacity_on(n * 4, bench_ctx());
874
875    let init = measure(config, n as u64, || {
876        cuda_check(unsafe { init_fp4(d_a.as_mut_raw_ptr(), d_a.as_ptr(), n) });
877    });
878    cuda_check(unsafe { init_fp4(d_b.as_mut_raw_ptr(), d_b.as_ptr(), n) });
879    sync();
880
881    let ops = n as u64 * reps as u64;
882
883    let add = measure(config, ops, || {
884        cuda_check(unsafe {
885            add_fp4(
886                d_out.as_mut_raw_ptr(),
887                d_a.as_raw_ptr(),
888                d_b.as_raw_ptr(),
889                n,
890                reps,
891            )
892        });
893    });
894
895    let mul = measure(config, ops, || {
896        cuda_check(unsafe {
897            mul_fp4(
898                d_out.as_mut_raw_ptr(),
899                d_a.as_raw_ptr(),
900                d_b.as_raw_ptr(),
901                n,
902                reps,
903            )
904        });
905    });
906
907    let inv = measure(config, ops, || {
908        cuda_check(unsafe { inv_fp4(d_out.as_mut_raw_ptr(), d_a.as_raw_ptr(), n, reps) });
909    });
910
911    FieldBenchResult {
912        field_name: "Fp4".into(),
913        bits: 124,
914        u32s_per_element: 4,
915        init,
916        add,
917        mul,
918        inv,
919    }
920}
921
922pub fn bench_fpext(config: &BenchConfig) -> FieldBenchResult {
923    let n = config.num_elements;
924    let reps = config.ops_per_element;
925
926    // Only 3 device buffers: a, b, out (init works in-place)
927    let d_a = random_u32s(n * 4, 12345).to_device_on(bench_ctx()).unwrap();
928    let d_b = random_u32s(n * 4, 67890).to_device_on(bench_ctx()).unwrap();
929    let d_out = DeviceBuffer::<u32>::with_capacity_on(n * 4, bench_ctx());
930
931    // Benchmark init (in-place: raw u32s -> FpExt) - also initializes d_a
932    let init = measure(config, n as u64, || {
933        cuda_check(unsafe { init_fpext(d_a.as_mut_raw_ptr(), d_a.as_ptr(), n) });
934    });
935    cuda_check(unsafe { init_fpext(d_b.as_mut_raw_ptr(), d_b.as_ptr(), n) });
936    sync();
937
938    let ops = n as u64 * reps as u64;
939
940    let add = measure(config, ops, || {
941        cuda_check(unsafe {
942            add_fpext(
943                d_out.as_mut_raw_ptr(),
944                d_a.as_raw_ptr(),
945                d_b.as_raw_ptr(),
946                n,
947                reps,
948            )
949        });
950    });
951
952    let mul = measure(config, ops, || {
953        cuda_check(unsafe {
954            mul_fpext(
955                d_out.as_mut_raw_ptr(),
956                d_a.as_raw_ptr(),
957                d_b.as_raw_ptr(),
958                n,
959                reps,
960            )
961        });
962    });
963
964    let inv = measure(config, ops, || {
965        cuda_check(unsafe { inv_fpext(d_out.as_mut_raw_ptr(), d_a.as_raw_ptr(), n, reps) });
966    });
967
968    FieldBenchResult {
969        field_name: "FpExt".into(),
970        bits: 124,
971        u32s_per_element: 4,
972        init,
973        add,
974        mul,
975        inv,
976    }
977}
978
979pub fn bench_fp5(config: &BenchConfig) -> FieldBenchResult {
980    let n = config.num_elements;
981    let reps = config.ops_per_element;
982
983    // Only 3 device buffers: a, b, out (init works in-place)
984    let d_a = random_u32s(n * 5, 12345).to_device_on(bench_ctx()).unwrap();
985    let d_b = random_u32s(n * 5, 67890).to_device_on(bench_ctx()).unwrap();
986    let d_out = DeviceBuffer::<u32>::with_capacity_on(n * 5, bench_ctx());
987
988    // Benchmark init (in-place: raw u32s -> Fp5) - also initializes d_a
989    let init = measure(config, n as u64, || {
990        cuda_check(unsafe { init_fp5(d_a.as_mut_raw_ptr(), d_a.as_ptr(), n) });
991    });
992    cuda_check(unsafe { init_fp5(d_b.as_mut_raw_ptr(), d_b.as_ptr(), n) });
993    sync();
994
995    let ops = n as u64 * reps as u64;
996
997    let add = measure(config, ops, || {
998        cuda_check(unsafe {
999            add_fp5(
1000                d_out.as_mut_raw_ptr(),
1001                d_a.as_raw_ptr(),
1002                d_b.as_raw_ptr(),
1003                n,
1004                reps,
1005            )
1006        });
1007    });
1008
1009    let mul = measure(config, ops, || {
1010        cuda_check(unsafe {
1011            mul_fp5(
1012                d_out.as_mut_raw_ptr(),
1013                d_a.as_raw_ptr(),
1014                d_b.as_raw_ptr(),
1015                n,
1016                reps,
1017            )
1018        });
1019    });
1020
1021    let inv = measure(config, ops, || {
1022        cuda_check(unsafe { inv_fp5(d_out.as_mut_raw_ptr(), d_a.as_raw_ptr(), n, reps) });
1023    });
1024
1025    FieldBenchResult {
1026        field_name: "Fp5".into(),
1027        bits: 155,
1028        u32s_per_element: 5,
1029        init,
1030        add,
1031        mul,
1032        inv,
1033    }
1034}
1035
1036pub fn bench_fp6(config: &BenchConfig) -> FieldBenchResult {
1037    let n = config.num_elements;
1038    let reps = config.ops_per_element;
1039
1040    // Only 3 device buffers: a, b, out (init works in-place)
1041    let d_a = random_u32s(n * 6, 12345).to_device_on(bench_ctx()).unwrap();
1042    let d_b = random_u32s(n * 6, 67890).to_device_on(bench_ctx()).unwrap();
1043    let d_out = DeviceBuffer::<u32>::with_capacity_on(n * 6, bench_ctx());
1044
1045    // Benchmark init (in-place: raw u32s -> Fp6) - also initializes d_a
1046    let init = measure(config, n as u64, || {
1047        cuda_check(unsafe { init_fp6(d_a.as_mut_raw_ptr(), d_a.as_ptr(), n) });
1048    });
1049    cuda_check(unsafe { init_fp6(d_b.as_mut_raw_ptr(), d_b.as_ptr(), n) });
1050    sync();
1051
1052    let ops = n as u64 * reps as u64;
1053
1054    let add = measure(config, ops, || {
1055        cuda_check(unsafe {
1056            add_fp6(
1057                d_out.as_mut_raw_ptr(),
1058                d_a.as_raw_ptr(),
1059                d_b.as_raw_ptr(),
1060                n,
1061                reps,
1062            )
1063        });
1064    });
1065
1066    let mul = measure(config, ops, || {
1067        cuda_check(unsafe {
1068            mul_fp6(
1069                d_out.as_mut_raw_ptr(),
1070                d_a.as_raw_ptr(),
1071                d_b.as_raw_ptr(),
1072                n,
1073                reps,
1074            )
1075        });
1076    });
1077
1078    let inv = measure(config, ops, || {
1079        cuda_check(unsafe { inv_fp6(d_out.as_mut_raw_ptr(), d_a.as_raw_ptr(), n, reps) });
1080    });
1081
1082    FieldBenchResult {
1083        field_name: "Fp6".into(),
1084        bits: 186,
1085        u32s_per_element: 6,
1086        init,
1087        add,
1088        mul,
1089        inv,
1090    }
1091}
1092
1093pub fn bench_fp2x3(config: &BenchConfig) -> FieldBenchResult {
1094    let n = config.num_elements;
1095    let reps = config.ops_per_element;
1096
1097    let d_a = random_u32s(n * 6, 12345).to_device_on(bench_ctx()).unwrap();
1098    let d_b = random_u32s(n * 6, 67890).to_device_on(bench_ctx()).unwrap();
1099    let d_out = DeviceBuffer::<u32>::with_capacity_on(n * 6, bench_ctx());
1100
1101    let init = measure(config, n as u64, || {
1102        cuda_check(unsafe { init_fp2x3(d_a.as_mut_raw_ptr(), d_a.as_ptr(), n) });
1103    });
1104    cuda_check(unsafe { init_fp2x3(d_b.as_mut_raw_ptr(), d_b.as_ptr(), n) });
1105    sync();
1106
1107    let ops = n as u64 * reps as u64;
1108
1109    let add = measure(config, ops, || {
1110        cuda_check(unsafe {
1111            add_fp2x3(
1112                d_out.as_mut_raw_ptr(),
1113                d_a.as_raw_ptr(),
1114                d_b.as_raw_ptr(),
1115                n,
1116                reps,
1117            )
1118        });
1119    });
1120
1121    let mul = measure(config, ops, || {
1122        cuda_check(unsafe {
1123            mul_fp2x3(
1124                d_out.as_mut_raw_ptr(),
1125                d_a.as_raw_ptr(),
1126                d_b.as_raw_ptr(),
1127                n,
1128                reps,
1129            )
1130        });
1131    });
1132
1133    let inv = measure(config, ops, || {
1134        cuda_check(unsafe { inv_fp2x3(d_out.as_mut_raw_ptr(), d_a.as_raw_ptr(), n, reps) });
1135    });
1136
1137    FieldBenchResult {
1138        field_name: "Fp2x3".into(),
1139        bits: 186,
1140        u32s_per_element: 6,
1141        init,
1142        add,
1143        mul,
1144        inv,
1145    }
1146}
1147
1148pub fn bench_fp3x2(config: &BenchConfig) -> FieldBenchResult {
1149    let n = config.num_elements;
1150    let reps = config.ops_per_element;
1151
1152    let d_a = random_u32s(n * 6, 12345).to_device_on(bench_ctx()).unwrap();
1153    let d_b = random_u32s(n * 6, 67890).to_device_on(bench_ctx()).unwrap();
1154    let d_out = DeviceBuffer::<u32>::with_capacity_on(n * 6, bench_ctx());
1155
1156    let init = measure(config, n as u64, || {
1157        cuda_check(unsafe { init_fp3x2(d_a.as_mut_raw_ptr(), d_a.as_ptr(), n) });
1158    });
1159    cuda_check(unsafe { init_fp3x2(d_b.as_mut_raw_ptr(), d_b.as_ptr(), n) });
1160    sync();
1161
1162    let ops = n as u64 * reps as u64;
1163
1164    let add = measure(config, ops, || {
1165        cuda_check(unsafe {
1166            add_fp3x2(
1167                d_out.as_mut_raw_ptr(),
1168                d_a.as_raw_ptr(),
1169                d_b.as_raw_ptr(),
1170                n,
1171                reps,
1172            )
1173        });
1174    });
1175
1176    let mul = measure(config, ops, || {
1177        cuda_check(unsafe {
1178            mul_fp3x2(
1179                d_out.as_mut_raw_ptr(),
1180                d_a.as_raw_ptr(),
1181                d_b.as_raw_ptr(),
1182                n,
1183                reps,
1184            )
1185        });
1186    });
1187
1188    let inv = measure(config, ops, || {
1189        cuda_check(unsafe { inv_fp3x2(d_out.as_mut_raw_ptr(), d_a.as_raw_ptr(), n, reps) });
1190    });
1191
1192    FieldBenchResult {
1193        field_name: "Fp3x2".into(),
1194        bits: 186,
1195        u32s_per_element: 6,
1196        init,
1197        add,
1198        mul,
1199        inv,
1200    }
1201}
1202
1203pub fn bench_kb(config: &BenchConfig) -> FieldBenchResult {
1204    let n = config.num_elements;
1205    let reps = config.ops_per_element;
1206
1207    // Only 3 device buffers: a, b, out (init works in-place)
1208    let d_a = random_u32s(n, 12345).to_device_on(bench_ctx()).unwrap();
1209    let d_b = random_u32s(n, 67890).to_device_on(bench_ctx()).unwrap();
1210    let d_out = DeviceBuffer::<u32>::with_capacity_on(n, bench_ctx());
1211
1212    // Benchmark init (in-place: raw u32 -> Kb) - also initializes d_a
1213    let init = measure(config, n as u64, || {
1214        cuda_check(unsafe { init_kb(d_a.as_mut_raw_ptr(), d_a.as_ptr(), n) });
1215    });
1216    cuda_check(unsafe { init_kb(d_b.as_mut_raw_ptr(), d_b.as_ptr(), n) });
1217    sync();
1218
1219    let ops = n as u64 * reps as u64;
1220
1221    let add = measure(config, ops, || {
1222        cuda_check(unsafe {
1223            add_kb(
1224                d_out.as_mut_raw_ptr(),
1225                d_a.as_raw_ptr(),
1226                d_b.as_raw_ptr(),
1227                n,
1228                reps,
1229            )
1230        });
1231    });
1232
1233    let mul = measure(config, ops, || {
1234        cuda_check(unsafe {
1235            mul_kb(
1236                d_out.as_mut_raw_ptr(),
1237                d_a.as_raw_ptr(),
1238                d_b.as_raw_ptr(),
1239                n,
1240                reps,
1241            )
1242        });
1243    });
1244
1245    let inv = measure(config, ops, || {
1246        cuda_check(unsafe { inv_kb(d_out.as_mut_raw_ptr(), d_a.as_raw_ptr(), n, reps) });
1247    });
1248
1249    FieldBenchResult {
1250        field_name: "Kb".into(),
1251        bits: 31,
1252        u32s_per_element: 1,
1253        init,
1254        add,
1255        mul,
1256        inv,
1257    }
1258}
1259
1260pub fn bench_kb5(config: &BenchConfig) -> FieldBenchResult {
1261    let n = config.num_elements;
1262    let reps = config.ops_per_element;
1263
1264    // Kb5 has 5 u32s per element
1265    let d_a = random_u32s(n * 5, 11111).to_device_on(bench_ctx()).unwrap();
1266    let d_b = random_u32s(n * 5, 22222).to_device_on(bench_ctx()).unwrap();
1267    let d_out = DeviceBuffer::<u32>::with_capacity_on(n * 5, bench_ctx());
1268
1269    let init = measure(config, n as u64, || {
1270        cuda_check(unsafe { init_kb5(d_a.as_mut_raw_ptr(), d_a.as_ptr(), n) });
1271    });
1272    cuda_check(unsafe { init_kb5(d_b.as_mut_raw_ptr(), d_b.as_ptr(), n) });
1273    sync();
1274
1275    let ops = n as u64 * reps as u64;
1276
1277    let add = measure(config, ops, || {
1278        cuda_check(unsafe {
1279            add_kb5(
1280                d_out.as_mut_raw_ptr(),
1281                d_a.as_raw_ptr(),
1282                d_b.as_raw_ptr(),
1283                n,
1284                reps,
1285            )
1286        });
1287    });
1288
1289    let mul = measure(config, ops, || {
1290        cuda_check(unsafe {
1291            mul_kb5(
1292                d_out.as_mut_raw_ptr(),
1293                d_a.as_raw_ptr(),
1294                d_b.as_raw_ptr(),
1295                n,
1296                reps,
1297            )
1298        });
1299    });
1300
1301    let inv = measure(config, ops, || {
1302        cuda_check(unsafe { inv_kb5(d_out.as_mut_raw_ptr(), d_a.as_raw_ptr(), n, reps) });
1303    });
1304
1305    FieldBenchResult {
1306        field_name: "Kb5".into(),
1307        bits: 155,
1308        u32s_per_element: 5,
1309        init,
1310        add,
1311        mul,
1312        inv,
1313    }
1314}
1315
1316pub fn bench_kb6(config: &BenchConfig) -> FieldBenchResult {
1317    let n = config.num_elements;
1318    let reps = config.ops_per_element;
1319
1320    // Kb6 has 6 u32s per element
1321    let d_a = random_u32s(n * 6, 33333).to_device_on(bench_ctx()).unwrap();
1322    let d_b = random_u32s(n * 6, 44444).to_device_on(bench_ctx()).unwrap();
1323    let d_out = DeviceBuffer::<u32>::with_capacity_on(n * 6, bench_ctx());
1324
1325    let init = measure(config, n as u64, || {
1326        cuda_check(unsafe { init_kb6(d_a.as_mut_raw_ptr(), d_a.as_ptr(), n) });
1327    });
1328    cuda_check(unsafe { init_kb6(d_b.as_mut_raw_ptr(), d_b.as_ptr(), n) });
1329    sync();
1330
1331    let ops = n as u64 * reps as u64;
1332
1333    let add = measure(config, ops, || {
1334        cuda_check(unsafe {
1335            add_kb6(
1336                d_out.as_mut_raw_ptr(),
1337                d_a.as_raw_ptr(),
1338                d_b.as_raw_ptr(),
1339                n,
1340                reps,
1341            )
1342        });
1343    });
1344
1345    let mul = measure(config, ops, || {
1346        cuda_check(unsafe {
1347            mul_kb6(
1348                d_out.as_mut_raw_ptr(),
1349                d_a.as_raw_ptr(),
1350                d_b.as_raw_ptr(),
1351                n,
1352                reps,
1353            )
1354        });
1355    });
1356
1357    let inv = measure(config, ops, || {
1358        cuda_check(unsafe { inv_kb6(d_out.as_mut_raw_ptr(), d_a.as_raw_ptr(), n, reps) });
1359    });
1360
1361    FieldBenchResult {
1362        field_name: "Kb6".into(),
1363        bits: 186,
1364        u32s_per_element: 6,
1365        init,
1366        add,
1367        mul,
1368        inv,
1369    }
1370}
1371
1372pub fn bench_kb2x3(config: &BenchConfig) -> FieldBenchResult {
1373    let n = config.num_elements;
1374    let reps = config.ops_per_element;
1375
1376    // Kb2x3 has 6 u32s per element
1377    let d_a = random_u32s(n * 6, 55555).to_device_on(bench_ctx()).unwrap();
1378    let d_b = random_u32s(n * 6, 66666).to_device_on(bench_ctx()).unwrap();
1379    let d_out = DeviceBuffer::<u32>::with_capacity_on(n * 6, bench_ctx());
1380
1381    let init = measure(config, n as u64, || {
1382        cuda_check(unsafe { init_kb2x3(d_a.as_mut_raw_ptr(), d_a.as_ptr(), n) });
1383    });
1384    cuda_check(unsafe { init_kb2x3(d_b.as_mut_raw_ptr(), d_b.as_ptr(), n) });
1385    sync();
1386
1387    let ops = n as u64 * reps as u64;
1388
1389    let add = measure(config, ops, || {
1390        cuda_check(unsafe {
1391            add_kb2x3(
1392                d_out.as_mut_raw_ptr(),
1393                d_a.as_raw_ptr(),
1394                d_b.as_raw_ptr(),
1395                n,
1396                reps,
1397            )
1398        });
1399    });
1400
1401    let mul = measure(config, ops, || {
1402        cuda_check(unsafe {
1403            mul_kb2x3(
1404                d_out.as_mut_raw_ptr(),
1405                d_a.as_raw_ptr(),
1406                d_b.as_raw_ptr(),
1407                n,
1408                reps,
1409            )
1410        });
1411    });
1412
1413    let inv = measure(config, ops, || {
1414        cuda_check(unsafe { inv_kb2x3(d_out.as_mut_raw_ptr(), d_a.as_raw_ptr(), n, reps) });
1415    });
1416
1417    FieldBenchResult {
1418        field_name: "Kb2x3".into(),
1419        bits: 186,
1420        u32s_per_element: 6,
1421        init,
1422        add,
1423        mul,
1424        inv,
1425    }
1426}
1427
1428pub fn bench_kb3x2(config: &BenchConfig) -> FieldBenchResult {
1429    let n = config.num_elements;
1430    let reps = config.ops_per_element;
1431
1432    // Kb3x2 has 6 u32s per element
1433    let d_a = random_u32s(n * 6, 77777).to_device_on(bench_ctx()).unwrap();
1434    let d_b = random_u32s(n * 6, 88888).to_device_on(bench_ctx()).unwrap();
1435    let d_out = DeviceBuffer::<u32>::with_capacity_on(n * 6, bench_ctx());
1436
1437    let init = measure(config, n as u64, || {
1438        cuda_check(unsafe { init_kb3x2(d_a.as_mut_raw_ptr(), d_a.as_ptr(), n) });
1439    });
1440    cuda_check(unsafe { init_kb3x2(d_b.as_mut_raw_ptr(), d_b.as_ptr(), n) });
1441    sync();
1442
1443    let ops = n as u64 * reps as u64;
1444
1445    let add = measure(config, ops, || {
1446        cuda_check(unsafe {
1447            add_kb3x2(
1448                d_out.as_mut_raw_ptr(),
1449                d_a.as_raw_ptr(),
1450                d_b.as_raw_ptr(),
1451                n,
1452                reps,
1453            )
1454        });
1455    });
1456
1457    let mul = measure(config, ops, || {
1458        cuda_check(unsafe {
1459            mul_kb3x2(
1460                d_out.as_mut_raw_ptr(),
1461                d_a.as_raw_ptr(),
1462                d_b.as_raw_ptr(),
1463                n,
1464                reps,
1465            )
1466        });
1467    });
1468
1469    let inv = measure(config, ops, || {
1470        cuda_check(unsafe { inv_kb3x2(d_out.as_mut_raw_ptr(), d_a.as_raw_ptr(), n, reps) });
1471    });
1472
1473    FieldBenchResult {
1474        field_name: "Kb3x2".into(),
1475        bits: 186,
1476        u32s_per_element: 6,
1477        init,
1478        add,
1479        mul,
1480        inv,
1481    }
1482}
1483
1484pub fn bench_gl(config: &BenchConfig) -> FieldBenchResult {
1485    let n = config.num_elements;
1486    let reps = config.ops_per_element;
1487
1488    // Goldilocks uses u64 elements (8 bytes)
1489    let h_a = random_u64s(n, 99999);
1490    let h_b = random_u64s(n, 88888);
1491
1492    let d_a = h_a.to_device_on(bench_ctx()).unwrap();
1493    let d_b = h_b.to_device_on(bench_ctx()).unwrap();
1494    let d_out = DeviceBuffer::<u64>::with_capacity_on(n, bench_ctx());
1495
1496    let init = measure(config, n as u64, || {
1497        cuda_check(unsafe { init_gl(d_a.as_mut_raw_ptr(), d_a.as_ptr(), n) });
1498    });
1499    cuda_check(unsafe { init_gl(d_b.as_mut_raw_ptr(), d_b.as_ptr(), n) });
1500    sync();
1501
1502    let ops = n as u64 * reps as u64;
1503
1504    let add = measure(config, ops, || {
1505        cuda_check(unsafe {
1506            add_gl(
1507                d_out.as_mut_raw_ptr(),
1508                d_a.as_raw_ptr(),
1509                d_b.as_raw_ptr(),
1510                n,
1511                reps,
1512            )
1513        });
1514    });
1515
1516    let mul = measure(config, ops, || {
1517        cuda_check(unsafe {
1518            mul_gl(
1519                d_out.as_mut_raw_ptr(),
1520                d_a.as_raw_ptr(),
1521                d_b.as_raw_ptr(),
1522                n,
1523                reps,
1524            )
1525        });
1526    });
1527
1528    let inv = measure(config, ops, || {
1529        cuda_check(unsafe { inv_gl(d_out.as_mut_raw_ptr(), d_a.as_raw_ptr(), n, reps) });
1530    });
1531
1532    // u32s_per_element = 2 for Goldilocks (64-bit)
1533    FieldBenchResult {
1534        field_name: "Gl".into(),
1535        bits: 64,
1536        u32s_per_element: 2,
1537        init,
1538        add,
1539        mul,
1540        inv,
1541    }
1542}
1543
1544pub fn bench_gl3(config: &BenchConfig) -> FieldBenchResult {
1545    let n = config.num_elements;
1546    let reps = config.ops_per_element;
1547
1548    // Gl3 uses 3 x u64 elements (24 bytes)
1549    let h_a = random_u64s(n * 3, 77777);
1550    let h_b = random_u64s(n * 3, 66666);
1551
1552    let d_a = h_a.to_device_on(bench_ctx()).unwrap();
1553    let d_b = h_b.to_device_on(bench_ctx()).unwrap();
1554    let d_out = DeviceBuffer::<u64>::with_capacity_on(n * 3, bench_ctx());
1555
1556    let init = measure(config, n as u64, || {
1557        cuda_check(unsafe { init_gl3(d_a.as_mut_raw_ptr(), d_a.as_ptr(), n) });
1558    });
1559    cuda_check(unsafe { init_gl3(d_b.as_mut_raw_ptr(), d_b.as_ptr(), n) });
1560    sync();
1561
1562    let ops = n as u64 * reps as u64;
1563
1564    let add = measure(config, ops, || {
1565        cuda_check(unsafe {
1566            add_gl3(
1567                d_out.as_mut_raw_ptr(),
1568                d_a.as_raw_ptr(),
1569                d_b.as_raw_ptr(),
1570                n,
1571                reps,
1572            )
1573        });
1574    });
1575
1576    let mul = measure(config, ops, || {
1577        cuda_check(unsafe {
1578            mul_gl3(
1579                d_out.as_mut_raw_ptr(),
1580                d_a.as_raw_ptr(),
1581                d_b.as_raw_ptr(),
1582                n,
1583                reps,
1584            )
1585        });
1586    });
1587
1588    let inv = measure(config, ops, || {
1589        cuda_check(unsafe { inv_gl3(d_out.as_mut_raw_ptr(), d_a.as_raw_ptr(), n, reps) });
1590    });
1591
1592    // u32s_per_element = 6 for Gl3 (3 x 64-bit = 24 bytes)
1593    FieldBenchResult {
1594        field_name: "Gl3".into(),
1595        bits: 192,
1596        u32s_per_element: 6,
1597        init,
1598        add,
1599        mul,
1600        inv,
1601    }
1602}
1603
1604// ============================================================================
1605// Poseidon2 Benchmark
1606// ============================================================================
1607
1608#[derive(Debug, Clone)]
1609pub struct Poseidon2BenchResult {
1610    pub name: String,
1611    pub avg_time_ms: f64,
1612    pub throughput_gops: f64,
1613}
1614
1615pub fn bench_poseidon2_bb(config: &BenchConfig) -> Poseidon2BenchResult {
1616    let n = config.num_elements;
1617    let reps = config.ops_per_element;
1618
1619    // 16 u32s per poseidon2 state
1620    let d_states = random_u32s(n * 16, 55555)
1621        .to_device_on(bench_ctx())
1622        .unwrap();
1623
1624    // Initialize: raw u32 -> Fp field elements
1625    cuda_check(unsafe { init_poseidon2_bb(d_states.as_mut_raw_ptr(), d_states.as_ptr(), n) });
1626    sync();
1627
1628    let ops = n as u64 * reps as u64;
1629    let result = measure(config, ops, || {
1630        cuda_check(unsafe { run_poseidon2_bb(d_states.as_mut_raw_ptr(), n, reps) });
1631    });
1632
1633    Poseidon2BenchResult {
1634        name: "BB Poseidon2".into(),
1635        avg_time_ms: result.avg_time_ms,
1636        throughput_gops: result.throughput_gops,
1637    }
1638}
1639
1640pub fn bench_poseidon2_kb(config: &BenchConfig) -> Poseidon2BenchResult {
1641    let n = config.num_elements;
1642    let reps = config.ops_per_element;
1643
1644    // 16 u32s per poseidon2 state
1645    let d_states = random_u32s(n * 16, 66666)
1646        .to_device_on(bench_ctx())
1647        .unwrap();
1648
1649    // Initialize: raw u32 -> Kb field elements
1650    cuda_check(unsafe { init_poseidon2_kb(d_states.as_mut_raw_ptr(), d_states.as_ptr(), n) });
1651    sync();
1652
1653    let ops = n as u64 * reps as u64;
1654    let result = measure(config, ops, || {
1655        cuda_check(unsafe { run_poseidon2_kb(d_states.as_mut_raw_ptr(), n, reps) });
1656    });
1657
1658    Poseidon2BenchResult {
1659        name: "KB Poseidon2".into(),
1660        avg_time_ms: result.avg_time_ms,
1661        throughput_gops: result.throughput_gops,
1662    }
1663}
1664
1665pub fn run_all_benchmarks(config: &BenchConfig) {
1666    println!("=== Extension Field Benchmark ===");
1667    println!();
1668    println!("Elements: {}", config.num_elements);
1669    println!("Ops per element: {}", config.ops_per_element);
1670    println!(
1671        "Warmup: {}, Bench iters: {}",
1672        config.warmup_iters, config.bench_iters
1673    );
1674    println!();
1675
1676    // Collect all BabyBear results
1677    let fp = bench_fp(config);
1678    let fpext = bench_fpext(config);
1679    let fp5 = bench_fp5(config);
1680    let fp6 = bench_fp6(config);
1681    let fp2x3 = bench_fp2x3(config);
1682    let fp3x2 = bench_fp3x2(config);
1683
1684    // Collect all KoalaBear results
1685    let kb = bench_kb(config);
1686    let kb5 = bench_kb5(config);
1687    let kb6 = bench_kb6(config);
1688    let kb2x3 = bench_kb2x3(config);
1689    let kb3x2 = bench_kb3x2(config);
1690
1691    // Collect Goldilocks results
1692    let gl = bench_gl(config);
1693    let gl3 = bench_gl3(config);
1694
1695    // Combine all results and sort by bits (then by name for consistent ordering)
1696    let mut all_results = vec![
1697        fp.clone(),
1698        fpext,
1699        fp5,
1700        fp6,
1701        fp2x3,
1702        fp3x2,
1703        kb,
1704        kb5,
1705        kb6,
1706        kb2x3,
1707        kb3x2,
1708        gl,
1709        gl3,
1710    ];
1711    all_results.sort_by(|a, b| {
1712        a.bits
1713            .cmp(&b.bits)
1714            .then_with(|| a.field_name.cmp(&b.field_name))
1715    });
1716    print_benchmark_tables(&all_results, &fp);
1717
1718    // Poseidon2 benchmarks
1719    println!("### Poseidon2 Permutation");
1720    println!();
1721    let p2_bb = bench_poseidon2_bb(config);
1722    let p2_kb = bench_poseidon2_kb(config);
1723    println!(
1724        "| {:15} | {:>10} | {:>12} |",
1725        "Permutation", "Time (ms)", "Gops/s"
1726    );
1727    println!("|{:-<17}|{:-<12}|{:-<14}|", "", "", "");
1728    println!(
1729        "| {:15} | {:>10.3} | {:>12.1} |",
1730        p2_bb.name, p2_bb.avg_time_ms, p2_bb.throughput_gops
1731    );
1732    println!(
1733        "| {:15} | {:>10.3} | {:>12.1} |",
1734        p2_kb.name, p2_kb.avg_time_ms, p2_kb.throughput_gops
1735    );
1736    if p2_bb.throughput_gops > 0.0 {
1737        println!();
1738        println!(
1739            "KB/BB speedup: {:.2}x",
1740            p2_kb.throughput_gops / p2_bb.throughput_gops
1741        );
1742    }
1743    println!();
1744}