lib/accy/src/kernel/test.zig
daab053ee43316e1809a84551d573ddd1e5bf3d2
1 const std = @import("std");
2 const gpu = @import("gpu");
3 const choir_abi = @import("choir_abi");
4 const kernel = @import("root.zig");
5
6 const build_options = @import("build_options");
7 const artifact_product = @import("../artifact/root.zig");
8 const exec_product = @import("../executable/root.zig");
9 const logical_mod = @import("logical/root.zig");
10 const artifact = kernel.artifact;
11 const builder = kernel.builder;
12 const domain = kernel.domain;
13 const oracle = kernel.oracle;
14 const plan = kernel.plan;
15 const schedule = kernel.schedule;
16 const typed = kernel.typed;
17 const vector = kernel.vector;
18 const view = kernel.view;
19 const Param = kernel.Param;
20 const Buffer = kernel.Buffer;
21 const Type = kernel.Type;
22 const Kernel = kernel.Kernel;
23 const RawBuilder = kernel.RawBuilder;
24 const Builder = kernel.Builder;
25 const MmaShape = kernel.MmaShape;
26 const Value = kernel.Value;
27 const Block = kernel.Block;
28 const If = kernel.If;
29 const For = kernel.For;
30 const ForScope = kernel.ForScope;
31 const While = kernel.While;
32 const WhileScope = kernel.WhileScope;
33 const InsertionScope = kernel.InsertionScope;
34 const VerificationDiagnostic = kernel.VerificationDiagnostic;
35 const Plan = kernel.Plan;
36 const PlanOptions = kernel.PlanOptions;
37 const PlanError = kernel.PlanError;
38 const PlanVersion = kernel.PlanVersion;
39 const BackendHandle = kernel.BackendHandle;
40 const KernelArtifact = kernel.KernelArtifact;
41 const KernelCallArtifactOptions = kernel.KernelCallArtifactOptions;
42 const OwnedKernelCallArtifact = kernel.OwnedKernelCallArtifact;
43 const FragmentCompilerOptions = exec_product.FragmentCompilerOptions;
44 const ArtifactKernelSource = artifact_product.KernelSource;
45 const Argument = kernel.Argument;
46 const Scalar = kernel.Scalar;
47 const ExecutionDiagnostic = kernel.ExecutionDiagnostic;
48 const Graph = kernel.Graph;
49 const Program = kernel.Program;
50 const Family = kernel.Family;
51 const program_product_name = kernel.program_product_name;
52 const interpret = kernel.interpret;
53 const library = kernel.library;
54 const logical = kernel.logical;
55 const DomainAxis = kernel.DomainAxis;
56 const Domain2D = kernel.Domain2D;
57 const Domain3D = kernel.Domain3D;
58 const Index1D = kernel.Index1D;
59 const VectorIndex1D = kernel.VectorIndex1D;
60 const Index2D = kernel.Index2D;
61 const Index3D = kernel.Index3D;
62 const BufferView = kernel.BufferView;
63 const KeyBufferView = kernel.KeyBufferView;
64 const TypedValue = kernel.TypedValue;
65 const KeyValue = kernel.KeyValue;
66 const TypedVec2 = kernel.TypedVec2;
67 const TypedVec3 = kernel.TypedVec3;
68 const Vec2 = kernel.Vec2;
69 const Vec3 = kernel.Vec3;
70 const Guard = kernel.Guard;
71 const Schedule = kernel.Schedule;
72 const ScheduleError = kernel.ScheduleError;
73 const Launch = kernel.Launch;
74 const Axis = kernel.Axis;
75 const AxisId = kernel.AxisId;
76 const AxisStep = kernel.AxisStep;
77 const BindTarget = kernel.BindTarget;
78 const ScheduleSnapshot = kernel.ScheduleSnapshot;
79 const ScheduleSnapshotVersion = kernel.ScheduleSnapshotVersion;
80 const Split = kernel.Split;
81 const Step = kernel.Step;
82 const scalar = kernel.scalar;
83 const buffer = kernel.buffer;
84 const dynamicBuffer = kernel.dynamicBuffer;
85 const domainAxis = kernel.domainAxis;
86 const createPlan = kernel.createPlan;
87 const createArtifactJob = kernel.createArtifactJob;
88 const createArtifactPlan = kernel.createArtifactPlan;
89 const createKernelCallArtifact = kernel.createKernelCallArtifact;
90 const createBackendArtifactFromKernelCallEntry = kernel.createBackendArtifactFromKernelCallEntry;
91 const kernelCallEntryLaunchGeometry = kernel.kernelCallEntryLaunchGeometry;
92 const compileFragment = kernel.compileFragment;
93 const createKernelArtifact = kernel.createKernelArtifact;
94 const argumentBuffer = kernel.argumentBuffer;
95 const argumentBool = kernel.argumentBool;
96 const argumentI32 = kernel.argumentI32;
97 const argumentU32 = kernel.argumentU32;
98 const argumentI64 = kernel.argumentI64;
99 const argumentF32 = kernel.argumentF32;
100 const argumentF64 = kernel.argumentF64;
101 const Compare = kernel.Compare;
102 const Dimension = kernel.Dimension;
103 const AddressSpace = kernel.AddressSpace;
104 const MemoryOrder = kernel.MemoryOrder;
105 const Scope = kernel.Scope;
106 const WarpOpKind = kernel.WarpOpKind;
107 const dynamic_shared_byte_offset_attr_name =
108 @import("choir").dialects.gpu.attr_names.dynamic_shared_byte_offset;
109
110 test {
111 _ = @import("dsl/test.zig");
112 _ = @import("library/test.zig");
113 _ = @import("logical/test.zig");
114 _ = @import("program/test.zig");
115 _ = @import("call.zig");
116 @import("test_discovery").discover(kernel);
117 }
118
119 fn executableCopyEach(inner: anytype, index: Index1D, each_args: anytype) !void {
120 const value = try each_args.param(.src).load(inner, index);
121 try each_args.param(.dst).store(inner, value, index);
122 }
123
124 fn executableCopyBody(k: anytype, args: anytype) !void {
125 _ = try k.forEach1D("i", 4, 4, args, executableCopyEach);
126 }
127
128 const ExecutableCopy = Program(.{
129 .name = "kernel_dsl_executable_copy_i32",
130 .parameters = .{
131 .src = dynamicBuffer(.i32),
132 .dst = dynamicBuffer(.i32),
133 },
134 .body = executableCopyBody,
135 });
136
137 fn kernelCallAddEach(inner: anytype, index: Index1D, each_args: anytype) !void {
138 const lhs = try each_args.param(.lhs).load(inner, index);
139 const rhs = try each_args.param(.rhs).load(inner, index);
140 const sum = try lhs.add(inner, rhs);
141 try each_args.param(.dst).store(inner, sum, index);
142 }
143
144 fn kernelCallAddBody(k: anytype, args: anytype) !void {
145 _ = try k.forEach1D("i", 8, 8, args, kernelCallAddEach);
146 }
147
148 const KernelCallAdd = Program(.{
149 .name = "kernel_call_add_f32",
150 .parameters = .{
151 .dst = dynamicBuffer(.f32),
152 .lhs = dynamicBuffer(.f32),
153 .rhs = dynamicBuffer(.f32),
154 },
155 .body = kernelCallAddBody,
156 });
157
158 fn logicalKernelCallAddEach(inner: anytype, index: Index1D, each_args: anytype) !void {
159 const lhs = try each_args.param(.lhs).load(inner, index);
160 const rhs = try each_args.param(.rhs).load(inner, index);
161 const sum = try lhs.add(inner, rhs);
162 try each_args.param(.dst).store(inner, sum, index);
163 }
164
165 fn logicalKernelCallAddBody(k: anytype, args: anytype) !void {
166 _ = try k.forEach1D("i", 8, args, logicalKernelCallAddEach);
167 }
168
169 const LogicalKernelCallAdd = logical_mod.Program(.{
170 .name = "logical_kernel_call_add_f32",
171 .parameters = .{
172 .dst = dynamicBuffer(.f32),
173 .lhs = dynamicBuffer(.f32),
174 .rhs = dynamicBuffer(.f32),
175 },
176 .body = logicalKernelCallAddBody,
177 }).withSchedule(logical_mod.schedule.threadBlocks(.{ .x = 4 }));
178
179 fn metalGeneratedScaleEach(inner: anytype, index: Index1D, each_args: anytype) !void {
180 const value = try each_args.param(.src).load(inner, index);
181 const scaled = try value.mul(inner, each_args.param(.scale));
182 try each_args.param(.dst).store(inner, scaled, index);
183 }
184
185 fn metalGeneratedScaleBody(k: anytype, args: anytype) !void {
186 _ = try k.forEach1D("i", 5, 4, args, metalGeneratedScaleEach);
187 }
188
189 const MetalGeneratedScale = Program(.{
190 .name = "kernel_metal_generated_scale_f32",
191 .parameters = .{
192 .src = dynamicBuffer(.f32),
193 .dst = dynamicBuffer(.f32),
194 .scale = scalar(.f32),
195 },
196 .body = metalGeneratedScaleBody,
197 });
198
199 test "kernel Program creates executable fragment through BackendHandle" {
200 const allocator = std.testing.allocator;
201 var state = gpu.recording.BackendState{
202 .allocator = allocator,
203 .kind = .vulkan,
204 .format = .vulkan_spirv,
205 };
206
207 var graph = try ExecutableCopy.build(allocator, ExecutableCopy.Limits.testing);
208 defer graph.deinit();
209 var schedule_snapshot = try graph.scheduleSnapshot(allocator);
210 defer schedule_snapshot.deinit(allocator);
211 var checked_plan = try graph.createCheckedPlan(allocator, .{});
212 defer checked_plan.deinit();
213 const body_fingerprint = checked_plan.body_fingerprint;
214 const options = FragmentCompilerOptions{ .authored_kernel_diagnostic_id = "choir/kernel/program-executable" };
215 const compiled = try compileFragment(
216 allocator,
217 state.handle(),
218 &graph,
219 options,
220 );
221 var fragment = try @import("../executable/root.zig").loadFragment(allocator, state.handle(), compiled, options);
222 defer fragment.deinit();
223
224 try std.testing.expectEqual(@as(usize, 1), state.create_count);
225 try std.testing.expect(state.last_create_had_payload);
226 try std.testing.expectEqual(@as(u32, 2), state.last_create_argument_count);
227 try std.testing.expectEqual(gpu.DTypeSet.init(&.{.i32}).bits, state.last_create_required_dtype_bits);
228 try std.testing.expectEqual(@as(usize, 1), state.load_count);
229 try std.testing.expectEqual(@as(usize, 1), fragment.kernelCount());
230
231 const summary = try fragment.kernelSummary(0);
232 try std.testing.expectEqual(ArtifactKernelSource.choir_kernel, summary.source);
233 try std.testing.expectEqualStrings("kernel_dsl_executable_copy_i32", summary.entry_name);
234 try std.testing.expectEqual(gpu.ArtifactFormat.vulkan_spirv, summary.artifact_format);
235 try std.testing.expectEqual(@as(u32, 2), summary.compile_argument_count);
236 try std.testing.expectEqual(gpu.DTypeSet.init(&.{.i32}).bits, summary.compile_required_dtype_bits);
237 try std.testing.expect(std.meta.eql(choir_abi.Features{}, summary.compile_required_features));
238 try std.testing.expect(std.meta.eql(choir_abi.SubgroupRequirements{}, summary.compile_required_subgroup));
239 try std.testing.expectEqual(.words_u32, summary.compile_payload);
240 try std.testing.expectEqual(.authored, summary.compile_launch);
241 try std.testing.expect(summary.fixed_threadgroup);
242 try std.testing.expectEqual(@as(u64, 0), summary.output_layout_fingerprint);
243 try std.testing.expectEqual(@as(u64, 0), summary.input_layout_fingerprint);
244 try std.testing.expectEqual(@as(usize, 1), schedule_snapshot.allAxes().len);
245 try std.testing.expectEqualStrings("i", schedule_snapshot.allAxes()[0].name);
246 try std.testing.expectEqual(BindTarget.thread_x, schedule_snapshot.allAxes()[0].bind.?);
247 try std.testing.expectEqual(checked_plan.schedule_fingerprint, schedule_snapshot.fingerprint());
248 try std.testing.expectEqual(body_fingerprint, try graph.bodyFingerprint(allocator));
249
250 const buffers = [_]gpu.BufferBinding{
251 testBufferBinding(31, .vulkan, 16),
252 testBufferBinding(32, .vulkan, 16),
253 };
254 try fragment.launchKernelWithArguments(0, &buffers, &.{}, .{});
255
256 try std.testing.expectEqual(@as(usize, 1), state.launch_count);
257 try std.testing.expectEqual(state.last_loaded_id.?, state.last_launch_loaded_id.?);
258 try std.testing.expectEqual(@as(usize, 2), state.last_launch_buffer_count);
259 try std.testing.expectEqual(@as(usize, 0), state.last_launch_scalar_count);
260 try std.testing.expectEqual(@as(u32, 1), state.last_launch_grid[0]);
261 try std.testing.expectEqual(@as(u32, 4), state.last_launch_threadgroup[0]);
262 try std.testing.expectEqual(@as(gpu.BackendObjectId, 31), state.last_buffer_ids[0]);
263 try std.testing.expectEqual(@as(gpu.BackendObjectId, 32), state.last_buffer_ids[1]);
264 }
265
266 test "kernel Program creates registry-ready kernel_call artifact" {
267 const allocator = std.testing.allocator;
268 var state = gpu.recording.BackendState{
269 .allocator = allocator,
270 .kind = .cuda,
271 .format = .cuda_ptx,
272 };
273
274 var call_artifact = try KernelCallAdd.createKernelCallArtifact(allocator, KernelCallAdd.Limits.testing, state.handle(), .{
275 .target = "accy.custom.add",
276 });
277 defer call_artifact.deinit();
278
279 const registry = call_artifact.registry();
280 const entry = registry.find("accy.custom.add", 1, .cuda_ptx) orelse return error.TestExpectedKernelCallArtifact;
281 try std.testing.expectEqualStrings("kernel_call_add_f32", entry.entry_name);
282 try std.testing.expectEqual(@as(u32, 3), entry.argument_count);
283 try std.testing.expectEqual(gpu.DTypeSet.init(&.{.f32}).bits, entry.required_dtypes.bits);
284 switch (entry.launch) {
285 .fixed => |geometry| {
286 try std.testing.expectEqual(@as(u32, 1), geometry.grid[0]);
287 try std.testing.expectEqual(@as(u32, 8), geometry.threadgroup[0]);
288 },
289 else => return error.TestExpectedFixedLaunch,
290 }
291 try std.testing.expectEqual(.none, entry.element_count_argument);
292 try std.testing.expectEqual(@as(usize, 0), entry.static_arguments.len);
293 switch (entry.payload) {
294 .text => |text| try std.testing.expect(std.mem.indexOf(u8, text, "kernel_call_add_f32") != null),
295 else => return error.TestExpectedPtxPayload,
296 }
297 }
298
299 test "logical Program creates scheduled registry-ready kernel_call artifact" {
300 const allocator = std.testing.allocator;
301 var state = gpu.recording.BackendState{
302 .allocator = allocator,
303 .kind = .cuda,
304 .format = .cuda_ptx,
305 };
306
307 var call_artifact = try LogicalKernelCallAdd.createKernelCallArtifact(allocator, LogicalKernelCallAdd.Limits.testing, state.handle(), .{
308 .target = "accy.custom.logical_add",
309 });
310 defer call_artifact.deinit();
311
312 const entry = call_artifact.entry();
313 try std.testing.expectEqualStrings("logical_kernel_call_add_f32", entry.entry_name);
314 try std.testing.expectEqual(@as(u32, 3), entry.argument_count);
315 switch (entry.launch) {
316 .fixed => |geometry| {
317 try std.testing.expectEqual(@as(u32, 2), geometry.grid[0]);
318 try std.testing.expectEqual(@as(u32, 4), geometry.threadgroup[0]);
319 },
320 else => return error.TestExpectedFixedLaunch,
321 }
322 }
323
324 test "kernel Program creates standalone kernel artifact through BackendHandle" {
325 const allocator = std.testing.allocator;
326 var state = gpu.recording.BackendState{
327 .allocator = allocator,
328 .kind = .vulkan,
329 .format = .vulkan_spirv,
330 };
331
332 var compiled = try ExecutableCopy.createKernelArtifact(
333 allocator,
334 ExecutableCopy.Limits.testing,
335 state.handle(),
336 .{ .authored_kernel_diagnostic_id = "choir/kernel/program-artifact" },
337 );
338 defer compiled.deinit();
339
340 try std.testing.expectEqual(@as(usize, 1), state.create_count);
341 try std.testing.expect(state.last_create_had_payload);
342 try std.testing.expectEqual(@as(u32, 2), state.last_create_argument_count);
343 try std.testing.expectEqual(gpu.DTypeSet.init(&.{.i32}).bits, state.last_create_required_dtype_bits);
344 try std.testing.expectEqual(@as(usize, 0), state.load_count);
345 try std.testing.expectEqual(gpu.BackendKind.vulkan, compiled.backend);
346 try std.testing.expectEqual(gpu.ArtifactFormat.vulkan_spirv, compiled.format);
347 try std.testing.expectEqual(@as(u32, 2), compiled.argument_count);
348 try std.testing.expectEqualStrings("kernel_dsl_executable_copy_i32", compiled.entry_name);
349 try std.testing.expectEqualStrings("choir/kernel/program-artifact", compiled.diagnostic_id.?);
350 }
351
352 test "kernel Program launches generated Choir artifact on live Metal" {
353 const allocator = std.testing.allocator;
354
355 var state = try initMetalStateOrSkip(allocator);
356 defer state.deinit();
357 const handle = state.handle();
358
359 const options = FragmentCompilerOptions{
360 .artifact_format = .metal_msl,
361 .authored_kernel_diagnostic_id = "choir/kernel/generated-metal-scale-live",
362 };
363 const compiled = try MetalGeneratedScale.compileFragment(
364 allocator,
365 MetalGeneratedScale.Limits.testing,
366 handle,
367 options,
368 );
369 var fragment = try @import("../executable/root.zig").loadFragment(allocator, handle, compiled, options);
370 defer fragment.deinit();
371
372 try std.testing.expectEqual(@as(usize, 1), fragment.kernelCount());
373 const summary = try fragment.kernelSummary(0);
374 try std.testing.expectEqual(ArtifactKernelSource.choir_kernel, summary.source);
375 try std.testing.expectEqual(gpu.ArtifactFormat.metal_msl, summary.artifact_format);
376 try std.testing.expectEqual(@as(u32, 3), summary.compile_argument_count);
377 try std.testing.expectEqual(gpu.DTypeSet.init(&.{.f32}).bits, summary.compile_required_dtype_bits);
378 try std.testing.expect(summary.fixed_threadgroup);
379 try std.testing.expectEqual(@as(u32, 2), summary.launch_geometry.grid[0]);
380 try std.testing.expectEqual(@as(u32, 4), summary.launch_geometry.threadgroup[0]);
381
382 var input = [_]f32{ 1.25, -2.0, 0.5, 4.0, -3.5 };
383 var output = @as([input.len]f32, @splat(0));
384 const scale: f32 = -3.0;
385
386 const input_buffer = try handle.allocateBuffer(.{
387 .byte_size = @sizeOf(@TypeOf(input)),
388 .alignment = 256,
389 .dtype = .f32,
390 .element_count = input.len,
391 });
392 const output_buffer = try handle.allocateBuffer(.{
393 .byte_size = @sizeOf(@TypeOf(output)),
394 .alignment = 256,
395 .dtype = .f32,
396 .element_count = output.len,
397 });
398
399 try handle.writeBuffer(.{
400 .handle = input_buffer,
401 .bytes = std.mem.sliceAsBytes(input[0..]),
402 });
403 try handle.writeBuffer(.{
404 .handle = output_buffer,
405 .bytes = std.mem.sliceAsBytes(output[0..]),
406 });
407
408 const stream = try handle.createStream(.{});
409 const event = try handle.createEvent(.{});
410 const bindings = [_]gpu.BufferBinding{
411 liveBufferBinding(input_buffer, .read_only),
412 liveBufferBinding(output_buffer, .write_only),
413 };
414 try fragment.launchKernelWithArguments(
415 0,
416 &bindings,
417 &.{.{ .f32 = scale }},
418 .{
419 .stream = stream,
420 .signal_event = event,
421 },
422 );
423
424 try handle.synchronize(.{ .scope = .event, .event = event });
425 try std.testing.expect(try handle.queryEvent(.{ .event = event }));
426 try handle.readBuffer(.{
427 .handle = output_buffer,
428 .bytes = std.mem.sliceAsBytes(output[0..]),
429 });
430
431 for (input, output) |source, observed| {
432 try std.testing.expectApproxEqAbs(source * scale, observed, 1e-6);
433 }
434 }
435
436 fn initMetalStateOrSkip(allocator: std.mem.Allocator) !gpu.metal.State {
437 if (!build_options.metal_tests) return error.SkipZigTest;
438 return gpu.metal.State.initDevice(allocator) catch |err| switch (err) {
439 error.RuntimeUnavailable => return error.SkipZigTest,
440 else => return err,
441 };
442 }
443
444 fn testBufferBinding(
445 id: gpu.BackendObjectId,
446 kind: gpu.BackendKind,
447 bytes: usize,
448 ) gpu.BufferBinding {
449 return .{
450 .handle = .{
451 .id = id,
452 .backend = kind,
453 .byte_size = bytes,
454 .ownership = .backend,
455 },
456 .access = .read_write,
457 .ownership = .backend,
458 .byte_size = bytes,
459 };
460 }
461
462 fn liveBufferBinding(handle: gpu.BufferHandle, access: gpu.BufferAccess) gpu.BufferBinding {
463 return .{
464 .handle = handle,
465 .access = access,
466 .ownership = handle.ownership,
467 .byte_size = handle.byte_size,
468 };
469 }
470
471 test "Schedule launch drives Choir-backed CPU kernel oracle" {
472 var schedule_plan = try Schedule.init(std.testing.allocator, Schedule.Limits.testing);
473 defer schedule_plan.deinit(std.testing.allocator);
474
475 const i = try schedule_plan.addAxis("i", 4);
476 try schedule_plan.bind(i, .thread_x);
477
478 var b = try builder.Builder.init(std.testing.allocator, builder.Builder.Limits.testing, "copy_i32_scheduled", &.{
479 dynamicBuffer(.i32),
480 dynamicBuffer(.i32),
481 }, .{});
482 errdefer b.deinit();
483
484 const src = b.argument(0);
485 const dst = b.argument(1);
486 const index = try b.globalId(.x);
487 const value = try b.load(src, index);
488 try b.store(value, dst, index);
489 try b.return_();
490
491 var kernel_graph = try b.finish();
492 defer kernel_graph.deinit();
493
494 var input = [_]i32{ 4, 3, 2, 1 };
495 var output = [_]i32{ 0, 0, 0, 0 };
496 const launch: Launch = try schedule_plan.launch();
497 try oracle.run(std.testing.allocator, &kernel_graph, &.{
498 oracle.argumentBuffer(i32, input[0..]),
499 oracle.argumentBuffer(i32, output[0..]),
500 }, launch);
501
502 try std.testing.expectEqualSlices(i32, input[0..], output[0..]);
503 }
504
505 test "accy kernel declaration coverage" {
506 std.testing.refAllDecls(kernel);
507 }