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 }