lib/accy/src/kernel/logical/test.zig

daab053ee43316e1809a84551d573ddd1e5bf3d2

  1 const std = @import("std");
  2 const logical = @import("root.zig");
  3 
  4 const kernel = @import("../root.zig");
  5 const Axis = logical.Axis;
  6 const ActivationKind = logical.ActivationKind;
  7 const ActivationSelectionRequest = logical.ActivationSelectionRequest;
  8 const Builder = logical.Builder;
  9 const Domain2D = logical.Domain2D;
 10 const Domain3D = logical.Domain3D;
 11 const EinsumKernelKind = logical.EinsumKernelKind;
 12 const EinsumOperand = logical.EinsumOperand;
 13 const EinsumSchedule = logical.EinsumSchedule;
 14 const EinsumSelectionRequest = logical.EinsumSelectionRequest;
 15 const Family = logical.Family;
 16 const FusedKernelKind = logical.FusedKernelKind;
 17 const FusedLinalgEpilogue = logical.FusedLinalgEpilogue;
 18 const FusedMatrixProductSchedule = logical.FusedMatrixProductSchedule;
 19 const FusedMatrixProductSelectionRequest = logical.FusedMatrixProductSelectionRequest;
 20 const FusedMatrixVectorProductSchedule = logical.FusedMatrixVectorProductSchedule;
 21 const FusedMatrixVectorProductSelectionRequest = logical.FusedMatrixVectorProductSelectionRequest;
 22 const FusedRowNormalizationKind = logical.FusedRowNormalizationKind;
 23 const FusedRowNormalizationSchedule = logical.FusedRowNormalizationSchedule;
 24 const FusedRowNormalizationSelectionRequest = logical.FusedRowNormalizationSelectionRequest;
 25 const FusedSelectionRequest = logical.FusedSelectionRequest;
 26 const FusedVectorKind = logical.FusedVectorKind;
 27 const FusedVectorSelectionRequest = logical.FusedVectorSelectionRequest;
 28 const Program = logical.Program;
 29 const RowNormalizationKind = logical.RowNormalizationKind;
 30 const RowNormalizationParameterization = logical.RowNormalizationParameterization;
 31 const RowNormalizationSchedule = logical.RowNormalizationSchedule;
 32 const RowNormalizationSelectionRequest = logical.RowNormalizationSelectionRequest;
 33 const SelectedActivationKernel = logical.SelectedActivationKernel;
 34 const SelectedEinsumKernel = logical.SelectedEinsumKernel;
 35 const OwnedSelectedEinsumKernel = logical.OwnedSelectedEinsumKernel;
 36 const SelectedFusedKernel = logical.SelectedFusedKernel;
 37 const SelectedRowNormalizationKernel = logical.SelectedRowNormalizationKernel;
 38 const axis = logical.axis;
 39 const schedule = logical.schedule;
 40 const selectActivationCatalog = logical.selectActivationCatalog;
 41 const selectEinsumCatalog = logical.selectEinsumCatalog;
 42 const selectOwnedEinsumCatalog = logical.selectOwnedEinsumCatalog;
 43 const selectFusedCatalog = logical.selectFusedCatalog;
 44 const selectRowNormalizationCatalog = logical.selectRowNormalizationCatalog;
 45 const wrap = logical.wrap;
 46 
 47 test {
 48     _ = @import("family.zig");
 49     _ = @import("selection/test.zig");
 50     @import("test_discovery").discover(logical);
 51 }
 52 
 53 fn scaleEach(inner: anytype, index: kernel.Index1D, each_args: anytype) !void {
 54     const value = try each_args.param(.src).load(inner, index);
 55     const scaled = try value.mul(inner, each_args.param(.scale));
 56     try each_args.param(.dst).store(inner, scaled, index);
 57 }
 58 
 59 fn scaleBody(k: anytype, args: anytype) !void {
 60     _ = try k.forEach1D("i", 5, args, scaleEach);
 61 }
 62 
 63 const Scale = Program(.{
 64     .name = "kernel_logical_scale_f32",
 65     .parameters = .{
 66         .src = kernel.dynamicBuffer(.f32),
 67         .dst = kernel.dynamicBuffer(.f32),
 68         .scale = kernel.scalar(.f32),
 69     },
 70     .body = scaleBody,
 71 });
 72 
 73 const OneThreadScale = Scale.withSchedule(schedule.threadBlocks(.{ .x = 1 }));
 74 
 75 const AxisThreads = struct {
 76     name: []const u8,
 77     x: u32,
 78 
 79     pub fn index1D(self: *@This(), ctx: anytype) !kernel.Index1D {
 80         if (std.mem.eql(u8, ctx.axis.name, self.name)) {
 81             return ctx.useThreads(self.x);
 82         }
 83         return ctx.default();
 84     }
 85 };
 86 
 87 const DelegatedStackScale = Scale.withSchedule(schedule.stack(.{
 88     schedule.with(AxisThreads{ .name = "j", .x = 2 }),
 89     schedule.useThreads(.{ .x = 1 }),
 90 }));
 91 
 92 const OverrideStackScale = Scale.withSchedule(schedule.stack(.{
 93     schedule.with(AxisThreads{ .name = "i", .x = 2 }),
 94     schedule.useThreads(.{ .x = 1 }),
 95 }));
 96 
 97 fn windowSumStep(fold_inner: anytype, offset: kernel.Value, acc: kernel.Value, fold_ctx: anytype) !kernel.Value {
 98     const input_index = try fold_inner.add(fold_ctx.base, offset);
 99     const value = try fold_ctx.src.load(fold_inner, input_index);
100     return fold_inner.add(acc, value.raw());
101 }
102 
103 fn windowSumEach(inner: anytype, index: kernel.Index1D, each_args: anytype) !void {
104     const zero = try inner.constantFloat(.f32, 0.0);
105     const width = try inner.constantIndex(3);
106     const base = try inner.mul(index.index, width);
107 
108     const sum = try inner.foldRange(0, 3, 1, zero, .{
109         .src = each_args.param(.src),
110         .base = base,
111     }, windowSumStep);
112 
113     try each_args.param(.dst).store(inner, sum, index);
114 }
115 
116 fn windowSumBody(k: anytype, args: anytype) !void {
117     _ = try k.forEach1D("window", 4, args, windowSumEach);
118 }
119 
120 const WindowSum = Program(.{
121     .name = "kernel_logical_window_sum_f32",
122     .parameters = .{
123         .src = kernel.dynamicBuffer(.f32),
124         .dst = kernel.dynamicBuffer(.f32),
125     },
126     .body = windowSumBody,
127 });
128 
129 const DropMul = struct {
130     pub fn mul(_: *@This(), ctx: anytype) !kernel.Value {
131         return ctx.lhs;
132     }
133 };
134 
135 const MulPresence = struct {
136     has_mul: bool = false,
137 
138     pub const Result = struct {
139         graph: kernel.Graph,
140         has_mul: bool,
141 
142         pub fn deinit(self: *@This()) void {
143             self.graph.deinit();
144             self.* = undefined;
145         }
146     };
147 
148     pub fn mul(self: *@This(), ctx: anytype) !kernel.Value {
149         self.has_mul = true;
150         return ctx.default();
151     }
152 
153     pub fn finish(self: *@This(), ctx: anytype) !Result {
154         return .{
155             .graph = try ctx.default(),
156             .has_mul = self.has_mul,
157         };
158     }
159 };
160 
161 const DropMulWhenPresent = struct {
162     enabled: bool,
163 
164     pub fn mul(self: *@This(), ctx: anytype) !kernel.Value {
165         if (self.enabled) return ctx.lhs;
166         return ctx.default();
167     }
168 };
169 
170 fn dropMulFromPresence(analysis: *const MulPresence.Result) DropMulWhenPresent {
171     return .{ .enabled = analysis.has_mul };
172 }
173 
174 const DropFold = struct {
175     pub fn fold(_: *@This(), ctx: anytype) !kernel.Value {
176         return ctx.initial;
177     }
178 };
179 
180 const Unscaled = Scale.transform(DropMul{});
181 const OneThreadUnscaled = OneThreadScale.transform(DropMul{});
182 const OverrideStackUnscaled = OverrideStackScale.transform(DropMul{});
183 const LateScheduledUnscaled = Unscaled.withSchedule(schedule.stack(.{
184     schedule.with(AxisThreads{ .name = "i", .x = 2 }),
185     schedule.useThreads(.{ .x = 1 }),
186 }));
187 const RescheduledOneThreadUnscaled = OneThreadUnscaled.withSchedule(schedule.threadBlocks(.{ .x = 5 }));
188 const AnalyzedUnscaled = Scale.analyze(MulPresence{}).transform(dropMulFromPresence);
189 const LateScheduledAnalyzedUnscaled = AnalyzedUnscaled.withSchedule(schedule.stack(.{
190     schedule.with(AxisThreads{ .name = "i", .x = 2 }),
191     schedule.useThreads(.{ .x = 1 }),
192 }));
193 const RescheduledAnalyzedUnscaled = OneThreadScale.analyze(MulPresence{}).transform(dropMulFromPresence).withSchedule(schedule.threadBlocks(.{ .x = 5 }));
194 const EmptyWindowSum = WindowSum.transform(DropFold{});
195 
196 test "logical Program compiles to a checked scheduled kernel graph" {
197     try std.testing.expectEqual(@as(usize, 0), Scale.arg(.src));
198     try std.testing.expectEqual(@as(usize, 1), Scale.arg(.dst));
199     try std.testing.expectEqual(@as(usize, 2), Scale.arg(.scale));
200 
201     try std.testing.expectEqual(@as(usize, 0), OneThreadScale.arg(.src));
202 
203     var graph = try Scale.build(std.testing.allocator, Scale.Limits.testing);
204     defer graph.deinit();
205     try graph.verify();
206 
207     const launch_value = try graph.launch();
208     try std.testing.expectEqual(@as(u32, 1), launch_value.grid[0]);
209     try std.testing.expectEqual(@as(u32, 5), launch_value.block[0]);
210 
211     var plan = try Scale.createCheckedPlan(std.testing.allocator, Scale.Limits.testing, .{});
212     defer plan.deinit();
213     try std.testing.expectEqual(@as(u32, 3), plan.argument_count);
214     try std.testing.expectEqualStrings("kernel_logical_scale_f32", plan.entry_name);
215 
216     var input = [_]f32{ 1.0, -2.0, 3.5, 4.0, -0.25 };
217     var output = [_]f32{ 0, 0, 0, 0, 0 };
218     try Scale.runCpu(std.testing.allocator, Scale.Limits.testing, &.{
219         kernel.argumentBuffer(f32, input[0..]),
220         kernel.argumentBuffer(f32, output[0..]),
221         kernel.argumentF32(2.0),
222     });
223     try std.testing.expectEqualSlices(f32, &.{ 2.0, -4.0, 7.0, 8.0, -0.5 }, output[0..]);
224 }
225 
226 test "logical Program schedules through a first-class policy" {
227     const default_launch = try Scale.launch(std.testing.allocator, Scale.Limits.testing);
228     const one_thread_launch = try OneThreadScale.launch(std.testing.allocator, OneThreadScale.Limits.testing);
229     const delegated_launch = try DelegatedStackScale.launch(std.testing.allocator, DelegatedStackScale.Limits.testing);
230     const override_launch = try OverrideStackScale.launch(std.testing.allocator, OverrideStackScale.Limits.testing);
231     var override_snapshot = try OverrideStackScale.scheduleSnapshot(std.testing.allocator, OverrideStackScale.Limits.testing);
232     defer override_snapshot.deinit(std.testing.allocator);
233 
234     try std.testing.expectEqual(@as(u32, 1), default_launch.grid[0]);
235     try std.testing.expectEqual(@as(u32, 5), default_launch.block[0]);
236     try std.testing.expectEqual(@as(u32, 5), one_thread_launch.grid[0]);
237     try std.testing.expectEqual(@as(u32, 1), one_thread_launch.block[0]);
238     try std.testing.expectEqual(@as(u32, 5), delegated_launch.grid[0]);
239     try std.testing.expectEqual(@as(u32, 1), delegated_launch.block[0]);
240     try std.testing.expectEqual(@as(u32, 3), override_launch.grid[0]);
241     try std.testing.expectEqual(@as(u32, 2), override_launch.block[0]);
242     try std.testing.expectEqual(@as(usize, 2), override_snapshot.allAxes().len);
243     try std.testing.expectEqual(@as(usize, 4), override_snapshot.allSteps().len);
244     try std.testing.expectEqualStrings("i_tile", override_snapshot.allAxes()[0].name);
245     try std.testing.expectEqualStrings("i_lane", override_snapshot.allAxes()[1].name);
246     try std.testing.expectEqual(kernel.BindTarget.block_x, override_snapshot.allAxes()[0].bind.?);
247     try std.testing.expectEqual(kernel.BindTarget.thread_x, override_snapshot.allAxes()[1].bind.?);
248 
249     var input = [_]f32{ 1.0, -2.0, 3.5, 4.0, -0.25 };
250     var output = [_]f32{ 0, 0, 0, 0, 0 };
251     try OverrideStackScale.runCpu(std.testing.allocator, OverrideStackScale.Limits.testing, &.{
252         kernel.argumentBuffer(f32, input[0..]),
253         kernel.argumentBuffer(f32, output[0..]),
254         kernel.argumentF32(3.0),
255     });
256     try std.testing.expectEqualSlices(f32, &.{ 3.0, -6.0, 10.5, 12.0, -0.75 }, output[0..]);
257 }
258 
259 test "logical Program folds scalar ranges inside scheduled bodies" {
260     const launch_value = try WindowSum.launch(std.testing.allocator, WindowSum.Limits.testing);
261     try std.testing.expectEqual(@as(u32, 1), launch_value.grid[0]);
262     try std.testing.expectEqual(@as(u32, 4), launch_value.block[0]);
263 
264     var input = [_]f32{ 1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0, 10.0, 11.0, 12.0 };
265     var output = [_]f32{ 0, 0, 0, 0 };
266     try WindowSum.runCpu(std.testing.allocator, WindowSum.Limits.testing, &.{
267         kernel.argumentBuffer(f32, input[0..]),
268         kernel.argumentBuffer(f32, output[0..]),
269     });
270     try std.testing.expectEqualSlices(f32, &.{ 6.0, 15.0, 24.0, 33.0 }, output[0..]);
271 
272     var empty_output = [_]f32{ -1.0, -1.0, -1.0, -1.0 };
273     try EmptyWindowSum.runCpu(std.testing.allocator, EmptyWindowSum.Limits.testing, &.{
274         kernel.argumentBuffer(f32, input[0..]),
275         kernel.argumentBuffer(f32, empty_output[0..]),
276     });
277     try std.testing.expectEqualSlices(f32, &.{ 0.0, 0.0, 0.0, 0.0 }, empty_output[0..]);
278 }
279 
280 test "logical Program composes with kernel interpreter transforms" {
281     const late_scheduled_launch = try LateScheduledUnscaled.launch(std.testing.allocator, LateScheduledUnscaled.Limits.testing);
282     const rescheduled_launch = try RescheduledOneThreadUnscaled.launch(std.testing.allocator, RescheduledOneThreadUnscaled.Limits.testing);
283     const analyzed_late_scheduled_launch = try LateScheduledAnalyzedUnscaled.launch(std.testing.allocator, LateScheduledAnalyzedUnscaled.Limits.testing);
284     const analyzed_rescheduled_launch = try RescheduledAnalyzedUnscaled.launch(std.testing.allocator, RescheduledAnalyzedUnscaled.Limits.testing);
285     var late_scheduled_snapshot = try LateScheduledUnscaled.scheduleSnapshot(std.testing.allocator, LateScheduledUnscaled.Limits.testing);
286     defer late_scheduled_snapshot.deinit(std.testing.allocator);
287     var analyzed_late_scheduled_snapshot = try LateScheduledAnalyzedUnscaled.scheduleSnapshot(std.testing.allocator, LateScheduledAnalyzedUnscaled.Limits.testing);
288     defer analyzed_late_scheduled_snapshot.deinit(std.testing.allocator);
289     var analyzed_late_scheduled_plan = try LateScheduledAnalyzedUnscaled.createCheckedPlan(std.testing.allocator, LateScheduledAnalyzedUnscaled.Limits.testing, .{});
290     defer analyzed_late_scheduled_plan.deinit();
291     try std.testing.expectEqual(@as(u32, 3), late_scheduled_launch.grid[0]);
292     try std.testing.expectEqual(@as(u32, 2), late_scheduled_launch.block[0]);
293     try std.testing.expectEqual(@as(u32, 1), rescheduled_launch.grid[0]);
294     try std.testing.expectEqual(@as(u32, 5), rescheduled_launch.block[0]);
295     try std.testing.expectEqual(@as(u32, 3), analyzed_late_scheduled_launch.grid[0]);
296     try std.testing.expectEqual(@as(u32, 2), analyzed_late_scheduled_launch.block[0]);
297     try std.testing.expectEqual(@as(u32, 1), analyzed_rescheduled_launch.grid[0]);
298     try std.testing.expectEqual(@as(u32, 5), analyzed_rescheduled_launch.block[0]);
299     try std.testing.expectEqual(late_scheduled_snapshot.fingerprint(), analyzed_late_scheduled_snapshot.fingerprint());
300     try std.testing.expectEqual(analyzed_late_scheduled_plan.schedule_fingerprint, analyzed_late_scheduled_snapshot.fingerprint());
301     try std.testing.expectEqual(@as(usize, 4), analyzed_late_scheduled_snapshot.allSteps().len);
302     try std.testing.expectEqualStrings("i", analyzed_late_scheduled_snapshot.allSteps()[0].axis.name);
303     try std.testing.expectEqual(@as(u64, 2), analyzed_late_scheduled_snapshot.allSteps()[1].tile.factor);
304 
305     var input = [_]f32{ 1.0, -2.0, 3.5, 4.0, -0.25 };
306     var output = [_]f32{ 0, 0, 0, 0, 0 };
307     try Unscaled.runCpu(std.testing.allocator, Unscaled.Limits.testing, &.{
308         kernel.argumentBuffer(f32, input[0..]),
309         kernel.argumentBuffer(f32, output[0..]),
310         kernel.argumentF32(2.0),
311     });
312     try std.testing.expectEqualSlices(f32, input[0..], output[0..]);
313 
314     var one_thread_output = [_]f32{ 0, 0, 0, 0, 0 };
315     try OneThreadUnscaled.runCpu(std.testing.allocator, OneThreadUnscaled.Limits.testing, &.{
316         kernel.argumentBuffer(f32, input[0..]),
317         kernel.argumentBuffer(f32, one_thread_output[0..]),
318         kernel.argumentF32(2.0),
319     });
320     try std.testing.expectEqualSlices(f32, input[0..], one_thread_output[0..]);
321 
322     var late_scheduled_output = [_]f32{ 0, 0, 0, 0, 0 };
323     try LateScheduledUnscaled.runCpu(std.testing.allocator, LateScheduledUnscaled.Limits.testing, &.{
324         kernel.argumentBuffer(f32, input[0..]),
325         kernel.argumentBuffer(f32, late_scheduled_output[0..]),
326         kernel.argumentF32(2.0),
327     });
328     try std.testing.expectEqualSlices(f32, input[0..], late_scheduled_output[0..]);
329 
330     var rescheduled_output = [_]f32{ 0, 0, 0, 0, 0 };
331     try RescheduledOneThreadUnscaled.runCpu(std.testing.allocator, RescheduledOneThreadUnscaled.Limits.testing, &.{
332         kernel.argumentBuffer(f32, input[0..]),
333         kernel.argumentBuffer(f32, rescheduled_output[0..]),
334         kernel.argumentF32(2.0),
335     });
336     try std.testing.expectEqualSlices(f32, input[0..], rescheduled_output[0..]);
337 
338     var analyzed_late_scheduled_output = [_]f32{ 0, 0, 0, 0, 0 };
339     try LateScheduledAnalyzedUnscaled.runCpu(std.testing.allocator, LateScheduledAnalyzedUnscaled.Limits.testing, &.{
340         kernel.argumentBuffer(f32, input[0..]),
341         kernel.argumentBuffer(f32, analyzed_late_scheduled_output[0..]),
342         kernel.argumentF32(2.0),
343     });
344     try std.testing.expectEqualSlices(f32, input[0..], analyzed_late_scheduled_output[0..]);
345 
346     var analyzed_rescheduled_output = [_]f32{ 0, 0, 0, 0, 0 };
347     try RescheduledAnalyzedUnscaled.runCpu(std.testing.allocator, RescheduledAnalyzedUnscaled.Limits.testing, &.{
348         kernel.argumentBuffer(f32, input[0..]),
349         kernel.argumentBuffer(f32, analyzed_rescheduled_output[0..]),
350         kernel.argumentF32(2.0),
351     });
352     try std.testing.expectEqualSlices(f32, input[0..], analyzed_rescheduled_output[0..]);
353 
354     var stacked_output = [_]f32{ 0, 0, 0, 0, 0 };
355     try OverrideStackUnscaled.runCpu(std.testing.allocator, OverrideStackUnscaled.Limits.testing, &.{
356         kernel.argumentBuffer(f32, input[0..]),
357         kernel.argumentBuffer(f32, stacked_output[0..]),
358         kernel.argumentF32(2.0),
359     });
360     try std.testing.expectEqualSlices(f32, input[0..], stacked_output[0..]);
361 }
362 
363 test "accy kernel logical declaration coverage" {
364     std.testing.refAllDecls(logical);
365 }