snarkrs_gpu_kernels/lib.rs
1//! The CUDA C kernel sources, and the include graph that assembles them.
2//!
3//! Neither NVRTC nor hipRTC has a filesystem, so a translation unit is assembled by
4//! concatenating the headers a kernel needs with the kernel itself. That is why the `.cu`
5//! files carry no `#include` of their own: the include graph lives here, in Rust, where it
6//! is checked by the compiler rather than by a preprocessor search path that only exists on
7//! the machine that happened to build it.
8//!
9//! These eight files are meant to serve both vendors, and today only `snarkrs-cuda` consumes
10//! them: hipRTC is the intended second consumer and nothing implements it yet. The
11//! arithmetic is not vendor-specific, which is what makes that possible: the CIOS Montgomery
12//! multiply, the NTT butterflies and the Pippenger buckets are the same code on NVIDIA and
13//! on AMD, and a second copy of them under a `hip/` directory would be two implementations
14//! of one algorithm that only a proof-verification failure could tell apart. One copy, one
15//! place to fix a carry.
16//!
17//! What differs per backend is the preprocessor prelude, and that stays the caller's
18//! business: [`unit_stages`] and its two twins take it as a string and paste it at the top.
19//! `snarkrs-cuda` splices its NVIDIA-only inline-PTX switch in there, which is a define no AMD
20//! compile should ever see.
21
22/// BN254 scalar field. Twin of `snarkrs_gpu_layout`'s `PackedFr`, guarded by `snarkrs-cuda`'s
23/// `tests::cuda_declares_the_same_constants`.
24pub const FR_CUH: &str = include_str!("kernels/bn254_fr.cuh");
25
26/// Stage 0: gather the A and B coefficient columns out of the CSR proving key.
27pub const GATHER_CU: &str = include_str!("kernels/gather.cu");
28
29/// Stages 1 to 3: the six transforms, iNTT then coset shift then forward NTT.
30pub const NTT_CU: &str = include_str!("kernels/ntt.cu");
31
32/// Stage 4: `H = A*B - C`, fused into the tail of the last transform where possible.
33pub const POINTWISE_CU: &str = include_str!("kernels/pointwise.cu");
34
35/// BN254 Fq, Fq2 and the G1/G2 point arithmetic, shared by the MSM and FFT units. Split
36/// out of `msm.cu` so the FFT unit does not recompile the seventeen MSM entry points;
37/// everything in it is a template or a `__forceinline__` helper, so a unit pays only for
38/// what its own kernels instantiate.
39pub const CURVE_CUH: &str = include_str!("kernels/bn254_curve.cuh");
40
41/// Stages 5 to 9: the CSR Pippenger MSM kernels.
42pub const MSM_CU: &str = include_str!("kernels/msm.cu");
43
44/// The group inverse FFT behind `ptau prepare`: the GLV ladder and the two mix kernels.
45pub const FFT_CU: &str = include_str!("kernels/fft.cu");
46
47/// Field correctness probe. Test-only, but compiled the same way as everything else so
48/// that a change which breaks compilation cannot hide behind a `cfg(test)`.
49pub const FIELD_PROBE_CU: &str = include_str!("kernels/field_probe.cu");
50
51/// Stages 0 to 4, one translation unit. They share `Fr` and nothing else, and compiling
52/// them together means one runtime-compiler invocation instead of three.
53///
54/// The order gather / ntt / pointwise is a correctness contract, not a preference:
55/// `ntt.cu` forward-declares `g16_store_h` and `pointwise.cu` defines it, so pointwise
56/// must come last. Both files say so in prose.
57pub fn unit_stages(defines: &str) -> String {
58 format!("{defines}{FR_CUH}\n{GATHER_CU}\n{NTT_CU}\n{POINTWISE_CU}")
59}
60
61/// Stages 5 to 9. Separate from the stages unit because it is by far the largest and
62/// compiling it is most of the prepare cost.
63pub fn unit_msm(defines: &str) -> String {
64 format!("{defines}{FR_CUH}\n{CURVE_CUH}\n{MSM_CU}")
65}
66
67/// The ceremony group FFT. Its own unit rather than a rider on the MSM one: a unit
68/// compiles every `extern "C" __global__` it contains whether or not the host loads it,
69/// the MSM unit already costs 113.6 s of NVRTC on a cold T4, and `ptau prepare` never
70/// launches an MSM kernel. The price is that the shared curve arithmetic is compiled
71/// twice on a box that runs both the prover and the ceremony, once per unit, which the
72/// on-disk PTX cache pays only on the first run of each.
73pub fn unit_fft(defines: &str) -> String {
74 format!("{defines}{FR_CUH}\n{CURVE_CUH}\n{FFT_CU}")
75}
76
77/// The `fr_probe` / `fr_constants` translation unit.
78pub fn unit_field_probe(defines: &str) -> String {
79 format!("{defines}{FR_CUH}\n{FIELD_PROBE_CU}")
80}