Skip to main content

MSM_CU

Constant MSM_CU 

Source
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.