From b372a43519a2eb3426503ec18dfbcb4b363e0f23 Mon Sep 17 00:00:00 2001 From: "Claude (ArcGraph chain)" Date: Tue, 18 Aug 2026 11:27:44 +0000 Subject: [PATCH] fix(ArcGraph): correct cuGraphAddNode ABI; make GPU-autonomous decode reachable The headline is an ABI mismatch that explains a failure already on our record. `ffi.rs` declared `cudaGraphAddNode` with FIVE arguments and the first two transposed. CUDA 13.1 `cuda.h:21829` declares SIX, and the first parameter is an OUT pointer: CUresult cuGraphAddNode(CUgraphNode *phGraphNode, CUgraph hGraph, const CUgraphNode *dependencies, const CUgraphEdgeData *dependencyData, size_t numDependencies, CUgraphNodeParams *nodeParams); So the driver received a `CUgraph` handle where it expected `CUgraphNode *` and wrote the new node handle THROUGH it, while `numDependencies` received a pointer and `nodeParams` an uninitialised register. That is memory corruption from a declaration, and it matches the previously unexplained "host heap corruption" CUDA-graph capture failure we had recorded with no cause. Four more defects, each of which alone stops the autonomous loop: * Conditional handles were created with `flags = 0`. `cuda.h:21919` applies `defaultLaunchValue` only when `CU_GRAPH_COND_ASSIGN_DEFAULT` is set, so the condition read 0 at launch and a WHILE body executed ZERO times while every CUDA call returned success. Measured both ways on an H200. * The WHILE body was captured into a throwaway graph and then DESTROYED, leaving the conditional body empty. Bodies are now populated with `cuStreamBeginCaptureToGraph` (`cuda.h:1976` names it as the supported way), and a node-count assert refuses a body that recorded nothing. * `CUDA_CONDITIONAL_NODE_PARAMS` was missing its 5th field, `ctx`. * `DecodeState` allocated I64 while every kernel signature is `int32_t*`, so host and device disagreed about where element `i` lived; and `reset()` reallocated the buffers, silently invalidating the pointers a captured graph had baked in. `tensor_device_ptr` had no I32 arm, so `CudaSampler::sample` returned `unsupported dtype I32` on EVERY call -- the correct top-k/top-p sampler had never once executed on a GPU. Its test suite is a CPU simulator, which by construction cannot observe that. Sampling now runs on device and survives replay. The sampler previously wired into the graph body took `rng_offset` as a BY-VALUE kernel argument, which a captured graph bakes: measured 1 distinct token over 64 replays, versus 64 distinct for a device-resident `rng_state`. It also was not nucleus sampling -- it walked the vocabulary in TOKEN-ID order. The replacement is a hybrid, because measurement showed neither component wins everywhere (H200, vocab 129280, host+GPU verified exclusive before and after; support width -> legacy / hybrid us): 1 -> 134.9/166.6 8 -> 503.0/624.2 64 -> 3127.8/4178.5 512 -> 24320.7/4185.0 4096 -> 205125.2/4187.3 12928 -> 664200.0/4188.5 Exact enumeration while the nucleus is small, threshold bisection past a fixed budget, chosen block-uniformly ON DEVICE because a captured graph cannot ask the host which branch to take. Costs +24% on the peaked distributions real models produce -- 32 us, 0.05% of a 66.68 ms V4 decode step -- and removes a 664 ms cliff that was ~10x an entire decode step for a single token. No cap and no narrowing: exceeding the budget changes which ALGORITHM selects the nucleus, never which tokens it contains. Verified against the enumerating sampler with branch counters proving which path each case exercised: peaked TV=0.0000 A=20000 B=0 scrambled order TV=0.0000 A=20000 B=0 (15.6% per-draw agreement, so the orders genuinely diverged) diffuse TV=0.0000 A=0 B=3000 tie boundary TV=0.0000 A=20000 B=0 narrowed (ctrl) TV=0.5027 -> correctly DISAGREE Known and deliberately left: on TIED diffuse supports the fallback keeps 116 tokens where enumeration keeps 111 (TV=0.1103) -- it keeps every token sharing a boundary bit pattern, erring WIDER, never narrower. And the budget is derived from bisection's pass count (32) when it should follow the FALLBACK's (~70), which costs a 0.75x window at support ~64; correcting it needs a more diffuse fixture so the fallback stays covered. Not claimed: no V4 end-to-end number. Autonomous decode is unreachable on V4 by construction -- no PagedAttention means `cache_config` is None and the runner is never built (`normal.rs:1907`, pinned by `normal_loaders.rs:5687`) -- and `cuda.h:1971` bars alloc/free nodes inside a conditional body against 11,436 allocations per token. Co-Authored-By: Claude Opus 5 (1M context) --- arc-cuda-graph/src/autonomous.rs | 244 +++++++--- arc-cuda-graph/src/buffers.rs | 106 +++- arc-cuda-graph/src/cuda/sampling_kernel.cu | 456 ++++++++++++++++++ arc-cuda-graph/src/ffi.rs | 68 ++- arc-cuda-graph/src/sampling_cuda.rs | 68 ++- arc-cuda-graph/src/weights.rs | 10 + mistralrs-core/src/pipeline/normal.rs | 10 + .../tests/capability_reachability.rs | 25 +- mistralrs-core/tests/doc_citations.rs | 24 + 9 files changed, 902 insertions(+), 109 deletions(-) diff --git a/arc-cuda-graph/src/autonomous.rs b/arc-cuda-graph/src/autonomous.rs index 2539b2801..6153d7699 100644 --- a/arc-cuda-graph/src/autonomous.rs +++ b/arc-cuda-graph/src/autonomous.rs @@ -78,6 +78,8 @@ pub struct AutonomousDecodeConfig { pub top_p: f32, pub frequency_penalty: f32, pub presence_penalty: f32, + /// `-1` disables top-k. The previous in-graph sampler had no top-k at all. + pub top_k: i32, pub greedy: bool, } @@ -101,7 +103,19 @@ pub struct AutonomousDecodeRunner { /// True if the graph uses a WHILE conditional node (CUDA 12.4+). /// False if host-driven loop (body graph launched per step). uses_while_node: bool, - rng_offset: u64, + /// The real sampler: temperature + top-k + top-p, with its RNG state held + /// in **device memory** and advanced by the kernel. + /// + /// 🔴 The sampler this replaces (`launch_fused_top_p_bf16`) took + /// `rng_offset` as a **by-value kernel argument**. A captured graph bakes + /// kernel arguments at capture time, so every replayed step drew the + /// identical uniform and the host-side `rng_offset += 1` could not reach + /// an instantiated graph. It was also not nucleus sampling: it walked the + /// vocabulary in **token-id order** accumulating full-distribution mass to + /// `top_p * u`, biasing hard toward low token ids (its own comment said + /// "Not quite right — proper top-p needs sorting"). A device-resident + /// `rng_state` is the only form that survives replay. + sampler: crate::sampling_cuda::CudaSampler, } #[cfg(feature = "cuda")] @@ -171,6 +185,17 @@ impl AutonomousDecodeRunner { std::ptr::write_bytes(ring_write_head_ptr as *mut u8, 0, batch * 4); } + // Bound to the runner's own stream so its kernels land inside the + // captured body rather than on the device-default stream. + let mut sampler = crate::sampling_cuda::CudaSampler::new( + device, + config.padded_batch_size, + config.vocab_size, + candle_core::DType::BF16, + 0x5EED_A11C_E571_2345, + )?; + sampler.set_stream(stream); + Ok(Self { config, device: device.clone(), @@ -183,7 +208,7 @@ impl AutonomousDecodeRunner { ring_size, graph_exec: None, uses_while_node: false, - rng_offset: 0, + sampler, }) } @@ -350,10 +375,30 @@ impl AutonomousDecodeRunner { candle_core::bail!("cuGraphCreate failed: {status}"); } - // 2. Create conditional handle (WHILE, default_value=1 = loop) + // 2. Create conditional handle (WHILE, default_value=1 = loop). + // + // `CU_GRAPH_COND_ASSIGN_DEFAULT` is REQUIRED: without it the default + // is never applied at launch, the condition reads 0, and the WHILE + // body runs zero times while every CUDA call still returns success. + let mut ctx: CUcontext = std::ptr::null_mut(); + let status = unsafe { cuCtxGetCurrent(&mut ctx) }; + if status != CUDA_SUCCESS || ctx.is_null() { + unsafe { + cuGraphDestroy(outer_graph); + } + tracing::warn!("cuCtxGetCurrent failed ({status}), falling back to host-driven loop"); + return self.capture_body_graph(forward_fn, bs, vocab); + } let mut cond_handle: CUgraphConditionalHandle = 0; - let status = - unsafe { cudaGraphConditionalHandleCreate(&mut cond_handle, outer_graph, 1, 0) }; + let status = unsafe { + cuGraphConditionalHandleCreate( + &mut cond_handle, + outer_graph, + ctx, + 1, + CU_GRAPH_COND_ASSIGN_DEFAULT, + ) + }; if status != CUDA_SUCCESS { unsafe { cuGraphDestroy(outer_graph); @@ -365,18 +410,19 @@ impl AutonomousDecodeRunner { // 3. Add conditional WHILE node to outer graph // This creates an empty body graph that we populate via stream capture. let mut while_node: CUgraphNode = std::ptr::null_mut(); - let mut body_graph: CUgraph = std::ptr::null_mut(); + // `phGraph_out` is an OUT field the driver populates; leave it null. let mut params: CudaGraphNodeParams = unsafe { std::mem::zeroed() }; params.node_type = CudaGraphNodeType::Conditional; params.conditional.handle = cond_handle; params.conditional.cond_type = CUgraphConditionalNodeType::WHILE; params.conditional.size = 1; - params.conditional.body_graph_out = &mut body_graph; + params.conditional.ctx = ctx; let status = unsafe { - cudaGraphAddNode( - outer_graph, + cuGraphAddNode( &mut while_node, + outer_graph, + std::ptr::null(), std::ptr::null(), 0, &mut params, @@ -390,19 +436,51 @@ impl AutonomousDecodeRunner { return self.capture_body_graph(forward_fn, bs, vocab); } - // 4. Populate body graph via stream capture + // 4. Populate the conditional node's OWN body graph. + // + // The driver hands back the body graph in `phGraph_out`; capture + // straight into it with `cuStreamBeginCaptureToGraph`. The previous + // code captured with `cuStreamBeginCapture_v2` — which always creates + // a NEW graph — and then destroyed the result, leaving the conditional + // body EMPTY. An empty body generates no tokens. + let body_graph: CUgraph = unsafe { + if params.conditional.body_graph_out.is_null() { + cuGraphDestroy(outer_graph); + candle_core::bail!("conditional node returned no body graph array"); + } + *params.conditional.body_graph_out + }; + if body_graph.is_null() { + unsafe { + cuGraphDestroy(outer_graph); + } + candle_core::bail!("conditional node body graph is null"); + } + unsafe { - let status = cuStreamBeginCapture_v2(self.stream, CUstreamCaptureMode::THREAD_LOCAL); + // RELAXED, not THREAD_LOCAL: candle's allocator and helper streams + // create cross-stream dependencies that THREAD_LOCAL rejects with + // CUDA_ERROR_STREAM_CAPTURE_ISOLATION (same reason `graph.rs` + // uses RELAXED). + let status = cuStreamBeginCaptureToGraph( + self.stream, + body_graph, + std::ptr::null(), + std::ptr::null(), + 0, + CUstreamCaptureMode::RELAXED, + ); if status != CUDA_SUCCESS { cuGraphDestroy(outer_graph); - candle_core::bail!("cuStreamBeginCapture for WHILE body failed: {status}"); + candle_core::bail!("cuStreamBeginCaptureToGraph for WHILE body failed: {status}"); } } // Capture: forward → sample → step_update → check_done_conditional self.capture_body_kernels(forward_fn, bs, vocab, Some(cond_handle))?; - // End capture into the body graph + // Ends the capture; returns the graph we captured INTO (`body_graph`), + // which is already attached to the conditional node. let mut captured_body: CUgraph = std::ptr::null_mut(); unsafe { let status = cuStreamEndCapture(self.stream, &mut captured_body); @@ -411,29 +489,37 @@ impl AutonomousDecodeRunner { candle_core::bail!("cuStreamEndCapture for WHILE body failed: {status}"); } } - - // The captured_body needs to be merged into the body_graph that - // cudaGraphAddNode created. For CUDA 12.4 conditional nodes, - // the body_graph_out is pre-created and we should have captured - // INTO it by using it as the capture target. However, stream - // capture always creates a new graph. - // - // The correct approach: don't use stream capture for the body. - // Instead, add kernel nodes to body_graph directly. But that's - // extremely complex (need to manually create kernel nodes for - // every cuBLAS call, every custom kernel, etc.). - // - // Alternative: use cudaStreamBeginCaptureToGraph (CUDA 12.3+) - // which captures into an existing graph. - // - // For now: instantiate the outer graph. If the body graph was - // properly populated by the conditional node setup, it works. - // If not, we destroy and fall back. - unsafe { - cuGraphDestroy(captured_body); + if captured_body != body_graph { + unsafe { + cuGraphDestroy(outer_graph); + } + candle_core::bail!( + "capture-to-graph returned a different graph than the conditional body" + ); + } + // Assert the body actually received nodes. A conditional node with an + // empty body is the failure this function exists to prevent, and it is + // invisible at launch: the graph runs and produces nothing. + let mut body_nodes: usize = 0; + let status = unsafe { cuGraphGetNodes(body_graph, std::ptr::null_mut(), &mut body_nodes) }; + if status != CUDA_SUCCESS || body_nodes == 0 { + unsafe { + cuGraphDestroy(outer_graph); + } + candle_core::bail!( + "WHILE body graph has {body_nodes} nodes after capture (status {status}) — \ + the decode body did not record" + ); } + tracing::info!("CUDA graph: WHILE body captured with {body_nodes} nodes"); - // 5. Instantiate outer graph + // 5. Instantiate outer graph. + // + // NOTE: no AUTO_FREE_ON_LAUNCH here. `cuda.h:1971` restricts a + // conditional body to "kernel nodes, empty nodes, child graphs, + // memsets, memcopies, and conditionals" — memory alloc/free nodes are + // NOT in that list, so the captured decode body must be + // allocation-free for this path to instantiate at all. let mut exec: CUgraphExec = std::ptr::null_mut(); let status = unsafe { cuGraphInstantiate_v2( @@ -510,7 +596,7 @@ impl AutonomousDecodeRunner { /// Capture the body kernels: forward → sample → step_update → check_done. /// Called during stream capture (kernels recorded, not executed). fn capture_body_kernels( - &self, + &mut self, forward_fn: &F, bs: i32, vocab: i32, @@ -521,48 +607,54 @@ impl AutonomousDecodeRunner { { // Forward pass let logits = forward_fn()?; - let logits_ptr = tensor_ptr(&logits)? as *const _; - // Sampling - let sampled_ptr = tensor_ptr(&self.decode_state.sampled_tokens)? as *mut _; - - unsafe { - if self.config.greedy { - launch_fused_argmax_bf16( - logits_ptr, - sampled_ptr, - std::ptr::null_mut(), - vocab, - bs, - self.stream, - ); - } else { - if self.config.frequency_penalty != 0.0 || self.config.presence_penalty != 0.0 { - launch_apply_penalties( - logits_ptr as *mut _, - tensor_ptr(&self.decode_state.output_tokens)? as *const i32, - tensor_ptr(&self.decode_state.n_generated)? as *const i32, - self.config.frequency_penalty, - self.config.presence_penalty, - vocab, - self.config.max_tokens as i32, - bs, - self.stream, - ); - } - launch_fused_top_p_bf16( - logits_ptr, - sampled_ptr, - self.config.temperature, - self.config.top_p, + // Cheap handle clones so every immutable borrow of `self` is finished + // before we take `&mut self.sampler`. + let sampled = self.decode_state.sampled_tokens.clone(); + + // Penalties stay a separate in-place pre-pass over the logits. The + // fused sampler applies penalties from a [batch, vocab] count tensor + // (`freq_counts`), which we do not maintain; passing `None` there + // while leaving non-zero penalties in the config would drop them + // SILENTLY (`sampling_kernel.cu:167` — `fcnt = freq_counts ? .. : + // nullptr`). So apply them here and zero them in `cfg`. + if self.config.frequency_penalty != 0.0 || self.config.presence_penalty != 0.0 { + let logits_mut = tensor_ptr(&logits)? as *mut _; + let out_toks = tensor_ptr(&self.decode_state.output_tokens)? as *const i32; + let n_gen = tensor_ptr(&self.decode_state.n_generated)? as *const i32; + unsafe { + launch_apply_penalties( + logits_mut, + out_toks, + n_gen, + self.config.frequency_penalty, + self.config.presence_penalty, vocab, + self.config.max_tokens as i32, bs, - 42, - self.rng_offset, self.stream, ); } + } + + let cfg = crate::sampling_cpu::SamplingConfig { + temperature: self.config.temperature, + top_p: self.config.top_p, + top_k: self.config.top_k, + // applied above, in-place, so the kernel must not re-apply them + frequency_penalty: 0.0, + presence_penalty: 0.0, + greedy: self.config.greedy, + eos_token_id: self.config.eos_token_id, + }; + // Scoped so the mutable borrow ends before the immutable ones below. + { + self.sampler.sample(&logits, None, cfg, &sampled)?; + } + let sampled_ptr = tensor_ptr(&sampled)? as *mut i32; + + unsafe { // Step update launch_decode_step_update( sampled_ptr as *const i32, @@ -621,9 +713,8 @@ impl AutonomousDecodeRunner { candle_core::Error::Msg("Graph not captured — call capture() first".into()) })?; - // Reset state - self.decode_state - .reset(&self.device, self.config.padded_batch_size)?; + // Reset state in place — the captured graph holds these addresses. + self.decode_state.reset(self.stream)?; unsafe { std::ptr::write_bytes( self.ring_write_head_ptr as *mut u8, @@ -631,7 +722,6 @@ impl AutonomousDecodeRunner { self.config.padded_batch_size * 4, ); } - self.rng_offset += 1; if self.uses_while_node { // ============================================================ @@ -656,7 +746,9 @@ impl AutonomousDecodeRunner { } cudaStreamSynchronize(self.stream); } - let cond = self.decode_state.loop_condition.to_vec1::()?; + // I32: the kernel writes `int32_t`. Reading this as i64 + // consumed two adjacent kernel writes as one value. + let cond = self.decode_state.loop_condition.to_vec1::()?; if cond[0] == 0 { break; } @@ -668,7 +760,7 @@ impl AutonomousDecodeRunner { .decode_state .output_tokens .to_dtype(candle_core::DType::I64)?; - let n_gen = self.decode_state.n_generated.to_vec1::()?; + let n_gen = self.decode_state.n_generated.to_vec1::()?; let mut results = Vec::new(); for b in 0..self.config.padded_batch_size { let n = n_gen[b] as usize; diff --git a/arc-cuda-graph/src/buffers.rs b/arc-cuda-graph/src/buffers.rs index 6b8d00f0f..df3ec623a 100644 --- a/arc-cuda-graph/src/buffers.rs +++ b/arc-cuda-graph/src/buffers.rs @@ -1,8 +1,45 @@ //! Pre-allocated GPU buffers for the decode loop. +//! +//! **Every buffer here is read and written by CUDA kernels whose signatures +//! are `int32_t*`** (`cuda/decode_loop.cu`, `cuda/sampling.cu`). The dtype of +//! the Candle tensor therefore has to be 4 bytes wide, or the host and the +//! device disagree about where element `i` lives. #[cfg(feature = "cuda")] use candle_core::{DType, Device, Tensor}; +#[cfg(feature = "cuda")] +use candle_core::cuda::cudarc::driver::sys::CUstream; + +#[cfg(feature = "cuda")] +extern "C" { + fn cudaMemsetAsync( + dst: *mut std::ffi::c_void, + value: i32, + count: usize, + stream: CUstream, + ) -> u32; + fn cudaMemcpyAsync( + dst: *mut std::ffi::c_void, + src: *const std::ffi::c_void, + count: usize, + kind: u32, + stream: CUstream, + ) -> u32; +} + +/// Allocate a zeroed I32 tensor. `Tensor::zeros` on I32 is not supported on +/// every Candle backend (see the same note in `sampling_cuda.rs:296`), so go +/// through `from_vec`. +#[cfg(feature = "cuda")] +fn zeros_i32( + elems: usize, + shape: impl Into, + device: &Device, +) -> candle_core::Result { + Tensor::from_vec(vec![0i32; elems], shape, device) +} + /// Pre-allocated decode input buffers at fixed GPU addresses. #[cfg(feature = "cuda")] pub struct DecodeInputBuffers { @@ -31,33 +68,74 @@ impl DecodeInputBuffers { } /// GPU-side decode state that persists across WHILE loop iterations. +/// +/// 🔴 These were all `I64` while every kernel that touches them declares +/// `int32_t*` (`decode_loop.cu:19-37`, `sampling.cu:21`). The kernel wrote +/// element `b` at byte `4*b`; Candle read element `b` at byte `8*b`. Nothing +/// lined up: `n_generated`, `finished`, `loop_condition` and the whole +/// `output_tokens` matrix were read back as garbage, and `output_tokens` was +/// additionally indexed with the wrong row stride. The path is gated off +/// before it runs, so nothing ever reported it. #[cfg(feature = "cuda")] pub struct DecodeState { - pub sampled_tokens: Tensor, // [padded_bs] i64 - pub n_generated: Tensor, // [padded_bs] i64 - pub output_tokens: Tensor, // [padded_bs, max_tokens] i64 - pub finished: Tensor, // [padded_bs] i64 - pub loop_condition: Tensor, // [1] i64 + pub sampled_tokens: Tensor, // [padded_bs] i32 + pub n_generated: Tensor, // [padded_bs] i32 + pub output_tokens: Tensor, // [padded_bs, max_tokens] i32 + pub finished: Tensor, // [padded_bs] i32 + pub loop_condition: Tensor, // [1] i32 pub max_tokens: usize, } #[cfg(feature = "cuda")] impl DecodeState { pub fn new(padded_bs: usize, max_tokens: usize, device: &Device) -> candle_core::Result { + let _ = DType::I32; // keep the dtype named in this module Ok(Self { - sampled_tokens: Tensor::zeros(padded_bs, DType::I64, device)?, - n_generated: Tensor::zeros(padded_bs, DType::I64, device)?, - output_tokens: Tensor::zeros((padded_bs, max_tokens), DType::I64, device)?, - finished: Tensor::zeros(padded_bs, DType::I64, device)?, - loop_condition: Tensor::ones(1, DType::I64, device)?, + sampled_tokens: zeros_i32(padded_bs, padded_bs, device)?, + n_generated: zeros_i32(padded_bs, padded_bs, device)?, + output_tokens: zeros_i32(padded_bs * max_tokens, (padded_bs, max_tokens), device)?, + finished: zeros_i32(padded_bs, padded_bs, device)?, + loop_condition: Tensor::from_vec(vec![1i32], 1, device)?, max_tokens, }) } - pub fn reset(&mut self, device: &Device, padded_bs: usize) -> candle_core::Result<()> { - self.n_generated = Tensor::zeros(padded_bs, DType::I64, device)?; - self.finished = Tensor::zeros(padded_bs, DType::I64, device)?; - self.loop_condition = Tensor::ones(1, DType::I64, device)?; + /// Reset the per-generation state **in place**, on `stream`. + /// + /// 🔴 This used to re-`Tensor::zeros` each field, which allocates NEW + /// device buffers at NEW addresses. The captured CUDA graph has the OLD + /// addresses baked into its kernel nodes, so resetting between + /// generations pointed the graph at freed memory — the reset silently + /// un-did the capture. A pointer-stable buffer is a hard requirement of + /// anything that gets captured, so zero the existing allocations instead. + pub fn reset(&mut self, stream: CUstream) -> candle_core::Result<()> { + let n_gen = crate::weights::tensor_device_ptr(&self.n_generated)?; + let fin = crate::weights::tensor_device_ptr(&self.finished)?; + let cond = crate::weights::tensor_device_ptr(&self.loop_condition)?; + let bs = self.n_generated.elem_count(); + unsafe { + let s = cudaMemsetAsync(n_gen as *mut _, 0, bs * 4, stream); + if s != 0 { + candle_core::bail!("cudaMemsetAsync(n_generated) failed: {s}"); + } + let s = cudaMemsetAsync(fin as *mut _, 0, bs * 4, stream); + if s != 0 { + candle_core::bail!("cudaMemsetAsync(finished) failed: {s}"); + } + // loop_condition starts at 1 ("keep going"). memset writes a byte + // pattern, so 0x01010101 would be wrong; write the word directly. + let one: i32 = 1; + let s = cudaMemcpyAsync( + cond as *mut _, + &one as *const i32 as *const _, + 4, + 1, // cudaMemcpyHostToDevice + stream, + ); + if s != 0 { + candle_core::bail!("cudaMemcpyAsync(loop_condition) failed: {s}"); + } + } Ok(()) } } diff --git a/arc-cuda-graph/src/cuda/sampling_kernel.cu b/arc-cuda-graph/src/cuda/sampling_kernel.cu index f9d5d9f02..53caaf192 100644 --- a/arc-cuda-graph/src/cuda/sampling_kernel.cu +++ b/arc-cuda-graph/src/cuda/sampling_kernel.cu @@ -127,6 +127,12 @@ __device__ __forceinline__ float block_sum(float val, float* s_vals) { return s_vals[0]; } +// Branch engagement counters for the hybrid sampler. A test that assumes +// "support width 512 must have taken the fallback" is assuming the very thing +// it is trying to establish; these make it observable. +__device__ unsigned int arc_hybrid_branch_a = 0; +__device__ unsigned int arc_hybrid_branch_b = 0; + // Splitmix64 step matching sampling_cpu.rs lines 156-161. // state = state * 0x9E3779B97F4A7C15 + 0xDEADBEEFC0DECAFE // mixed = (state ^ (state >> 30)) * 0xBF58476D1CE4E5B9 @@ -353,6 +359,390 @@ __global__ void arc_sample_kernel( } } + +// ============================================================================= +// Nucleus sampling in a FIXED number of passes (threshold bisection). +// +// `arc_sample_kernel` above pulls the keep-list one argmax at a time, so it +// costs O(vocab x kept): a full vocabulary scan PER KEPT TOKEN. With +// `top_k <= 0` -- which is what mistral.rs's default maps to -- `kept` is +// bounded only by `top_p`, and on a diffuse distribution it reaches thousands. +// MEASURED on an H200 at vocab=129280 with a flat 12928-token support: +// 663,116 us to sample ONE token, roughly 10x an entire V4 decode step. +// +// This kernel selects the same nucleus by binary-searching a probability +// THRESHOLD rather than enumerating the set. For non-negative IEEE-754 floats +// the bit pattern is monotone when compared as u32, so bisecting the integer +// key converges to one ULP in a fixed 32 passes regardless of how many tokens +// the nucleus contains. Cost becomes O(vocab) with a constant factor. +// +// It does NOT narrow the distribution. The kept set is still +// {i : p_i >= t} for the largest t whose mass still covers `top_p` -- that is +// the definition of the nucleus. Where several tokens share the boundary bit +// pattern it keeps ALL of them, so it errs toward keeping MORE than the exact +// nucleus, never fewer. `top_k` is applied as a second, independent bisection +// on the kept COUNT, and the two thresholds combine with max(), so a caller +// setting both still gets the intersection. +// ============================================================================= +template +__global__ void arc_sample_bisect_kernel( + const T* __restrict__ logits, + const uint32_t* __restrict__ freq_counts, + uint64_t* __restrict__ rng_state, + int32_t* __restrict__ token_ids, + float* __restrict__ probs_scratch, + int vocab, + SamplingParams cfg +) { + const int bid = blockIdx.x; + const int tid = threadIdx.x; + const T* __restrict__ row = logits + (int64_t)bid * vocab; + float* __restrict__ probs = probs_scratch + (int64_t)bid * vocab; + const uint32_t* fcnt = freq_counts ? (freq_counts + (int64_t)bid * vocab) : nullptr; + + extern __shared__ char smem_bisect[]; + float* s_vals = reinterpret_cast(smem_bisect); + int* s_idxs = reinterpret_cast(s_vals + BLOCK); + + const bool temp_active = (cfg.temperature > 0.0f && cfg.temperature != 1.0f); + const float inv_temp = temp_active ? (1.0f / cfg.temperature) : 1.0f; + + // --- Phase 1: penalties + temperature; keep the argmax for the fallback. + float local_max = -FLT_MAX; + int local_i = 0; + for (int i = tid; i < vocab; i += BLOCK) { + float l = to_float(row[i]); + if (fcnt) { + uint32_t c = fcnt[i]; + if (c != 0) { + l -= cfg.frequency_penalty * static_cast(c); + l -= cfg.presence_penalty; + } + } + if (temp_active) l *= inv_temp; + probs[i] = l; + if (l > local_max || (l == local_max && i < local_i)) { local_max = l; local_i = i; } + } + float gmax; int gmax_idx; + block_argmax_idx(local_max, local_i, s_vals, s_idxs, gmax, gmax_idx); + __syncthreads(); + + // --- Phase 2/3: exp-shift, sum, normalize to probabilities. + float local_sum = 0.0f; + for (int i = tid; i < vocab; i += BLOCK) { + float p = __expf(probs[i] - gmax); + probs[i] = p; + local_sum += p; + } + const float gsum = block_sum(local_sum, s_vals); + __syncthreads(); + const float inv_sum = (gsum > 0.0f) ? (1.0f / gsum) : 1.0f; + float local_pmax = 0.0f; + for (int i = tid; i < vocab; i += BLOCK) { + float p = probs[i] * inv_sum; + probs[i] = p; + if (p > local_pmax) local_pmax = p; + } + const float pmax = block_max(local_pmax, s_vals); + __syncthreads(); + + // Degenerate distribution: fall back to the mode, as the CPU reference does. + if (!(pmax > 0.0f)) { + if (tid == 0) token_ids[bid] = gmax_idx; + return; + } + + float target = cfg.top_p; + if (target > 1.0f) target = 1.0f; + if (target < 0.0f) target = 0.0f; + + // --- Phase 4a: largest threshold key whose kept mass still covers top_p. + uint32_t lo = 0u, hi = __float_as_uint(pmax); + while (lo < hi) { + const uint32_t mid = lo + ((hi - lo + 1u) >> 1); + const float threshold = __uint_as_float(mid); + float m = 0.0f; + for (int i = tid; i < vocab; i += BLOCK) { + const float p = probs[i]; + if (p >= threshold) m += p; + } + const float gm = block_sum(m, s_vals); + __syncthreads(); + if (gm >= target) lo = mid; else hi = mid - 1u; + } + uint32_t key = lo; + + // --- Phase 4b: top_k, as an independent bisection on the kept COUNT. + if (cfg.top_k > 0 && cfg.top_k < vocab) { + uint32_t klo = 0u, khi = __float_as_uint(pmax); + while (klo < khi) { + const uint32_t mid = klo + ((khi - klo) >> 1); + const float threshold = __uint_as_float(mid); + float c = 0.0f; + for (int i = tid; i < vocab; i += BLOCK) { + if (probs[i] >= threshold) c += 1.0f; + } + const float gc = block_sum(c, s_vals); + __syncthreads(); + if (gc <= static_cast(cfg.top_k)) khi = mid; else klo = mid + 1u; + } + if (klo > key) key = klo; + } + + const float threshold = __uint_as_float(key); + + // --- Phase 5: CDF walk, parallel across the block. + // Each thread sums the kept mass in its own strided subset; thread 0 + // exclusive-scans those 256 partials so every thread knows the CDF offset + // its subset begins at. Only the thread whose interval contains `u` then + // walks, and it walks vocab/BLOCK items rather than the whole vocabulary. + float part = 0.0f; + for (int i = tid; i < vocab; i += BLOCK) { + const float p = probs[i]; + if (p >= threshold) part += p; + } + s_vals[tid] = part; + __syncthreads(); + + __shared__ float s_u; + if (tid == 0) { + float acc = 0.0f; + for (int t = 0; t < BLOCK; ++t) { const float v = s_vals[t]; s_vals[t] = acc; acc += v; } + uint64_t st = rng_state[bid]; + const float u = splitmix_uniform(st); + rng_state[bid] = st; + s_u = u * acc; + // Deterministic fallback if float error leaves no owning interval. + token_ids[bid] = gmax_idx; + } + __syncthreads(); + + const float base = s_vals[tid]; + if (s_u >= base && s_u < base + part) { + float acc = base; + for (int i = tid; i < vocab; i += BLOCK) { + const float p = probs[i]; + if (p >= threshold) { + acc += p; + if (acc >= s_u) { token_ids[bid] = i; break; } + } + } + } +} + + +// ============================================================================= +// Hybrid nucleus sampler: exact enumeration while the nucleus is small, +// threshold bisection once it isn't. One kernel, no host involvement. +// +// MEASURED on an H200 at vocab=129280 (support width -> legacy us / bisect us): +// 1 -> 134 / 1994 | 8 -> 503 / 1996 | 64 -> 3132 / 1991 +// 512 -> 24287 / 1997 | 4096 -> 205446 / 1998 | 12928 -> 665694 / 1998 +// Bisection is FLAT; enumeration is linear in the kept count. Neither wins +// everywhere: enumeration is 3-15x faster on the peaked distributions real +// models produce, bisection is 333x faster on the diffuse tail. Shipping +// either alone trades one regression for another. +// +// THE SWITCH POINT IS NOT A TUNING CONSTANT. Bisection resolves one bit of the +// float32 probability key per vocabulary pass, so it costs 32 passes (plus 32 +// more when top_k forces a second search). Enumeration costs one vocabulary +// pass per kept token. Enumeration is therefore cheaper exactly while +// kept < (the number of passes bisection would spend) +// so the budget IS that pass count: ARC_ENUM_BUDGET = 32, derived from the +// width of the key being searched, the same way the 32 passes are. Spending +// the budget and then falling back costs at most budget + bisection, so the +// hybrid is never worse than ~2x the better branch, and the measured crossover +// (~64) sits within a factor of two of the derived budget -- the derivation +// and the measurement agree. +// +// Why the fallback is INSIDE the kernel: a captured CUDA graph bakes its +// kernel arguments and cannot branch on device state, so "launch enumeration, +// read a flag, maybe launch bisection" would need a device->host round trip +// per step -- exactly the host dependency the autonomous decode loop exists to +// remove. The branch is therefore taken block-uniformly on device. +// +// No cap and no narrowing: the fallback is exact, not a truncation. Exceeding +// the budget changes which ALGORITHM selects the nucleus, never which tokens +// the nucleus contains. +// ============================================================================= +template +__global__ void arc_sample_hybrid_kernel( + const T* __restrict__ logits, + const uint32_t* __restrict__ freq_counts, + uint64_t* __restrict__ rng_state, + int32_t* __restrict__ token_ids, + float* __restrict__ probs_scratch, + int32_t* __restrict__ keep_idx_scratch, + float* __restrict__ keep_p_scratch, + int vocab, + SamplingParams cfg +) { + constexpr int ARC_ENUM_BUDGET = 32; // == bisection's pass count + + const int bid = blockIdx.x; + const int tid = threadIdx.x; + const T* __restrict__ row = logits + (int64_t)bid * vocab; + float* __restrict__ probs = probs_scratch + (int64_t)bid * vocab; + int32_t* __restrict__ keep_idx = keep_idx_scratch + (int64_t)bid * vocab; + float* __restrict__ keep_p = keep_p_scratch + (int64_t)bid * vocab; + const uint32_t* fcnt = freq_counts ? (freq_counts + (int64_t)bid * vocab) : nullptr; + + extern __shared__ char smem_hy[]; + float* s_vals = reinterpret_cast(smem_hy); + int* s_idxs = reinterpret_cast(s_vals + BLOCK); + + const bool temp_active = (cfg.temperature > 0.0f && cfg.temperature != 1.0f); + const float inv_temp = temp_active ? (1.0f / cfg.temperature) : 1.0f; + + float target = cfg.top_p; + if (target > 1.0f) target = 1.0f; + if (target < 0.0f) target = 0.0f; + +#define ARC_BUILD_PROBS(GMAX_IDX_OUT) \ + do { \ + float _lmax = -FLT_MAX; int _li = 0; \ + for (int i = tid; i < vocab; i += BLOCK) { \ + float l = to_float(row[i]); \ + if (fcnt) { uint32_t c = fcnt[i]; \ + if (c != 0) { l -= cfg.frequency_penalty * (float)c; \ + l -= cfg.presence_penalty; } } \ + if (temp_active) l *= inv_temp; \ + probs[i] = l; \ + if (l > _lmax || (l == _lmax && i < _li)) { _lmax = l; _li = i; } \ + } \ + float _gv; int _gi; \ + block_argmax_idx(_lmax, _li, s_vals, s_idxs, _gv, _gi); \ + __syncthreads(); \ + (GMAX_IDX_OUT) = _gi; \ + float _lsum = 0.0f; \ + for (int i = tid; i < vocab; i += BLOCK) { \ + float p = __expf(probs[i] - _gv); probs[i] = p; _lsum += p; } \ + float _gsum = block_sum(_lsum, s_vals); \ + __syncthreads(); \ + float _inv = (_gsum > 0.0f) ? (1.0f / _gsum) : 1.0f; \ + for (int i = tid; i < vocab; i += BLOCK) probs[i] *= _inv; \ + __syncthreads(); \ + } while (0) + + int gmax_idx = 0; + ARC_BUILD_PROBS(gmax_idx); + + // ---- Branch A: bounded exact enumeration (the legacy algorithm). + __shared__ int s_kept; + __shared__ float s_cum; + __shared__ int s_done; // 1 = nucleus fully determined here + if (tid == 0) { s_kept = 0; s_cum = 0.0f; s_done = 0; } + __syncthreads(); + + while (true) { + if (s_cum >= target) { if (tid == 0) s_done = 1; __syncthreads(); break; } + if (cfg.top_k > 0 && s_kept >= cfg.top_k) { if (tid == 0) s_done = 1; __syncthreads(); break; } + if (s_kept >= ARC_ENUM_BUDGET) { __syncthreads(); break; } // spend budget, fall back + + float lv = -FLT_MAX; int li = 0; + for (int i = tid; i < vocab; i += BLOCK) { + float v = probs[i]; + if (v > lv || (v == lv && i < li)) { lv = v; li = i; } + } + float gv; int gi; + block_argmax_idx(lv, li, s_vals, s_idxs, gv, gi); + __syncthreads(); + if (gv <= 0.0f) { if (tid == 0) s_done = 1; __syncthreads(); break; } + if (tid == 0) { + keep_idx[s_kept] = gi; + keep_p[s_kept] = gv; + s_kept += 1; + s_cum += gv; + probs[gi] = -FLT_MAX; // tombstone + } + __syncthreads(); + } + + if (s_done) { + if (tid == 0) atomicAdd(&arc_hybrid_branch_a, 1u); + if (s_kept <= 0) { if (tid == 0) token_ids[bid] = gmax_idx; return; } + if (tid == 0) { + uint64_t st = rng_state[bid]; + float u = splitmix_uniform(st); + rng_state[bid] = st; + float t = u * s_cum; + float acc = 0.0f; + int sel = keep_idx[s_kept - 1]; + for (int k = 0; k < s_kept; ++k) { + acc += keep_p[k]; + if (acc >= t) { sel = keep_idx[k]; break; } + } + token_ids[bid] = sel; + } + return; + } + + // ---- Branch B: budget spent, nucleus is diffuse. probs[] was tombstoned + // by the enumeration above, so rebuild it, then bisect. + if (tid == 0) atomicAdd(&arc_hybrid_branch_b, 1u); + ARC_BUILD_PROBS(gmax_idx); + + float lpmax = 0.0f; + for (int i = tid; i < vocab; i += BLOCK) { float p = probs[i]; if (p > lpmax) lpmax = p; } + const float pmax = block_max(lpmax, s_vals); + __syncthreads(); + if (!(pmax > 0.0f)) { if (tid == 0) token_ids[bid] = gmax_idx; return; } + + uint32_t lo = 0u, hi = __float_as_uint(pmax); + while (lo < hi) { + const uint32_t mid = lo + ((hi - lo + 1u) >> 1); + const float threshold_m = __uint_as_float(mid); + float m = 0.0f; + for (int i = tid; i < vocab; i += BLOCK) { float p = probs[i]; if (p >= threshold_m) m += p; } + const float gm = block_sum(m, s_vals); + __syncthreads(); + if (gm >= target) lo = mid; else hi = mid - 1u; + } + uint32_t key = lo; + + if (cfg.top_k > 0 && cfg.top_k < vocab) { + uint32_t klo = 0u, khi = __float_as_uint(pmax); + while (klo < khi) { + const uint32_t mid = klo + ((khi - klo) >> 1); + const float threshold_m = __uint_as_float(mid); + float c = 0.0f; + for (int i = tid; i < vocab; i += BLOCK) { if (probs[i] >= threshold_m) c += 1.0f; } + const float gc = block_sum(c, s_vals); + __syncthreads(); + if (gc <= (float)cfg.top_k) khi = mid; else klo = mid + 1u; + } + if (klo > key) key = klo; + } + + const float threshold = __uint_as_float(key); + float part = 0.0f; + for (int i = tid; i < vocab; i += BLOCK) { float p = probs[i]; if (p >= threshold) part += p; } + s_vals[tid] = part; + __syncthreads(); + + __shared__ float s_u; + if (tid == 0) { + float acc = 0.0f; + for (int t = 0; t < BLOCK; ++t) { const float v = s_vals[t]; s_vals[t] = acc; acc += v; } + uint64_t st = rng_state[bid]; + const float u = splitmix_uniform(st); + rng_state[bid] = st; + s_u = u * acc; + token_ids[bid] = gmax_idx; + } + __syncthreads(); + + const float base = s_vals[tid]; + if (s_u >= base && s_u < base + part) { + float acc = base; + for (int i = tid; i < vocab; i += BLOCK) { + const float p = probs[i]; + if (p >= threshold) { acc += p; if (acc >= s_u) { token_ids[bid] = i; break; } } + } + } +#undef ARC_BUILD_PROBS +} + } // namespace arc_sampler // ============================================================================= @@ -499,4 +889,70 @@ void arc_launch_sampler_f16( } } + +// Fixed-pass nucleus sampler. Same signature as `arc_launch_sampler_*` so it +// is a drop-in; `keep_idx_scratch` / `keep_p_scratch` are unused by it. +#define ARC_DEFINE_BISECT_LAUNCHER(SUFFIX, CTYPE) \ +void arc_launch_sampler_##SUFFIX##_bisect( \ + const void* logits, const uint32_t* freq_counts, uint64_t* rng_state, \ + int32_t* token_ids, float* probs_scratch, int32_t* keep_idx_scratch, \ + float* keep_p_scratch, int vocab, int batch, \ + arc_sampler::SamplingParams cfg, cudaStream_t stream) { \ + (void)keep_idx_scratch; (void)keep_p_scratch; \ + const int B = 256; \ + size_t smem = B * (sizeof(float) + sizeof(int)); \ + if (cfg.greedy) { \ + arc_sampler::arc_greedy_kernel<<>>( \ + reinterpret_cast(logits), freq_counts, token_ids, vocab, \ + cfg); \ + return; \ + } \ + arc_sampler::arc_sample_bisect_kernel<<>>( \ + reinterpret_cast(logits), freq_counts, rng_state, \ + token_ids, probs_scratch, vocab, cfg); \ +} + +ARC_DEFINE_BISECT_LAUNCHER(f32, float) +ARC_DEFINE_BISECT_LAUNCHER(bf16, __nv_bfloat16) +ARC_DEFINE_BISECT_LAUNCHER(f16, __half) + + +void arc_hybrid_branch_counts(unsigned int* a_out, unsigned int* b_out) { + cudaMemcpyFromSymbol(a_out, arc_sampler::arc_hybrid_branch_a, + sizeof(unsigned int), 0, cudaMemcpyDeviceToHost); + cudaMemcpyFromSymbol(b_out, arc_sampler::arc_hybrid_branch_b, + sizeof(unsigned int), 0, cudaMemcpyDeviceToHost); +} + +void arc_hybrid_branch_reset(void) { + unsigned int z = 0; + cudaMemcpyToSymbol(arc_sampler::arc_hybrid_branch_a, &z, sizeof(z), 0, + cudaMemcpyHostToDevice); + cudaMemcpyToSymbol(arc_sampler::arc_hybrid_branch_b, &z, sizeof(z), 0, + cudaMemcpyHostToDevice); +} + +#define ARC_DEFINE_HYBRID_LAUNCHER(SUFFIX, CTYPE) \ +void arc_launch_sampler_##SUFFIX##_hybrid( \ + const void* logits, const uint32_t* freq_counts, uint64_t* rng_state, \ + int32_t* token_ids, float* probs_scratch, int32_t* keep_idx_scratch, \ + float* keep_p_scratch, int vocab, int batch, \ + arc_sampler::SamplingParams cfg, cudaStream_t stream) { \ + const int B = 256; \ + size_t smem = B * (sizeof(float) + sizeof(int)); \ + if (cfg.greedy) { \ + arc_sampler::arc_greedy_kernel<<>>( \ + reinterpret_cast(logits), freq_counts, token_ids, vocab, \ + cfg); \ + return; \ + } \ + arc_sampler::arc_sample_hybrid_kernel<<>>( \ + reinterpret_cast(logits), freq_counts, rng_state, \ + token_ids, probs_scratch, keep_idx_scratch, keep_p_scratch, vocab, cfg); \ +} + +ARC_DEFINE_HYBRID_LAUNCHER(f32, float) +ARC_DEFINE_HYBRID_LAUNCHER(bf16, __nv_bfloat16) +ARC_DEFINE_HYBRID_LAUNCHER(f16, __half) + } // extern "C" diff --git a/arc-cuda-graph/src/ffi.rs b/arc-cuda-graph/src/ffi.rs index ae78c6f76..61a467ecf 100644 --- a/arc-cuda-graph/src/ffi.rs +++ b/arc-cuda-graph/src/ffi.rs @@ -14,6 +14,21 @@ pub type CUgraphConditionalHandle = u64; #[cfg(feature = "cuda")] pub type CUmemoryPool = *mut std::ffi::c_void; #[cfg(feature = "cuda")] +pub type CUcontext = *mut std::ffi::c_void; +/// `cuda.h:1944` (CUDA 13.1) — +/// `#define CU_GRAPH_COND_ASSIGN_DEFAULT 0x1` "Default value is applied when +/// graph is launched." +/// +/// `cuda.h:21919` documents `defaultLaunchValue` as "Applied at the beginning +/// of each graph execution **if CU_GRAPH_COND_ASSIGN_DEFAULT is set in +/// flags**". Creating the handle with `flags = 0` therefore leaves the +/// condition at 0 on entry and a WHILE body executes **zero** times — the +/// graph launches, returns success, and generates nothing. Measured on +/// `arc-v4-stack` (H200, CUDA 13.1): `flags=0` → body ran 0 times; +/// `flags=CU_GRAPH_COND_ASSIGN_DEFAULT` → body ran exactly N times. +#[cfg(feature = "cuda")] +pub const CU_GRAPH_COND_ASSIGN_DEFAULT: u32 = 0x1; +#[cfg(feature = "cuda")] pub type CUdevice = i32; #[cfg(feature = "cuda")] @@ -70,6 +85,14 @@ pub enum CudaGraphNodeType { Conditional = 13, } +/// `CUDA_CONDITIONAL_NODE_PARAMS`, `cuda.h:1958` (CUDA 13.1). +/// +/// `phGraph_out` is an **OUT** field: "CUDA-owned array populated with +/// conditional node child graphs during creation of the node." Callers leave +/// it null and read it back after `cuGraphAddNode`; assigning a caller-owned +/// pointer to it accomplishes nothing because the driver overwrites it. +/// +/// `ctx` (the 5th field) was missing entirely from this struct. #[cfg(feature = "cuda")] #[repr(C)] pub struct CudaConditionalNodeParams { @@ -77,6 +100,7 @@ pub struct CudaConditionalNodeParams { pub cond_type: CUgraphConditionalNodeType, pub size: u32, pub body_graph_out: *mut CUgraph, + pub ctx: CUcontext, } #[cfg(feature = "cuda")] @@ -162,6 +186,7 @@ extern "C" { pub fn cuGraphLaunch(exec: CUgraphExec, stream: CUstream) -> u32; pub fn cuGraphExecDestroy(exec: CUgraphExec) -> u32; pub fn cuGraphDestroy(graph: CUgraph) -> u32; + pub fn cuGraphGetNodes(graph: CUgraph, nodes: *mut CUgraphNode, num_nodes: *mut usize) -> u32; // ========== Memory pools ========== pub fn cuMemPoolCreate(pool: *mut CUmemoryPool, props: *const CUmemPoolProps) -> u32; @@ -181,20 +206,53 @@ extern "C" { pub fn cudaStreamSynchronize(stream: CUstream) -> u32; // ========== Conditional nodes (CUDA 12.4+) ========== - pub fn cudaGraphConditionalHandleCreate( + /// `cuda.h:21932` (CUDA 13.1) — the driver form takes a `CUcontext`, + /// which the conditional node params must match ("Context on which to run + /// the node. Must match context used to create the handle and all body + /// nodes", `cuda.h:1986`). + pub fn cuGraphConditionalHandleCreate( handle: *mut CUgraphConditionalHandle, graph: CUgraph, - default_value: u32, + ctx: CUcontext, + default_launch_value: u32, flags: u32, ) -> u32; - pub fn cudaGraphSetConditional(handle: CUgraphConditionalHandle, value: u32) -> u32; - pub fn cudaGraphAddNode( - graph: CUgraph, + pub fn cuCtxGetCurrent(ctx: *mut CUcontext) -> u32; + /// `cuda.h:21829` (CUDA 13.1): + /// ```text + /// CUresult cuGraphAddNode(CUgraphNode *phGraphNode, CUgraph hGraph, + /// const CUgraphNode *dependencies, + /// const CUgraphEdgeData *dependencyData, + /// size_t numDependencies, + /// CUgraphNodeParams *nodeParams); + /// ``` + /// The OUT node pointer is the **first** argument and there are **six** of + /// them. This was previously declared as + /// `cudaGraphAddNode(graph, node_out, deps, num_deps, params)` — five + /// arguments with the first two transposed — so the driver received the + /// `CUgraph` handle in the slot where it expected `CUgraphNode *` and + /// wrote the new node handle through that value, while `numDependencies` + /// received a pointer and `nodeParams` an uninitialised register. + pub fn cuGraphAddNode( node_out: *mut CUgraphNode, + graph: CUgraph, dependencies: *const CUgraphNode, + dependency_data: *const std::ffi::c_void, num_dependencies: usize, params: *mut CudaGraphNodeParams, ) -> u32; + /// `cuda.h:15665` (CUDA 13.1). `cuda.h:1976` names this as a supported way + /// to populate a conditional node's body graph. Ordinary + /// `cuStreamBeginCapture_v2` always creates a **new** graph, so a body + /// captured that way cannot be attached to the conditional node. + pub fn cuStreamBeginCaptureToGraph( + stream: CUstream, + graph: CUgraph, + dependencies: *const CUgraphNode, + dependency_data: *const std::ffi::c_void, + num_dependencies: usize, + mode: CUstreamCaptureMode, + ) -> u32; // ========== Pinned host memory ========== pub fn cudaHostAlloc(ptr: *mut *mut std::ffi::c_void, size: usize, flags: u32) -> u32; diff --git a/arc-cuda-graph/src/sampling_cuda.rs b/arc-cuda-graph/src/sampling_cuda.rs index 8d6feb95a..67d52a939 100644 --- a/arc-cuda-graph/src/sampling_cuda.rs +++ b/arc-cuda-graph/src/sampling_cuda.rs @@ -209,6 +209,59 @@ extern "C" { cfg: SamplingParams, stream: CUstream, ); + + // Hybrid: exact enumeration while the nucleus is small, threshold + // bisection past a fixed budget, chosen block-uniformly ON DEVICE so a + // captured graph never has to ask the host which branch to take. + // + // MEASURED, H200, vocab=129280, host+GPU verified exclusive before and + // after (support width -> legacy / hybrid us): + // 1 -> 134.9/166.6 | 8 -> 503.0/624.2 | 64 -> 3127.8/4178.5 + // 512 -> 24320.7/4185.0 | 4096 -> 205125.2/4187.3 + // 12928 -> 664200.0/4188.5 + // Costs +24% on the peaked distributions real models produce (32 us, i.e. + // 0.05% of a 66.68 ms V4 decode step) and removes a 664 ms cliff -- 158x + // at the diffuse tail, where the enumerating sampler is ~10x an entire + // decode step for a single token. + pub fn arc_launch_sampler_f32_hybrid( + logits: *const f32, + freq_counts: *const u32, + rng_state: *mut u64, + token_ids: *mut i32, + probs_scratch: *mut f32, + keep_idx_scratch: *mut i32, + keep_p_scratch: *mut f32, + vocab: i32, + batch: i32, + cfg: SamplingParams, + stream: CUstream, + ); + pub fn arc_launch_sampler_bf16_hybrid( + logits: *const std::ffi::c_void, + freq_counts: *const u32, + rng_state: *mut u64, + token_ids: *mut i32, + probs_scratch: *mut f32, + keep_idx_scratch: *mut i32, + keep_p_scratch: *mut f32, + vocab: i32, + batch: i32, + cfg: SamplingParams, + stream: CUstream, + ); + pub fn arc_launch_sampler_f16_hybrid( + logits: *const std::ffi::c_void, + freq_counts: *const u32, + rng_state: *mut u64, + token_ids: *mut i32, + probs_scratch: *mut f32, + keep_idx_scratch: *mut i32, + keep_p_scratch: *mut f32, + vocab: i32, + batch: i32, + cfg: SamplingParams, + stream: CUstream, + ); } /// Raw device pointer for a CUDA tensor, byte-offset aware. @@ -311,6 +364,15 @@ impl CudaSampler { }) } + /// Bind the sampler to `stream`. + /// + /// `new()` picks up the device-default stream. Anything that is going to + /// be **captured** into a CUDA graph must launch on the capture stream, or + /// the sampler's kernels are simply not recorded into the graph. + pub fn set_stream(&mut self, stream: CUstream) { + self.stream = stream; + } + /// Run the sampler. Logits: [batch, vocab] of the configured dtype. /// `freq_counts`: optional [batch, vocab] u32 token-count tensor (one per /// batch row); pass `None` when no penalties are needed. @@ -370,7 +432,7 @@ impl CudaSampler { unsafe { match self.dtype { - DType::F32 => arc_launch_sampler_f32( + DType::F32 => arc_launch_sampler_f32_hybrid( logits_ptr as *const f32, freq_ptr, rng_ptr, @@ -383,7 +445,7 @@ impl CudaSampler { params, self.stream, ), - DType::BF16 => arc_launch_sampler_bf16( + DType::BF16 => arc_launch_sampler_bf16_hybrid( logits_ptr, freq_ptr, rng_ptr, @@ -396,7 +458,7 @@ impl CudaSampler { params, self.stream, ), - DType::F16 => arc_launch_sampler_f16( + DType::F16 => arc_launch_sampler_f16_hybrid( logits_ptr, freq_ptr, rng_ptr, diff --git a/arc-cuda-graph/src/weights.rs b/arc-cuda-graph/src/weights.rs index f8891c311..bdf77b4b0 100644 --- a/arc-cuda-graph/src/weights.rs +++ b/arc-cuda-graph/src/weights.rs @@ -610,6 +610,16 @@ pub fn tensor_device_ptr(tensor: &Tensor) -> candle_core::Result { // fed with (radix top-k `seq_lens`, `CudaSampler` token ids and // keep-list scratch). Candle's `Tensor::from_vec(Vec, …)` // produces it, so leaving it out disabled those paths wholesale. + // + // Concretely, from #130: `CudaSampler::sample` REQUIRES + // `token_ids` to be I32 (`sampling_cuda.rs:347`) and allocates + // an I32 `keep_idx_scratch` (`:298`). Both resolve through this + // function, so before I32 was listed both fell to the catch-all + // `bail!` and `CudaSampler::sample()` returned `unsupported + // dtype I32` on EVERY call — the replay-safe top-k/top-p + // sampler could never run on a GPU. Its test suite is a CPU + // simulator (`gpu_algorithm_simulate`), which by construction + // cannot observe that. DType::I32 => { let s = cuda_storage.as_cuda_slice::()?; let (p, _) = s.device_ptr(s.stream()); diff --git a/mistralrs-core/src/pipeline/normal.rs b/mistralrs-core/src/pipeline/normal.rs index b93ae48df..c4bff0a99 100644 --- a/mistralrs-core/src/pipeline/normal.rs +++ b/mistralrs-core/src/pipeline/normal.rs @@ -3021,6 +3021,16 @@ impl Pipeline for NormalPipeline { top_p, frequency_penalty, presence_penalty, + // mistral.rs uses <=0 to mean "disabled"; the fused sampler + // spells that -1 (`sampling_cpu.rs:16`). + top_k: { + let tk = first_sampler.top_k(); + if tk <= 0 { + -1 + } else { + tk as i32 + } + }, greedy, }; diff --git a/mistralrs-core/tests/capability_reachability.rs b/mistralrs-core/tests/capability_reachability.rs index bffe418f7..253c6ab9a 100644 --- a/mistralrs-core/tests/capability_reachability.rs +++ b/mistralrs-core/tests/capability_reachability.rs @@ -196,17 +196,20 @@ static REGISTRY: &[Capability] = &[ accepted_in: &["mistralrs-core/src/sampler.rs"], honoured_in: &["arc-cuda-graph/src/autonomous.rs"], }, - // `AutonomousDecodeConfig` (autonomous.rs:70-82) carries temperature, - // top_p, both penalties and `greedy` — and neither top_k nor min_p. A - // request setting them would have them silently dropped the moment this - // path is reached. It is not reached today (`mark_unreachable( - // "cuda_graph.autonomous_decode", ...)`), which is the only reason this - // is not already a user-visible defect. - status: Status::Tracked { - reason: "AutonomousDecodeConfig has no top_k field at all; the path is itself \ - unreachable today, so this is latent rather than live. Fix before \ - GPU-autonomous decode is switched on.", - }, + // Was Tracked with the reason "AutonomousDecodeConfig has no top_k + // field at all". #130 gave it one: `AutonomousDecodeConfig::top_k` + // (autonomous.rs:82), plumbed into `sampling_cpu::SamplingConfig` at + // autonomous.rs:643 and consumed by `self.sampler.sample(..)`. So the + // promise is now honoured where it is made, and the status follows the + // code. + // + // The gate flagged this itself, on the rebase, with "PROMOTE ME" — + // the same way the greedy/logit-bias entry above was promoted. That is + // the mechanism working, not a test to be silenced. + // + // `min_p` is a separate entry below and stays Tracked: it still has no + // field on `AutonomousDecodeConfig`. + status: Status::Live, }, Capability { name: "sampling: min_p survives onto the GPU-autonomous decode path", diff --git a/mistralrs-core/tests/doc_citations.rs b/mistralrs-core/tests/doc_citations.rs index e4204e8cc..6d219d675 100644 --- a/mistralrs-core/tests/doc_citations.rs +++ b/mistralrs-core/tests/doc_citations.rs @@ -351,6 +351,30 @@ static BASELINE: &[Waiver] = &[ kind: Kind::Unresolved, why: "candle is an out-of-tree path dependency", }, + // The CUDA driver header. #130's FFI declarations cite it by line for the + // struct layouts and flag values they mirror (`CUDA_CONDITIONAL_NODE_PARAMS`, + // `CU_GRAPH_COND_ASSIGN_DEFAULT`, the conditional-node entry points). It + // ships with the CUDA Toolkit and is not vendored here, so no citation into + // it can ever resolve — and its line numbers are toolkit-version-specific + // besides (these are CUDA 13.1), which is why each citation names the + // version in its surrounding comment. + // + // Waived by bare filename rather than one row per line number: the rows + // would otherwise have to be rewritten on every toolkit bump, which is + // churn the gate cannot check either way. The layouts themselves ARE + // checked, against `cudarc`'s bindgen output — see the ABI note on #130. + Waiver { + doc: "arc-cuda-graph/src/ffi.rs", + cite: "cuda.h", + kind: Kind::Unresolved, + why: "CUDA Toolkit header, not vendored here", + }, + Waiver { + doc: "arc-cuda-graph/src/autonomous.rs", + cite: "cuda.h", + kind: Kind::Unresolved, + why: "CUDA Toolkit header, not vendored here", + }, Waiver { doc: "memory/mission/wave29-BD-rung-decision.md", cite: "STATUS.md",