diff --git a/Cargo.lock b/Cargo.lock index 7c13591..f119e7c 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -267,7 +267,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "39cab71617ae0d63f51a36d69f866391735b51691dbda63cf6f96d042b63efeb" dependencies = [ "libc", - "windows-sys", + "windows-sys 0.61.2", ] [[package]] @@ -320,6 +320,15 @@ version = "0.5.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "fc0fef456e4baa96da950455cd02c081ca953b141298e41db3fc7e36b1da849c" +[[package]] +name = "home" +version = "0.5.12" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "cc627f471c528ff0c4a49e1d5e60450c8f6461dd6d10ba9dcd3a61d3dff7728d" +dependencies = [ + "windows-sys 0.61.2", +] + [[package]] name = "is-terminal" version = "0.4.17" @@ -328,7 +337,7 @@ checksum = "3640c1c38b8e4e43584d8df18be5fc6b0aa314ce6ebf51b53313d4306cca8e46" dependencies = [ "hermit-abi", "libc", - "windows-sys", + "windows-sys 0.61.2", ] [[package]] @@ -392,6 +401,12 @@ version = "0.2.16" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "b6d2cec3eae94f9f509c767b45932f1ada8350c4bdb85af2fcab4a3c14807981" +[[package]] +name = "linux-raw-sys" +version = "0.4.15" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d26c52dbd32dccf2d10cac7725f8eae5296885fb5703b261f7d0a0739ec807ab" + [[package]] name = "linux-raw-sys" version = "0.12.1" @@ -466,6 +481,7 @@ dependencies = [ "rand", "rayon", "serde", + "sppark", "tracing", "tracing-subscriber", "tracing-texray", @@ -477,7 +493,7 @@ version = "0.50.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "7957b9740744892f114936ab4a57b3f487491bbeafaf8083688b16841a4240e5" dependencies = [ - "windows-sys", + "windows-sys 0.61.2", ] [[package]] @@ -968,6 +984,19 @@ version = "0.8.10" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "dc897dd8d9e8bd1ed8cdad82b5966c3e0ecae09fb1907d58efaa013543185d0a" +[[package]] +name = "rustix" +version = "0.38.44" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "fdb5bc1ae2baa591800df16c9ca78619bf65c0488b41b96ccec5d11220d8c154" +dependencies = [ + "bitflags", + "errno", + "libc", + "linux-raw-sys 0.4.15", + "windows-sys 0.59.0", +] + [[package]] name = "rustix" version = "1.1.4" @@ -977,8 +1006,8 @@ dependencies = [ "bitflags", "errno", "libc", - "linux-raw-sys", - "windows-sys", + "linux-raw-sys 0.12.1", + "windows-sys 0.61.2", ] [[package]] @@ -1081,6 +1110,15 @@ dependencies = [ "lock_api", ] +[[package]] +name = "sppark" +version = "0.1.15" +source = "git+https://github.com/argumentcomputer/sppark?rev=e10e107673aa22861f0f8b9758fc62169ab919ae#e10e107673aa22861f0f8b9758fc62169ab919ae" +dependencies = [ + "cc", + "which", +] + [[package]] name = "syn" version = "2.0.117" @@ -1098,8 +1136,8 @@ version = "0.4.4" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "230a1b821ccbd75b185820a1f1ff7b14d21da1e442e22c0863ea5f08771a8874" dependencies = [ - "rustix", - "windows-sys", + "rustix 1.1.4", + "windows-sys 0.61.2", ] [[package]] @@ -1312,13 +1350,25 @@ dependencies = [ "wasm-bindgen", ] +[[package]] +name = "which" +version = "4.4.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "87ba24419a2078cd2b0f2ede2691b6c66d8e47836da3b6db8265ebad47afbfc7" +dependencies = [ + "either", + "home", + "once_cell", + "rustix 0.38.44", +] + [[package]] name = "winapi-util" version = "0.1.11" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "c2a7b1c03c876122aa43f3020e6c3c3ee5c05081c9a00739faf7503aeba10d22" dependencies = [ - "windows-sys", + "windows-sys 0.61.2", ] [[package]] @@ -1336,6 +1386,15 @@ dependencies = [ "windows-link", ] +[[package]] +name = "windows-sys" +version = "0.59.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1e38bc4d79ed67fd075bcc251a1c39b32a1776bbe92e5bef1f0bf1f8c531853b" +dependencies = [ + "windows-targets", +] + [[package]] name = "windows-sys" version = "0.61.2" @@ -1345,6 +1404,70 @@ dependencies = [ "windows-link", ] +[[package]] +name = "windows-targets" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9b724f72796e036ab90c1021d4780d4d3d648aca59e491e6b98e725b84e99973" +dependencies = [ + "windows_aarch64_gnullvm", + "windows_aarch64_msvc", + "windows_i686_gnu", + "windows_i686_gnullvm", + "windows_i686_msvc", + "windows_x86_64_gnu", + "windows_x86_64_gnullvm", + "windows_x86_64_msvc", +] + +[[package]] +name = "windows_aarch64_gnullvm" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "32a4622180e7a0ec044bb555404c800bc9fd9ec262ec147edd5989ccd0c02cd3" + +[[package]] +name = "windows_aarch64_msvc" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "09ec2a7bb152e2252b53fa7803150007879548bc709c039df7627cabbd05d469" + +[[package]] +name = "windows_i686_gnu" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8e9b5ad5ab802e97eb8e295ac6720e509ee4c243f69d781394014ebfe8bbfa0b" + +[[package]] +name = "windows_i686_gnullvm" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0eee52d38c090b3caa76c563b86c3a4bd71ef1a819287c19d586d7334ae8ed66" + +[[package]] +name = "windows_i686_msvc" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "240948bc05c5e7c6dabba28bf89d89ffce3e303022809e73deaefe4f6ec56c66" + +[[package]] +name = "windows_x86_64_gnu" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "147a5c80aabfbf0c7d901cb5895d1de30ef2907eb21fbbab29ca94c5b08b1a78" + +[[package]] +name = "windows_x86_64_gnullvm" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "24d5b23dc417412679681396f2b49f3de8c1473deb516bd34410872eff51ed0d" + +[[package]] +name = "windows_x86_64_msvc" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "589f6da84c646204747d1270a2a5661ea66ed1cced2631d546fdfb155959f9ec" + [[package]] name = "zerocopy" version = "0.8.48" diff --git a/Cargo.toml b/Cargo.toml index e3aac95..7e042d6 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -5,6 +5,7 @@ edition = "2024" authors = ["Argument Engineering "] license = "MIT OR Apache-2.0" rust-version = "1.98" +links = "multi_stark_cuda" # See more keys and their definitions at https://doc.rust-lang.org/cargo/reference/manifest.html @@ -32,6 +33,14 @@ p3-merkle-tree = { git = "https://github.com/Plonky3/Plonky3", rev = "3152b14a89 p3-symmetric = { git = "https://github.com/Plonky3/Plonky3", rev = "3152b14a89067c83775a8076cc262ffc48a1fd7c" } p3-util = { git = "https://github.com/Plonky3/Plonky3", rev = "3152b14a89067c83775a8076cc262ffc48a1fd7c" } +# The CUDA transform engine. The fork supplies column batching, borrowed +# streams and a runtime compatible with the Lean executable's C++ linkage. +[dependencies.sppark] +git = "https://github.com/argumentcomputer/sppark" +rev = "e10e107673aa22861f0f8b9758fc62169ab919ae" +optional = true +features = ["cuda"] + [dev-dependencies] criterion = "0.5" p3-baby-bear = { git = "https://github.com/Plonky3/Plonky3", rev = "3152b14a89067c83775a8076cc262ffc48a1fd7c" } @@ -45,10 +54,9 @@ harness = false [features] parallel = ["p3-maybe-rayon/parallel"] -# Use the first-party CUDA Goldilocks DFT/LDE backend. Enabling this feature -# requires a CUDA toolkit at build time and an NVIDIA GPU at runtime; the -# default CPU build never invokes nvcc or links the CUDA runtime. -cuda = ["dep:itertools", "dep:rayon"] +# CUDA requires the toolkit at build time and an NVIDIA GPU at runtime. +# CPU builds neither compile sppark nor link CUDA. +cuda = ["dep:itertools", "dep:rayon", "dep:sppark"] # Similar to `release`, but preserves debug info [profile.dev-ci] diff --git a/README.md b/README.md index 66d2482..641add6 100644 --- a/README.md +++ b/README.md @@ -17,7 +17,7 @@ lookup arguments for shared state. a batteries-included Goldilocks/Blake3 instantiation is provided - **Serialization** — `Proof::to_bytes` / `Proof::from_bytes` via bincode - **Parallel proving** — opt-in via the `parallel` feature flag -- **CUDA transforms** — opt-in first-party Goldilocks DFT/LDE backend via the +- **CUDA transforms** — opt-in sppark Goldilocks DFT/LDE backend via the `cuda` feature; normal builds remain independent of CUDA ## Reference configuration @@ -80,7 +80,7 @@ are enabled by default via `.cargo/config.toml`. ## CUDA acceleration The optional `cuda` feature routes the production Goldilocks configuration's -PCS and quotient transforms through first-party CUDA kernels. CUDA and CPU +PCS and quotient transforms through sppark. CUDA and CPU proofs generated by the same revision are byte-identical and use the same CPU verifier. Canonical field serialization changes newly generated CPU proof bytes relative to revisions before this backend; previously generated proofs remain @@ -104,6 +104,29 @@ hardware and downstream Ix measurements live in See [cuda/README.md](cuda/README.md) for architecture, build settings, the NVIDIA correctness harness, platform limitations, and benchmark commands. +### sppark transforms + +Every CUDA build uses [sppark](https://github.com/argumentcomputer/sppark)'s +Goldilocks NTT, pinned to the fork's `dev` branch. Its borrowed streams keep +panel allocation, transforms and freeing on the caller's stream. Tiny generic +host matrices retain the CPU DFT. CPU builds do not depend on sppark or CUDA. + +An immutable per-device plan fixes each transform's panel size, launch groups +and coset powers. Admission and execution share the plan. Dimensions must fit +the compiled domain (currently 2^28 rows), checked byte arithmetic and at least +one column within the panel budget; invalid shapes fail before allocation. + +`MULTI_STARK_SPPARK_PANEL_BYTES` caps the transform scratch, 4 GiB by default; +a budget smaller than one column is a configuration error. Columns are +transformed in launch groups sized to the device's L2 cache. Both are captured +when the device DFT is constructed. The fork's `SPPARK_NO_CXX_RUNTIME` mode +aborts on CUDA errors with diagnostics. + +```sh +cargo test --release --features parallel,cuda --lib cuda::sppark::tests:: +MULTI_STARK_CUDA_BENCH_SHAPES="20,533,2" cargo run --release --features parallel,cuda --example cuda_resident_lde_bench +``` + ## License MIT or Apache-2.0 diff --git a/build.rs b/build.rs index d3641c7..514b734 100644 --- a/build.rs +++ b/build.rs @@ -11,6 +11,14 @@ use std::process::Command; fn main() { println!("cargo:rerun-if-changed=build.rs"); println!("cargo:rerun-if-changed=cuda/kernels.cu"); + println!("cargo:rerun-if-changed=cuda/goldilocks.cuh"); + println!("cargo:rerun-if-changed=cuda/sppark_ntt.cu"); + println!("cargo:rerun-if-changed=cuda/ntt.cuh"); + if let Some(root) = env::var_os("DEP_SPPARK_ROOT") { + let root = PathBuf::from(root); + println!("cargo:rerun-if-changed={}", root.join("ntt").display()); + println!("cargo:rerun-if-changed={}", root.join("util").display()); + } println!("cargo:rerun-if-env-changed=NVCC"); println!("cargo:rerun-if-env-changed=CUDA_HOME"); println!("cargo:rerun-if-env-changed=CUDA_PATH"); @@ -20,6 +28,9 @@ fn main() { return; } + let include = PathBuf::from(env::var_os("CARGO_MANIFEST_DIR").unwrap()).join("cuda"); + println!("cargo:include={}", include.display()); + assert_eq!( env::var("CARGO_CFG_TARGET_OS").as_deref(), Ok("linux"), @@ -36,6 +47,57 @@ fn main() { let library = out_dir.join("libmulti_stark_cuda.a"); let architectures = cuda_architectures(&nvcc); + // sppark's units are compiled first, on their own, without the + // `_GNU_SOURCE` undefinition below: its runtime includes libstdc++'s + // ``, whose GNU-only pthread functions that flag would hide. Their + // objects then join the archive. + let mut sppark_objects = Vec::new(); + { + let root = PathBuf::from( + env::var_os("DEP_SPPARK_ROOT").expect("sppark's build script exports DEP_SPPARK_ROOT"), + ); + for (source, object) in [ + (PathBuf::from("cuda/sppark_ntt.cu"), "sppark_ntt.o"), + (root.join("util/all_gpus.cpp"), "sppark_all_gpus.o"), + ] { + let object = out_dir.join(object); + let mut compile = Command::new(&nvcc); + compile + .arg("-c") + .arg("--std=c++17") + .arg("--cudart=static") + .arg("--default-stream=per-thread") + .arg("-O3") + .arg("-lineinfo") + .arg("--compiler-options=-fPIC") + .arg(format!("-I{}", root.display())) + .arg("-Icuda") + .arg("-DFEATURE_GOLDILOCKS") + // The fork's runtime without exceptions or its thread pool: + // no C++ runtime library symbols, so the archive links into + // the Lean executable, which carries libc++ rather than + // libstdc++. + .arg("-DSPPARK_NO_CXX_RUNTIME") + .arg("-o") + .arg(&object) + .arg(&source); + for architecture in &architectures { + compile.arg(format!( + "-gencode=arch=compute_{architecture},code=sm_{architecture}" + )); + } + let status = compile + .status() + .unwrap_or_else(|error| panic!("failed to execute {:?}: {error}", nvcc)); + assert!( + status.success(), + "nvcc failed on {} with status {status}", + source.display() + ); + sppark_objects.push(object); + } + } + let mut command = Command::new(&nvcc); command .arg("--lib") @@ -45,9 +107,18 @@ fn main() { .arg("-O3") .arg("-lineinfo") .arg("--compiler-options=-fPIC") + // Host code is linked by whatever toolchain links the final binary, + // and Lean's bundled clang links against a sysroot older than + // glibc 2.38. g++ predefines _GNU_SOURCE, under which glibc 2.38+ + // renames strtol and friends to their C23 variants (__isoc23_*), + // which that sysroot lacks. The POSIX and default feature sets keep + // everything the kernels' host code uses (clock_gettime, pthreads) + // under the plain names. + .arg("--compiler-options=-U_GNU_SOURCE,-D_DEFAULT_SOURCE,-D_POSIX_C_SOURCE=200809L") .arg("-o") .arg(&library) .arg("cuda/kernels.cu"); + command.args(&sppark_objects); for architecture in &architectures { command.arg(format!( diff --git a/cuda/README.md b/cuda/README.md index 8d203f1..4fb5f93 100644 --- a/cuda/README.md +++ b/cuda/README.md @@ -1,9 +1,9 @@ # CUDA backend -This directory contains multi-stark's first-party CUDA implementation of -Goldilocks arithmetic, DFTs, coset low-degree extensions, BLAKE3 Merkle -commitments, lookup construction, quotient evaluation, and FRI proving. It -does not use ICICLE or copy code from it. +This directory contains multi-stark's CUDA prover: Goldilocks arithmetic, +BLAKE3 Merkle commitments, lookup construction, quotient evaluation, and FRI +proving. All GPU DFTs and coset low-degree extensions use the pinned sppark +fork through `sppark_ntt.cu`. The `cuda` Cargo feature selects `CudaDft` for the production `GoldilocksBlake3Config`. BabyBear tests remain on their CPU DFT. Without the @@ -15,11 +15,13 @@ independent of CUDA. The backend preserves Plonky3's public PCS interfaces while keeping the hot prover pipeline device-resident: -1. Rust validates dimensions and builds exact P3-compatible twiddle tables. -2. Trace matrices are uploaded and transformed with fused radix-4/radix-8 - DIF kernels into resident coset LDEs. +1. Rust validates dimensions and caches immutable transform plans, including + panel scratch, batch groups and coset powers. Admission uses the same plan. +2. Trace matrices are uploaded and transformed through sppark on the caller's + stream into resident coset LDEs. 3. First-party BLAKE3 kernels commit mixed-height matrices without copying - LDEs back to the host. + LDEs back to the host. Rows of at most one BLAKE3 chunk (1024 bytes) hash + with one thread per row; longer rows take a warp per row. 4. Lookup traces and quotient LDEs are constructed from resident commitments. 5. Batched openings, reductions, FRI folding, and Merkle authentication remain resident; only protocol-visible openings and proofs return to Rust. @@ -97,8 +99,8 @@ Ix measurements are documented in [`docs/cuda-benchmarks.md`](../docs/cuda-bench - CUDA outputs are canonical field representatives. - Rust checks power-of-two sizes, two-adicity, integer overflow, and buffer lengths before each synchronous FFI call. -- The CUDA source is licensed under the repository's MIT/Apache-2.0 terms and - depends only on the CUDA runtime/toolkit when enabled. +- First-party CUDA sources use the repository's MIT/Apache-2.0 terms. The + sppark dependency retains its own license; see the pinned fork's `LICENSE`. - CUDA affects prover execution only. Fields, BLAKE3 hashing, transcripts, proof format, and the CPU verifier are unchanged. Canonical representation changed newly generated proof bytes relative to pre-CUDA revisions as diff --git a/cuda/goldilocks.cuh b/cuda/goldilocks.cuh new file mode 100644 index 0000000..abebe64 --- /dev/null +++ b/cuda/goldilocks.cuh @@ -0,0 +1,80 @@ +// SPDX-License-Identifier: MIT OR Apache-2.0 +#pragma once + +#include +#include + +namespace multi_stark_cuda { + +constexpr uint64_t GOLDILOCKS_P = 0xffffffff00000001ULL; +constexpr uint64_t GOLDILOCKS_EPSILON = 0x00000000ffffffffULL; + +__device__ __forceinline__ uint64_t canonicalize(uint64_t value) { + return value >= GOLDILOCKS_P ? value - GOLDILOCKS_P : value; +} + +__device__ __forceinline__ uint64_t goldilocks_add(uint64_t left, uint64_t right) { + left = canonicalize(left); + right = canonicalize(right); + // Written this way to avoid relying on overflow behavior in the source + // language. `GOLDILOCKS_P - right` is always representable. + const uint64_t gap = GOLDILOCKS_P - right; + return left >= gap ? left - gap : left + right; +} + +__device__ __forceinline__ uint64_t goldilocks_sub(uint64_t left, uint64_t right) { + left = canonicalize(left); + right = canonicalize(right); + return left >= right ? left - right : GOLDILOCKS_P - (right - left); +} + +__device__ __forceinline__ uint64_t goldilocks_mul(uint64_t left, uint64_t right) { + left = canonicalize(left); + right = canonicalize(right); + + const uint64_t low = left * right; + const uint64_t high = __umul64hi(left, right); + + // Reduce low + high * 2^64 using 2^64 = 2^32 - 1 (mod p). + const uint64_t high_high = high >> 32; + const uint64_t high_low = high & GOLDILOCKS_EPSILON; + uint64_t reduced_low = low - high_high; + if (low < high_high) { + // The wrapped subtraction added 2^64; replace that with +p. + reduced_low -= GOLDILOCKS_EPSILON; + } + const uint64_t reduced_high = high_low * GOLDILOCKS_EPSILON; + return goldilocks_add(reduced_low, reduced_high); +} + +__device__ __forceinline__ uint64_t goldilocks_pow(uint64_t base, uint64_t exponent) { + uint64_t result = 1; + while (exponent != 0) { + if ((exponent & 1U) != 0) { + result = goldilocks_mul(result, base); + } + base = goldilocks_mul(base, base); + exponent >>= 1; + } + return result; +} + +// The resident PCS canonicalizes every committed LDE, and interpolation +// tables are serialized with `as_canonical_u64`. Restrict the faster +// canonical-input arithmetic to this boundary instead of weakening the +// representation guarantees of lookup/quotient code, whose host inputs may +// legitimately use lazy representatives. +__device__ __forceinline__ uint64_t canonical_add(uint64_t a,uint64_t b){ + const uint64_t sum=a+b; + if(sum=GOLDILOCKS_P?sum-GOLDILOCKS_P:sum; +} +__device__ __forceinline__ uint64_t canonical_mul(uint64_t a,uint64_t b){ + const uint64_t low=a*b,high=__umul64hi(a,b),high_high=high>>32; + uint64_t reduced_low=low-high_high; + if(low +#include "goldilocks.cuh" #include #include @@ -16,13 +17,14 @@ #include #include #include +#include "ntt.cuh" namespace { -constexpr uint64_t GOLDILOCKS_P = 0xffffffff00000001ULL; -constexpr uint64_t GOLDILOCKS_EPSILON = 0x00000000ffffffffULL; +using namespace multi_stark_cuda; constexpr unsigned int THREADS = 256; constexpr unsigned int MAX_BLOCKS = 65535; +constexpr size_t BLAKE3_CHUNK_BYTES = 1024; constexpr uint32_t BLAKE3_CHUNK_START = 1U << 0; constexpr uint32_t BLAKE3_CHUNK_END = 1U << 1; constexpr uint32_t BLAKE3_PARENT = 1U << 2; @@ -36,6 +38,18 @@ inline cudaError_t stream_malloc(void** pointer,size_t bytes){return cudaMallocA inline cudaError_t stream_free(void* pointer){return pointer?cudaFreeAsync(pointer,cudaStreamPerThread):cudaSuccess;} #define cudaMalloc stream_malloc #define cudaFree stream_free +// Host-side pools are kept per device, so provers on different devices in +// one process never lease each other's buffers or control blocks. Every +// entry point selects its device first, so the pools index by the calling +// thread's current device. +constexpr int MAX_CUDA_DEVICES = 64; +inline int current_device_index() { + int device = 0; + if (cudaGetDevice(&device) != cudaSuccess || device < 0 || device >= MAX_CUDA_DEVICES) { + return 0; + } + return device; +} struct ConstantCacheEntry{int device=-1;size_t count=0;uint64_t kind=0,key0=0,key1=0;uint64_t* values=nullptr;ConstantCacheEntry* next=nullptr;}; ConstantCacheEntry* constant_cache=nullptr;volatile int constant_cache_lock=0; @@ -62,56 +76,6 @@ __device__ __constant__ unsigned int BLAKE3_PERMUTATION[16] = { 2, 6, 3, 10, 7, 0, 4, 13, 1, 11, 12, 5, 9, 14, 15, 8, }; -__device__ __forceinline__ uint64_t canonicalize(uint64_t value) { - return value >= GOLDILOCKS_P ? value - GOLDILOCKS_P : value; -} - -__device__ __forceinline__ uint64_t goldilocks_add(uint64_t left, uint64_t right) { - left = canonicalize(left); - right = canonicalize(right); - // Written this way to avoid relying on overflow behavior in the source - // language. `GOLDILOCKS_P - right` is always representable. - const uint64_t gap = GOLDILOCKS_P - right; - return left >= gap ? left - gap : left + right; -} - -__device__ __forceinline__ uint64_t goldilocks_sub(uint64_t left, uint64_t right) { - left = canonicalize(left); - right = canonicalize(right); - return left >= right ? left - right : GOLDILOCKS_P - (right - left); -} - -__device__ __forceinline__ uint64_t goldilocks_mul(uint64_t left, uint64_t right) { - left = canonicalize(left); - right = canonicalize(right); - - const uint64_t low = left * right; - const uint64_t high = __umul64hi(left, right); - - // Reduce low + high * 2^64 using 2^64 = 2^32 - 1 (mod p). - const uint64_t high_high = high >> 32; - const uint64_t high_low = high & GOLDILOCKS_EPSILON; - uint64_t reduced_low = low - high_high; - if (low < high_high) { - // The wrapped subtraction added 2^64; replace that with +p. - reduced_low -= GOLDILOCKS_EPSILON; - } - const uint64_t reduced_high = high_low * GOLDILOCKS_EPSILON; - return goldilocks_add(reduced_low, reduced_high); -} - -__device__ __forceinline__ uint64_t goldilocks_pow(uint64_t base, uint64_t exponent) { - uint64_t result = 1; - while (exponent != 0) { - if ((exponent & 1U) != 0) { - result = goldilocks_mul(result, base); - } - base = goldilocks_mul(base, base); - exponent >>= 1; - } - return result; -} - __device__ __forceinline__ uint32_t rotate_right(uint32_t value, unsigned int count) { return __funnelshift_r(value, value, count); @@ -281,13 +245,6 @@ __device__ __forceinline__ void blake3_hash_digest_pair( } } -__device__ __forceinline__ size_t reverse_index_bits(size_t index, unsigned int bits) { - if (bits == 0) { - return 0; - } - return static_cast(__brevll(static_cast(index)) >> (64 - bits)); -} - unsigned int blocks_for(size_t work_items) { if (work_items == 0) { return 0; @@ -385,7 +342,7 @@ struct PinnedHostPoolEntry { }; constexpr size_t PINNED_HOST_POOL_SIZE = 2; -PinnedHostPoolEntry pinned_host_pool[PINNED_HOST_POOL_SIZE]; +PinnedHostPoolEntry pinned_host_pool[MAX_CUDA_DEVICES][PINNED_HOST_POOL_SIZE]; pthread_mutex_t pinned_host_pool_mutex = PTHREAD_MUTEX_INITIALIZER; class PinnedHostBuffer { @@ -397,7 +354,7 @@ class PinnedHostBuffer { ~PinnedHostBuffer() { if (pool_index_ < PINNED_HOST_POOL_SIZE) { pthread_mutex_lock(&pinned_host_pool_mutex); - pinned_host_pool[pool_index_].in_use = false; + pinned_host_pool[device_][pool_index_].in_use = false; pthread_mutex_unlock(&pinned_host_pool_mutex); } else if (pointer_ != nullptr) { cudaFreeHost(pointer_); @@ -408,19 +365,21 @@ class PinnedHostBuffer { if (bytes == 0 || pointer_ != nullptr) { return cudaErrorInvalidValue; } + device_ = current_device_index(); + PinnedHostPoolEntry* pool = pinned_host_pool[device_]; pthread_mutex_lock(&pinned_host_pool_mutex); size_t reusable = PINNED_HOST_POOL_SIZE; for (size_t i = 0; i < PINNED_HOST_POOL_SIZE; ++i) { - const auto& entry = pinned_host_pool[i]; + const auto& entry = pool[i]; if (!entry.in_use && entry.pointer != nullptr && entry.capacity >= bytes && (reusable == PINNED_HOST_POOL_SIZE || - entry.capacity < pinned_host_pool[reusable].capacity)) { + entry.capacity < pool[reusable].capacity)) { reusable = i; } } if (reusable < PINNED_HOST_POOL_SIZE) { - auto& entry = pinned_host_pool[reusable]; + auto& entry = pool[reusable]; entry.in_use = true; pointer_ = entry.pointer; pool_index_ = reusable; @@ -430,9 +389,9 @@ class PinnedHostBuffer { size_t available = PINNED_HOST_POOL_SIZE; for (size_t i = 0; i < PINNED_HOST_POOL_SIZE; ++i) { - if (!pinned_host_pool[i].in_use && + if (!pool[i].in_use && (available == PINNED_HOST_POOL_SIZE || - pinned_host_pool[i].capacity < pinned_host_pool[available].capacity)) { + pool[i].capacity < pool[available].capacity)) { available = i; } } @@ -441,7 +400,7 @@ class PinnedHostBuffer { return cudaMallocHost(reinterpret_cast(&pointer_), bytes); } - auto& entry = pinned_host_pool[available]; + auto& entry = pool[available]; entry.in_use = true; uint64_t* old_pointer = entry.pointer; const size_t old_capacity = entry.capacity; @@ -474,6 +433,7 @@ class PinnedHostBuffer { private: uint64_t* pointer_ = nullptr; size_t pool_index_ = PINNED_HOST_POOL_SIZE; + int device_ = 0; }; struct ResidentMerkleTree { @@ -509,21 +469,73 @@ struct ResidentMixedMerkleTree { // copy costs one host memcpy per chunk. Slots are leased under a mutex and // callers wait for a free one, which bounds concurrent uploads to the slot // count (the PCIe link is shared anyway). +// A slot is a ring of chunks, each with the event of its last transfer, so +// the host copy of one chunk overlaps the DMA of the chunks before it. constexpr size_t UPLOAD_STAGING_SLOTS = 4; constexpr size_t UPLOAD_STAGING_BYTES = size_t(64) << 20; -uint64_t* upload_staging[UPLOAD_STAGING_SLOTS] = {nullptr}; -bool upload_staging_in_use[UPLOAD_STAGING_SLOTS] = {false}; +constexpr size_t UPLOAD_STAGING_RING = 4; +constexpr size_t UPLOAD_STAGING_CHUNK = UPLOAD_STAGING_BYTES / UPLOAD_STAGING_RING; +uint64_t* upload_staging[MAX_CUDA_DEVICES][UPLOAD_STAGING_SLOTS] = {{nullptr}}; +cudaEvent_t upload_staging_events[MAX_CUDA_DEVICES][UPLOAD_STAGING_SLOTS][UPLOAD_STAGING_RING] = {{{nullptr}}}; +bool upload_staging_in_use[MAX_CUDA_DEVICES][UPLOAD_STAGING_SLOTS] = {{false}}; pthread_mutex_t upload_staging_mutex = PTHREAD_MUTEX_INITIALIZER; pthread_cond_t upload_staging_cond = PTHREAD_COND_INITIALIZER; +// One thread copies into pinned memory at a fraction of the PCIe rate, so a +// serial copy, not the DMA, would bound a staged upload; large chunks are +// copied by a few threads. +struct MemcpyPart { + unsigned char* destination; + const unsigned char* source; + size_t bytes; +}; + +void* memcpy_part(void* argument) { + const auto* part = static_cast(argument); + memcpy(part->destination, part->source, part->bytes); + return nullptr; +} + +// pthreads rather than std::thread: the executable is linked by the Lean +// toolchain without the C++ runtime library. +void parallel_memcpy(void* destination, const void* source, size_t bytes) { + constexpr size_t THREADS = 4; + constexpr size_t MIN_PARALLEL = size_t(8) << 20; + if (bytes < MIN_PARALLEL) { + memcpy(destination, source, bytes); + return; + } + const size_t part = (bytes + THREADS - 1) / THREADS; + auto* dst = static_cast(destination); + const auto* src = static_cast(source); + MemcpyPart parts[THREADS]; + pthread_t workers[THREADS]; + bool started[THREADS] = {false}; + for (size_t t = 1; t < THREADS; ++t) { + const size_t offset = t * part; + if (offset >= bytes) break; + parts[t] = {dst + offset, src + offset, std::min(part, bytes - offset)}; + started[t] = pthread_create(&workers[t], nullptr, memcpy_part, &parts[t]) == 0; + if (!started[t]) memcpy_part(&parts[t]); + } + memcpy(dst, src, std::min(part, bytes)); + for (size_t t = 1; t < THREADS; ++t) { + if (started[t]) pthread_join(workers[t], nullptr); + } +} + // Copies `bytes` of pageable host memory to the device on the per-thread // stream through a leased staging slot; returns with the copy complete. cudaError_t staged_upload(void* device, const void* host, size_t bytes) { + const int device_index = current_device_index(); + cudaError_t status = cudaSuccess; + uint64_t** slots = upload_staging[device_index]; + bool* in_use = upload_staging_in_use[device_index]; pthread_mutex_lock(&upload_staging_mutex); size_t slot = UPLOAD_STAGING_SLOTS; while (slot == UPLOAD_STAGING_SLOTS) { for (size_t i = 0; i < UPLOAD_STAGING_SLOTS; ++i) { - if (!upload_staging_in_use[i]) { + if (!in_use[i]) { slot = i; break; } @@ -532,34 +544,50 @@ cudaError_t staged_upload(void* device, const void* host, size_t bytes) { pthread_cond_wait(&upload_staging_cond, &upload_staging_mutex); } } - upload_staging_in_use[slot] = true; - cudaError_t status = cudaSuccess; - if (upload_staging[slot] == nullptr) { - status = cudaMallocHost(reinterpret_cast(&upload_staging[slot]), + in_use[slot] = true; + cudaEvent_t* events = upload_staging_events[device_index][slot]; + if (slots[slot] == nullptr) { + status = cudaMallocHost(reinterpret_cast(&slots[slot]), UPLOAD_STAGING_BYTES); - if (status != cudaSuccess) upload_staging[slot] = nullptr; + if (status != cudaSuccess) slots[slot] = nullptr; + for (size_t i = 0; status == cudaSuccess && i < UPLOAD_STAGING_RING; ++i) { + status = cudaEventCreateWithFlags(&events[i], cudaEventDisableTiming); + } } pthread_mutex_unlock(&upload_staging_mutex); - uint64_t* staging = upload_staging[slot]; + auto* staging = reinterpret_cast(slots[slot]); if (status == cudaSuccess) { const auto* source = static_cast(host); auto* target = static_cast(device); + size_t index = 0; for (size_t offset = 0; offset < bytes && status == cudaSuccess; - offset += UPLOAD_STAGING_BYTES) { - const size_t chunk = std::min(UPLOAD_STAGING_BYTES, bytes - offset); - memcpy(staging, source + offset, chunk); - status = cudaMemcpyAsync(target + offset, staging, chunk, - cudaMemcpyHostToDevice, cudaStreamPerThread); - if (status == cudaSuccess) status = cudaStreamSynchronize(cudaStreamPerThread); + offset += UPLOAD_STAGING_CHUNK, ++index) { + const size_t ring = index % UPLOAD_STAGING_RING; + unsigned char* buffer = staging + ring * UPLOAD_STAGING_CHUNK; + // The buffer is free once its previous transfer has landed. + if (index >= UPLOAD_STAGING_RING) status = cudaEventSynchronize(events[ring]); + const size_t chunk = std::min(UPLOAD_STAGING_CHUNK, bytes - offset); + if (status == cudaSuccess) { + parallel_memcpy(buffer, source + offset, chunk); + status = cudaMemcpyAsync(target + offset, buffer, chunk, + cudaMemcpyHostToDevice, cudaStreamPerThread); + } + if (status == cudaSuccess) status = cudaEventRecord(events[ring], cudaStreamPerThread); } + // The slot is handed to the next caller on release. + const cudaError_t synced = cudaStreamSynchronize(cudaStreamPerThread); + if (status == cudaSuccess) status = synced; } pthread_mutex_lock(&upload_staging_mutex); - upload_staging_in_use[slot] = false; - pthread_cond_signal(&upload_staging_cond); + in_use[slot] = false; + pthread_cond_broadcast(&upload_staging_cond); pthread_mutex_unlock(&upload_staging_mutex); return status; } +using TraceWriter = int (*)(void*, int, uint64_t*, size_t, size_t); +using TraceDestroy = void (*)(void*); + struct ResidentLde { uint64_t* values = nullptr; // Original row-major evaluations, retained for prover stages (notably @@ -569,12 +597,19 @@ struct ResidentLde { const uint64_t* host_trace_values = nullptr; bool host_trace_registered = false; size_t trace_height = 0; + void* trace_context = nullptr; + TraceWriter trace_writer = nullptr; + TraceDestroy trace_destroy = nullptr; size_t height = 0; size_t width = 0; uint8_t* interpolation_scratch = nullptr; size_t interpolation_scratch_bytes = 0; + // The device whose slab this control block came from, so it returns + // there when destroyed from a thread on another device. + int device = 0; ~ResidentLde() { + if (trace_destroy) trace_destroy(trace_context); if (values != nullptr) { cudaFree(values); } @@ -591,14 +626,17 @@ struct ResidentLde { // whole device, so control blocks come from slabs allocated once per process // and are recycled through a free list rather than allocated per LDE. constexpr size_t RESIDENT_LDE_SLAB = 256; -void* resident_lde_free = nullptr; +// One free list per device: a slab is allocated with its device current, so +// its pages prefer that device and are never read by kernels on another. +void* resident_lde_free[MAX_CUDA_DEVICES] = {nullptr}; volatile int resident_lde_lock = 0; cudaError_t create_resident_lde(ResidentLde** output) { if (output == nullptr) return cudaErrorInvalidValue; *output = nullptr; + const int device = current_device_index(); while (__sync_lock_test_and_set(&resident_lde_lock, 1)) {} - if (resident_lde_free == nullptr) { + if (resident_lde_free[device] == nullptr) { void* slab = nullptr; const cudaError_t status = cudaMallocManaged(&slab, RESIDENT_LDE_SLAB * sizeof(ResidentLde)); @@ -609,23 +647,25 @@ cudaError_t create_resident_lde(ResidentLde** output) { auto* bytes = static_cast(slab); for (size_t i = 0; i < RESIDENT_LDE_SLAB; ++i) { void* slot = bytes + i * sizeof(ResidentLde); - *static_cast(slot) = resident_lde_free; - resident_lde_free = slot; + *static_cast(slot) = resident_lde_free[device]; + resident_lde_free[device] = slot; } } - void* slot = resident_lde_free; - resident_lde_free = *static_cast(slot); + void* slot = resident_lde_free[device]; + resident_lde_free[device] = *static_cast(slot); __sync_lock_release(&resident_lde_lock); *output = new (slot) ResidentLde; + (*output)->device = device; return cudaSuccess; } cudaError_t destroy_resident_lde(ResidentLde* lde) { if (lde == nullptr) return cudaSuccess; + const int device = lde->device; lde->~ResidentLde(); while (__sync_lock_test_and_set(&resident_lde_lock, 1)) {} - *reinterpret_cast(lde) = resident_lde_free; - resident_lde_free = lde; + *reinterpret_cast(lde) = resident_lde_free[device]; + resident_lde_free[device] = lde; __sync_lock_release(&resident_lde_lock); return cudaSuccess; } @@ -680,23 +720,6 @@ struct PendingLookupLde { } }; -// The resident PCS canonicalizes every committed LDE, and interpolation -// tables are serialized with `as_canonical_u64`. Restrict the faster -// canonical-input arithmetic to this boundary instead of weakening the -// representation guarantees of lookup/quotient code, whose host inputs may -// legitimately use lazy representatives. -__device__ __forceinline__ uint64_t canonical_add(uint64_t a,uint64_t b){ - const uint64_t sum=a+b; - if(sum=GOLDILOCKS_P?sum-GOLDILOCKS_P:sum; -} -__device__ __forceinline__ uint64_t canonical_mul(uint64_t a,uint64_t b){ - const uint64_t low=a*b,high=__umul64hi(a,b),high_high=high>>32; - uint64_t reduced_low=low-high_high; - if(low> 1) * width; - const size_t stride = height / (2 * half); - const size_t grid_stride = static_cast(blockDim.x) * gridDim.x; - for (size_t index = static_cast(blockIdx.x) * blockDim.x + threadIdx.x; - index < total; index += grid_stride) { - const size_t butterfly = index / width; - const size_t column = index - butterfly * width; - const size_t offset = butterfly % half; - const size_t group = butterfly / half; - const size_t row_0 = group * (2 * half) + offset; - const size_t row_1 = row_0 + half; - const size_t index_0 = row_0 * width + column; - const size_t index_1 = row_1 * width + column; - const uint64_t left = values[index_0]; - const uint64_t right = values[index_1]; - values[index_0] = goldilocks_add(left, right); - values[index_1] = goldilocks_mul( - goldilocks_sub(left, right), twiddles[offset * stride]); - } -} - -// Fuse two consecutive radix-2 DIF stages. Each thread owns one column of -// one four-row butterfly, so row-major accesses remain coalesced while the -// intermediate values never return to global memory. -__global__ void radix4_dif_stage(uint64_t* values,size_t height,size_t width, - size_t half,const uint64_t* twiddles){ - const size_t quarter=half>>1,total=(height>>2)*width,stride=height/(2*half); - const size_t grid_stride=static_cast(blockDim.x)*gridDim.x; - for(size_t index=static_cast(blockIdx.x)*blockDim.x+threadIdx.x;index>2,total=(height>>3)*width,stride=height/(2*half); - const size_t grid_stride=static_cast(blockDim.x)*gridDim.x; - for(size_t index=static_cast(blockIdx.x)*blockDim.x+threadIdx.x;index>= 1) { - const size_t butterflies = start_half * width; - const size_t stride = height / (2 * half); - for (size_t index = threadIdx.x; index < butterflies; - index += blockDim.x) { - const size_t butterfly = index / width; - const size_t column = index - butterfly * width; - const size_t offset = butterfly % half; - const size_t subgroup = butterfly / half; - const size_t row_0 = subgroup * (2 * half) + offset; - const size_t row_1 = row_0 + half; - const size_t index_0 = row_0 * width + column; - const size_t index_1 = row_1 * width + column; - const uint64_t left = local_values[index_0]; - const uint64_t right = local_values[index_1]; - local_values[index_0] = goldilocks_add(left, right); - local_values[index_1] = goldilocks_mul( - goldilocks_sub(left, right), twiddles[offset * stride]); - } - __syncthreads(); - if (half == 1) { - break; - } - } - - for (size_t index = threadIdx.x; index < elements_per_group; - index += blockDim.x) { - values[global_base + index] = local_values[index]; - } - __syncthreads(); - } -} - -// Wide row-major matrices cannot fit all columns of a row group in shared -// memory. Tile the columns as well, preserving coalesced row-major loads while -// fusing the final DIF stages independently for each column tile. -__global__ void radix2_dif_tail_tiled(uint64_t* values,size_t height,size_t width, - size_t start_half,const uint64_t* twiddles,size_t columns_per_tile){ - extern __shared__ uint64_t local_values[];const size_t rows=2*start_half; - const size_t groups=height/rows,column_tiles=(width+columns_per_tile-1)/columns_per_tile; - const size_t tile_count=groups*column_tiles; - for(size_t tile_index=blockIdx.x;tile_index>=1){const size_t butterflies=start_half*columns,stride=height/(2*half); - for(size_t index=threadIdx.x;index> 1) * width; - const unsigned int blocks = blocks_for(total); - // Tall, wide row-major batches amortize the heavier fused kernel. Width-2 - // FRI codewords retain the shared-memory tail specialized below. - if (width >= 8 && height >= (size_t(1) << 18)) { - size_t half = height >> 1; - while (half >= 4) { - radix8_dif_stage<<> 3) * width), THREADS>>>( - values, height, width, half, twiddles); - const cudaError_t status = cudaGetLastError(); - if (status != cudaSuccess) return status; - half >>= 3; - } - if (half == 2) { - radix4_dif_stage<<> 2) * width), THREADS>>>( - values, height, width, half, twiddles); - return cudaGetLastError(); - } - if (half == 1) { - radix2_dif_stage<<>>(values, height, width, 1, - twiddles); - } - return cudaGetLastError(); - } - if (width > 2) { - for (size_t half = height >> 1;; half >>= 1) { - radix2_dif_stage<<>>(values, height, width, half, - twiddles); - const cudaError_t status = cudaGetLastError(); - if (status != cudaSuccess) return status; - if (half == 1) return cudaSuccess; - } - } - constexpr size_t TAIL_HALF = 128; - for (size_t half = height >> 1; half > TAIL_HALF; half >>= 1) { - radix2_dif_stage<<>>(values, height, width, half, - twiddles); - const cudaError_t status = cudaGetLastError(); - if (status != cudaSuccess) { - return status; - } - } - - const size_t start_half = - height < 2 * TAIL_HALF ? height >> 1 : TAIL_HALF; - const size_t groups = height / (2 * start_half); - const unsigned int tail_blocks = static_cast( - groups < MAX_BLOCKS ? groups : MAX_BLOCKS); - const size_t shared_bytes = 2 * start_half * width * sizeof(uint64_t); - radix2_dif_tail<<>>( - values, height, width, start_half, twiddles); - return cudaGetLastError(); -} - -__global__ void bit_reverse_scale_and_shift(uint64_t* values, size_t height, - size_t width, unsigned int log_height, - uint64_t height_inverse, - const uint64_t* shift_powers) { - const size_t total = height * width; - const size_t grid_stride = static_cast(blockDim.x) * gridDim.x; - for (size_t index = static_cast(blockIdx.x) * blockDim.x + threadIdx.x; - index < total; index += grid_stride) { - const size_t row = index / width; - const size_t column = index - row * width; - const size_t reverse_row = reverse_index_bits(row, log_height); - if (row > reverse_row) { - continue; - } - - const size_t reverse_index = reverse_row * width + column; - const uint64_t left = values[index]; - if (row == reverse_row) { - values[index] = goldilocks_mul( - left, goldilocks_mul(height_inverse, shift_powers[row])); - continue; - } - - const uint64_t right = values[reverse_index]; - values[index] = goldilocks_mul( - right, goldilocks_mul(height_inverse, shift_powers[row])); - values[reverse_index] = goldilocks_mul( - left, goldilocks_mul(height_inverse, shift_powers[reverse_row])); - } -} - __global__ void goldilocks_ops_kernel(uint64_t* sums, uint64_t* differences, uint64_t* products, uint64_t* inverses, const uint64_t* left, const uint64_t* right, @@ -1793,6 +1561,32 @@ __global__ void goldilocks_ops_kernel(uint64_t* sums, uint64_t* differences, } } +// A single-chunk message has no chunk tree to reduce. Assigning one thread +// per row keeps every lane doing useful compression instead of leaving 31 +// lanes idle in the warp-per-row kernel, with the same chunk primitive and +// little-endian digest encoding. +__global__ void blake3_hash_short_rows_kernel(uint8_t* digests, + const uint8_t* messages, + size_t message_bytes, + size_t message_count) { + const size_t grid_stride = static_cast(blockDim.x) * gridDim.x; + for (size_t index = static_cast(blockIdx.x) * blockDim.x + threadIdx.x; + index < message_count; index += grid_stride) { + uint32_t output[8]; + blake3_chunk(messages + index * message_bytes, message_bytes, 0, true, + output); + uint8_t* digest = digests + index * 32; +#pragma unroll + for (unsigned int word = 0; word < 8; ++word) { + const uint32_t value = output[word]; + digest[word * 4] = static_cast(value); + digest[word * 4 + 1] = static_cast(value >> 8); + digest[word * 4 + 2] = static_cast(value >> 16); + digest[word * 4 + 3] = static_cast(value >> 24); + } + } +} + // One warp owns one message. Chunk compression is distributed across lanes; // lane zero then reduces the (at most 32) chunk chaining values into the root. // This matches the BLAKE3 tree shape while exposing parallelism within wide @@ -1890,6 +1684,12 @@ __global__ void blake3_hash_digest_pairs_kernel( cudaError_t launch_blake3_rows(uint8_t* digests, const uint8_t* messages, size_t message_bytes, size_t message_count) { + if (message_bytes <= BLAKE3_CHUNK_BYTES) { + blake3_hash_short_rows_kernel<<>>( + digests, messages, message_bytes, message_count); + return cudaGetLastError(); + } constexpr unsigned int WARPS_PER_BLOCK = THREADS / 32; const size_t required = (message_count + WARPS_PER_BLOCK - 1) / WARPS_PER_BLOCK; @@ -2163,10 +1963,41 @@ cudaError_t copy_to_host(uint64_t* destination, const DeviceBuffer& source, } // namespace +static cudaError_t plan_constants(int device, const MultiStarkNttPlan* plan, + const uint64_t** powers) { + if (!plan || !plan->shift_powers) return cudaErrorInvalidValue; + const auto* host = plan->shift_powers; + return cached_device_constants(device, host, plan->height, 3, host[0], + plan->height > 1 ? host[1] : 0, powers); +} + +static bool plan_matches(const MultiStarkNttPlan* plan, size_t height, + size_t width, size_t extended_height) { + return plan && plan->height == height && plan->width == width && + plan->extended_height == extended_height; +} + +static cudaError_t coset_lde(int device, const uint64_t* trace, uint64_t* values, + size_t height, size_t width, size_t extended_height, + const MultiStarkNttPlan* plan) { + if (!plan_matches(plan, height, width, extended_height)) return cudaErrorInvalidValue; + const uint64_t* powers = nullptr; + cudaError_t status = plan_constants(device, plan, &powers); + if (status != cudaSuccess) return status; + return static_cast(multi_stark_sppark_coset_lde(device, trace, values, plan, powers)); +} + +static cudaError_t forward_in_place(int device, uint64_t* values, size_t height, + size_t width, const MultiStarkNttPlan* plan) { + if (!plan_matches(plan, height, width, height) || plan->shift_powers) + return cudaErrorInvalidValue; + return static_cast(multi_stark_sppark_forward(device, values, plan)); +} + extern "C" int multi_stark_cuda_dft_batch(int device_id, uint64_t* values, size_t height, size_t width, - const uint64_t* twiddles) { - if (values == nullptr || twiddles == nullptr || !is_power_of_two(height) || + const MultiStarkNttPlan* plan) { + if (values == nullptr || !is_power_of_two(height) || width == 0 || !product_fits(height, width)) { return static_cast(cudaErrorInvalidValue); } @@ -2178,13 +2009,9 @@ extern "C" int multi_stark_cuda_dft_batch(int device_id, uint64_t* values, const size_t elements = height * width; HostRegistration registered_values(values, elements * sizeof(uint64_t)); DeviceBuffer device_values; - DeviceBuffer device_twiddles; status = copy_to_device(device_values, values, elements); if (status == cudaSuccess) { - status = copy_to_device(device_twiddles, twiddles, height / 2); - } - if (status == cudaSuccess) { - status = launch_dif(device_values.get(), height, width, device_twiddles.get()); + status = forward_in_place(device_id, device_values.get(), height, width, plan); } if (status == cudaSuccess) { status = copy_to_host(values, device_values, elements); @@ -2194,11 +2021,8 @@ extern "C" int multi_stark_cuda_dft_batch(int device_id, uint64_t* values, extern "C" int multi_stark_cuda_coset_lde_batch( int device_id, uint64_t* output, const uint64_t* input, size_t height, - size_t width, size_t added_bits, const uint64_t* inverse_twiddles, - const uint64_t* shift_powers, const uint64_t* forward_twiddles, - uint64_t height_inverse) { - if (output == nullptr || input == nullptr || inverse_twiddles == nullptr || - shift_powers == nullptr || forward_twiddles == nullptr || + size_t width, size_t added_bits, const MultiStarkNttPlan* plan) { + if (output == nullptr || input == nullptr || plan == nullptr || plan->shift_powers == nullptr || !is_power_of_two(height) || width == 0 || added_bits >= sizeof(size_t) * 8 || height > (SIZE_MAX >> added_bits)) { return static_cast(cudaErrorInvalidValue); @@ -2220,42 +2044,14 @@ extern "C" int multi_stark_cuda_coset_lde_batch( HostRegistration registered_output(output, output_elements * sizeof(uint64_t)); DeviceBuffer device_values; - DeviceBuffer device_inverse_twiddles; - DeviceBuffer device_shift_powers; - DeviceBuffer device_forward_twiddles; status = device_values.allocate(output_elements); - if (status == cudaSuccess) { - status = cudaMemset(device_values.get(), 0, output_elements * sizeof(uint64_t)); - } if (status == cudaSuccess) { status = cudaMemcpy(device_values.get(), input, input_elements * sizeof(uint64_t), cudaMemcpyHostToDevice); } if (status == cudaSuccess) { - status = copy_to_device(device_inverse_twiddles, inverse_twiddles, height / 2); - } - if (status == cudaSuccess) { - status = copy_to_device(device_shift_powers, shift_powers, height); - } - if (status == cudaSuccess) { - status = copy_to_device(device_forward_twiddles, forward_twiddles, - extended_height / 2); - } - if (status == cudaSuccess) { - status = launch_dif(device_values.get(), height, width, - device_inverse_twiddles.get()); - } - if (status == cudaSuccess) { - const size_t total = input_elements; - bit_reverse_scale_and_shift<<>>( - device_values.get(), height, width, strict_log2(height), - height_inverse, device_shift_powers.get()); - status = cudaGetLastError(); - } - if (status == cudaSuccess) { - status = launch_dif(device_values.get(), extended_height, width, - device_forward_twiddles.get()); + status = coset_lde(device_id, device_values.get(), device_values.get(), height, width, extended_height, plan); } if (status == cudaSuccess) { status = copy_to_host(output, device_values, output_elements); @@ -2263,13 +2059,11 @@ extern "C" int multi_stark_cuda_coset_lde_batch( return static_cast(status); } -extern "C" int multi_stark_cuda_coset_lde_create( +static int coset_lde_create( int device_id, void** handle, const uint64_t* input, size_t height, - size_t width, size_t added_bits, const uint64_t* inverse_twiddles, - const uint64_t* shift_powers, const uint64_t* forward_twiddles, - uint64_t height_inverse) { - if (handle == nullptr || input == nullptr || (height > 1 && inverse_twiddles == nullptr) || - shift_powers == nullptr || forward_twiddles == nullptr || + size_t width, size_t added_bits, const MultiStarkNttPlan* plan, void* context, TraceWriter writer) { + if (handle == nullptr || (input == nullptr && writer == nullptr) || + plan == nullptr || plan->shift_powers == nullptr || !is_power_of_two(height) || width == 0 || added_bits >= sizeof(size_t) * 8 || height > (SIZE_MAX >> added_bits)) { return static_cast(cudaErrorInvalidValue); @@ -2287,6 +2081,8 @@ extern "C" int multi_stark_cuda_coset_lde_create( ResidentLde* lde = nullptr; status = create_resident_lde(&lde); if (status != cudaSuccess) return static_cast(status); + lde->trace_context = context; + lde->trace_writer = writer; lde->height = extended_height; lde->width = width; const size_t input_elements = height * width; @@ -2300,17 +2096,15 @@ extern "C" int multi_stark_cuda_coset_lde_create( } if (status == cudaSuccess) lde->trace_height = height; - // Large pageable uploads otherwise serialize through the driver's hidden - // staging pool; they go through the persistent staging slots instead. - // Small ones take the direct pageable path. - const uint64_t *device_inverse_twiddles=nullptr,*device_shift_powers=nullptr,*device_forward_twiddles=nullptr; - if (status == cudaSuccess) { - status = cudaMemsetAsync(lde->values, 0, - output_elements * sizeof(uint64_t), - cudaStreamPerThread); - } if (status == cudaSuccess) { - if (input_bytes >= (size_t(8) << 20)) { + if (writer) { + constexpr size_t TILE_ROWS = size_t(1) << 16; + for (size_t first = 0; status == cudaSuccess && first < height; first += TILE_ROWS) { + const size_t rows = height - first < TILE_ROWS ? height - first : TILE_ROWS; + status = static_cast(writer(context, device_id, + lde->trace_values + first * width, first, rows)); + } + } else if (input_bytes >= (size_t(8) << 20)) { status = staged_upload(lde->trace_values, input, input_bytes); } else { status = cudaMemcpyAsync(lde->trace_values, input, @@ -2319,38 +2113,8 @@ extern "C" int multi_stark_cuda_coset_lde_create( cudaStreamPerThread); } } - if (status == cudaSuccess) { - status = cudaMemcpyAsync(lde->values, lde->trace_values, - input_elements * sizeof(uint64_t), - cudaMemcpyDeviceToDevice, - cudaStreamPerThread); - } - if (status == cudaSuccess && height > 1) { - status = cached_device_constants(device_id,inverse_twiddles,height/2,1,0,0,&device_inverse_twiddles); - } - if (status == cudaSuccess) { - status = cached_device_constants(device_id,shift_powers,height,3,shift_powers[0],height>1?shift_powers[1]:0,&device_shift_powers); - } - if (status == cudaSuccess) { - status = cached_device_constants(device_id,forward_twiddles,extended_height/2,2,0,0,&device_forward_twiddles); - } - if (status == cudaSuccess) { - status = launch_dif(lde->values, height, width,device_inverse_twiddles); - } - if (status == cudaSuccess) { - bit_reverse_scale_and_shift<<>>( - lde->values, height, width, strict_log2(height), height_inverse, - device_shift_powers); - status = cudaGetLastError(); - } - if (status == cudaSuccess) { - status = launch_dif(lde->values, extended_height, width,device_forward_twiddles); - } - if (status == cudaSuccess) { - canonicalize_goldilocks<<>>( - lde->values, output_elements); - status = cudaGetLastError(); - } + if (status == cudaSuccess) + status = coset_lde(device_id, lde->trace_values, lde->values, height, width, extended_height, plan); if (status == cudaSuccess) status = cudaStreamSynchronize(cudaStreamPerThread); if (status != cudaSuccess) { destroy_resident_lde(lde); @@ -2360,19 +2124,40 @@ extern "C" int multi_stark_cuda_coset_lde_create( return static_cast(cudaSuccess); } -extern "C" int multi_stark_cuda_prepare_lde_constants( - int device_id,const uint64_t* inverse_twiddles,size_t inverse_count, - const uint64_t* shift_powers,size_t height,const uint64_t* forward_twiddles, - size_t forward_count) { - if((inverse_count&&!inverse_twiddles)||!shift_powers||!height|| - !forward_twiddles||!forward_count)return static_cast(cudaErrorInvalidValue); - cudaError_t status=cudaSetDevice(device_id);const uint64_t* ignored=nullptr; - if(status==cudaSuccess&&inverse_count) - status=cached_device_constants(device_id,inverse_twiddles,inverse_count,1,0,0,&ignored); - if(status==cudaSuccess) - status=cached_device_constants(device_id,shift_powers,height,3,shift_powers[0],height>1?shift_powers[1]:0,&ignored); - if(status==cudaSuccess) - status=cached_device_constants(device_id,forward_twiddles,forward_count,2,0,0,&ignored); +extern "C" int multi_stark_cuda_coset_lde_create( + int device_id, void** handle, const uint64_t* input, size_t height, + size_t width, size_t added_bits, const MultiStarkNttPlan* plan) { + return coset_lde_create(device_id, handle, input, height, width, added_bits, + plan, nullptr, nullptr); +} + +// Context ownership transfers only on success. Its lifetime covers trace +// release and LDE eviction because lookup recovery still needs original rows. +extern "C" int multi_stark_cuda_coset_lde_generate( + int device_id, void** handle, size_t height, size_t width, size_t added_bits, + const MultiStarkNttPlan* plan, + void* context, TraceWriter writer, TraceDestroy destroy) { + if (!context || !writer || !destroy) return static_cast(cudaErrorInvalidValue); + const int status = coset_lde_create(device_id, handle, nullptr, height, width, added_bits, + plan, context, writer); + if (status == 0) static_cast(*handle)->trace_destroy = destroy; + return status; +} + +extern "C" bool multi_stark_cuda_lde_has_generator(const void* handle) { + return handle && static_cast(handle)->trace_writer; +} + +extern "C" void* multi_stark_cuda_lde_generator_context(const void* handle) { + if (!handle) return nullptr; + const auto* lde = static_cast(handle); + return lde->trace_writer ? lde->trace_context : nullptr; +} + +extern "C" int multi_stark_cuda_prepare_lde_constants(int device_id, const MultiStarkNttPlan* plan) { + cudaError_t status = cudaSetDevice(device_id); + const uint64_t* ignored = nullptr; + if (status == cudaSuccess) status = plan_constants(device_id, plan, &ignored); return static_cast(status); } @@ -2709,12 +2494,11 @@ extern "C" int multi_stark_cuda_quotient_lde( const uint64_t* alpha, size_t constraint_count, const uint64_t* delta, uint64_t ext_w, size_t quotient_size, size_t next_step, size_t quotient_degree, - size_t log_blowup, const uint64_t* quotient_twiddles, - const uint64_t* lde_twiddles, const uint64_t* slice_weights) { + size_t log_blowup, const MultiStarkNttPlan* quotient_plan, + const MultiStarkNttPlan* lde_plan, const uint64_t* slice_weights) { if (output_handle == nullptr || nodes == nullptr || roots == nullptr || main_handle == nullptr || stage2_handle == nullptr || publics == nullptr || alpha == nullptr || delta == nullptr || - quotient_twiddles == nullptr || lde_twiddles == nullptr || slice_weights == nullptr || node_count == 0 || slot_count == 0 || group_size == 0 || !is_power_of_two(quotient_size) || !is_power_of_two(quotient_degree) || quotient_degree > quotient_size || @@ -2758,9 +2542,7 @@ extern "C" int multi_stark_cuda_quotient_lde( if(status==cudaSuccess)status=generate_coset_selectors(ds,quotient_size,next_step, coset_shift,coset_generator,trace_last,vanishing_start,vanishing_step); - const uint64_t *device_quotient_twiddles=nullptr,*device_lde_twiddles=nullptr,*device_weights=nullptr; - if(status==cudaSuccess)status=cached_device_constants(device_id,quotient_twiddles,quotient_size/2,2,0,0,&device_quotient_twiddles); - if(status==cudaSuccess)status=cached_device_constants(device_id,lde_twiddles,lde_height/2,2,0,0,&device_lde_twiddles); + const uint64_t* device_weights=nullptr; if(status==cudaSuccess)status=cached_device_constants(device_id,slice_weights,quotient_degree,4,slice_weights[0],quotient_degree>1?slice_weights[1]:0,&device_weights); size_t budget=0;if(status==cudaSuccess)status=quotient_shared_memory_budget(device_id,&budget); @@ -2782,7 +2564,7 @@ extern "C" int multi_stark_cuda_quotient_lde( dp,ds,dal,dd,ext_w,quotient_size,next_step,scratch,0,quotient_size,false); status=cudaGetLastError(); } - if(status==cudaSuccess)status=launch_dif(quotient,quotient_size,2,device_quotient_twiddles); + if(status==cudaSuccess)status=forward_in_place(device_id, quotient,quotient_size,2,quotient_plan); ResidentLde* lde = nullptr; if(status==cudaSuccess) { @@ -2799,7 +2581,7 @@ extern "C" int multi_stark_cuda_quotient_lde( quotient_degree,2); status=cudaGetLastError(); } - if(status==cudaSuccess)status=launch_dif(lde->values,lde_height,width,device_lde_twiddles); + if(status==cudaSuccess)status=forward_in_place(device_id, lde->values,lde_height,width,lde_plan); if(status==cudaSuccess)status=cudaStreamSynchronize(0); if(status==cudaSuccess) { *output_handle=lde; @@ -2832,16 +2614,15 @@ extern "C" int multi_stark_cuda_quotient_lde_mixed( uint64_t vanishing_start, uint64_t vanishing_step, const uint64_t* alpha, size_t constraint_count, const uint64_t* delta, uint64_t ext_w, size_t quotient_size, size_t next_step, size_t quotient_degree, - size_t log_blowup, const uint64_t* quotient_twiddles, - const uint64_t* lde_twiddles, const uint64_t* slice_weights) { + size_t log_blowup, const MultiStarkNttPlan* quotient_plan, + const MultiStarkNttPlan* lde_plan, const uint64_t* slice_weights) { const auto one_source = [](const void* handle, const uint64_t* host) { return (handle != nullptr) != (host != nullptr); }; if (output_handle == nullptr || nodes == nullptr || roots == nullptr || !one_source(main_handle, main_host) || !one_source(stage2_handle, stage2_host) || publics == nullptr || - alpha == nullptr || delta == nullptr || quotient_twiddles == nullptr || - lde_twiddles == nullptr || slice_weights == nullptr || node_count == 0 || + alpha == nullptr || delta == nullptr || slice_weights == nullptr || node_count == 0 || slot_count == 0 || group_size == 0 || !is_power_of_two(quotient_size) || !is_power_of_two(quotient_degree) || quotient_degree > quotient_size || quotient_size % quotient_degree != 0 || !is_power_of_two(next_step) || @@ -2887,9 +2668,7 @@ extern "C" int multi_stark_cuda_quotient_lde_mixed( if(status==cudaSuccess)status=generate_coset_selectors(ds,quotient_size,next_step, coset_shift,coset_generator,trace_last,vanishing_start,vanishing_step); - const uint64_t *device_quotient_twiddles=nullptr,*device_lde_twiddles=nullptr,*device_weights=nullptr; - if(status==cudaSuccess)status=cached_device_constants(device_id,quotient_twiddles,quotient_size/2,2,0,0,&device_quotient_twiddles); - if(status==cudaSuccess)status=cached_device_constants(device_id,lde_twiddles,lde_height/2,2,0,0,&device_lde_twiddles); + const uint64_t* device_weights=nullptr; if(status==cudaSuccess)status=cached_device_constants(device_id,slice_weights,quotient_degree,4,slice_weights[0],quotient_degree>1?slice_weights[1]:0,&device_weights); const auto* prep_resident=static_cast(preprocessed_handle); @@ -3011,7 +2790,7 @@ extern "C" int multi_stark_cuda_quotient_lde_mixed( } for(size_t i=0;i<2;++i)if(status==cudaSuccess&&stream_busy[i])status=cudaStreamSynchronize(streams[i]); - if(status==cudaSuccess)status=launch_dif(quotient,quotient_size,2,device_quotient_twiddles); + if(status==cudaSuccess)status=forward_in_place(device_id, quotient,quotient_size,2,quotient_plan); ResidentLde* lde=nullptr; if(status==cudaSuccess)status=create_resident_lde(&lde); if(status==cudaSuccess){lde->height=lde_height;lde->width=width;status=cudaMalloc(reinterpret_cast(&lde->values),lde_height*width*sizeof(uint64_t));} @@ -3021,7 +2800,7 @@ extern "C" int multi_stark_cuda_quotient_lde_mixed( lde->values,quotient,device_weights,quotient_size,trace_height,quotient_degree,2); status=cudaGetLastError(); } - if(status==cudaSuccess)status=launch_dif(lde->values,lde_height,width,device_lde_twiddles); + if(status==cudaSuccess)status=forward_in_place(device_id, lde->values,lde_height,width,lde_plan); if(status==cudaSuccess)status=cudaStreamSynchronize(0); if(status==cudaSuccess)*output_handle=lde;else if(lde)destroy_resident_lde(lde); for(size_t i=0;i<2;++i){ @@ -3060,8 +2839,9 @@ extern "C" int multi_stark_cuda_fri_workspace_create(int device_id,void** handle cudaError_t status=cudaSetDevice(device_id);auto* w=new(std::nothrow) ResidentFriWorkspace;if(!w)return static_cast(cudaErrorMemoryAllocation); w->inv_count=inv_count;w->coset_count=coset_count; if(status==cudaSuccess)status=cudaMalloc(reinterpret_cast(&w->inv_denoms),inv_count*sizeof(Ext2)); - if(status==cudaSuccess)status=cudaMalloc(reinterpret_cast(&w->coset),coset_count*sizeof(uint64_t)); - if(status==cudaSuccess)status=cudaMemcpy(w->coset,coset,coset_count*sizeof(uint64_t),cudaMemcpyHostToDevice); + // One upload per device and size: every shard's opening indexes the same + // bit-reversed coset, 512 MiB at 2^26 rows. + if(status==cudaSuccess)status=cached_device_constants(device_id,coset,coset_count,5,coset[0],coset_count>1?coset[1]:0,&w->coset); uint64_t *norms=nullptr,*inverses=nullptr; if(status==cudaSuccess)status=cudaMalloc(reinterpret_cast(&norms),max_count*sizeof(uint64_t)); if(status==cudaSuccess)status=cudaMalloc(reinterpret_cast(&inverses),max_count*sizeof(uint64_t)); @@ -3231,14 +3011,13 @@ extern "C" int multi_stark_cuda_lookup_graph_lde(int device_id,void** output_han const void* nodes,size_t node_count,size_t slot_count,const void* lookups,size_t lookup_count, const uint32_t* lookup_args,size_t lookup_arg_count,const void* preprocessed_handle, const void* main_handle,size_t group_size,const uint64_t* beta,const uint64_t* gamma, - uint64_t ext_w,size_t added_bits,const uint64_t* inverse_twiddles, - const uint64_t* shift_powers,const uint64_t* forward_twiddles,uint64_t height_inverse){ + uint64_t ext_w,size_t added_bits,const MultiStarkNttPlan* plan){ auto* main=const_cast(static_cast(main_handle)); auto* prep=static_cast(preprocessed_handle); if(!output_handle||!total||!nodes||!node_count||!slot_count||!lookups||!lookup_count|| (lookup_arg_count&&!lookup_args)||!main|| - (!main->trace_values&&!main->host_trace_values)||!main->trace_height|| - !group_size||!beta||!gamma||!inverse_twiddles||!shift_powers||!forward_twiddles|| + (!main->trace_values&&!main->host_trace_values&&!main->trace_writer)||!main->trace_height|| + !group_size||!beta||!gamma||(!plan || !plan->shift_powers)|| !is_power_of_two(main->trace_height)||added_bits>=sizeof(size_t)*8|| main->trace_height>(SIZE_MAX>>added_bits)||(prep&&!prep->trace_values)) { return static_cast(cudaErrorInvalidValue); @@ -3272,14 +3051,19 @@ extern "C" int multi_stark_cuda_lookup_graph_lde(int device_id,void** output_han const size_t budget=48*1024;size_t tile=budget/(slot_count*sizeof(uint64_t));const bool global=tile<32; if(global)tile=128;else if(tile>32)tile=32;size_t blocks=(chunk_rows+tile-1)/tile;const size_t cap=global?256:1024;if(blocks>cap)blocks=cap; if(status==cudaSuccess&&global)status=cudaMalloc(reinterpret_cast(&scratch),blocks*slot_count*tile*sizeof(uint64_t)); - if(status==cudaSuccess&&main->host_trace_values)status=cudaMalloc(reinterpret_cast(&trace_chunk),(chunk_rows+1)*main->width*sizeof(uint64_t)); + const bool generated = !main->trace_values && main->trace_writer; + const bool tiled = main->host_trace_values || generated; + if(status==cudaSuccess&&tiled)status=cudaMalloc(reinterpret_cast(&trace_chunk),(chunk_rows+1)*main->width*sizeof(uint64_t)); for(size_t row_start=0;status==cudaSuccess&&row_starttrace_values; - if(main->host_trace_values){ + if(generated){ + status=static_cast(main->trace_writer(main->trace_context,device_id,trace_chunk,row_start,rows+1)); + active_trace=trace_chunk; + } else if(main->host_trace_values){ status=cudaMemcpy(trace_chunk,main->host_trace_values+row_start*main->width,rows*main->width*sizeof(uint64_t),cudaMemcpyHostToDevice); const size_t next_row=(row_start+rows)&(height-1); if(status==cudaSuccess)status=cudaMemcpy(trace_chunk+rows*main->width,main->host_trace_values+next_row*main->width,main->width*sizeof(uint64_t),cudaMemcpyHostToDevice); @@ -3287,7 +3071,7 @@ extern "C" int multi_stark_cuda_lookup_graph_lde(int device_id,void** output_han } if(status==cudaSuccess){lookup_messages_graph<<(chunk_blocks),static_cast(tile),global?0:slot_count*tile*sizeof(uint64_t)>>>( conjugates,norms,multiplicities,dn,node_count,slot_count,dl,lookup_count,da,prep,active_trace,main->width, - main->host_trace_values!=nullptr,{beta[0],beta[1]},{gamma[0],gamma[1]},ext_w,height,row_start,rows,scratch);status=cudaGetLastError();} + tiled,{beta[0],beta[1]},{gamma[0],gamma[1]},ext_w,height,row_start,rows,scratch);status=cudaGetLastError();} if(status==cudaSuccess)status=batch_inverse_norms(norm_inverses,norms,messages); if(status==cudaSuccess){lookup_group_deltas_batched<<>>( deltas+row_start*groups,multiplicities,conjugates,norm_inverses,rows,lookup_count,group_size,ext_w);status=cudaGetLastError();} @@ -3295,14 +3079,7 @@ extern "C" int multi_stark_cuda_lookup_graph_lde(int device_id,void** output_han if(status==cudaSuccess)status=exclusive_scan_ext2(reinterpret_cast(lde->values),deltas,count); if(status==cudaSuccess)status=cudaMemcpy(total,lde->values+2*(count-1),sizeof(Ext2),cudaMemcpyDeviceToHost); if(status==cudaSuccess)status=cudaMemcpy(total+2,deltas+count-1,sizeof(Ext2),cudaMemcpyDeviceToHost); - const uint64_t *dit=nullptr,*dshift=nullptr,*dft=nullptr; - if(status==cudaSuccess)status=cached_device_constants(device_id,inverse_twiddles,height/2,1,0,0,&dit); - if(status==cudaSuccess)status=cached_device_constants(device_id,shift_powers,height,3,shift_powers[0],height>1?shift_powers[1]:0,&dshift); - if(status==cudaSuccess)status=cached_device_constants(device_id,forward_twiddles,extended_height/2,2,0,0,&dft); - if(status==cudaSuccess)status=launch_dif(lde->values,height,width,dit); - if(status==cudaSuccess){bit_reverse_scale_and_shift<<>>(lde->values,height,width,strict_log2(height),height_inverse,dshift);status=cudaGetLastError();} - if(status==cudaSuccess)status=launch_dif(lde->values,extended_height,width,dft); - if(status==cudaSuccess){canonicalize_goldilocks<<>>(lde->values,extended_height*width);status=cudaGetLastError();} + if(status==cudaSuccess)status=coset_lde(device_id,lde->values,lde->values,height,width,extended_height,plan); if(status==cudaSuccess)status=cudaStreamSynchronize(0); if(status==cudaSuccess)*output_handle=lde;else destroy_resident_lde(lde); cudaFree(trace_chunk);cudaFree(scratch);cudaFree(deltas);cudaFree(multiplicities);cudaFree(norm_inverses);cudaFree(norms);cudaFree(conjugates);cudaFree(metadata); @@ -3313,11 +3090,10 @@ extern "C" int multi_stark_cuda_lookup_lde(int device_id,void** output_handle,ui const uint64_t* multiplicities,const uint64_t* args,const size_t* arg_offsets, size_t height,size_t num_lookups,size_t args_width,size_t group_size, const uint64_t* beta,const uint64_t* gamma,uint64_t ext_w,size_t added_bits, - const uint64_t* inverse_twiddles,const uint64_t* shift_powers, - const uint64_t* forward_twiddles,uint64_t height_inverse){ + const MultiStarkNttPlan* plan){ if(!output_handle||!total||!multiplicities||!arg_offsets||!height||!num_lookups|| - !group_size||!beta||!gamma||!inverse_twiddles||!shift_powers|| - !forward_twiddles||(args_width&&!args)||!is_power_of_two(height)|| + !group_size||!beta||!gamma||(!plan || !plan->shift_powers)|| + (args_width&&!args)||!is_power_of_two(height)|| added_bits>=sizeof(size_t)*8||height>(SIZE_MAX>>added_bits)) return static_cast(cudaErrorInvalidValue); *output_handle=nullptr;const size_t slots=(num_lookups+group_size-1)/group_size; @@ -3361,14 +3137,7 @@ extern "C" int multi_stark_cuda_lookup_lde(int device_id,void** output_handle,ui if(status==cudaSuccess)status=cudaMemcpy(total,lde->values+2*(count-1),sizeof(Ext2),cudaMemcpyDeviceToHost); if(status==cudaSuccess)status=cudaMemcpy(total+2,deltas+count-1,sizeof(Ext2),cudaMemcpyDeviceToHost); scan_done=now(); - const uint64_t *dit=nullptr,*dshift=nullptr,*dft=nullptr; - if(status==cudaSuccess)status=cached_device_constants(device_id,inverse_twiddles,height/2,1,0,0,&dit); - if(status==cudaSuccess)status=cached_device_constants(device_id,shift_powers,height,3,shift_powers[0],height>1?shift_powers[1]:0,&dshift); - if(status==cudaSuccess)status=cached_device_constants(device_id,forward_twiddles,extended_height/2,2,0,0,&dft); - if(status==cudaSuccess)status=launch_dif(lde->values,height,width,dit); - if(status==cudaSuccess){bit_reverse_scale_and_shift<<>>(lde->values,height,width,strict_log2(height),height_inverse,dshift);status=cudaGetLastError();} - if(status==cudaSuccess)status=launch_dif(lde->values,extended_height,width,dft); - if(status==cudaSuccess){canonicalize_goldilocks<<>>(lde->values,extended_height*width);status=cudaGetLastError();} + if(status==cudaSuccess)status=coset_lde(device_id,lde->values,lde->values,height,width,extended_height,plan); if(status==cudaSuccess)status=cudaStreamSynchronize(0); if(profile){const double finished=now();fprintf(stderr, "[multi-stark/cuda] lookup phases: height=%zu lookups=%zu slots=%zu args_width=%zu allocate=%.3fs rows=%.3fs scan=%.3fs dft=%.3fs\n", @@ -3520,11 +3289,9 @@ extern "C" int multi_stark_cuda_lookup_lde_cpu_rows_partitioned( extern "C" int multi_stark_cuda_lookup_lde_finish_partitioned( int device_id, void* pending_handle, void** output_handle, uint64_t* total, - const uint64_t* inverse_twiddles, const uint64_t* shift_powers, - const uint64_t* forward_twiddles, uint64_t height_inverse) { + const MultiStarkNttPlan* plan) { auto* pending = static_cast(pending_handle); - if (!pending || !output_handle || !total || !inverse_twiddles || - !shift_powers || !forward_twiddles) { + if (!pending || !output_handle || !total || (!plan || !plan->shift_powers)) { return static_cast(cudaErrorInvalidValue); } *output_handle = nullptr; @@ -3540,38 +3307,9 @@ extern "C" int multi_stark_cuda_lookup_lde_finish_partitioned( status = cudaMemcpy(total + 2, pending->deltas + count - 1, sizeof(Ext2), cudaMemcpyDeviceToHost); - const uint64_t* inverse = nullptr; - const uint64_t* shifts = nullptr; - const uint64_t* forward = nullptr; - if (status == cudaSuccess) - status = cached_device_constants(device_id, inverse_twiddles, - pending->height / 2, 1, 0, 0, &inverse); if (status == cudaSuccess) - status = cached_device_constants(device_id, shift_powers, - pending->height, 3, shift_powers[0], - pending->height > 1 ? shift_powers[1] : 0, - &shifts); - if (status == cudaSuccess) - status = cached_device_constants(device_id, forward_twiddles, - pending->extended_height / 2, 2, 0, 0, - &forward); - const size_t width = 2 * pending->slots; - if (status == cudaSuccess) - status = launch_dif(pending->lde->values, pending->height, width, inverse); - if (status == cudaSuccess) { - bit_reverse_scale_and_shift<<height * width), THREADS>>>( - pending->lde->values, pending->height, width, - strict_log2(pending->height), height_inverse, shifts); - status = cudaGetLastError(); - } - if (status == cudaSuccess) - status = launch_dif(pending->lde->values, pending->extended_height, - width, forward); - if (status == cudaSuccess) { - canonicalize_goldilocks<<extended_height * width), THREADS>>>( - pending->lde->values, pending->extended_height * width); - status = cudaGetLastError(); - } + status = coset_lde(device_id, pending->lde->values, pending->lde->values, + pending->height, pending->lde->width, pending->lde->height, plan); if (status == cudaSuccess) status = cudaStreamSynchronize(cudaStreamPerThread); if (status == cudaSuccess) { *output_handle = pending->lde; @@ -3615,7 +3353,7 @@ extern "C" int multi_stark_cuda_lde_release_trace(int device_id,void* handle){ if(!handle)return static_cast(cudaSuccess);cudaError_t status=cudaSetDevice(device_id); auto* lde=static_cast(handle);if(status==cudaSuccess&&lde->trace_values){status=cudaFree(lde->trace_values);lde->trace_values=nullptr;} if(status==cudaSuccess&&lde->host_trace_registered){status=cudaHostUnregister(const_cast(lde->host_trace_values));lde->host_trace_registered=false;} - if(status==cudaSuccess){lde->host_trace_values=nullptr;lde->trace_height=0;} + if(status==cudaSuccess){lde->host_trace_values=nullptr;if(!lde->trace_writer)lde->trace_height=0;} return static_cast(status); } @@ -3648,7 +3386,7 @@ extern "C" int multi_stark_cuda_lde_release_values(int device_id, void* handle) } if (status == cudaSuccess) { lde->host_trace_values = nullptr; - lde->trace_height = 0; + if (!lde->trace_writer) lde->trace_height = 0; } return static_cast(status); } @@ -4445,3 +4183,4 @@ extern "C" int multi_stark_cuda_memory_info(int device_id, size_t* free_bytes, } return static_cast(status); } + diff --git a/cuda/ntt.cuh b/cuda/ntt.cuh new file mode 100644 index 0000000..c23654e --- /dev/null +++ b/cuda/ntt.cuh @@ -0,0 +1,23 @@ +// SPDX-License-Identifier: MIT OR Apache-2.0 +#pragma once +#include +#include + +// Borrowed from the Rust plan; its host constants outlive the synchronous +// prover entry that uploads them. All sizes and groups are fixed at admission. +struct MultiStarkNttPlan { + size_t height; + size_t width; + size_t extended_height; + size_t columns; + size_t inverse_group; + size_t forward_group; + size_t scratch_bytes; + const uint64_t* shift_powers; +}; + +extern "C" int multi_stark_sppark_forward(int device, uint64_t* values, + const MultiStarkNttPlan* plan); +extern "C" int multi_stark_sppark_coset_lde(int device, const uint64_t* trace, + uint64_t* values, const MultiStarkNttPlan* plan, + const uint64_t* device_shift_powers); diff --git a/cuda/smoke.sh b/cuda/smoke.sh index 0ad061a..ff752b0 100755 --- a/cuda/smoke.sh +++ b/cuda/smoke.sh @@ -17,12 +17,12 @@ echo "MULTI_STARK_CUDA_ARCHS=${MULTI_STARK_CUDA_ARCHS}" cargo clippy --release --locked --all-targets --features parallel,cuda -- -D warnings cargo test --release --locked --features parallel,cuda \ - cuda::tests::goldilocks_field_kernels_match_cpu -- --test-threads=1 + cuda::tests::goldilocks_field_kernels_match_cpu cargo test --release --locked --features parallel,cuda \ - cuda::tests::batched_dft_matches_cpu -- --test-threads=1 + cuda::tests::batched_dft_matches_cpu cargo test --release --locked --features parallel,cuda \ - cuda::tests::coset_lde_matches_cpu_including_storage_layout -- --test-threads=1 -cargo test --release --locked --features parallel,cuda -- --test-threads=1 + cuda::tests::coset_lde_matches_cpu_including_storage_layout +cargo test --release --locked --features parallel,cuda compat_dir="$(mktemp -d)" trap 'rm -rf "$compat_dir"' EXIT diff --git a/cuda/sppark_ntt.cu b/cuda/sppark_ntt.cu new file mode 100644 index 0000000..5a8416f --- /dev/null +++ b/cuda/sppark_ntt.cu @@ -0,0 +1,328 @@ +// Goldilocks transforms on the caller's stream, with immutable panel plans. +#include +#include +#include "goldilocks.cuh" +#include "ntt.cuh" + +#include +#include + +namespace { + +bool valid_arguments(const void* d_inout, uint32_t lg, int order, int direction, int coset) { + return d_inout && lg > 0 && lg <= MAX_LG_DOMAIN_SIZE && order >= 0 && order <= 3 && + direction >= 0 && direction <= 1 && coset >= 0 && coset <= 1; +} + +} // namespace + +extern "C" int multi_stark_sppark_max_lg_domain() { return MAX_LG_DOMAIN_SIZE; } + +// `batch` in-place transforms of 2^lg field elements each, `stride` +// elements apart from `d_inout` on `device`, in one launch sequence. +// `order` is NTT::InputOutputOrder (NN, NR, RN, RR), `direction` 0 forward +// or 1 inverse, `coset` 1 for the multiplicative coset by the field +// generator. Inverse transforms are normalized by 1/2^lg upstream. +static int ntt_batch_on_stream(int device, uint64_t* d_inout, uint32_t lg, + int order, int direction, int coset, + uint32_t batch, size_t stride, cudaStream_t caller) { + // Upstream indexes with index_t, 32 bits at this domain limit. + if (!valid_arguments(d_inout, lg, order, direction, coset) || batch == 0 || batch > 65535 || + (batch > 1 && stride < (size_t(1) << lg)) || stride > size_t(std::numeric_limits::max())) + return static_cast(cudaErrorInvalidValue); + cudaError_t status = cudaSetDevice(device); + if (status != cudaSuccess) return static_cast(status); + // Upstream keeps only the devices it supports, so its logical index can + // differ from the CUDA ordinal; find the entry by ordinal. + const gpu_t* found = nullptr; + for (const gpu_t* candidate : all_gpus()) + if (candidate->cid() == device) found = candidate; + if (!found) return static_cast(cudaErrorInvalidDevice); + const gpu_t& gpu = select_gpu(found->id()); + stream_t stream(gpu.id(), caller); + NTT::Base_dev_ptr_batch(stream, reinterpret_cast(d_inout), lg, + static_cast(order), + static_cast(direction), + static_cast(coset), batch, stride); + return static_cast(cudaSuccess); +} + +static int ntt_batch_device(int device, uint64_t* data, uint32_t lg, int order, int direction, + int coset, uint32_t batch, size_t stride) { + return ntt_batch_on_stream(device, data, lg, order, direction, coset, batch, stride, cudaStreamPerThread); +} + +// A fresh caller-owned stream carrying producer, transforms and consumer. +// Reusing it after each borrowed wrapper is dropped checks its ownership. +extern "C" int multi_stark_sppark_borrowed_round_trip(int device, uint64_t* values, uint32_t lg) { + cudaError_t status = cudaSetDevice(device); + if (status != cudaSuccess) return static_cast(status); + cudaStream_t caller = nullptr; + status = cudaStreamCreateWithFlags(&caller, cudaStreamNonBlocking); + if (status != cudaSuccess) return static_cast(status); + uint64_t* data = nullptr; + const size_t count = size_t(1) << lg, bytes = count * sizeof(uint64_t); + status = cudaMallocAsync(reinterpret_cast(&data), bytes, caller); + if (status == cudaSuccess) status = cudaMemcpyAsync(data, values, bytes, cudaMemcpyHostToDevice, caller); + int result = static_cast(status); + if (result == 0) result = ntt_batch_on_stream(device, data, lg, 1, 0, 0, 1, count, caller); + if (result == 0) result = ntt_batch_on_stream(device, data, lg, 2, 1, 0, 1, count, caller); + if (result == 0) result = static_cast(cudaMemcpyAsync(values, data, bytes, cudaMemcpyDeviceToHost, caller)); + if (data) cudaFreeAsync(data, caller); + const auto synced = cudaStreamSynchronize(caller); + if (result == 0) result = static_cast(synced); + cudaStreamDestroy(caller); + return result; +} + +// The batched transform on host memory: `batch * stride` words uploaded, +// transformed, downloaded and synchronized. For contract checks and small +// inputs, not the prover. +extern "C" int multi_stark_sppark_ntt_batch_host(int device, uint64_t* inout, uint32_t lg, int order, + int direction, int coset, uint32_t batch, + size_t stride) { + if (!valid_arguments(inout, lg, order, direction, coset) || batch == 0 || stride < (size_t(1) << lg)) + return static_cast(cudaErrorInvalidValue); + const size_t bytes = batch * stride * sizeof(uint64_t); + cudaError_t status = cudaSetDevice(device); + if (status != cudaSuccess) return static_cast(status); + uint64_t* d_inout = nullptr; + status = cudaMalloc(reinterpret_cast(&d_inout), bytes); + if (status != cudaSuccess) return static_cast(status); + status = cudaMemcpyAsync(d_inout, inout, bytes, cudaMemcpyHostToDevice, cudaStreamPerThread); + int result = static_cast(status); + if (result == 0) + result = ntt_batch_device(device, d_inout, lg, order, direction, coset, batch, stride); + if (result == 0) + result = static_cast( + cudaMemcpyAsync(inout, d_inout, bytes, cudaMemcpyDeviceToHost, cudaStreamPerThread)); + const cudaError_t synced = cudaStreamSynchronize(cudaStreamPerThread); + if (result == 0) result = static_cast(synced); + cudaFree(d_inout); + return result; +} + +// One transform on host memory: the batch entry with a single vector. +extern "C" int multi_stark_sppark_ntt_host(int device, uint64_t* inout, uint32_t lg, int order, + int direction, int coset) { + return multi_stark_sppark_ntt_batch_host(device, inout, lg, order, direction, coset, 1, size_t(1) << lg); +} + +// Row-major matrices are gathered into contiguous column panels for the +// inverse NTT, coefficient-order restoration with coset shift and expansion, +// and forward NR transform. Scatter preserves the bit-reversed row storage +// consumed by commitments. Gather canonicalizes lazy field representatives. + +namespace { + +constexpr unsigned PANEL_THREADS = 256; +constexpr size_t MAX_BLOCKS = 65535; +constexpr unsigned TILE = 32; +constexpr unsigned TILE_ROWS = 8; +// Matrices narrower than this take the row-per-thread gather and scatter, +// whose per-thread row segments are then a few contiguous words; from here +// up the tiled transposes win. +constexpr size_t NARROW_MATRIX = 8; + +unsigned blocks_for_total(size_t total) { + const size_t blocks = (total + PANEL_THREADS - 1) / PANEL_THREADS; + return static_cast(blocks < MAX_BLOCKS ? blocks : MAX_BLOCKS); +} + +// Tiles over rows on the grid's first dimension, which is wide enough for +// 2^26 rows, and over columns on the second. +dim3 tile_grid(size_t rows, size_t columns) { + return dim3(static_cast((rows + TILE - 1) / TILE), static_cast((columns + TILE - 1) / TILE)); +} + +__device__ __forceinline__ size_t reverse_bits(size_t index, unsigned log) { + return log == 0 ? 0 : static_cast(__brev(static_cast(index)) >> (32 - log)); +} + +// Row-major rows [0, height) of columns [first, first + count) into the +// column-major panel (columns `column_stride` apart), canonical, through +// a shared-memory tile so both sides are coalesced. +__global__ void gather_tiles(const uint64_t* __restrict__ source, size_t height, size_t width, size_t first, + size_t count, size_t column_stride, uint64_t* __restrict__ panel) { + __shared__ uint64_t tile[TILE][TILE + 1]; + const size_t row0 = static_cast(blockIdx.x) * TILE; + const size_t col0 = static_cast(blockIdx.y) * TILE; + for (unsigned r = threadIdx.y; r < TILE; r += TILE_ROWS) { + const size_t row = row0 + r, col = col0 + threadIdx.x; + if (row < height && col < count) + tile[r][threadIdx.x] = multi_stark_cuda::canonicalize(source[row * width + first + col]); + } + __syncthreads(); + for (unsigned c = threadIdx.y; c < TILE; c += TILE_ROWS) { + const size_t col = col0 + c, row = row0 + threadIdx.x; + if (row < height && col < count) panel[col * column_stride + row] = tile[threadIdx.x][c]; + } +} + +// The column panel back into row-major storage without a row permutation. +__global__ void scatter_tiles(const uint64_t* __restrict__ panel, size_t rows, size_t column_stride, size_t width, + size_t first, size_t count, + uint64_t* __restrict__ values) { + __shared__ uint64_t tile[TILE][TILE + 1]; + const size_t row0 = static_cast(blockIdx.x) * TILE; + const size_t col0 = static_cast(blockIdx.y) * TILE; + for (unsigned c = threadIdx.y; c < TILE; c += TILE_ROWS) { + const size_t col = col0 + c, row = row0 + threadIdx.x; + if (row < rows && col < count) tile[c][threadIdx.x] = panel[col * column_stride + row]; + } + __syncthreads(); + for (unsigned r = threadIdx.y; r < TILE; r += TILE_ROWS) { + const size_t row = row0 + r, col = col0 + threadIdx.x; + if (row < rows && col < count) { + values[row * width + first + col] = tile[threadIdx.x][r]; + } + } +} + +// The expansion with the coefficient order restored: B[c][i] = +// A[c][rev(i)] * shift^i for i < N, zero beyond, so the forward transform +// runs in NR order and the scatter stays natural. The permuted reads are +// cheap while the compact panel fits the L2 cache. One column per grid row. +__global__ void shift_columns(const uint64_t* __restrict__ coefficients, const uint64_t* __restrict__ shift_powers, + size_t height, unsigned log_height, size_t extended_height, + uint64_t* __restrict__ panel) { + const size_t column = blockIdx.y; + const size_t stride = static_cast(blockDim.x) * gridDim.x; + for (size_t i = static_cast(blockIdx.x) * blockDim.x + threadIdx.x; i < extended_height; i += stride) { + uint64_t value = 0; + if (i < height) + value = multi_stark_cuda::goldilocks_mul(coefficients[column * height + reverse_bits(i, log_height)], + shift_powers[i]); + panel[column * extended_height + i] = value; + } +} + +// The gather and scatter for a narrow panel: one thread per row moves its +// few words, the panel side coalesced across the block; the tiled +// transpose would leave most of its block idle on such panels. +__global__ void gather_rows(const uint64_t* __restrict__ source, size_t height, size_t width, size_t first, + unsigned count, size_t column_stride, uint64_t* __restrict__ panel) { + const size_t stride = static_cast(blockDim.x) * gridDim.x; + for (size_t row = static_cast(blockIdx.x) * blockDim.x + threadIdx.x; row < height; row += stride) { + const uint64_t* in = source + row * width + first; + for (unsigned c = 0; c < count; ++c) panel[c * column_stride + row] = multi_stark_cuda::canonicalize(in[c]); + } +} + +__global__ void scatter_rows(const uint64_t* __restrict__ panel, size_t rows, size_t column_stride, size_t width, + size_t first, unsigned count, + uint64_t* __restrict__ values) { + const size_t stride = static_cast(blockDim.x) * gridDim.x; + for (size_t row = static_cast(blockIdx.x) * blockDim.x + threadIdx.x; row < rows; row += stride) { + uint64_t* out = values + row * width + first; + for (unsigned c = 0; c < count; ++c) out[c] = panel[c * column_stride + row]; + } +} + +// The panel's gather and scatter by the matrix width. +cudaError_t gather_panel(const uint64_t* source, size_t height, size_t width, size_t first, size_t count, + size_t column_stride, uint64_t* panel) { + if (width < NARROW_MATRIX) + gather_rows<<>>( + source, height, width, first, static_cast(count), column_stride, panel); + else + gather_tiles<<>>( + source, height, width, first, count, column_stride, panel); + return cudaGetLastError(); +} + +cudaError_t scatter_panel(const uint64_t* panel, size_t rows, size_t column_stride, size_t width, size_t first, + size_t count, uint64_t* values) { + if (width < NARROW_MATRIX) + scatter_rows<<>>( + panel, rows, column_stride, width, first, static_cast(count), values); + else + scatter_tiles<<>>( + panel, rows, column_stride, width, first, count, values); + return cudaGetLastError(); +} + +unsigned log2_exact(size_t value) { + unsigned log = 0; + while ((size_t(1) << log) < value) ++log; + return log; +} + +// The `count` columns of a panel, `stride` elements apart, through batched +// launch sequences of as many columns as the group budget holds. +int transform_columns(int device, uint64_t* panel, uint32_t lg, int order, int direction, size_t count, + size_t stride, size_t group) { + if (lg == 0) return 0; + int result = 0; + for (size_t first = 0; result == 0 && first < count; first += group) { + const size_t batch = count - first < group ? count - first : group; + result = ntt_batch_device(device, panel + first * stride, lg, order, direction, 0, + static_cast(batch), stride); + } + return result; +} + +} // namespace + +extern "C" int multi_stark_sppark_l2_bytes(int device, size_t* bytes) { + int l2 = 0; + const cudaError_t status = cudaDeviceGetAttribute(&l2, cudaDevAttrL2CacheSize, device); + if (status == cudaSuccess) *bytes = l2 > 0 ? size_t(l2) : size_t(64) << 20; + return static_cast(status); +} + +extern "C" int multi_stark_sppark_coset_lde(int device, const uint64_t* trace, uint64_t* values, + const MultiStarkNttPlan* plan, + const uint64_t* shift_powers) { + const size_t height = plan->height, width = plan->width; + const size_t extended_height = plan->extended_height, columns = plan->columns; + const unsigned log_height = log2_exact(height), log_extended = log2_exact(extended_height); + cudaError_t status = cudaSetDevice(device); + if (status != cudaSuccess) return static_cast(status); + // Pool allocation and freeing follow all panel kernels on the same stream. + // CUPTI 2026.2.1 faults when tracing an async free of non-pool memory. + uint64_t* scratch = nullptr; + status = cudaMallocAsync(reinterpret_cast(&scratch), plan->scratch_bytes, cudaStreamPerThread); + if (status != cudaSuccess) return static_cast(status); + uint64_t* a = scratch; + uint64_t* b = scratch + columns * height; + int result = 0; + for (size_t first = 0; result == 0 && first < width; first += columns) { + const size_t count = width - first < columns ? width - first : columns; + result = static_cast(gather_panel(trace, height, width, first, count, height, a)); + if (result == 0) result = transform_columns(device, a, log_height, 1, 1, count, height, plan->inverse_group); + if (result == 0) { + const dim3 column_grid(blocks_for_total(extended_height), static_cast(count)); + shift_columns<<>>( + a, shift_powers, height, log_height, extended_height, b); + result = static_cast(cudaGetLastError()); + } + if (result == 0) result = transform_columns(device, b, log_extended, 1, 0, count, extended_height, plan->forward_group); + if (result == 0) + result = static_cast(scatter_panel(b, extended_height, extended_height, width, first, count, values)); + } + const cudaError_t freed = cudaFreeAsync(scratch, cudaStreamPerThread); + if (result == 0) result = static_cast(freed); + return result; +} + +extern "C" int multi_stark_sppark_forward(int device, uint64_t* values, const MultiStarkNttPlan* plan) { + const size_t height = plan->height, width = plan->width, columns = plan->columns; + const unsigned log_height = log2_exact(height); + cudaError_t status = cudaSetDevice(device); + if (status != cudaSuccess) return static_cast(status); + uint64_t* panel = nullptr; + status = cudaMallocAsync(reinterpret_cast(&panel), plan->scratch_bytes, cudaStreamPerThread); + if (status != cudaSuccess) return static_cast(status); + int result = 0; + for (size_t first = 0; result == 0 && first < width; first += columns) { + const size_t count = width - first < columns ? width - first : columns; + result = static_cast(gather_panel(values, height, width, first, count, height, panel)); + if (result == 0) result = transform_columns(device, panel, log_height, 1, 0, count, height, plan->forward_group); + if (result == 0) + result = static_cast(scatter_panel(panel, height, height, width, first, count, values)); + } + const cudaError_t freed = cudaFreeAsync(panel, cudaStreamPerThread); + if (result == 0) result = static_cast(freed); + return result; +} diff --git a/examples/cuda_commit_bench.rs b/examples/cuda_commit_bench.rs new file mode 100644 index 0000000..5b22110 --- /dev/null +++ b/examples/cuda_commit_bench.rs @@ -0,0 +1,90 @@ +//! Complete PCS commitment timings: host upload, LDEs, and mixed-height Merkle tree. +//! Run identical binaries before/after a kernel change and compare commitments. +//! Witness allocation is outside the timer; cold iteration zero is reported separately. +//! +//! cargo run --release --locked --features parallel,cuda --example cuda_commit_bench +//! MULTI_STARK_CUDA_BENCH_ITERATIONS controls the number of warm iterations (default 5). + +use std::hint::black_box; +use std::time::Instant; + +use multi_stark::config::StarkGenericConfig; +use multi_stark::types::{ + Challenger, CommitmentParameters, ExtVal, FriParameters, GoldilocksBlake3Config, Pcs, Val, +}; +use p3_commit::Pcs as PcsTrait; +use p3_field::PrimeCharacteristicRing; +use p3_matrix::dense::RowMajorMatrix; + +fn main() { + let iterations: usize = std::env::var("MULTI_STARK_CUDA_BENCH_ITERATIONS") + .map_or(5, |s| s.parse().expect("invalid iteration count")); + let config = GoldilocksBlake3Config::new( + CommitmentParameters { + log_blowup: 1, + cap_height: 0, + }, + FriParameters { + log_final_poly_len: 0, + max_log_arity: 1, + num_queries: 2, + commit_proof_of_work_bits: 0, + query_proof_of_work_bits: 0, + }, + ); + let pcs = config.pcs(); + // Equal-height matrices are hashed as concatenated rows. The mixed case + // covers both dispatches and lower-height row injections in one commitment. + let shapes: &[(&str, &[(usize, usize)])] = &[ + ("narrow", &[(20, 8)]), + ("medium", &[(18, 40)]), + ("chunk_boundary", &[(18, 128)]), + ("above_boundary", &[(18, 129)]), + ("wide", &[(16, 925)]), + ( + "mixed", + &[(18, 4), (18, 12), (18, 24), (17, 129), (16, 533)], + ), + ]; + println!("shape,iteration,seconds,commitment"); + for &(name, shape) in shapes { + let mut expected = None; + for iteration in 0..=iterations { + let inputs = shape + .iter() + .enumerate() + .map(|(matrix_index, &(log_height, width))| { + let height = 1 << log_height; + let values = (0..height * width) + .map(|index| { + let value = (index as u64) + .wrapping_mul(0x9e37_79b9_7f4a_7c15) + .wrapping_add(matrix_index as u64); + Val::from_u64(value) + }) + .collect(); + let domain = >::natural_domain_for_degree( + pcs, height, + ); + (domain, RowMajorMatrix::new(values, width)) + }) + .collect::>(); + let started = Instant::now(); + let (commitment, data) = >::commit(pcs, inputs); + let seconds = started.elapsed().as_secs_f64(); + if let Some(expected) = &expected { + assert_eq!(&commitment, expected); + } else { + expected = Some(commitment.clone()); + } + let bytes = + bincode::serde::encode_to_vec(&commitment, bincode::config::standard()).unwrap(); + let hex: String = bytes.iter().map(|byte| format!("{byte:02x}")).collect(); + println!("{name},{iteration},{seconds:.9},{hex}"); + // Retain all prover data until after timing: committing must not + // silently include materialization or opening work from a consumer. + black_box(&data); + drop(data); + } + } +} diff --git a/examples/cuda_resident_lde_bench.rs b/examples/cuda_resident_lde_bench.rs new file mode 100644 index 0000000..6e8ec0e --- /dev/null +++ b/examples/cuda_resident_lde_bench.rs @@ -0,0 +1,90 @@ +//! Resident LDE timings including upload, excluding host input construction. +//! Correctness is covered by the `cuda::sppark` tests and `proof_compatibility`. + +#[cfg(not(feature = "cuda"))] +fn main() { + eprintln!("enable --features cuda"); +} + +#[cfg(feature = "cuda")] +fn main() { + use multi_stark::cuda::{CudaDft, pcs::CudaPcsDft}; + use p3_field::{Field, PrimeCharacteristicRing}; + use p3_goldilocks::Goldilocks; + use p3_matrix::dense::RowMajorMatrix; + use std::hint::black_box; + use std::time::Instant; + + let iterations: usize = std::env::var("MULTI_STARK_CUDA_BENCH_ITERATIONS") + .map_or(7, |s| s.parse().expect("invalid iteration count")); + let gpu = CudaDft::default(); + // MULTI_STARK_CUDA_BENCH_SHAPES="20,533,2;24,6,2" runs those shapes only. + let only: Vec<(usize, usize, usize)> = std::env::var("MULTI_STARK_CUDA_BENCH_SHAPES") + .map(|list| { + list.split(';') + .map(|shape| { + let mut parts = shape + .split(',') + .map(|part| part.trim().parse().expect("invalid shape")); + ( + parts.next().expect("log height"), + parts.next().expect("width"), + parts.next().expect("added bits"), + ) + }) + .collect() + }) + .unwrap_or_default(); + println!("log_height,width,added_bits,iteration,seconds"); + // The last five are the shapes the profiled Init proofs commit most: + // BLAKE3 pieces, wide and narrow IxVM circuits at the height cap, and + // the width-2 codewords of the quotient and narrow lookups. + for (log_height, width, added_bits) in [ + (20, 1, 1), + (20, 2, 1), + (20, 8, 1), + (18, 40, 1), + (18, 128, 1), + (18, 129, 1), + (16, 925, 1), + (18, 40, 2), + (20, 533, 2), + (24, 6, 2), + (24, 17, 2), + (22, 2, 2), + (20, 2, 2), + // Short aggregation tables and narrow lookup and quotient codewords. + (8, 20, 2), + (10, 33, 2), + (12, 8, 2), + (14, 40, 2), + (15, 129, 2), + (16, 2, 2), + (17, 16, 2), + (18, 2, 2), + (19, 2, 2), + (19, 49, 2), + ] { + if !only.is_empty() && !only.contains(&(log_height, width, added_bits)) { + continue; + } + let height = 1 << log_height; + let values = (0..height * width) + .map(|index| Goldilocks::from_u64((index as u64).wrapping_mul(0x9e37_79b9_7f4a_7c15))) + .collect(); + let matrix = RowMajorMatrix::new(values, width); + for iteration in 0..=iterations { + let started = Instant::now(); + let lde = CudaPcsDft::coset_lde_batch_resident( + &gpu, + &matrix, + added_bits, + Goldilocks::GENERATOR, + ); + let seconds = started.elapsed().as_secs_f64(); + println!("{log_height},{width},{added_bits},{iteration},{seconds:.9}"); + black_box(&lde); + drop(lde); + } + } +} diff --git a/src/batch.rs b/src/batch.rs index abbefa4..54b6da7 100644 --- a/src/batch.rs +++ b/src/batch.rs @@ -478,7 +478,7 @@ where /// # Panics /// Panics if `claims` is empty, if any shard's traces are all empty, or /// if a regenerated shard does not reproduce its round-one header. - pub fn prove_batch_with( + pub fn prove_batch_with( &self, key: &ProverKey, claims: &[Vec>>], @@ -488,7 +488,8 @@ where ) -> BatchProof where Com: PartialEq, - W: FnMut(usize) -> SystemWitness>, + W: FnMut(usize) -> T, + T: Into>> + Send, Self: Sync, ProverKey: Sync, BatchPreamble: Sync, @@ -518,9 +519,10 @@ where /// the next shard's witness on the calling thread; at most one witness /// is held ahead of the prover. #[tracing::instrument(level = "info", skip_all, name = "stark/batch_round_one")] - pub fn batch_round_one(&self, round_one: I, retention: Retention) -> BatchBarrier + pub fn batch_round_one(&self, round_one: I, retention: Retention) -> BatchBarrier where - I: IntoIterator>>, SystemWitness>)>, + I: IntoIterator>>, T)>, + T: Into>> + Send, Self: Sync, SystemWitness>: Send, Stage1: Send, @@ -559,7 +561,7 @@ where /// on the calling thread for every shard round one did not retain; at /// most one witness is held ahead of the prover. #[tracing::instrument(level = "info", skip_all, name = "stark/batch_round_two")] - pub fn batch_round_two( + pub fn batch_round_two( &self, key: &ProverKey, barrier: BatchBarrier, @@ -568,7 +570,8 @@ where ) -> BatchProof where Com: PartialEq, - W: FnMut(usize) -> SystemWitness>, + W: FnMut(usize) -> T, + T: Into>> + Send, Self: Sync, ProverKey: Sync, BatchPreamble: Sync, diff --git a/src/config.rs b/src/config.rs index bca90b2..6dcd85b 100644 --- a/src/config.rs +++ b/src/config.rs @@ -175,6 +175,21 @@ pub trait StarkGenericConfig { /// evaluations. fn log_blowup(&self) -> usize; + /// Commit deterministic main-trace sources. + fn commit_main( + &self, + evaluations: Vec<(Domain, crate::witness::TraceSource>)>, + ) -> (Com, PcsData) + where + Self: Sized, + { + self.pcs().commit( + evaluations + .into_iter() + .map(|(domain, trace)| (domain, trace.materialize())), + ) + } + /// Normalize any backend-dependent field representatives before proof /// serialization. Most fields have canonical in-memory representations; /// configurations whose field permits lazy reduction can override this. diff --git a/src/cuda/mmcs.rs b/src/cuda/mmcs.rs index 7b419b6..41e7d7f 100644 --- a/src/cuda/mmcs.rs +++ b/src/cuda/mmcs.rs @@ -108,6 +108,38 @@ fn hybrid_resident_candidates( .collect() } +/// Frees the device caches of generator-backed resident LDEs, one at a time +/// until `target_free_bytes` is met, and returns the free bytes then. A cache +/// is the cheapest headroom there is: the generator rebuilds it from host +/// seeds on demand, whereas an evicted LDE must be materialized and uploaded +/// again. +fn release_generator_caches( + device_id: i32, + data: &CudaMmcsData, + protected_index: Option, + target_free_bytes: usize, +) -> usize { + let mut free_bytes = super::device_memory_info(device_id).0; + let CudaMmcsData::Hybrid { resident, .. } = data else { + return free_bytes; + }; + for (index, lde) in resident.iter().enumerate() { + if free_bytes >= target_free_bytes { + break; + } + if protected_index == Some(index) { + continue; + } + let Some(lde) = lde else { continue }; + if !lde.has_generator() { + continue; + } + lde.release_generator_device(); + free_bytes = super::device_memory_info(device_id).0; + } + free_bytes +} + fn evict_hybrid_resident(data: &CudaMmcsData, index: usize) { let CudaMmcsData::Hybrid { resident, @@ -126,6 +158,11 @@ fn evict_hybrid_resident(data: &CudaMmcsData, index: usize) { // SAFETY: admission transitions run between proving stages, when no CUDA // operation can access this LDE or its trace. unsafe { lde.release_values() }; + lde.release_generator_device(); + super::witness::record_lde_spill( + lde.device_id, + lde.height() * lde.width() * size_of::(), + ); resident_active[index].store(false, std::sync::atomic::Ordering::Release); } @@ -153,6 +190,46 @@ impl CudaMmcsData { } impl CudaMmcsData> { + /// Whether [`Self::resident_with_trace`] returns an LDE for this matrix, + /// without attaching anything. A hybrid LDE qualifies while its values are + /// resident, and after a spill as long as it can regenerate its trace from + /// a generator or a retained host trace. Admission decisions that size a + /// budget for the graph kernel must use this predicate, so that the budget + /// and the kernel that then runs agree. + pub(crate) fn has_resident_with_trace(&self, index: usize) -> bool { + match self { + Self::Cuda { + resident, + retained_traces, + .. + } => { + index < resident.len() + && retained_traces + .get() + .is_none_or(|retained| index < retained.len()) + } + Self::Hybrid { + resident, + resident_active, + retained_traces, + .. + } => { + let Some(Some(lde)) = resident.get(index) else { + return false; + }; + let (Some(active), Some(retained)) = + (resident_active.get(index), retained_traces.get(index)) + else { + return false; + }; + active.load(std::sync::atomic::Ordering::Acquire) + || lde.has_generator() + || retained.is_some() + } + Self::Cpu(_) => false, + } + } + pub(crate) fn resident_with_trace(&self, index: usize) -> Option<&CudaLde> { match self { Self::Cuda { @@ -175,13 +252,15 @@ impl CudaMmcsData> { retained_traces, .. } => { + let lde = resident.get(index)?.as_ref()?; if !resident_active .get(index)? .load(std::sync::atomic::Ordering::Acquire) + && !lde.has_generator() + && retained_traces.get(index)?.is_none() { return None; } - let lde = resident.get(index)?.as_ref()?; if let Some(trace) = retained_traces.get(index)?.as_ref() { // SAFETY: the trace is owned by this prover data and drops // after the resident LDE which holds the registered pointer. @@ -217,9 +296,33 @@ pub(crate) fn hash_cpu_height_groups( .collect() } +/// The CUDA BLAKE3 leaf kernel hashes one message of at most 32 KiB per +/// row, and a mixed tree's leaf row at a height is every matrix of that +/// height side by side. A height group whose rows are wider than that is +/// hashed on the host and enters the tree as a digest group; the device +/// still builds the rest of the tree. +pub(crate) const MAX_DEVICE_LEAF_ROW_BYTES: usize = 32 * 1024; + +/// The heights whose concatenated leaf row exceeds +/// [`MAX_DEVICE_LEAF_ROW_BYTES`], and so must be hashed on the host. +pub(crate) fn host_hashed_heights( + dimensions: impl IntoIterator, +) -> std::collections::BTreeSet { + let mut row_bytes = std::collections::BTreeMap::::new(); + for dimensions in dimensions { + *row_bytes.entry(dimensions.height).or_default() += + dimensions.width * size_of::(); + } + row_bytes + .into_iter() + .filter(|(_, bytes)| *bytes > MAX_DEVICE_LEAF_ROW_BYTES) + .map(|(height, _)| height) + .collect() +} + pub(crate) fn hash_host_only_height_groups( matrices: &[Option>], - resident: &[Option], + resident: &[Option<&CudaLde>], deferred_dimensions: &[Option], prehashed_heights: &std::collections::BTreeSet, ) -> Vec<(usize, Vec<[u8; 32]>)> { @@ -242,7 +345,7 @@ pub(crate) fn hash_host_only_height_groups( let height = matrix.as_ref().map_or_else( || { lde.as_ref() - .map_or_else(|| deferred.unwrap().height, CudaLde::height) + .map_or_else(|| deferred.unwrap().height, |lde| lde.height()) }, Matrix::height, ); @@ -389,11 +492,13 @@ pub struct CudaMmcs { } impl CudaMmcs { - pub(crate) fn new(cpu: CpuMmcs) -> Self { - Self { - cpu, - device_id: super::configured_device(), - } + /// A commitment scheme resident on the given CUDA device. Every kernel + /// it launches, allocation it makes and buffer it stages through belongs + /// to that device, so several of these in one process can each own a + /// device of their own. + pub(crate) fn with_device(cpu: CpuMmcs, device_id: i32) -> Self { + assert!(device_id >= 0, "CUDA device id must be non-negative"); + Self { cpu, device_id } } } @@ -430,6 +535,7 @@ pub trait CudaCommitMmcs: Mmcs { data: &Self::ProverData, target_free_bytes: usize, protected_index: Option, + phase: &'static str, ) -> usize; /// Selects spill candidates across all supplied commitments instead of @@ -438,6 +544,7 @@ pub trait CudaCommitMmcs: Mmcs { &self, data: &[&Self::ProverData>], target_free_bytes: usize, + phase: &'static str, ) -> usize; fn matrix_dimensions(&self, data: &Self::ProverData>) -> Vec; @@ -587,17 +694,28 @@ impl CudaCommitMmcs for CudaMmcs { data: &Self::ProverData, target_free_bytes: usize, protected_index: Option, + phase: &'static str, ) -> usize { let (mut free_bytes, _) = super::device_memory_info(self.device_id); if super::memory_diagnostics_enabled() { eprintln!( - "[multi-stark/cuda] admission start: target={} free={}", + "[multi-stark/cuda] {phase} admission start: target={} free={}", target_free_bytes, free_bytes ); } if free_bytes >= target_free_bytes { return free_bytes; } + free_bytes = + release_generator_caches(self.device_id, data, protected_index, target_free_bytes); + if free_bytes >= target_free_bytes { + if super::memory_diagnostics_enabled() { + eprintln!( + "[multi-stark/cuda] {phase} admitted after releasing generator caches: free={free_bytes}" + ); + } + return free_bytes; + } let mut candidates = hybrid_resident_candidates(data, protected_index); while free_bytes < target_free_bytes && !candidates.is_empty() { let deficit_cells = target_free_bytes @@ -636,17 +754,30 @@ impl CudaCommitMmcs for CudaMmcs { &self, data: &[&Self::ProverData>], target_free_bytes: usize, + phase: &'static str, ) -> usize { let (mut measured_free_bytes, _) = super::device_memory_info(self.device_id); if super::memory_diagnostics_enabled() { eprintln!( - "[multi-stark/cuda] batch admission start: target={} free={}", + "[multi-stark/cuda] {phase} batch admission start: target={} free={}", target_free_bytes, measured_free_bytes ); } if measured_free_bytes >= target_free_bytes { return measured_free_bytes; } + for data in data { + measured_free_bytes = + release_generator_caches(self.device_id, data, None, target_free_bytes); + if measured_free_bytes >= target_free_bytes { + if super::memory_diagnostics_enabled() { + eprintln!( + "[multi-stark/cuda] {phase} batch admitted after releasing generator caches: free={measured_free_bytes}" + ); + } + return measured_free_bytes; + } + } let initial_free_bytes = measured_free_bytes; let mut released_bytes = 0usize; let mut candidates = data @@ -934,6 +1065,38 @@ impl CudaCommitMmcs for CudaMmcs { Self::Commitment, Self::ProverData>, ) { + // A height group whose leaf row is wider than the device kernel + // hashes cannot be part of an all-resident tree: its LDEs come back + // to the host, its digests are hashed there, and the commitment is + // the hybrid one every other stage already handles. + let wide_heights = host_hashed_heights(ldes.iter().map(|lde| Dimensions { + width: lde.width(), + height: lde.height(), + })); + if !wide_heights.is_empty() { + let mut resident: Vec> = Vec::with_capacity(ldes.len()); + let mut host: Vec>> = Vec::with_capacity(ldes.len()); + for lde in ldes { + if wide_heights.contains(&lde.height()) { + host.push(Some(lde.to_row_major_matrix())); + resident.push(None); + } else { + host.push(None); + resident.push(Some(lde)); + } + } + let deferred: Vec>> = + (0..resident.len()).map(|_| None).collect(); + let deferred_dimensions = vec![None; resident.len()]; + let digests = hash_host_only_height_groups( + &host, + &resident.iter().map(Option::as_ref).collect::>(), + &deferred_dimensions, + &std::collections::BTreeSet::new(), + ); + let retained = (0..resident.len()).map(|_| None).collect(); + return self.commit_cuda_hybrid(resident, host, deferred, retained, digests); + } let ldes = std::sync::Arc::new(ldes); let tree = CudaMixedMerkleTree::from_ldes(self.device_id, &ldes); let commitment = MerkleCap::new(vec![tree.root()]); @@ -1094,11 +1257,20 @@ impl Mmcs for CudaMmcs { &self, inputs: Vec, ) -> (Self::Commitment, Self::ProverData) { - if inputs - .iter() - .map(|matrix| matrix.height().saturating_mul(matrix.width())) - .sum::() - > 1 << 18 + // Small commitments stay on the CPU, and so does one with a height + // group wider than the device leaf kernel hashes (32 KiB per row); + // the CPU MMCS is the reference the device tree reproduces. + let device_hashable = host_hashed_heights(inputs.iter().map(|matrix| Dimensions { + width: matrix.width(), + height: matrix.height(), + })) + .is_empty(); + if device_hashable + && inputs + .iter() + .map(|matrix| matrix.height().saturating_mul(matrix.width())) + .sum::() + > 1 << 18 { let resident: Vec<_> = inputs .iter() @@ -1317,6 +1489,73 @@ mod tests { ) } + #[test] + fn wide_height_groups_are_hashed_on_the_host() { + let d = |width, height| Dimensions { width, height }; + // 4096 columns of 8 bytes is exactly the 32 KiB the leaf kernel takes. + assert!(host_hashed_heights([d(4096, 8)]).is_empty()); + assert_eq!( + host_hashed_heights([d(4097, 8)]) + .into_iter() + .collect::>(), + vec![8] + ); + // Matrices at one height share a leaf row, so their widths add up; + // other heights are judged on their own. + assert_eq!( + host_hashed_heights([d(3000, 16), d(1500, 16), d(9282, 4), d(10, 4), d(14, 4096)]) + .into_iter() + .collect::>(), + vec![4, 16] + ); + } + + /// The dimensions of the Aiur BLAKE3 hashing test's stage-one traces at + /// size 64: the height-4 group is 9,731 columns wide, 78 KB per leaf row. + #[test] + fn resident_commit_with_a_wide_height_group_matches_the_cpu() { + let dims = [ + (4096, 14), + (2048, 14), + (4096, 17), + (4, 37), + (4, 35), + (4, 371), + (4, 9282), + (256, 110), + (4, 6), + (1024, 3), + (262144, 10), + ]; + let matrices: Vec<_> = dims + .iter() + .enumerate() + .map(|(i, &(height, width))| matrix(height, width, 7 * i + 1)) + .collect(); + let cpu = CpuMmcs::new( + SerializingHasher::new(Blake3), + Blake3CompressionFunction::new(Blake3), + 0, + ); + let (expected_commitment, expected_data) = cpu.commit(matrices.clone()); + let indices = [0, 3, 1000, 262143]; + let expected_openings: Vec<_> = indices + .iter() + .map(|&i| cpu.open_batch(i, &expected_data)) + .collect(); + let mmcs = CudaMmcs::with_device(cpu, 0); + let ldes = matrices + .iter() + .map(|m| CudaLde::from_row_major_matrix(mmcs.device_id, m)) + .collect(); + let (commitment, data) = mmcs.commit_cuda_resident(ldes); + assert_eq!(commitment, expected_commitment); + for (&index, expected) in indices.iter().zip(&expected_openings) { + let opening = mmcs.open_batch(index, &data); + assert_eq!(opening.opened_values, expected.opened_values); + } + } + #[test] fn hybrid_openings_survive_device_spill() { let matrices = vec![matrix(16, 2, 3), matrix(8, 3, 71), matrix(16, 1, 109)]; @@ -1328,7 +1567,7 @@ mod tests { let (expected_commitment, expected_data) = cpu.commit(matrices.clone()); let expected_opening = cpu.open_batch(5, &expected_data); - let mmcs = CudaMmcs::new(cpu); + let mmcs = CudaMmcs::with_device(cpu, 0); let resident = vec![ Some(CudaLde::from_row_major_matrix(mmcs.device_id, &matrices[0])), None, @@ -1337,7 +1576,7 @@ mod tests { let host_matrices = vec![None, Some(matrices[1].clone()), Some(matrices[2].clone())]; let host_digest_groups = hash_host_only_height_groups( &host_matrices, - &resident, + &resident.iter().map(Option::as_ref).collect::>(), &[None, None, None], &std::collections::BTreeSet::new(), ); @@ -1362,7 +1601,7 @@ mod tests { expected_opening.opening_proof ); - mmcs.ensure_device_headroom(&data, usize::MAX, None); + mmcs.ensure_device_headroom(&data, usize::MAX, None, "test"); assert!(!(0..matrices.len()).any(|index| mmcs.is_matrix_cuda_resident(&data, index))); let spilled_opening = mmcs.open_batch(5, &data); assert_eq!( diff --git a/src/cuda/mod.rs b/src/cuda/mod.rs index a59619b..4d12da4 100644 --- a/src/cuda/mod.rs +++ b/src/cuda/mod.rs @@ -8,12 +8,14 @@ pub(crate) mod mmcs; #[doc(hidden)] pub mod pcs; +pub(crate) mod sppark; +pub(crate) mod witness; use core::ffi::{CStr, c_char, c_void}; use core::mem::{align_of, size_of}; use core::ptr::NonNull; -use std::collections::BTreeMap; -use std::sync::{Arc, RwLock}; +use sppark::{RawPlan, TransformPlan}; +use std::sync::Arc; use crate::expr::{RowOffset, Source}; use crate::graph::{ConstraintGraph, Node}; @@ -29,20 +31,14 @@ use p3_util::log2_strict_usize; const _: () = assert!(size_of::() == size_of::()); const _: () = assert!(align_of::() == align_of::()); -type CachedPowers = Arc<[Goldilocks]>; -type SharedPowerCache = Arc>>; - /// CUDA-backed batched DFT for the Goldilocks field. /// -/// Clones share the host twiddle cache. This is important because the -/// production configuration gives one clone to the PCS and retains another -/// for quotient transforms. +/// Clones share immutable transform plans and their host coset powers. #[derive(Clone, Debug)] pub struct CudaDft { device_id: i32, cpu: Radix2DitParallel, - twiddles: SharedPowerCache<(usize, bool)>, - shift_powers: SharedPowerCache<(usize, u64)>, + planner: sppark::Planner, } impl Default for CudaDft { @@ -68,8 +64,7 @@ impl CudaDft { Self { device_id, cpu: Radix2DitParallel::default(), - twiddles: Arc::default(), - shift_powers: Arc::default(), + planner: sppark::Planner::new(device_id), } } @@ -79,45 +74,18 @@ impl CudaDft { self.device_id } - fn twiddles(&self, log_height: usize, inverse: bool) -> Arc<[Goldilocks]> { - let key = (log_height, inverse); - if let Some(twiddles) = self - .twiddles - .read() - .expect("twiddle cache poisoned") - .get(&key) - { - return Arc::clone(twiddles); - } - - let mut cache = self.twiddles.write().expect("twiddle cache poisoned"); - Arc::clone(cache.entry(key).or_insert_with(|| { - let root = Goldilocks::two_adic_generator(log_height); - let root = if inverse { root.inverse() } else { root }; - root.powers().take((1 << log_height) / 2).collect().into() - })) - } - - fn shift_powers(&self, height: usize, shift: Goldilocks) -> Arc<[Goldilocks]> { - let key = (height, shift.as_canonical_u64()); - if let Some(powers) = self - .shift_powers - .read() - .expect("shift-power cache poisoned") - .get(&key) - { - return Arc::clone(powers); - } - - let mut cache = self - .shift_powers - .write() - .expect("shift-power cache poisoned"); - Arc::clone( - cache - .entry(key) - .or_insert_with(|| shift.powers().take(height).collect().into()), - ) + pub(crate) fn forward_plan(&self, height: usize, width: usize) -> Arc { + self.planner.plan(height, width, 0, None) + } + + pub(crate) fn lde_plan( + &self, + height: usize, + width: usize, + added_bits: usize, + shift: Goldilocks, + ) -> Arc { + self.planner.plan(height, width, added_bits, Some(shift)) } fn validate_dimensions(height: usize, width: usize) { @@ -173,6 +141,15 @@ impl CudaDft { ) -> CudaLde { let height = matrix.height(); let width = matrix.width(); + let _span = tracing::info_span!( + "cuda/lde", + kind = "host", + device = self.device_id, + height, + width, + added_bits + ) + .entered(); Self::validate_dimensions(height, width); assert!(width > 0, "resident CUDA LDE requires at least one column"); let extended_height = height @@ -180,11 +157,7 @@ impl CudaDft { .expect("LDE height overflows usize"); Self::validate_dimensions(extended_height, width); - let log_height = log2_strict_usize(height); - let inverse_twiddles = self.twiddles(log_height, true); - let forward_twiddles = self.twiddles(log2_strict_usize(extended_height), false); - let shift_powers = self.shift_powers(height, shift); - let height_inverse = Goldilocks::ONE.div_2exp_u64(log_height as u64); + let plan = self.lde_plan(height, width, added_bits, shift); let mut handle = core::ptr::null_mut(); // SAFETY: every host buffer has the exact dimensions validated above; // successful creation transfers the device allocation to `CudaLde`. @@ -196,10 +169,7 @@ impl CudaDft { height, width, added_bits, - inverse_twiddles.as_ptr().cast(), - shift_powers.as_ptr().cast(), - forward_twiddles.as_ptr().cast(), - raw_u64(height_inverse), + plan.raw(), ) }; check_cuda(status, "resident coset LDE"); @@ -211,6 +181,7 @@ impl CudaDft { } } + /// Uploads the coset powers before taking a device-memory snapshot. pub(crate) fn prepare_coset_lde_constants( &self, height: usize, @@ -222,24 +193,141 @@ impl CudaDft { .expect("LDE height overflows usize"); Self::validate_dimensions(height, 1); Self::validate_dimensions(extended_height, 1); - let inverse_twiddles = self.twiddles(log2_strict_usize(height), true); - let shift_powers = self.shift_powers(height, shift); - let forward_twiddles = self.twiddles(log2_strict_usize(extended_height), false); + let plan = self.lde_plan(height, 1, added_bits, shift); + let status = unsafe { multi_stark_cuda_prepare_lde_constants(self.device_id, plan.raw()) }; + check_cuda(status, "prepare resident LDE constants"); + } +} + +/// A borrowed output tile on the selected CUDA device. The writer must finish +/// using it before returning; rows wrap at the source's padded height. +pub struct DeviceTraceView<'a> { + device_id: i32, + output: *mut u64, + first: usize, + rows: usize, + width: usize, + _borrow: core::marker::PhantomData<&'a mut [u64]>, +} + +impl DeviceTraceView<'_> { + pub fn device_id(&self) -> i32 { + self.device_id + } + pub fn as_mut_ptr(&self) -> *mut u64 { + self.output + } + pub fn first_row(&self) -> usize { + self.first + } + pub fn rows(&self) -> usize { + self.rows + } + pub fn width(&self) -> usize { + self.width + } +} + +type Generator = Arc>; + +unsafe extern "C" fn write_generated_trace( + context: *mut c_void, + device_id: i32, + output: *mut u64, + first: usize, + rows: usize, +) -> i32 { + let generator = unsafe { &*context.cast::() }; + let result = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| { + generator.write_device_rows(DeviceTraceView { + device_id, + output, + first, + rows, + width: generator.width(), + _borrow: core::marker::PhantomData, + }) + })); + match result { + Ok(Ok(())) => 0, + Ok(Err(error)) => { + eprintln!("[multi-stark/cuda] trace generation failed: {error}"); + 999 + } + Err(_) => 999, + } +} + +unsafe extern "C" fn destroy_generated_trace(context: *mut c_void) { + drop(unsafe { Box::from_raw(context.cast::()) }); +} + +impl CudaDft { + pub(crate) fn generate_coset_lde( + &self, + generator: Generator, + added_bits: usize, + shift: Goldilocks, + ) -> CudaLde { + let height = generator.height(); + let width = generator.width(); + let _span = tracing::info_span!( + "cuda/lde", + kind = "generated", + device = self.device_id, + height, + width, + added_bits + ) + .entered(); + Self::validate_dimensions(height, width); + let extended_height = height + .checked_shl(added_bits.try_into().unwrap()) + .expect("LDE height overflow"); + Self::validate_dimensions(extended_height, width); + let plan = self.lde_plan(height, width, added_bits, shift); + let mut context = Box::new(generator); + let mut handle = core::ptr::null_mut(); let status = unsafe { - multi_stark_cuda_prepare_lde_constants( + multi_stark_cuda_coset_lde_generate( self.device_id, - inverse_twiddles.as_ptr().cast(), - inverse_twiddles.len(), - shift_powers.as_ptr().cast(), + &mut handle, height, - forward_twiddles.as_ptr().cast(), - forward_twiddles.len(), + width, + added_bits, + plan.raw(), + (&mut *context as *mut Generator).cast(), + write_generated_trace, + destroy_generated_trace, ) }; - check_cuda(status, "prepare resident LDE constants"); + check_cuda(status, "generated trace LDE"); + let _ = Box::into_raw(context); + CudaLde { + device_id: self.device_id, + handle: NonNull::new(handle).expect("null generated LDE"), + height: extended_height, + width, + } } } +unsafe extern "C" { + fn multi_stark_cuda_coset_lde_generate( + device: i32, + output: *mut *mut c_void, + height: usize, + width: usize, + added_bits: usize, + plan: *const RawPlan, + context: *mut c_void, + writer: unsafe extern "C" fn(*mut c_void, i32, *mut u64, usize, usize) -> i32, + destroy: unsafe extern "C" fn(*mut c_void), + ) -> i32; + fn multi_stark_cuda_lde_has_generator(handle: *const c_void) -> bool; + fn multi_stark_cuda_lde_generator_context(handle: *const c_void) -> *mut c_void; +} + /// Bit-reversed coset-LDE storage owned by a CUDA device allocation. #[doc(hidden)] pub struct CudaLde { @@ -253,6 +341,23 @@ unsafe impl Send for CudaLde {} unsafe impl Sync for CudaLde {} impl CudaLde { + pub(crate) fn has_generator(&self) -> bool { + unsafe { multi_stark_cuda_lde_has_generator(self.raw_handle()) } + } + + /// Frees what the trace generator, if any, caches on this device. The + /// generator stays attached and keeps serving tiles from the host. + pub(crate) fn release_generator_device(&self) { + let context = unsafe { multi_stark_cuda_lde_generator_context(self.raw_handle()) }; + if context.is_null() { + return; + } + // SAFETY: the handle owns the boxed generator until it is destroyed, + // and the box holds exactly the `Generator` the LDE was created with. + let generator = unsafe { &*context.cast::() }; + generator.release_device(self.device_id); + } + pub(crate) const fn raw_handle(&self) -> *const c_void { self.handle.as_ptr() } @@ -306,6 +411,7 @@ impl CudaLde { g_inv: Goldilocks, ext_w: Goldilocks, ) -> Self { + let _span = tracing::info_span!("cuda/fri_fold").entered(); let mut handle = core::ptr::null_mut(); let status = unsafe { multi_stark_cuda_fri_fold_resident( @@ -846,20 +952,13 @@ pub(crate) fn lookup_graph_lde_memory_upper_bound( .saturating_add(1) .saturating_mul(main_width) .saturating_mul(size_of::()); - // Device-cached twiddles and shift powers may be cold for this height. - let constant_bytes = height - .saturating_div(2) - .saturating_add(height) - .saturating_add(extended_height / 2) - .saturating_mul(size_of::()); Some(( output_bytes, metadata_bytes .saturating_add(message_bytes) .saturating_add(delta_bytes) .saturating_add(scratch_bytes) - .saturating_add(trace_chunk_bytes) - .saturating_add(constant_bytes), + .saturating_add(trace_chunk_bytes), )) } @@ -1021,6 +1120,14 @@ pub(crate) fn quotient_lde_mixed( quotient_degree: usize, log_blowup: usize, ) -> CudaLde { + let _span = tracing::info_span!( + "cuda/quotient_lde", + device = dft.device_id, + quotient_size, + quotient_degree, + added_bits = log_blowup + ) + .entered(); quotient_lde_sources( dft, graph, @@ -1092,8 +1199,8 @@ fn quotient_lde_sources( assert_eq!(alpha.len(), 2 * expected_constraints); let trace_height = quotient_size / quotient_degree; let lde_height = trace_height << log_blowup; - let quotient_twiddles = dft.twiddles(log2_strict_usize(quotient_size), false); - let lde_twiddles = dft.twiddles(log2_strict_usize(lde_height), false); + let quotient_plan = dft.forward_plan(quotient_size, 2); + let lde_plan = dft.forward_plan(lde_height, 2 * quotient_degree); let height_inverse = Goldilocks::ONE.div_2exp_u64(log2_strict_usize(quotient_size) as u64); let weight_step = Goldilocks::GENERATOR.exp_u64(trace_height as u64).inverse(); let weights: Vec<_> = weight_step @@ -1169,8 +1276,8 @@ fn quotient_lde_sources( next_step, quotient_degree, log_blowup, - quotient_twiddles.as_ptr().cast(), - lde_twiddles.as_ptr().cast(), + quotient_plan.raw(), + lde_plan.raw(), weights.as_ptr().cast(), ) } @@ -1209,8 +1316,8 @@ fn quotient_lde_sources( next_step, quotient_degree, log_blowup, - quotient_twiddles.as_ptr().cast(), - lde_twiddles.as_ptr().cast(), + quotient_plan.raw(), + lde_plan.raw(), weights.as_ptr().cast(), ) } @@ -1564,6 +1671,16 @@ pub(crate) fn lookup_lde_resident( ext_w: Goldilocks, log_blowup: usize, ) -> (CudaLde, [Goldilocks; 2]) { + let _span = tracing::info_span!( + "cuda/lookup_lde", + path = "direct", + device = dft.device_id, + height, + num_lookups, + group_size, + added_bits = log_blowup + ) + .entered(); assert!((1..=8).contains(&group_size)); assert_eq!(arg_offsets.len(), num_lookups + 1); assert_eq!(arg_offsets.first(), Some(&0)); @@ -1590,10 +1707,12 @@ pub(crate) fn lookup_lde_resident( assert_eq!(multiplicities.len(), height * num_lookups); assert_eq!(args.len(), height * args_width); let extended_height = height << log_blowup; - let inverse_twiddles = dft.twiddles(log2_strict_usize(height), true); - let shift_powers = dft.shift_powers(height, Goldilocks::GENERATOR); - let forward_twiddles = dft.twiddles(log2_strict_usize(extended_height), false); - let height_inverse = Goldilocks::ONE.div_2exp_u64(log2_strict_usize(height) as u64); + let plan = dft.lde_plan( + height, + 2 * num_lookups.div_ceil(group_size), + log_blowup, + Goldilocks::GENERATOR, + ); let mut tail = [Goldilocks::ZERO; 4]; let mut handle = core::ptr::null_mut(); let status = unsafe { @@ -1612,10 +1731,7 @@ pub(crate) fn lookup_lde_resident( gamma.as_ptr().cast(), raw_u64(ext_w), log_blowup, - inverse_twiddles.as_ptr().cast(), - shift_powers.as_ptr().cast(), - forward_twiddles.as_ptr().cast(), - raw_u64(height_inverse), + plan.raw(), ) }; check_cuda(status, "resident CUDA lookup LDE"); @@ -1645,6 +1761,16 @@ pub(crate) fn lookup_lde_resident_partitioned( log_blowup: usize, cpu_deltas: impl Fn(core::ops::Range) -> Vec<[Goldilocks; 2]> + Sync, ) -> (CudaLde, [Goldilocks; 2]) { + let _span = tracing::info_span!( + "cuda/lookup_lde", + path = "partitioned", + device = dft.device_id, + height, + num_lookups, + group_size, + added_bits = log_blowup + ) + .entered(); assert!((1..=8).contains(&group_size)); assert_eq!(arg_offsets.len(), num_lookups + 1); assert_eq!(arg_offsets.first(), Some(&0)); @@ -1655,10 +1781,12 @@ pub(crate) fn lookup_lde_resident_partitioned( assert_eq!(multiplicities.len(), height * num_lookups); assert_eq!(args.len(), height * args_width); let extended_height = height << log_blowup; - let inverse_twiddles = dft.twiddles(log2_strict_usize(height), true); - let shift_powers = dft.shift_powers(height, Goldilocks::GENERATOR); - let forward_twiddles = dft.twiddles(log2_strict_usize(extended_height), false); - let height_inverse = Goldilocks::ONE.div_2exp_u64(log2_strict_usize(height) as u64); + let plan = dft.lde_plan( + height, + 2 * num_lookups.div_ceil(group_size), + log_blowup, + Goldilocks::GENERATOR, + ); let total_started = std::time::Instant::now(); let create_started = std::time::Instant::now(); let mut pending_handle = core::ptr::null_mut(); @@ -1788,10 +1916,7 @@ pub(crate) fn lookup_lde_resident_partitioned( pending_handle.as_ptr(), &mut handle, tail.as_mut_ptr().cast(), - inverse_twiddles.as_ptr().cast(), - shift_powers.as_ptr().cast(), - forward_twiddles.as_ptr().cast(), - raw_u64(height_inverse), + plan.raw(), ) }; pending.handle = None; @@ -1832,14 +1957,26 @@ pub(crate) fn lookup_graph_lde_resident( log_blowup: usize, ) -> Option<(CudaLde, [Goldilocks; 2])> { assert!((1..=8).contains(&group_size)); + let _span = tracing::info_span!( + "cuda/lookup_lde", + path = "graph", + device = dft.device_id, + height, + num_lookups = graph.lookups.len(), + group_size, + added_bits = log_blowup + ) + .entered(); let (nodes, slot_count, lookups, args) = encode_lookup_nodes(graph)?; let num_lookups = lookups.len(); let groups = num_lookups.div_ceil(group_size.max(1)); let extended_height = height << log_blowup; - let inverse_twiddles = dft.twiddles(log2_strict_usize(height), true); - let shift_powers = dft.shift_powers(height, Goldilocks::GENERATOR); - let forward_twiddles = dft.twiddles(log2_strict_usize(extended_height), false); - let height_inverse = Goldilocks::ONE.div_2exp_u64(log2_strict_usize(height) as u64); + let plan = dft.lde_plan( + height, + 2 * num_lookups.div_ceil(group_size), + log_blowup, + Goldilocks::GENERATOR, + ); let mut tail = [Goldilocks::ZERO; 4]; let mut handle = core::ptr::null_mut(); let status = unsafe { @@ -1861,10 +1998,7 @@ pub(crate) fn lookup_graph_lde_resident( gamma.as_ptr().cast(), raw_u64(ext_w), log_blowup, - inverse_twiddles.as_ptr().cast(), - shift_powers.as_ptr().cast(), - forward_twiddles.as_ptr().cast(), - raw_u64(height_inverse), + plan.raw(), ) }; check_cuda(status, "resident CUDA graph lookup LDE"); @@ -1934,6 +2068,20 @@ impl TwoAdicSubgroupDft for CudaDft { let height = matrix.height(); let width = matrix.width(); Self::validate_dimensions(height, width); + let _span = tracing::info_span!( + "cuda/dft_batch", + device = self.device_id, + height, + width, + backend = if height == 1 || width == 0 { + "noop" + } else if Self::use_cuda_dft(height, width) { + "sppark" + } else { + "cpu" + } + ) + .entered(); if height == 1 || width == 0 { return BitReversalPerm::new_view(matrix); } @@ -1941,23 +2089,23 @@ impl TwoAdicSubgroupDft for CudaDft { return self.cpu.dft_batch(matrix); } - let twiddles = self.twiddles(log2_strict_usize(height), false); + let plan = self.forward_plan(height, width); // SAFETY: Goldilocks is repr(transparent) over u64 (asserted above), // every u64 bit pattern is a valid Goldilocks value, all buffers have // the element counts implied by height/width, and the FFI call is - // synchronous so the borrowed twiddle buffer outlives device use. + // synchronous so the matrix and transform plan outlive device use. let status = unsafe { multi_stark_cuda_dft_batch( self.device_id, matrix.values.as_mut_ptr().cast(), height, width, - twiddles.as_ptr().cast(), + plan.raw(), ) }; check_cuda(status, "batched DFT"); - // The CUDA DIF kernel writes bit-reversed rows. Wrap that storage so + // The CUDA transform writes bit-reversed rows. Wrap that storage so // callers observe the natural-order evaluations required by the trait. BitReversalPerm::new_view(matrix) } @@ -1975,6 +2123,21 @@ impl TwoAdicSubgroupDft for CudaDft { .checked_shl(u32::try_from(added_bits).expect("LDE blowup exceeds u32")) .expect("LDE height overflows usize"); Self::validate_dimensions(extended_height, width); + let _span = tracing::info_span!( + "cuda/coset_lde_batch", + device = self.device_id, + height, + width, + added_bits, + backend = if width == 0 { + "noop" + } else if height > 1 && Self::use_cuda_coset_lde(extended_height, width) { + "sppark" + } else { + "cpu" + } + ) + .entered(); if width == 0 { return BitReversalPerm::new_view(RowMajorMatrix::new(Vec::new(), width)); @@ -1990,11 +2153,7 @@ impl TwoAdicSubgroupDft for CudaDft { return self.cpu.coset_lde_batch(matrix, added_bits, shift); } - let log_height = log2_strict_usize(height); - let inverse_twiddles = self.twiddles(log_height, true); - let forward_twiddles = self.twiddles(log2_strict_usize(extended_height), false); - let shift_powers = self.shift_powers(height, shift); - let height_inverse = Goldilocks::ONE.div_2exp_u64(log_height as u64); + let plan = self.lde_plan(height, width, added_bits, shift); let mut output = Goldilocks::zero_vec(extended_height * width); // SAFETY: the input/output and cached tables have the exact lengths @@ -2008,10 +2167,7 @@ impl TwoAdicSubgroupDft for CudaDft { height, width, added_bits, - inverse_twiddles.as_ptr().cast(), - shift_powers.as_ptr().cast(), - forward_twiddles.as_ptr().cast(), - raw_u64(height_inverse), + plan.raw(), ) }; check_cuda(status, "coset LDE"); @@ -2034,6 +2190,15 @@ pub(crate) fn device_memory_info(device_id: i32) -> (usize, usize) { (free_bytes, total_bytes) } +/// Device bytes the prover keeps free for its later stages: a quarter of the +/// device unless `MULTI_STARK_CUDA_MIN_FREE_BYTES` says otherwise. +pub(crate) fn minimum_free_bytes(total_bytes: usize) -> usize { + std::env::var("MULTI_STARK_CUDA_MIN_FREE_BYTES") + .ok() + .and_then(|value| value.parse().ok()) + .unwrap_or(total_bytes / 4) +} + pub(crate) fn memory_diagnostics_enabled() -> bool { std::env::var_os("MULTI_STARK_CUDA_MEMORY_LOG").is_some() } @@ -2370,7 +2535,14 @@ impl CudaMixedMerkleTree { handles.len(), ) }; - check_cuda(status, "resident LDE Merkle tree creation"); + if status != 0 { + let dimensions: Vec<(usize, usize)> = + ldes.iter().map(|lde| (lde.height(), lde.width())).collect(); + check_cuda( + status, + &format!("resident LDE Merkle tree creation over (height, width) {dimensions:?}"), + ); + } Self { device_id, handle: NonNull::new(handle).expect("CUDA returned a null mixed Merkle handle"), @@ -2486,7 +2658,28 @@ impl CudaMixedMerkleTree { host_digest_groups.len(), ) }; - check_cuda(status, "hybrid CPU/CUDA mixed-height Merkle tree creation"); + if status != 0 { + let dims: Vec> = ldes + .iter() + .zip(host_matrices) + .zip(deferred_dimensions) + .map(|((lde, host), deferred)| { + lde.map(|l| (l.height(), l.width())) + .or_else(|| host.map(|m| (m.height(), m.width()))) + .or_else(|| deferred.map(|d| (d.height, d.width))) + }) + .collect(); + check_cuda( + status, + &format!( + "hybrid Merkle tree creation over (height, width) {dims:?}, host digest groups at heights {:?}", + host_digest_groups + .iter() + .map(|(h, _)| *h) + .collect::>() + ), + ); + } let row_count = heights.into_iter().max().unwrap(); Self { device_id, @@ -2691,7 +2884,7 @@ unsafe extern "C" { values: *mut u64, height: usize, width: usize, - twiddles: *const u64, + plan: *const RawPlan, ) -> i32; fn multi_stark_cuda_coset_lde_batch( @@ -2701,10 +2894,7 @@ unsafe extern "C" { height: usize, width: usize, added_bits: usize, - inverse_twiddles: *const u64, - shift_powers: *const u64, - forward_twiddles: *const u64, - height_inverse: u64, + plan: *const RawPlan, ) -> i32; fn multi_stark_cuda_coset_lde_create( @@ -2714,22 +2904,10 @@ unsafe extern "C" { height: usize, width: usize, added_bits: usize, - inverse_twiddles: *const u64, - shift_powers: *const u64, - forward_twiddles: *const u64, - height_inverse: u64, - ) -> i32; - - fn multi_stark_cuda_prepare_lde_constants( - device_id: i32, - inverse_twiddles: *const u64, - inverse_count: usize, - shift_powers: *const u64, - height: usize, - forward_twiddles: *const u64, - forward_count: usize, + plan: *const RawPlan, ) -> i32; + fn multi_stark_cuda_prepare_lde_constants(device_id: i32, plan: *const RawPlan) -> i32; fn multi_stark_cuda_lde_create_from_host( device_id: i32, handle: *mut *mut c_void, @@ -2855,8 +3033,8 @@ unsafe extern "C" { next_step: usize, quotient_degree: usize, log_blowup: usize, - quotient_twiddles: *const u64, - lde_twiddles: *const u64, + quotient_plan: *const RawPlan, + lde_plan: *const RawPlan, slice_weights: *const u64, ) -> i32; fn multi_stark_cuda_quotient_lde_mixed( @@ -2899,8 +3077,8 @@ unsafe extern "C" { next_step: usize, quotient_degree: usize, log_blowup: usize, - quotient_twiddles: *const u64, - lde_twiddles: *const u64, + quotient_plan: *const RawPlan, + lde_plan: *const RawPlan, slice_weights: *const u64, ) -> i32; fn multi_stark_cuda_mixed_lde_open_row( @@ -2995,10 +3173,7 @@ unsafe extern "C" { gamma: *const u64, ext_w: u64, added_bits: usize, - inverse_twiddles: *const u64, - shift_powers: *const u64, - forward_twiddles: *const u64, - height_inverse: u64, + plan: *const RawPlan, ) -> i32; fn multi_stark_cuda_lookup_lde( device_id: i32, @@ -3015,10 +3190,7 @@ unsafe extern "C" { gamma: *const u64, ext_w: u64, added_bits: usize, - inverse_twiddles: *const u64, - shift_powers: *const u64, - forward_twiddles: *const u64, - height_inverse: u64, + plan: *const RawPlan, ) -> i32; fn multi_stark_cuda_lookup_lde_begin_partitioned( device_id: i32, @@ -3053,10 +3225,7 @@ unsafe extern "C" { pending_handle: *mut c_void, output_handle: *mut *mut c_void, total: *mut u64, - inverse_twiddles: *const u64, - shift_powers: *const u64, - forward_twiddles: *const u64, - height_inverse: u64, + plan: *const RawPlan, ) -> i32; fn multi_stark_cuda_lookup_lde_cancel_partitioned( device_id: i32, @@ -3527,36 +3696,41 @@ mod tests { } #[test] - fn blake3_rows_match_cpu_across_chunk_boundaries() { + fn blake3_rows_match_cpu_across_chunk_and_launch_boundaries() { use p3_blake3::Blake3; - for message_bytes in [1usize, 63, 64, 65, 1023, 1024, 1025, 4264, 7400] { - let message_count = 17; - let messages: Vec = (0..message_bytes * message_count) - .map(|index| (index as u64).wrapping_mul(0x9e37_79b9).to_le_bytes()[0]) - .collect(); - let mut digests = vec![0u8; 32 * message_count]; - // SAFETY: the input contains `message_count` fixed-size messages, - // the output has one 32-byte digest per message, and the call is - // synchronous. - let status = unsafe { - multi_stark_cuda_blake3_hash_rows( - 0, - digests.as_mut_ptr(), - messages.as_ptr(), - message_bytes, - message_count, - ) - }; - check_cuda(status, "BLAKE3 row hashing contract"); + let mut rng = SmallRng::seed_from_u64(0xb1a3e3); + for message_bytes in [ + 1usize, 3, 4, 7, 8, 16, 31, 32, 63, 64, 65, 127, 128, 129, 511, 512, 513, 1023, 1024, + 1025, 2048, 3072, 4264, 7400, 32768, + ] { + for message_count in [1, 31, 32, 33, 255, 256, 257, 513] { + let messages: Vec = (0..message_bytes * message_count) + .map(|_| rng.random()) + .collect(); + let mut digests = vec![0u8; 32 * message_count]; + // SAFETY: the input contains `message_count` fixed-size messages, + // the output has one 32-byte digest per message, and the call is + // synchronous. + let status = unsafe { + multi_stark_cuda_blake3_hash_rows( + 0, + digests.as_mut_ptr(), + messages.as_ptr(), + message_bytes, + message_count, + ) + }; + check_cuda(status, "BLAKE3 row hashing contract"); - for (index, message) in messages.chunks_exact(message_bytes).enumerate() { - let expected: [u8; 32] = Blake3.hash_iter(message.iter().copied()); - assert_eq!( - &digests[index * 32..(index + 1) * 32], - &expected, - "message_bytes={message_bytes}, index={index}" - ); + for (index, message) in messages.chunks_exact(message_bytes).enumerate() { + let expected: [u8; 32] = Blake3.hash_iter(message.iter().copied()); + assert_eq!( + &digests[index * 32..(index + 1) * 32], + &expected, + "message_bytes={message_bytes}, count={message_count}, index={index}" + ); + } } } } @@ -3565,7 +3739,14 @@ mod tests { fn blake3_merkle_root_matches_cpu() { use p3_blake3::Blake3; - for (row_bytes, row_count) in [(16usize, 1usize), (64, 8), (4264, 1024)] { + for (row_bytes, row_count) in [ + (16usize, 1usize), + (64, 8), + (1023, 512), + (1024, 512), + (1025, 512), + (4264, 1024), + ] { let rows: Vec = (0..row_bytes * row_count) .map(|index| (index as u64).wrapping_mul(0x517c_c1b7).to_le_bytes()[0]) .collect(); @@ -3713,6 +3894,55 @@ mod tests { } } + #[test] + fn resident_coset_lde_padding_and_raw_representatives() { + let cpu = Radix2DitParallel::::default(); + let gpu = CudaDft::default(); + let p = Goldilocks::ORDER_U64; + let representatives = [0, 1, p - 1, p, p + 1, u64::MAX]; + // Exercise absent padding, one padded half, multiple padded halves, + // and height-one transforms with lazy input field representatives. + for log_height in [0usize, 1, 2, 5, 8] { + for width in [1usize, 2, 3, 8, 17] { + let height = 1 << log_height; + let matrix = RowMajorMatrix::new( + (0..height * width) + .map(|i| Goldilocks::new(representatives[i % representatives.len()])) + .collect(), + width, + ); + for added_bits in [0usize, 1, 2, 3] { + for shift in [ + Goldilocks::ONE, + Goldilocks::GENERATOR, + Goldilocks::from_u64(11), + ] { + let expected = cpu + .coset_lde_batch(matrix.clone(), added_bits, shift) + .bit_reverse_rows() + .to_row_major_matrix(); + let actual = gpu + .coset_lde_batch_resident(&matrix, added_bits, shift) + .to_row_major_matrix(); + assert_eq!( + actual, expected, + "height=2^{log_height}, width={width}, added_bits={added_bits}, shift={shift}" + ); + assert!( + actual.values.iter().all(|&value| { + // SAFETY: Goldilocks is repr(transparent) over u64. + // Inspect storage without canonicalizing via a field accessor. + let raw = unsafe { core::mem::transmute::(value) }; + raw < p + }), + "LDE digest input must contain canonical field bytes" + ); + } + } + } + } + } + #[test] fn resident_coset_lde_matches_cpu_storage() { let mut rng = SmallRng::seed_from_u64(0x51de); @@ -3724,6 +3954,12 @@ mod tests { (12, 2, 1), (14, 2, 2), (16, 1, 1), + // Fused radix-8 dispatch and all stage-count residues mod 3. + (17, 8, 1), + (18, 8, 1), + (18, 8, 2), + (19, 8, 1), + (18, 3, 1), ] { let height = 1 << log_height; let matrix = diff --git a/src/cuda/pcs.rs b/src/cuda/pcs.rs index 250f24d..81767dd 100644 --- a/src/cuda/pcs.rs +++ b/src/cuda/pcs.rs @@ -72,6 +72,14 @@ fn goldilocks_quadratic_inverse_denominators( } pub trait CudaPcsDft: TwoAdicSubgroupDft { + fn coset_lde_workspace_bytes( + &self, + height: usize, + width: usize, + added_bits: usize, + shift: T, + ) -> usize; + fn prepare_coset_lde_constants(&self, height: usize, added_bits: usize, shift: T); fn coset_lde_batch_resident( @@ -263,6 +271,59 @@ pub struct CudaTwoAdicFriPcs { _phantom: PhantomData, } +/// The coset `GENERATOR * ` in bit-reversed order, as the opening's +/// denominators and interpolation index it, with at least `2^log_height` +/// elements. Element `i` of the bit-reversed coset of size `2^k` is the +/// shift times `g_k^{rev_k(i)}`, and `rev_{k+1}(i) = 2 rev_k(i)` for +/// `i < 2^k`, so every smaller coset is a prefix of a larger one: one vector +/// per process serves every height, and grows only when a larger height is +/// opened. Every shard of a proof opens on the same coset. +fn bit_reversed_coset( + log_height: usize, +) -> std::sync::Arc> { + static CACHE: std::sync::Mutex>>> = + std::sync::Mutex::new(None); + if let Some(coset) = cache_at_least(&CACHE, 1 << log_height) { + return coset; + } + let to_gold = |v: Val| Goldilocks::from_u64(v.as_canonical_u64()); + let generator: Goldilocks = TwoAdicField::two_adic_generator(log_height); + let shift = ::GENERATOR; + assert_eq!(to_gold(Val::two_adic_generator(log_height)), generator); + assert_eq!(to_gold(Val::GENERATOR), shift); + let size = 1usize << log_height; + let mut coset = vec![Goldilocks::ZERO; size]; + const CHUNK: usize = 1 << 14; + coset + .par_chunks_mut(CHUNK) + .enumerate() + .for_each(|(index, chunk)| { + let mut x: Goldilocks = shift * generator.exp_u64((index * CHUNK) as u64); + for value in chunk { + *value = x; + x *= generator; + } + }); + reverse_slice_index_bits(&mut coset); + let mut slot = CACHE.lock().unwrap(); + match &*slot { + Some(cached) if cached.len() >= size => std::sync::Arc::clone(cached), + _ => std::sync::Arc::clone(slot.insert(std::sync::Arc::new(coset))), + } +} + +fn cache_at_least( + cache: &std::sync::Mutex>>>, + size: usize, +) -> Option>> { + cache + .lock() + .unwrap() + .as_ref() + .filter(|coset| coset.len() >= size) + .map(std::sync::Arc::clone) +} + fn prove_fri_cuda_resident( params: &FriParameters, mut inputs: Vec, @@ -341,10 +402,14 @@ where params.max_log_arity, ); let arity = 1 << log_arity; - let (commitment, round) = params.mmcs.commit_cuda_fri(codeword, log_arity); + let (commitment, round) = tracing::info_span!("stark/fri_round_commit") + .in_scope(|| params.mmcs.commit_cuda_fri(codeword, log_arity)); challenger.observe(commitment.clone()); commits.push(commitment); - commit_pow_witnesses.push(challenger.grind(params.commit_proof_of_work_bits)); + commit_pow_witnesses.push( + tracing::info_span!("stark/fri_commit_grind") + .in_scope(|| challenger.grind(params.commit_proof_of_work_bits)), + ); let beta: Challenge = challenger.sample_algebra_element(); let mut beta_step = beta; let betas = (0..log_arity) @@ -396,10 +461,13 @@ where for &log_arity in &log_arities { challenger.observe(Val::from_usize(log_arity)); } - let query_pow_witness = challenger.grind(params.query_proof_of_work_bits); + let _query_span = tracing::info_span!("stark/fri_queries").entered(); + let query_pow_witness = tracing::info_span!("stark/fri_query_grind") + .in_scope(|| challenger.grind(params.query_proof_of_work_bits)); let query_indices = iter::repeat_with(|| challenger.sample_bits(log_global_max_height)) .take(params.num_queries) .collect_vec(); + let input_openings_span = tracing::info_span!("stark/fri_input_openings").entered(); let input_openings = prover_data_with_opening_points .iter() .map(|(data, _)| { @@ -415,6 +483,7 @@ where } }) .collect_vec(); + drop(input_openings_span); let mut current_indices = query_indices; let commit_phase_openings = rounds .iter() @@ -430,7 +499,8 @@ where .map(|&index| index >> log_arity) .collect_vec(); let (opened_rows, opening_proof) = - params.mmcs.open_cuda_fri_batch(round, &group_indices); + tracing::info_span!("stark/fri_commit_phase_opening") + .in_scope(|| params.mmcs.open_cuda_fri_batch(round, &group_indices)); current_indices = group_indices; let sibling_values = positions .into_iter() @@ -548,12 +618,21 @@ where .sum::(); let dft = &self.dft; let log_blowup = self.fri.log_blowup; + let transform_workspaces: Vec<_> = evaluations + .iter() + .map(|(domain, matrix)| { + dft.coset_lde_workspace_bytes( + matrix.height(), + matrix.width(), + log_blowup, + Val::GENERATOR / domain.shift(), + ) + }) + .collect(); + let max_transform_workspace = transform_workspaces.iter().copied().max().unwrap_or(0); let (initial_free, total_bytes) = crate::cuda::device_memory_info(self.mmcs.cuda_device_id()); - let minimum_free = std::env::var("MULTI_STARK_CUDA_MIN_FREE_BYTES") - .ok() - .and_then(|value| value.parse().ok()) - .unwrap_or(total_bytes / 4); + let minimum_free = crate::cuda::minimum_free_bytes(total_bytes); let source_bytes = source_cells.saturating_mul(size_of::()); let lde_bytes = source_bytes .checked_shl(u32::try_from(log_blowup).expect("LDE blowup exceeds u32")) @@ -561,8 +640,13 @@ where let release_traces_during_construction = source_bytes .saturating_add(lde_bytes) .saturating_add(minimum_free) + .saturating_add(max_transform_workspace) > initial_free; - if lde_bytes.saturating_add(minimum_free) > initial_free { + if lde_bytes + .saturating_add(minimum_free) + .saturating_add(max_transform_workspace) + > initial_free + { let max_lde_height = evaluations .iter() .map(|(_, matrix)| matrix.height() << log_blowup) @@ -574,7 +658,8 @@ where let tree_workspace_bytes = max_lde_height.saturating_mul(96); let gpu_lde_budget = initial_free .saturating_sub(minimum_free) - .saturating_sub(tree_workspace_bytes); + .saturating_sub(tree_workspace_bytes) + .saturating_sub(max_transform_workspace); let matrix_resources = evaluations .iter() .map(|(_, matrix)| { @@ -604,6 +689,19 @@ where // Split allocatable device capacity between persistent LDEs and // later phase workspace. This prevents stage one from filling VRAM // with values which lookup immediately has to copy back and evict. + // Height groups wider than the device leaf kernel can hash are + // neither durable nor transient: their LDEs and digests come + // from the host, whatever the memory budget says. + let wide_heights = + super::mmcs::host_hashed_heights(evaluations.iter().map(|(_, matrix)| { + p3_matrix::Dimensions { + width: matrix.width(), + height: matrix.height(), + } + })); + for height in &wide_heights { + height_groups.remove(height); + } let durable_budget = gpu_lde_budget / 2; // A Merkle leaf combines every matrix at a given height. Keeping // height groups intact avoids streaming the CPU half of a split @@ -627,11 +725,16 @@ where // GPU lane during commitment. Its selected matrices are transformed // temporarily for hashing, then recomputed by the CPU for durable // host storage while CUDA moves on to the retained groups. - let transient_reserve = max_lde_height.saturating_mul(32).saturating_add(64 << 20); + let transient_reserve = max_lde_height + .saturating_mul(32) + .saturating_add(64 << 20) + .saturating_add(max_transform_workspace); let select_transient_plan = |transient_budget| { height_indices .iter() - .filter(|(height, _)| !durable_heights.contains(height)) + .filter(|(height, _)| { + !durable_heights.contains(height) && !wide_heights.contains(height) + }) .filter_map(|(&height, indices)| { let resources = indices .iter() @@ -891,7 +994,7 @@ where .collect_vec(); host_digest_groups.extend(hash_host_only_height_groups( &host_matrices, - &resident, + &resident.iter().map(Option::as_ref).collect::>(), &deferred_dimensions, &prehashed_heights, )); @@ -954,7 +1057,11 @@ where let wave_size = if release_traces_during_construction { 1 } else { - CUDA_LDE_WAVE + let workspace_budget = initial_free + .saturating_sub(source_bytes) + .saturating_sub(lde_bytes) + .saturating_sub(minimum_free); + (workspace_budget / max_transform_workspace.max(1)).clamp(1, CUDA_LDE_WAVE) }; for wave in evaluations.chunks(wave_size) { let transform = @@ -1184,7 +1291,7 @@ where .iter() .all(|(data, _)| self.mmcs.is_cuda_resident(data)) { - let resident_rounds = debug_span!("cuda prepare resident rounds").in_scope(|| { + let resident_rounds = tracing::info_span!("stark/fri_prepare_rounds").in_scope(|| { commitment_data_with_opening_points .iter() .map(|(data, points)| (self.mmcs.resident_or_upload(data), points)) @@ -1198,7 +1305,7 @@ where .unwrap_or(0); let final_fri_height = self.fri.blowup() * self.fri.final_poly_len(); if resident_max_height > 1024 && resident_max_height > final_fri_height { - let _resident_guard = debug_span!("cuda resident fri").entered(); + let _resident_guard = tracing::info_span!("stark/fri_resident").entered(); let rounds = resident_rounds; let device_id = self.mmcs.cuda_device_id(); assert_eq!(>::DIMENSION, 2); @@ -1225,10 +1332,7 @@ where .max() .unwrap(); let log_global_max_height = log2_strict_usize(global_max_height); - let coset_domain = - TwoAdicMultiplicativeCoset::new(Val::GENERATOR, log_global_max_height).unwrap(); - let mut coset: Vec = coset_domain.iter().collect(); - reverse_slice_index_bits(&mut coset); + let coset_gold = bit_reversed_coset::(log_global_max_height); let mut max_log: LinearMap = LinearMap::new(); for (ldes, points) in &rounds { for (lde, ps) in ldes.iter().zip(points.iter()) { @@ -1245,7 +1349,6 @@ where Challenge::from_basis_coefficients_slice(&[Val::ZERO, Val::ONE]).unwrap(); let ext_w = to_pair(ext_x * ext_x)[0]; assert_eq!(ext_w, Goldilocks::from_u64(7)); - let coset_gold: Vec<_> = coset.iter().copied().map(to_gold).collect(); let mut inv_offsets = LinearMap::new(); let mut inverse_points = Vec::new(); let mut inverse_counts = Vec::new(); @@ -1294,7 +1397,7 @@ where .collect_vec() }) .collect_vec(); - let interpolated = debug_span!("cuda interpolate openings") + let interpolated = tracing::info_span!("stark/fri_interpolate") .in_scope(|| workspace.interpolate(&interpolation_tasks, output_count, ext_w)); let all_opened_values = layouts .into_iter() @@ -1352,10 +1455,10 @@ where } } } - debug_span!("cuda reduce openings") + tracing::info_span!("stark/fri_reduce") .in_scope(|| workspace.reduce(&reduction_tasks, &alpha_pairs, ext_w)); let fri_input = reduced.into_iter().rev().flatten().collect_vec(); - let fri_proof = debug_span!("cuda prove fri").in_scope(|| { + let fri_proof = tracing::info_span!("stark/fri_prove").in_scope(|| { prove_fri_cuda_resident( &self.fri, fri_input, @@ -1401,7 +1504,8 @@ where .expect("No Matrices Supplied?"); let final_fri_height = self.fri.blowup() * self.fri.final_poly_len(); if cuda_max_height > 1024 && cuda_max_height > final_fri_height { - let _resident_guard = debug_span!("cuda streamed fri").entered(); + let _resident_guard = tracing::info_span!("stark/fri_streamed").entered(); + let prepare_span = tracing::info_span!("stark/fri_prepare").entered(); let phase_started = std::time::Instant::now(); let device_id = self.mmcs.cuda_device_id(); assert_eq!(>::DIMENSION, 2); @@ -1419,10 +1523,7 @@ where .expect("quadratic extension element") }; let log_global_max_height = log2_strict_usize(cuda_max_height); - let coset_domain = - TwoAdicMultiplicativeCoset::new(Val::GENERATOR, log_global_max_height).unwrap(); - let mut coset: Vec = coset_domain.iter().collect(); - reverse_slice_index_bits(&mut coset); + let coset_gold = bit_reversed_coset::(log_global_max_height); let mut max_log: LinearMap = LinearMap::new(); for ((_, points), round_dimensions) in commitment_data_with_opening_points @@ -1445,7 +1546,6 @@ where .expect("quadratic extension generator"); let ext_w = to_pair(ext_x * ext_x)[0]; assert_eq!(ext_w, Goldilocks::from_u64(7)); - let coset_gold: Vec<_> = coset.iter().copied().map(to_gold).collect(); let mut inv_offsets = LinearMap::new(); let mut inverse_points = Vec::new(); let mut inverse_counts = Vec::new(); @@ -1480,7 +1580,7 @@ where .map(|entry| entry.0) .collect_vec(); self.mmcs - .ensure_device_headroom_batch(&admission_data, fri_workspace_bytes); + .ensure_device_headroom_batch(&admission_data, fri_workspace_bytes, "fri"); if crate::cuda::memory_diagnostics_enabled() { eprintln!( "[multi-stark/cuda] FRI admission: {:.3}s", @@ -1570,6 +1670,8 @@ where ); } + drop(prepare_span); + let interpolate_span = tracing::info_span!("stark/fri_interpolate").entered(); let interpolation_started = std::time::Instant::now(); let ((cpu_opened, cpu_interpolation_seconds), gpu_opened, gpu_interpolation_seconds) = std::thread::scope(|scope| { @@ -1699,6 +1801,8 @@ where ); } + drop(interpolate_span); + let observe_span = tracing::info_span!("stark/fri_observe_openings").entered(); for round in &all_opened_values { for matrix in round { for values in matrix { @@ -1735,6 +1839,8 @@ where }) .collect_vec(); + drop(observe_span); + let _reduce_span = tracing::info_span!("stark/fri_reduce").entered(); let reduction_started = std::time::Instant::now(); let ((cpu_reduced, cpu_reduction_seconds), mut gpu_reduced, gpu_reduction_seconds) = std::thread::scope(|scope| { @@ -1867,7 +1973,7 @@ where } let fri_input = gpu_reduced.into_iter().rev().flatten().collect_vec(); let folding_started = std::time::Instant::now(); - let fri_proof = debug_span!("cuda prove streamed fri").in_scope(|| { + let fri_proof = tracing::info_span!("stark/fri_prove").in_scope(|| { prove_fri_cuda_resident( &self.fri, fri_input, diff --git a/src/cuda/sppark.rs b/src/cuda/sppark.rs new file mode 100644 index 0000000..8d49b4e --- /dev/null +++ b/src/cuda/sppark.rs @@ -0,0 +1,674 @@ +//! sppark's Goldilocks NTT through the device-pointer adapter in +//! `cuda/sppark_ntt.cu`, with shared immutable transform plans. +//! +//! The contracts the prover relies on, checked by the tests below against +//! the CPU reference: a forward transform in natural order equals the DFT, +//! the `R` orders are the bit-reversed permutation of the `N` orders, +//! inverse transforms are normalized upstream, the coset variant is the +//! DFT on the coset shifted by the field generator, and outputs are stored +//! as canonical words. Inputs must be canonical: upstream does not reduce +//! representatives at or above the modulus. + +use core::ffi::c_int; + +use p3_goldilocks::Goldilocks; + +use super::check_cuda; + +/// The allocation and launch shape shared by admission and CUDA execution. +#[derive(Debug)] +#[repr(C)] +pub(crate) struct RawPlan { + height: usize, + width: usize, + extended_height: usize, + columns: usize, + inverse_group: usize, + forward_group: usize, + scratch_bytes: usize, + shift_powers: *const u64, +} + +#[derive(Debug)] +pub(crate) struct TransformPlan { + raw: RawPlan, + powers: Option>, +} + +// SAFETY: the pointer references immutable host powers owned by this plan. +// CUDA entry points borrow it only for their synchronous host-side upload. +unsafe impl Send for TransformPlan {} +unsafe impl Sync for TransformPlan {} + +impl TransformPlan { + pub(crate) fn raw(&self) -> &RawPlan { + &self.raw + } + + pub(crate) fn scratch_bytes(&self) -> usize { + self.raw.scratch_bytes + } + + pub(crate) fn constant_bytes(&self) -> usize { + self.powers.as_ref().map_or(0, |powers| powers.len() * 8) + } +} + +type PlanKey = (usize, usize, usize, Option); +type PowerCache = std::collections::BTreeMap<(usize, u64), std::sync::Arc<[Goldilocks]>>; + +#[derive(Default, Debug)] +struct Plans { + shapes: std::collections::BTreeMap>, + powers: PowerCache, +} + +/// Immutable per-device settings and shared plans; environment changes cannot +/// change the workspace between admission and execution of a transform. +#[derive(Clone, Debug)] +pub(crate) struct Planner { + panel_bytes: usize, + batch_bytes: usize, + plans: std::sync::Arc>, +} + +impl Planner { + pub(crate) fn new(device: i32) -> Self { + let panel_bytes = std::env::var("MULTI_STARK_SPPARK_PANEL_BYTES") + .ok() + .map(|value| { + value.parse().unwrap_or_else(|_| { + panic!("MULTI_STARK_SPPARK_PANEL_BYTES must be a non-negative byte count") + }) + }) + .filter(|&n| n != 0) + .unwrap_or(4usize << 30); + // Launch groups sized to the L2 cache keep a group's columns resident + // across the stages of a batched transform. + let mut l2_bytes = 0; + check_cuda( + unsafe { multi_stark_sppark_l2_bytes(device, &mut l2_bytes) }, + "NTT device properties", + ); + Self::with_budgets(panel_bytes, l2_bytes) + } + + pub(crate) fn with_budgets(panel_bytes: usize, batch_bytes: usize) -> Self { + Self { + panel_bytes, + batch_bytes, + plans: Default::default(), + } + } + + pub(crate) fn plan( + &self, + height: usize, + width: usize, + added_bits: usize, + shift: Option, + ) -> std::sync::Arc { + use p3_field::{PrimeCharacteristicRing, PrimeField64}; + use std::sync::Arc; + assert!( + height.is_power_of_two(), + "NTT height must be a power of two" + ); + assert!(width > 0, "NTT width must be positive"); + let log = height.trailing_zeros() as usize; + assert!( + added_bits <= max_log_domain() && log + added_bits <= max_log_domain(), + "NTT exceeds sppark's compiled domain" + ); + assert!(shift.is_some() || added_bits == 0); + let extended_height = height << added_bits; + extended_height + .checked_mul(width) + .and_then(|n| n.checked_mul(8)) + .expect("NTT matrix byte count overflows usize"); + let column_bytes = extended_height + .checked_add(if shift.is_some() { height } else { 0 }) + .and_then(|n| n.checked_mul(8)) + .expect("NTT column byte count overflows usize"); + assert!( + column_bytes <= self.panel_bytes, + "MULTI_STARK_SPPARK_PANEL_BYTES cannot hold one NTT column: need {column_bytes} bytes, budget {}", + self.panel_bytes + ); + let key = ( + height, + width, + added_bits, + shift.map(|s| s.as_canonical_u64()), + ); + if let Some(plan) = self + .plans + .read() + .expect("NTT plan cache poisoned") + .shapes + .get(&key) + { + return Arc::clone(plan); + } + let mut cache = self.plans.write().expect("NTT plan cache poisoned"); + if let Some(plan) = cache.shapes.get(&key) { + return Arc::clone(plan); + } + let powers = shift.map(|shift| { + Arc::clone( + cache + .powers + .entry((height, shift.as_canonical_u64())) + .or_insert_with(|| shift.powers().take(height).collect().into()), + ) + }); + let columns = (self.panel_bytes / column_bytes).min(width).min(65535); + let group = |rows: usize| (self.batch_bytes / (rows * 8)).max(1).min(columns); + let raw = RawPlan { + height, + width, + extended_height, + columns, + inverse_group: group(height), + forward_group: group(extended_height), + scratch_bytes: columns * column_bytes, + shift_powers: powers + .as_ref() + .map_or(core::ptr::null(), |p| p.as_ptr().cast()), + }; + let plan = Arc::new(TransformPlan { raw, powers }); + cache.shapes.insert(key, Arc::clone(&plan)); + plan + } +} + +/// `NTT::InputOutputOrder`: whether the input and the output are in natural +/// (`N`) or bit-reversed (`R`) order. +#[cfg(test)] +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +#[repr(i32)] +enum Order { + NN = 0, + NR = 1, + RN = 2, + RR = 3, +} + +#[cfg(test)] +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +#[repr(i32)] +enum Direction { + Forward = 0, + Inverse = 1, +} + +unsafe extern "C" { + fn multi_stark_sppark_max_lg_domain() -> c_int; + fn multi_stark_sppark_l2_bytes(device: c_int, bytes: *mut usize) -> c_int; + #[cfg(test)] + fn multi_stark_sppark_borrowed_round_trip(device: c_int, values: *mut u64, lg: u32) -> c_int; + #[cfg(test)] + fn multi_stark_sppark_ntt_host( + device: c_int, + inout: *mut u64, + lg: u32, + order: c_int, + direction: c_int, + coset: c_int, + ) -> c_int; + #[cfg(test)] + fn multi_stark_sppark_ntt_batch_host( + device: c_int, + inout: *mut u64, + lg: u32, + order: c_int, + direction: c_int, + coset: c_int, + batch: u32, + stride: usize, + ) -> c_int; +} + +/// The largest log domain size the compiled upstream parameters support. +pub(crate) fn max_log_domain() -> usize { + usize::try_from(unsafe { multi_stark_sppark_max_lg_domain() }).expect("domain limit") +} + +#[cfg(test)] +fn log_len(values: usize) -> u32 { + assert!( + values.is_power_of_two() && values > 1, + "sppark transforms take 2^lg elements, lg > 0" + ); + values.trailing_zeros() +} + +/// Transforms `values` in place on `device`: uploaded, transformed and +/// downloaded within the call. `coset` selects the coset by the field +/// generator. An error is the CUDA status the adapter returned; `device` +/// must be the CUDA ordinal of a device upstream supports. +#[cfg(test)] +fn try_ntt_host( + device: i32, + values: &mut [Goldilocks], + order: Order, + direction: Direction, + coset: bool, +) -> Result<(), i32> { + let lg = log_len(values.len()); + let status = unsafe { + multi_stark_sppark_ntt_host( + device, + values.as_mut_ptr().cast(), + lg, + order as c_int, + direction as c_int, + c_int::from(coset), + ) + }; + if status == 0 { Ok(()) } else { Err(status) } +} + +/// `count` vectors of `2^lg` elements laid out `stride` elements apart. +#[cfg(test)] +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +struct Batch { + lg: u32, + count: u32, + stride: usize, +} + +/// Transforms the vectors of `batch` in `values` (`count * stride` long) in +/// one batched launch sequence, uploaded, transformed and downloaded within +/// the call. +#[cfg(test)] +fn try_ntt_batch_host( + device: i32, + values: &mut [Goldilocks], + batch: Batch, + order: Order, + direction: Direction, + coset: bool, +) -> Result<(), i32> { + let words = (batch.count as usize) + .checked_mul(batch.stride) + .expect("batch words overflow usize"); + assert_eq!(values.len(), words); + assert!( + u32::try_from(batch.stride).is_ok(), + "the adapter indexes vectors with 32 bits" + ); + let status = unsafe { + multi_stark_sppark_ntt_batch_host( + device, + values.as_mut_ptr().cast(), + batch.lg, + order as c_int, + direction as c_int, + c_int::from(coset), + batch.count, + batch.stride, + ) + }; + if status == 0 { Ok(()) } else { Err(status) } +} + +/// [`try_ntt_host`], panicking on a CUDA status like the other backends. +#[cfg(test)] +fn ntt_host( + device: i32, + values: &mut [Goldilocks], + order: Order, + direction: Direction, + coset: bool, +) { + if let Err(status) = try_ntt_host(device, values, order, direction, coset) { + check_cuda(status, "sppark host transform"); + } +} + +/// The raw stored words of `values`; Goldilocks is `repr(transparent)`. +#[cfg(test)] +fn raw_words(values: &[Goldilocks]) -> &[u64] { + // SAFETY: Goldilocks is a transparent wrapper over u64. + unsafe { core::slice::from_raw_parts(values.as_ptr().cast(), values.len()) } +} + +#[cfg(test)] +mod tests { + use super::*; + use p3_dft::{Radix2DitParallel, TwoAdicSubgroupDft}; + use p3_field::{Field, PrimeCharacteristicRing, PrimeField64}; + use p3_matrix::Matrix; + use p3_matrix::bitrev::BitReversibleMatrix; + use p3_matrix::dense::RowMajorMatrix; + use p3_util::reverse_slice_index_bits; + use rand::{RngExt, SeedableRng, rngs::SmallRng}; + + const LOGS: [usize; 6] = [1, 5, 10, 11, 16, 20]; + + fn random(lg: usize, seed: u64) -> Vec { + let mut rng = SmallRng::seed_from_u64(seed); + (0..1usize << lg).map(|_| rng.random()).collect() + } + + fn cpu_dft(values: &[Goldilocks]) -> Vec { + Radix2DitParallel::::default() + .dft_batch(RowMajorMatrix::new(values.to_vec(), 1)) + .to_row_major_matrix() + .values + } + + fn cpu_idft(values: &[Goldilocks]) -> Vec { + Radix2DitParallel::::default() + .idft_batch(RowMajorMatrix::new(values.to_vec(), 1)) + .values + } + + fn cpu_coset_dft(values: &[Goldilocks]) -> Vec { + Radix2DitParallel::::default() + .coset_dft_batch( + RowMajorMatrix::new(values.to_vec(), 1), + Goldilocks::GENERATOR, + ) + .to_row_major_matrix() + .values + } + + fn assert_canonical(values: &[Goldilocks], what: &str) { + for (i, &word) in raw_words(values).iter().enumerate() { + assert!( + word < Goldilocks::ORDER_U64, + "{what}: word {i} is not canonical: {word:#x}" + ); + } + } + + #[test] + fn an_unknown_device_ordinal_is_rejected_before_any_launch() { + // cudaErrorInvalidDevice is 101; the adapter answers it for an + // ordinal upstream did not enumerate, without touching the buffer. + let input = random(8, 0xde); + let mut values = input.clone(); + let status = try_ntt_host(1 << 20, &mut values, Order::NN, Direction::Forward, false); + assert_eq!(status, Err(101)); + assert_eq!(values, input); + } + + fn resident_matches_cpu( + dft: &super::super::CudaDft, + matrix: &RowMajorMatrix, + added_bits: usize, + shift: Goldilocks, + ) { + let expected = Radix2DitParallel::::default() + .coset_lde_batch(matrix.clone(), added_bits, shift) + .bit_reverse_rows() + .to_row_major_matrix(); + let actual = dft + .coset_lde_batch_resident(matrix, added_bits, shift) + .to_row_major_matrix(); + assert_eq!(actual, expected); + assert_canonical(&actual.values, "resident LDE"); + } + + #[test] + fn resident_lde_reduces_raw_representatives() { + let p = Goldilocks::ORDER_U64; + let words = [0u64, 1, p - 1, p, p + 1, u64::MAX, 7, p + 7]; + let height = 1usize << 10; + let width = 9; + // SAFETY: the test reinterprets raw words as field elements on purpose. + let values: Vec = (0..height * width) + .map(|i| unsafe { core::mem::transmute::(words[i % words.len()]) }) + .collect(); + for shift in [ + Goldilocks::GENERATOR, + Goldilocks::ONE, + Goldilocks::from_u64(11), + ] { + resident_matches_cpu( + &super::super::CudaDft::new(0), + &RowMajorMatrix::new(values.clone(), width), + 2, + shift, + ); + } + } + + #[test] + fn panel_and_batch_boundaries_match_cpu_without_global_settings() { + let mut rng = SmallRng::seed_from_u64(0x9a7e); + let matrix = RowMajorMatrix::new((0..(1 << 12) * 33).map(|_| rng.random()).collect(), 33); + for batch_bytes in [0, 64 << 10, 64 << 20] { + let mut dft = super::super::CudaDft::new(0); + dft.planner = Planner::with_budgets(256 << 10, batch_bytes); + let plan = dft.lde_plan(matrix.height(), matrix.width(), 2, Goldilocks::GENERATOR); + assert!(plan.raw.columns < matrix.width()); + resident_matches_cpu(&dft, &matrix, 2, Goldilocks::GENERATOR); + } + } + + /// The fork's batched launch sequence transforms every vector of a + /// batch exactly as the single-vector entry does, for each order, + /// direction and coset setting, with padding between the vectors. + #[test] + fn a_batch_of_vectors_matches_the_vectors_one_by_one() { + let mut rng = SmallRng::seed_from_u64(0xba7c); + for (lg, count, padding) in [(4u32, 3u32, 0usize), (10, 7, 8), (12, 5, 0), (16, 3, 16)] { + let length = 1usize << lg; + let stride = length + padding; + let batch = Batch { lg, count, stride }; + let values: Vec = + (0..count as usize * stride).map(|_| rng.random()).collect(); + for order in [Order::NN, Order::NR, Order::RN, Order::RR] { + for direction in [Direction::Forward, Direction::Inverse] { + for coset in [false, true] { + let mut batched = values.clone(); + try_ntt_batch_host(0, &mut batched, batch, order, direction, coset) + .unwrap(); + let mut expected = values.clone(); + for vector in expected.chunks_exact_mut(stride) { + ntt_host(0, &mut vector[..length], order, direction, coset); + } + assert_eq!( + raw_words(&batched), + raw_words(&expected), + "2^{lg} x {count} stride {stride} {order:?} {direction:?} coset {coset}" + ); + } + } + } + } + } + + #[test] + fn plan_reuses_constants_and_rejects_unrepresentable_shapes() { + use std::sync::Arc; + let planner = Planner::with_budgets(256 << 10, 64 << 10); + let plan = planner.plan(1 << 12, 33, 2, Some(Goldilocks::GENERATOR)); + assert!(Arc::ptr_eq( + &plan, + &planner.plan(1 << 12, 33, 2, Some(Goldilocks::GENERATOR)) + )); + assert!(plan.scratch_bytes() <= 256 << 10); + let narrow = planner.plan(1 << 12, 1, 2, Some(Goldilocks::GENERATOR)); + assert_eq!(plan.raw.shift_powers, narrow.raw.shift_powers); + for (height, width, bits) in [ + (1 << 29, 1, 0), + (1 << 24, 1, 5), + (1 << 20, 1, 2), + (2, usize::MAX, 0), + ] { + assert!( + std::panic::catch_unwind(|| planner.plan( + height, + width, + bits, + Some(Goldilocks::ONE) + )) + .is_err() + ); + } + } + + #[test] + fn ffi_rejects_mismatched_transform_plans() { + let dft = super::super::CudaDft::new(0); + let input = vec![Goldilocks::ONE; 8 * 3]; + for plan in [ + dft.forward_plan(4, 3), + dft.forward_plan(8, 4), + dft.lde_plan(8, 3, 1, Goldilocks::GENERATOR), + dft.lde_plan(8, 3, 0, Goldilocks::GENERATOR), + ] { + let mut values = input.clone(); + let status = unsafe { + super::super::multi_stark_cuda_dft_batch( + 0, + values.as_mut_ptr().cast(), + 8, + 3, + plan.raw(), + ) + }; + assert_eq!(status, 1, "cudaErrorInvalidValue"); + assert_eq!(values, input); + } + for plan in [ + dft.lde_plan(4, 3, 2, Goldilocks::GENERATOR), + dft.lde_plan(8, 4, 1, Goldilocks::GENERATOR), + dft.lde_plan(8, 3, 2, Goldilocks::GENERATOR), + ] { + let mut output = vec![Goldilocks::TWO; 16 * 3]; + let status = unsafe { + super::super::multi_stark_cuda_coset_lde_batch( + 0, + output.as_mut_ptr().cast(), + input.as_ptr().cast(), + 8, + 3, + 1, + plan.raw(), + ) + }; + assert_eq!(status, 1, "cudaErrorInvalidValue"); + assert!(output.iter().all(|v| *v == Goldilocks::TWO)); + let mut handle = core::ptr::null_mut(); + let status = unsafe { + super::super::multi_stark_cuda_coset_lde_create( + 0, + &mut handle, + input.as_ptr().cast(), + 8, + 3, + 1, + plan.raw(), + ) + }; + assert_eq!(status, 1, "cudaErrorInvalidValue"); + assert!(handle.is_null()); + } + } + + #[test] + fn borrowed_stream_preserves_order_and_caller_ownership() { + let mut values = random(12, 0xb0770); + let expected = values.clone(); + check_cuda( + unsafe { multi_stark_sppark_borrowed_round_trip(0, values.as_mut_ptr().cast(), 12) }, + "borrowed stream round trip", + ); + assert_eq!(values, expected); + } + + #[test] + fn domain_limit_covers_the_trace_cap_and_its_lde() { + assert!( + max_log_domain() >= 26, + "the prover extends 2^24 traces by 4" + ); + } + + #[test] + fn forward_natural_matches_the_cpu_dft() { + for lg in LOGS { + let input = random(lg, 0xf0 + lg as u64); + let expected = cpu_dft(&input); + let mut actual = input.clone(); + ntt_host(0, &mut actual, Order::NN, Direction::Forward, false); + assert_eq!(actual, expected, "lg {lg}"); + assert_canonical(&actual, "forward NN"); + } + } + + #[test] + fn reversed_orders_are_the_bit_reversal_of_natural_ones() { + for lg in LOGS { + let input = random(lg, 0xb1 + lg as u64); + let mut natural = cpu_dft(&input); + let mut nr = input.clone(); + ntt_host(0, &mut nr, Order::NR, Direction::Forward, false); + reverse_slice_index_bits(&mut natural); + assert_eq!(nr, natural, "NR at lg {lg}"); + // RN: bit-reversed input yields the natural-order output. + let mut reversed_input = input.clone(); + reverse_slice_index_bits(&mut reversed_input); + ntt_host(0, &mut reversed_input, Order::RN, Direction::Forward, false); + assert_eq!(reversed_input, cpu_dft(&input), "RN at lg {lg}"); + } + } + + #[test] + fn inverse_is_normalized_upstream() { + for lg in LOGS { + let input = random(lg, 0x1d + lg as u64); + let mut actual = input.clone(); + ntt_host(0, &mut actual, Order::NN, Direction::Inverse, false); + assert_eq!(actual, cpu_idft(&input), "inverse NN at lg {lg}"); + // A forward NR followed by an inverse RN is the identity, so the + // prover's inverse-then-forward pair needs no scaling of its own. + let mut round_trip = input.clone(); + ntt_host(0, &mut round_trip, Order::NR, Direction::Forward, false); + ntt_host(0, &mut round_trip, Order::RN, Direction::Inverse, false); + assert_eq!(round_trip, input, "round trip at lg {lg}"); + } + } + + #[test] + fn coset_transform_uses_the_field_generator() { + for lg in LOGS { + let input = random(lg, 0xc0 + lg as u64); + let mut actual = input.clone(); + ntt_host(0, &mut actual, Order::NN, Direction::Forward, true); + assert_eq!(actual, cpu_coset_dft(&input), "coset forward NN at lg {lg}"); + } + } + + #[test] + fn extreme_canonical_values_transform_and_stay_canonical() { + // Upstream takes canonical words only: raw representatives at or + // above the modulus, which the first-party kernels reduce lazily, + // transform to different values, so the adapter's callers reduce + // first. Canonical extremes must survive unchanged. + let p = Goldilocks::ORDER_U64; + let words: Vec = [0, 1, p - 1, 7, p - 7, 1 << 32, (1 << 32) - 1, p - (1 << 32)] + .into_iter() + .cycle() + .take(1 << 10) + .collect(); + let input: Vec = words.iter().map(|&w| Goldilocks::from_u64(w)).collect(); + assert_eq!( + raw_words(&input), + &words[..], + "the inputs are stored canonically" + ); + let mut actual = input.clone(); + ntt_host(0, &mut actual, Order::NN, Direction::Forward, false); + assert_eq!(actual, cpu_dft(&input)); + assert_canonical(&actual, "forward of canonical extremes"); + ntt_host(0, &mut actual, Order::NN, Direction::Inverse, false); + assert_eq!(actual, input); + assert_canonical(&actual, "inverse of canonical extremes"); + } +} diff --git a/src/cuda/witness.rs b/src/cuda/witness.rs new file mode 100644 index 0000000..badf474 --- /dev/null +++ b/src/cuda/witness.rs @@ -0,0 +1,204 @@ +//! Device main-trace generation and commitment with bounded construction workspace. + +use std::sync::{OnceLock, atomic::AtomicBool}; + +use p3_commit::Pcs; +use p3_field::Field; +use p3_matrix::{Dimensions, dense::RowMajorMatrix}; +use p3_symmetric::MerkleCap; + +use super::mmcs::CudaMmcsData; +use super::{CudaLde, CudaMixedMerkleTree, device_memory_info}; +use crate::config::{Com, Domain, PcsData}; +use crate::types::{GoldilocksBlake3Config as Config, Pcs as CudaPcs, Val}; +use crate::witness::TraceSource; + +pub(crate) fn record_lde_spill(device: i32, bytes: usize) { + tracing::debug!(device, bytes, "spilled active LDE"); +} + +fn spill_lde(lde: &CudaLde) -> RowMajorMatrix { + let matrix = lde.to_row_major_matrix(); + // Materialization synchronizes every use of the released values. + unsafe { lde.release_values() }; + record_lde_spill(lde.device_id, lde.height() * lde.width() * 8); + matrix +} + +pub(crate) fn commit( + pcs: &CudaPcs, + evaluations: Vec<(Domain, TraceSource)>, +) -> (Com, PcsData) { + if evaluations + .iter() + .all(|(_, m)| matches!(m, TraceSource::Host(_))) + { + tracing::debug!( + device = pcs.dft.device_id(), + path = "host_pcs", + "main commitment path" + ); + return >::commit( + pcs, + evaluations.into_iter().map(|(d, m)| (d, m.materialize())), + ); + } + let device = pcs.dft.device_id(); + tracing::debug!(device, path = "prepared", "main commitment path"); + let blowup = pcs.fri.log_blowup; + let dimensions: Vec<_> = evaluations + .iter() + .map(|(_, m)| Dimensions { + width: m.width(), + height: m + .height() + .checked_shl(blowup.try_into().expect("LDE blowup exceeds u32")) + .expect("LDE height overflow"), + }) + .collect(); + for (domain, source) in &evaluations { + assert_eq!( + domain.size(), + source.height(), + "main-trace domain height mismatch" + ); + } + let max_height = dimensions.iter().map(|d| d.height).max().unwrap(); + let (_, total) = device_memory_info(device); + let reserve = super::minimum_free_bytes(total).saturating_add(128 << 20); + let mut resident: Vec> = Vec::with_capacity(evaluations.len()); + let mut host: Vec>> = Vec::with_capacity(evaluations.len()); + let mut retained = Vec::with_capacity(evaluations.len()); + for (domain, source) in evaluations { + let shift = Val::GENERATOR / domain.shift(); + let plan = pcs + .dft + .lde_plan(source.height(), source.width(), blowup, shift); + let needed = source + .height() + .saturating_mul(source.width()) + .saturating_mul(8) + .saturating_mul((1 << blowup) + 1) + .saturating_add(plan.scratch_bytes()) + .saturating_add(plan.constant_bytes()) + .saturating_add(reserve); + // Generator caches go before any LDE spills: they are rebuilt from + // host seeds on demand, a spilled LDE is uploaded again. + for slot in resident.iter() { + if device_memory_info(device).0 >= needed { + break; + } + slot.as_ref().unwrap().release_generator_device(); + } + for index in 0..resident.len() { + if device_memory_info(device).0 >= needed { + break; + } + if host[index].is_none() { + let lde = resident[index].as_ref().unwrap(); + host[index] = Some(spill_lde(lde)); + } + } + assert!( + device_memory_info(device).0 >= needed, + "generated trace commitment exceeds device admission; reduce the shard cell budget" + ); + let source_span = tracing::info_span!( + "stark/commit_source", + kind = match &source { + TraceSource::Host(_) => "host", + TraceSource::Generated(_) => "generated", + }, + height = source.height(), + width = source.width() + ) + .entered(); + pcs.dft + .prepare_coset_lde_constants(source.height(), blowup, shift); + let (lde, trace) = match source { + TraceSource::Host(matrix) => ( + pcs.dft.coset_lde_batch_resident(&matrix, blowup, shift), + Some(matrix), + ), + TraceSource::Generated(source) => { + (pcs.dft.generate_coset_lde(source, blowup, shift), None) + } + }; + // Both kinds retain a bounded recovery source. Raw device rows need + // not coexist with later lookup/quotient workspace. + unsafe { lde.release_trace() }; + resident.push(Some(lde)); + host.push(None); + retained.push(trace); + drop(source_span); + } + let _merkle_span = tracing::info_span!("stark/commit_merkle").entered(); + let needed = max_height.saturating_mul(96).saturating_add(reserve); + for index in 0..resident.len() { + if device_memory_info(device).0 >= needed { + break; + } + if host[index].is_none() { + host[index] = Some(spill_lde(resident[index].as_ref().unwrap())); + } + } + assert!( + device_memory_info(device).0 >= needed, + "main Merkle tree exceeds device admission; reduce the shard cell budget" + ); + let mut spilled_heights: std::collections::BTreeSet<_> = host + .iter() + .enumerate() + .filter_map(|(index, h)| h.is_some().then_some(dimensions[index].height)) + .collect(); + // A height group wider than the device leaf kernel hashes is spilled + // whole and hashed on the host, like a group spilled for memory. + spilled_heights.extend(super::mmcs::host_hashed_heights(dimensions.iter().copied())); + for index in 0..host.len() { + if host[index].is_none() && spilled_heights.contains(&dimensions[index].height) { + let lde = resident[index].as_ref().unwrap(); + host[index] = Some(spill_lde(lde)); + } + } + let resident_refs: Vec<_> = resident + .iter() + .zip(&host) + .map(|(r, h)| if h.is_none() { r.as_ref() } else { None }) + .collect(); + let host_refs: Vec<_> = host.iter().map(Option::as_ref).collect(); + let deferred = vec![None; dimensions.len()]; + let digests = super::mmcs::hash_host_only_height_groups( + &host, + &resident_refs, + &deferred, + &std::collections::BTreeSet::new(), + ); + let tree = + CudaMixedMerkleTree::from_hybrid(device, &resident_refs, &host_refs, &deferred, &digests); + let commitment = MerkleCap::new(vec![tree.root()]); + let active = host.iter().map(|h| AtomicBool::new(h.is_none())).collect(); + let committed = host + .into_iter() + .map(|h| { + let cell = OnceLock::new(); + if let Some(h) = h { + cell.set(h).unwrap(); + } + cell + }) + .collect(); + let count = dimensions.len(); + ( + commitment, + CudaMmcsData::Hybrid { + resident, + resident_active: active, + materialize: Box::new(CudaLde::to_row_major_matrix), + committed_matrices: committed, + deferred_matrices: (0..count).map(|_| None).collect(), + dimensions, + retained_traces: retained, + tree, + }, + ) +} diff --git a/src/lib.rs b/src/lib.rs index 4172a75..c1f1aed 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -16,6 +16,7 @@ pub mod system; mod test_circuits; pub mod types; pub mod verifier; +pub mod witness; pub use p3_air; pub use p3_field; diff --git a/src/prover.rs b/src/prover.rs index bbc6d27..4bf1f01 100644 --- a/src/prover.rs +++ b/src/prover.rs @@ -434,7 +434,11 @@ where /// # Panics /// Panics if every circuit's trace is empty (nothing to prove). #[tracing::instrument(level = "info", skip_all, name = "stark/stage1_commit")] - pub fn prove_stage_1(&self, witness: SystemWitness>) -> Stage1 { + pub fn prove_stage_1( + &self, + witness: impl Into>>, + ) -> Stage1 { + let witness = witness.into(); let pcs = self.config.pcs(); // Sparse activation: a circuit whose stage-1 trace is empty is @@ -474,7 +478,8 @@ where log_degrees.push(log_degree); (trace_domain, trace) }); - let (stage_1_trace_commit, stage_1_trace_data) = pcs.commit(evaluations); + let (stage_1_trace_commit, stage_1_trace_data) = + self.config.commit_main(evaluations.collect()); // Only active circuits enter the accumulator chain; the chain (and // `intermediate_accumulators`) is indexed by active position. @@ -1210,45 +1215,6 @@ mod tests { use p3_dft::{NaiveDft, Radix2DitParallel}; use rand::{RngExt, SeedableRng, rngs::SmallRng}; - /// Host model of the CUDA radix-2 DIF kernel. It intentionally mirrors - /// `cuda/kernels.cu` stage/index/twiddle ordering and returns raw - /// bit-reversed storage. - fn cuda_dif_model(mut matrix: RowMajorMatrix, inverse: bool) -> RowMajorMatrix { - let height = matrix.height(); - if height <= 1 { - return matrix; - } - let width = matrix.width(); - let log_height = log2_strict_usize(height); - let root = Val::two_adic_generator(log_height); - let root = if inverse { root.inverse() } else { root }; - let twiddles: Vec<_> = root.powers().take(height / 2).collect(); - - let mut half = height / 2; - loop { - let stride = height / (2 * half); - for butterfly in 0..height / 2 { - let offset = butterfly % half; - let group = butterfly / half; - let row_0 = group * (2 * half) + offset; - let row_1 = row_0 + half; - for column in 0..width { - let index_0 = row_0 * width + column; - let index_1 = row_1 * width + column; - let left = matrix.values[index_0]; - let right = matrix.values[index_1]; - matrix.values[index_0] = left + right; - matrix.values[index_1] = (left - right) * twiddles[offset * stride]; - } - } - if half == 1 { - break; - } - half >>= 1; - } - matrix - } - /// `lde_from_coefficients` must reproduce, value for value, the matrix /// `TwoAdicFriPcs::commit` stores for the same polynomials given as /// trace-domain evaluations: `coset_lde_batch` with the generator shift @@ -1360,50 +1326,4 @@ mod tests { } } } - - /// Pins the first-party CUDA kernel's DIF ordering and fused coset-LDE - /// pipeline even on machines without a CUDA toolkit or GPU. - #[test] - fn cuda_dif_and_coset_lde_model_match_cpu() { - let mut rng = SmallRng::seed_from_u64(3); - let radix = Radix2DitParallel::::default(); - for log_height in [0usize, 1, 2, 5, 8] { - for width in [1usize, 2, 7] { - let height = 1 << log_height; - let matrix = - RowMajorMatrix::new((0..height * width).map(|_| rng.random()).collect(), width); - let expected_dft = radix - .dft_batch(matrix.clone()) - .bit_reverse_rows() - .to_row_major_matrix(); - assert_eq!(cuda_dif_model(matrix.clone(), false), expected_dft); - - for added_bits in [0usize, 1, 2, 3] { - // CUDA pipeline: inverse DIF (bit-reversed coefficients), - // bit-reverse + normalize + shift, zero-pad, forward DIF. - let mut coefficients = cuda_dif_model(matrix.clone(), true) - .bit_reverse_rows() - .to_row_major_matrix(); - let height_inverse = Val::ONE.div_2exp_u64(log_height as u64); - let mut shift = Val::ONE; - for row in 0..height { - for value in &mut coefficients.values[row * width..(row + 1) * width] { - *value *= height_inverse * shift; - } - shift *= Val::GENERATOR; - } - coefficients.pad_to_height(height << added_bits, Val::ZERO); - let actual = cuda_dif_model(coefficients, false); - let expected = radix - .coset_lde_batch(matrix.clone(), added_bits, Val::GENERATOR) - .bit_reverse_rows() - .to_row_major_matrix(); - assert_eq!( - actual, expected, - "height=2^{log_height}, blowup=2^{added_bits}, width={width}" - ); - } - } - } - } } diff --git a/src/types.rs b/src/types.rs index 09f8e7d..10ad8bc 100644 --- a/src/types.rs +++ b/src/types.rs @@ -277,6 +277,18 @@ pub struct GoldilocksBlake3Config { impl GoldilocksBlake3Config { pub fn new(commitment_parameters: CommitmentParameters, fri_parameters: FriParameters) -> Self { + Self::with_device(commitment_parameters, fri_parameters, None) + } + + /// A configuration whose prover lives on the given CUDA device (`None`: + /// the device `MULTI_STARK_CUDA_DEVICE` names, or 0). Without the + /// `cuda` feature only device 0 exists. Configurations on distinct + /// devices prove concurrently in one process. + pub fn with_device( + commitment_parameters: CommitmentParameters, + fri_parameters: FriParameters, + device_id: Option, + ) -> Self { #[cfg(feature = "cuda")] { assert_eq!( @@ -288,7 +300,20 @@ impl GoldilocksBlake3Config { "the CUDA backend currently supports only binary FRI folds" ); } - let (pcs, dft) = new_pcs(commitment_parameters, fri_parameters); + #[cfg(feature = "cuda")] + let (pcs, dft) = new_pcs( + commitment_parameters, + fri_parameters, + device_id.unwrap_or_else(crate::cuda::configured_device), + ); + #[cfg(not(feature = "cuda"))] + let (pcs, dft) = { + assert!( + device_id.is_none_or(|device| device == 0), + "no CUDA backend to place a prover on device {device_id:?}" + ); + new_pcs(commitment_parameters, fri_parameters) + }; // Seed the challenger with a protocol tag for domain separation, // followed by every protocol parameter. Binding the parameters into // the seed means transcripts produced under different parameters @@ -350,6 +375,17 @@ impl StarkGenericConfig for GoldilocksBlake3Config { self.log_blowup } + #[cfg(feature = "cuda")] + fn commit_main( + &self, + evaluations: Vec<( + crate::config::Domain, + crate::witness::TraceSource, + )>, + ) -> (crate::config::Com, crate::config::PcsData) { + crate::cuda::witness::commit(&self.pcs, evaluations) + } + fn canonicalize_proof(proof: &mut crate::prover::Proof) { fn canonical_base(value: &mut Val) { *value = Val::from_u64(value.as_canonical_u64()); @@ -529,11 +565,11 @@ impl StarkGenericConfig for GoldilocksBlake3Config { self.log_blowup, ); let current_staging = staging_bytes(input); - let constant_bytes = quotient_size - .saturating_add(lde_height) - .saturating_div(2) - .saturating_add(quotient_degree) - .saturating_mul(size_of::()); + let constant_bytes = quotient_degree.saturating_mul(size_of::()); + let quotient_plan = self.pcs.dft.forward_plan(quotient_size, 2); + let lde_plan = self.pcs.dft.forward_plan(lde_height, 2 * quotient_degree); + let kernel_workspace = kernel_workspace + .saturating_add(quotient_plan.scratch_bytes().max(lde_plan.scratch_bytes())); ( index, output_bytes, @@ -586,18 +622,24 @@ impl StarkGenericConfig for GoldilocksBlake3Config { .saturating_add(total_device_bytes / 64) }; let mut target = required(); - self.pcs - .mmcs - .ensure_device_headroom(input.stage_1.0, target, Some(input.stage_1.1)); + self.pcs.mmcs.ensure_device_headroom( + input.stage_1.0, + target, + Some(input.stage_1.1), + "quotient", + ); target = required(); - self.pcs - .mmcs - .ensure_device_headroom(input.stage_2.0, target, Some(input.stage_2.1)); + self.pcs.mmcs.ensure_device_headroom( + input.stage_2.0, + target, + Some(input.stage_2.1), + "quotient", + ); if let Some((data, matrix)) = input.preprocessed { target = required(); self.pcs .mmcs - .ensure_device_headroom(data, target, Some(matrix)); + .ensure_device_headroom(data, target, Some(matrix), "quotient"); } target = required(); let free_bytes = crate::cuda::device_memory_info(self.pcs.mmcs.cuda_device_id()).0; @@ -692,17 +734,23 @@ impl StarkGenericConfig for GoldilocksBlake3Config { .saturating_mul(96) .saturating_add(total_device_bytes / 64); if let Some(input) = inputs.first() { - self.pcs - .mmcs - .ensure_device_headroom(input.stage_1.0, tree_headroom, None); - self.pcs - .mmcs - .ensure_device_headroom(input.stage_2.0, tree_headroom, None); + self.pcs.mmcs.ensure_device_headroom( + input.stage_1.0, + tree_headroom, + None, + "quotient_tree", + ); + self.pcs.mmcs.ensure_device_headroom( + input.stage_2.0, + tree_headroom, + None, + "quotient_tree", + ); } if let Some((data, _)) = inputs.iter().find_map(|input| input.preprocessed) { self.pcs .mmcs - .ensure_device_headroom(data, tree_headroom, None); + .ensure_device_headroom(data, tree_headroom, None, "quotient_tree"); } if crate::cuda::device_memory_info(self.pcs.mmcs.cuda_device_id()).0 < tree_headroom { return None; @@ -765,14 +813,6 @@ impl StarkGenericConfig for GoldilocksBlake3Config { .saturating_mul(2 * size_of::()), ) .saturating_add(arg_offsets.len().saturating_mul(size_of::())) - // Device-cached inverse/forward twiddles and coset - // powers may be cold for this height. - .saturating_add( - height - .saturating_add(height / 2) - .saturating_add(extended_height / 2) - .saturating_mul(size_of::()), - ) }; let main_width = self .pcs @@ -792,11 +832,20 @@ impl StarkGenericConfig for GoldilocksBlake3Config { ) }) .flatten(); + let transform_bytes = if num_lookups == 0 { + 0 + } else { + let plan = + self.pcs + .dft + .lde_plan(height, 2 * groups, self.log_blowup, Val::GENERATOR); + plan.scratch_bytes().saturating_add(plan.constant_bytes()) + }; ( index, output_bytes, - direct_temporary_bytes, - graph_memory.map(|(_, temporary)| temporary), + direct_temporary_bytes.saturating_add(transform_bytes), + graph_memory.map(|(_, temporary)| temporary.saturating_add(transform_bytes)), extended_height, ) }) @@ -806,10 +855,10 @@ impl StarkGenericConfig for GoldilocksBlake3Config { // resident, then allow that trace to become an eviction candidate for // later jobs. lookup_jobs.sort_unstable_by_key(|&(index, output, direct, graph, _)| { - let temporary = if self - .pcs - .mmcs - .is_matrix_cuda_resident(inputs[index].stage_1.0, inputs[index].stage_1.1) + let temporary = if inputs[index] + .stage_1 + .0 + .has_resident_with_trace(inputs[index].stage_1.1) { graph.unwrap_or(direct) } else { @@ -846,16 +895,27 @@ impl StarkGenericConfig for GoldilocksBlake3Config { let ext_w = pair(extension_generator * extension_generator)[0]; let lookup_pool = std::sync::Arc::new(cuda_host_pool("lookup-rows", cuda_lookup_worker_count())); - let evaluate = |input: &crate::config::LookupCommitInput<'_, Self>, cooperative: bool| { + // `graph_path` is decided by admission below, which budgets the graph + // kernel from the same residency predicate; the kernel must not be + // re-chosen here or it can run under the other path's budget. + let evaluate = |input: &crate::config::LookupCommitInput<'_, Self>, + graph_path: bool, + cooperative: bool| { let (height, num_lookups, multiplicities, args, arg_offsets) = input.lookup_values.cuda_parts(); let group_size = input.circuit.lookup_group_size.max(1); - let main = input.stage_1.0.resident_with_trace(input.stage_1.1); - let result = if let Some(main) = main.filter(|_| num_lookups != 0) { - let preprocessed = match input.preprocessed { - Some((data, index)) => Some(data.resident(index)?), - None => None, - }; + let main = graph_path.then(|| { + input + .stage_1 + .0 + .resident_with_trace(input.stage_1.1) + .expect("lookup graph admission requires a trace-backed resident LDE") + }); + let result = if let Some(main) = main { + let preprocessed = input.preprocessed.map(|(data, index)| { + data.resident(index) + .expect("lookup graph admission requires a resident preprocessed LDE") + }); crate::cuda::lookup_graph_lde_resident( &self.pcs.dft, &input.circuit.graph, @@ -912,6 +972,7 @@ impl StarkGenericConfig for GoldilocksBlake3Config { // retained matrix remains owned by the prover data. if let Some(main) = main { unsafe { main.release_trace() }; + main.release_generator_device(); } result }; @@ -919,22 +980,32 @@ impl StarkGenericConfig for GoldilocksBlake3Config { for (index, output_bytes, direct_temporary_bytes, graph_temporary_bytes, _) in lookup_jobs { let job_started = std::time::Instant::now(); let input = &inputs[index]; - let graph_path = self - .pcs - .mmcs - .is_matrix_cuda_resident(input.stage_1.0, input.stage_1.1) - && graph_temporary_bytes.is_some() - && input.preprocessed.is_none_or(|(data, matrix)| { - self.pcs.mmcs.is_matrix_cuda_resident(data, matrix) - }); - let temporary_bytes = if graph_path { - graph_temporary_bytes.unwrap() - } else { - direct_temporary_bytes + // The graph kernel reads the trace through `resident_with_trace`, + // which also serves spilled LDEs that regenerate from a generator + // or a retained host trace. Admit it under the same predicate so + // the budget matches the kernel that runs. + let admits_graph = || { + graph_temporary_bytes.is_some() + && input.stage_1.0.has_resident_with_trace(input.stage_1.1) + && input + .preprocessed + .is_none_or(|(data, matrix)| data.resident(matrix).is_some()) }; - let target = output_bytes - .saturating_add(temporary_bytes) - .saturating_add(total_device_bytes / 64); + let target_for = |graph_path: bool| { + let temporary_bytes = if graph_path { + graph_temporary_bytes.unwrap() + } else { + direct_temporary_bytes + }; + ( + temporary_bytes, + output_bytes + .saturating_add(temporary_bytes) + .saturating_add(total_device_bytes / 64), + ) + }; + let mut graph_path = admits_graph(); + let (mut temporary_bytes, mut target) = target_for(graph_path); if target > total_device_bytes { return None; } @@ -942,15 +1013,28 @@ impl StarkGenericConfig for GoldilocksBlake3Config { input.stage_1.0, target, Some(input.stage_1.1), + "lookup", ); if free_bytes < target { // This circuit alone does not fit beside its resident trace. - // Spill it as a last resort and use the direct lookup-values - // path, which remains protocol-identical. - free_bytes = self - .pcs - .mmcs - .ensure_device_headroom(input.stage_1.0, target, None); + // Spill it as a last resort. If the spilled LDE can still + // regenerate its trace it stays on the graph path; otherwise + // re-budget for the direct lookup-values path, which remains + // protocol-identical. + free_bytes = + self.pcs + .mmcs + .ensure_device_headroom(input.stage_1.0, target, None, "lookup"); + graph_path = admits_graph(); + (temporary_bytes, target) = target_for(graph_path); + if free_bytes < target { + free_bytes = self.pcs.mmcs.ensure_device_headroom( + input.stage_1.0, + target, + None, + "lookup", + ); + } } if free_bytes < target { return None; @@ -965,7 +1049,7 @@ impl StarkGenericConfig for GoldilocksBlake3Config { && output_bytes >= (8usize << 30) && num_lookups >= 64 && arg_offsets.last().copied().unwrap_or(0) >= 256; - results[index] = Some(evaluate(&inputs[index], cooperative)?); + results[index] = Some(evaluate(&inputs[index], graph_path, cooperative)?); if crate::cuda::memory_diagnostics_enabled() { eprintln!( "[multi-stark/cuda] lookup job {index} complete: {:.3}s", @@ -977,10 +1061,12 @@ impl StarkGenericConfig for GoldilocksBlake3Config { .saturating_mul(96) .saturating_add(total_device_bytes / 64); if let Some(input) = inputs.first() { - let free_bytes = - self.pcs - .mmcs - .ensure_device_headroom(input.stage_1.0, tree_headroom, None); + let free_bytes = self.pcs.mmcs.ensure_device_headroom( + input.stage_1.0, + tree_headroom, + None, + "lookup_tree", + ); if free_bytes < tree_headroom { return None; } @@ -1032,6 +1118,20 @@ pub(crate) type Blake3CompressionFunction = CompressionFunctionFromHasher for CudaDft { + fn coset_lde_workspace_bytes( + &self, + height: usize, + width: usize, + added_bits: usize, + shift: Val, + ) -> usize { + if width == 0 { + return 0; + } + let plan = self.lde_plan(height, width, added_bits, shift); + plan.scratch_bytes().saturating_add(plan.constant_bytes()) + } + fn prepare_coset_lde_constants(&self, height: usize, added_bits: usize, shift: Val) { self.prepare_coset_lde_constants(height, added_bits, shift); } @@ -1052,7 +1152,7 @@ type PcsDft = Dft; #[cfg(feature = "cuda")] type PcsDft = CudaDft; -fn new_mmcs(cap_height: usize) -> Mmcs { +fn new_mmcs(cap_height: usize, #[cfg(feature = "cuda")] device_id: i32) -> Mmcs { let byte_hash = Blake3; let field_hash = SerializingHasher::new(byte_hash); let compress = Blake3CompressionFunction::new(byte_hash); @@ -1060,13 +1160,17 @@ fn new_mmcs(cap_height: usize) -> Mmcs { #[cfg(not(feature = "cuda"))] return cpu; #[cfg(feature = "cuda")] - crate::cuda::mmcs::CudaMmcs::new(cpu) + crate::cuda::mmcs::CudaMmcs::with_device(cpu, device_id) } fn new_pcs( commitment_parameters: CommitmentParameters, fri_parameters: FriParameters, + #[cfg(feature = "cuda")] device_id: i32, ) -> (Pcs, Dft) { + #[cfg(feature = "cuda")] + let val_mmcs = new_mmcs(commitment_parameters.cap_height, device_id); + #[cfg(not(feature = "cuda"))] let val_mmcs = new_mmcs(commitment_parameters.cap_height); let mmcs = ExtensionMmcs::new(val_mmcs.clone()); let inner_parameters = InnerFriParameters { @@ -1079,7 +1183,11 @@ fn new_pcs( mmcs, }; let dft = Dft::default(); - let pcs = Pcs::new(PcsDft::default(), val_mmcs, inner_parameters); + #[cfg(feature = "cuda")] + let pcs_dft = CudaDft::new(device_id); + #[cfg(not(feature = "cuda"))] + let pcs_dft = PcsDft::default(); + let pcs = Pcs::new(pcs_dft, val_mmcs, inner_parameters); (pcs, dft) } @@ -1273,7 +1381,11 @@ mod pcs_ref_gen { m1[8] = f(109); // row 2 = [107, 108, 109] let mut m2 = vec![f(0); 2]; m2[1] = f(202); // row 1 = [202] - let mmcs = new_mmcs(0); + let mmcs = new_mmcs( + 0, + #[cfg(feature = "cuda")] + 0, + ); let (commit, pd) = mmcs.commit(vec![ RowMajorMatrix::new(m0.clone(), 2), RowMajorMatrix::new(m1.clone(), 3), @@ -1323,7 +1435,11 @@ mod pcs_ref_gen { ]] ); - let mmcs = new_mmcs(2); + let mmcs = new_mmcs( + 2, + #[cfg(feature = "cuda")] + 0, + ); let (commit, pd) = mmcs.commit(vec![ RowMajorMatrix::new(m0, 2), RowMajorMatrix::new(m1, 3), diff --git a/src/witness.rs b/src/witness.rs new file mode 100644 index 0000000..3f29286 --- /dev/null +++ b/src/witness.rs @@ -0,0 +1,77 @@ +//! Owned, deterministic main-trace sources. Producers do not allocate device memory. + +use std::sync::Arc; + +use p3_field::Field; +use p3_matrix::{Matrix, dense::RowMajorMatrix}; + +/// A frozen source which reproduces the same canonical cells on every call. +pub trait TraceGenerator: Send + Sync { + fn height(&self) -> usize; + fn width(&self) -> usize; + /// Fill contiguous rows, wrapping at the padded height for lookup halos. + fn write_rows(&self, first: usize, output: &mut [F]); + + /// Fill a device tile synchronously on its owning device and calling stream. + #[cfg(feature = "cuda")] + fn write_device_rows(&self, output: crate::cuda::DeviceTraceView<'_>) -> Result<(), String>; + + /// Free whatever the source keeps on `device_id` to serve tiles faster, + /// such as seeds left resident between a commitment and the lookup pass + /// that regenerates rows from them. Tiles must still be served afterwards. + #[cfg(feature = "cuda")] + fn release_device(&self, _device_id: i32) {} +} + +#[derive(Clone)] +pub enum TraceSource { + Host(RowMajorMatrix), + Generated(Arc>), +} + +impl TraceSource { + pub fn height(&self) -> usize { + match self { + Self::Host(m) => m.height(), + Self::Generated(g) => g.height(), + } + } + + pub fn width(&self) -> usize { + match self { + Self::Host(m) => m.width(), + Self::Generated(g) => g.width(), + } + } + + pub fn materialize(self) -> RowMajorMatrix { + match self { + Self::Host(m) => m, + Self::Generated(g) => { + let mut values = vec![ + F::ZERO; + g.height() + .checked_mul(g.width()) + .expect("trace size overflow") + ]; + g.write_rows(0, &mut values); + RowMajorMatrix::new(values, g.width()) + } + } + } +} + +#[derive(Clone)] +pub struct PreparedWitness { + pub traces: Vec>, + pub lookups: Vec>, +} + +impl From> for PreparedWitness { + fn from(witness: crate::system::SystemWitness) -> Self { + Self { + traces: witness.traces.into_iter().map(TraceSource::Host).collect(), + lookups: witness.lookups, + } + } +}