use super::*;
use crate::types::{BufferResizeCost, DeviceType};
use anyhow::{Context as _, Result};
use cudarc::driver::{
CudaContext, CudaFunction, CudaModule, CudaSlice, CudaStream, DevicePtr, DeviceRepr, LaunchConfig, PushKernelArg,
};
use cudarc::nvrtc::Ptx;
use std::collections::HashMap;
use std::path::PathBuf;
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::{Arc, Mutex};
#[repr(C)]
#[derive(Clone, Copy)]
struct CudaBufferArg {
data: u64,
count: usize,
}
unsafe impl DeviceRepr for CudaBufferArg {}
pub(crate) struct CudaBackend {
adapter_info: Vec<AdapterInfo>,
devices: HashMap<DeviceHandle, CudaDevice>,
contexts: HashMap<ContextHandle, Arc<CudaSubmitContext>>,
buffers: HashMap<BufferHandle, CudaBuffer>,
buffer_slots: HashMap<u32, BufferHandle>,
shaders: HashMap<ShaderHandle, CudaShader>,
compute_pipelines: HashMap<ComputePipelineHandle, CudaComputePipeline>,
next_device: DeviceHandle,
next_context: ContextHandle,
next_buffer: BufferHandle,
next_slot: u32,
next_shader: ShaderHandle,
next_compute_pipeline: ComputePipelineHandle,
}
struct CudaDevice {
ctx: Arc<CudaContext>,
stream: Arc<CudaStream>,
next_timeline: Arc<AtomicU64>,
retired: Arc<AtomicU64>,
}
struct CudaSubmitContext {
device: DeviceHandle,
completed: AtomicU64,
signal_queue: crate::signal::SignalQueue,
}
struct CudaProgress {
context: Arc<CudaSubmitContext>,
}
impl ContextGpuProgress for CudaProgress {
fn gpu_progress(&self) -> crate::timeline::TimelineValue {
self.context.completed.load(Ordering::Acquire)
}
}
struct CudaDestroyContext;
impl ContextDestroyHandle for CudaDestroyContext {
fn wait(&self) -> Result<()> {
Ok(())
}
fn finish(self: Box<Self>) -> Result<()> {
Ok(())
}
}
struct CudaBuffer {
device: DeviceHandle,
memory: Arc<Mutex<CudaSlice<u8>>>,
offset: u64,
size: u64,
capacity: u64,
element_stride: Option<u32>,
slot: Option<u32>,
readback: bool,
}
struct CudaShader {
device: DeviceHandle,
source: String,
search_paths: Vec<String>,
defines: Vec<(String, String)>,
optimization_level: crate::types::OptimizationLevel,
}
struct CudaComputePipeline {
device: DeviceHandle,
#[allow(dead_code)]
module: Arc<CudaModule>,
function: CudaFunction,
workgroup_size: [u32; 3],
slot_access: Vec<Option<ResourceAccess>>,
}
impl CudaBackend {
pub(crate) fn new() -> Result<Self> {
ensure_cuda_toolkit_on_path();
cudarc::driver::result::init().context("CUDA: driver init failed")?;
let count = CudaContext::device_count().context("CUDA: enumerate devices")?;
if count <= 0 {
anyhow::bail!("CUDA: no devices found");
}
let mut adapter_info = Vec::with_capacity(count as usize);
for ordinal in 0..count {
let ctx = CudaContext::new(ordinal as usize).with_context(|| format!("CUDA: open device {ordinal}"))?;
let name = ctx.name().unwrap_or_else(|_| format!("CUDA device {ordinal}"));
adapter_info.push(AdapterInfo {
id: ordinal as u32,
name,
vendor: "NVIDIA".to_string(),
backend: BackendType::Cuda,
device_type: DeviceType::DiscreteGpu,
});
}
Ok(Self {
adapter_info,
devices: HashMap::new(),
contexts: HashMap::new(),
buffers: HashMap::new(),
buffer_slots: HashMap::new(),
shaders: HashMap::new(),
compute_pipelines: HashMap::new(),
next_device: 1,
next_context: 1,
next_buffer: 1,
next_slot: 0,
next_shader: 1,
next_compute_pipeline: 1,
})
}
fn device(&self, handle: DeviceHandle) -> Result<&CudaDevice> {
self.devices.get(&handle).context("CUDA: invalid device handle")
}
fn context(&self, handle: ContextHandle) -> Result<&Arc<CudaSubmitContext>> {
self.contexts.get(&handle).context("CUDA: invalid context handle")
}
fn unsupported<T>(operation: &str) -> Result<T> {
anyhow::bail!("CUDA compute-only backend does not support {operation}")
}
fn create_storage_buffer(
&mut self,
device: DeviceHandle,
logical_size: u64,
capacity: u64,
element_stride: Option<u32>,
) -> Result<BufferHandle> {
let capacity = capacity.max(logical_size).max(4);
let gpu = self.device(device)?;
let memory = Arc::new(Mutex::new(
gpu.stream
.alloc_zeros::<u8>(capacity as usize)
.context("CUDA: alloc buffer")?,
));
let handle = self.next_buffer;
self.next_buffer += 1;
let slot = self.next_slot;
self.next_slot = self
.next_slot
.checked_add(1)
.context("CUDA buffer registry exhausted")?;
self.buffer_slots.insert(slot, handle);
self.buffers.insert(
handle,
CudaBuffer {
device,
memory,
offset: 0,
size: logical_size,
capacity,
element_stride,
slot: Some(slot),
readback: false,
},
);
Ok(handle)
}
fn compile_compute_ptx(&self, shader: &CudaShader) -> Result<(String, Vec<Option<ResourceAccess>>, [u32; 3])> {
ensure_cuda_toolkit_on_path();
let compiler = crate::slang::SlangCompiler::new().context("CUDA: initialize Slang")?;
let paths: Vec<&str> = shader.search_paths.iter().map(String::as_str).collect();
let defines: Vec<(&str, &str)> = shader
.defines
.iter()
.map(|(name, value)| (name.as_str(), value.as_str()))
.collect();
let cuda_source = crate::slang::virtual_main::transform_virtual_main_cuda_compute(&shader.source)
.map_err(|error| anyhow::anyhow!("CUDA shader lowering failed: {error}"))?;
let workgroup_size = crate::slang::parse_numthreads(&shader.source).unwrap_or([1, 1, 1]);
let compiled = compiler.compile_bindless_with_reflection_and_defines(
&cuda_source,
crate::slang::ShaderTarget::Ptx,
&[("cs_main", crate::slang::SlangStage::Compute)],
&paths,
&defines,
&[],
shader.optimization_level,
)?;
let mut ptx = compiled
.shader
.as_str()
.context("CUDA: Slang returned non-text PTX output")?
.to_owned();
while ptx.ends_with('\0') {
ptx.pop();
}
let access = crate::slang::virtual_main::extract_push_constant_categories(&shader.source)
.iter()
.map(|category| {
category.map(|category| match category {
crate::types::ResourceCategory::Broadcast
| crate::types::ResourceCategory::Texture
| crate::types::ResourceCategory::Sampler => ResourceAccess::Read,
crate::types::ResourceCategory::Scattered | crate::types::ResourceCategory::StorageImage => {
ResourceAccess::ReadWrite
}
})
})
.collect();
Ok((ptx, access, workgroup_size))
}
fn buffer_arg(&self, stream: &Arc<CudaStream>, buffer: &CudaBuffer) -> Result<CudaBufferArg> {
let memory = buffer.memory.lock().unwrap();
let start = buffer.offset as usize;
let end = (buffer.offset + buffer.size) as usize;
let view = memory.try_slice(start..end).context("CUDA: buffer view out of range")?;
let (ptr, _sync) = view.device_ptr(stream);
let stride = buffer.element_stride.unwrap_or(1).max(1) as u64;
let count = if buffer.size == 0 {
0
} else {
(buffer.size / stride) as usize
};
Ok(CudaBufferArg { data: ptr, count })
}
fn write_buffer_region(stream: &Arc<CudaStream>, buffer: &CudaBuffer, offset: u64, data: &[u8]) -> Result<()> {
if offset + data.len() as u64 > buffer.size {
anyhow::bail!("CUDA: write exceeds logical buffer size");
}
let mut memory = buffer.memory.lock().unwrap();
let start = (buffer.offset + offset) as usize;
let end = start + data.len();
let mut view = memory
.try_slice_mut(start..end)
.context("CUDA: write range out of bounds")?;
stream.memcpy_htod(data, &mut view).context("CUDA: HtoD write failed")
}
fn clear_buffer_region(stream: &Arc<CudaStream>, buffer: &CudaBuffer, offset: u64, size: u64) -> Result<()> {
let clear_size = if size == 0 {
buffer.size.saturating_sub(offset)
} else {
size
};
let mut memory = buffer.memory.lock().unwrap();
let start = (buffer.offset + offset) as usize;
let end = start + clear_size as usize;
let mut view = memory
.try_slice_mut(start..end)
.context("CUDA: clear range out of bounds")?;
stream.memset_zeros(&mut view).context("CUDA: memset failed")
}
fn copy_buffer_region(
stream: &Arc<CudaStream>,
src: &CudaBuffer,
src_offset: u64,
dst: &CudaBuffer,
dst_offset: u64,
size: u64,
) -> Result<()> {
let mut host = vec![0u8; size as usize];
{
let memory = src.memory.lock().unwrap();
let src_view = memory
.try_slice((src.offset + src_offset) as usize..(src.offset + src_offset + size) as usize)
.context("CUDA: copy source out of bounds")?;
stream.memcpy_dtoh(&src_view, &mut host).context("CUDA: copy DtoH")?;
}
{
let mut memory = dst.memory.lock().unwrap();
let mut dst_view = memory
.try_slice_mut((dst.offset + dst_offset) as usize..(dst.offset + dst_offset + size) as usize)
.context("CUDA: copy destination out of bounds")?;
stream.memcpy_htod(&host, &mut dst_view).context("CUDA: copy HtoD")?;
}
Ok(())
}
fn submit_commands(
&mut self,
ctx: ContextHandle,
commands: &[GpuCommand],
) -> Result<crate::timeline::TimelineValue> {
let context = Arc::clone(self.context(ctx)?);
let device_handle = context.device;
let stream = Arc::clone(&self.device(device_handle)?.stream);
let next_timeline = Arc::clone(&self.device(device_handle)?.next_timeline);
let retired = Arc::clone(&self.device(device_handle)?.retired);
let mut current_pipeline: Option<ComputePipelineHandle> = None;
let mut current_indices: Vec<u32> = Vec::new();
for command in commands {
match command {
GpuCommand::SetPipeline(pipeline) => current_pipeline = Some(*pipeline),
GpuCommand::BindResourcesRaw { indices, user, .. } => {
if !user.is_empty() {
anyhow::bail!(
"CUDA: user scalar dispatch parameters are not implemented; use a bound broadcast buffer"
);
}
current_indices.clone_from(indices);
}
GpuCommand::Dispatch {
workgroups_x,
workgroups_y,
workgroups_z,
..
} => {
let pipeline_handle = current_pipeline.context("CUDA: dispatch without a compute pipeline")?;
let pipeline = self
.compute_pipelines
.get(&pipeline_handle)
.context("CUDA: invalid compute pipeline")?;
let mut args = Vec::with_capacity(current_indices.len());
for (binding, index) in current_indices.iter().copied().enumerate() {
let handle = self.buffer_slots.get(&index).with_context(|| {
format!("CUDA: binding {binding} references unknown registry key {index}")
})?;
let buffer = self
.buffers
.get(handle)
.with_context(|| format!("CUDA: registry key {index} references a destroyed buffer"))?;
args.push(self.buffer_arg(&stream, buffer)?);
}
let cfg = LaunchConfig {
grid_dim: (*workgroups_x, *workgroups_y, *workgroups_z),
block_dim: (
pipeline.workgroup_size[0],
pipeline.workgroup_size[1],
pipeline.workgroup_size[2],
),
shared_mem_bytes: 0,
};
unsafe {
let mut builder = stream.launch_builder(&pipeline.function);
for arg in &args {
builder.arg(arg);
}
builder.launch(cfg).context("CUDA: cuLaunchKernel failed")?;
}
}
GpuCommand::DispatchIndirect { .. } => {
anyhow::bail!(
"CUDA compute-only PoC does not support indirect dispatch; \
use the graphics-companion fallback in the full CUDA backend"
)
}
GpuCommand::ClearBuffer { buffer, offset, size } => {
let buffer = self.buffers.get(buffer).context("CUDA: invalid clear buffer")?;
Self::clear_buffer_region(&stream, buffer, *offset, *size)?;
}
GpuCommand::WriteBuffer { buffer, offset, data } => {
let buffer = self.buffers.get(buffer).context("CUDA: invalid write buffer")?;
Self::write_buffer_region(&stream, buffer, *offset, data)?;
}
GpuCommand::CopyBuffer {
src,
src_offset,
dst,
dst_offset,
size,
} => {
let src_buf = self.buffers.get(src).context("CUDA: invalid copy source")?.clone_meta();
let dst_buf = self
.buffers
.get(dst)
.context("CUDA: invalid copy destination")?
.clone_meta();
Self::copy_buffer_region(&stream, &src_buf, *src_offset, &dst_buf, *dst_offset, *size)?;
}
GpuCommand::FrameTableStaging { .. } | GpuCommand::ResourceBarrier { .. } => {
}
GpuCommand::DispatchBatch { .. } => {
anyhow::bail!("CUDA: native dispatch batching is not supported")
}
GpuCommand::WriteTexture { .. }
| GpuCommand::WriteTextureRegion { .. }
| GpuCommand::CopyTexture { .. }
| GpuCommand::CopyRenderTarget { .. }
| GpuCommand::CopyBufferToTexture { .. }
| GpuCommand::CopyTextureToReadback { .. } => {
anyhow::bail!("CUDA compute-only backend: texture command is not supported")
}
}
}
stream.synchronize().context("CUDA: stream synchronize failed")?;
let value = next_timeline.fetch_add(1, Ordering::AcqRel);
context.completed.store(value, Ordering::Release);
retired.fetch_max(value, Ordering::AcqRel);
context.signal_queue.push_boundary_crossed(value);
Ok(value)
}
}
impl CudaBuffer {
fn clone_meta(&self) -> Self {
Self {
device: self.device,
memory: Arc::clone(&self.memory),
offset: self.offset,
size: self.size,
capacity: self.capacity,
element_stride: self.element_stride,
slot: self.slot,
readback: self.readback,
}
}
}
fn ensure_cuda_toolkit_on_path() {
let path = std::env::var_os("PATH").unwrap_or_default();
let candidates = [
std::env::var_os("CUDA_PATH")
.map(PathBuf::from)
.map(|p| p.join("bin/x64")),
std::env::var_os("CUDA_PATH").map(PathBuf::from).map(|p| p.join("bin")),
Some(PathBuf::from(
r"C:\Program Files\NVIDIA GPU Computing Toolkit\CUDA\v13.1\bin\x64",
)),
Some(PathBuf::from(
r"C:\Program Files\NVIDIA GPU Computing Toolkit\CUDA\v13.1\bin",
)),
Some(PathBuf::from("/usr/local/cuda/bin")),
];
for cand in candidates.into_iter().flatten() {
if !cand.is_dir() {
continue;
}
let cand_os = cand.as_os_str();
if path.to_string_lossy().contains(cand.to_string_lossy().as_ref()) {
return;
}
let mut new_path = cand_os.to_os_string();
#[cfg(windows)]
new_path.push(";");
#[cfg(not(windows))]
new_path.push(":");
new_path.push(&path);
unsafe { std::env::set_var("PATH", new_path) };
return;
}
}
impl GpuBackendSubmitSession for CudaBackend {
fn clone_context_submit_session(
&self,
_ctx: ContextHandle,
backend: std::sync::Arc<std::sync::Mutex<Box<dyn GpuBackend>>>,
) -> std::sync::Arc<dyn ContextSubmitSession> {
LockedSubmitSession::with_backend_type(backend, BackendType::Cuda)
}
}
impl GpuBackendTimelineWait for CudaBackend {
fn take_timeline_submission_epoch_wait(
&self,
_ctx: ContextHandle,
_value: crate::timeline::TimelineValue,
) -> Result<Option<submission_worker::SubmissionEpochWait>> {
Ok(None)
}
fn take_timeline_blocking_wait(
&self,
_ctx: ContextHandle,
_value: crate::timeline::TimelineValue,
) -> Result<Option<Box<dyn TimelineBlockingWait>>> {
Ok(None)
}
fn finish_timeline_wait(&mut self, ctx: ContextHandle, value: crate::timeline::TimelineValue) -> Result<()> {
if self.gpu_progress(ctx) < value {
anyhow::bail!("CUDA: timeline value {value} was not submitted on context {ctx}");
}
Ok(())
}
}
impl GpuBackendPresentSplit for CudaBackend {
fn take_present_gpu_work(
&mut self,
_frame: FrameToken,
_submit_tv: crate::timeline::TimelineValue,
) -> Result<Box<dyn PresentGpuWork>> {
Self::unsupported("presentation")
}
fn finish_present(
&mut self,
_finish: PresentFinishState,
_submit_tv: crate::timeline::TimelineValue,
) -> Result<crate::timeline::TimelineValue> {
Self::unsupported("presentation")
}
}
impl GpuBackend for CudaBackend {
fn as_any_mut(&mut self) -> &mut dyn std::any::Any {
self
}
fn backend_type(&self) -> BackendType {
BackendType::Cuda
}
fn enumerate_adapters(&self) -> Vec<AdapterInfo> {
self.adapter_info.clone()
}
fn adapter_capabilities(&self, _adapter_id: u32) -> crate::device::DeviceCapabilities {
crate::device::DeviceCapabilities {
has_zero_copy_storage_readback: false,
buffer_resize_cost: BufferResizeCost::Copy,
buffer_decommit_supported: false,
host_sidecar_on_submit_worker: false,
split_compute_partitions_on_barrier_cost: false,
fuse_upload_with_compute_partitions: true,
..crate::device::DeviceCapabilities::default()
}
}
fn create_device(&mut self, adapter_id: u32) -> Result<DeviceHandle> {
ensure_cuda_toolkit_on_path();
let ctx = CudaContext::new(adapter_id as usize)
.with_context(|| format!("CUDA: create device for adapter {adapter_id}"))?;
let stream = ctx.default_stream();
let handle = self.next_device;
self.next_device += 1;
self.devices.insert(
handle,
CudaDevice {
ctx,
stream,
next_timeline: Arc::new(AtomicU64::new(1)),
retired: Arc::new(AtomicU64::new(0)),
},
);
Ok(handle)
}
fn destroy_device(&mut self, device: DeviceHandle) {
if let Some(gpu) = self.devices.remove(&device) {
let _ = gpu.stream.synchronize();
}
self.contexts.retain(|_, context| context.device != device);
self.buffers.retain(|_, buffer| buffer.device != device);
self.shaders.retain(|_, shader| shader.device != device);
self.compute_pipelines.retain(|_, pipeline| pipeline.device != device);
self.buffer_slots.retain(|_, handle| self.buffers.contains_key(handle));
}
fn is_device_valid(&self, device: DeviceHandle) -> bool {
self.devices.contains_key(&device)
}
fn device_wait_idle(&mut self, device: DeviceHandle) -> Result<()> {
self.device(device)?
.stream
.synchronize()
.context("CUDA: device wait idle failed")
}
fn create_context(&mut self, device: DeviceHandle) -> Result<ContextHandle> {
self.device(device)?;
if self.contexts.values().any(|context| context.device == device) {
anyhow::bail!("CUDA prototype supports one submission context per device");
}
let handle = self.next_context;
self.next_context += 1;
self.contexts.insert(
handle,
Arc::new(CudaSubmitContext {
device,
completed: AtomicU64::new(0),
signal_queue: crate::signal::SignalQueue::new(),
}),
);
Ok(handle)
}
fn detach_context_for_destroy(&mut self, ctx: ContextHandle) -> Option<Box<dyn ContextDestroyHandle>> {
self.contexts.remove(&ctx)?;
Some(Box::new(CudaDestroyContext))
}
fn clone_context_deletion_flush(
&self,
ctx: ContextHandle,
) -> Option<std::sync::Arc<dyn ContextDeferredDeletionFlush>> {
self.contexts
.contains_key(&ctx)
.then(|| Arc::new(NoOpDeferredDeletionFlush) as Arc<dyn ContextDeferredDeletionFlush>)
}
fn clone_context_gpu_progress(&self, ctx: ContextHandle) -> Option<std::sync::Arc<dyn ContextGpuProgress>> {
Some(Arc::new(CudaProgress {
context: Arc::clone(self.contexts.get(&ctx)?),
}))
}
fn context_device(&self, ctx: ContextHandle) -> DeviceHandle {
self.contexts.get(&ctx).map(|context| context.device).unwrap_or(0)
}
fn create_buffer(
&mut self,
device: DeviceHandle,
size: u64,
_access: BufferKind,
element_stride: Option<u32>,
_flags: BufferFlags,
) -> Result<BufferHandle> {
self.create_storage_buffer(device, size, size, element_stride)
}
fn create_buffer_with_capacity(
&mut self,
device: DeviceHandle,
initial_size: u64,
capacity: u64,
_access: BufferKind,
element_stride: Option<u32>,
_flags: BufferFlags,
) -> Result<(BufferHandle, u64)> {
let capacity = capacity.max(initial_size);
Ok((
self.create_storage_buffer(device, initial_size, capacity, element_stride)?,
capacity,
))
}
fn destroy_buffer(&mut self, buffer: BufferHandle) {
if let Some(buffer) = self.buffers.remove(&buffer) {
if let Some(slot) = buffer.slot {
self.buffer_slots.remove(&slot);
}
}
}
fn write_buffer(&mut self, buffer: BufferHandle, offset: u64, data: &[u8]) -> Result<()> {
let buffer = self.buffers.get(&buffer).context("CUDA: invalid buffer handle")?;
let stream = Arc::clone(&self.device(buffer.device)?.stream);
Self::write_buffer_region(&stream, buffer, offset, data)
}
fn alloc_readback_buffer(&mut self, device: DeviceHandle, size: u64) -> Result<BufferHandle> {
let gpu = self.device(device)?;
let capacity = size.max(4);
let memory = Arc::new(Mutex::new(
gpu.stream
.alloc_zeros::<u8>(capacity as usize)
.context("CUDA: alloc readback")?,
));
let handle = self.next_buffer;
self.next_buffer += 1;
self.buffers.insert(
handle,
CudaBuffer {
device,
memory,
offset: 0,
size,
capacity,
element_stride: None,
slot: None,
readback: true,
},
);
Ok(handle)
}
fn read_readback_buffer(&self, buffer: BufferHandle, output: &mut [u8]) -> Result<()> {
let buffer = self.buffers.get(&buffer).context("CUDA: invalid readback buffer")?;
if !buffer.readback {
anyhow::bail!("CUDA: buffer is not readback staging");
}
if output.len() as u64 > buffer.size {
anyhow::bail!("CUDA: read exceeds readback buffer size");
}
let stream = Arc::clone(&self.device(buffer.device)?.stream);
let memory = buffer.memory.lock().unwrap();
let view = memory
.try_slice(buffer.offset as usize..(buffer.offset as usize + output.len()))
.context("CUDA: readback range out of bounds")?;
stream.memcpy_dtoh(&view, output).context("CUDA: DtoH readback failed")
}
fn free_readback_buffer(&mut self, buffer: BufferHandle) {
self.destroy_buffer(buffer);
}
fn query_texture_copy_footprint(
&self,
_device: DeviceHandle,
_width: u32,
_height: u32,
_format: TextureFormat,
) -> Result<TextureCopyFootprint> {
Self::unsupported("texture readback")
}
fn alloc_texture_readback_staging(
&mut self,
_device: DeviceHandle,
_layout: TextureCopyFootprint,
) -> Result<BufferHandle> {
Self::unsupported("texture readback")
}
fn read_texture_readback_staging(
&self,
_buffer: BufferHandle,
_layout: TextureCopyFootprint,
_output: &mut [u8],
) -> Result<()> {
Self::unsupported("texture readback")
}
fn texture_copy_retention_tag(&self, _texture: TextureHandle) -> u64 {
0
}
fn clear_buffer(&mut self, device: DeviceHandle, buffer: BufferHandle, offset: u64, size: u64) -> Result<()> {
let stream = Arc::clone(&self.device(device)?.stream);
let target = self.buffers.get(&buffer).context("CUDA: invalid buffer handle")?;
Self::clear_buffer_region(&stream, target, offset, size)
}
fn buffer_size(&self, buffer: BufferHandle) -> u64 {
self.buffers.get(&buffer).map(|buffer| buffer.size).unwrap_or(0)
}
fn buffer_capacity(&self, buffer: BufferHandle) -> u64 {
self.buffers.get(&buffer).map(|buffer| buffer.capacity).unwrap_or(0)
}
fn set_buffer_logical_size(
&mut self,
_device: DeviceHandle,
buffer: BufferHandle,
new_logical_size: u64,
) -> Result<()> {
let buffer = self.buffers.get_mut(&buffer).context("CUDA: invalid buffer handle")?;
if new_logical_size == 0 || new_logical_size > buffer.capacity {
anyhow::bail!("CUDA: logical size must be in 1..=capacity");
}
buffer.size = new_logical_size;
Ok(())
}
fn buffer_bindless_index(&self, buffer: BufferHandle) -> Option<u32> {
self.buffers.get(&buffer)?.slot
}
fn buffer_bindless_srv_index(&self, buffer: BufferHandle) -> Option<u32> {
self.buffer_bindless_index(buffer)
}
fn create_buffer_view(
&mut self,
parent: BufferHandle,
offset: u64,
size: u64,
element_stride: Option<u32>,
) -> Result<BufferHandle> {
let parent = self
.buffers
.get(&parent)
.context("CUDA: invalid parent buffer")?
.clone_meta();
if offset + size > parent.size {
anyhow::bail!("CUDA: buffer view exceeds parent");
}
let handle = self.next_buffer;
self.next_buffer += 1;
let slot = self.next_slot;
self.next_slot += 1;
self.buffer_slots.insert(slot, handle);
self.buffers.insert(
handle,
CudaBuffer {
device: parent.device,
memory: parent.memory,
offset: parent.offset + offset,
size,
capacity: size,
element_stride: element_stride.or(parent.element_stride),
slot: Some(slot),
readback: false,
},
);
Ok(handle)
}
fn resize_buffer(
&mut self,
device: DeviceHandle,
buffer: BufferHandle,
new_size: u64,
preserve_contents: bool,
) -> Result<()> {
let old = self
.buffers
.get(&buffer)
.context("CUDA: invalid buffer handle")?
.clone_meta();
if old.device != device {
anyhow::bail!("CUDA: buffer belongs to another device");
}
let stream = Arc::clone(&self.device(device)?.stream);
let capacity = new_size.max(4);
let replacement = Arc::new(Mutex::new(
stream
.alloc_zeros::<u8>(capacity as usize)
.context("CUDA: resize alloc")?,
));
if preserve_contents {
let copy_size = old.size.min(new_size);
if copy_size > 0 {
let mut host = vec![0u8; copy_size as usize];
{
let memory = old.memory.lock().unwrap();
let src = memory
.try_slice(old.offset as usize..(old.offset + copy_size) as usize)
.context("CUDA: resize src")?;
stream.memcpy_dtoh(&src, &mut host).context("CUDA: resize DtoH")?;
}
{
let mut memory = replacement.lock().unwrap();
let mut dst = memory
.try_slice_mut(0..copy_size as usize)
.context("CUDA: resize dst")?;
stream.memcpy_htod(&host, &mut dst).context("CUDA: resize HtoD")?;
}
stream.synchronize().context("CUDA: resize sync")?;
}
}
let target = self.buffers.get_mut(&buffer).expect("validated above");
target.memory = replacement;
target.offset = 0;
target.size = new_size;
target.capacity = capacity;
Ok(())
}
fn create_shader_with_paths(
&mut self,
device: DeviceHandle,
slang_source: &str,
search_paths: &[&str],
defines: &[(&str, &str)],
optimization_level: crate::types::OptimizationLevel,
) -> Result<ShaderHandle> {
self.device(device)?;
let handle = self.next_shader;
self.next_shader += 1;
self.shaders.insert(
handle,
CudaShader {
device,
source: slang_source.to_owned(),
search_paths: search_paths.iter().map(|value| (*value).to_owned()).collect(),
defines: defines
.iter()
.map(|(name, value)| ((*name).to_owned(), (*value).to_owned()))
.collect(),
optimization_level,
},
);
Ok(handle)
}
fn destroy_shader(&mut self, shader: ShaderHandle) {
self.shaders.remove(&shader);
}
fn create_pipeline(
&mut self,
_device: DeviceHandle,
_vertex_shader: ShaderHandle,
_fragment_shader: ShaderHandle,
_vertex_layout: &VertexBufferLayout,
_topology: PrimitiveTopology,
_target_format: TextureFormat,
) -> Result<PipelineHandle> {
Self::unsupported("graphics pipelines")
}
fn destroy_pipeline(&mut self, _pipeline: PipelineHandle) {}
fn create_pipeline_with_depth(
&mut self,
_device: DeviceHandle,
_vertex_shader: ShaderHandle,
_fragment_shader: ShaderHandle,
_vertex_layout: &VertexBufferLayout,
_topology: PrimitiveTopology,
_target_format: TextureFormat,
_depth_stencil: Option<&DepthStencilState>,
) -> Result<PipelineHandle> {
Self::unsupported("graphics pipelines")
}
fn create_render_target_with_depth(
&mut self,
_device: DeviceHandle,
_width: u32,
_height: u32,
_color_format: TextureFormat,
_depth_format: Option<DepthFormat>,
) -> Result<RenderTargetHandle> {
Self::unsupported("render targets")
}
fn render_to_target(
&mut self,
_device: DeviceHandle,
_target: RenderTargetHandle,
_color_load: crate::types::TargetLoad,
_commands: &[RenderCommand],
) -> Result<()> {
Self::unsupported("rendering")
}
fn create_texture(
&mut self,
_device: DeviceHandle,
_width: u32,
_height: u32,
_format: TextureFormat,
_access: TextureKind,
_flags: TextureFlags,
) -> Result<TextureHandle> {
Self::unsupported("textures")
}
fn write_texture(&mut self, _texture: TextureHandle, _data: &[u8], _width: u32, _height: u32) -> Result<()> {
Self::unsupported("textures")
}
fn write_texture_region(
&mut self,
_texture: TextureHandle,
_x: u32,
_y: u32,
_width: u32,
_height: u32,
_data: &[u8],
) -> Result<()> {
Self::unsupported("textures")
}
fn destroy_texture(&mut self, _texture: TextureHandle) {}
fn texture_bindless_index(&self, _texture: TextureHandle) -> Option<u32> {
None
}
fn texture_bindless_sampled_index(&self, _texture: TextureHandle) -> Option<u32> {
None
}
fn create_sampler(&mut self, _device: DeviceHandle, _desc: &SamplerDesc) -> Result<SamplerHandle> {
Self::unsupported("samplers")
}
fn destroy_sampler(&mut self, _sampler: SamplerHandle) {}
fn sampler_bindless_index(&self, _sampler: SamplerHandle) -> Option<u32> {
None
}
fn create_surface(
&mut self,
_device: DeviceHandle,
_window: &dyn raw_window_handle::HasWindowHandle,
_display: &dyn raw_window_handle::HasDisplayHandle,
_depth_format: Option<DepthFormat>,
) -> Result<SurfaceHandle> {
Self::unsupported("surfaces")
}
fn destroy_surface(&mut self, _surface: SurfaceHandle) {}
fn surface_resize(&mut self, _surface: SurfaceHandle, _width: u32, _height: u32) -> Result<()> {
Self::unsupported("surfaces")
}
fn surface_size(&self, _surface: SurfaceHandle) -> (u32, u32) {
(0, 0)
}
fn surface_format(&self, _surface: SurfaceHandle) -> TextureFormat {
TextureFormat::Bgra8UnormSrgb
}
fn gpu_progress(&self, ctx: ContextHandle) -> crate::timeline::TimelineValue {
self.contexts
.get(&ctx)
.map(|context| context.completed.load(Ordering::Acquire))
.unwrap_or(0)
}
fn device_timeline_retired(&self, device: DeviceHandle) -> crate::timeline::TimelineValue {
self.devices
.get(&device)
.map(|device| device.retired.load(Ordering::Acquire))
.unwrap_or(0)
}
fn device_wait_until(&mut self, device: DeviceHandle, value: crate::timeline::TimelineValue) -> Result<()> {
self.device_wait_idle(device)?;
if self.device_timeline_retired(device) < value {
anyhow::bail!("CUDA: timeline value {value} has not been submitted");
}
Ok(())
}
fn poll_signals(
&mut self,
ctx: ContextHandle,
_progress: crate::timeline::TimelineValue,
) -> Vec<crate::signal::QueuedSignal> {
self.contexts
.get(&ctx)
.map(|context| crate::signal::drain_all_queued_signals(&context.signal_queue))
.unwrap_or_default()
}
fn submit_standalone(
&mut self,
ctx: ContextHandle,
commands: &[GpuCommand],
sync: Option<&SubmitSync>,
) -> Result<crate::timeline::TimelineValue> {
if let Some(sync) = sync {
for epoch in sync
.waits
.iter()
.chain(sync.cpu_waits.iter())
.chain(sync.host_observed_waits.iter())
{
self.device_wait_until(self.context_device(ctx), epoch.value)?;
}
for write in &sync.deferred_host_writes {
self.write_buffer(write.buffer, write.offset, &write.data)?;
}
}
let effective = commands_with_sync_prologue(commands, sync);
self.submit_commands(ctx, &effective)
}
fn begin_frame(&mut self, _surface: SurfaceHandle, _ctx: ContextHandle) -> Result<(FrameToken, TextureHandle)> {
Self::unsupported("frames")
}
fn submit_frame(&mut self, _frame: &FrameToken) -> Result<crate::timeline::TimelineValue> {
Self::unsupported("frames")
}
fn create_compute_pipeline(
&mut self,
device: DeviceHandle,
compute_shader: ShaderHandle,
_debug_name: Option<&str>,
) -> Result<ComputePipelineHandle> {
let shader = self
.shaders
.get(&compute_shader)
.context("CUDA: invalid shader handle")?;
if shader.device != device {
anyhow::bail!("CUDA: shader belongs to another device");
}
let (ptx, slot_access, workgroup_size) = self.compile_compute_ptx(shader)?;
let gpu = self.device(device)?;
let module = gpu
.ctx
.load_module(Ptx::from_src(ptx))
.context("CUDA: cuModuleLoadData failed")?;
let function = module
.load_function("cs_main")
.context("CUDA: cuModuleGetFunction(cs_main) failed")?;
let handle = self.next_compute_pipeline;
self.next_compute_pipeline += 1;
self.compute_pipelines.insert(
handle,
CudaComputePipeline {
device,
module,
function,
workgroup_size,
slot_access,
},
);
Ok(handle)
}
fn destroy_compute_pipeline(&mut self, pipeline: ComputePipelineHandle) {
self.compute_pipelines.remove(&pipeline);
}
fn compute_pipeline_slot_access(&self, pipeline: ComputePipelineHandle) -> Vec<Option<ResourceAccess>> {
self.compute_pipelines
.get(&pipeline)
.map(|pipeline| pipeline.slot_access.clone())
.unwrap_or_default()
}
fn max_bindless_slots_per_category(&self, _device: DeviceHandle, category: crate::types::ResourceCategory) -> u32 {
if matches!(
category,
crate::types::ResourceCategory::Scattered | crate::types::ResourceCategory::Broadcast
) {
4096
} else {
0
}
}
fn available_bindless_slots(&self, device: DeviceHandle, category: crate::types::ResourceCategory) -> u32 {
self.max_bindless_slots_per_category(device, category).saturating_sub(
self.buffers
.values()
.filter(|buffer| buffer.device == device && buffer.slot.is_some())
.count() as u32,
)
}
fn max_submission_contexts(&self, _device: DeviceHandle) -> u32 {
1
}
}
#[cfg(test)]
mod tests {
use super::*;
const DOUBLE_SLANG: &str = r#"
[shader("compute")]
[numthreads(1, 1, 1)]
void cs_main(uniform RWStructuredBuffer<uint> values, uint3 id : SV_DispatchThreadID) {
values[id.x] = values[id.x] * 2;
}
"#;
const DOUBLE_GOLDY_SLANG: &str = r#"
import goldy_exp;
[goldy_compute]
[numthreads(1, 1, 1)]
void cs_main(Scattered<uint> values, ThreadId id) {
values[id.x] = values[id.x] * 2;
}
"#;
const DOUBLE_GOLDY_TWO_BUFFER_SLANG: &str = r#"
import goldy_exp;
[goldy_compute]
[numthreads(1, 1, 1)]
void cs_main(BufRO<uint> input, Scattered<uint> output, ThreadId id) {
output[id.x] = input[id.x] * 2;
}
"#;
fn run_compute_dispatch_and_readback(shader_source: &str) -> Result<()> {
let mut backend = match CudaBackend::new() {
Ok(backend) => backend,
Err(error) => {
eprintln!("skipping CUDA compute test: {error:#}");
return Ok(());
}
};
let device = backend.create_device(0)?;
let ctx = backend.create_context(device)?;
let buffer = backend.create_buffer(
device,
16,
BufferKind::Scattered,
Some(4),
BufferFlags::COPY_SRC | BufferFlags::COPY_DST,
)?;
backend.write_buffer(buffer, 0, bytemuck::cast_slice(&[1u32, 2, 3, 4]))?;
let shader = backend.create_shader_with_paths(
device,
shader_source,
&[],
&[],
crate::types::OptimizationLevel::Default,
)?;
let pipeline = backend.create_compute_pipeline(device, shader, Some("double"))?;
let slot = backend.buffer_bindless_index(buffer).context("missing registry key")?;
let submitted = backend.submit_standalone(
ctx,
&[
GpuCommand::SetPipeline(pipeline),
GpuCommand::BindResourcesRaw {
indices: vec![slot],
user: vec![],
frame_table_base: 0,
},
GpuCommand::Dispatch {
label: Some("double"),
workgroups_x: 4,
workgroups_y: 1,
workgroups_z: 1,
},
],
None,
)?;
assert_eq!(backend.gpu_progress(ctx), submitted);
let readback = backend.alloc_readback_buffer(device, 16)?;
backend.submit_standalone(
ctx,
&[GpuCommand::CopyBuffer {
src: buffer,
src_offset: 0,
dst: readback,
dst_offset: 0,
size: 16,
}],
None,
)?;
let mut bytes = [0u8; 16];
backend.read_readback_buffer(readback, &mut bytes)?;
assert_eq!(bytemuck::cast_slice::<u8, u32>(&bytes), &[2, 4, 6, 8]);
Ok(())
}
#[test]
fn slang_compute_dispatch_and_readback() -> Result<()> {
run_compute_dispatch_and_readback(DOUBLE_SLANG)
}
fn run_scheme_compute_and_withdraw(shader_source: &str) -> Result<()> {
let backend = match CudaBackend::new() {
Ok(backend) => backend,
Err(error) => {
eprintln!("skipping CUDA scheme test: {error:#}");
return Ok(());
}
};
let device = Arc::new(crate::Device::from_backend(Box::new(backend))?);
let ctx = device.create_context()?;
let mut pool = crate::RetainedPool::new(Arc::clone(&device));
let buffer = pool.acquire_buffer_with_data(&[1u32, 2, 3, 4], BufferKind::Scattered)?;
let shader = crate::ShaderModule::from_slang(&device, shader_source)?;
let pipeline = crate::ComputePipeline::new(&device, &shader)?;
let mut scheme = crate::Scheme::new(&ctx);
scheme
.node("double", &pipeline)
.with_parcel(&buffer, crate::NodeAccess::ReadWrite)
.dispatch(4, 1, 1);
let withdraw = crate::MemoryExchange::new(&ctx).bind_withdraw(&mut scheme, &buffer)?;
let mut submission = scheme.submit()?;
let bytes = withdraw.claim(&mut submission)?.consume()?;
assert_eq!(bytemuck::cast_slice::<u8, u32>(&bytes), &[2, 4, 6, 8]);
Ok(())
}
#[test]
fn scheme_dispatches_goldy_virtual_compute_and_withdraws() -> Result<()> {
run_scheme_compute_and_withdraw(DOUBLE_GOLDY_SLANG)
}
#[test]
fn scheme_binds_two_goldy_buffers_in_parameter_order() -> Result<()> {
let backend = match CudaBackend::new() {
Ok(backend) => backend,
Err(error) => {
eprintln!("skipping CUDA scheme test: {error:#}");
return Ok(());
}
};
let device = Arc::new(crate::Device::from_backend(Box::new(backend))?);
let ctx = device.create_context()?;
let mut pool = crate::RetainedPool::new(Arc::clone(&device));
let input = pool.acquire_buffer_with_data(&[1u32, 2, 3, 4], BufferKind::Scattered)?;
let output = pool.acquire_buffer_sized::<u32>(4, BufferKind::Scattered, BufferFlags::empty())?;
let shader = crate::ShaderModule::from_slang(&device, DOUBLE_GOLDY_TWO_BUFFER_SLANG)?;
let pipeline = crate::ComputePipeline::new(&device, &shader)?;
let mut scheme = crate::Scheme::new(&ctx);
scheme
.node("double", &pipeline)
.with_parcel(&input, crate::NodeAccess::Read)
.with_parcel(&output, crate::NodeAccess::Write)
.dispatch(4, 1, 1);
let withdraw = crate::MemoryExchange::new(&ctx).bind_withdraw(&mut scheme, &output)?;
let mut submission = scheme.submit()?;
let bytes = withdraw.claim(&mut submission)?.consume()?;
assert_eq!(bytemuck::cast_slice::<u8, u32>(&bytes), &[2, 4, 6, 8]);
Ok(())
}
#[test]
fn slang_emits_ptx_for_compute() -> Result<()> {
ensure_cuda_toolkit_on_path();
let compiler = match crate::slang::SlangCompiler::new() {
Ok(compiler) => compiler,
Err(error) => {
eprintln!("skipping CUDA PTX emission test: {error:#}");
return Ok(());
}
};
let compiled = match compiler.compile_bindless_with_reflection(
DOUBLE_SLANG,
crate::slang::ShaderTarget::Ptx,
&[("cs_main", crate::slang::SlangStage::Compute)],
&[],
) {
Ok(compiled) => compiled,
Err(error) => {
eprintln!("skipping CUDA PTX emission test (Slang/NVRTC): {error:#}");
return Ok(());
}
};
let ptx = compiled.shader.as_str().context("expected text PTX")?;
assert!(
ptx.contains(".entry") || ptx.contains("cs_main"),
"Slang output did not look like PTX:\n{ptx}"
);
Ok(())
}
#[test]
fn ptx_cache_key_differs_from_wgsl() {
use crate::shader_cache::compile_cache_key;
use crate::slang::{ffi::SlangStage, ShaderTarget};
use crate::types::OptimizationLevel;
let src = "void cs_main() {}";
let eps = [("cs_main", SlangStage::Compute)];
let ptx = compile_cache_key(src, ShaderTarget::Ptx, &eps, &[], &[], &[], OptimizationLevel::Default);
let wgsl = compile_cache_key(src, ShaderTarget::Wgsl, &eps, &[], &[], &[], OptimizationLevel::Default);
assert_ne!(ptx, wgsl);
}
}