pub const MSM_CU: &str = "// Stages 5-9: Pippenger multi-scalar multiplication over BN254 G1 and G2, in CUDA.\n//\n// Port of crates/metal/src/shaders/msm.metal. Same algorithm, same kernel names, same\n// semantics, so a Metal-versus-CUDA number measures the hardware and the compiler and not\n// two different MSMs.\n//\n// This file is compiled at run time by NVRTC, concatenated AFTER `bn254_fr.cuh` and\n// `bn254_curve.cuh`, which supply `Fr`, the two curve fields and the point arithmetic.\n// NVRTC has no filesystem, so there is no `#include` anywhere here;\n// snarkrs_gpu_kernels::unit_msm pastes the three sources together and the `#ifndef` guards\n// make that safe.\n//\n// WHAT THIS FILE DECIDES, AND WHY\n//\n// 1. BUCKET WRITE CONFLICTS ARE REMOVED BY CONSTRUCTION, NOT BY ATOMICS.\n// There is no 32-byte atomic on any GPU, so the usual `buckets[d] += P` cannot be\n// written directly. Three approaches exist in the literature: sort the (digit, point)\n// pairs, build a sparse-matrix transpose in the cuZK style, or give every thread a\n// private bucket array and merge afterwards. Private buckets are out on arithmetic\n// grounds alone: `threads * 2^(c-1)` accumulators at 128 bytes each is hundreds of\n// megabytes and the merge costs more point additions than the accumulation it\n// parallelises. Sorting the pairs is what zkonduit\'s Metal MSM does, on the CPU.\n//\n// What is implemented here is the third shape, the cuZK one, and it is cheaper than\n// a sort because the keys are already dense small integers: count how many points\n// land in every bucket (32-bit `atomicAdd`, on plain counters, never on a field\n// element), prefix-sum the counts into row offsets, scatter each point index into its\n// bucket\'s run (again a 32-bit `atomicAdd`, this time on a cursor), and then give one\n// thread exclusive ownership of one bucket. That is a counting sort by bucket index,\n// O(n) rather than O(n log n), and after the scatter every bucket is written by\n// exactly one thread, so the accumulation needs no synchronisation of any kind. Order\n// within a bucket is not preserved, which does not matter because bucket accumulation\n// is commutative.\n//\n// If you ever find yourself wanting an atomic on a point here, something upstream has\n// been mis-ported. Every atomic in this file is `atomicAdd` on a `u32`.\n//\n// 2. IN VARIABLE WORK, SCALARS 0 AND 1 NEVER REACH A BUCKET.\n// Four of the five MSMs take the witness as scalars, and in a bit-heavy circuit over\n// 99% of those are 0 or 1. `msm_count` and `msm_scatter` test for both before doing\n// anything else: a zero scalar reads its 8 limbs, fails the test and exits, costing\n// the digit scan and nothing more; a one scalar is routed to `msm_ones_*`, which\n// performs exactly one mixed addition for it.\n//\n// Leaving the ones in Pippenger would have been correct but pathological, and this is\n// the specific trap worth naming: the signed recoding sends every scalar equal to 1 to\n// digit +1 of window 0, so bucket (0, 0) would collect *every* one-scalar, and since\n// one thread owns one bucket that single thread would serially accumulate 100k points\n// while the other quarter-million threads sat idle. Special-casing 1 is not a\n// micro-optimisation here, it is what stops the dispatch degenerating to one thread.\n//\n// 3. POINTS ACCUMULATE IN XYZZ, BASES STAY AFFINE.\n// Extended Jacobian (X, Y, ZZ, ZZZ) with x = X/ZZ, y = Y/ZZZ and ZZ^3 = ZZZ^2. Mixed\n// addition (madd-2008-s) is 7M + 2S against 7M + 4S for Jacobian madd-2007-bl, and\n// the general addition (add-2008-s) is 12M + 2S against 11M + 5S, with roughly a\n// third of the field additions. Bucket accumulation is essentially all mixed\n// additions, so this is the operation that decides the kernel.\n//\n// The alternative is affine batch addition, which is what rapidsnark and sppark use\n// on CPU and which is why our own CPU MSM leaves 11-19% on the table. It is not taken\n// here, for the reason given in the MSL twin: a batch inversion is three more\n// dispatches and two more passes over the bucket array per accumulation round. The\n// honest statement is that it is untested here, not that it loses.\n//\n// Identity is ZZ == 0. On Metal that comes free, because a freshly allocated MTLBuffer\n// is zero-filled. ON CUDA IT DOES NOT: cudaMalloc hands back whatever was there. Every\n// bucket array must be run through `msm_clear_g1` / `msm_clear_g2` and every count\n// array through `zero_u32` before first use. Getting this wrong yields a wrong proof,\n// not a crash, and it will look like a flaky failure. This is the single most likely\n// way to break the port, which is why it is said twice.\n//\n// 4. THE WINDOW REDUCTION ENDS ON THE HOST.\n// `msm_reduce_*` collapses each window\'s 2^(c-1) buckets to a single point through a\n// per-thread segment plus a block-level tree, so the host reads back `n_windows`\n// points per MSM and does only the Horner combination. The serial tail stays where\n// serial tails belong.\n//\n// CONSTANT WORK (opt in): every scalar emits in every window. Zero digits accumulate\n// actual bases in dummy rows, discarded only by the mathematical reduction. Ones are\n// ordinary digits. Fixed slice and fold geometry depends on n, not classification.\n// Fold and reduce use Metal\'s complete homogeneous formulas. This is NOT constant time:\n// atomics, addresses, run boundaries, mixed-add exceptions and host math depend on data.\n// NVRTC/PTX and arithmetic validation on NVIDIA are required for this source port.\n//\n// NVRTC HOUSE RULES OBSERVED HERE\n// * No `#include`, so no `uint2` and no `<cmath>`. The entry pair is a POD struct\n// declared below, and the two helpers that would otherwise come from a header (`min`,\n// `clz`) are written out, so nothing depends on which declarations NVRTC happens to\n// pre-inject. nvcc supplies both silently, which is exactly why the compile check\n// cannot be trusted to catch their absence.\n// * Every entry point is `extern \"C\"`. Without it NVRTC mangles the name and\n// `load_function` fails at run time with a name lookup error, not at compile time.\n// * Every conditional reduction stays a ternary. See the prelude for the measurement.\n//\n// COMPILE TIME IS A REAL COST HERE, MEASURED, AND IT LANDS ON THE FIRST PROOF.\n// NVRTC only emits PTX; the driver then runs the same ptxas at cuModuleLoad, so whatever\n// ptxas costs is paid inside `prepare`. On the T4 box, `nvcc -arch=sm_75 -cubin -Xptxas -v`\n// over this unit takes about 4m50s wall, and it is all in the Fq2 instantiations:\n// msm_reduce_g2 85.6 s, msm_merge_g2 24.8 s, msm_ones_g2 16.6 s, against 4.7 ms for\n// msm_scan. `-Xptxas -O1` only brings the total to 3m48s, so the optimisation level is\n// not the lever. The cause is that every `f_*` and `pt_*` here is `__forceinline__`, so\n// `pt_mul_small<Fq2>` inlines pt_dbl and pt_add, which inline six fq_mul each, and\n// msm_reduce_g2 becomes one enormous function; ptxas is superlinear in that.\n//\n// It is left as it is, because the mandate for this file is a faithful port and because\n// the two candidate fixes are both benchmarks rather than obvious wins: mark the Fq2\n// curve routines `__noinline__` (which would also cut the 255-register, stack-spilling\n// G2 kernels, but pays an ABI call with a 256-byte point by value), or cache the compiled\n// module on disk so only the first run ever pays. Whoever measures the prepare stage will\n// find this at the top of the profile; that is the point of writing the numbers down.\n//\n// Register footprint at -O3, same run, for whoever sizes the launches: zero_u32 and\n// msm_clear_* 6, msm_scan 17, msm_count 28, msm_scatter 30, then the point kernels at\n// 122-217 for G1 and pinned at the 255 ceiling with a small stack spill for every G2 one.\n// 255 registers means at most 8 resident warps per SM out of 32, so the G2 kernels are\n// occupancy-bound by construction and a bigger block size will not help them.\n\n#ifndef G16_MSM_CU\n#define G16_MSM_CU\n\n// Renes, Costello and Batina (2016), algorithms 7 and 9, ported from\n// metal/src/shaders/msm.metal without changing the complete formulas.\n// Local to the MSM unit; shared curve/FFT arithmetic is unchanged.\ntemplate <typename F>\nstruct Proj {\n F x;\n F y;\n F z;\n};\n\ntypedef Proj<Fq> ProjG1;\ntypedef Proj<Fq2> ProjG2;\n\n// 3b for each group: 9 on G1 (y^2 = x^3 + 3), and 3 * 3 / (9 + u) on G2, the twist\'s\n// b, in Montgomery form. Cross-checked against arkworks by the CUDA CPU field-oracle tests.\n__constant__ u32 FQ2_G2_B3_C0[8] = { 0xb62e0d6au, 0x3baa927cu, 0xd1b664fdu, 0xd71e7c52u, 0xd95d4664u, 0x03873e63u, 0x082ab8f4u, 0x0e75b5b1u };\n__constant__ u32 FQ2_G2_B3_C1[8] = { 0x7596fe35u, 0xaab7c666u, 0xbb6a27bau, 0x31d21a78u, 0x680401ffu, 0x85dd7297u, 0xdf39a7e9u, 0x03c52d6au };\n\n// 9a as four additions.\n__device__ __forceinline__ Fq f_mul_b3(Fq a) {\n Fq a2 = fq_add(a, a);\n Fq a4 = fq_add(a2, a2);\n Fq a8 = fq_add(a4, a4);\n return fq_add(a8, a);\n}\n\n__device__ __forceinline__ Fq2 f_mul_b3(Fq2 a) {\n Fq2 b3;\n for (u32 i = 0; i < FQ_LIMBS; i++) {\n b3.c0.v[i] = FQ2_G2_B3_C0[i];\n b3.c1.v[i] = FQ2_G2_B3_C1[i];\n }\n return fq2_mul(a, b3);\n}\n\n__device__ __forceinline__ Fq fq_select(bool take_a, Fq a, Fq b) {\n Fq out;\n for (u32 i = 0; i < FQ_LIMBS; i++) {\n out.v[i] = (take_a ? a.v[i] : b.v[i]);\n }\n return out;\n}\n\n__device__ __forceinline__ Fq2 fq2_select(bool take_a, Fq2 a, Fq2 b) {\n Fq2 out;\n out.c0 = fq_select(take_a, a.c0, b.c0);\n out.c1 = fq_select(take_a, a.c1, b.c1);\n return out;\n}\n\n__device__ __forceinline__ Fq f_select(bool take_a, Fq a, Fq b) { return fq_select(take_a, a, b); }\n__device__ __forceinline__ Fq2 f_select(bool take_a, Fq2 a, Fq2 b) { return fq2_select(take_a, a, b); }\n\ntemplate <typename F>\n__device__ __forceinline__ Proj<F> proj_zero() {\n Proj<F> r;\n f_set_zero(r.x);\n f_set_one(r.y);\n f_set_zero(r.z);\n return r;\n}\n\n// (X/ZZ, Y/ZZZ) = (X ZZZ : Y ZZ : ZZ ZZZ). The XYZZ identity is all zero, which maps\n// to (0 : 0 : 0), not a point, so Y becomes 1 for it; the select reads what the\n// products just consumed, so nothing stays live.\ntemplate <typename F>\n__device__ __forceinline__ Proj<F> proj_from_xyzz(Xyzz<F> p) {\n Proj<F> r;\n r.x = f_mul(p.x, p.zzz);\n r.y = f_mul(p.y, p.zz);\n r.z = f_mul(p.zz, p.zzz);\n F one;\n f_set_one(one);\n r.y = f_select(f_is_zero(p.zz), one, r.y);\n return r;\n}\n\n// (X/Z, Y/Z) = (X Z, Y Z^2, Z^2, Z^3) in XYZZ. Z = 0 gives all zero, the XYZZ identity.\ntemplate <typename F>\n__device__ __forceinline__ Xyzz<F> xyzz_from_proj(Proj<F> p) {\n Xyzz<F> r;\n r.zz = f_sqr(p.z);\n r.zzz = f_mul(r.zz, p.z);\n r.x = f_mul(p.x, p.z);\n r.y = f_mul(p.y, r.zz);\n return r;\n}\n\n// Algorithm 7 of the paper, a = 0: 12M + 2 products by 3b + 19 additions.\n// X3 = (X1 Y2 + X2 Y1)(Y1 Y2 - 3b Z1 Z2) - 3b (Y1 Z2 + Y2 Z1)(X1 Z2 + X2 Z1)\n// Y3 = (Y1 Y2 + 3b Z1 Z2)(Y1 Y2 - 3b Z1 Z2) + 9b X1 X2 (X1 Z2 + X2 Z1)\n// Z3 = (Y1 Z2 + Y2 Z1)(Y1 Y2 + 3b Z1 Z2) + 3 X1 X2 (X1 Y2 + X2 Y1)\ntemplate <typename F>\n__device__ __forceinline__ Proj<F> pt_add(Proj<F> a, Proj<F> b) {\n F t0 = f_mul(a.x, b.x);\n F t1 = f_mul(a.y, b.y);\n F t2 = f_mul(a.z, b.z);\n F t3 = f_mul(f_add(a.x, a.y), f_add(b.x, b.y));\n t3 = f_sub(t3, f_add(t0, t1)); // X1 Y2 + X2 Y1\n F t4 = f_mul(f_add(a.y, a.z), f_add(b.y, b.z));\n t4 = f_sub(t4, f_add(t1, t2)); // Y1 Z2 + Y2 Z1\n F t5 = f_mul(f_add(a.x, a.z), f_add(b.x, b.z));\n t5 = f_sub(t5, f_add(t0, t2)); // X1 Z2 + X2 Z1\n t0 = f_add(f_add(t0, t0), t0); // 3 X1 X2\n t2 = f_mul_b3(t2); // 3b Z1 Z2\n F z3 = f_add(t1, t2); // Y1 Y2 + 3b Z1 Z2\n t1 = f_sub(t1, t2); // Y1 Y2 - 3b Z1 Z2\n t5 = f_mul_b3(t5); // 3b (X1 Z2 + X2 Z1)\n Proj<F> r;\n r.x = f_sub(f_mul(t3, t1), f_mul(t4, t5));\n r.y = f_add(f_mul(z3, t1), f_mul(t5, t0));\n r.z = f_add(f_mul(t4, z3), f_mul(t0, t3));\n return r;\n}\n\n// Algorithm 9, a = 0: 6M + 2S + 1 product by 3b, and the identity maps to itself.\n// X3 = 2 X Y (Y^2 - 9b Z^2)\n// Y3 = (Y^2 - 9b Z^2)(Y^2 + 3b Z^2) + 24b Y^2 Z^2\n// Z3 = 8 Y^3 Z\ntemplate <typename F>\n__device__ __forceinline__ Proj<F> pt_dbl(Proj<F> p) {\n F t0 = f_sqr(p.y); // Y^2\n F z3 = f_add(t0, t0);\n z3 = f_add(z3, z3);\n z3 = f_add(z3, z3); // 8 Y^2\n F t1 = f_mul(p.y, p.z); // Y Z\n F t2 = f_mul_b3(f_sqr(p.z)); // 3b Z^2\n F x3 = f_mul(t2, z3); // 24b Y^2 Z^2\n F y3 = f_add(t0, t2); // Y^2 + 3b Z^2\n z3 = f_mul(t1, z3); // 8 Y^3 Z\n t1 = f_add(t2, t2);\n t2 = f_add(t1, t2); // 9b Z^2\n t0 = f_sub(t0, t2); // Y^2 - 9b Z^2\n y3 = f_add(f_mul(t0, y3), x3);\n t1 = f_mul(p.x, p.y);\n x3 = f_mul(t0, t1);\n Proj<F> r;\n r.x = f_add(x3, x3);\n r.y = y3;\n r.z = z3;\n return r;\n}\n\n// `pt_mul_small` on the complete formulas: only `k`, a bucket offset the key fixes,\n// shapes the ladder, and an identity `p` rides it to the identity.\ntemplate <typename F>\n__device__ __forceinline__ Proj<F> pt_mul_small(Proj<F> p, u32 k) {\n Proj<F> acc = proj_zero<F>();\n if (k == 0u) {\n return acc;\n }\n u32 hi = msm_hibit(k);\n for (int i = (int)hi; i >= 0; i--) {\n acc = pt_dbl(acc);\n if ((k >> (u32)i) & 1u) {\n acc = pt_add(acc, p);\n }\n }\n return acc;\n}\n\n// The reduce is written once over its point type `P`: `Xyzz<F>` with the shortcuts,\n// or `Proj<F>` under constant work. Loads and stores are XYZZ either way.\ntemplate <typename P>\nstruct PtOps;\n\ntemplate <typename F>\nstruct PtOps<Xyzz<F>> {\n static __device__ __forceinline__ Xyzz<F> zero() { return pt_zero<F>(); }\n static __device__ __forceinline__ Xyzz<F> load(Xyzz<F> p) { return p; }\n static __device__ __forceinline__ Xyzz<F> store(Xyzz<F> p) { return p; }\n};\n\ntemplate <typename F>\nstruct PtOps<Proj<F>> {\n static __device__ __forceinline__ Proj<F> zero() { return proj_zero<F>(); }\n static __device__ __forceinline__ Proj<F> load(Xyzz<F> p) { return proj_from_xyzz(p); }\n static __device__ __forceinline__ Xyzz<F> store(Proj<F> p) { return xyzz_from_proj(p); }\n};\n\n\n// ---------------------------------------------------------------------------\n// Kernel parameters. Mirrors `msm::MsmParams` on the host, which must be `#[repr(C)]`\n// with these eleven `u32` fields in this order. Passed by value: CUDA kernel arguments live\n// in a small parameter space, so there is no constant buffer to bind and no\n// `[[buffer(n)]]` index to keep in sync, which removes a whole class of Metal-side\n// mistake. The host declares it locally and can therefore implement cudarc\'s `DeviceRepr`\n// for it directly.\n// ---------------------------------------------------------------------------\n\nstruct MsmParams {\n u32 n; // scalars in this MSM\n u32 c; // window width in bits\n u32 n_windows; // ceil(255 / c)\n u32 n_buckets; // 2^(c-1)\n u32 cap; // entries per window, general count or n under constant work\n u32 scalar_off; // element offset into the scalar buffer\n u32 base_off; // element offset into the base buffer\n u32 ones_groups; // blocks in msm_ones_*\n u32 slice_len; // entries per thread in the segmented accumulation\n u32 slices; // ceil(cap / slice_len), threads per window there\n u32 dummy_rows; // constant work: zero-digit rows per window\n};\n\nstatic_assert(sizeof(MsmParams) == 44, \"MsmParams must be eleven u32 to match the host struct\");\n\nstruct FoldParams {\n u32 slots;\n u32 groups;\n u32 len;\n u32 last;\n};\nstatic_assert(sizeof(FoldParams) == 16, \"FoldParams must be four u32\");\n\n__device__ __forceinline__ u32 msm_dummy_row(MsmParams p, u32 w, u32 gid) {\n return p.n_windows * p.n_buckets + w * p.dummy_rows + (gid & (p.dummy_rows - 1u));\n}\n\n// The scatter entry, which is MSL\'s `uint2` in the twin. NVRTC has no `vector_types.h`\n// and this file may not include one, so the pair is declared here as a POD struct; the\n// host mirrors it as a `#[repr(C)]` pair of `u32`, since only the size and field order\n// matter.\n//\n// x = row = w * n_buckets + (|digit| - 1)\n// y = (point index << 1) | (digit is negative)\n//\n// The row is stored rather than recomputed because the segmented accumulation walks a\n// fixed-length slice of the entry array and has to discover where one bucket\'s run ends,\n// which it cannot do from the point index alone.\nstruct MsmEntry {\n u32 x;\n u32 y;\n};\n\nstatic_assert(sizeof(MsmEntry) == 8, \"MsmEntry must be two u32, the twin of MSL uint2\");\n\n// ---------------------------------------------------------------------------\n// Kernels.\n// ---------------------------------------------------------------------------\n\n// One thread per word. This is not an optimisation, it is a correctness requirement on\n// CUDA: cudaMalloc does not zero, and the count histogram is accumulated with atomicAdd\n// on top of whatever the allocation already held. Metal gets that zeroing free from a\n// fresh MTLBuffer; this backend must clear every counter array itself, every time, before\n// msm_count.\nextern \"C\" __global__ void zero_u32(u32* buf, u32 len) {\n u32 gid = blockIdx.x * blockDim.x + threadIdx.x;\n if (gid < len) {\n buf[gid] = 0u;\n }\n}\n\n// Montgomery limbs (layout::PackedFr) to standard limbs (layout::PackedScalar).\n// The NTT leaves H in Montgomery form on the device; Pippenger needs the integer.\n// Byte-identical twin of g16_mont_to_std in pointwise.cu, and the duplication is forced by\n// the unit split: unit_msm does not carry pointwise.cu. A change to fr_from_mont\'s callers\n// is made in both.\nextern \"C\" __global__ void fr_mont_to_std(const Fr* in, Fr* out, u32 len) {\n u32 gid = blockIdx.x * blockDim.x + threadIdx.x;\n if (gid < len) {\n out[gid] = fr_from_mont(in[gid]);\n }\n}\n\n// Stage 1 of the counting sort: how many points land in each (window, bucket).\n// One thread per scalar. `counts` must have been zeroed first, see zero_u32.\nextern \"C\" __global__ void msm_count(const u32* scalars, u32* counts, MsmParams p) {\n u32 gid = blockIdx.x * blockDim.x + threadIdx.x;\n if (gid >= p.n) {\n return;\n }\n u32 s[8];\n u32 base = (p.scalar_off + gid) * 8u;\n#pragma unroll\n for (u32 i = 0; i < 8u; i++) {\n s[i] = scalars[base + i];\n }\n // Variable work skips zeros and routes ones to msm_ones_*. Constant work\n // emits every scalar in every window, including zero digits.\n if (p.dummy_rows == 0u && (sc_is_zero(s) || sc_is_one(s))) {\n return;\n }\n for (u32 w = 0; w < p.n_windows; w++) {\n u32 mag;\n bool neg;\n sc_signed_digit(s, w, p.c, mag, neg);\n u32 row;\n if (mag != 0u) {\n row = w * p.n_buckets + (mag - 1u);\n } else if (p.dummy_rows != 0u) {\n row = msm_dummy_row(p, w, gid);\n } else {\n continue;\n }\n // A plain 32-bit atomic on a plain counter. No atomic anywhere in this file\n // touches a field element, and none ever should.\n atomicAdd(&counts[row], 1u);\n }\n}\n\n// Stage 2: exclusive prefix sum of the counts inside each window, biased by that window\'s\n// base offset in the entry array. ONE BLOCK PER WINDOW.\n//\n// SCAN_TG is the shared array size, not the launch size: the host passes up to SCAN_TG\n// threads and the kernel reads the real count from blockDim.x, so a smaller block still\n// works. Launching more than SCAN_TG threads per block would run off the end of `tmp`.\n#define SCAN_TG 256\n\nextern \"C\" __global__ void msm_scan(const u32* counts, u32* cursor, MsmParams p) {\n __shared__ u32 tmp[SCAN_TG];\n u32 w = blockIdx.x;\n u32 tid = threadIdx.x;\n u32 tcount = blockDim.x;\n\n u32 running = w * p.cap;\n u32 chunks = (p.n_buckets + tcount - 1u) / tcount;\n for (u32 ch = 0; ch < chunks; ch++) {\n u32 idx = ch * tcount + tid;\n u32 v = (idx < p.n_buckets) ? counts[w * p.n_buckets + idx] : 0u;\n tmp[tid] = v;\n __syncthreads();\n // Hillis-Steele inclusive scan. The read is separated from the write by a barrier\n // on both sides, which is what makes the in-place update safe. Every thread of the\n // block reaches every barrier: the trip count depends only on tcount, which is\n // uniform, so there is no divergent __syncthreads here.\n for (u32 d = 1; d < tcount; d <<= 1) {\n u32 x = (tid >= d) ? tmp[tid - d] : 0u;\n __syncthreads();\n tmp[tid] += x;\n __syncthreads();\n }\n if (idx < p.n_buckets) {\n cursor[w * p.n_buckets + idx] = running + tmp[tid] - v;\n }\n u32 total = tmp[tcount - 1u];\n // Not decoration: this is what stops the next chunk\'s `tmp[tid] = v` landing\n // before a slower lane has read tmp[tcount - 1].\n __syncthreads();\n running += total;\n }\n if (tid == 0u) {\n for (u32 j = 0; j < p.dummy_rows; j++) {\n u32 row = p.n_windows * p.n_buckets + w * p.dummy_rows + j;\n cursor[row] = running;\n running += counts[row];\n }\n }\n}\n\n// Stage 3: scatter each point index into its bucket\'s run. One thread per scalar. The\n// cursor is bumped with a relaxed 32-bit atomicAdd, so points land inside their run in an\n// arbitrary order, which is fine because bucket accumulation is commutative. After this\n// kernel `cursor` holds each run\'s END, and the run\'s start is `cursor - counts`.\nextern \"C\" __global__ void msm_scatter(const u32* scalars, u32* cursor, MsmEntry* entries, MsmParams p) {\n u32 gid = blockIdx.x * blockDim.x + threadIdx.x;\n if (gid >= p.n) {\n return;\n }\n u32 s[8];\n u32 base = (p.scalar_off + gid) * 8u;\n#pragma unroll\n for (u32 i = 0; i < 8u; i++) {\n s[i] = scalars[base + i];\n }\n if (p.dummy_rows == 0u && (sc_is_zero(s) || sc_is_one(s))) {\n return;\n }\n for (u32 w = 0; w < p.n_windows; w++) {\n u32 mag;\n bool neg;\n sc_signed_digit(s, w, p.c, mag, neg);\n u32 row;\n if (mag != 0u) {\n row = w * p.n_buckets + (mag - 1u);\n } else if (p.dummy_rows != 0u) {\n row = msm_dummy_row(p, w, gid);\n } else {\n continue;\n }\n u32 slot = atomicAdd(&cursor[row], 1u);\n MsmEntry e;\n e.x = row;\n e.y = (gid << 1) | (neg ? 1u : 0u);\n entries[slot] = e;\n }\n}\n\n// Stage 4, the simple form: one thread owns one bucket, exclusively, so there is nothing\n// to synchronise. Correct, and a third of the code of the segmented form below, but the\n// dispatch finishes when the fattest bucket does.\n//\n// MEASURED on the Metal side, and the numbers are why the segmented kernel exists. On\n// js_2x2_d32 at c=11 the busiest bucket holds 11,758 entries against a mean of 17.02, a\n// 691x imbalance; this kernel takes 200.99 ms for that MSM and the segmented one takes\n// 8.00 ms on identical inputs. End to end over all five MSMs at the tuned window widths,\n// this kernel gives 455.9 ms on js_16x16_d32 against 126.8 ms segmented.\n//\n// Kept so the same comparison can be re-run here rather than assumed to carry over from\n// Apple silicon. The imbalance is a property of the witness and does carry over; what it\n// costs is a property of the machine and does not.\ntemplate <typename F>\n__device__ __forceinline__ void msm_accumulate_impl(const MsmEntry* entries,\n const Aff<F>* bases,\n const u32* counts,\n const u32* cursor,\n Xyzz<F>* buckets,\n MsmParams p,\n u32 gid) {\n u32 total = p.n_windows * p.n_buckets;\n if (gid >= total) {\n return;\n }\n u32 end = cursor[gid];\n u32 cnt = counts[gid];\n u32 start = end - cnt;\n Xyzz<F> acc = pt_zero<F>();\n for (u32 i = start; i < end; i++) {\n u32 e = entries[i].y;\n Aff<F> b = bases[p.base_off + (e >> 1)];\n if ((e & 1u) != 0u) {\n b.y = f_neg(b.y);\n }\n acc = pt_madd(acc, b);\n }\n buckets[gid] = acc;\n}\n\n// A templated function cannot itself be `extern \"C\" __global__`, which is why every entry\n// point below is a thin instantiating wrapper. Same pattern as the MSL twin.\nextern \"C\" __global__ void msm_accumulate_g1(const MsmEntry* entries,\n const AffG1* bases,\n const u32* counts,\n const u32* cursor,\n PtG1* buckets,\n MsmParams p) {\n u32 gid = blockIdx.x * blockDim.x + threadIdx.x;\n msm_accumulate_impl<Fq>(entries, bases, counts, cursor, buckets, p, gid);\n}\n\nextern \"C\" __global__ void msm_accumulate_g2(const MsmEntry* entries,\n const AffG2* bases,\n const u32* counts,\n const u32* cursor,\n PtG2* buckets,\n MsmParams p) {\n u32 gid = blockIdx.x * blockDim.x + threadIdx.x;\n msm_accumulate_impl<Fq2>(entries, bases, counts, cursor, buckets, p, gid);\n}\n\n// Stage 4, the load-balanced form. This is the kernel the backend actually uses.\n//\n// THE PROBLEM IT SOLVES, measured rather than assumed. One thread per bucket makes\n// per-thread work proportional to bucket occupancy, and occupancy is not uniform: a\n// witness contains repeated values, and every copy of one value lands in the same bucket\n// of every window. On js_2x2_d32 the fattest bucket holds 11,758 entries against a mean\n// of 17. That single thread runs 11,758 serial mixed additions while the other 24,575\n// threads finish in about 17 and then wait, and it costs 201 ms of a 253 ms MSM. On CUDA\n// the shape of the problem is if anything worse: a straggler lane holds its whole warp,\n// and the warp holds its registers and its slot on the SM until it retires.\n//\n// THE FIX, which is zkmopro\'s segmented SMVP adapted to our layout. Slice the entry array\n// into fixed-length runs of `slice_len` and give one thread each slice, so per-thread work\n// is uniform BY CONSTRUCTION rather than by hoping the digits spread. Within its slice a\n// thread finds bucket boundaries by watching `entry.x` change.\n//\n// * A run that neither starts at the slice\'s first entry nor ends at its last is\n// wholly contained here, so no other thread will ever touch that bucket, and it is\n// written straight to `buckets[row]` with no synchronisation.\n// * The first and last runs may continue into the neighbouring slices, so they are\n// written to that slice\'s two spill slots, tagged with their row. That is at most\n// two spills per thread, and a slice containing a single run spills once.\n//\n// `msm_merge_*` then adds a bucket\'s spills to whatever was direct-written. It only has\n// to look at the slices its own run overlaps, which it computes from the run\'s start and\n// end, so there is no search. The worst-case merge cost is `count / slice_len` additions\n// for the fattest bucket, which turns the 11,758-step serial loop into about 180.\n//\n// `slice_len` trades the two against each other: accumulation is `slice_len` mixed\n// additions per thread and the merge is `max_count / slice_len` full additions, so the\n// balance point is near sqrt(max_count). 64 also keeps the thread count high enough to\n// fill the machine at our smaller domains, and that half of the argument is stronger here\n// than on Metal because there is more machine to fill.\n\n__constant__ u32 MSM_NO_ROW = 0xffffffffu;\n\n// Every bucket has to start at the identity, and on CUDA a fresh cudaMalloc is not zero,\n// never mind a pooled buffer still holding the previous proof\'s points. Only `zz` is\n// written: `pt_add`, `pt_madd` and the host conversion all test `zz` alone, and a bucket\n// that is direct-written is overwritten in full anyway, so clearing the other three\n// coordinates would be pure memory traffic. One thread per bucket row.\n//\n// Constant work clears every coordinate: homogeneous conversion consumes them\n// even for the identity.\n// THIS KERNEL IS MANDATORY HERE. On Metal it only matters for a reused buffer.\ntemplate <typename F>\n__device__ __forceinline__ void msm_clear_impl(Xyzz<F>* buckets, MsmParams p, u32 gid) {\n if (gid < p.n_windows * (p.n_buckets + p.dummy_rows)) {\n if (p.dummy_rows != 0u) {\n buckets[gid] = pt_zero<F>();\n return;\n }\n F z;\n f_set_zero(z);\n buckets[gid].zz = z;\n }\n}\n\nextern \"C\" __global__ void msm_clear_g1(PtG1* buckets, MsmParams p) {\n u32 gid = blockIdx.x * blockDim.x + threadIdx.x;\n msm_clear_impl<Fq>(buckets, p, gid);\n}\n\nextern \"C\" __global__ void msm_clear_g2(PtG2* buckets, MsmParams p) {\n u32 gid = blockIdx.x * blockDim.x + threadIdx.x;\n msm_clear_impl<Fq2>(buckets, p, gid);\n}\n\n// One thread per slice, gid = w * slices + k. Both spill slots are written before any\n// early return that can reach them, so the spill arrays need no pre-zeroing of their own;\n// the bucket array still does, through msm_clear_*.\ntemplate <typename F>\n__device__ __forceinline__ void msm_segmented_impl(const MsmEntry* entries,\n const Aff<F>* bases,\n const u32* cursor,\n Xyzz<F>* buckets,\n Xyzz<F>* spill_pts,\n u32* spill_rows,\n MsmParams p,\n u32 gid) {\n u32 w = gid / p.slices;\n u32 k = gid - w * p.slices;\n if (w >= p.n_windows) {\n return;\n }\n u32 base = w * p.cap;\n // The scatter left every cursor at its run\'s end, so the last bucket\'s cursor is the\n // end of the whole window region. Constant work includes the dummy rows.\n u32 last = p.dummy_rows == 0u\n ? w * p.n_buckets + p.n_buckets - 1u\n : p.n_windows * p.n_buckets + w * p.dummy_rows + p.dummy_rows - 1u;\n u32 used = cursor[last] - base;\n\n u32 head_slot = 2u * gid;\n u32 tail_slot = head_slot + 1u;\n spill_rows[head_slot] = MSM_NO_ROW;\n spill_rows[tail_slot] = MSM_NO_ROW;\n\n u32 lo = k * p.slice_len;\n if (lo >= used) {\n return;\n }\n u32 hi = msm_min(lo + p.slice_len, used);\n\n u32 cur_row = entries[base + lo].x;\n Xyzz<F> acc = pt_zero<F>();\n bool is_first_run = true;\n\n for (u32 i = lo; i < hi; i++) {\n MsmEntry e = entries[base + i];\n if (e.x != cur_row) {\n if (is_first_run) {\n spill_rows[head_slot] = cur_row;\n spill_pts[head_slot] = acc;\n is_first_run = false;\n } else {\n // Strictly interior: this thread is the only one that will ever see this\n // bucket, so the write needs no synchronisation and no spill slot.\n buckets[cur_row] = acc;\n }\n acc = pt_zero<F>();\n cur_row = e.x;\n }\n Aff<F> b = bases[p.base_off + (e.y >> 1)];\n if ((e.y & 1u) != 0u) {\n b.y = f_neg(b.y);\n }\n acc = pt_madd(acc, b);\n }\n\n // The run that ends at the slice boundary always spills, whether or not it actually\n // continues. Spilling one run that did not need to costs the merge one addition;\n // failing to spill one that did would lose it.\n if (is_first_run) {\n spill_rows[head_slot] = cur_row;\n spill_pts[head_slot] = acc;\n } else {\n spill_rows[tail_slot] = cur_row;\n spill_pts[tail_slot] = acc;\n }\n}\n\nextern \"C\" __global__ void msm_segmented_g1(const MsmEntry* entries,\n const AffG1* bases,\n const u32* cursor,\n PtG1* buckets,\n PtG1* spill_pts,\n u32* spill_rows,\n MsmParams p) {\n u32 gid = blockIdx.x * blockDim.x + threadIdx.x;\n msm_segmented_impl<Fq>(entries, bases, cursor, buckets, spill_pts, spill_rows, p, gid);\n}\n\nextern \"C\" __global__ void msm_segmented_g2(const MsmEntry* entries,\n const AffG2* bases,\n const u32* cursor,\n PtG2* buckets,\n PtG2* spill_pts,\n u32* spill_rows,\n MsmParams p) {\n u32 gid = blockIdx.x * blockDim.x + threadIdx.x;\n msm_segmented_impl<Fq2>(entries, bases, cursor, buckets, spill_pts, spill_rows, p, gid);\n}\n\n// Fold each bucket\'s spilled partials into it. One thread per bucket row, and it looks\n// only at the slices its own run overlaps, so there is no search and no atomic.\ntemplate <typename F>\n__device__ __forceinline__ void msm_merge_impl(Xyzz<F>* buckets,\n const Xyzz<F>* spill_pts,\n const u32* spill_rows,\n const u32* counts,\n const u32* cursor,\n MsmParams p,\n u32 row) {\n if (row >= p.n_windows * p.n_buckets) {\n return;\n }\n u32 cnt = counts[row];\n if (cnt == 0u) {\n return;\n }\n u32 w = row / p.n_buckets;\n u32 base = w * p.cap;\n u32 start = cursor[row] - cnt - base;\n u32 end = cursor[row] - base;\n u32 k_lo = start / p.slice_len;\n u32 k_hi = (end - 1u) / p.slice_len;\n\n Xyzz<F> acc = buckets[row];\n for (u32 k = k_lo; k <= k_hi; k++) {\n u32 slot = 2u * (w * p.slices + k);\n if (spill_rows[slot] == row) {\n acc = pt_add(acc, spill_pts[slot]);\n }\n if (spill_rows[slot + 1u] == row) {\n acc = pt_add(acc, spill_pts[slot + 1u]);\n }\n }\n buckets[row] = acc;\n}\n\nextern \"C\" __global__ void msm_merge_g1(PtG1* buckets,\n const PtG1* spill_pts,\n const u32* spill_rows,\n const u32* counts,\n const u32* cursor,\n MsmParams p) {\n u32 row = blockIdx.x * blockDim.x + threadIdx.x;\n msm_merge_impl<Fq>(buckets, spill_pts, spill_rows, counts, cursor, p, row);\n}\n\nextern \"C\" __global__ void msm_merge_g2(PtG2* buckets,\n const PtG2* spill_pts,\n const u32* spill_rows,\n const u32* counts,\n const u32* cursor,\n MsmParams p) {\n u32 row = blockIdx.x * blockDim.x + threadIdx.x;\n msm_merge_impl<Fq2>(buckets, spill_pts, spill_rows, counts, cursor, p, row);\n}\n\n// Fixed spill tree, ported from Metal msm_fold_impl. No occupancy-sized merge.\ntemplate <typename F>\n__device__ __forceinline__ void msm_fold_impl(const Xyzz<F>* in_pts,\n const u32* in_rows,\n Xyzz<F>* out_pts,\n u32* out_rows,\n Xyzz<F>* buckets,\n MsmParams p,\n FoldParams f,\n u32 gid) {\n u32 m_in = f.slots;\n u32 groups = f.groups;\n u32 len = f.len;\n bool last = f.last != 0u;\n u32 w = gid / groups;\n u32 g = gid - w * groups;\n if (w >= p.n_windows) {\n return;\n }\n u32 head_slot = 2u * gid;\n u32 tail_slot = head_slot + 1u;\n if (!last) {\n out_rows[head_slot] = MSM_NO_ROW;\n out_rows[tail_slot] = MSM_NO_ROW;\n }\n\n u32 lo = g * len;\n u32 hi = msm_min(lo + len, m_in);\n u32 cur = MSM_NO_ROW;\n bool is_first_run = true;\n Proj<F> acc = proj_zero<F>();\n // `len` iterations in every group, the window\'s short last one included, so a\n // thread\'s additions do not depend on where it sits either.\n for (u32 i = lo; i < lo + len; i++) {\n u32 j = w * m_in + msm_min(i, hi - 1u);\n u32 r = (i < hi) ? in_rows[j] : MSM_NO_ROW;\n Xyzz<F> acc_out = xyzz_from_proj(acc);\n bool empty = r == MSM_NO_ROW;\n if (!empty && r != cur) {\n if (cur != MSM_NO_ROW) {\n if (is_first_run && !last) {\n out_pts[head_slot] = acc_out;\n out_rows[head_slot] = cur;\n } else {\n buckets[cur] = acc_out;\n }\n is_first_run = false;\n }\n acc = proj_zero<F>();\n cur = r;\n }\n Proj<F> q = proj_from_xyzz(in_pts[j]);\n if (empty) {\n q = proj_zero<F>();\n }\n acc = pt_add(acc, q);\n }\n\n if (cur == MSM_NO_ROW) {\n return;\n }\n Xyzz<F> acc_out = xyzz_from_proj(acc);\n if (last) {\n buckets[cur] = acc_out;\n } else if (is_first_run) {\n out_pts[head_slot] = acc_out;\n out_rows[head_slot] = cur;\n } else {\n out_pts[tail_slot] = acc_out;\n out_rows[tail_slot] = cur;\n }\n}\n\nextern \"C\" __global__ void msm_fold_g1(const PtG1* in_pts,\n const u32* in_rows,\n PtG1* out_pts,\n u32* out_rows,\n PtG1* buckets,\n MsmParams p,\n FoldParams f) {\n u32 gid = blockIdx.x * blockDim.x + threadIdx.x;\n msm_fold_impl<Fq>(in_pts, in_rows, out_pts, out_rows, buckets, p, f, gid);\n}\n\nextern \"C\" __global__ void msm_fold_g2(const PtG2* in_pts,\n const u32* in_rows,\n PtG2* out_pts,\n u32* out_rows,\n PtG2* buckets,\n MsmParams p,\n FoldParams f) {\n u32 gid = blockIdx.x * blockDim.x + threadIdx.x;\n msm_fold_impl<Fq2>(in_pts, in_rows, out_pts, out_rows, buckets, p, f, gid);\n}\n\n// Stage 5: collapse one window\'s 2^(c-1) buckets to one point. ONE BLOCK PER WINDOW.\n//\n// The window sum is sum_j (j+1) B_j. Split the buckets into one segment per thread, at\n// [lo, hi). Inside a segment the reverse running sum gives\n// P = sum_j (j - lo + 1) B_j and Q = sum_j B_j in two additions per bucket, and the\n// segment contributes P + lo * Q. The per-thread results are then tree-reduced in shared\n// memory, so the host reads back one point per window and does nothing but the Horner\n// combination.\n//\n// REDUCE_TG is the shared array size and therefore the largest block the host may launch\n// at these kernels. 64 rather than 128 is an occupancy choice: at 64 the G2 array is\n// 64 * 256 = 16 KB. An sm_75 SM has 64 KB of shared memory, so 16 KB per block leaves\n// room for four resident blocks; 128 threads would be 32 KB and halve that. The MSL twin\n// picks the same 64 off the same reasoning applied to a 32 KB threadgroup budget, which\n// is a coincidence of the numbers rather than a shared derivation.\n#define REDUCE_TG 64\n\ntemplate <typename F, typename P>\n__device__ __forceinline__ void msm_reduce_impl(const Xyzz<F>* buckets,\n Xyzz<F>* window_sums,\n MsmParams p,\n P* shared,\n u32 w,\n u32 tid,\n u32 tcount) {\n u32 seg_len = (p.n_buckets + tcount - 1u) / tcount;\n u32 lo = tid * seg_len;\n u32 hi = msm_min(lo + seg_len, p.n_buckets);\n\n P mine = PtOps<P>::zero();\n if (lo < hi) {\n P run = PtOps<P>::zero();\n P tot = PtOps<P>::zero();\n for (u32 j = hi; j > lo; j--) {\n run = pt_add(run, PtOps<P>::load(buckets[w * p.n_buckets + (j - 1u)]));\n tot = pt_add(tot, run);\n }\n mine = pt_add(tot, pt_mul_small(run, lo));\n }\n shared[tid] = mine;\n\n // Every thread of the block runs this loop the same number of times, so no\n // __syncthreads() is reached by only part of the block. A thread whose segment was\n // empty still contributes its identity and still has to arrive at the barriers.\n for (u32 s = 1; s < tcount; s <<= 1) {\n __syncthreads();\n if ((tid & ((s << 1) - 1u)) == 0u && tid + s < tcount) {\n shared[tid] = pt_add(shared[tid], shared[tid + s]);\n }\n }\n __syncthreads();\n if (tid == 0u) {\n window_sums[w] = PtOps<P>::store(shared[0]);\n }\n}\n\nextern \"C\" __global__ void msm_reduce_g1(const PtG1* buckets, PtG1* window_sums, MsmParams p) {\n __shared__ union {\n PtG1 variable[REDUCE_TG];\n ProjG1 constant[REDUCE_TG];\n } shared;\n if (p.dummy_rows != 0u) {\n msm_reduce_impl<Fq, ProjG1>(buckets, window_sums, p, shared.constant, blockIdx.x, threadIdx.x, blockDim.x);\n } else {\n msm_reduce_impl<Fq, PtG1>(buckets, window_sums, p, shared.variable, blockIdx.x, threadIdx.x, blockDim.x);\n }\n}\n\nextern \"C\" __global__ void msm_reduce_g2(const PtG2* buckets, PtG2* window_sums, MsmParams p) {\n __shared__ union {\n PtG2 variable[REDUCE_TG];\n ProjG2 constant[REDUCE_TG];\n } shared;\n if (p.dummy_rows != 0u) {\n msm_reduce_impl<Fq2, ProjG2>(buckets, window_sums, p, shared.constant, blockIdx.x, threadIdx.x, blockDim.x);\n } else {\n msm_reduce_impl<Fq2, PtG2>(buckets, window_sums, p, shared.variable, blockIdx.x, threadIdx.x, blockDim.x);\n }\n}\n\n// The scalar-of-1 path: sum the bases whose scalar is exactly 1, one mixed addition each.\n// Strided so consecutive lanes read consecutive scalars, which on CUDA also means a warp\'s\n// 8-limb reads coalesce into contiguous 1 KB transactions rather than 32 scattered ones.\n// Then the same block-level tree as the reduction, so the host adds only `ones_groups`\n// points.\n//\n// In variable work this kernel is required: count and scatter drop every scalar\n// equal to 1, so without it those terms are simply missing from the proof.\ntemplate <typename F>\n__device__ __forceinline__ void msm_ones_impl(const u32* scalars,\n const Aff<F>* bases,\n Xyzz<F>* out,\n MsmParams p,\n Xyzz<F>* shared,\n u32 g,\n u32 tid,\n u32 tcount) {\n u32 stride = p.ones_groups * tcount;\n Xyzz<F> acc = pt_zero<F>();\n for (u32 i = g * tcount + tid; i < p.n; i += stride) {\n u32 s[8];\n u32 base = (p.scalar_off + i) * 8u;\n#pragma unroll\n for (u32 k = 0; k < 8u; k++) {\n s[k] = scalars[base + k];\n }\n if (!sc_is_one(s)) {\n continue;\n }\n acc = pt_madd(acc, bases[p.base_off + i]);\n }\n shared[tid] = acc;\n for (u32 s = 1; s < tcount; s <<= 1) {\n __syncthreads();\n if ((tid & ((s << 1) - 1u)) == 0u && tid + s < tcount) {\n shared[tid] = pt_add(shared[tid], shared[tid + s]);\n }\n }\n __syncthreads();\n if (tid == 0u) {\n out[g] = shared[0];\n }\n}\n\nextern \"C\" __global__ void msm_ones_g1(const u32* scalars, const AffG1* bases, PtG1* out, MsmParams p) {\n __shared__ PtG1 shared[REDUCE_TG];\n msm_ones_impl<Fq>(scalars, bases, out, p, shared, blockIdx.x, threadIdx.x, blockDim.x);\n}\n\nextern \"C\" __global__ void msm_ones_g2(const u32* scalars, const AffG2* bases, PtG2* out, MsmParams p) {\n __shared__ PtG2 shared[REDUCE_TG];\n msm_ones_impl<Fq2>(scalars, bases, out, p, shared, blockIdx.x, threadIdx.x, blockDim.x);\n}\n\n#endif // G16_MSM_CU\n";Expand description
Stages 5 to 9: the CSR Pippenger MSM kernels.