Skip to main content

sp1_gpu_sys/
kernels.rs

1use std::ffi::c_void;
2
3use crate::runtime::{CudaRustError, CudaStreamHandle, KernelPtr};
4
5extern "C" {
6    // Sum kernels
7    pub fn sum_kernel_u32() -> KernelPtr;
8    pub fn sum_kernel_felt() -> KernelPtr;
9    pub fn sum_kernel_ext() -> KernelPtr;
10
11    // Tracegen kernels
12    pub fn generate_col_index() -> KernelPtr;
13    pub fn generate_start_indices() -> KernelPtr;
14    pub fn fill_buffer() -> KernelPtr;
15    pub fn count_and_add_kernel() -> KernelPtr;
16    pub fn sum_to_trace_kernel() -> KernelPtr;
17
18    // Reduce kernels
19    pub fn reduce_kernel_felt() -> KernelPtr;
20    pub fn reduce_kernel_ext() -> KernelPtr;
21
22    // JaggedMLE kernels
23    pub fn jagged_eval_kernel_chunked_felt() -> KernelPtr;
24    pub fn jagged_eval_kernel_chunked_ext() -> KernelPtr;
25
26    // JaggedInfo kernels
27    pub fn initialize_jagged_info() -> KernelPtr;
28    pub fn fix_last_variable_jagged_info() -> KernelPtr;
29
30    // Basic jagged fix last variable
31    pub fn fix_last_variable_jagged_felt() -> KernelPtr;
32    pub fn fix_last_variable_jagged_ext() -> KernelPtr;
33
34    // Fused two-variable jagged fold (base trace → twice-folded ext trace in
35    // one pass), used by the zerocheck fused first-two-rounds.
36    pub fn fix_last_two_variables_jagged_felt() -> KernelPtr;
37
38    // Fused dispatch: one launch per non-empty tier handles every
39    // Sequential chunk in a round. The launcher's per-block dispatch
40    // descriptor maps each block to its `(chunk_id, row_offset, n_rows)`.
41    pub fn zerocheck_fused_sequential_kb_32_kernel() -> KernelPtr;
42    pub fn zerocheck_fused_sequential_kb_64_kernel() -> KernelPtr;
43    pub fn zerocheck_fused_sequential_kb_128_kernel() -> KernelPtr;
44    pub fn zerocheck_fused_sequential_kb_256_kernel() -> KernelPtr;
45    pub fn zerocheck_fused_sequential_kb_512_kernel() -> KernelPtr;
46    pub fn zerocheck_fused_sequential_kb_1024_kernel() -> KernelPtr;
47    pub fn zerocheck_fused_sequential_ext_32_kernel() -> KernelPtr;
48    pub fn zerocheck_fused_sequential_ext_64_kernel() -> KernelPtr;
49    pub fn zerocheck_fused_sequential_ext_128_kernel() -> KernelPtr;
50    pub fn zerocheck_fused_sequential_ext_256_kernel() -> KernelPtr;
51    pub fn zerocheck_fused_sequential_ext_512_kernel() -> KernelPtr;
52    pub fn zerocheck_fused_sequential_ext_1024_kernel() -> KernelPtr;
53
54    // zerocheck (DAG-native): bivariate variants for the fused
55    // first-two-rounds evaluation. Eval nodes on blockIdx.z (12 non-boolean
56    // grid nodes of {0,1,2,4}^2), quadruple row consumption, output stride
57    // 12. Round 0 only — base-field trace.
58    pub fn zerocheck_fused_sequential_bivariate_kb_32_kernel() -> KernelPtr;
59    pub fn zerocheck_fused_sequential_bivariate_kb_64_kernel() -> KernelPtr;
60    pub fn zerocheck_fused_sequential_bivariate_kb_128_kernel() -> KernelPtr;
61    pub fn zerocheck_fused_sequential_bivariate_kb_256_kernel() -> KernelPtr;
62    pub fn zerocheck_fused_sequential_bivariate_kb_512_kernel() -> KernelPtr;
63    pub fn zerocheck_fused_sequential_bivariate_kb_1024_kernel() -> KernelPtr;
64
65    // zerocheck (DAG-native): ColumnTile lowering kernels.
66    pub fn zerocheck_column_tile_kb_kernel() -> KernelPtr;
67    pub fn zerocheck_column_tile_ext_kernel() -> KernelPtr;
68
69    // zerocheck (DAG-native): per-chip geq correction. One block per chip,
70    // writes 3 ext_t partials per chip (one per eval point) that the host
71    // aggregation sums into the round's totals.
72    pub fn zerocheck_geq_corrections_kernel() -> KernelPtr;
73
74    // zerocheck (DAG-native): bivariate geq correction for the fused
75    // first-two-rounds. One block per geq chip, 12 ext_t partials per chip
76    // (one per non-boolean grid node).
77    pub fn zerocheck_geq_corrections_bivariate_kernel() -> KernelPtr;
78
79    // zerocheck (DAG-native): apply `VirtualGeq::fix_last_variable(alpha)`
80    // in place to each chip's geq state. One thread per chip.
81    pub fn zerocheck_fix_geq_state_kernel() -> KernelPtr;
82
83    // zerocheck (DAG-native): aggregate per-block partials into the 3
84    // per-eval-point totals via a single-block grid-stride reduction. The
85    // host then only downloads the 3 totals instead of the full partials.
86    pub fn zerocheck_aggregate_partials_kernel() -> KernelPtr;
87
88    // zerocheck (DAG-native): strided aggregation for the fused
89    // first-two-rounds partials ([group][e] layout with `stride` slots per
90    // group). Launched with gridDim.x == stride.
91    pub fn zerocheck_aggregate_partials_strided_kernel() -> KernelPtr;
92
93    // zerocheck (DAG-native): per-chip GKR column sweep. Decoupled from the
94    // sequential constraint kernel so wide chips can parallelise the column
95    // reduction across a warp's lanes. One block per (chip, row-tile).
96    pub fn zerocheck_gkr_sweep_kb_kernel() -> KernelPtr;
97    pub fn zerocheck_gkr_sweep_ext_kernel() -> KernelPtr;
98
99    // zerocheck (DAG-native): GKR corner sweep for the fused
100    // first-two-rounds — the opening batch at the four boolean grid corners
101    // (raw rows of each quadruple, no interpolation). Output stride 4.
102    pub fn zerocheck_gkr_corner_sweep_kb_kernel() -> KernelPtr;
103
104    // zerocheck (DAG-native): per-chunk padded_row_adjustment via the
105    // bytecode interpreter at the all-zero trace. One thread per chunk;
106    // output is one ext_t per chunk, summed by chip on the host into the
107    // per-chip `padded_row_adjustment`. Tiered by MAX_REGS like
108    // `fused_sequential`.
109    pub fn zerocheck_pad_adj_32_kernel() -> KernelPtr;
110    pub fn zerocheck_pad_adj_64_kernel() -> KernelPtr;
111    pub fn zerocheck_pad_adj_128_kernel() -> KernelPtr;
112    pub fn zerocheck_pad_adj_256_kernel() -> KernelPtr;
113    pub fn zerocheck_pad_adj_512_kernel() -> KernelPtr;
114    pub fn zerocheck_pad_adj_1024_kernel() -> KernelPtr;
115
116    // JaggedMle fold-metadata: one fused multi-block kernel reads
117    // `column_heights`, writes `new_column_heights` (= `h.div_ceil(4)*2`
118    // element-wise) and `new_start_indices` (= exclusive prefix sum) — all
119    // on device, no host round-trip. Uses decoupled-lookback to handle any
120    // n_columns. See `include/jagged_assist/fold_metadata.cuh` for the
121    // caller-init contract on `block_counter`, `flags`, `scan_values`.
122    pub fn jagged_fold_metadata_kernel() -> KernelPtr;
123    pub fn jagged_fold_metadata_block_dim() -> u32;
124    pub fn jagged_fold_metadata_section_size() -> u32;
125
126    // JaggedMle chip-layouts: reads `start_indices` + `column_heights` at
127    // the sparse per-chip positions described by `ChipColumnLayoutEntry`,
128    // writes per-chip `ChipLayout[chip_idx]` + `chip_heights[chip_idx]`.
129    // One thread per chip. See `include/jagged_assist/chip_layouts.cuh`.
130    pub fn jagged_chip_layouts_kernel() -> KernelPtr;
131
132    // Jagged Zerocheck Kernels
133    pub fn jagged_constraint_poly_eval_32_koala_bear_kernel() -> KernelPtr;
134    pub fn jagged_constraint_poly_eval_64_koala_bear_kernel() -> KernelPtr;
135    pub fn jagged_constraint_poly_eval_128_koala_bear_kernel() -> KernelPtr;
136    pub fn jagged_constraint_poly_eval_256_koala_bear_kernel() -> KernelPtr;
137    pub fn jagged_constraint_poly_eval_512_koala_bear_kernel() -> KernelPtr;
138    pub fn jagged_constraint_poly_eval_1024_koala_bear_kernel() -> KernelPtr;
139
140    pub fn jagged_constraint_poly_eval_32_koala_bear_extension_kernel() -> KernelPtr;
141    pub fn jagged_constraint_poly_eval_64_koala_bear_extension_kernel() -> KernelPtr;
142    pub fn jagged_constraint_poly_eval_128_koala_bear_extension_kernel() -> KernelPtr;
143    pub fn jagged_constraint_poly_eval_256_koala_bear_extension_kernel() -> KernelPtr;
144    pub fn jagged_constraint_poly_eval_512_koala_bear_extension_kernel() -> KernelPtr;
145    pub fn jagged_constraint_poly_eval_1024_koala_bear_extension_kernel() -> KernelPtr;
146
147    // Zerocheck kernels
148    pub fn zerocheck_sum_as_poly_base_ext_kernel() -> KernelPtr;
149    pub fn zerocheck_sum_as_poly_ext_ext_kernel() -> KernelPtr;
150
151    pub fn zerocheck_fix_last_variable_and_sum_as_poly_base_ext_kernel() -> KernelPtr;
152    pub fn zerocheck_fix_last_variable_and_sum_as_poly_ext_ext_kernel() -> KernelPtr;
153
154    // Hadamard kernels
155    pub fn hadamard_sum_as_poly_base_ext_kernel() -> KernelPtr;
156    pub fn hadamard_sum_as_poly_ext_ext_kernel() -> KernelPtr;
157
158    pub fn hadamard_fix_last_variable_and_sum_as_poly_base_ext_kernel() -> KernelPtr;
159    pub fn hadamard_fix_last_variable_and_sum_as_poly_ext_ext_kernel() -> KernelPtr;
160
161    pub fn fix_last_variable_felt_ext_kernel() -> KernelPtr;
162    pub fn fix_last_variable_ext_ext_kernel() -> KernelPtr;
163    pub fn mle_fix_last_variable_koala_bear_base_base_constant_padding() -> KernelPtr;
164    pub fn mle_fix_last_variable_koala_bear_base_extension_constant_padding() -> KernelPtr;
165    pub fn mle_fix_last_variable_koala_bear_ext_ext_constant_padding() -> KernelPtr;
166
167    pub fn mle_fix_last_variable_koala_bear_ext_ext_zero_padding() -> KernelPtr;
168
169    // ******** LogUp GKR kernels - Round operations ********
170    pub fn logup_gkr_sum_as_poly_circuit_layer() -> KernelPtr;
171    pub fn logup_gkr_first_sum_as_poly_circuit_layer() -> KernelPtr;
172    pub fn logup_gkr_fix_last_variable_circuit_layer() -> KernelPtr;
173    pub fn logup_gkr_fix_last_variable_last_circuit_layer() -> KernelPtr;
174    pub fn logup_gkr_sum_as_poly_interactions_layer() -> KernelPtr;
175    pub fn logup_gkr_fix_last_variable_interactions_layer() -> KernelPtr;
176
177    // LogUp GKR kernels - First layer operations
178    pub fn logup_gkr_fix_last_variable_first_layer() -> KernelPtr;
179    pub fn logup_gkr_fix_and_sum_first_layer() -> KernelPtr;
180    pub fn logup_gkr_sum_as_poly_first_layer() -> KernelPtr;
181    pub fn logup_gkr_first_layer_transition() -> KernelPtr;
182
183    // LogUp GKR kernels - Execution operations
184    pub fn logup_gkr_circuit_transition() -> KernelPtr;
185    pub fn logup_gkr_populate_last_circuit_layer() -> KernelPtr;
186    pub fn logup_gkr_extract_output() -> KernelPtr;
187
188    // Logup GKR kernels - Fused fix and sum kernels
189    pub fn logup_gkr_fix_and_sum_circuit_layer() -> KernelPtr;
190    pub fn logup_gkr_fix_and_sum_last_circuit_layer() -> KernelPtr;
191    pub fn logup_gkr_fix_and_sum_interactions_layer() -> KernelPtr;
192
193    // Logup GKR kernels - Two-round lookahead kernels
194    pub fn logup_gkr_two_round_sum_circuit_layer() -> KernelPtr;
195    pub fn logup_gkr_two_round_sum_first_layer() -> KernelPtr;
196    pub fn logup_gkr_two_round_fix_and_sum_circuit_layer() -> KernelPtr;
197    pub fn logup_gkr_two_round_fix_and_sum_first_layer() -> KernelPtr;
198
199    // ******** Jagged sumcheck kernels ********
200    pub fn jagged_two_round_sum_as_poly() -> KernelPtr;
201    pub fn jagged_two_round_fix_and_sum() -> KernelPtr;
202    pub fn padded_hadamard_fix_and_sum() -> KernelPtr;
203
204    // Populate restrict eq
205    pub fn populate_restrict_eq_host(
206        src: *const c_void,
207        len: usize,
208        stream: CudaStreamHandle,
209    ) -> CudaRustError;
210    pub fn populate_restrict_eq_device(
211        src: *const c_void,
212        len: usize,
213        stream: CudaStreamHandle,
214    ) -> CudaRustError;
215
216    // ******** Hadamard look ahead kernels ********
217    // Look ahead kernels - FIX_TILE=32
218    pub fn round_kernel_1_32_2_2_false() -> KernelPtr;
219    pub fn round_kernel_2_32_2_2_true() -> KernelPtr;
220    pub fn round_kernel_2_32_2_2_false() -> KernelPtr;
221    pub fn round_kernel_4_32_2_2_true() -> KernelPtr;
222    pub fn round_kernel_4_32_2_2_false() -> KernelPtr;
223    pub fn round_kernel_8_32_2_2_true() -> KernelPtr;
224    pub fn round_kernel_8_32_2_2_false() -> KernelPtr;
225
226    // Look ahead kernels - FIX_TILE=64
227    pub fn round_kernel_1_64_2_2_false() -> KernelPtr;
228    pub fn round_kernel_2_64_2_2_true() -> KernelPtr;
229    pub fn round_kernel_2_64_2_2_false() -> KernelPtr;
230    pub fn round_kernel_4_64_2_2_true() -> KernelPtr;
231    pub fn round_kernel_4_64_2_2_false() -> KernelPtr;
232    pub fn round_kernel_8_64_2_2_true() -> KernelPtr;
233    pub fn round_kernel_8_64_2_2_false() -> KernelPtr;
234
235    // Look ahead kernels - NUM_POINTS=3, FIX_TILE=32
236    pub fn round_kernel_1_32_2_3_false() -> KernelPtr;
237    pub fn round_kernel_2_32_2_3_true() -> KernelPtr;
238    pub fn round_kernel_2_32_2_3_false() -> KernelPtr;
239    pub fn round_kernel_4_32_2_3_true() -> KernelPtr;
240    pub fn round_kernel_4_32_2_3_false() -> KernelPtr;
241    pub fn round_kernel_8_32_2_3_true() -> KernelPtr;
242    pub fn round_kernel_8_32_2_3_false() -> KernelPtr;
243
244    // Look ahead kernels - NUM_POINTS=3, FIX_TILE=64
245    pub fn round_kernel_1_64_2_3_false() -> KernelPtr;
246    pub fn round_kernel_1_64_4_8_false() -> KernelPtr;
247    pub fn round_kernel_2_64_2_3_true() -> KernelPtr;
248    pub fn round_kernel_2_64_2_3_false() -> KernelPtr;
249    pub fn round_kernel_4_64_2_3_true() -> KernelPtr;
250    pub fn round_kernel_4_64_2_3_false() -> KernelPtr;
251    pub fn round_kernel_4_64_4_8_true() -> KernelPtr;
252    pub fn round_kernel_4_64_4_8_false() -> KernelPtr;
253    pub fn round_kernel_8_64_2_3_true() -> KernelPtr;
254    pub fn round_kernel_8_64_2_3_false() -> KernelPtr;
255
256    // Look ahead kernels - FIX_TILE=128
257    pub fn round_kernel_1_128_4_8_false() -> KernelPtr;
258}