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 }