Skip to main content

ringkernel_cuda/
lib.rs

1//! CUDA Backend for RingKernel
2//!
3//! This crate provides NVIDIA CUDA GPU support for RingKernel using cudarc.
4//!
5//! # Features
6//!
7//! - Persistent kernel execution (cooperative groups)
8//! - Lock-free message queues in GPU global memory
9//! - PTX compilation via NVRTC
10//! - Multi-GPU support
11//!
12//! # Requirements
13//!
14//! - NVIDIA GPU with Compute Capability 6.0+ (Pascal and newer) for the core
15//!   persistent-actor/cooperative-groups path. Some features have higher
16//!   floors (e.g. Thread Block Clusters/DSMEM/TMA/Green Contexts require
17//!   Hopper, CC 9.0+) — see the top-level README's "Feature to minimum
18//!   compute capability" table for the full breakdown.
19//! - CUDA Toolkit 11.0+
20//! - Native Linux (persistent kernels) or WSL2 (event-driven fallback)
21//!
22//! # Example
23//!
24//! ```ignore
25//! use ringkernel_cuda::CudaRuntime;
26//! use ringkernel_core::runtime::RingKernelRuntime;
27//!
28//! #[tokio::main]
29//! async fn main() -> Result<(), Box<dyn std::error::Error>> {
30//!     let runtime = CudaRuntime::new().await?;
31//!     let kernel = runtime.launch("vector_add", Default::default()).await?;
32//!     kernel.activate().await?;
33//!     Ok(())
34//! }
35//! ```
36
37#![warn(missing_docs)]
38#![warn(clippy::unwrap_used)]
39
40#[cfg(feature = "ptx-cache")]
41pub mod compile;
42#[cfg(feature = "cooperative")]
43pub mod cooperative;
44#[cfg(feature = "cuda")]
45mod device;
46#[cfg(feature = "cuda")]
47pub mod driver_api;
48#[cfg(feature = "cuda")]
49pub mod hopper;
50#[cfg(feature = "cuda")]
51pub mod k2k_gpu;
52#[cfg(feature = "cuda")]
53mod kernel;
54#[cfg(feature = "cuda")]
55pub mod launch_config;
56#[cfg(feature = "cuda")]
57mod memory;
58#[cfg(feature = "cuda")]
59pub mod memory_pool;
60#[cfg(feature = "cuda")]
61pub mod multi_gpu;
62#[cfg(feature = "cuda")]
63pub mod persistent;
64#[cfg(feature = "cuda")]
65pub mod phases;
66#[cfg(feature = "profiling")]
67pub mod profiling;
68#[cfg(feature = "cuda")]
69pub mod reduction;
70#[cfg(feature = "cuda")]
71mod runtime;
72#[cfg(feature = "cuda")]
73mod stencil;
74#[cfg(feature = "cuda")]
75pub mod stream;
76
77#[cfg(feature = "cuda")]
78pub use device::CudaDevice;
79#[cfg(feature = "cuda")]
80pub use kernel::CudaKernel;
81#[cfg(feature = "cuda")]
82pub use memory::{CudaBuffer, CudaControlBlock, CudaMemoryPool, CudaMessageQueue};
83#[cfg(feature = "cuda")]
84pub use persistent::CudaMappedBuffer;
85#[cfg(feature = "cuda")]
86pub use phases::{
87    InterPhaseReduction, KernelPhase, MultiPhaseConfig, MultiPhaseExecutor, PhaseExecutionStats,
88    SyncMode,
89};
90#[cfg(feature = "cuda")]
91pub use reduction::{
92    generate_block_reduce_code, generate_grid_reduce_code, generate_reduce_and_broadcast_code,
93    CacheKey, CacheStats, CachedReductionBuffer, ReductionBuffer, ReductionBufferBuilder,
94    ReductionBufferCache,
95};
96#[cfg(feature = "cuda")]
97pub use runtime::CudaRuntime;
98#[cfg(feature = "cuda")]
99pub use stencil::{CompiledStencilKernel, LaunchConfig, StencilKernelLoader};
100
101// Profiling re-exports
102#[cfg(feature = "profiling")]
103pub use profiling::{
104    CudaEvent, CudaEventFlags, CudaMemoryKind, CudaMemoryTracker, CudaNvtxProfiler,
105    GpuChromeTraceBuilder, GpuEventArgs, GpuTimer, GpuTimerPool, GpuTraceEvent, KernelMetrics,
106    ProfilingSession, TrackedAllocation, TransferDirection, TransferMetrics,
107};
108
109// PTX cache re-exports
110#[cfg(feature = "ptx-cache")]
111pub use compile::{PtxCache, PtxCacheError, PtxCacheResult, PtxCacheStats, CACHE_VERSION};
112
113// GPU memory pool re-exports
114#[cfg(feature = "cuda")]
115pub use memory_pool::{
116    GpuBucketStats, GpuPoolConfig, GpuPoolDiagnostics, GpuSizeClass, GpuStratifiedPool,
117};
118
119// Stream manager re-exports
120#[cfg(feature = "cuda")]
121pub use stream::{
122    OverlapMetrics, StreamConfig, StreamConfigBuilder, StreamError, StreamId, StreamManager,
123    StreamPool, StreamPoolStats, StreamResult,
124};
125
126/// Re-export memory module for advanced usage.
127#[cfg(feature = "cuda")]
128pub mod memory_exports {
129    pub use super::memory::{CudaBuffer, CudaControlBlock, CudaMemoryPool, CudaMessageQueue};
130}
131
132// Placeholder implementations when CUDA is not available
133#[cfg(not(feature = "cuda"))]
134mod stub {
135    ringkernel_core::unavailable_backend!(
136        CudaRuntime,
137        ringkernel_core::runtime::Backend::Cuda,
138        "CUDA"
139    );
140}
141
142#[cfg(not(feature = "cuda"))]
143pub use stub::CudaRuntime;
144
145/// Check if CUDA is available at runtime.
146///
147/// This function returns false if:
148/// - CUDA feature is not enabled
149/// - CUDA libraries are not installed on the system
150/// - No CUDA devices are present
151///
152/// It safely catches panics from cudarc when CUDA is not installed.
153pub fn is_cuda_available() -> bool {
154    #[cfg(feature = "cuda")]
155    {
156        // cudarc panics if CUDA libraries are not found, so we catch that
157        std::panic::catch_unwind(|| {
158            cudarc::driver::CudaContext::device_count()
159                .map(|c| c > 0)
160                .unwrap_or(false)
161        })
162        .unwrap_or(false)
163    }
164    #[cfg(not(feature = "cuda"))]
165    {
166        false
167    }
168}
169
170/// Get CUDA device count.
171///
172/// Returns 0 if CUDA is not available or libraries are not installed.
173pub fn cuda_device_count() -> usize {
174    #[cfg(feature = "cuda")]
175    {
176        // cudarc panics if CUDA libraries are not found, so we catch that
177        std::panic::catch_unwind(|| {
178            cudarc::driver::CudaContext::device_count().unwrap_or(0) as usize
179        })
180        .unwrap_or(0)
181    }
182    #[cfg(not(feature = "cuda"))]
183    {
184        0
185    }
186}
187
188/// Compile CUDA C source code to PTX using NVRTC.
189///
190/// This wraps `cudarc::nvrtc::compile_ptx` to provide PTX compilation
191/// without requiring downstream crates to depend on cudarc directly.
192///
193/// # Arguments
194///
195/// * `cuda_source` - CUDA C source code string
196///
197/// # Returns
198///
199/// PTX assembly as a string, or an error if compilation fails.
200///
201/// # Example
202///
203/// ```ignore
204/// use ringkernel_cuda::compile_ptx;
205///
206/// let cuda_source = r#"
207///     extern "C" __global__ void add(float* a, float* b, float* c, int n) {
208///         int i = blockIdx.x * blockDim.x + threadIdx.x;
209///         if (i < n) c[i] = a[i] + b[i];
210///     }
211/// "#;
212///
213/// let ptx = compile_ptx(cuda_source)?;
214/// ```
215#[cfg(feature = "cuda")]
216pub fn compile_ptx(cuda_source: &str) -> ringkernel_core::error::Result<String> {
217    use ringkernel_core::error::RingKernelError;
218
219    let ptx = cudarc::nvrtc::compile_ptx(cuda_source).map_err(|e| {
220        RingKernelError::CompilationError(format!("NVRTC compilation failed: {}", e))
221    })?;
222
223    Ok(ptx.to_src().to_string())
224}
225
226/// Stub compile_ptx when CUDA is not available.
227#[cfg(not(feature = "cuda"))]
228pub fn compile_ptx(_cuda_source: &str) -> ringkernel_core::error::Result<String> {
229    Err(ringkernel_core::error::RingKernelError::BackendUnavailable(
230        "CUDA feature not enabled".to_string(),
231    ))
232}
233
234/// PTX kernel source template for persistent ring kernel, generic over the
235/// target compute capability (major, minor).
236///
237/// This is a minimal kernel that immediately marks itself as terminated.
238///
239/// PTX `.target` directives are only forward-compatible (PTX built for
240/// `sm_X` runs on `sm_X` and newer, never older) — so this must be
241/// generated per-device rather than hardcoded to a single value. Maps the
242/// device's actual compute capability to the corresponding PTX ISA
243/// version, matching the minimum ISA version that introduced each `sm_XX`
244/// target (per NVIDIA's PTX ISA documentation).
245pub fn ring_kernel_ptx_template_for(major: u32, minor: u32) -> String {
246    let ptx_version = match (major, minor) {
247        // Pascal
248        (6, 0) | (6, 1) | (6, 2) => "5.0",
249        // Volta / Turing
250        (7, 0) | (7, 2) => "6.0",
251        (7, 5) => "6.3",
252        // Ampere
253        (8, 0) => "7.0",
254        (8, 6) | (8, 7) => "7.1",
255        (8, 9) => "8.0",
256        // Hopper
257        (9, 0) => "8.0",
258        // Blackwell and newer — fall through to the newest known ISA
259        (major, _) if major >= 10 => "8.5",
260        // Below Pascal (Maxwell/Kepler) or anything unrecognized: fall back
261        // to the oldest ISA version this template's instructions need.
262        _ => "5.0",
263    };
264    format!(
265        r#"
266.version {ptx_version}
267.target sm_{major}{minor}
268.address_size 64
269
270.visible .entry ring_kernel_main(
271    .param .u64 control_block_ptr,
272    .param .u64 input_queue_ptr,
273    .param .u64 output_queue_ptr,
274    .param .u64 shared_state_ptr
275) {{
276    .reg .u64 %cb_ptr;
277    .reg .u32 %one;
278
279    // Load control block pointer
280    ld.param.u64 %cb_ptr, [control_block_ptr];
281
282    // Mark as terminated immediately (offset 8)
283    mov.u32 %one, 1;
284    st.global.u32 [%cb_ptr + 8], %one;
285
286    ret;
287}}
288"#
289    )
290}
291
292/// PTX kernel source template for persistent ring kernel — **deprecated**
293/// fixed-`sm_75` fallback, kept only for API compatibility with existing
294/// callers that don't have a device handle available.
295///
296/// PTX built for `sm_75` will fail to load (`CUDA_ERROR_INVALID_PTX`) on any
297/// GPU with a compute capability below 7.5 (e.g. Pascal, sm_61) — PTX
298/// forward-compatibility only extends to equal-or-newer architectures.
299/// Prefer [`ring_kernel_ptx_template_for`] with the actual target device's
300/// compute capability (`CudaDevice::compute_capability`, not part of this
301/// crate's public API surface).
302#[deprecated(
303    since = "1.1.1",
304    note = "hardcodes sm_75, breaking on Pascal and older GPUs — use ring_kernel_ptx_template_for(major, minor) with the real device's compute capability instead"
305)]
306pub const RING_KERNEL_PTX_TEMPLATE: &str = r#"
307.version 8.0
308.target sm_75
309.address_size 64
310
311.visible .entry ring_kernel_main(
312    .param .u64 control_block_ptr,
313    .param .u64 input_queue_ptr,
314    .param .u64 output_queue_ptr,
315    .param .u64 shared_state_ptr
316) {
317    .reg .u64 %cb_ptr;
318    .reg .u32 %one;
319
320    // Load control block pointer
321    ld.param.u64 %cb_ptr, [control_block_ptr];
322
323    // Mark as terminated immediately (offset 8)
324    mov.u32 %one, 1;
325    st.global.u32 [%cb_ptr + 8], %one;
326
327    ret;
328}
329"#;