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 }