lib/gpu/src/cuda.zig

daab053ee43316e1809a84551d573ddd1e5bf3d2

   1 const std = @import("std");
   2 const choir_abi = @import("choir_abi");
   3 const builtin = @import("builtin");
   4 const build_options = @import("build_options");
   5 const sys = @import("sys");
   6 
   7 const backend = @import("root.zig");
   8 const runtime_root = @import("runtime/root.zig");
   9 
  10 const driver_mod = sys.cuda;
  11 const runtime_mod = runtime_root.cuda.runtime;
  12 
  13 const Allocator = std.mem.Allocator;
  14 const BackendObjectId = backend.BackendObjectId;
  15 const Runtime = runtime_mod.Runtime;
  16 
  17 const RuntimeStorage = union(enum) {
  18     none,
  19     borrowed: *Runtime,
  20     owned: Runtime,
  21 
  22     fn ptr(self: *RuntimeStorage) ?*Runtime {
  23         return switch (self.*) {
  24             .none => null,
  25             .borrowed => |runtime| runtime,
  26             .owned => |*runtime| runtime,
  27         };
  28     }
  29 
  30     fn constPtr(self: *const RuntimeStorage) ?*const Runtime {
  31         return switch (self.*) {
  32             .none => null,
  33             .borrowed => |runtime| runtime,
  34             .owned => |*runtime| runtime,
  35         };
  36     }
  37 
  38     fn deinit(self: *RuntimeStorage) void {
  39         switch (self.*) {
  40             .none, .borrowed => {},
  41             .owned => |*runtime| runtime.deinit(),
  42         }
  43         self.* = .none;
  44     }
  45 };
  46 
  47 pub fn platformSupported() bool {
  48     return driver_mod.platformSupported();
  49 }
  50 
  51 pub const State = struct {
  52     allocator: Allocator,
  53     runtime: RuntimeStorage = .none,
  54     next_id: BackendObjectId = 1,
  55     objects: std.AutoHashMapUnmanaged(BackendObjectId, Object) = .{},
  56 
  57     fn init(allocator: Allocator) State {
  58         return .{
  59             .allocator = allocator,
  60         };
  61     }
  62 
  63     pub fn initDevice(allocator: Allocator, device_ordinal: i32) backend.BackendError!State {
  64         const runtime = Runtime.init(allocator, device_ordinal) catch |err| return mapRuntimeInitError(err);
  65         return .{
  66             .allocator = allocator,
  67             .runtime = .{ .owned = runtime },
  68         };
  69     }
  70 
  71     fn initWithRuntime(allocator: Allocator, runtime: ?*Runtime) State {
  72         return .{
  73             .allocator = allocator,
  74             .runtime = if (runtime) |rt| .{ .borrowed = rt } else .none,
  75         };
  76     }
  77 
  78     pub fn deinit(self: *State) void {
  79         var it = self.objects.iterator();
  80         while (it.next()) |entry| {
  81             deinitObject(self.allocator, entry.value_ptr);
  82         }
  83         self.objects.deinit(self.allocator);
  84         self.objects = .{};
  85         self.next_id = 1;
  86         self.runtime.deinit();
  87     }
  88 
  89     pub fn handle(self: *State) backend.BackendHandle {
  90         return .{
  91             .ptr = self,
  92             .vtable = &vtable,
  93             .kind = .cuda,
  94         };
  95     }
  96 
  97     fn putObject(self: *State, object: Object) backend.BackendError!BackendObjectId {
  98         const id = self.next_id;
  99         if (id == std.math.maxInt(BackendObjectId)) return error.OutOfMemory;
 100         self.next_id += 1;
 101         self.objects.put(self.allocator, id, object) catch return error.OutOfMemory;
 102         return id;
 103     }
 104 
 105     fn getLoaded(self: *State, loaded: backend.LoadedArtifact) backend.BackendError!*LoadedKernel {
 106         if (loaded.backend != .cuda or loaded.format != .cuda_ptx) return error.InvalidArtifact;
 107         const object = self.objects.getPtr(loaded.id) orelse return error.InvalidArtifact;
 108         return switch (object.*) {
 109             .loaded_artifact => |*kernel| kernel,
 110             else => error.InvalidArtifact,
 111         };
 112     }
 113 
 114     fn getBuffer(self: *State, buffer_handle: backend.BufferHandle) backend.BackendError!*runtime_mod.DeviceBuffer {
 115         if (buffer_handle.backend != .cuda) return error.InvalidBuffer;
 116         const object = self.objects.getPtr(buffer_handle.id) orelse return error.InvalidBuffer;
 117         return switch (object.*) {
 118             .buffer => |*buffer| buffer,
 119             else => error.InvalidBuffer,
 120         };
 121     }
 122 
 123     fn getStream(self: *State, stream_handle: backend.StreamHandle) backend.BackendError!*runtime_mod.Stream {
 124         if (stream_handle.backend != .cuda) return error.InvalidStream;
 125         const object = self.objects.getPtr(stream_handle.id) orelse return error.InvalidStream;
 126         return switch (object.*) {
 127             .stream => |stream| stream,
 128             else => error.InvalidStream,
 129         };
 130     }
 131 
 132     fn getEvent(self: *State, event_handle: backend.EventHandle) backend.BackendError!*runtime_mod.Event {
 133         if (event_handle.backend != .cuda) return error.InvalidEvent;
 134         const object = self.objects.getPtr(event_handle.id) orelse return error.InvalidEvent;
 135         return switch (object.*) {
 136             .event => |event| event,
 137             else => error.InvalidEvent,
 138         };
 139     }
 140 };
 141 
 142 const LoadedKernel = struct {
 143     module: *runtime_mod.Module,
 144     kernel: runtime_mod.Kernel,
 145     entry_name: [:0]u8,
 146     argument_count: u32,
 147 };
 148 
 149 const Object = union(enum) {
 150     loaded_artifact: LoadedKernel,
 151     buffer: runtime_mod.DeviceBuffer,
 152     stream: *runtime_mod.Stream,
 153     event: *runtime_mod.Event,
 154 };
 155 
 156 fn capabilitiesForRuntime(runtime: ?*const Runtime) backend.BackendCapabilities {
 157     const has_runtime = runtime != null;
 158     const device_name = if (runtime) |rt| rt.device.name() else "";
 159     const supports_tf32_tensor = if (runtime) |rt| rt.device.supportsTf32TensorOps() else false;
 160     const shared_memory_bytes = if (runtime) |rt| rt.device.dynamicSharedMemoryLimit() else 48 * 1024;
 161     return .{
 162         .identity = .{
 163             .backend = .cuda,
 164             .family = .nvidia_cuda,
 165             .name = if (device_name.len == 0) "cuda" else device_name,
 166         },
 167         .memory = .{
 168             .shared_memory_per_threadgroup_bytes = shared_memory_bytes,
 169             .constant_memory_bytes = 64 * 1024,
 170             .min_buffer_alignment = 256,
 171         },
 172         .subgroup = .{
 173             .supported = true,
 174             .size_min = 32,
 175             .size_max = 32,
 176             .shuffle = true,
 177             .ballot = true,
 178             .vote = true,
 179             .arithmetic = true,
 180             .scan = true,
 181         },
 182         .threadgroup = .{
 183             .max_threads = 1024,
 184             .max_blocks = .{ 2_147_483_647, 65_535, 65_535 },
 185             .max_threads_per_dim = .{ 1024, 1024, 64 },
 186             .max_grid_per_dim = .{ 2_147_483_647, 65_535, 65_535 },
 187             .shared_memory_bytes = shared_memory_bytes,
 188         },
 189         .dtypes = backend.DTypeSet.init(&.{ .i1, .i8, .i16, .i32, .u8, .u16, .u32, .i64, .u64, .f32, .f16, .bf16, .f64, .key }),
 190         .layouts = .{
 191             .row_major = true,
 192             .compact_strides = true,
 193             .broadcast_strides = true,
 194             .tiled = true,
 195             .vectorized = true,
 196             .opaque_backend_layouts = true,
 197         },
 198         .runtime = .{
 199             .driver_loaded = has_runtime,
 200             .device_context = has_runtime,
 201             .streams = true,
 202             .events = true,
 203         },
 204         .features = .{
 205             .atomic_i32 = true,
 206             .atomic_u32 = true,
 207             .atomic_index = true,
 208             .atomic_f32_add_device = true,
 209             .atomic_f32_add_shared = true,
 210             .tensor_cores = supports_tf32_tensor,
 211             .dynamic_shared_memory = true,
 212         },
 213         .artifact_formats = backend.ArtifactFormatSet.init(&.{.cuda_ptx}),
 214     };
 215 }
 216 
 217 fn queryCapabilities(ptr: *anyopaque) backend.BackendError!backend.BackendCapabilities {
 218     const state: *State = @ptrCast(@alignCast(ptr));
 219     return capabilitiesForRuntime(state.runtime.constPtr());
 220 }
 221 
 222 fn createArtifact(ptr: *anyopaque, request: backend.CompileRequest) backend.BackendError!backend.KernelArtifact {
 223     if (request.requested_format != .cuda_ptx) return error.UnsupportedOperation;
 224     const state: *State = @ptrCast(@alignCast(ptr));
 225     return switch (request.payload) {
 226         .text => |ptx| createPtxArtifact(state, request, ptx),
 227         .bytes => |ptx| createPtxArtifact(state, request, ptx),
 228         .none => error.UnsupportedOperation,
 229         .words_u32 => error.UnsupportedArtifactFormat,
 230     };
 231 }
 232 
 233 fn createPtxArtifact(
 234     state: *State,
 235     request: backend.CompileRequest,
 236     ptx: []const u8,
 237 ) backend.BackendError!backend.KernelArtifact {
 238     if (request.kernel_name.len == 0) return error.InvalidArtifact;
 239     if (ptx.len == 0) return error.InvalidArtifact;
 240 
 241     var artifact = backend.KernelArtifact.init(state.allocator, .{
 242         .backend = .cuda,
 243         .format = .cuda_ptx,
 244         .entry_name = request.kernel_name,
 245         .argument_count = request.argument_count,
 246         .scalar_argument_count = request.scalar_argument_count,
 247         .diagnostic_id = request.diagnostic_id,
 248     }) catch return error.OutOfMemory;
 249     errdefer artifact.deinit();
 250     try artifact.setOwnedText(ptx);
 251     return artifact;
 252 }
 253 
 254 fn loadArtifact(ptr: *anyopaque, artifact: *const backend.KernelArtifact) backend.BackendError!backend.LoadedArtifact {
 255     const state: *State = @ptrCast(@alignCast(ptr));
 256     if (artifact.backend != .cuda) return error.CapabilityMismatch;
 257     if (artifact.format != .cuda_ptx) return error.UnsupportedArtifactFormat;
 258     if (artifact.entry_name.len == 0) return error.InvalidArtifact;
 259 
 260     const ptx = switch (artifact.payload) {
 261         .text => |text| text,
 262         .bytes => |bytes| bytes,
 263         else => return error.InvalidArtifact,
 264     };
 265     if (ptx.len == 0) return error.InvalidArtifact;
 266 
 267     const rt = state.runtime.ptr() orelse return error.RuntimeUnavailable;
 268     const entry_name = state.allocator.dupeSentinel(u8, artifact.entry_name, 0) catch return error.OutOfMemory;
 269     errdefer state.allocator.free(entry_name);
 270 
 271     const module = rt.loadPtx(ptx) catch |err| return mapArtifactError(err);
 272     errdefer module.release();
 273     const kernel = module.getKernel(entry_name) catch |err| return mapArtifactError(err);
 274 
 275     const id = try state.putObject(.{ .loaded_artifact = .{
 276         .module = module,
 277         .kernel = kernel,
 278         .entry_name = entry_name,
 279         .argument_count = artifact.argument_count,
 280     } });
 281     return .{
 282         .id = id,
 283         .backend = .cuda,
 284         .format = .cuda_ptx,
 285     };
 286 }
 287 
 288 fn allocateBuffer(ptr: *anyopaque, request: backend.BufferAllocation) backend.BackendError!backend.BufferHandle {
 289     const state: *State = @ptrCast(@alignCast(ptr));
 290     if (request.byte_size == 0) return error.InvalidBuffer;
 291     const rt = state.runtime.ptr() orelse return error.RuntimeUnavailable;
 292 
 293     var buffer = runtime_mod.DeviceBuffer.alloc(rt, request.byte_size) catch |err| return mapAllocationError(err);
 294     errdefer buffer.deinit();
 295 
 296     const id = try state.putObject(.{ .buffer = buffer });
 297     return .{
 298         .id = id,
 299         .backend = .cuda,
 300         .byte_size = request.byte_size,
 301         .ownership = .backend,
 302     };
 303 }
 304 
 305 fn createStream(ptr: *anyopaque, _: backend.StreamAllocation) backend.BackendError!backend.StreamHandle {
 306     const state: *State = @ptrCast(@alignCast(ptr));
 307     const rt = state.runtime.ptr() orelse return error.RuntimeUnavailable;
 308 
 309     const stream = rt.createStream() catch |err| return mapStreamError(err);
 310     errdefer rt.destroyStream(stream);
 311 
 312     const id = try state.putObject(.{ .stream = stream });
 313     return .{
 314         .id = id,
 315         .backend = .cuda,
 316     };
 317 }
 318 
 319 fn createEvent(ptr: *anyopaque, _: backend.EventAllocation) backend.BackendError!backend.EventHandle {
 320     const state: *State = @ptrCast(@alignCast(ptr));
 321     const rt = state.runtime.ptr() orelse return error.RuntimeUnavailable;
 322 
 323     const event = rt.createEvent() catch |err| return mapEventError(err);
 324     errdefer rt.destroyEvent(event);
 325 
 326     const id = try state.putObject(.{ .event = event });
 327     return .{
 328         .id = id,
 329         .backend = .cuda,
 330     };
 331 }
 332 
 333 fn writeBuffer(ptr: *anyopaque, request: backend.BufferWriteRequest) backend.BackendError!void {
 334     const state: *State = @ptrCast(@alignCast(ptr));
 335     const buffer = try state.getBuffer(request.handle);
 336     if (request.bytes.len > buffer.bytes) return error.InvalidBuffer;
 337     buffer.copyFromHost(request.bytes) catch |err| return mapAllocationError(err);
 338 }
 339 
 340 fn fillBuffer(ptr: *anyopaque, request: backend.BufferFillRequest) backend.BackendError!void {
 341     const state: *State = @ptrCast(@alignCast(ptr));
 342     const buffer = try state.getBuffer(request.handle);
 343     buffer.fillWords(request.pattern) catch |err| return mapAllocationError(err);
 344 }
 345 
 346 fn readBuffer(ptr: *anyopaque, request: backend.BufferReadRequest) backend.BackendError!void {
 347     const state: *State = @ptrCast(@alignCast(ptr));
 348     const buffer = try state.getBuffer(request.handle);
 349     if (request.bytes.len < buffer.bytes) return error.ReadBufferDestinationTooSmall;
 350     buffer.copyToHost(request.bytes) catch |err| return mapAllocationError(err);
 351 }
 352 
 353 fn launch(ptr: *anyopaque, request: backend.LaunchRequest) backend.BackendError!void {
 354     const state: *State = @ptrCast(@alignCast(ptr));
 355     if (request.artifact.backend != .cuda or request.artifact.format != .cuda_ptx) {
 356         return error.CapabilityMismatch;
 357     }
 358     const rt = state.runtime.ptr() orelse return error.RuntimeUnavailable;
 359     const loaded_handle = request.loaded_artifact orelse return error.InvalidArtifact;
 360     const loaded = try state.getLoaded(loaded_handle);
 361     const argument_count = request.buffers.len + request.scalar_arguments.len;
 362     if (argument_count != @as(usize, @intCast(loaded.argument_count))) {
 363         return error.LaunchArgumentMismatch;
 364     }
 365     if (argument_count > runtime_mod.max_kernel_args) return error.LaunchArgumentMismatch;
 366 
 367     var args: [runtime_mod.max_kernel_args]runtime_mod.KernelArg = undefined;
 368     var arg_count: usize = 0;
 369     for (request.buffers) |binding| {
 370         if (binding.handle.backend != .cuda or binding.ownership != .backend) return error.InvalidBuffer;
 371         const buffer = try state.getBuffer(binding.handle);
 372         if (binding.byte_size > buffer.bytes or binding.handle.byte_size != buffer.bytes) {
 373             return error.InvalidBuffer;
 374         }
 375         args[arg_count] = .{ .device_ptr = buffer.ptr };
 376         arg_count += 1;
 377     }
 378     for (request.scalar_arguments) |arg| {
 379         args[arg_count] = lowerScalarArgument(arg);
 380         arg_count += 1;
 381     }
 382 
 383     const stream = if (request.stream) |stream_handle|
 384         try state.getStream(stream_handle)
 385     else
 386         rt.defaultStream() catch |err| return mapStreamError(err);
 387     const signal_event = if (request.signal_event) |event_handle|
 388         try state.getEvent(event_handle)
 389     else
 390         null;
 391     for (request.wait_events) |event_handle| {
 392         _ = try state.getEvent(event_handle);
 393     }
 394 
 395     for (request.wait_events) |event_handle| {
 396         const event = try state.getEvent(event_handle);
 397         stream.waitEvent(event) catch |err| return mapStreamError(err);
 398     }
 399 
 400     loaded.kernel.launchOnStream(
 401         stream,
 402         request.geometry.grid,
 403         request.geometry.threadgroup,
 404         request.geometry.dynamic_shared_memory_bytes,
 405         args[0..arg_count],
 406     ) catch |err| return mapLaunchError(err);
 407 
 408     if (signal_event) |event| {
 409         event.record(stream) catch |err| return mapEventError(err);
 410     }
 411 }
 412 
 413 fn lowerScalarArgument(arg: choir_abi.ScalarArgument) runtime_mod.KernelArg {
 414     return switch (arg) {
 415         .i32 => |value| .{ .i32 = value },
 416         .u32 => |value| .{ .u32 = value },
 417         .i64 => |value| .{ .i64 = value },
 418         .u64 => |value| .{ .u64 = value },
 419         .f32 => |value| .{ .f32 = value },
 420         .f64 => |value| .{ .f64 = value },
 421     };
 422 }
 423 
 424 fn synchronize(ptr: *anyopaque, request: backend.SyncRequest) backend.BackendError!void {
 425     const state: *State = @ptrCast(@alignCast(ptr));
 426     switch (request.scope) {
 427         .default_stream => {
 428             const rt = state.runtime.ptr() orelse return error.RuntimeUnavailable;
 429             const stream = rt.defaultStream() catch |err| return mapStreamError(err);
 430             stream.synchronize() catch |err| return mapStreamError(err);
 431         },
 432         .device => {
 433             const rt = state.runtime.ptr() orelse return error.RuntimeUnavailable;
 434             rt.synchronize() catch |err| return mapSyncError(err);
 435         },
 436         .stream => {
 437             const stream = try state.getStream(request.stream.?);
 438             stream.synchronize() catch |err| return mapStreamError(err);
 439         },
 440         .event => {
 441             const event = try state.getEvent(request.event.?);
 442             event.synchronize() catch |err| return mapEventError(err);
 443         },
 444     }
 445 }
 446 
 447 fn queryEvent(ptr: *anyopaque, request: backend.EventQueryRequest) backend.BackendError!bool {
 448     const state: *State = @ptrCast(@alignCast(ptr));
 449     const event = try state.getEvent(request.event);
 450     return event.query() catch |err| return mapEventError(err);
 451 }
 452 
 453 fn recordEvent(ptr: *anyopaque, request: backend.EventRecordRequest) backend.BackendError!void {
 454     const state: *State = @ptrCast(@alignCast(ptr));
 455     const stream = try state.getStream(request.stream);
 456     const event = try state.getEvent(request.event);
 457     event.record(stream) catch |err| return mapEventError(err);
 458 }
 459 
 460 fn elapsedEventNs(ptr: *anyopaque, request: backend.EventElapsedRequest) backend.BackendError!u64 {
 461     const state: *State = @ptrCast(@alignCast(ptr));
 462     const start = try state.getEvent(request.start);
 463     const end = try state.getEvent(request.end);
 464     const elapsed_ms = end.elapsedMs(start) catch |err| return mapEventError(err);
 465     if (!std.math.isFinite(elapsed_ms) or elapsed_ms < 0) return error.RuntimeUnavailable;
 466     const elapsed_ns = @as(f64, elapsed_ms) * 1_000_000.0;
 467     if (elapsed_ns > @as(f64, @floatFromInt(std.math.maxInt(u64)))) return error.RuntimeUnavailable;
 468     return @intFromFloat(@round(elapsed_ns));
 469 }
 470 
 471 fn destroyObject(ptr: *anyopaque, id: BackendObjectId) void {
 472     const state: *State = @ptrCast(@alignCast(ptr));
 473     if (state.objects.fetchRemove(id)) |entry| {
 474         var object = entry.value;
 475         deinitObject(state.allocator, &object);
 476     }
 477 }
 478 
 479 fn deinitHandle(ptr: *anyopaque, allocator: Allocator) void {
 480     _ = allocator;
 481     const state: *State = @ptrCast(@alignCast(ptr));
 482     state.deinit();
 483 }
 484 
 485 fn deinitObject(allocator: Allocator, object: *Object) void {
 486     switch (object.*) {
 487         .loaded_artifact => |*loaded| {
 488             loaded.module.release();
 489             allocator.free(loaded.entry_name);
 490         },
 491         .buffer => |*buffer| buffer.deinit(),
 492         .stream => |stream| stream.runtime.destroyStream(stream),
 493         .event => |event| event.runtime.destroyEvent(event),
 494     }
 495     object.* = undefined;
 496 }
 497 
 498 fn mapRuntimeInitError(err: runtime_mod.Error) backend.BackendError {
 499     return switch (err) {
 500         error.CudaOutOfMemory => error.OutOfMemory,
 501         error.CudaDriverNotFound,
 502         error.CudaSymbolMissing,
 503         error.CudaNotInitialized,
 504         error.CudaDeinitialized,
 505         error.CudaNoDevice,
 506         error.CudaInvalidDevice,
 507         error.CudaInvalidContext,
 508         error.CudaNotSupported,
 509         error.CudaSharedObjectSymbolNotFound,
 510         error.CudaSharedObjectInitFailed,
 511         => error.RuntimeUnavailable,
 512         else => error.RuntimeUnavailable,
 513     };
 514 }
 515 
 516 fn mapArtifactError(err: runtime_mod.Error) backend.BackendError {
 517     return switch (err) {
 518         error.CudaOutOfMemory => error.OutOfMemory,
 519         error.CudaInvalidPtx,
 520         error.CudaUnsupportedPtxVersion,
 521         error.CudaJitCompilerNotFound,
 522         error.CudaInvalidImage,
 523         error.CudaNoBinaryForGpu,
 524         error.CudaInvalidValue,
 525         error.CudaInvalidHandle,
 526         => error.InvalidArtifact,
 527         error.CudaDriverNotFound,
 528         error.CudaSymbolMissing,
 529         error.CudaNotInitialized,
 530         error.CudaDeinitialized,
 531         error.CudaNoDevice,
 532         error.CudaInvalidDevice,
 533         error.CudaInvalidContext,
 534         error.CudaNotSupported,
 535         error.CudaSharedObjectSymbolNotFound,
 536         error.CudaSharedObjectInitFailed,
 537         => error.RuntimeUnavailable,
 538         error.CudaIllegalAddress => error.DeviceLost,
 539         else => error.LaunchFailed,
 540     };
 541 }
 542 
 543 fn mapAllocationError(err: runtime_mod.Error) backend.BackendError {
 544     return switch (err) {
 545         error.CudaOutOfMemory => error.OutOfMemory,
 546         error.CudaInvalidValue, error.CudaInvalidHandle => error.InvalidBuffer,
 547         error.CudaIllegalAddress => error.DeviceLost,
 548         error.CudaDriverNotFound,
 549         error.CudaSymbolMissing,
 550         error.CudaNotInitialized,
 551         error.CudaDeinitialized,
 552         error.CudaNoDevice,
 553         error.CudaInvalidDevice,
 554         error.CudaInvalidContext,
 555         => error.RuntimeUnavailable,
 556         else => error.LaunchFailed,
 557     };
 558 }
 559 
 560 fn mapLaunchError(err: runtime_mod.Error) backend.BackendError {
 561     return switch (err) {
 562         error.CudaOutOfMemory => error.OutOfMemory,
 563         error.CudaInvalidValue, error.CudaLaunchOutOfResources => error.LaunchArgumentMismatch,
 564         error.CudaIllegalAddress => error.DeviceLost,
 565         error.CudaLaunchFailed,
 566         error.CudaLaunchTimeout,
 567         error.CudaLaunchIncompatibleTexturing,
 568         error.CudaStreamPoisoned,
 569         => error.LaunchFailed,
 570         error.CudaDriverNotFound,
 571         error.CudaSymbolMissing,
 572         error.CudaNotInitialized,
 573         error.CudaDeinitialized,
 574         error.CudaNoDevice,
 575         error.CudaInvalidDevice,
 576         error.CudaInvalidContext,
 577         => error.RuntimeUnavailable,
 578         else => error.LaunchFailed,
 579     };
 580 }
 581 
 582 fn mapStreamError(err: runtime_mod.Error) backend.BackendError {
 583     return switch (err) {
 584         error.CudaOutOfMemory => error.OutOfMemory,
 585         error.CudaInvalidValue, error.CudaInvalidHandle => error.InvalidStream,
 586         error.CudaIllegalAddress => error.DeviceLost,
 587         error.CudaLaunchFailed,
 588         error.CudaLaunchTimeout,
 589         error.CudaLaunchIncompatibleTexturing,
 590         error.CudaStreamPoisoned,
 591         => error.LaunchFailed,
 592         error.CudaDriverNotFound,
 593         error.CudaSymbolMissing,
 594         error.CudaNotInitialized,
 595         error.CudaDeinitialized,
 596         error.CudaNoDevice,
 597         error.CudaInvalidDevice,
 598         error.CudaInvalidContext,
 599         => error.RuntimeUnavailable,
 600         else => error.RuntimeUnavailable,
 601     };
 602 }
 603 
 604 fn mapEventError(err: runtime_mod.Error) backend.BackendError {
 605     return switch (err) {
 606         error.CudaOutOfMemory => error.OutOfMemory,
 607         error.CudaInvalidValue, error.CudaInvalidHandle => error.InvalidEvent,
 608         error.CudaIllegalAddress => error.DeviceLost,
 609         error.CudaLaunchFailed,
 610         error.CudaLaunchTimeout,
 611         error.CudaLaunchIncompatibleTexturing,
 612         error.CudaStreamPoisoned,
 613         => error.LaunchFailed,
 614         error.CudaDriverNotFound,
 615         error.CudaSymbolMissing,
 616         error.CudaNotInitialized,
 617         error.CudaDeinitialized,
 618         error.CudaNoDevice,
 619         error.CudaInvalidDevice,
 620         error.CudaInvalidContext,
 621         => error.RuntimeUnavailable,
 622         else => error.RuntimeUnavailable,
 623     };
 624 }
 625 
 626 fn mapSyncError(err: runtime_mod.Error) backend.BackendError {
 627     return switch (err) {
 628         error.CudaOutOfMemory => error.OutOfMemory,
 629         error.CudaIllegalAddress => error.DeviceLost,
 630         error.CudaLaunchFailed,
 631         error.CudaLaunchTimeout,
 632         error.CudaLaunchIncompatibleTexturing,
 633         error.CudaStreamPoisoned,
 634         => error.LaunchFailed,
 635         error.CudaDriverNotFound,
 636         error.CudaSymbolMissing,
 637         error.CudaNotInitialized,
 638         error.CudaDeinitialized,
 639         error.CudaNoDevice,
 640         error.CudaInvalidDevice,
 641         error.CudaInvalidContext,
 642         => error.RuntimeUnavailable,
 643         else => error.RuntimeUnavailable,
 644     };
 645 }
 646 
 647 const vtable = backend.BackendVTable{
 648     .query_capabilities = queryCapabilities,
 649     .create_artifact = createArtifact,
 650     .load_artifact = loadArtifact,
 651     .allocate_buffer = allocateBuffer,
 652     .create_stream = createStream,
 653     .create_event = createEvent,
 654     .write_buffer = writeBuffer,
 655     .fill_buffer = fillBuffer,
 656     .read_buffer = readBuffer,
 657     .launch = launch,
 658     .synchronize = synchronize,
 659     .query_event = queryEvent,
 660     .record_event = recordEvent,
 661     .elapsed_event_ns = elapsedEventNs,
 662     .destroy_object = destroyObject,
 663     .deinit = deinitHandle,
 664 };
 665 
 666 pub const testing = if (builtin.is_test) struct {
 667     pub const FakeSnapshot: type = runtime_mod.testing.FakeSnapshot;
 668 
 669     pub fn resetFakeRuntime() void {
 670         runtime_mod.testing.resetFake();
 671     }
 672 
 673     pub fn fakeRuntime(allocator: Allocator) Runtime {
 674         return runtime_mod.testing.fakeRuntime(allocator);
 675     }
 676 
 677     pub fn snapshot() FakeSnapshot {
 678         return runtime_mod.testing.snapshot();
 679     }
 680 
 681     pub fn setFakeEventElapsedMs(ms: f32) void {
 682         runtime_mod.testing.setFakeEventElapsedMs(ms);
 683     }
 684 
 685     pub fn stateWithRuntime(allocator: Allocator, runtime: *Runtime) State {
 686         return State.initWithRuntime(allocator, runtime);
 687     }
 688 } else struct {};
 689 
 690 test "cuda contract reports capabilities without a live runtime" {
 691     var state = State.init(std.testing.allocator);
 692     defer state.deinit();
 693     const handle = state.handle();
 694 
 695     const caps = try handle.queryCapabilities();
 696     try std.testing.expectEqual(backend.BackendKind.cuda, caps.identity.backend);
 697     try std.testing.expectEqual(backend.DeviceFamily.nvidia_cuda, caps.identity.family);
 698     try std.testing.expect(caps.supportsDType(.i1));
 699     try std.testing.expect(caps.supportsDType(.f32));
 700     try std.testing.expect(caps.supportsDType(.i32));
 701     try std.testing.expect(caps.supportsDType(.f16));
 702     try std.testing.expect(caps.supportsDType(.bf16));
 703     try std.testing.expect(caps.supportsDType(.u32));
 704     try std.testing.expect(caps.supportsDType(.u64));
 705     try std.testing.expect(caps.supportsDType(.f64));
 706     try std.testing.expect(caps.supportsDType(.i64));
 707     try std.testing.expect(caps.features.atomic_i32);
 708     try std.testing.expect(caps.features.atomic_u32);
 709     try std.testing.expect(caps.features.atomic_index);
 710     try std.testing.expect(caps.features.atomic_f32_add_device);
 711     try std.testing.expect(caps.features.atomic_f32_add_shared);
 712     try std.testing.expect(caps.features.dynamic_shared_memory);
 713     try std.testing.expect(!caps.features.tensor_cores);
 714     try std.testing.expect(caps.supportsArtifactFormat(.cuda_ptx));
 715     try std.testing.expect(!caps.runtime.driver_loaded);
 716     try std.testing.expect(caps.runtime.streams);
 717 }
 718 
 719 test "cuda contract reports tensor cores for fake ampere runtime" {
 720     var runtime = testing.fakeRuntime(std.testing.allocator);
 721     defer runtime.module_cache.deinit(&runtime);
 722 
 723     var state = testing.stateWithRuntime(std.testing.allocator, &runtime);
 724     defer state.deinit();
 725 
 726     const caps = try state.handle().queryCapabilities();
 727     try std.testing.expectEqualStrings("fake cuda", caps.identity.name);
 728     try std.testing.expect(caps.features.tensor_cores);
 729     try std.testing.expectEqual(@as(u32, 100 * 1024), caps.memory.shared_memory_per_threadgroup_bytes);
 730     try std.testing.expectEqual(@as(u32, 100 * 1024), caps.threadgroup.shared_memory_bytes);
 731 }
 732 
 733 test "cuda contract records unsupported and unavailable phases" {
 734     var state = State.init(std.testing.allocator);
 735     defer state.deinit();
 736     const handle = state.handle();
 737 
 738     try std.testing.expectError(error.UnsupportedOperation, handle.createArtifact(.{
 739         .kernel_name = "add",
 740         .requested_format = .cuda_ptx,
 741     }));
 742     try std.testing.expectError(error.RuntimeUnavailable, handle.allocateBuffer(.{
 743         .byte_size = 16,
 744         .alignment = 256,
 745     }));
 746     try std.testing.expectError(error.RuntimeUnavailable, handle.createStream(.{}));
 747     try std.testing.expectError(error.RuntimeUnavailable, handle.createEvent(.{}));
 748     try std.testing.expectError(error.RuntimeUnavailable, handle.synchronize(.{
 749         .scope = .device,
 750     }));
 751     try std.testing.expectError(error.InvalidEvent, handle.queryEvent(.{
 752         .event = .{
 753             .id = 1,
 754             .backend = .cuda,
 755         },
 756     }));
 757     try std.testing.expectError(error.InvalidEvent, handle.elapsedEventNs(.{
 758         .start = .{
 759             .id = 1,
 760             .backend = .cuda,
 761         },
 762         .end = .{
 763             .id = 2,
 764             .backend = .cuda,
 765         },
 766     }));
 767 
 768     var artifact = try backend.KernelArtifact.init(std.testing.allocator, .{
 769         .backend = .cuda,
 770         .format = .cuda_ptx,
 771         .entry_name = "main",
 772         .argument_count = 0,
 773     });
 774     defer artifact.deinit();
 775     artifact.setBorrowedText("// ptx");
 776 
 777     try std.testing.expectError(error.RuntimeUnavailable, handle.loadArtifact(&artifact));
 778 }
 779 
 780 test "cuda contract lowers event elapsed time to nanoseconds" {
 781     testing.resetFakeRuntime();
 782     testing.setFakeEventElapsedMs(2.5);
 783     var runtime = testing.fakeRuntime(std.testing.allocator);
 784     defer runtime.live_streams.deinit(runtime.allocator);
 785     var state = testing.stateWithRuntime(std.testing.allocator, &runtime);
 786     defer state.deinit();
 787     const handle = state.handle();
 788 
 789     const start = try handle.createEvent(.{});
 790     const end = try handle.createEvent(.{});
 791 
 792     try std.testing.expectEqual(@as(u64, 2_500_000), try handle.elapsedEventNs(.{
 793         .start = start,
 794         .end = end,
 795     }));
 796     try std.testing.expectEqual(@as(u32, 1), testing.snapshot().event_elapsed_calls);
 797 }
 798 
 799 test "cuda contract creates PTX artifact before runtime access" {
 800     var state = State.init(std.testing.allocator);
 801     defer state.deinit();
 802     const handle = state.handle();
 803 
 804     var artifact = try handle.createArtifact(.{
 805         .kernel_name = "accy_choir_contract_add",
 806         .requested_format = .cuda_ptx,
 807         .argument_count = 4,
 808         .required_dtypes = backend.DTypeSet.init(&.{.f32}),
 809         .diagnostic_id = "choir/kernel/0",
 810         .payload = .{ .text = ".version 7.0\n.visible .entry accy_choir_contract_add() { ret; }\n" },
 811     });
 812     defer artifact.deinit();
 813 
 814     try std.testing.expectEqual(backend.BackendKind.cuda, artifact.backend);
 815     try std.testing.expectEqual(backend.ArtifactFormat.cuda_ptx, artifact.format);
 816     try std.testing.expectEqualStrings("accy_choir_contract_add", artifact.entry_name);
 817     try std.testing.expectEqual(@as(u32, 4), artifact.argument_count);
 818     try std.testing.expectEqualStrings("choir/kernel/0", artifact.diagnostic_id.?);
 819     try std.testing.expect(std.mem.indexOf(u8, artifact.payload.text, ".visible .entry accy_choir_contract_add") != null);
 820     try std.testing.expectError(error.RuntimeUnavailable, handle.loadArtifact(&artifact));
 821 }
 822 
 823 fn initRuntimeOrSkip(
 824     allocator: Allocator,
 825     device_ordinal: i32,
 826 ) (driver_mod.Error || error{SkipZigTest})!Runtime {
 827     if (!build_options.gpu_tests) return error.SkipZigTest;
 828     if (!driver_mod.platformSupported()) return error.SkipZigTest;
 829 
 830     return Runtime.init(allocator, device_ordinal) catch |err| switch (err) {
 831         error.CudaDriverNotFound,
 832         error.CudaNoDevice,
 833         error.CudaInvalidDevice,
 834         => return error.SkipZigTest,
 835         else => return err,
 836     };
 837 }
 838 
 839 test "cuda contract validates artifact format before runtime access" {
 840     var state = State.init(std.testing.allocator);
 841     defer state.deinit();
 842     const handle = state.handle();
 843 
 844     var artifact = try backend.KernelArtifact.init(std.testing.allocator, .{
 845         .backend = .cuda,
 846         .format = .cuda_cubin,
 847         .entry_name = "main",
 848         .argument_count = 0,
 849     });
 850     defer artifact.deinit();
 851     try artifact.setOwnedBytes(&.{ 0xca, 0xfe });
 852 
 853     try std.testing.expectError(error.UnsupportedArtifactFormat, handle.loadArtifact(&artifact));
 854 }
 855 
 856 test "cuda destroyObject releases loaded artifact modules after final shared reference" {
 857     runtime_mod.testing.resetFake();
 858     var rt = runtime_mod.testing.fakeRuntime(std.testing.allocator);
 859     defer rt.module_cache.deinit(&rt);
 860 
 861     var state = State.initWithRuntime(std.testing.allocator, &rt);
 862     defer state.deinit();
 863     const handle = state.handle();
 864 
 865     var artifact = try backend.KernelArtifact.init(std.testing.allocator, .{
 866         .backend = .cuda,
 867         .format = .cuda_ptx,
 868         .entry_name = "shared",
 869         .argument_count = 0,
 870     });
 871     defer artifact.deinit();
 872     artifact.setBorrowedText(".version 6.0\n.visible .entry shared() { ret; }\n");
 873 
 874     const first = try handle.loadArtifact(&artifact);
 875     const second = try handle.loadArtifact(&artifact);
 876     var snapshot = runtime_mod.testing.snapshot();
 877     try std.testing.expectEqual(@as(u32, 1), snapshot.module_loads);
 878     try std.testing.expectEqual(@as(u32, 1), snapshot.module_load_data_ex_calls);
 879     try std.testing.expectEqual(@as(u32, 0), snapshot.module_unloads);
 880 
 881     handle.destroyObject(first.id);
 882     snapshot = runtime_mod.testing.snapshot();
 883     try std.testing.expectEqual(@as(u32, 0), snapshot.module_unloads);
 884 
 885     handle.destroyObject(second.id);
 886     snapshot = runtime_mod.testing.snapshot();
 887     try std.testing.expectEqual(@as(u32, 1), snapshot.module_unloads);
 888 }
 889 
 890 test "cuda destroyObject releases device buffers before state deinit" {
 891     runtime_mod.testing.resetFake();
 892     var rt = runtime_mod.testing.fakeRuntime(std.testing.allocator);
 893     defer rt.module_cache.deinit(&rt);
 894 
 895     var state = State.initWithRuntime(std.testing.allocator, &rt);
 896     defer state.deinit();
 897     const handle = state.handle();
 898 
 899     const buffer = try handle.allocateBuffer(.{
 900         .byte_size = 256,
 901         .alignment = 256,
 902         .dtype = .f32,
 903         .element_count = 64,
 904     });
 905     var snapshot = runtime_mod.testing.snapshot();
 906     try std.testing.expectEqual(@as(u32, 1), snapshot.mem_allocs);
 907     try std.testing.expectEqual(@as(u32, 0), snapshot.mem_frees);
 908 
 909     handle.destroyObject(buffer.id);
 910     snapshot = runtime_mod.testing.snapshot();
 911     try std.testing.expectEqual(@as(u32, 1), snapshot.mem_frees);
 912 
 913     handle.destroyObject(buffer.id);
 914     snapshot = runtime_mod.testing.snapshot();
 915     try std.testing.expectEqual(@as(u32, 1), snapshot.mem_frees);
 916 }
 917 
 918 test "cuda contract live stream event launch can be queried" {
 919     var rt = try initRuntimeOrSkip(std.testing.allocator, 0);
 920     defer rt.deinit();
 921 
 922     var state = State.initWithRuntime(std.testing.allocator, &rt);
 923     defer state.deinit();
 924     const handle = state.handle();
 925 
 926     const element_count: usize = 256;
 927     var host_in: [element_count]f32 = undefined;
 928     for (&host_in, 0..) |*x, i| x.* = @as(f32, @floatFromInt(i)) * 0.125;
 929 
 930     var host_out: [element_count]f32 = @as([element_count]f32, @splat(0));
 931 
 932     var artifact = try backend.KernelArtifact.init(std.testing.allocator, .{
 933         .backend = .cuda,
 934         .format = .cuda_ptx,
 935         .entry_name = "add_one",
 936         .argument_count = 1,
 937     });
 938     defer artifact.deinit();
 939     artifact.setBorrowedText(
 940         \\.version 6.0
 941         \\.target sm_52
 942         \\.address_size 64
 943         \\
 944         \\.visible .entry add_one(
 945         \\    .param .u64 in_ptr
 946         \\)
 947         \\{
 948         \\    .reg .b32    %r<4>;
 949         \\    .reg .f32    %f<3>;
 950         \\    .reg .b64    %rd<5>;
 951         \\
 952         \\    ld.param.u64       %rd1, [in_ptr];
 953         \\
 954         \\    mov.u32            %r1, %ntid.x;
 955         \\    mov.u32            %r2, %ctaid.x;
 956         \\    mov.u32            %r3, %tid.x;
 957         \\    mad.lo.s32         %r1, %r1, %r2, %r3;
 958         \\
 959         \\    cvta.to.global.u64 %rd2, %rd1;
 960         \\    mul.wide.s32       %rd3, %r1, 4;
 961         \\    add.s64            %rd4, %rd2, %rd3;
 962         \\
 963         \\    ld.global.f32      %f1, [%rd4];
 964         \\    add.f32            %f2, %f1, 0f3F800000;
 965         \\    st.global.f32      [%rd4], %f2;
 966         \\
 967         \\    ret;
 968         \\}
 969     );
 970 
 971     const loaded = try handle.loadArtifact(&artifact);
 972     const buffer = try handle.allocateBuffer(.{
 973         .byte_size = element_count * @sizeOf(f32),
 974         .alignment = 256,
 975         .dtype = .f32,
 976         .element_count = element_count,
 977     });
 978     try handle.writeBuffer(.{
 979         .handle = buffer,
 980         .bytes = std.mem.sliceAsBytes(host_in[0..]),
 981     });
 982 
 983     const stream = try handle.createStream(.{});
 984     const event = try handle.createEvent(.{});
 985     const binding = backend.BufferBinding{
 986         .handle = buffer,
 987         .access = .read_write,
 988         .ownership = buffer.ownership,
 989         .byte_size = buffer.byte_size,
 990     };
 991 
 992     try handle.launch(.{
 993         .artifact = &artifact,
 994         .loaded_artifact = loaded,
 995         .buffers = &.{binding},
 996         .geometry = .{
 997             .grid = .{ 1, 1, 1 },
 998             .threadgroup = .{ @intCast(element_count), 1, 1 },
 999         },
1000         .stream = stream,
1001         .signal_event = event,
1002     });
1003 
1004     _ = try handle.queryEvent(.{ .event = event });
1005     try handle.synchronize(.{
1006         .scope = .event,
1007         .event = event,
1008     });
1009     try std.testing.expect(try handle.queryEvent(.{ .event = event }));
1010 
1011     try handle.readBuffer(.{
1012         .handle = buffer,
1013         .bytes = std.mem.sliceAsBytes(host_out[0..]),
1014     });
1015 
1016     for (host_in, host_out) |in, out| {
1017         try std.testing.expectApproxEqAbs(in + 1.0, out, 1e-6);
1018     }
1019 }