pub const FFT_CU: &str = "// The group inverse FFT behind `ptau prepare`, in CUDA C.\n//\n// Port of crates/metal/src/shaders/fft.metal, kernel for kernel. Compiled at run\n// time by NVRTC after `bn254_fr.cuh` and `bn254_curve.cuh`, in that order. Every field\n// and point operation comes from the curve header unchanged: `Xyzz<F>`, `pt_add`,\n// `pt_neg`, `pt_zero` and the `jac_*` Jacobian representation. What this file adds on\n// top of them is the GLV ladder, and it is the only formula here that a mismatch\n// against the CPU can be blamed on.\n//\n// The design arguments are in the Metal twin\'s banner and are not repeated: the dense\n// grid over the `j != 0` butterflies, the twiddle table in standard form, GLV with the\n// lattice on the host, and the two permutations that are deliberately not kernels. Two\n// places this file knowingly differs from that twin:\n//\n// 1. THE OUT-OF-PLACE, ONE-PASS-AT-A-TIME SHAPE IS INHERITED, NOT REQUIRED HERE. The\n// Metal design is forced by a macOS behaviour: a submission that keeps the GPU busy\n// while the GPU also drives a display is killed with\n// `kIOGPUCommandBufferCallbackErrorImpactingInteractivity`, so every pass has to be\n// re-runnable and the host retries. A headless Linux compute card has no analogue,\n// and the CUDA host driver (crates/cuda/src/fft.rs) already drops the retry loop and\n// the ladder budget that exist only to survive that kill. The out-of-place ping-pong\n// is kept anyway, because it is the shape proven byte-identical to the CPU at powers\n// 13 to 20 on Metal; whoever collapses it in place is relaxing a Metal constraint,\n// not a CUDA one, and buys device memory rather than time with it.\n//\n// 2. ONLY ONE WINDOW WIDTH IS COMPILED BY DEFAULT. Metal compiles four widths per group\n// in 54 ms of runtime MSL; NVRTC costs 113.6 s of source-to-PTX plus about 175 s of\n// driver JIT for the MSM unit alone on a cold T4 (crates/cuda/src/context.rs), and\n// `pt_mul_glv` fully inlines `jac_dbl` and `jac_add`, which fully inline the CIOS\n// multiply. Eight instantiations would tax every fresh machine to keep a sweep knob\n// the shipped configuration does not use, so only the shipped c = 5 (swept on Metal\n// through a whole power-16 prepare, fft.rs `FFT_WINDOW`) is always instantiated.\n//\n// 3. THE EXPERIMENT VARIANTS AT THE BOTTOM, BEHIND `G16_FFT_VARIANTS`. Any edit to this\n// file costs a full NVRTC + ptxas rebuild on the measuring machine (the PTX cache in\n// crates/cuda/src/context.rs keys on the source), so a sweep is affordable only if one\n// compile carries every candidate. The `#ifdef G16_FFT_VARIANTS` block instantiates\n// the candidate set as extra entry points in this same unit; the host defines the\n// guard when `G16_CUDA_FFT_VARIANTS=1` (crates/cuda/src/kernels.rs) and picks an entry\n// point per run via `G16_CUDA_FFT_VARIANT`, so the whole sweep is two cached units\n// (one more with `G16_CUDA_FF_PTX=1`) and every run after the first is free. A user\n// who never sets the env compiles the two shipped entry-point pairs and nothing else.\n\n#ifndef G16_FFT_CU\n#define G16_FFT_CU\n\n// ---------------------------------------------------------------------------\n// Kernel parameters. Mirrors `fft::FftParams` in crates/cuda/src/fft.rs, passed by value\n// the way `MsmParams` is; same field names as the MSL twin, `span` and not `half`, so\n// the two files stay a readable diff.\n// ---------------------------------------------------------------------------\n\nstruct FftParams {\n u32 n; // points in this block, a power of two\n u32 span; // butterflies per group in this pass, 1 << (exp - 1)\n u32 log_span; // exp - 1\n u32 tw_shift; // bits - exp, the stride into the twiddle table\n u32 gid_off; // index of the first butterfly this dispatch owns\n};\n\nstatic_assert(sizeof(FftParams) == 20, \"FftParams must be five u32 to match the host struct\");\n\n// ---------------------------------------------------------------------------\n// The GLV ladder\n// ---------------------------------------------------------------------------\n\n// Words one twiddle takes in the table: `|k1|`, `|k2|`, then the sign word. Mirrors\n// `snarkrs_gpu_layout::PackedGlv`, which a Rust test greps for this exact line.\n#define GLV_WORDS 9\n\n// Bits a magnitude is recoded over. The lattice bounds `|k1|` and `|k2|` by\n// `|n11| + |n21|`, which is 1.4795e38, so 127 bits is the true width; 128 is what the\n// window count is taken from, and every compiled `C` divides into it with `C * nw >= 128`,\n// which is what stops the top window from borrowing off the end of the value.\n// A Rust test greps this line.\n#define GLV_RECODE_BITS 128\n\n// `beta`, a primitive cube root of one in Fq, Montgomery form. Equal to\n// `ark_bn254::g1::Config::ENDO_COEFFS[0]` and to the G2 one\'s `c0`, which a Rust test\n// asserts rather than trusts. The G2 endomorphism is the same `beta` embedded in Fq2\n// with `c1 == 0`, so it is two Fq multiplies there and one here, never a full Fq2 one.\n__constant__ u32 FQ_BETA[8] = { 0x13e80b9cu, 0x3350c88eu, 0xdb5e56b9u, 0x7dce557cu, 0xb615564au, 0x6001b4b8u, 0x020217e0u, 0x2682e617u };\n\n__device__ __forceinline__ Fq fq_beta() {\n Fq b;\n#pragma unroll\n for (u32 i = 0; i < 8u; i++) {\n b.v[i] = FQ_BETA[i];\n }\n return b;\n}\n\n__device__ __forceinline__ Fq f_beta_mul(Fq x) { return fq_mul(x, fq_beta()); }\n__device__ __forceinline__ Fq2 f_beta_mul(Fq2 x) {\n Fq b = fq_beta();\n Fq2 r;\n r.c0 = fq_mul(x.c0, b);\n r.c1 = fq_mul(x.c1, b);\n return r;\n}\n\n// `phi(X, Y, Z) = (beta X, Y, Z)`. In Jacobian the point is `(X/Z^2, Y/Z^3)`, so scaling\n// X by beta scales the affine x by beta and leaves y alone, which is the endomorphism\n// exactly. Negation is on Y, so `phi` and `jac_neg` commute and the digit\'s sign can be\n// applied on either side.\ntemplate <typename F>\n__device__ __forceinline__ Jac<F> jac_endo(Jac<F> p) {\n Jac<F> r = p;\n r.x = f_beta_mul(p.x);\n return r;\n}\n\n// `width` bits at `bit_off` of the 128-bit magnitude in `v[0..4]`, reading past the top\n// as zero.\n//\n// `sc_read_bits` cannot serve: it spans all eight limbs, and here the two magnitudes are\n// consecutive halves of one array, so a top window of `|k1|` would pull in the bottom bits\n// of `|k2|` and silently produce a different scalar.\n__device__ __forceinline__ u32 glv_read_bits(const u32* v, u32 bit_off, u32 width) {\n u32 idx = bit_off >> 5;\n if (idx >= 4u) {\n return 0u;\n }\n u32 sh = bit_off & 31u;\n u64 buf = (u64)v[idx] >> sh;\n if (sh + width > 32u && idx + 1u < 4u) {\n buf |= (u64)v[idx + 1u] << (32u - sh);\n }\n return (u32)(buf & (((u64)1 << width) - (u64)1));\n}\n\n// `sc_signed_digit` over a 128-bit magnitude, digit for digit. Split out only because of\n// the read above; the recoding itself is the same one the MSM uses.\n__device__ __forceinline__ void glv_signed_digit(const u32* v, u32 i, u32 c, u32& mag, bool& neg) {\n u32 off = i * c;\n u32 b = glv_read_bits(v, off, c);\n u32 carry = (off == 0u) ? 0u : glv_read_bits(v, off - 1u, 1u);\n bool borrow = ((b >> (c - 1u)) & 1u) != 0u;\n if (borrow) {\n u32 m = (1u << c) - b - carry;\n mag = m;\n neg = true;\n if (m == 0u) {\n neg = false;\n }\n } else {\n mag = b + carry;\n neg = false;\n }\n}\n\n// `[k] p` for a twiddle the host has already put through the lattice: `k[0..4]` is\n// `|k1|`, `k[4..8]` is `|k2|`, `sign` bit 0 is `k1 < 0` and bit 1 is `k2 < 0`, and\n// `k == +-k1 + lambda * (+-k2)` for that group\'s `lambda`.\n//\n// Interleaved, not joint: one accumulator, `C` doublings a window, then up to two\n// additions off ONE table of `2^(C-1)` multiples of `p`. The `phi` side reads that same\n// table because `phi` is a homomorphism, so the register footprint is a plain fixed-window\n// ladder\'s at the same `C` while the doubling chain is half as long.\n//\n// The two signs cost nothing. `k1`\'s is folded into the base point once, before the table\n// is built, and `k2`\'s then rides as `flip`: the table is multiples of `s1 p`, so a `phi`\n// entry carries an extra factor of `s1` and the digit\'s own sign is XORed with\n// `s1 != s2`. Both are a negated Y.\n//\n// The point arrives and leaves in XYZZ, and the ladder runs in Jacobian with a = 0, for\n// the reason the Jac section of `bn254_curve.cuh` gives: the doubling chain dominates and\n// dbl-2009-l is 7 multiplies against XYZZ\'s 9. That argument is weaker here, since GLV\n// halves the doublings and adds an addition per window, but it is still the right way\n// round.\ntemplate <typename F, u32 C>\n__device__ __forceinline__ Xyzz<F> pt_mul_glv(Xyzz<F> p, const u32* k, u32 sign) {\n if (pt_is_zero(p)) {\n return pt_zero<F>();\n }\n Jac<F> q = jac_from_xyzz(p);\n if ((sign & 1u) != 0u) {\n q = jac_neg(q);\n }\n bool flip = ((sign ^ (sign >> 1u)) & 1u) != 0u;\n Jac<F> acc = jac_zero<F>();\n\n Jac<F> tbl[1u << (C - 1u)];\n tbl[0] = q;\n for (u32 i = 1u; i < (1u << (C - 1u)); i++) {\n tbl[i] = ((i & 1u) != 0u) ? jac_dbl(tbl[(i - 1u) >> 1u]) : jac_add(tbl[i - 1u], q);\n }\n\n const u32 nw = (GLV_RECODE_BITS + C - 1u) / C;\n for (u32 w = 0; w < nw; w++) {\n for (u32 j = 0; j < C; j++) {\n // A no-op while `acc` is still the identity, which `jac_dbl` exits on.\n acc = jac_dbl(acc);\n }\n u32 mag;\n bool neg;\n glv_signed_digit(k, nw - 1u - w, C, mag, neg);\n if (mag != 0u) {\n Jac<F> t = tbl[mag - 1u];\n acc = jac_add(acc, neg ? jac_neg(t) : t);\n }\n glv_signed_digit(k + 4u, nw - 1u - w, C, mag, neg);\n if (mag != 0u) {\n Jac<F> t = jac_endo(tbl[mag - 1u]);\n acc = jac_add(acc, (neg != flip) ? jac_neg(t) : t);\n }\n }\n return xyzz_from_jac(acc);\n}\n\n// `tbl[mag - 1]` as an unrolled compare-select chain over constant indices, so the array\n// is never dynamically indexed and ptxas may keep it in registers. `mag` is in\n// [1, 2^(C-1)]; the chain is 2^(C-1) - 1 predicated struct copies against a ~60-multiply\n// window body.\ntemplate <typename F, u32 C>\n__device__ __forceinline__ Jac<F> glv_tbl_select(const Jac<F>* tbl, u32 mag) {\n Jac<F> t = tbl[0];\n#pragma unroll\n for (u32 i = 1u; i < (1u << (C - 1u)); i++) {\n if (mag == i + 1u) {\n t = tbl[i];\n }\n }\n return t;\n}\n\n// `pt_mul_glv` with the window table pinned out of local memory. `tbl[mag - 1u]` above\n// is dynamically indexed, which forces the table into local memory: 1.5 KB per thread on\n// G1 and 3 KB on G2 at c = 5, so the tables of a single resident block already exceed\n// Turing\'s whole L1 and every window rereads an entry through L2. Here the build loop\n// and both lookups are fully unrolled with constant indices, so the table is eligible\n// for the register file instead. That only has a chance of fitting at a narrow width,\n// which is why only C == 3 is instantiated: four G1 entries are 96 registers of table,\n// plus 48 for the accumulator and base, under the 255 cap with room for the multiply;\n// the G2 table alone is 192 and ptxas will spill some of it, which the ptxas -v report\n// says before any run does. The price is 43 windows instead of 26, ~34 more window\n// additions per ladder. Which side of that trade the T4 takes is exactly what the\n// `c3r` variant exists to measure.\ntemplate <typename F, u32 C>\n__device__ __forceinline__ Xyzz<F> pt_mul_glv_reg(Xyzz<F> p, const u32* k, u32 sign) {\n if (pt_is_zero(p)) {\n return pt_zero<F>();\n }\n Jac<F> q = jac_from_xyzz(p);\n if ((sign & 1u) != 0u) {\n q = jac_neg(q);\n }\n bool flip = ((sign ^ (sign >> 1u)) & 1u) != 0u;\n Jac<F> acc = jac_zero<F>();\n\n Jac<F> tbl[1u << (C - 1u)];\n tbl[0] = q;\n#pragma unroll\n for (u32 i = 1u; i < (1u << (C - 1u)); i++) {\n tbl[i] = ((i & 1u) != 0u) ? jac_dbl(tbl[(i - 1u) >> 1u]) : jac_add(tbl[i - 1u], q);\n }\n\n const u32 nw = (GLV_RECODE_BITS + C - 1u) / C;\n for (u32 w = 0; w < nw; w++) {\n for (u32 j = 0; j < C; j++) {\n acc = jac_dbl(acc);\n }\n u32 mag;\n bool neg;\n glv_signed_digit(k, nw - 1u - w, C, mag, neg);\n if (mag != 0u) {\n Jac<F> t = glv_tbl_select<F, C>(tbl, mag);\n acc = jac_add(acc, neg ? jac_neg(t) : t);\n }\n glv_signed_digit(k + 4u, nw - 1u - w, C, mag, neg);\n if (mag != 0u) {\n Jac<F> t = jac_endo(glv_tbl_select<F, C>(tbl, mag));\n acc = jac_add(acc, (neg != flip) ? jac_neg(t) : t);\n }\n }\n return xyzz_from_jac(acc);\n}\n\n// Compile-time ladder choice for the kernels below. `if constexpr` so an entry point\n// instantiates only the ladder it uses; a plain `if` would compile both into every one.\ntemplate <typename F, u32 C, bool REG>\n__device__ __forceinline__ Xyzz<F> fft_ladder(Xyzz<F> p, const u32* k, u32 sign) {\n if constexpr (REG) {\n return pt_mul_glv_reg<F, C>(p, k, sign);\n } else {\n return pt_mul_glv<F, C>(p, k, sign);\n }\n}\n\n// ---------------------------------------------------------------------------\n// One `fftMix` pass\n// ---------------------------------------------------------------------------\n\n// `out[lo], out[hi] <- in[lo] + [w^j] in[hi], in[lo] - [w^j] in[hi]` for one butterfly.\n//\n// The grid is DENSE over the `j != 0` butterflies, not over all of them. The obvious\n// one-thread-per-butterfly map lets the `j == 0` lanes skip the ladder, and that skip\n// buys nothing below span 32: an idle lane rides inside a live warp for the whole pass,\n// so a span-2 pass costs every thread-slot full ladder time and half of them do no work\n// (measured on the Metal twin, whose SIMD group is the same 32 lanes: a 2^16 block\'s mix\n// pass took 31 ms whether 16,384 or 32,767 of its 32,768 slots actually laddered). So\n// thread `d` here owns ladder butterfly `d` of the `groups * (span - 1)` that exist:\n// group `d / (span - 1)`, position `1 + d % (span - 1)`. The division is against ~3,000\n// field multiplies per thread and does not show.\n//\n// The `n / 2^exp` ladderless `j == 0` butterflies ride along on the first `groups`\n// threads, two additions each next to a full ladder. At `exp == 1` there is nothing to\n// compact, every butterfly is `j == 0` and the pass is dispatched over all of them.\n//\n// Every index of `out` is still written by exactly one thread (the two families pair\n// disjoint indices) and `in` is never written, so any sub-range of a pass can be\n// dispatched again after a failure and produce the same answer.\ntemplate <typename F, u32 C, bool REG>\n__device__ __forceinline__ void fft_mix_impl(const Xyzz<F>* __restrict__ in,\n Xyzz<F>* __restrict__ out,\n const u32* __restrict__ tw,\n FftParams p,\n u32 tid) {\n u32 d = tid + p.gid_off;\n u32 groups = (p.n >> 1u) >> p.log_span;\n\n if (p.span == 1u) {\n if (d >= groups) {\n return;\n }\n Xyzz<F> u = in[d << 1u];\n Xyzz<F> t = in[(d << 1u) + 1u];\n out[d << 1u] = pt_add(u, t);\n out[(d << 1u) + 1u] = pt_add(u, pt_neg(t));\n return;\n }\n\n if (d >= (p.n >> 1u) - groups) {\n return;\n }\n u32 g = d / (p.span - 1u);\n u32 j = 1u + (d - g * (p.span - 1u));\n u32 lo = (g << (p.log_span + 1u)) + j;\n u32 hi = lo + p.span;\n\n u32 k[8];\n u32 base = (j << p.tw_shift) * GLV_WORDS;\n#pragma unroll\n for (u32 i = 0; i < 8u; i++) {\n k[i] = tw[base + i];\n }\n Xyzz<F> t = fft_ladder<F, C, REG>(in[hi], k, tw[base + 8u]);\n Xyzz<F> u = in[lo];\n out[lo] = pt_add(u, t);\n out[hi] = pt_add(u, pt_neg(t));\n\n if (d < groups) {\n u32 lo0 = d << (p.log_span + 1u);\n u32 hi0 = lo0 + p.span;\n Xyzz<F> u0 = in[lo0];\n Xyzz<F> t0 = in[hi0];\n out[lo0] = pt_add(u0, t0);\n out[hi0] = pt_add(u0, pt_neg(t0));\n }\n}\n\n// ---------------------------------------------------------------------------\n// The last mix pass, with the `1/n` scaling fused in\n// ---------------------------------------------------------------------------\n\n// `out[lo], out[hi] <- [s] in[lo] + [s w^j] in[hi], [s] in[lo] - [s w^j] in[hi]` with\n// `s = 1/n`, for pass `exp == bits` only. Exact by distributivity:\n// `[s](u + [w^j]v) == [s]u + [s w^j]v`, and the exit through `batch_to_affine` is what\n// keeps the different representative from mattering, as ever.\n//\n// The host folds `s` into this pass\'s own twiddle table (`stw[j] = s * W^j`, so\n// `stw[0] == s` covers the `u` side), which is what saves the separate scaling pass:\n// see `fft_mix_scale_impl` in the Metal twin for the ladder count.\n//\n// The last pass has exactly one group, so `lo` is `gid` and `hi` is `gid + n/2` with no\n// group arithmetic. Still out of place and read-only on `in`; each thread is two\n// ladders.\ntemplate <typename F, u32 C, bool REG>\n__device__ __forceinline__ void fft_mix_scale_impl(const Xyzz<F>* __restrict__ in,\n Xyzz<F>* __restrict__ out,\n const u32* __restrict__ stw,\n FftParams p,\n u32 tid) {\n u32 gid = tid + p.gid_off;\n u32 half_n = p.n >> 1u;\n if (gid >= half_n) {\n return;\n }\n u32 s[8];\n u32 k[8];\n u32 base = gid * GLV_WORDS;\n#pragma unroll\n for (u32 i = 0; i < 8u; i++) {\n s[i] = stw[i];\n k[i] = stw[base + i];\n }\n Xyzz<F> t = fft_ladder<F, C, REG>(in[gid + half_n], k, stw[base + 8u]);\n Xyzz<F> u = fft_ladder<F, C, REG>(in[gid], s, stw[8]);\n out[gid] = pt_add(u, t);\n out[gid + half_n] = pt_add(u, pt_neg(t));\n}\n\n// ---------------------------------------------------------------------------\n// The kernels, one pair per group per compiled variant.\n//\n// A macro rather than a template because a templated function cannot itself be\n// `extern \"C\" __global__`, the same pattern as the MSM entry points; the host selects a\n// function by name. crates/cuda/src/fft.rs\'s `VARIANTS` builds the same names from the\n// same list, and a grep test on each side keeps them aligned.\n// ---------------------------------------------------------------------------\n\n#define FFT_KERNELS(SUF, FT, PTT, C, REG) \\\n extern \"C\" __global__ void fft_mix_##SUF(const PTT* __restrict__ in, \\\n PTT* __restrict__ out, \\\n const u32* __restrict__ tw, \\\n FftParams p) { \\\n u32 tid = blockIdx.x * blockDim.x + threadIdx.x; \\\n fft_mix_impl<FT, C, REG>(in, out, tw, p, tid); \\\n } \\\n extern \"C\" __global__ void fft_mix_scale_##SUF(const PTT* __restrict__ in,\\\n PTT* __restrict__ out, \\\n const u32* __restrict__ stw,\\\n FftParams p) { \\\n u32 tid = blockIdx.x * blockDim.x + threadIdx.x; \\\n fft_mix_scale_impl<FT, C, REG>(in, out, stw, p, tid); \\\n }\n\n// The same pair under a `__launch_bounds__` cap. `MAXT` bounds the block and `MINB`\n// blocks must co-reside per SM, so ptxas is forced down to 64K / (MAXT * MINB) registers\n// per thread and spills the difference. Whether trading spill for occupancy pays on a\n// latency-heavy dependent chain is a measurement, not a judgement call.\n#define FFT_KERNELS_LB(SUF, FT, PTT, C, REG, MAXT, MINB) \\\n extern \"C\" __global__ void __launch_bounds__(MAXT, MINB) \\\n fft_mix_##SUF(const PTT* __restrict__ in, \\\n PTT* __restrict__ out, \\\n const u32* __restrict__ tw, \\\n FftParams p) { \\\n u32 tid = blockIdx.x * blockDim.x + threadIdx.x; \\\n fft_mix_impl<FT, C, REG>(in, out, tw, p, tid); \\\n } \\\n extern \"C\" __global__ void __launch_bounds__(MAXT, MINB) \\\n fft_mix_scale_##SUF(const PTT* __restrict__ in, \\\n PTT* __restrict__ out, \\\n const u32* __restrict__ stw, \\\n FftParams p) { \\\n u32 tid = blockIdx.x * blockDim.x + threadIdx.x; \\\n fft_mix_scale_impl<FT, C, REG>(in, out, stw, p, tid); \\\n }\n\nFFT_KERNELS(g1_c5, Fq, PtG1, 5u, false)\nFFT_KERNELS(g2_c5, Fq2, PtG2, 5u, false)\n\n// The experiment set, one hypothesis per pair (see the banner\'s point 3 for why they are\n// all in this one unit and how the host reaches them):\n//\n// * `c4`: the c = 5 table is what forces local-memory traffic; halving it to 768 B / 1.5 KB\n// per thread buys more than the arithmetic it costs (32 windows instead of 26, ~12\n// more window additions, 8 fewer table-build entries).\n// * `c3r`: local memory is the wrong home for the table altogether; four entries selected\n// by predicated moves keep it in registers and beat both widths, or the ~34 extra\n// additions bury the saving.\n// * `c5r128`: the default register allocation (up to 255 a thread, ~2 blocks of 128 per\n// SM) starves the SM of warps; capping at 128 registers doubles residency and the\n// added spill costs less than the latency it hides.\n#ifdef G16_FFT_VARIANTS\nFFT_KERNELS(g1_c4, Fq, PtG1, 4u, false)\nFFT_KERNELS(g2_c4, Fq2, PtG2, 4u, false)\nFFT_KERNELS(g1_c3r, Fq, PtG1, 3u, true)\nFFT_KERNELS(g2_c3r, Fq2, PtG2, 3u, true)\nFFT_KERNELS_LB(g1_c5r128, Fq, PtG1, 5u, false, 128, 4)\nFFT_KERNELS_LB(g2_c5r128, Fq2, PtG2, 5u, false, 128, 4)\n#endif // G16_FFT_VARIANTS\n\n#endif // G16_FFT_CU\n";Expand description
The group inverse FFT behind ptau prepare: the GLV ladder and the two mix kernels.