lib/accy/src/kernel/logical/selection/fused.zig
daab053ee43316e1809a84551d573ddd1e5bf3d2
1 const std = @import("std");
2 const choir_abi = @import("choir_abi");
3
4 const accy = @import("../../../root.zig");
5 const library = @import("../../library/root.zig");
6
7 const DType = choir_abi.DType;
8
9 pub const FusedVectorKind = library.FusedVectorKind;
10 pub const FusedRowNormalizationKind = library.FusedRowNormalizationKind;
11 pub const FusedMatrixProductSchedule = library.MatrixProductSchedule;
12 pub const FusedMatrixVectorProductSchedule = library.MatrixVectorProductSchedule;
13 pub const FusedRowNormalizationSchedule = library.RowNormalizationSchedule;
14
15 pub const FusedVectorSelectionRequest = struct {
16 dtype: DType,
17 kind: FusedVectorKind,
18 extent: u64,
19 };
20
21 pub const FusedLinalgEpilogue = library.FusedLinalgEpilogue;
22
23 pub const FusedMatrixProductSelectionRequest = struct {
24 dtype: DType,
25 epilogue: FusedLinalgEpilogue,
26 lhs_indices: []const u8,
27 rhs_indices: []const u8,
28 bias_indices: []const u8,
29 output_indices: []const u8,
30 lhs_dims: []const i64,
31 rhs_dims: []const i64,
32 bias_dims: []const i64,
33 output_dims: []const i64,
34 schedule: ?FusedMatrixProductSchedule = null,
35 };
36
37 pub const FusedMatrixVectorProductSelectionRequest = struct {
38 dtype: DType,
39 epilogue: FusedLinalgEpilogue,
40 matrix_indices: []const u8,
41 vector_indices: []const u8,
42 bias_indices: []const u8,
43 output_indices: []const u8,
44 matrix_dims: []const i64,
45 vector_dims: []const i64,
46 bias_dims: []const i64,
47 output_dims: []const i64,
48 schedule: ?FusedMatrixVectorProductSchedule = null,
49 };
50
51 pub const FusedRowNormalizationSelectionRequest = struct {
52 dtype: DType,
53 kind: FusedRowNormalizationKind,
54 rows: u64,
55 cols: u64,
56 schedule: ?FusedRowNormalizationSchedule = null,
57 };
58
59 pub const FusedSelectionRequest = union(enum) {
60 vector: FusedVectorSelectionRequest,
61 matrix_product: FusedMatrixProductSelectionRequest,
62 matrix_vector_product: FusedMatrixVectorProductSelectionRequest,
63 row_normalization: FusedRowNormalizationSelectionRequest,
64 };
65
66 pub const FusedKernelKind = union(enum) {
67 vector: FusedVectorKind,
68 matrix_product: FusedLinalgEpilogue,
69 matrix_vector_product: FusedLinalgEpilogue,
70 row_normalization: FusedRowNormalizationKind,
71 };
72
73 pub const SelectedFusedKernel = struct {
74 kind: FusedKernelKind,
75 descriptor: library.CatalogDescriptor,
76 };
77
78 pub fn selectCatalog(request: FusedSelectionRequest) ?SelectedFusedKernel {
79 return switch (request) {
80 .vector => |vector| selectVector(vector),
81 .matrix_product => |matrix_product| selectMatrixProduct(matrix_product),
82 .matrix_vector_product => |matrix_vector_product| selectMatrixVectorProduct(matrix_vector_product),
83 .row_normalization => |row_normalization| selectRowNormalization(row_normalization),
84 };
85 }
86
87 fn selectVector(request: FusedVectorSelectionRequest) ?SelectedFusedKernel {
88 const descriptor = library.select(.{ .fused = .{ .vector = .{
89 .dtype = request.dtype,
90 .kind = request.kind,
91 .extent = request.extent,
92 } } }) orelse return null;
93 return .{
94 .kind = .{ .vector = request.kind },
95 .descriptor = descriptor,
96 };
97 }
98
99 fn selectMatrixProduct(request: FusedMatrixProductSelectionRequest) ?SelectedFusedKernel {
100 const descriptor = library.select(.{ .fused = .{ .matrix_product = .{
101 .dtype = request.dtype,
102 .epilogue = request.epilogue,
103 .lhs_indices = request.lhs_indices,
104 .rhs_indices = request.rhs_indices,
105 .bias_indices = request.bias_indices,
106 .output_indices = request.output_indices,
107 .lhs_dims = request.lhs_dims,
108 .rhs_dims = request.rhs_dims,
109 .bias_dims = request.bias_dims,
110 .output_dims = request.output_dims,
111 .schedule = request.schedule,
112 } } }) orelse return null;
113 return .{
114 .kind = .{ .matrix_product = request.epilogue },
115 .descriptor = descriptor,
116 };
117 }
118
119 fn selectMatrixVectorProduct(request: FusedMatrixVectorProductSelectionRequest) ?SelectedFusedKernel {
120 const descriptor = library.select(.{ .fused = .{ .matrix_vector_product = .{
121 .dtype = request.dtype,
122 .epilogue = request.epilogue,
123 .matrix_indices = request.matrix_indices,
124 .vector_indices = request.vector_indices,
125 .bias_indices = request.bias_indices,
126 .output_indices = request.output_indices,
127 .matrix_dims = request.matrix_dims,
128 .vector_dims = request.vector_dims,
129 .bias_dims = request.bias_dims,
130 .output_dims = request.output_dims,
131 .schedule = request.schedule,
132 } } }) orelse return null;
133 return .{
134 .kind = .{ .matrix_vector_product = request.epilogue },
135 .descriptor = descriptor,
136 };
137 }
138
139 fn selectRowNormalization(request: FusedRowNormalizationSelectionRequest) ?SelectedFusedKernel {
140 const descriptor = library.select(.{ .fused = .{ .row_normalization = .{
141 .dtype = request.dtype,
142 .kind = request.kind,
143 .rows = request.rows,
144 .cols = request.cols,
145 .schedule = request.schedule,
146 } } }) orelse return null;
147 return .{
148 .kind = .{ .row_normalization = request.kind },
149 .descriptor = descriptor,
150 };
151 }
152
153 test "logical fused selection chooses vector bias gelu catalog entry" {
154 const selected = selectCatalog(.{ .vector = .{
155 .dtype = .f32,
156 .kind = .{ .bias_activation = .gelu },
157 .extent = 8,
158 } }) orelse return error.TestExpectedBiasGeluSelection;
159
160 try std.testing.expect(std.meta.eql(FusedKernelKind{ .vector = .{ .bias_activation = .gelu } }, selected.kind));
161 try std.testing.expectEqualStrings(library.fused.BiasGelu8F32.target, selected.descriptor.metadata.target);
162 }
163
164 test "logical fused selection chooses vector bias relu and silu catalog entries" {
165 const relu = selectCatalog(.{ .vector = .{
166 .dtype = .f32,
167 .kind = .{ .bias_activation = .relu },
168 .extent = 8,
169 } }) orelse return error.TestExpectedBiasReluSelection;
170 const silu = selectCatalog(.{ .vector = .{
171 .dtype = .f32,
172 .kind = .{ .bias_activation = .silu },
173 .extent = 8,
174 } }) orelse return error.TestExpectedBiasSiluSelection;
175
176 try std.testing.expect(std.meta.eql(FusedKernelKind{ .vector = .{ .bias_activation = .relu } }, relu.kind));
177 try std.testing.expectEqualStrings(library.fused.BiasRelu8F32.target, relu.descriptor.metadata.target);
178 try std.testing.expect(std.meta.eql(FusedKernelKind{ .vector = .{ .bias_activation = .silu } }, silu.kind));
179 try std.testing.expectEqualStrings(library.fused.BiasSilu8F32.target, silu.descriptor.metadata.target);
180 }
181
182 test "logical fused selection chooses vector swiglu catalog entry" {
183 const selected = selectCatalog(.{ .vector = .{
184 .dtype = .f32,
185 .kind = .{ .gated_activation = .silu },
186 .extent = 8,
187 } }) orelse return error.TestExpectedSwiGluSelection;
188
189 try std.testing.expect(std.meta.eql(FusedKernelKind{ .vector = .{ .gated_activation = .silu } }, selected.kind));
190 try std.testing.expectEqualStrings(library.fused.SwiGlu8F32.target, selected.descriptor.metadata.target);
191 }
192
193 test "logical fused selection chooses vector geglu catalog entry" {
194 const selected = selectCatalog(.{ .vector = .{
195 .dtype = .f32,
196 .kind = .{ .gated_activation = .gelu },
197 .extent = 8,
198 } }) orelse return error.TestExpectedGeGluSelection;
199
200 try std.testing.expect(std.meta.eql(FusedKernelKind{ .vector = .{ .gated_activation = .gelu } }, selected.kind));
201 try std.testing.expectEqualStrings(library.fused.GeGlu8F32.target, selected.descriptor.metadata.target);
202 }
203
204 test "logical fused selection chooses vector reglu catalog entry" {
205 const selected = selectCatalog(.{ .vector = .{
206 .dtype = .f32,
207 .kind = .{ .gated_activation = .relu },
208 .extent = 8,
209 } }) orelse return error.TestExpectedReGluSelection;
210
211 try std.testing.expect(std.meta.eql(FusedKernelKind{ .vector = .{ .gated_activation = .relu } }, selected.kind));
212 try std.testing.expectEqualStrings(library.fused.ReGlu8F32.target, selected.descriptor.metadata.target);
213 }
214
215 test "logical fused selection chooses matrix product bias gelu catalog entry" {
216 const lhs_dims = [_]i64{ 2, 4 };
217 const rhs_dims = [_]i64{ 4, 3 };
218 const bias_dims = [_]i64{3};
219 const output_dims = [_]i64{ 2, 3 };
220
221 const selected = selectCatalog(.{ .matrix_product = .{
222 .dtype = .f32,
223 .epilogue = .{ .bias_activation = .gelu },
224 .lhs_indices = "mk",
225 .rhs_indices = "kn",
226 .bias_indices = "n",
227 .output_indices = "mn",
228 .lhs_dims = &lhs_dims,
229 .rhs_dims = &rhs_dims,
230 .bias_dims = &bias_dims,
231 .output_dims = &output_dims,
232 } }) orelse return error.TestExpectedMatrixProductBiasGeluSelection;
233
234 try std.testing.expect(std.meta.eql(FusedKernelKind{ .matrix_product = .{ .bias_activation = .gelu } }, selected.kind));
235 try std.testing.expectEqualStrings(library.fused.MatrixProductBiasGelu2x3x4F32.target, selected.descriptor.metadata.target);
236 }
237
238 test "logical fused selection chooses schedule-specialized matrix product catalog entry" {
239 const lhs_dims = [_]i64{ 2, 4 };
240 const rhs_dims = [_]i64{ 4, 3 };
241 const bias_dims = [_]i64{3};
242 const output_dims = [_]i64{ 2, 3 };
243
244 const selected = selectCatalog(.{ .matrix_product = .{
245 .dtype = .f32,
246 .epilogue = .{ .bias_activation = .gelu },
247 .lhs_indices = "mk",
248 .rhs_indices = "kn",
249 .bias_indices = "n",
250 .output_indices = "mn",
251 .lhs_dims = &lhs_dims,
252 .rhs_dims = &rhs_dims,
253 .bias_dims = &bias_dims,
254 .output_dims = &output_dims,
255 .schedule = .{ .thread_blocks = .{ .x = 1, .y = 2 } },
256 } }) orelse return error.TestExpectedMatrixProductBiasGeluSelection;
257
258 try std.testing.expect(std.meta.eql(FusedKernelKind{ .matrix_product = .{ .bias_activation = .gelu } }, selected.kind));
259 try std.testing.expectEqualStrings(library.fused.MatrixProductBiasGelu2x3x4ThreadBlocks1x2F32.target, selected.descriptor.metadata.target);
260 try std.testing.expectEqual(@as(u32, 1), selected.descriptor.metadata.specialization.launch.?.threadgroup[0]);
261 try std.testing.expectEqual(@as(u32, 2), selected.descriptor.metadata.specialization.launch.?.threadgroup[1]);
262 }
263
264 test "logical fused selection chooses matrix product bias relu and silu catalog entries" {
265 const lhs_dims = [_]i64{ 2, 4 };
266 const rhs_dims = [_]i64{ 4, 3 };
267 const bias_dims = [_]i64{3};
268 const output_dims = [_]i64{ 2, 3 };
269
270 const relu = selectCatalog(.{ .matrix_product = .{
271 .dtype = .f32,
272 .epilogue = .{ .bias_activation = .relu },
273 .lhs_indices = "mk",
274 .rhs_indices = "kn",
275 .bias_indices = "n",
276 .output_indices = "mn",
277 .lhs_dims = &lhs_dims,
278 .rhs_dims = &rhs_dims,
279 .bias_dims = &bias_dims,
280 .output_dims = &output_dims,
281 } }) orelse return error.TestExpectedMatrixProductBiasReluSelection;
282 const silu = selectCatalog(.{ .matrix_product = .{
283 .dtype = .f32,
284 .epilogue = .{ .bias_activation = .silu },
285 .lhs_indices = "mk",
286 .rhs_indices = "kn",
287 .bias_indices = "n",
288 .output_indices = "mn",
289 .lhs_dims = &lhs_dims,
290 .rhs_dims = &rhs_dims,
291 .bias_dims = &bias_dims,
292 .output_dims = &output_dims,
293 } }) orelse return error.TestExpectedMatrixProductBiasSiluSelection;
294
295 try std.testing.expect(std.meta.eql(FusedKernelKind{ .matrix_product = .{ .bias_activation = .relu } }, relu.kind));
296 try std.testing.expectEqualStrings(library.fused.MatrixProductBiasRelu2x3x4F32.target, relu.descriptor.metadata.target);
297 try std.testing.expect(std.meta.eql(FusedKernelKind{ .matrix_product = .{ .bias_activation = .silu } }, silu.kind));
298 try std.testing.expectEqualStrings(library.fused.MatrixProductBiasSilu2x3x4F32.target, silu.descriptor.metadata.target);
299 }
300
301 test "logical fused selection chooses matrix vector product bias gelu catalog entry" {
302 const matrix_dims = [_]i64{ 4, 8 };
303 const vector_dims = [_]i64{8};
304 const bias_dims = [_]i64{4};
305 const output_dims = [_]i64{4};
306
307 const selected = selectCatalog(.{ .matrix_vector_product = .{
308 .dtype = .f32,
309 .epilogue = .{ .bias_activation = .gelu },
310 .matrix_indices = "mk",
311 .vector_indices = "k",
312 .bias_indices = "m",
313 .output_indices = "m",
314 .matrix_dims = &matrix_dims,
315 .vector_dims = &vector_dims,
316 .bias_dims = &bias_dims,
317 .output_dims = &output_dims,
318 } }) orelse return error.TestExpectedMatrixVectorProductBiasGeluSelection;
319
320 try std.testing.expect(std.meta.eql(FusedKernelKind{ .matrix_vector_product = .{ .bias_activation = .gelu } }, selected.kind));
321 try std.testing.expectEqualStrings(library.fused.MatrixVectorProductBiasGelu4x8F32.target, selected.descriptor.metadata.target);
322 }
323
324 test "logical fused selection chooses matrix vector product bias relu and silu catalog entries" {
325 const matrix_dims = [_]i64{ 4, 8 };
326 const vector_dims = [_]i64{8};
327 const bias_dims = [_]i64{4};
328 const output_dims = [_]i64{4};
329
330 const relu = selectCatalog(.{ .matrix_vector_product = .{
331 .dtype = .f32,
332 .epilogue = .{ .bias_activation = .relu },
333 .matrix_indices = "mk",
334 .vector_indices = "k",
335 .bias_indices = "m",
336 .output_indices = "m",
337 .matrix_dims = &matrix_dims,
338 .vector_dims = &vector_dims,
339 .bias_dims = &bias_dims,
340 .output_dims = &output_dims,
341 } }) orelse return error.TestExpectedMatrixVectorProductBiasReluSelection;
342 const silu = selectCatalog(.{ .matrix_vector_product = .{
343 .dtype = .f32,
344 .epilogue = .{ .bias_activation = .silu },
345 .matrix_indices = "mk",
346 .vector_indices = "k",
347 .bias_indices = "m",
348 .output_indices = "m",
349 .matrix_dims = &matrix_dims,
350 .vector_dims = &vector_dims,
351 .bias_dims = &bias_dims,
352 .output_dims = &output_dims,
353 } }) orelse return error.TestExpectedMatrixVectorProductBiasSiluSelection;
354
355 try std.testing.expect(std.meta.eql(FusedKernelKind{ .matrix_vector_product = .{ .bias_activation = .relu } }, relu.kind));
356 try std.testing.expectEqualStrings(library.fused.MatrixVectorProductBiasRelu4x8F32.target, relu.descriptor.metadata.target);
357 try std.testing.expect(std.meta.eql(FusedKernelKind{ .matrix_vector_product = .{ .bias_activation = .silu } }, silu.kind));
358 try std.testing.expectEqualStrings(library.fused.MatrixVectorProductBiasSilu4x8F32.target, silu.descriptor.metadata.target);
359 }
360
361 test "logical fused selection chooses row residual rmsnorm catalog entry" {
362 const selected = selectCatalog(.{ .row_normalization = .{
363 .dtype = .f32,
364 .kind = .{ .residual = .{ .rmsnorm = .scale } },
365 .rows = 2,
366 .cols = 4,
367 } }) orelse return error.TestExpectedRowResidualRmsNormSelection;
368
369 try std.testing.expect(std.meta.eql(FusedKernelKind{ .row_normalization = .{ .residual = .{ .rmsnorm = .scale } } }, selected.kind));
370 try std.testing.expectEqualStrings(library.normalization.RowResidualRmsNorm2x4F32.target, selected.descriptor.metadata.target);
371 }
372
373 test "logical fused selection rejects unavailable catalog entries" {
374 const lhs_dims = [_]i64{ 2, 4 };
375 const rhs_dims = [_]i64{ 4, 3 };
376 const bias_dims = [_]i64{3};
377 const wrong_bias_dims = [_]i64{4};
378 const output_dims = [_]i64{ 2, 3 };
379 const matrix_dims = [_]i64{ 4, 8 };
380 const vector_dims = [_]i64{8};
381 const matvec_bias_dims = [_]i64{4};
382 const wrong_matvec_bias_dims = [_]i64{5};
383 const matvec_output_dims = [_]i64{4};
384
385 try std.testing.expect(selectCatalog(.{ .vector = .{
386 .dtype = .i32,
387 .kind = .{ .bias_activation = .gelu },
388 .extent = 8,
389 } }) == null);
390 try std.testing.expect(selectCatalog(.{ .vector = .{
391 .dtype = .f32,
392 .kind = .{ .bias_activation = .gelu },
393 .extent = 16,
394 } }) == null);
395 try std.testing.expect(selectCatalog(.{ .vector = .{
396 .dtype = .i32,
397 .kind = .{ .gated_activation = .gelu },
398 .extent = 8,
399 } }) == null);
400 try std.testing.expect(selectCatalog(.{ .vector = .{
401 .dtype = .f32,
402 .kind = .{ .gated_activation = .gelu },
403 .extent = 16,
404 } }) == null);
405 try std.testing.expect(selectCatalog(.{ .vector = .{
406 .dtype = .i32,
407 .kind = .{ .gated_activation = .relu },
408 .extent = 8,
409 } }) == null);
410 try std.testing.expect(selectCatalog(.{ .vector = .{
411 .dtype = .f32,
412 .kind = .{ .gated_activation = .relu },
413 .extent = 16,
414 } }) == null);
415 try std.testing.expect(selectCatalog(.{ .vector = .{
416 .dtype = .i32,
417 .kind = .{ .gated_activation = .silu },
418 .extent = 8,
419 } }) == null);
420 try std.testing.expect(selectCatalog(.{ .vector = .{
421 .dtype = .f32,
422 .kind = .{ .gated_activation = .silu },
423 .extent = 16,
424 } }) == null);
425 try std.testing.expect(selectCatalog(.{ .matrix_product = .{
426 .dtype = .f32,
427 .epilogue = .{ .bias_activation = .gelu },
428 .lhs_indices = "mk",
429 .rhs_indices = "kn",
430 .bias_indices = "n",
431 .output_indices = "mn",
432 .lhs_dims = &lhs_dims,
433 .rhs_dims = &rhs_dims,
434 .bias_dims = &wrong_bias_dims,
435 .output_dims = &output_dims,
436 } }) == null);
437 try std.testing.expect(selectCatalog(.{ .matrix_product = .{
438 .dtype = .f32,
439 .epilogue = .{ .bias_activation = .gelu },
440 .lhs_indices = "mk",
441 .rhs_indices = "kn",
442 .bias_indices = "m",
443 .output_indices = "mn",
444 .lhs_dims = &lhs_dims,
445 .rhs_dims = &rhs_dims,
446 .bias_dims = &bias_dims,
447 .output_dims = &output_dims,
448 } }) == null);
449 try std.testing.expect(selectCatalog(.{ .matrix_product = .{
450 .dtype = .f32,
451 .epilogue = .{ .bias_activation = .gelu },
452 .lhs_indices = "mk",
453 .rhs_indices = "kn",
454 .bias_indices = "n",
455 .output_indices = "mn",
456 .lhs_dims = &lhs_dims,
457 .rhs_dims = &rhs_dims,
458 .bias_dims = &bias_dims,
459 .output_dims = &output_dims,
460 .schedule = .{ .thread_blocks = .{ .x = 8, .y = 2 } },
461 } }) == null);
462 try std.testing.expect(selectCatalog(.{ .matrix_vector_product = .{
463 .dtype = .i32,
464 .epilogue = .{ .bias_activation = .gelu },
465 .matrix_indices = "mk",
466 .vector_indices = "k",
467 .bias_indices = "m",
468 .output_indices = "m",
469 .matrix_dims = &matrix_dims,
470 .vector_dims = &vector_dims,
471 .bias_dims = &matvec_bias_dims,
472 .output_dims = &matvec_output_dims,
473 } }) == null);
474 try std.testing.expect(selectCatalog(.{ .matrix_vector_product = .{
475 .dtype = .f32,
476 .epilogue = .{ .bias_activation = .gelu },
477 .matrix_indices = "mk",
478 .vector_indices = "k",
479 .bias_indices = "k",
480 .output_indices = "m",
481 .matrix_dims = &matrix_dims,
482 .vector_dims = &vector_dims,
483 .bias_dims = &matvec_bias_dims,
484 .output_dims = &matvec_output_dims,
485 } }) == null);
486 try std.testing.expect(selectCatalog(.{ .matrix_vector_product = .{
487 .dtype = .f32,
488 .epilogue = .{ .bias_activation = .gelu },
489 .matrix_indices = "mk",
490 .vector_indices = "k",
491 .bias_indices = "m",
492 .output_indices = "m",
493 .matrix_dims = &matrix_dims,
494 .vector_dims = &vector_dims,
495 .bias_dims = &wrong_matvec_bias_dims,
496 .output_dims = &matvec_output_dims,
497 } }) == null);
498 try std.testing.expect(selectCatalog(.{ .matrix_vector_product = .{
499 .dtype = .f32,
500 .epilogue = .{ .bias_activation = .gelu },
501 .matrix_indices = "mk",
502 .vector_indices = "k",
503 .bias_indices = "m",
504 .output_indices = "m",
505 .matrix_dims = &matrix_dims,
506 .vector_dims = &vector_dims,
507 .bias_dims = &matvec_bias_dims,
508 .output_dims = &matvec_output_dims,
509 .schedule = .{ .thread_blocks = 8 },
510 } }) == null);
511 try std.testing.expect(selectCatalog(.{ .row_normalization = .{
512 .dtype = .i32,
513 .kind = .{ .residual = .{ .rmsnorm = .scale } },
514 .rows = 2,
515 .cols = 4,
516 } }) == null);
517 try std.testing.expect(selectCatalog(.{ .row_normalization = .{
518 .dtype = .f32,
519 .kind = .{ .residual = .{ .rmsnorm = .none } },
520 .rows = 2,
521 .cols = 4,
522 } }) == null);
523 try std.testing.expect(selectCatalog(.{ .row_normalization = .{
524 .dtype = .f32,
525 .kind = .{ .residual = .{ .rmsnorm = .scale } },
526 .rows = 2,
527 .cols = 5,
528 } }) == null);
529 try std.testing.expect(selectCatalog(.{ .row_normalization = .{
530 .dtype = .f32,
531 .kind = .{ .residual = .{ .rmsnorm = .scale } },
532 .rows = 2,
533 .cols = 4,
534 .schedule = .{ .thread_blocks = .{ .x = 8, .y = 2 } },
535 } }) == null);
536 }