lib/accy/src/kernel/library/catalog/test.zig
daab053ee43316e1809a84551d573ddd1e5bf3d2
1 const std = @import("std");
2 const gpu = @import("gpu");
3 const choir_abi = @import("choir_abi");
4 const catalog = @import("root.zig");
5
6 const accy_artifact = @import("../../../artifact/root.zig");
7 const library = @import("../root.zig");
8
9 const attention = library.attention;
10 const compaction = library.compaction;
11 const elementwise = library.elementwise;
12 const entry = library.entry;
13 const factor_mod = library.factor;
14 const fused = library.fused;
15 const histogram_mod = library.histogram;
16 const image = library.image;
17 const indexing = library.indexing;
18 const layout = library.layout;
19 const normalization = library.normalization;
20 const random = library.random;
21 const reduction = library.reduction;
22 const scan_mod = library.scan;
23 const segmented = library.segmented;
24 const sort_mod = library.sort;
25 const sparse_mod = library.sparse;
26 const spatial_mod = library.spatial;
27 const tuning = library.tuning;
28
29 const Descriptor = catalog.Descriptor;
30 const OwnedDescriptor = catalog.OwnedDescriptor;
31 const OwnedKernelCallPipelinePackage = catalog.OwnedKernelCallPipelinePackage;
32 const Query = catalog.Query;
33 const ReductionOperand = catalog.ReductionOperand;
34 const MatrixProductSchedule = catalog.MatrixProductSchedule;
35 const MatrixProductCandidateDescriptors = catalog.MatrixProductCandidateDescriptors;
36 const entries = catalog.entries;
37 const descriptors = catalog.descriptors;
38 const findEntry = catalog.findEntry;
39 const createKernelCallArtifact = catalog.createKernelCallArtifact;
40 const createOwnedKernelCallArtifact = catalog.createOwnedKernelCallArtifact;
41 const createOwnedKernelCallPipelinePackage = catalog.createOwnedKernelCallPipelinePackage;
42 const createOwnedKernelCallArtifactRegistry = catalog.createOwnedKernelCallArtifactRegistry;
43 const select = catalog.select;
44 const selectOwned = catalog.selectOwned;
45 const selectOwnedMatrixProductCandidates = catalog.selectOwnedMatrixProductCandidates;
46 const selectOwnedStencilWindowCandidates = catalog.selectOwnedStencilWindowCandidates;
47 const selectOwnedGatherCandidates = catalog.selectOwnedGatherCandidates;
48 const selectOwnedScatterCandidates = catalog.selectOwnedScatterCandidates;
49 const selectOwnedScatterAddCandidates = catalog.selectOwnedScatterAddCandidates;
50 const selectOwnedHistogramCandidates = catalog.selectOwnedHistogramCandidates;
51 const selectOwnedFactor = catalog.selectOwnedFactor;
52 const selectOwnedSparse = catalog.selectOwnedSparse;
53 const selectOwnedSparseCandidates = catalog.selectOwnedSparseCandidates;
54 const selectSparseWithTuning = catalog.selectSparseWithTuning;
55 const selectOwnedSpatial = catalog.selectOwnedSpatial;
56 const selectOwnedImageCandidates = catalog.selectOwnedImageCandidates;
57 const selectOwnedRandomCandidates = catalog.selectOwnedRandomCandidates;
58 const selectOwnedFilterCandidates = catalog.selectOwnedFilterCandidates;
59 const selectOwnedPrefixSumCandidates = catalog.selectOwnedPrefixSumCandidates;
60 const selectRadixSortWithTuning = catalog.selectRadixSortWithTuning;
61 const linalg = library.linalg;
62 const stencil = library.stencil;
63
64 const MatrixProductShape = struct {
65 m: u64,
66 n: u64,
67 k: u64,
68 };
69
70 fn matrixProductDescriptorMatches(
71 descriptor: Descriptor,
72 shape: MatrixProductShape,
73 schedule: ?MatrixProductSchedule,
74 dtype: choir_abi.DType,
75 ) bool {
76 if (descriptor.metadata.category != .linalg) return false;
77 const specialization = descriptor.metadata.specialization;
78 if (!specialization.scheduleMatchesLaunch()) return false;
79 if (!specialization.operationIs(.{ .linalg = .matrix_product })) return false;
80 if (specialization.dtype != dtype) return false;
81 if (specialization.accumulation_dtype != linalg.matrixProductAccumulationDType(dtype)) return false;
82 const equation = specialization.equation orelse return false;
83 if (!std.mem.eql(u8, equation, "mk,kn->mn")) return false;
84 if (specialization.inputs.len != 2 or specialization.outputs.len != 1) return false;
85 if (specialization.reductions.len != 1) return false;
86 return specialization.inputHasExtents(0, &.{ shape.m, shape.k }) and
87 specialization.inputHasExtents(1, &.{ shape.k, shape.n }) and
88 specialization.outputHasExtents(0, &.{ shape.m, shape.n }) and
89 specialization.reductionMatches(0, .{
90 .name = "dot",
91 .operator = .dot_product,
92 .extents = &.{shape.k},
93 }) and matrixProductScheduleMatches(specialization, schedule);
94 }
95
96 fn stencilWindowDescriptorMatches(descriptor: Descriptor, query: catalog.StencilQuery) bool {
97 if (descriptor.metadata.category != .stencil) return false;
98 const specialization = descriptor.metadata.specialization;
99 if (!specialization.scheduleMatchesLaunch()) return false;
100 if (!specialization.operationIs(.{ .stencil = .window })) return false;
101 if (specialization.dtype != query.dtype) return false;
102 if (specialization.accumulation_dtype != stencil.windowAccumulationDType(query.dtype)) return false;
103 if (specialization.inputs.len != 2 or specialization.outputs.len != 1) return false;
104 if (specialization.reductions.len != 1) return false;
105 if (!stencil.windowRadiusValid(query.radius)) return false;
106 const side = 2 * @as(u64, query.radius) + 1;
107 const taps = side * side;
108 const padded_rows = query.rows + 2 * @as(u64, query.radius);
109 const padded_cols = query.cols + 2 * @as(u64, query.radius);
110 return specialization.inputHasExtents(0, &.{ padded_rows, padded_cols }) and
111 specialization.inputHasExtents(1, &.{taps}) and
112 specialization.outputHasExtents(0, &.{ query.rows, query.cols }) and
113 specialization.reductionMatches(0, .{
114 .name = "window",
115 .operator = .weighted_sum,
116 .extents = &.{taps},
117 }) and stencilWindowScheduleMatches(specialization, query.schedule);
118 }
119
120 fn matrixProductScheduleMatches(specialization: entry.Specialization, schedule: ?MatrixProductSchedule) bool {
121 const requested = schedule orelse return true;
122 return switch (requested) {
123 .thread_blocks => |threads| specializationThreadgroup2DMatches(specialization, threads),
124 };
125 }
126
127 fn stencilWindowScheduleMatches(specialization: entry.Specialization, schedule: ?catalog.StencilSchedule) bool {
128 const requested = schedule orelse return true;
129 return switch (requested) {
130 .thread_blocks => |threads| specializationThreadgroup2DMatches(specialization, threads),
131 };
132 }
133
134 fn specializationThreadgroup2DMatches(specialization: entry.Specialization, threads: entry.Threads2D) bool {
135 const launch = specialization.launch orelse return false;
136 return launch.threadgroup[0] == threads.x and
137 launch.threadgroup[1] == threads.y;
138 }
139
140 const registry_tests = struct {
141 test "kernel library catalog exposes concrete entry descriptors" {
142 try std.testing.expectEqual(entries.len, descriptors.len);
143 for (descriptors) |descriptor| {
144 try std.testing.expect(descriptor.metadata.specialization.reductionDependenciesAreValid());
145 try std.testing.expect(descriptor.metadata.specialization.reductionReuseScopesAreValid());
146 }
147
148 const sdpa = findEntry(attention.ScaledDotProductAttention2x2x3x2x2F32.target, attention.ScaledDotProductAttention2x2x3x2x2F32.version) orelse {
149 return error.TestExpectedCatalogEntry;
150 };
151 try std.testing.expectEqualStrings(attention.ScaledDotProductAttention2x2x3x2x2F32.name, sdpa.name);
152 try std.testing.expectEqual(entry.Category.attention, sdpa.metadata.category);
153 try std.testing.expect(sdpa.metadata.specialization.operationIs(.{ .attention = .scaled_dot_product }));
154
155 const vector_add = findEntry(elementwise.VectorAdd8F32.target, elementwise.VectorAdd8F32.version) orelse {
156 return error.TestExpectedCatalogEntry;
157 };
158 try std.testing.expectEqualStrings(elementwise.VectorAdd8F32.name, vector_add.name);
159 try std.testing.expectEqual(entry.Category.elementwise, vector_add.metadata.category);
160
161 const gelu = findEntry(elementwise.Gelu8F32.target, elementwise.Gelu8F32.version) orelse {
162 return error.TestExpectedCatalogEntry;
163 };
164 try std.testing.expectEqualStrings(elementwise.Gelu8F32.name, gelu.name);
165 try std.testing.expect(gelu.metadata.specialization.operationIs(.{ .activation = .gelu }));
166
167 const relu = findEntry(elementwise.Relu8F32.target, elementwise.Relu8F32.version) orelse {
168 return error.TestExpectedCatalogEntry;
169 };
170 try std.testing.expectEqualStrings(elementwise.Relu8F32.name, relu.name);
171 try std.testing.expect(relu.metadata.specialization.operationIs(.{ .activation = .relu }));
172
173 const silu = findEntry(elementwise.Silu8F32.target, elementwise.Silu8F32.version) orelse {
174 return error.TestExpectedCatalogEntry;
175 };
176 try std.testing.expectEqualStrings(elementwise.Silu8F32.name, silu.name);
177 try std.testing.expect(silu.metadata.specialization.operationIs(.{ .activation = .silu }));
178
179 const bias_gelu = findEntry(fused.BiasGelu8F32.target, fused.BiasGelu8F32.version) orelse {
180 return error.TestExpectedCatalogEntry;
181 };
182 try std.testing.expectEqualStrings(fused.BiasGelu8F32.name, bias_gelu.name);
183 try std.testing.expectEqual(entry.Category.fused, bias_gelu.metadata.category);
184 try std.testing.expect(bias_gelu.metadata.specialization.operationIs(.{ .elementwise = .add }));
185
186 const bias_relu = findEntry(fused.BiasRelu8F32.target, fused.BiasRelu8F32.version) orelse {
187 return error.TestExpectedCatalogEntry;
188 };
189 try std.testing.expectEqualStrings(fused.BiasRelu8F32.name, bias_relu.name);
190 try std.testing.expectEqual(entry.Category.fused, bias_relu.metadata.category);
191 try std.testing.expect(bias_relu.metadata.specialization.epilogueMatches(0, .{ .operator = .{ .activation = .relu } }));
192
193 const bias_silu = findEntry(fused.BiasSilu8F32.target, fused.BiasSilu8F32.version) orelse {
194 return error.TestExpectedCatalogEntry;
195 };
196 try std.testing.expectEqualStrings(fused.BiasSilu8F32.name, bias_silu.name);
197 try std.testing.expectEqual(entry.Category.fused, bias_silu.metadata.category);
198 try std.testing.expect(bias_silu.metadata.specialization.epilogueMatches(0, .{ .operator = .{ .activation = .silu } }));
199
200 const geglu = findEntry(fused.GeGlu8F32.target, fused.GeGlu8F32.version) orelse {
201 return error.TestExpectedCatalogEntry;
202 };
203 try std.testing.expectEqualStrings(fused.GeGlu8F32.name, geglu.name);
204 try std.testing.expectEqual(entry.Category.fused, geglu.metadata.category);
205 try std.testing.expect(geglu.metadata.specialization.operationIs(.{ .elementwise = .mul }));
206
207 const reglu = findEntry(fused.ReGlu8F32.target, fused.ReGlu8F32.version) orelse {
208 return error.TestExpectedCatalogEntry;
209 };
210 try std.testing.expectEqualStrings(fused.ReGlu8F32.name, reglu.name);
211 try std.testing.expectEqual(entry.Category.fused, reglu.metadata.category);
212 try std.testing.expect(reglu.metadata.specialization.operationIs(.{ .elementwise = .mul }));
213
214 const swiglu = findEntry(fused.SwiGlu8F32.target, fused.SwiGlu8F32.version) orelse {
215 return error.TestExpectedCatalogEntry;
216 };
217 try std.testing.expectEqualStrings(fused.SwiGlu8F32.name, swiglu.name);
218 try std.testing.expectEqual(entry.Category.fused, swiglu.metadata.category);
219 try std.testing.expect(swiglu.metadata.specialization.operationIs(.{ .elementwise = .mul }));
220
221 const matmul_bias_gelu = findEntry(fused.MatrixProductBiasGelu2x3x4F32.target, fused.MatrixProductBiasGelu2x3x4F32.version) orelse {
222 return error.TestExpectedCatalogEntry;
223 };
224 try std.testing.expectEqualStrings(fused.MatrixProductBiasGelu2x3x4F32.name, matmul_bias_gelu.name);
225 try std.testing.expectEqual(entry.Category.fused, matmul_bias_gelu.metadata.category);
226 try std.testing.expect(matmul_bias_gelu.metadata.specialization.operationIs(.{ .linalg = .matrix_product }));
227
228 const matmul_bias_gelu_thread_blocks = findEntry(fused.MatrixProductBiasGelu2x3x4ThreadBlocks1x2F32.target, fused.MatrixProductBiasGelu2x3x4ThreadBlocks1x2F32.version) orelse {
229 return error.TestExpectedCatalogEntry;
230 };
231 try std.testing.expectEqualStrings(fused.MatrixProductBiasGelu2x3x4ThreadBlocks1x2F32.name, matmul_bias_gelu_thread_blocks.name);
232 try std.testing.expectEqual(entry.Category.fused, matmul_bias_gelu_thread_blocks.metadata.category);
233 try std.testing.expect(matmul_bias_gelu_thread_blocks.metadata.specialization.operationIs(.{ .linalg = .matrix_product }));
234
235 const matmul_bias_relu = findEntry(fused.MatrixProductBiasRelu2x3x4F32.target, fused.MatrixProductBiasRelu2x3x4F32.version) orelse {
236 return error.TestExpectedCatalogEntry;
237 };
238 try std.testing.expectEqualStrings(fused.MatrixProductBiasRelu2x3x4F32.name, matmul_bias_relu.name);
239 try std.testing.expectEqual(entry.Category.fused, matmul_bias_relu.metadata.category);
240 try std.testing.expect(matmul_bias_relu.metadata.specialization.epilogueMatches(1, .{ .operator = .{ .activation = .relu } }));
241
242 const matmul_bias_silu = findEntry(fused.MatrixProductBiasSilu2x3x4F32.target, fused.MatrixProductBiasSilu2x3x4F32.version) orelse {
243 return error.TestExpectedCatalogEntry;
244 };
245 try std.testing.expectEqualStrings(fused.MatrixProductBiasSilu2x3x4F32.name, matmul_bias_silu.name);
246 try std.testing.expectEqual(entry.Category.fused, matmul_bias_silu.metadata.category);
247 try std.testing.expect(matmul_bias_silu.metadata.specialization.epilogueMatches(1, .{ .operator = .{ .activation = .silu } }));
248
249 const matvec_bias_gelu = findEntry(fused.MatrixVectorProductBiasGelu4x8F32.target, fused.MatrixVectorProductBiasGelu4x8F32.version) orelse {
250 return error.TestExpectedCatalogEntry;
251 };
252 try std.testing.expectEqualStrings(fused.MatrixVectorProductBiasGelu4x8F32.name, matvec_bias_gelu.name);
253 try std.testing.expectEqual(entry.Category.fused, matvec_bias_gelu.metadata.category);
254 try std.testing.expect(matvec_bias_gelu.metadata.specialization.operationIs(.{ .linalg = .matrix_vector_product }));
255
256 const matvec_bias_relu = findEntry(fused.MatrixVectorProductBiasRelu4x8F32.target, fused.MatrixVectorProductBiasRelu4x8F32.version) orelse {
257 return error.TestExpectedCatalogEntry;
258 };
259 try std.testing.expectEqualStrings(fused.MatrixVectorProductBiasRelu4x8F32.name, matvec_bias_relu.name);
260 try std.testing.expectEqual(entry.Category.fused, matvec_bias_relu.metadata.category);
261 try std.testing.expect(matvec_bias_relu.metadata.specialization.epilogueMatches(1, .{ .operator = .{ .activation = .relu } }));
262
263 const matvec_bias_silu = findEntry(fused.MatrixVectorProductBiasSilu4x8F32.target, fused.MatrixVectorProductBiasSilu4x8F32.version) orelse {
264 return error.TestExpectedCatalogEntry;
265 };
266 try std.testing.expectEqualStrings(fused.MatrixVectorProductBiasSilu4x8F32.name, matvec_bias_silu.name);
267 try std.testing.expectEqual(entry.Category.fused, matvec_bias_silu.metadata.category);
268 try std.testing.expect(matvec_bias_silu.metadata.specialization.epilogueMatches(1, .{ .operator = .{ .activation = .silu } }));
269
270 const transpose = findEntry(layout.Transpose8x16F32.target, layout.Transpose8x16F32.version) orelse {
271 return error.TestExpectedCatalogEntry;
272 };
273 try std.testing.expectEqualStrings(layout.Transpose8x16F32.name, transpose.name);
274 try std.testing.expectEqual(entry.Category.layout, transpose.metadata.category);
275 try std.testing.expect(transpose.metadata.specialization.operationIs(.{ .layout = .transpose }));
276
277 const matrix_product = findEntry(linalg.MatrixProduct4x16x8F32.target, linalg.MatrixProduct4x16x8F32.version) orelse {
278 return error.TestExpectedCatalogEntry;
279 };
280 try std.testing.expectEqualStrings(linalg.MatrixProduct4x16x8F32.name, matrix_product.name);
281 try std.testing.expectEqual(entry.Category.linalg, matrix_product.metadata.category);
282 try std.testing.expect(matrix_product.metadata.specialization.operationIs(.{ .linalg = .matrix_product }));
283
284 const matrix_product_thread_blocks = findEntry(linalg.MatrixProduct4x16x8ThreadBlocks4x2F32.target, linalg.MatrixProduct4x16x8ThreadBlocks4x2F32.version) orelse {
285 return error.TestExpectedCatalogEntry;
286 };
287 try std.testing.expectEqualStrings(linalg.MatrixProduct4x16x8ThreadBlocks4x2F32.name, matrix_product_thread_blocks.name);
288 try std.testing.expectEqual(entry.Category.linalg, matrix_product_thread_blocks.metadata.category);
289 try std.testing.expect(matrix_product_thread_blocks.metadata.specialization.operationIs(.{ .linalg = .matrix_product }));
290 try std.testing.expect(findEntry(linalg.BatchedMatrixProduct2x2x3x4F32.target, linalg.BatchedMatrixProduct2x2x3x4F32.version) == null);
291 try std.testing.expect(findEntry(linalg.MatrixProduct2x3x4F32.target, linalg.MatrixProduct2x3x4F32.version) == null);
292 try std.testing.expect(findEntry(linalg.MatrixProduct8x12x16F32.target, linalg.MatrixProduct8x12x16F32.version) == null);
293 try std.testing.expect(findEntry(linalg.MatrixVectorProduct4x8F32.target, linalg.MatrixVectorProduct4x8F32.version) == null);
294 try std.testing.expect(findEntry(linalg.OuterProduct4x3F32.target, linalg.OuterProduct4x3F32.version) == null);
295
296 const sum = findEntry(reduction.Sum8F32.target, reduction.Sum8F32.version) orelse {
297 return error.TestExpectedCatalogEntry;
298 };
299 try std.testing.expectEqualStrings(reduction.Sum8F32.name, sum.name);
300 try std.testing.expectEqual(entry.Category.reduction, sum.metadata.category);
301 try std.testing.expect(sum.metadata.specialization.operationIs(.{ .reduction = .sum }));
302
303 const dot = findEntry(reduction.Dot8F32.target, reduction.Dot8F32.version) orelse {
304 return error.TestExpectedCatalogEntry;
305 };
306 try std.testing.expectEqualStrings(reduction.Dot8F32.name, dot.name);
307 try std.testing.expectEqual(entry.Category.reduction, dot.metadata.category);
308 try std.testing.expect(dot.metadata.specialization.operationIs(.{ .reduction = .dot_product }));
309
310 const row_softmax = findEntry(normalization.RowSoftmax2x4F32.target, normalization.RowSoftmax2x4F32.version) orelse {
311 return error.TestExpectedCatalogEntry;
312 };
313 try std.testing.expectEqualStrings(normalization.RowSoftmax2x4F32.name, row_softmax.name);
314 try std.testing.expectEqual(entry.Category.normalization, row_softmax.metadata.category);
315 try std.testing.expect(row_softmax.metadata.specialization.operationIs(.{ .row_normalization = .softmax }));
316
317 const row_softmax_thread_blocks = findEntry(normalization.RowSoftmax2x4ThreadBlocks2x2F32.target, normalization.RowSoftmax2x4ThreadBlocks2x2F32.version) orelse {
318 return error.TestExpectedCatalogEntry;
319 };
320 try std.testing.expectEqualStrings(normalization.RowSoftmax2x4ThreadBlocks2x2F32.name, row_softmax_thread_blocks.name);
321 try std.testing.expectEqual(entry.Category.normalization, row_softmax_thread_blocks.metadata.category);
322 try std.testing.expect(row_softmax_thread_blocks.metadata.specialization.operationIs(.{ .row_normalization = .softmax }));
323
324 const row_log_softmax = findEntry(normalization.RowLogSoftmax2x4F32.target, normalization.RowLogSoftmax2x4F32.version) orelse {
325 return error.TestExpectedCatalogEntry;
326 };
327 try std.testing.expectEqualStrings(normalization.RowLogSoftmax2x4F32.name, row_log_softmax.name);
328 try std.testing.expectEqual(entry.Category.normalization, row_log_softmax.metadata.category);
329 try std.testing.expect(row_log_softmax.metadata.specialization.operationIs(.{ .row_normalization = .log_softmax }));
330
331 const row_rmsnorm = findEntry(normalization.RowRmsNorm2x4F32.target, normalization.RowRmsNorm2x4F32.version) orelse {
332 return error.TestExpectedCatalogEntry;
333 };
334 try std.testing.expectEqualStrings(normalization.RowRmsNorm2x4F32.name, row_rmsnorm.name);
335 try std.testing.expectEqual(entry.Category.normalization, row_rmsnorm.metadata.category);
336 try std.testing.expect(row_rmsnorm.metadata.specialization.operationIs(.{ .row_normalization = .{ .rmsnorm = .scale } }));
337
338 const row_residual_rmsnorm = findEntry(normalization.RowResidualRmsNorm2x4F32.target, normalization.RowResidualRmsNorm2x4F32.version) orelse {
339 return error.TestExpectedCatalogEntry;
340 };
341 try std.testing.expectEqualStrings(normalization.RowResidualRmsNorm2x4F32.name, row_residual_rmsnorm.name);
342 try std.testing.expectEqual(entry.Category.fused, row_residual_rmsnorm.metadata.category);
343 try std.testing.expect(row_residual_rmsnorm.metadata.specialization.operationIs(.{ .row_normalization = .{ .rmsnorm = .scale } }));
344 try std.testing.expect(row_residual_rmsnorm.metadata.specialization.inputTransformMatches(0, .{
345 .operator = .residual_add,
346 .input_index = 1,
347 .extents = &.{ 2, 4 },
348 }));
349
350 const row_layernorm = findEntry(normalization.RowLayerNorm2x4F32.target, normalization.RowLayerNorm2x4F32.version) orelse {
351 return error.TestExpectedCatalogEntry;
352 };
353 try std.testing.expectEqualStrings(normalization.RowLayerNorm2x4F32.name, row_layernorm.name);
354 try std.testing.expectEqual(entry.Category.normalization, row_layernorm.metadata.category);
355 try std.testing.expect(row_layernorm.metadata.specialization.operationIs(.{ .row_normalization = .{ .layernorm = .none } }));
356
357 const row_affine_layernorm = findEntry(normalization.RowAffineLayerNorm2x4F32.target, normalization.RowAffineLayerNorm2x4F32.version) orelse {
358 return error.TestExpectedCatalogEntry;
359 };
360 try std.testing.expectEqualStrings(normalization.RowAffineLayerNorm2x4F32.name, row_affine_layernorm.name);
361 try std.testing.expectEqual(entry.Category.normalization, row_affine_layernorm.metadata.category);
362 try std.testing.expect(row_affine_layernorm.metadata.specialization.operationIs(.{ .row_normalization = .{ .layernorm = .scale_bias } }));
363 }
364 };
365
366 const artifact_tests = struct {
367 test "kernel library catalog builds histogram scheduled artifact registry" {
368 const allocator = std.testing.allocator;
369 var state = gpu.recording.BackendState{
370 .allocator = allocator,
371 .kind = .cuda,
372 .format = .cuda_ptx,
373 };
374 var candidates = try selectOwnedHistogramCandidates(allocator, .{
375 .dtype = .f32,
376 .bins = 128,
377 .count = 8192,
378 .schedule = .{ .thread_blocks = 256 },
379 });
380 defer candidates.deinit();
381 try std.testing.expectEqual(@as(usize, 1), candidates.count);
382
383 var owned_registry = try createOwnedKernelCallArtifactRegistry(allocator, state.handle(), candidates.slice(), .{ .limits = .testing });
384 defer owned_registry.deinit();
385 const registry = owned_registry.registry();
386 try std.testing.expectEqual(candidates.count, registry.entries.len);
387
388 for (candidates.slice()) |candidate| {
389 const found = registry.find(
390 candidate.descriptor.metadata.target,
391 histogram_mod.histogram_family_version,
392 .cuda_ptx,
393 ) orelse return error.TestExpectedKernelCallArtifact;
394 try std.testing.expectEqual(@as(u32, 4), found.runtime_scalar_argument_count);
395 const launch = switch (found.launch) {
396 .derived => |derived| derived,
397 .fixed => return error.TestExpectedDerivedLaunch,
398 };
399 const specialization_launch = candidate.descriptor.metadata.specialization.launch orelse {
400 return error.TestExpectedSpecializationLaunch;
401 };
402 try std.testing.expectEqual(specialization_launch.threadgroup[0], launch.threadgroup[0]);
403 const profile = found.shape_profile orelse return error.TestExpectedShapeProfile;
404 try std.testing.expectEqual(@as(usize, 2), profile.dimensions.len);
405 }
406 }
407
408 test "kernel library catalog builds image descriptor artifact registry" {
409 const allocator = std.testing.allocator;
410 var state = gpu.recording.BackendState{
411 .allocator = allocator,
412 .kind = .cuda,
413 .format = .cuda_ptx,
414 };
415
416 var descriptors_owned: [2]OwnedDescriptor = undefined;
417 descriptors_owned[0] = (try selectOwned(allocator, .{ .image = .{
418 .dtype = .u32,
419 .kind = .{ .blur_pass = .{ .radius = 2, .axis = .horizontal } },
420 .width = 24,
421 .height = 10,
422 .schedule = .{ .thread_blocks = .{ .x = 8, .y = 4 } },
423 } })) orelse return error.TestExpectedImageDescriptor;
424 var owned_count: usize = 1;
425 defer for (descriptors_owned[0..owned_count]) |*descriptor| descriptor.deinit();
426
427 descriptors_owned[1] = (try selectOwned(allocator, .{ .image = .{
428 .dtype = .u32,
429 .kind = .{ .resize_bilinear = .{ .src_width = 20, .src_height = 12 } },
430 .width = 11,
431 .height = 7,
432 .schedule = .{ .thread_blocks = .{ .x = 8, .y = 4 } },
433 } })) orelse return error.TestExpectedImageDescriptor;
434 owned_count = 2;
435
436 var owned_registry = try createOwnedKernelCallArtifactRegistry(
437 allocator,
438 state.handle(),
439 descriptors_owned[0..owned_count],
440 .{ .limits = .testing },
441 );
442 defer owned_registry.deinit();
443 const registry = owned_registry.registry();
444 try std.testing.expectEqual(@as(usize, 2), registry.entries.len);
445
446 const blur_artifact = registry.find("image_blur_pass_family_r2h_8x4_rgba8", image.blur_family_version, .cuda_ptx) orelse {
447 return error.TestExpectedKernelCallArtifact;
448 };
449 try std.testing.expectEqualStrings("accy_image_blur_pass_r2h_rgba8", blur_artifact.entry_name);
450 try std.testing.expectEqual(@as(u32, 5), blur_artifact.argument_count);
451 try std.testing.expectEqual(@as(u32, 2), blur_artifact.runtime_scalar_argument_count);
452 try std.testing.expect(blur_artifact.required_dtypes.contains(.u32));
453 try std.testing.expect(blur_artifact.required_dtypes.contains(.f32));
454 try std.testing.expect(blur_artifact.required_dtypes.contains(.i32));
455 const blur_launch = switch (blur_artifact.launch) {
456 .derived => |derived| derived,
457 .fixed => return error.TestExpectedDerivedLaunch,
458 };
459 try std.testing.expectEqual(@as(u32, 8), blur_launch.threadgroup[0]);
460 try std.testing.expectEqual(@as(u32, 4), blur_launch.threadgroup[1]);
461 const blur_profile = blur_artifact.shape_profile orelse return error.TestExpectedShapeProfile;
462 try std.testing.expectEqualStrings("image_blur_pass", blur_profile.name);
463 try std.testing.expectEqual(@as(usize, 2), blur_profile.dimensions.len);
464 try std.testing.expect(blur_profile.runtimeScalarDimension(0) != null);
465 try std.testing.expect(blur_profile.runtimeScalarDimension(1) != null);
466
467 const resize_artifact = registry.find("image_resize_bilinear_family_8x4_rgba8", image.resize_family_version, .cuda_ptx) orelse {
468 return error.TestExpectedKernelCallArtifact;
469 };
470 try std.testing.expectEqualStrings("accy_image_resize_bilinear_rgba8", resize_artifact.entry_name);
471 try std.testing.expectEqual(@as(u32, 9), resize_artifact.argument_count);
472 try std.testing.expectEqual(@as(u32, 7), resize_artifact.runtime_scalar_argument_count);
473 try std.testing.expect(resize_artifact.required_dtypes.contains(.u32));
474 try std.testing.expect(resize_artifact.required_dtypes.contains(.f32));
475 try std.testing.expect(resize_artifact.required_dtypes.contains(.i32));
476 const resize_launch = switch (resize_artifact.launch) {
477 .derived => |derived| derived,
478 .fixed => return error.TestExpectedDerivedLaunch,
479 };
480 try std.testing.expectEqual(@as(u32, 8), resize_launch.threadgroup[0]);
481 try std.testing.expectEqual(@as(u32, 4), resize_launch.threadgroup[1]);
482 const resize_profile = resize_artifact.shape_profile orelse return error.TestExpectedShapeProfile;
483 try std.testing.expectEqualStrings("image_resize_bilinear", resize_profile.name);
484 try std.testing.expectEqual(@as(usize, 3), resize_profile.dimensions.len);
485 try std.testing.expect(resize_profile.runtimeScalarDimension(0) != null);
486 try std.testing.expect(resize_profile.runtimeScalarDimension(1) != null);
487 try std.testing.expect(resize_profile.runtimeScalarDimension(2) != null);
488 }
489
490 test "kernel library catalog builds factor descriptor artifact registry" {
491 const allocator = std.testing.allocator;
492 var state = gpu.recording.BackendState{
493 .allocator = allocator,
494 .kind = .cuda,
495 .format = .cuda_ptx,
496 };
497
498 var descriptors_owned: [4]OwnedDescriptor = undefined;
499 descriptors_owned[0] = (try selectOwnedFactor(allocator, .{
500 .dtype = .f32,
501 .kind = .{ .batched_cholesky = .{ .n = 3, .batch = 5000 } },
502 .schedule = .{ .thread_blocks = 64 },
503 })) orelse return error.TestExpectedFactorDescriptor;
504 var owned_count: usize = 1;
505 defer for (descriptors_owned[0..owned_count]) |*descriptor| descriptor.deinit();
506
507 descriptors_owned[1] = (try selectOwnedFactor(allocator, .{
508 .dtype = .f32,
509 .kind = .{ .batched_cholesky_solve = .{ .n = 3, .batch = 5000 } },
510 .schedule = .{ .thread_blocks = 64 },
511 })) orelse return error.TestExpectedFactorDescriptor;
512 owned_count = 2;
513
514 descriptors_owned[2] = (try selectOwnedFactor(allocator, .{
515 .dtype = .f32,
516 .kind = .{ .batched_cholesky_solve = .{ .n = 3, .batch = 5000 } },
517 .layout = .interleaved,
518 .schedule = .{ .thread_blocks = 64 },
519 })) orelse return error.TestExpectedFactorDescriptor;
520 owned_count = 3;
521 try std.testing.expectEqualStrings(
522 "accy.kernel.linalg.batched_cholesky_solve_family_3_64_il_f32",
523 descriptors_owned[2].descriptor.metadata.target,
524 );
525 try std.testing.expect(descriptors_owned[2].descriptor.metadata.specialization.layoutIs("interleaved"));
526
527 descriptors_owned[3] = (try selectOwnedFactor(allocator, .{
528 .dtype = .f32,
529 .kind = .{ .batched_inverse = .{ .n = 3, .batch = 5000 } },
530 .schedule = .{ .thread_blocks = 64 },
531 })) orelse return error.TestExpectedFactorDescriptor;
532 owned_count = 4;
533
534 var owned_registry = try createOwnedKernelCallArtifactRegistry(
535 allocator,
536 state.handle(),
537 descriptors_owned[0..owned_count],
538 .{ .limits = .testing },
539 );
540 defer owned_registry.deinit();
541 const registry = owned_registry.registry();
542 try std.testing.expectEqual(@as(usize, 4), registry.entries.len);
543
544 const versions = [_]u32{
545 factor_mod.batched_cholesky_family_version,
546 factor_mod.batched_cholesky_solve_family_version,
547 factor_mod.batched_cholesky_solve_family_version,
548 factor_mod.batched_inverse_family_version,
549 };
550 for (descriptors_owned[0..owned_count], versions) |descriptor, version| {
551 const found = registry.find(
552 descriptor.descriptor.metadata.target,
553 version,
554 .cuda_ptx,
555 ) orelse return error.TestExpectedKernelCallArtifact;
556 try std.testing.expectEqual(@as(u32, 1), found.runtime_scalar_argument_count);
557 const launch = switch (found.launch) {
558 .derived => |derived| derived,
559 .fixed => return error.TestExpectedDerivedLaunch,
560 };
561 try std.testing.expectEqual(@as(u32, 64), launch.threadgroup[0]);
562 const profile = found.shape_profile orelse return error.TestExpectedShapeProfile;
563 try std.testing.expectEqual(@as(usize, 1), profile.dimensions.len);
564 }
565 }
566
567 test "kernel library catalog rejects malformed factor descriptor artifacts" {
568 const allocator = std.testing.allocator;
569 var state = gpu.recording.BackendState{
570 .allocator = allocator,
571 .kind = .cuda,
572 .format = .cuda_ptx,
573 };
574
575 {
576 var descriptor = (try selectOwnedFactor(allocator, .{
577 .dtype = .f32,
578 .kind = .{ .batched_cholesky = .{ .n = 3, .batch = 5000 } },
579 .schedule = .{ .thread_blocks = 64 },
580 })) orelse return error.TestExpectedFactorDescriptor;
581 defer descriptor.deinit();
582
583 if (descriptor.specialization) |*specialization| {
584 const lifetime_allocator = specialization.allocator();
585 specialization.value.inputs = &.{
586 try entry.runtimeShape3D(lifetime_allocator, "b", 4096, "i", 3, "j", 3),
587 };
588 descriptor.descriptor.metadata.specialization = specialization.value;
589 } else return error.TestExpectedOwnedSpecialization;
590
591 try std.testing.expectError(
592 error.UnknownKernelLibraryEntry,
593 createOwnedKernelCallArtifact(allocator, state.handle(), descriptor, .{ .limits = .testing }),
594 );
595 }
596
597 {
598 var descriptor = (try selectOwnedFactor(allocator, .{
599 .dtype = .f32,
600 .kind = .{ .batched_cholesky_solve = .{ .n = 3, .batch = 5000 } },
601 .schedule = .{ .thread_blocks = 64 },
602 })) orelse return error.TestExpectedFactorDescriptor;
603 defer descriptor.deinit();
604
605 if (descriptor.specialization) |*specialization| {
606 const lifetime_allocator = specialization.allocator();
607 const inputs = try lifetime_allocator.alloc(entry.Shape, 2);
608 inputs[0] = specialization.value.inputs[0];
609 inputs[1] = try entry.runtimeShape2D(lifetime_allocator, "b", 5000, "i", 4);
610 specialization.value.inputs = inputs;
611 descriptor.descriptor.metadata.specialization = specialization.value;
612 } else return error.TestExpectedOwnedSpecialization;
613
614 try std.testing.expectError(
615 error.UnknownKernelLibraryEntry,
616 createOwnedKernelCallArtifact(allocator, state.handle(), descriptor, .{ .limits = .testing }),
617 );
618 }
619 }
620
621 test "kernel library catalog builds sparse descriptor artifact registry" {
622 const allocator = std.testing.allocator;
623 var state = gpu.recording.BackendState{
624 .allocator = allocator,
625 .kind = .cuda,
626 .format = .cuda_ptx,
627 };
628
629 var descriptors_owned: [10]OwnedDescriptor = undefined;
630 descriptors_owned[0] = (try selectOwnedSparse(allocator, .{
631 .dtype = .f32,
632 .kind = .{ .csr_spmv = .{ .rows = 70, .nnz = 560, .x_extent = 40 } },
633 .schedule = .{ .thread_blocks = 64 },
634 })) orelse return error.TestExpectedSparseDescriptor;
635 var owned_count: usize = 1;
636 defer for (descriptors_owned[0..owned_count]) |*descriptor| descriptor.deinit();
637
638 descriptors_owned[1] = (try selectOwnedSparse(allocator, .{
639 .dtype = .f32,
640 .kind = .{ .csr_spmv = .{ .rows = 70, .nnz = 90, .x_extent = 40 } },
641 .structure = .row_thread,
642 .schedule = .{ .thread_blocks = 64 },
643 })) orelse return error.TestExpectedSparseDescriptor;
644 owned_count = 2;
645
646 descriptors_owned[2] = (try selectOwnedSparse(allocator, .{
647 .dtype = .f16,
648 .kind = .{ .csr_spmv = .{ .rows = 70, .nnz = 560, .x_extent = 40 } },
649 .schedule = .{ .thread_blocks = 64 },
650 })) orelse return error.TestExpectedSparseDescriptor;
651 owned_count = 3;
652
653 descriptors_owned[3] = (try selectOwnedSparse(allocator, .{
654 .dtype = .f64,
655 .kind = .{ .csr_spmv = .{ .rows = 70, .nnz = 560, .x_extent = 40 } },
656 .schedule = .{ .thread_blocks = 64 },
657 })) orelse return error.TestExpectedSparseDescriptor;
658 owned_count = 4;
659
660 descriptors_owned[4] = (try selectOwnedSparse(allocator, .{
661 .dtype = .f32,
662 .kind = .{ .coo_spmv = .{ .rows = 70, .nnz = 560, .x_extent = 40 } },
663 .schedule = .{ .thread_blocks = 64 },
664 })) orelse return error.TestExpectedSparseDescriptor;
665 owned_count = 5;
666
667 descriptors_owned[5] = (try selectOwnedSparse(allocator, .{
668 .dtype = .f16,
669 .kind = .{ .coo_spmv = .{ .rows = 70, .nnz = 560, .x_extent = 40 } },
670 .schedule = .{ .thread_blocks = 64 },
671 })) orelse return error.TestExpectedSparseDescriptor;
672 owned_count = 6;
673
674 descriptors_owned[6] = (try selectOwnedSparse(allocator, .{
675 .dtype = .f64,
676 .kind = .{ .coo_spmv = .{ .rows = 70, .nnz = 560, .x_extent = 40 } },
677 .schedule = .{ .thread_blocks = 64 },
678 })) orelse return error.TestExpectedSparseDescriptor;
679 owned_count = 7;
680
681 descriptors_owned[7] = (try selectOwnedSparse(allocator, .{
682 .dtype = .f32,
683 .kind = .{ .csr_spmm = .{ .rows = 70, .columns = 45, .nnz = 512, .x_extent = 40 } },
684 .schedule = .{ .thread_blocks_2d = .{ .x = 8, .y = 4 } },
685 })) orelse return error.TestExpectedSparseDescriptor;
686 owned_count = 8;
687
688 descriptors_owned[8] = (try selectOwnedSparse(allocator, .{
689 .dtype = .f32,
690 .kind = .{ .ell_spmv = .{ .rows = 70, .slots = 8, .x_extent = 40 } },
691 .schedule = .{ .thread_blocks = 64 },
692 })) orelse return error.TestExpectedSparseDescriptor;
693 owned_count = 9;
694
695 descriptors_owned[9] = (try selectOwnedSparse(allocator, .{
696 .dtype = .f32,
697 .kind = .{ .sell_spmv = .{ .rows = 70, .slice_size = 8, .values_size = 400, .x_extent = 40 } },
698 .schedule = .{ .thread_blocks = 64 },
699 })) orelse return error.TestExpectedSparseDescriptor;
700 owned_count = 10;
701
702 const warp_instance = sparse_mod.spmvCsrInstanceFromSpecialization(
703 descriptors_owned[0].descriptor.metadata.specialization,
704 ) orelse return error.TestExpectedSparseInstance;
705 try std.testing.expectEqual(@as(u64, 70), warp_instance.rows);
706 try std.testing.expectEqual(@as(u64, 560), warp_instance.nnz);
707 try std.testing.expectEqual(@as(u64, 40), warp_instance.x_extent);
708 try std.testing.expectEqual(sparse_mod.SpmvCsrStructure.row_warp, warp_instance.structure);
709
710 const thread_instance = sparse_mod.spmvCsrInstanceFromSpecialization(
711 descriptors_owned[1].descriptor.metadata.specialization,
712 ) orelse return error.TestExpectedSparseInstance;
713 try std.testing.expectEqual(@as(u64, 70), thread_instance.rows);
714 try std.testing.expectEqual(@as(u64, 90), thread_instance.nnz);
715 try std.testing.expectEqual(@as(u64, 40), thread_instance.x_extent);
716 try std.testing.expectEqual(sparse_mod.SpmvCsrStructure.row_thread, thread_instance.structure);
717
718 const half_instance = sparse_mod.spmvCsrInstanceFromSpecialization(
719 descriptors_owned[2].descriptor.metadata.specialization,
720 ) orelse return error.TestExpectedSparseInstance;
721 try std.testing.expectEqual(choir_abi.DType.f16, half_instance.dtype);
722 try std.testing.expectEqual(choir_abi.DType.f32, half_instance.accumulation_dtype);
723 try std.testing.expectEqual(sparse_mod.SpmvCsrStructure.row_warp, half_instance.structure);
724
725 const double_instance = sparse_mod.spmvCsrInstanceFromSpecialization(
726 descriptors_owned[3].descriptor.metadata.specialization,
727 ) orelse return error.TestExpectedSparseInstance;
728 try std.testing.expectEqual(choir_abi.DType.f64, double_instance.dtype);
729 try std.testing.expectEqual(choir_abi.DType.f64, double_instance.accumulation_dtype);
730 try std.testing.expectEqual(sparse_mod.SpmvCsrStructure.row_warp, double_instance.structure);
731
732 const coo_instance = sparse_mod.spmvCooInstanceFromSpecialization(
733 descriptors_owned[4].descriptor.metadata.specialization,
734 ) orelse return error.TestExpectedSparseInstance;
735 try std.testing.expectEqual(@as(u64, 70), coo_instance.rows);
736 try std.testing.expectEqual(@as(u64, 560), coo_instance.nnz);
737 try std.testing.expectEqual(@as(u64, 40), coo_instance.x_extent);
738 try std.testing.expectEqual(sparse_mod.SpmvCooStructure.element_thread, coo_instance.structure);
739
740 const coo_half_instance = sparse_mod.spmvCooInstanceFromSpecialization(
741 descriptors_owned[5].descriptor.metadata.specialization,
742 ) orelse return error.TestExpectedSparseInstance;
743 try std.testing.expectEqual(choir_abi.DType.f16, coo_half_instance.dtype);
744 try std.testing.expectEqual(choir_abi.DType.f32, coo_half_instance.accumulation_dtype);
745 try std.testing.expectEqual(sparse_mod.SpmvCooStructure.row_thread, coo_half_instance.structure);
746
747 const coo_double_instance = sparse_mod.spmvCooInstanceFromSpecialization(
748 descriptors_owned[6].descriptor.metadata.specialization,
749 ) orelse return error.TestExpectedSparseInstance;
750 try std.testing.expectEqual(choir_abi.DType.f64, coo_double_instance.dtype);
751 try std.testing.expectEqual(choir_abi.DType.f64, coo_double_instance.accumulation_dtype);
752 try std.testing.expectEqual(sparse_mod.SpmvCooStructure.row_thread, coo_double_instance.structure);
753
754 const spmm_instance = sparse_mod.spmmCsrInstanceFromSpecialization(
755 descriptors_owned[7].descriptor.metadata.specialization,
756 ) orelse return error.TestExpectedSparseInstance;
757 try std.testing.expectEqual(@as(u64, 70), spmm_instance.rows);
758 try std.testing.expectEqual(@as(u64, 45), spmm_instance.columns);
759 try std.testing.expectEqual(@as(u64, 512), spmm_instance.nnz);
760 try std.testing.expectEqual(@as(u64, 40), spmm_instance.x_extent);
761 try std.testing.expectEqual(sparse_mod.SpmmCsrStructure.row_column_thread, spmm_instance.structure);
762
763 const ell_instance = sparse_mod.spmvEllInstanceFromSpecialization(
764 descriptors_owned[8].descriptor.metadata.specialization,
765 ) orelse return error.TestExpectedSparseInstance;
766 try std.testing.expectEqual(@as(u64, 70), ell_instance.rows);
767 try std.testing.expectEqual(@as(u64, 8), ell_instance.slots);
768 try std.testing.expectEqual(@as(u64, 40), ell_instance.x_extent);
769 try std.testing.expectEqual(sparse_mod.SpmvEllStructure.row_thread, ell_instance.structure);
770
771 const sell_instance = sparse_mod.spmvSellInstanceFromSpecialization(
772 descriptors_owned[9].descriptor.metadata.specialization,
773 ) orelse return error.TestExpectedSparseInstance;
774 try std.testing.expectEqual(@as(u64, 70), sell_instance.rows);
775 try std.testing.expectEqual(@as(u64, 8), sell_instance.slice_size);
776 try std.testing.expectEqual(@as(u64, 400), sell_instance.values_size);
777 try std.testing.expectEqual(@as(u64, 40), sell_instance.x_extent);
778 try std.testing.expectEqual(sparse_mod.SpmvSellStructure.row_thread, sell_instance.structure);
779 try std.testing.expect(descriptors_owned[9].descriptor.metadata.specialization.staticParameterMatches(sparse_mod.spmv_sell_slice_size_parameter, 8));
780
781 var owned_registry = try createOwnedKernelCallArtifactRegistry(
782 allocator,
783 state.handle(),
784 descriptors_owned[0..owned_count],
785 .{ .limits = .testing },
786 );
787 defer owned_registry.deinit();
788 const registry = owned_registry.registry();
789 try std.testing.expectEqual(@as(usize, 10), registry.entries.len);
790
791 for (descriptors_owned[0..4]) |descriptor| {
792 const found = registry.find(
793 descriptor.descriptor.metadata.target,
794 sparse_mod.spmv_csr_family_version,
795 .cuda_ptx,
796 ) orelse return error.TestExpectedKernelCallArtifact;
797 try std.testing.expectEqual(@as(u32, 3), found.runtime_scalar_argument_count);
798 const launch = switch (found.launch) {
799 .derived => |derived| derived,
800 .fixed => return error.TestExpectedDerivedLaunch,
801 };
802 try std.testing.expectEqual(@as(u32, 64), launch.threadgroup[0]);
803 const profile = found.shape_profile orelse return error.TestExpectedShapeProfile;
804 try std.testing.expectEqual(@as(usize, 1), profile.dimensions.len);
805 }
806
807 for (descriptors_owned[4..7], 0..) |coo_descriptor, coo_index| {
808 const coo_artifact = registry.find(
809 coo_descriptor.descriptor.metadata.target,
810 sparse_mod.spmv_coo_family_version,
811 .cuda_ptx,
812 ) orelse return error.TestExpectedKernelCallArtifact;
813 try std.testing.expectEqual(@as(u32, 3), coo_artifact.runtime_scalar_argument_count);
814 const coo_launch = switch (coo_artifact.launch) {
815 .derived => |derived| derived,
816 .fixed => return error.TestExpectedDerivedLaunch,
817 };
818 try std.testing.expectEqual(@as(u32, 64), coo_launch.threadgroup[0]);
819 const coo_profile = coo_artifact.shape_profile orelse return error.TestExpectedShapeProfile;
820 try std.testing.expectEqual(@as(usize, 2), coo_profile.dimensions.len);
821 switch (coo_launch.grid[0]) {
822 .runtime_u32_ceil_div => |ceil| {
823 try std.testing.expectEqual(@as(u32, if (coo_index == 0) 1 else 0), ceil.argument_index);
824 try std.testing.expectEqual(@as(u32, 64), ceil.divisor);
825 },
826 else => return error.TestExpectedDerivedGrid,
827 }
828 }
829
830 const spmm_descriptor = descriptors_owned[7];
831 const spmm_artifact = registry.find(
832 spmm_descriptor.descriptor.metadata.target,
833 sparse_mod.spmm_csr_family_version,
834 .cuda_ptx,
835 ) orelse return error.TestExpectedKernelCallArtifact;
836 try std.testing.expectEqual(@as(u32, 4), spmm_artifact.runtime_scalar_argument_count);
837 const spmm_launch = switch (spmm_artifact.launch) {
838 .derived => |derived| derived,
839 .fixed => return error.TestExpectedDerivedLaunch,
840 };
841 try std.testing.expectEqual(@as(u32, 8), spmm_launch.threadgroup[0]);
842 try std.testing.expectEqual(@as(u32, 4), spmm_launch.threadgroup[1]);
843 const spmm_profile = spmm_artifact.shape_profile orelse return error.TestExpectedShapeProfile;
844 try std.testing.expectEqual(@as(usize, 2), spmm_profile.dimensions.len);
845
846 const ell_descriptor = descriptors_owned[8];
847 const ell_artifact = registry.find(
848 ell_descriptor.descriptor.metadata.target,
849 sparse_mod.spmv_ell_family_version,
850 .cuda_ptx,
851 ) orelse return error.TestExpectedKernelCallArtifact;
852 try std.testing.expectEqual(@as(u32, 3), ell_artifact.runtime_scalar_argument_count);
853 const ell_launch = switch (ell_artifact.launch) {
854 .derived => |derived| derived,
855 .fixed => return error.TestExpectedDerivedLaunch,
856 };
857 try std.testing.expectEqual(@as(u32, 64), ell_launch.threadgroup[0]);
858 const ell_profile = ell_artifact.shape_profile orelse return error.TestExpectedShapeProfile;
859 try std.testing.expectEqual(@as(usize, 1), ell_profile.dimensions.len);
860
861 const sell_descriptor = descriptors_owned[9];
862 const sell_artifact = registry.find(
863 sell_descriptor.descriptor.metadata.target,
864 sparse_mod.spmv_sell_family_version,
865 .cuda_ptx,
866 ) orelse return error.TestExpectedKernelCallArtifact;
867 try std.testing.expectEqual(@as(u32, 3), sell_artifact.runtime_scalar_argument_count);
868 const sell_launch = switch (sell_artifact.launch) {
869 .derived => |derived| derived,
870 .fixed => return error.TestExpectedDerivedLaunch,
871 };
872 try std.testing.expectEqual(@as(u32, 64), sell_launch.threadgroup[0]);
873 const sell_profile = sell_artifact.shape_profile orelse return error.TestExpectedShapeProfile;
874 try std.testing.expectEqual(@as(usize, 1), sell_profile.dimensions.len);
875
876 const warp_descriptor = descriptors_owned[0];
877 const warp_artifact = registry.find(
878 warp_descriptor.descriptor.metadata.target,
879 sparse_mod.spmv_csr_family_version,
880 .cuda_ptx,
881 ) orelse return error.TestExpectedKernelCallArtifact;
882 switch (warp_artifact.launch.derived.grid[0]) {
883 .runtime_u32_ceil_div => |ceil| {
884 try std.testing.expectEqual(@as(u32, 0), ceil.argument_index);
885 try std.testing.expectEqual(@as(u32, 2), ceil.divisor);
886 },
887 else => return error.TestExpectedDerivedGrid,
888 }
889 switch (spmm_artifact.launch.derived.grid[0]) {
890 .runtime_u32_ceil_div => |ceil| {
891 try std.testing.expectEqual(@as(u32, 3), ceil.argument_index);
892 try std.testing.expectEqual(@as(u32, 8), ceil.divisor);
893 },
894 else => return error.TestExpectedDerivedGrid,
895 }
896 switch (spmm_artifact.launch.derived.grid[1]) {
897 .runtime_u32_ceil_div => |ceil| {
898 try std.testing.expectEqual(@as(u32, 0), ceil.argument_index);
899 try std.testing.expectEqual(@as(u32, 4), ceil.divisor);
900 },
901 else => return error.TestExpectedDerivedGrid,
902 }
903 switch (ell_artifact.launch.derived.grid[0]) {
904 .runtime_u32_ceil_div => |ceil| {
905 try std.testing.expectEqual(@as(u32, 0), ceil.argument_index);
906 try std.testing.expectEqual(@as(u32, 64), ceil.divisor);
907 },
908 else => return error.TestExpectedDerivedGrid,
909 }
910 switch (sell_artifact.launch.derived.grid[0]) {
911 .runtime_u32_ceil_div => |ceil| {
912 try std.testing.expectEqual(@as(u32, 0), ceil.argument_index);
913 try std.testing.expectEqual(@as(u32, 64), ceil.divisor);
914 },
915 else => return error.TestExpectedDerivedGrid,
916 }
917 }
918
919 test "kernel library catalog builds tiny csr spmv candidate artifact registry" {
920 const allocator = std.testing.allocator;
921 var state = gpu.recording.BackendState{
922 .allocator = allocator,
923 .kind = .cuda,
924 .format = .cuda_ptx,
925 };
926
927 var candidates = try selectOwnedSparseCandidates(allocator, .{
928 .dtype = .f32,
929 .kind = .{ .csr_spmv = .{ .rows = 8, .nnz = 64, .x_extent = 8 } },
930 });
931 defer candidates.deinit();
932 try std.testing.expectEqual(@as(usize, 2), candidates.count);
933
934 var owned_registry = try createOwnedKernelCallArtifactRegistry(
935 allocator,
936 state.handle(),
937 candidates.slice(),
938 .{ .limits = .testing },
939 );
940 defer owned_registry.deinit();
941 const registry = owned_registry.registry();
942 try std.testing.expectEqual(candidates.count, registry.entries.len);
943
944 for (candidates.slice()) |candidate| {
945 const found = registry.find(
946 candidate.descriptor.metadata.target,
947 sparse_mod.spmv_csr_family_version,
948 .cuda_ptx,
949 ) orelse return error.TestExpectedKernelCallArtifact;
950 try std.testing.expectEqual(@as(u32, 3), found.runtime_scalar_argument_count);
951 }
952 }
953
954 test "kernel library catalog rejects malformed sparse descriptor artifacts" {
955 const allocator = std.testing.allocator;
956 var state = gpu.recording.BackendState{
957 .allocator = allocator,
958 .kind = .cuda,
959 .format = .cuda_ptx,
960 };
961
962 {
963 var descriptor = (try selectOwnedSparse(allocator, .{
964 .dtype = .f32,
965 .kind = .{ .csr_spmv = .{ .rows = 70, .nnz = 560, .x_extent = 40 } },
966 .schedule = .{ .thread_blocks = 64 },
967 })) orelse return error.TestExpectedSparseDescriptor;
968 defer descriptor.deinit();
969
970 if (descriptor.specialization) |*specialization| {
971 const lifetime_allocator = specialization.allocator();
972 const inputs = try lifetime_allocator.alloc(entry.Shape, 4);
973 inputs[0] = specialization.value.inputs[0];
974 inputs[1] = specialization.value.inputs[1];
975 inputs[2] = specialization.value.inputs[2];
976 inputs[3] = try entry.runtimeShape1D(lifetime_allocator, "bad_x", 40);
977 specialization.value.inputs = inputs;
978 descriptor.descriptor.metadata.specialization = specialization.value;
979 } else return error.TestExpectedOwnedSpecialization;
980
981 try std.testing.expectError(
982 error.UnknownKernelLibraryEntry,
983 createOwnedKernelCallArtifact(allocator, state.handle(), descriptor, .{ .limits = .testing }),
984 );
985 }
986
987 {
988 var descriptor = (try selectOwnedSparse(allocator, .{
989 .dtype = .f32,
990 .kind = .{ .csr_spmv = .{ .rows = 70, .nnz = 560, .x_extent = 40 } },
991 .structure = .row_warp,
992 .schedule = .{ .thread_blocks = 64 },
993 })) orelse return error.TestExpectedSparseDescriptor;
994 defer descriptor.deinit();
995
996 if (descriptor.specialization) |*specialization| {
997 const lifetime_allocator = specialization.allocator();
998 specialization.value.schedule = try entry.runtimeThreadBlocks1D(lifetime_allocator, "r", 70, 64);
999 specialization.value.launch = specialization.value.schedule.?.launch();
1000 descriptor.descriptor.metadata.specialization = specialization.value;
1001 } else return error.TestExpectedOwnedSpecialization;
1002
1003 try std.testing.expectError(
1004 error.UnknownKernelLibraryEntry,
1005 createOwnedKernelCallArtifact(allocator, state.handle(), descriptor, .{ .limits = .testing }),
1006 );
1007 }
1008
1009 {
1010 var descriptor = (try selectOwnedSparse(allocator, .{
1011 .dtype = .f32,
1012 .kind = .{ .coo_spmv = .{ .rows = 70, .nnz = 560, .x_extent = 40 } },
1013 .schedule = .{ .thread_blocks = 64 },
1014 })) orelse return error.TestExpectedSparseDescriptor;
1015 defer descriptor.deinit();
1016
1017 if (descriptor.specialization) |*specialization| {
1018 const lifetime_allocator = specialization.allocator();
1019 specialization.value.schedule = try entry.runtimeThreadBlocks1D(lifetime_allocator, "r", 70, 64);
1020 specialization.value.launch = specialization.value.schedule.?.launch();
1021 descriptor.descriptor.metadata.specialization = specialization.value;
1022 } else return error.TestExpectedOwnedSpecialization;
1023
1024 try std.testing.expectError(
1025 error.UnknownKernelLibraryEntry,
1026 createOwnedKernelCallArtifact(allocator, state.handle(), descriptor, .{ .limits = .testing }),
1027 );
1028 }
1029
1030 {
1031 var descriptor = (try selectOwnedSparse(allocator, .{
1032 .dtype = .f32,
1033 .kind = .{ .sell_spmv = .{ .rows = 70, .slice_size = 8, .values_size = 400, .x_extent = 40 } },
1034 .schedule = .{ .thread_blocks = 64 },
1035 })) orelse return error.TestExpectedSparseDescriptor;
1036 defer descriptor.deinit();
1037
1038 if (descriptor.specialization) |*specialization| {
1039 specialization.value.static_parameters = &.{};
1040 descriptor.descriptor.metadata.specialization = specialization.value;
1041 } else return error.TestExpectedOwnedSpecialization;
1042
1043 try std.testing.expectError(
1044 error.UnknownKernelLibraryEntry,
1045 createOwnedKernelCallArtifact(allocator, state.handle(), descriptor, .{ .limits = .testing }),
1046 );
1047 }
1048 }
1049
1050 test "kernel library catalog builds spatial descriptor artifact registry" {
1051 const allocator = std.testing.allocator;
1052 var state = gpu.recording.BackendState{
1053 .allocator = allocator,
1054 .kind = .cuda,
1055 .format = .cuda_ptx,
1056 };
1057
1058 var descriptors_owned: [3]OwnedDescriptor = undefined;
1059 descriptors_owned[0] = (try selectOwnedSpatial(allocator, .{
1060 .dtype = .f32,
1061 .kind = .grid_cells,
1062 .count = 5000,
1063 .schedule = .{ .thread_blocks = 64 },
1064 })) orelse return error.TestExpectedSpatialDescriptor;
1065 var owned_count: usize = 1;
1066 defer for (descriptors_owned[0..owned_count]) |*descriptor| descriptor.deinit();
1067
1068 descriptors_owned[1] = (try selectOwnedSpatial(allocator, .{
1069 .dtype = .i32,
1070 .kind = .{ .grid_count = .{ .cells = 64 } },
1071 .count = 5000,
1072 .schedule = .{ .thread_blocks = 64 },
1073 })) orelse return error.TestExpectedSpatialDescriptor;
1074 owned_count = 2;
1075
1076 descriptors_owned[2] = (try selectOwnedSpatial(allocator, .{
1077 .dtype = .f32,
1078 .kind = .{ .grid_neighbor_count = .{ .cells = 64, .stride = 79 } },
1079 .count = 5000,
1080 .schedule = .{ .thread_blocks = 64 },
1081 })) orelse return error.TestExpectedSpatialDescriptor;
1082 owned_count = 3;
1083
1084 var owned_registry = try createOwnedKernelCallArtifactRegistry(
1085 allocator,
1086 state.handle(),
1087 descriptors_owned[0..],
1088 .{ .limits = .testing },
1089 );
1090 defer owned_registry.deinit();
1091 const registry = owned_registry.registry();
1092 try std.testing.expectEqual(@as(usize, 3), registry.entries.len);
1093
1094 const expected_scalars = [_]u32{ 6, 2, 8 };
1095 const versions = [_]u32{
1096 spatial_mod.grid_cells_family_version,
1097 spatial_mod.grid_count_family_version,
1098 spatial_mod.grid_neighbor_count_family_version,
1099 };
1100 for (descriptors_owned[0..], expected_scalars, versions) |descriptor, scalars, version| {
1101 const found = registry.find(
1102 descriptor.descriptor.metadata.target,
1103 version,
1104 .cuda_ptx,
1105 ) orelse return error.TestExpectedKernelCallArtifact;
1106 try std.testing.expectEqual(scalars, found.runtime_scalar_argument_count);
1107 const launch = switch (found.launch) {
1108 .derived => |derived| derived,
1109 .fixed => return error.TestExpectedDerivedLaunch,
1110 };
1111 try std.testing.expectEqual(@as(u32, 64), launch.threadgroup[0]);
1112 const profile = found.shape_profile orelse return error.TestExpectedShapeProfile;
1113 try std.testing.expect(profile.dimensions.len >= 1);
1114 }
1115 }
1116
1117 test "kernel library catalog rejects malformed spatial descriptor artifacts" {
1118 const allocator = std.testing.allocator;
1119 var state = gpu.recording.BackendState{
1120 .allocator = allocator,
1121 .kind = .cuda,
1122 .format = .cuda_ptx,
1123 };
1124
1125 {
1126 var descriptor = (try selectOwnedSpatial(allocator, .{
1127 .dtype = .f32,
1128 .kind = .grid_cells,
1129 .count = 5000,
1130 .schedule = .{ .thread_blocks = 64 },
1131 })) orelse return error.TestExpectedSpatialDescriptor;
1132 defer descriptor.deinit();
1133
1134 if (descriptor.specialization) |*specialization| {
1135 const lifetime_allocator = specialization.allocator();
1136 const inputs = try lifetime_allocator.alloc(entry.Shape, 2);
1137 inputs[0] = try entry.runtimeShape1D(lifetime_allocator, "bad_p", 5000);
1138 inputs[1] = specialization.value.inputs[1];
1139 specialization.value.inputs = inputs;
1140 descriptor.descriptor.metadata.specialization = specialization.value;
1141 } else return error.TestExpectedOwnedSpecialization;
1142
1143 try std.testing.expectError(
1144 error.UnknownKernelLibraryEntry,
1145 createOwnedKernelCallArtifact(allocator, state.handle(), descriptor, .{ .limits = .testing }),
1146 );
1147 }
1148
1149 {
1150 var descriptor = (try selectOwnedSpatial(allocator, .{
1151 .dtype = .i32,
1152 .kind = .{ .grid_count = .{ .cells = 64 } },
1153 .count = 5000,
1154 .schedule = .{ .thread_blocks = 64 },
1155 })) orelse return error.TestExpectedSpatialDescriptor;
1156 defer descriptor.deinit();
1157
1158 if (descriptor.specialization) |*specialization| {
1159 const lifetime_allocator = specialization.allocator();
1160 const outputs = try lifetime_allocator.alloc(entry.Shape, 1);
1161 outputs[0] = try entry.runtimeShape2D(lifetime_allocator, spatial_mod.grid_count_cell_axis, 64, "bad_b", spatial_mod.gridCellsBlockCount(5000, 64));
1162 specialization.value.outputs = outputs;
1163 descriptor.descriptor.metadata.specialization = specialization.value;
1164 } else return error.TestExpectedOwnedSpecialization;
1165
1166 try std.testing.expectError(
1167 error.UnknownKernelLibraryEntry,
1168 createOwnedKernelCallArtifact(allocator, state.handle(), descriptor, .{ .limits = .testing }),
1169 );
1170 }
1171
1172 {
1173 var descriptor = (try selectOwnedSpatial(allocator, .{
1174 .dtype = .f32,
1175 .kind = .{ .grid_neighbor_count = .{ .cells = 64, .stride = 79 } },
1176 .count = 5000,
1177 .schedule = .{ .thread_blocks = 64 },
1178 })) orelse return error.TestExpectedSpatialDescriptor;
1179 defer descriptor.deinit();
1180
1181 if (descriptor.specialization) |*specialization| {
1182 const lifetime_allocator = specialization.allocator();
1183 const inputs = try lifetime_allocator.alloc(entry.Shape, 4);
1184 inputs[0] = specialization.value.inputs[0];
1185 inputs[1] = specialization.value.inputs[1];
1186 inputs[2] = try entry.runtimeShape1D(lifetime_allocator, "p", 4999);
1187 inputs[3] = specialization.value.inputs[3];
1188 specialization.value.inputs = inputs;
1189 descriptor.descriptor.metadata.specialization = specialization.value;
1190 } else return error.TestExpectedOwnedSpecialization;
1191
1192 try std.testing.expectError(
1193 error.UnknownKernelLibraryEntry,
1194 createOwnedKernelCallArtifact(allocator, state.handle(), descriptor, .{ .limits = .testing }),
1195 );
1196 }
1197 }
1198
1199 test "kernel library catalog builds static descriptor artifact registry" {
1200 const allocator = std.testing.allocator;
1201 var state = gpu.recording.BackendState{
1202 .allocator = allocator,
1203 .kind = .cuda,
1204 .format = .cuda_ptx,
1205 };
1206 const descriptor = findEntry(elementwise.VectorAdd8F32.target, elementwise.VectorAdd8F32.version) orelse {
1207 return error.TestExpectedCatalogDescriptor;
1208 };
1209 const owned = [_]OwnedDescriptor{.{ .descriptor = descriptor }};
1210 var registry = try createOwnedKernelCallArtifactRegistry(allocator, state.handle(), owned[0..], .{ .limits = .testing });
1211 defer registry.deinit();
1212 const artifact_entry = registry.registry().find(elementwise.VectorAdd8F32.target, elementwise.VectorAdd8F32.version, .cuda_ptx) orelse {
1213 return error.TestExpectedKernelCallArtifact;
1214 };
1215
1216 try std.testing.expectEqual(@as(usize, 1), registry.registry().entries.len);
1217 try std.testing.expectEqualStrings(elementwise.VectorAdd8F32.target, artifact_entry.target);
1218 try std.testing.expectEqualStrings(elementwise.VectorAdd8F32.name, artifact_entry.entry_name);
1219 }
1220 };
1221
1222 const variants_tests = struct {
1223 test "kernel library catalog matrix product candidates preserve explicit schedule" {
1224 const lhs_dims = [_]i64{ 5, 3 };
1225 const rhs_dims = [_]i64{ 3, 7 };
1226 const output_dims = [_]i64{ 5, 7 };
1227 const schedule = MatrixProductSchedule{ .thread_blocks = .{ .x = 4, .y = 2 } };
1228 var candidates = try selectOwnedMatrixProductCandidates(std.testing.allocator, .{
1229 .dtype = .f32,
1230 .lhs_indices = "mk",
1231 .rhs_indices = "kn",
1232 .output_indices = "mn",
1233 .lhs_dims = &lhs_dims,
1234 .rhs_dims = &rhs_dims,
1235 .output_dims = &output_dims,
1236 .schedule = schedule,
1237 });
1238 defer candidates.deinit();
1239
1240 try std.testing.expectEqual(@as(usize, 1), candidates.count);
1241 try std.testing.expectEqualStrings("accy.kernel.linalg.matmul_family_4x2_f32", candidates.items[0].descriptor.metadata.target);
1242 try std.testing.expect(matrixProductDescriptorMatches(candidates.items[0].descriptor, .{ .m = 5, .n = 7, .k = 3 }, schedule, .f32));
1243 }
1244
1245 test "kernel library catalog selects matrix product family schedule specialization" {
1246 const lhs_dims = [_]i64{ 5, 3 };
1247 const rhs_dims = [_]i64{ 3, 7 };
1248 const output_dims = [_]i64{ 5, 7 };
1249 const schedule = MatrixProductSchedule{ .thread_blocks = .{ .x = 4, .y = 2 } };
1250 var selected = (try selectOwned(std.testing.allocator, .{ .matrix_product = .{
1251 .dtype = .f32,
1252 .lhs_indices = "mk",
1253 .rhs_indices = "kn",
1254 .output_indices = "mn",
1255 .lhs_dims = &lhs_dims,
1256 .rhs_dims = &rhs_dims,
1257 .output_dims = &output_dims,
1258 .schedule = schedule,
1259 } })) orelse return error.TestExpectedMatrixProductFamilyScheduleSelection;
1260 defer selected.deinit();
1261 const specialization = selected.descriptor.metadata.specialization;
1262
1263 try std.testing.expectEqualStrings("accy.kernel.linalg.matmul_family_4x2_f32", selected.descriptor.metadata.target);
1264 try std.testing.expect(matrixProductDescriptorMatches(selected.descriptor, .{ .m = 5, .n = 7, .k = 3 }, schedule, .f32));
1265 try std.testing.expect(specialization.shape_family != null);
1266 try std.testing.expectEqual(@as(u32, 2), specialization.launch.?.grid[0]);
1267 try std.testing.expectEqual(@as(u32, 3), specialization.launch.?.grid[1]);
1268 try std.testing.expectEqual(@as(u32, 4), specialization.launch.?.threadgroup[0]);
1269 try std.testing.expectEqual(@as(u32, 2), specialization.launch.?.threadgroup[1]);
1270 }
1271
1272 test "kernel library catalog rejects unavailable matrix product family variants" {
1273 const lhs_dims = [_]i64{ 5, 3 };
1274 const rhs_dims = [_]i64{ 3, 7 };
1275 const output_dims = [_]i64{ 5, 7 };
1276 const bad_output_dims = [_]i64{ 5, 8 };
1277
1278 try std.testing.expect((try selectOwned(std.testing.allocator, .{ .matrix_product = .{
1279 .dtype = .bf16,
1280 .lhs_indices = "mk",
1281 .rhs_indices = "kn",
1282 .output_indices = "mn",
1283 .lhs_dims = &lhs_dims,
1284 .rhs_dims = &rhs_dims,
1285 .output_dims = &output_dims,
1286 } })) == null);
1287 try std.testing.expect((try selectOwned(std.testing.allocator, .{ .matrix_product = .{
1288 .dtype = .f32,
1289 .lhs_indices = "mk",
1290 .rhs_indices = "kn",
1291 .output_indices = "mn",
1292 .lhs_dims = &lhs_dims,
1293 .rhs_dims = &rhs_dims,
1294 .output_dims = &bad_output_dims,
1295 } })) == null);
1296 try std.testing.expect((try selectOwned(std.testing.allocator, .{ .matrix_product = .{
1297 .dtype = .f32,
1298 .lhs_indices = "mk",
1299 .rhs_indices = "kn",
1300 .output_indices = "mn",
1301 .lhs_dims = &lhs_dims,
1302 .rhs_dims = &rhs_dims,
1303 .output_dims = &output_dims,
1304 .schedule = .{ .thread_blocks = .{ .x = 0, .y = 2 } },
1305 } })) == null);
1306 }
1307
1308 test "kernel library catalog selects batched matrix product family for fresh extents" {
1309 const lhs_dims = [_]i64{ 3, 5, 4 };
1310 const rhs_dims = [_]i64{ 3, 4, 6 };
1311 const output_dims = [_]i64{ 3, 5, 6 };
1312 var selected = (try selectOwned(std.testing.allocator, .{ .batched_matrix_product = .{
1313 .dtype = .f32,
1314 .lhs_indices = "bmk",
1315 .rhs_indices = "bkn",
1316 .output_indices = "bmn",
1317 .lhs_dims = &lhs_dims,
1318 .rhs_dims = &rhs_dims,
1319 .output_dims = &output_dims,
1320 } })) orelse return error.TestExpectedBatchedMatrixProductFamilySelection;
1321 defer selected.deinit();
1322 const specialization = selected.descriptor.metadata.specialization;
1323
1324 try std.testing.expectEqualStrings("accy.kernel.linalg.batched_matmul_family_6x5x3_f32", selected.descriptor.metadata.target);
1325 try std.testing.expect(specialization.shape_family != null);
1326 try std.testing.expect(specialization.operationIs(.{ .linalg = .batched_matrix_product }));
1327 try std.testing.expect(specialization.inputHasExtents(0, &.{ 3, 5, 4 }));
1328 try std.testing.expect(specialization.inputHasExtents(1, &.{ 3, 4, 6 }));
1329 try std.testing.expect(specialization.outputHasExtents(0, &.{ 3, 5, 6 }));
1330 try std.testing.expect(specialization.reductionMatches(0, .{ .name = "dot", .operator = .dot_product, .extents = &.{4} }));
1331 try std.testing.expectEqual(@as(u32, 6), specialization.launch.?.threadgroup[0]);
1332 try std.testing.expectEqual(@as(u32, 5), specialization.launch.?.threadgroup[1]);
1333 try std.testing.expectEqual(@as(u32, 3), specialization.launch.?.threadgroup[2]);
1334 }
1335
1336 test "kernel library catalog creates batched matrix product family artifact from owned descriptor" {
1337 const lhs_dims = [_]i64{ 3, 5, 4 };
1338 const rhs_dims = [_]i64{ 3, 4, 6 };
1339 const output_dims = [_]i64{ 3, 5, 6 };
1340 var selected = (try selectOwned(std.testing.allocator, .{ .batched_matrix_product = .{
1341 .dtype = .f32,
1342 .lhs_indices = "bmk",
1343 .rhs_indices = "bkn",
1344 .output_indices = "bmn",
1345 .lhs_dims = &lhs_dims,
1346 .rhs_dims = &rhs_dims,
1347 .output_dims = &output_dims,
1348 .schedule = .{ .thread_blocks = .{ .x = 4, .y = 2, .z = 2 } },
1349 } })) orelse return error.TestExpectedBatchedMatrixProductFamilyScheduleSelection;
1350 defer selected.deinit();
1351
1352 var state = gpu.recording.BackendState{
1353 .allocator = std.testing.allocator,
1354 .kind = .cuda,
1355 .format = .cuda_ptx,
1356 };
1357 var artifact = try createOwnedKernelCallArtifact(std.testing.allocator, state.handle(), selected, .{ .limits = .testing });
1358 defer artifact.deinit();
1359 const entry_value = artifact.entry();
1360
1361 try std.testing.expectEqualStrings("accy.kernel.linalg.batched_matmul_family_4x2x2_f32", selected.descriptor.metadata.target);
1362 try std.testing.expectEqualStrings("accy.kernel.linalg.batched_matmul_family_4x2x2_f32", entry_value.target);
1363 try std.testing.expectEqual(@as(u32, 7), entry_value.argument_count);
1364 try std.testing.expectEqual(@as(u32, 4), entry_value.runtime_scalar_argument_count);
1365 switch (entry_value.launch) {
1366 .derived => |launch| {
1367 try std.testing.expectEqual(@as(u32, 4), launch.threadgroup[0]);
1368 try std.testing.expectEqual(@as(u32, 2), launch.threadgroup[1]);
1369 try std.testing.expectEqual(@as(u32, 2), launch.threadgroup[2]);
1370 },
1371 else => return error.TestExpectedDerivedLaunch,
1372 }
1373 }
1374
1375 test "kernel library catalog selects matrix vector product family by canonical einsum roles" {
1376 const matrix_dims = [_]i64{ 4, 8 };
1377 const vector_dims = [_]i64{8};
1378 const output_dims = [_]i64{4};
1379 const query = Query{ .matrix_vector_product = .{
1380 .dtype = .f32,
1381 .matrix_indices = "mk",
1382 .vector_indices = "k",
1383 .output_indices = "m",
1384 .matrix_dims = &matrix_dims,
1385 .vector_dims = &vector_dims,
1386 .output_dims = &output_dims,
1387 } };
1388 try std.testing.expect(select(query) == null);
1389 var selected = (try selectOwned(std.testing.allocator, query)) orelse return error.TestExpectedMatrixVectorProductSelection;
1390 defer selected.deinit();
1391 const specialization = selected.descriptor.metadata.specialization;
1392
1393 try std.testing.expect(selected.specialization != null);
1394 try std.testing.expectEqualStrings("accy.kernel.linalg.matvec_family_4x_f32", selected.descriptor.metadata.target);
1395 try std.testing.expectEqualStrings("accy_kernel_linalg_matvec_family_4x_f32", selected.descriptor.name);
1396 try std.testing.expect(specialization.operationIs(.{ .linalg = .matrix_vector_product }));
1397 try std.testing.expect(specialization.shape_family != null);
1398 try std.testing.expectEqual(@as(u32, 4), specialization.launch.?.threadgroup[0]);
1399 }
1400
1401 test "kernel library catalog selects outer product family by canonical einsum roles" {
1402 const lhs_dims = [_]i64{4};
1403 const rhs_dims = [_]i64{3};
1404 const output_dims = [_]i64{ 4, 3 };
1405 const query = Query{ .outer_product = .{
1406 .dtype = .f32,
1407 .lhs_indices = "m",
1408 .rhs_indices = "n",
1409 .output_indices = "mn",
1410 .lhs_dims = &lhs_dims,
1411 .rhs_dims = &rhs_dims,
1412 .output_dims = &output_dims,
1413 } };
1414 try std.testing.expect(select(query) == null);
1415 var selected = (try selectOwned(std.testing.allocator, query)) orelse return error.TestExpectedOuterProductSelection;
1416 defer selected.deinit();
1417 const specialization = selected.descriptor.metadata.specialization;
1418
1419 try std.testing.expect(selected.specialization != null);
1420 try std.testing.expectEqualStrings("accy.kernel.linalg.outer_family_3x4_f32", selected.descriptor.metadata.target);
1421 try std.testing.expectEqualStrings("accy_kernel_linalg_outer_family_3x4_f32", selected.descriptor.name);
1422 try std.testing.expect(specialization.operationIs(.{ .linalg = .outer_product }));
1423 try std.testing.expect(specialization.shape_family != null);
1424 try std.testing.expectEqual(@as(usize, 0), specialization.reductions.len);
1425 try std.testing.expectEqual(@as(u32, 3), specialization.launch.?.threadgroup[0]);
1426 try std.testing.expectEqual(@as(u32, 4), specialization.launch.?.threadgroup[1]);
1427 }
1428
1429 test "kernel library catalog selects outer product family schedule specialization" {
1430 const lhs_dims = [_]i64{4};
1431 const rhs_dims = [_]i64{3};
1432 const output_dims = [_]i64{ 4, 3 };
1433 const query = Query{ .outer_product = .{
1434 .dtype = .f32,
1435 .lhs_indices = "m",
1436 .rhs_indices = "n",
1437 .output_indices = "mn",
1438 .lhs_dims = &lhs_dims,
1439 .rhs_dims = &rhs_dims,
1440 .output_dims = &output_dims,
1441 .schedule = .{ .thread_blocks = .{ .x = 3, .y = 2 } },
1442 } };
1443 try std.testing.expect(select(query) == null);
1444 var selected = (try selectOwned(std.testing.allocator, query)) orelse return error.TestExpectedOuterProductScheduleSelection;
1445 defer selected.deinit();
1446 const specialization = selected.descriptor.metadata.specialization;
1447
1448 try std.testing.expect(selected.specialization != null);
1449 try std.testing.expectEqualStrings("accy.kernel.linalg.outer_family_3x2_f32", selected.descriptor.metadata.target);
1450 try std.testing.expectEqual(@as(u32, 3), specialization.launch.?.threadgroup[0]);
1451 try std.testing.expectEqual(@as(u32, 2), specialization.launch.?.threadgroup[1]);
1452 }
1453
1454 test "kernel library catalog selects outer product family for fresh extents" {
1455 const lhs_dims = [_]i64{5};
1456 const rhs_dims = [_]i64{6};
1457 const output_dims = [_]i64{ 5, 6 };
1458 var selected = (try selectOwned(std.testing.allocator, .{ .outer_product = .{
1459 .dtype = .f32,
1460 .lhs_indices = "m",
1461 .rhs_indices = "n",
1462 .output_indices = "mn",
1463 .lhs_dims = &lhs_dims,
1464 .rhs_dims = &rhs_dims,
1465 .output_dims = &output_dims,
1466 } })) orelse return error.TestExpectedOuterProductFamilySelection;
1467 defer selected.deinit();
1468 const specialization = selected.descriptor.metadata.specialization;
1469
1470 try std.testing.expectEqualStrings("accy.kernel.linalg.outer_family_6x5_f32", selected.descriptor.metadata.target);
1471 try std.testing.expect(specialization.shape_family != null);
1472 try std.testing.expect(specialization.operationIs(.{ .linalg = .outer_product }));
1473 try std.testing.expectEqual(@as(usize, 0), specialization.reductions.len);
1474 try std.testing.expect(specialization.inputHasExtents(0, &.{5}));
1475 try std.testing.expect(specialization.inputHasExtents(1, &.{6}));
1476 try std.testing.expect(specialization.outputHasExtents(0, &.{ 5, 6 }));
1477 try std.testing.expectEqual(@as(u32, 6), specialization.launch.?.threadgroup[0]);
1478 try std.testing.expectEqual(@as(u32, 5), specialization.launch.?.threadgroup[1]);
1479 }
1480
1481 test "kernel library catalog creates outer product family artifact from owned descriptor" {
1482 const lhs_dims = [_]i64{5};
1483 const rhs_dims = [_]i64{6};
1484 const output_dims = [_]i64{ 5, 6 };
1485 var selected = (try selectOwned(std.testing.allocator, .{ .outer_product = .{
1486 .dtype = .f32,
1487 .lhs_indices = "m",
1488 .rhs_indices = "n",
1489 .output_indices = "mn",
1490 .lhs_dims = &lhs_dims,
1491 .rhs_dims = &rhs_dims,
1492 .output_dims = &output_dims,
1493 .schedule = .{ .thread_blocks = .{ .x = 4, .y = 2 } },
1494 } })) orelse return error.TestExpectedOuterProductFamilyScheduleSelection;
1495 defer selected.deinit();
1496
1497 var state = gpu.recording.BackendState{
1498 .allocator = std.testing.allocator,
1499 .kind = .cuda,
1500 .format = .cuda_ptx,
1501 };
1502 var artifact = try createOwnedKernelCallArtifact(std.testing.allocator, state.handle(), selected, .{ .limits = .testing });
1503 defer artifact.deinit();
1504 const entry_value = artifact.entry();
1505
1506 try std.testing.expectEqualStrings("accy.kernel.linalg.outer_family_4x2_f32", selected.descriptor.metadata.target);
1507 try std.testing.expectEqualStrings("accy.kernel.linalg.outer_family_4x2_f32", entry_value.target);
1508 try std.testing.expectEqual(@as(u32, 5), entry_value.argument_count);
1509 try std.testing.expectEqual(@as(u32, 2), entry_value.runtime_scalar_argument_count);
1510 switch (entry_value.launch) {
1511 .derived => |launch| {
1512 try std.testing.expectEqual(@as(u32, 4), launch.threadgroup[0]);
1513 try std.testing.expectEqual(@as(u32, 2), launch.threadgroup[1]);
1514 },
1515 else => return error.TestExpectedDerivedLaunch,
1516 }
1517 }
1518
1519 test "kernel library catalog selects transpose by canonical einsum roles" {
1520 const input_dims = [_]i64{ 8, 16 };
1521 const output_dims = [_]i64{ 16, 8 };
1522 const selected = select(.{ .layout = .{
1523 .dtype = .f32,
1524 .kind = .transpose,
1525 .input_indices = "ij",
1526 .output_indices = "ji",
1527 .input_dims = &input_dims,
1528 .output_dims = &output_dims,
1529 } }) orelse return error.TestExpectedTransposeSelection;
1530
1531 try std.testing.expectEqualStrings(layout.Transpose8x16F32.target, selected.metadata.target);
1532 try std.testing.expectEqualStrings(layout.Transpose8x16F32.name, selected.name);
1533 }
1534 };
1535
1536 const fused_tests = struct {
1537 test "kernel library catalog selects activation by operation contract" {
1538 const gelu = select(.{ .activation = .{
1539 .dtype = .f32,
1540 .kind = .gelu,
1541 .extent = 8,
1542 } }) orelse return error.TestExpectedGeluSelection;
1543 const relu = select(.{ .activation = .{
1544 .dtype = .f32,
1545 .kind = .relu,
1546 .extent = 8,
1547 } }) orelse return error.TestExpectedReluSelection;
1548 const silu = select(.{ .activation = .{
1549 .dtype = .f32,
1550 .kind = .silu,
1551 .extent = 8,
1552 } }) orelse return error.TestExpectedSiluSelection;
1553
1554 try std.testing.expectEqualStrings(elementwise.Gelu8F32.target, gelu.metadata.target);
1555 try std.testing.expectEqualStrings(elementwise.Gelu8F32.name, gelu.name);
1556 try std.testing.expectEqualStrings(elementwise.Relu8F32.target, relu.metadata.target);
1557 try std.testing.expectEqualStrings(elementwise.Relu8F32.name, relu.name);
1558 try std.testing.expectEqualStrings(elementwise.Silu8F32.target, silu.metadata.target);
1559 try std.testing.expectEqualStrings(elementwise.Silu8F32.name, silu.name);
1560 }
1561
1562 test "kernel library catalog selects fused bias gelu by operation contract" {
1563 const selected = select(.{ .fused = .{ .vector = .{
1564 .dtype = .f32,
1565 .kind = .{ .bias_activation = .gelu },
1566 .extent = 8,
1567 } } }) orelse return error.TestExpectedBiasGeluSelection;
1568
1569 try std.testing.expectEqualStrings(fused.BiasGelu8F32.target, selected.metadata.target);
1570 try std.testing.expectEqualStrings(fused.BiasGelu8F32.name, selected.name);
1571 }
1572
1573 test "kernel library catalog selects fused bias relu and silu by operation contract" {
1574 const relu = select(.{ .fused = .{ .vector = .{
1575 .dtype = .f32,
1576 .kind = .{ .bias_activation = .relu },
1577 .extent = 8,
1578 } } }) orelse return error.TestExpectedBiasReluSelection;
1579 const silu = select(.{ .fused = .{ .vector = .{
1580 .dtype = .f32,
1581 .kind = .{ .bias_activation = .silu },
1582 .extent = 8,
1583 } } }) orelse return error.TestExpectedBiasSiluSelection;
1584
1585 try std.testing.expectEqualStrings(fused.BiasRelu8F32.target, relu.metadata.target);
1586 try std.testing.expectEqualStrings(fused.BiasRelu8F32.name, relu.name);
1587 try std.testing.expectEqualStrings(fused.BiasSilu8F32.target, silu.metadata.target);
1588 try std.testing.expectEqualStrings(fused.BiasSilu8F32.name, silu.name);
1589 }
1590
1591 test "kernel library catalog selects fused swiglu by operation contract" {
1592 const selected = select(.{ .fused = .{ .vector = .{
1593 .dtype = .f32,
1594 .kind = .{ .gated_activation = .silu },
1595 .extent = 8,
1596 } } }) orelse return error.TestExpectedSwiGluSelection;
1597
1598 try std.testing.expectEqualStrings(fused.SwiGlu8F32.target, selected.metadata.target);
1599 try std.testing.expectEqualStrings(fused.SwiGlu8F32.name, selected.name);
1600 }
1601
1602 test "kernel library catalog selects fused geglu by operation contract" {
1603 const selected = select(.{ .fused = .{ .vector = .{
1604 .dtype = .f32,
1605 .kind = .{ .gated_activation = .gelu },
1606 .extent = 8,
1607 } } }) orelse return error.TestExpectedGeGluSelection;
1608
1609 try std.testing.expectEqualStrings(fused.GeGlu8F32.target, selected.metadata.target);
1610 try std.testing.expectEqualStrings(fused.GeGlu8F32.name, selected.name);
1611 }
1612
1613 test "kernel library catalog selects fused reglu by operation contract" {
1614 const selected = select(.{ .fused = .{ .vector = .{
1615 .dtype = .f32,
1616 .kind = .{ .gated_activation = .relu },
1617 .extent = 8,
1618 } } }) orelse return error.TestExpectedReGluSelection;
1619
1620 try std.testing.expectEqualStrings(fused.ReGlu8F32.target, selected.metadata.target);
1621 try std.testing.expectEqualStrings(fused.ReGlu8F32.name, selected.name);
1622 }
1623
1624 test "kernel library catalog selects fused matrix product bias gelu by operation contract" {
1625 const lhs_dims = [_]i64{ 2, 4 };
1626 const rhs_dims = [_]i64{ 4, 3 };
1627 const bias_dims = [_]i64{3};
1628 const output_dims = [_]i64{ 2, 3 };
1629 const selected = select(.{ .fused = .{ .matrix_product = .{
1630 .dtype = .f32,
1631 .epilogue = .{ .bias_activation = .gelu },
1632 .lhs_indices = "mk",
1633 .rhs_indices = "kn",
1634 .bias_indices = "n",
1635 .output_indices = "mn",
1636 .lhs_dims = &lhs_dims,
1637 .rhs_dims = &rhs_dims,
1638 .bias_dims = &bias_dims,
1639 .output_dims = &output_dims,
1640 } } }) orelse return error.TestExpectedMatrixProductBiasGeluSelection;
1641
1642 try std.testing.expectEqualStrings(fused.MatrixProductBiasGelu2x3x4F32.target, selected.metadata.target);
1643 try std.testing.expectEqualStrings(fused.MatrixProductBiasGelu2x3x4F32.name, selected.name);
1644 }
1645
1646 test "kernel library catalog selects fused matrix product schedule specialization" {
1647 const lhs_dims = [_]i64{ 2, 4 };
1648 const rhs_dims = [_]i64{ 4, 3 };
1649 const bias_dims = [_]i64{3};
1650 const output_dims = [_]i64{ 2, 3 };
1651 const default_selected = select(.{ .fused = .{ .matrix_product = .{
1652 .dtype = .f32,
1653 .epilogue = .{ .bias_activation = .gelu },
1654 .lhs_indices = "mk",
1655 .rhs_indices = "kn",
1656 .bias_indices = "n",
1657 .output_indices = "mn",
1658 .lhs_dims = &lhs_dims,
1659 .rhs_dims = &rhs_dims,
1660 .bias_dims = &bias_dims,
1661 .output_dims = &output_dims,
1662 } } }) orelse return error.TestExpectedMatrixProductBiasGeluSelection;
1663 const scheduled = select(.{ .fused = .{ .matrix_product = .{
1664 .dtype = .f32,
1665 .epilogue = .{ .bias_activation = .gelu },
1666 .lhs_indices = "mk",
1667 .rhs_indices = "kn",
1668 .bias_indices = "n",
1669 .output_indices = "mn",
1670 .lhs_dims = &lhs_dims,
1671 .rhs_dims = &rhs_dims,
1672 .bias_dims = &bias_dims,
1673 .output_dims = &output_dims,
1674 .schedule = .{ .thread_blocks = .{ .x = 1, .y = 2 } },
1675 } } }) orelse return error.TestExpectedMatrixProductBiasGeluSelection;
1676
1677 try std.testing.expectEqualStrings(fused.MatrixProductBiasGelu2x3x4F32.target, default_selected.metadata.target);
1678 try std.testing.expectEqualStrings(fused.MatrixProductBiasGelu2x3x4ThreadBlocks1x2F32.target, scheduled.metadata.target);
1679 try std.testing.expectEqual(@as(u32, 1), scheduled.metadata.specialization.launch.?.threadgroup[0]);
1680 try std.testing.expectEqual(@as(u32, 2), scheduled.metadata.specialization.launch.?.threadgroup[1]);
1681 }
1682
1683 test "kernel library catalog selects fused matrix product bias relu and silu by operation contract" {
1684 const lhs_dims = [_]i64{ 2, 4 };
1685 const rhs_dims = [_]i64{ 4, 3 };
1686 const bias_dims = [_]i64{3};
1687 const output_dims = [_]i64{ 2, 3 };
1688 const relu = select(.{ .fused = .{ .matrix_product = .{
1689 .dtype = .f32,
1690 .epilogue = .{ .bias_activation = .relu },
1691 .lhs_indices = "mk",
1692 .rhs_indices = "kn",
1693 .bias_indices = "n",
1694 .output_indices = "mn",
1695 .lhs_dims = &lhs_dims,
1696 .rhs_dims = &rhs_dims,
1697 .bias_dims = &bias_dims,
1698 .output_dims = &output_dims,
1699 } } }) orelse return error.TestExpectedMatrixProductBiasReluSelection;
1700 const silu = select(.{ .fused = .{ .matrix_product = .{
1701 .dtype = .f32,
1702 .epilogue = .{ .bias_activation = .silu },
1703 .lhs_indices = "mk",
1704 .rhs_indices = "kn",
1705 .bias_indices = "n",
1706 .output_indices = "mn",
1707 .lhs_dims = &lhs_dims,
1708 .rhs_dims = &rhs_dims,
1709 .bias_dims = &bias_dims,
1710 .output_dims = &output_dims,
1711 } } }) orelse return error.TestExpectedMatrixProductBiasSiluSelection;
1712
1713 try std.testing.expectEqualStrings(fused.MatrixProductBiasRelu2x3x4F32.target, relu.metadata.target);
1714 try std.testing.expectEqualStrings(fused.MatrixProductBiasRelu2x3x4F32.name, relu.name);
1715 try std.testing.expectEqualStrings(fused.MatrixProductBiasSilu2x3x4F32.target, silu.metadata.target);
1716 try std.testing.expectEqualStrings(fused.MatrixProductBiasSilu2x3x4F32.name, silu.name);
1717 }
1718
1719 test "kernel library catalog selects fused matrix vector product bias gelu by operation contract" {
1720 const matrix_dims = [_]i64{ 4, 8 };
1721 const vector_dims = [_]i64{8};
1722 const bias_dims = [_]i64{4};
1723 const output_dims = [_]i64{4};
1724 const selected = select(.{ .fused = .{ .matrix_vector_product = .{
1725 .dtype = .f32,
1726 .epilogue = .{ .bias_activation = .gelu },
1727 .matrix_indices = "mk",
1728 .vector_indices = "k",
1729 .bias_indices = "m",
1730 .output_indices = "m",
1731 .matrix_dims = &matrix_dims,
1732 .vector_dims = &vector_dims,
1733 .bias_dims = &bias_dims,
1734 .output_dims = &output_dims,
1735 } } }) orelse return error.TestExpectedMatrixVectorProductBiasGeluSelection;
1736
1737 try std.testing.expectEqualStrings(fused.MatrixVectorProductBiasGelu4x8F32.target, selected.metadata.target);
1738 try std.testing.expectEqualStrings(fused.MatrixVectorProductBiasGelu4x8F32.name, selected.name);
1739 }
1740
1741 test "kernel library catalog selects fused matrix vector product bias relu and silu by operation contract" {
1742 const matrix_dims = [_]i64{ 4, 8 };
1743 const vector_dims = [_]i64{8};
1744 const bias_dims = [_]i64{4};
1745 const output_dims = [_]i64{4};
1746 const relu = select(.{ .fused = .{ .matrix_vector_product = .{
1747 .dtype = .f32,
1748 .epilogue = .{ .bias_activation = .relu },
1749 .matrix_indices = "mk",
1750 .vector_indices = "k",
1751 .bias_indices = "m",
1752 .output_indices = "m",
1753 .matrix_dims = &matrix_dims,
1754 .vector_dims = &vector_dims,
1755 .bias_dims = &bias_dims,
1756 .output_dims = &output_dims,
1757 } } }) orelse return error.TestExpectedMatrixVectorProductBiasReluSelection;
1758 const silu = select(.{ .fused = .{ .matrix_vector_product = .{
1759 .dtype = .f32,
1760 .epilogue = .{ .bias_activation = .silu },
1761 .matrix_indices = "mk",
1762 .vector_indices = "k",
1763 .bias_indices = "m",
1764 .output_indices = "m",
1765 .matrix_dims = &matrix_dims,
1766 .vector_dims = &vector_dims,
1767 .bias_dims = &bias_dims,
1768 .output_dims = &output_dims,
1769 } } }) orelse return error.TestExpectedMatrixVectorProductBiasSiluSelection;
1770
1771 try std.testing.expectEqualStrings(fused.MatrixVectorProductBiasRelu4x8F32.target, relu.metadata.target);
1772 try std.testing.expectEqualStrings(fused.MatrixVectorProductBiasRelu4x8F32.name, relu.name);
1773 try std.testing.expectEqualStrings(fused.MatrixVectorProductBiasSilu4x8F32.target, silu.metadata.target);
1774 try std.testing.expectEqualStrings(fused.MatrixVectorProductBiasSilu4x8F32.name, silu.name);
1775 }
1776
1777 test "kernel library catalog selects fused row residual rmsnorm by operation contract" {
1778 const selected = select(.{ .fused = .{ .row_normalization = .{
1779 .dtype = .f32,
1780 .kind = .{ .residual = .{ .rmsnorm = .scale } },
1781 .rows = 2,
1782 .cols = 4,
1783 } } }) orelse return error.TestExpectedRowResidualRmsNormSelection;
1784
1785 try std.testing.expectEqualStrings(normalization.RowResidualRmsNorm2x4F32.target, selected.metadata.target);
1786 try std.testing.expectEqualStrings(normalization.RowResidualRmsNorm2x4F32.name, selected.name);
1787 }
1788 };
1789
1790 const reject_tests = struct {
1791 test "kernel library catalog rejects unavailable scalar reductions" {
1792 const input_dims = [_]i64{8};
1793 const wrong_dims = [_]i64{16};
1794 const output_dims = [_]i64{};
1795 const vector_output_dims = [_]i64{1};
1796 const sum_inputs = [_]ReductionOperand{.{ .indices = "i", .dims = &input_dims }};
1797 const dot_inputs = [_]ReductionOperand{
1798 .{ .indices = "i", .dims = &input_dims },
1799 .{ .indices = "i", .dims = &input_dims },
1800 };
1801 const mismatched_dot_inputs = [_]ReductionOperand{
1802 .{ .indices = "i", .dims = &input_dims },
1803 .{ .indices = "i", .dims = &wrong_dims },
1804 };
1805 const wrong_axis_dot_inputs = [_]ReductionOperand{
1806 .{ .indices = "i", .dims = &input_dims },
1807 .{ .indices = "j", .dims = &input_dims },
1808 };
1809
1810 try std.testing.expect(select(.{ .reduction = .{
1811 .dtype = .i32,
1812 .kind = .sum,
1813 .inputs = &sum_inputs,
1814 .output_indices = "",
1815 .output_dims = &output_dims,
1816 } }) == null);
1817 try std.testing.expect(select(.{ .reduction = .{
1818 .dtype = .f32,
1819 .kind = .sum,
1820 .inputs = &sum_inputs,
1821 .output_indices = "",
1822 .output_dims = &vector_output_dims,
1823 } }) == null);
1824 try std.testing.expect(select(.{ .reduction = .{
1825 .dtype = .f32,
1826 .kind = .sum,
1827 .inputs = &.{.{ .indices = "ij", .dims = &input_dims }},
1828 .output_indices = "",
1829 .output_dims = &output_dims,
1830 } }) == null);
1831 try std.testing.expect(select(.{ .reduction = .{
1832 .dtype = .f32,
1833 .kind = .dot_product,
1834 .inputs = &dot_inputs,
1835 .output_indices = "i",
1836 .output_dims = &input_dims,
1837 } }) == null);
1838 try std.testing.expect(select(.{ .reduction = .{
1839 .dtype = .f32,
1840 .kind = .dot_product,
1841 .inputs = &mismatched_dot_inputs,
1842 .output_indices = "",
1843 .output_dims = &output_dims,
1844 } }) == null);
1845 try std.testing.expect(select(.{ .reduction = .{
1846 .dtype = .f32,
1847 .kind = .dot_product,
1848 .inputs = &wrong_axis_dot_inputs,
1849 .output_indices = "",
1850 .output_dims = &output_dims,
1851 } }) == null);
1852 try std.testing.expect(select(.{ .reduction = .{
1853 .dtype = .f32,
1854 .kind = .dot_product,
1855 .inputs = &dot_inputs,
1856 .output_indices = "",
1857 .output_dims = &wrong_dims,
1858 } }) == null);
1859 }
1860
1861 test "kernel library catalog rejects unavailable fused variants" {
1862 try std.testing.expect(select(.{ .fused = .{ .vector = .{
1863 .dtype = .i32,
1864 .kind = .{ .bias_activation = .gelu },
1865 .extent = 8,
1866 } } }) == null);
1867 try std.testing.expect(select(.{ .fused = .{ .vector = .{
1868 .dtype = .f32,
1869 .kind = .{ .bias_activation = .gelu },
1870 .extent = 0,
1871 } } }) == null);
1872 try std.testing.expect(select(.{ .fused = .{ .vector = .{
1873 .dtype = .f32,
1874 .kind = .{ .bias_activation = .gelu },
1875 .extent = 16,
1876 } } }) == null);
1877 try std.testing.expect(select(.{ .fused = .{ .vector = .{
1878 .dtype = .i32,
1879 .kind = .{ .gated_activation = .gelu },
1880 .extent = 8,
1881 } } }) == null);
1882 try std.testing.expect(select(.{ .fused = .{ .vector = .{
1883 .dtype = .f32,
1884 .kind = .{ .gated_activation = .gelu },
1885 .extent = 0,
1886 } } }) == null);
1887 try std.testing.expect(select(.{ .fused = .{ .vector = .{
1888 .dtype = .f32,
1889 .kind = .{ .gated_activation = .gelu },
1890 .extent = 16,
1891 } } }) == null);
1892 try std.testing.expect(select(.{ .fused = .{ .vector = .{
1893 .dtype = .i32,
1894 .kind = .{ .gated_activation = .relu },
1895 .extent = 8,
1896 } } }) == null);
1897 try std.testing.expect(select(.{ .fused = .{ .vector = .{
1898 .dtype = .f32,
1899 .kind = .{ .gated_activation = .relu },
1900 .extent = 0,
1901 } } }) == null);
1902 try std.testing.expect(select(.{ .fused = .{ .vector = .{
1903 .dtype = .f32,
1904 .kind = .{ .gated_activation = .relu },
1905 .extent = 16,
1906 } } }) == null);
1907 try std.testing.expect(select(.{ .fused = .{ .vector = .{
1908 .dtype = .i32,
1909 .kind = .{ .gated_activation = .silu },
1910 .extent = 8,
1911 } } }) == null);
1912 try std.testing.expect(select(.{ .fused = .{ .vector = .{
1913 .dtype = .f32,
1914 .kind = .{ .gated_activation = .silu },
1915 .extent = 0,
1916 } } }) == null);
1917 try std.testing.expect(select(.{ .fused = .{ .vector = .{
1918 .dtype = .f32,
1919 .kind = .{ .gated_activation = .silu },
1920 .extent = 16,
1921 } } }) == null);
1922 }
1923
1924 test "kernel library catalog rejects unavailable fused matrix product variants" {
1925 const lhs_dims = [_]i64{ 2, 4 };
1926 const rhs_dims = [_]i64{ 4, 3 };
1927 const bias_dims = [_]i64{3};
1928 const wrong_bias_dims = [_]i64{4};
1929 const unavailable_output_dims = [_]i64{ 3, 3 };
1930 const output_dims = [_]i64{ 2, 3 };
1931
1932 try std.testing.expect(select(.{ .fused = .{ .matrix_product = .{
1933 .dtype = .i32,
1934 .epilogue = .{ .bias_activation = .gelu },
1935 .lhs_indices = "mk",
1936 .rhs_indices = "kn",
1937 .bias_indices = "n",
1938 .output_indices = "mn",
1939 .lhs_dims = &lhs_dims,
1940 .rhs_dims = &rhs_dims,
1941 .bias_dims = &bias_dims,
1942 .output_dims = &output_dims,
1943 } } }) == null);
1944 try std.testing.expect(select(.{ .fused = .{ .matrix_product = .{
1945 .dtype = .f32,
1946 .epilogue = .{ .bias_activation = .gelu },
1947 .lhs_indices = "mk",
1948 .rhs_indices = "kn",
1949 .bias_indices = "m",
1950 .output_indices = "mn",
1951 .lhs_dims = &lhs_dims,
1952 .rhs_dims = &rhs_dims,
1953 .bias_dims = &bias_dims,
1954 .output_dims = &output_dims,
1955 } } }) == null);
1956 try std.testing.expect(select(.{ .fused = .{ .matrix_product = .{
1957 .dtype = .f32,
1958 .epilogue = .{ .bias_activation = .gelu },
1959 .lhs_indices = "mk",
1960 .rhs_indices = "kn",
1961 .bias_indices = "n",
1962 .output_indices = "mn",
1963 .lhs_dims = &lhs_dims,
1964 .rhs_dims = &rhs_dims,
1965 .bias_dims = &wrong_bias_dims,
1966 .output_dims = &output_dims,
1967 } } }) == null);
1968 try std.testing.expect(select(.{ .fused = .{ .matrix_product = .{
1969 .dtype = .f32,
1970 .epilogue = .{ .bias_activation = .gelu },
1971 .lhs_indices = "mk",
1972 .rhs_indices = "kn",
1973 .bias_indices = "n",
1974 .output_indices = "mn",
1975 .lhs_dims = &lhs_dims,
1976 .rhs_dims = &rhs_dims,
1977 .bias_dims = &bias_dims,
1978 .output_dims = &unavailable_output_dims,
1979 } } }) == null);
1980 try std.testing.expect(select(.{ .fused = .{ .matrix_product = .{
1981 .dtype = .f32,
1982 .epilogue = .{ .bias_activation = .gelu },
1983 .lhs_indices = "mk",
1984 .rhs_indices = "kn",
1985 .bias_indices = "n",
1986 .output_indices = "mn",
1987 .lhs_dims = &lhs_dims,
1988 .rhs_dims = &rhs_dims,
1989 .bias_dims = &bias_dims,
1990 .output_dims = &output_dims,
1991 .schedule = .{ .thread_blocks = .{ .x = 8, .y = 2 } },
1992 } } }) == null);
1993 }
1994
1995 test "kernel library catalog rejects unavailable fused matrix vector product variants" {
1996 const matrix_dims = [_]i64{ 4, 8 };
1997 const vector_dims = [_]i64{8};
1998 const bias_dims = [_]i64{4};
1999 const wrong_bias_dims = [_]i64{5};
2000 const unavailable_matrix_dims = [_]i64{ 5, 8 };
2001 const unavailable_bias_dims = [_]i64{5};
2002 const unavailable_output_dims = [_]i64{5};
2003 const output_dims = [_]i64{4};
2004
2005 try std.testing.expect(select(.{ .fused = .{ .matrix_vector_product = .{
2006 .dtype = .i32,
2007 .epilogue = .{ .bias_activation = .gelu },
2008 .matrix_indices = "mk",
2009 .vector_indices = "k",
2010 .bias_indices = "m",
2011 .output_indices = "m",
2012 .matrix_dims = &matrix_dims,
2013 .vector_dims = &vector_dims,
2014 .bias_dims = &bias_dims,
2015 .output_dims = &output_dims,
2016 } } }) == null);
2017 try std.testing.expect(select(.{ .fused = .{ .matrix_vector_product = .{
2018 .dtype = .f32,
2019 .epilogue = .{ .bias_activation = .gelu },
2020 .matrix_indices = "mk",
2021 .vector_indices = "k",
2022 .bias_indices = "k",
2023 .output_indices = "m",
2024 .matrix_dims = &matrix_dims,
2025 .vector_dims = &vector_dims,
2026 .bias_dims = &bias_dims,
2027 .output_dims = &output_dims,
2028 } } }) == null);
2029 try std.testing.expect(select(.{ .fused = .{ .matrix_vector_product = .{
2030 .dtype = .f32,
2031 .epilogue = .{ .bias_activation = .gelu },
2032 .matrix_indices = "mk",
2033 .vector_indices = "k",
2034 .bias_indices = "m",
2035 .output_indices = "m",
2036 .matrix_dims = &matrix_dims,
2037 .vector_dims = &vector_dims,
2038 .bias_dims = &wrong_bias_dims,
2039 .output_dims = &output_dims,
2040 } } }) == null);
2041 try std.testing.expect(select(.{ .fused = .{ .matrix_vector_product = .{
2042 .dtype = .f32,
2043 .epilogue = .{ .bias_activation = .gelu },
2044 .matrix_indices = "mk",
2045 .vector_indices = "k",
2046 .bias_indices = "m",
2047 .output_indices = "m",
2048 .matrix_dims = &unavailable_matrix_dims,
2049 .vector_dims = &vector_dims,
2050 .bias_dims = &unavailable_bias_dims,
2051 .output_dims = &unavailable_output_dims,
2052 } } }) == null);
2053 try std.testing.expect(select(.{ .fused = .{ .matrix_vector_product = .{
2054 .dtype = .f32,
2055 .epilogue = .{ .bias_activation = .gelu },
2056 .matrix_indices = "mk",
2057 .vector_indices = "k",
2058 .bias_indices = "m",
2059 .output_indices = "m",
2060 .matrix_dims = &matrix_dims,
2061 .vector_dims = &vector_dims,
2062 .bias_dims = &bias_dims,
2063 .output_dims = &output_dims,
2064 .schedule = .{ .thread_blocks = 8 },
2065 } } }) == null);
2066 }
2067
2068 test "kernel library catalog rejects unavailable fused row normalization variants" {
2069 try std.testing.expect(select(.{ .fused = .{ .row_normalization = .{
2070 .dtype = .i32,
2071 .kind = .{ .residual = .{ .rmsnorm = .scale } },
2072 .rows = 2,
2073 .cols = 4,
2074 } } }) == null);
2075 try std.testing.expect(select(.{ .fused = .{ .row_normalization = .{
2076 .dtype = .f32,
2077 .kind = .{ .residual = .{ .rmsnorm = .none } },
2078 .rows = 2,
2079 .cols = 4,
2080 } } }) == null);
2081 try std.testing.expect(select(.{ .fused = .{ .row_normalization = .{
2082 .dtype = .f32,
2083 .kind = .{ .residual = .softmax },
2084 .rows = 2,
2085 .cols = 4,
2086 } } }) == null);
2087 try std.testing.expect(select(.{ .fused = .{ .row_normalization = .{
2088 .dtype = .f32,
2089 .kind = .{ .residual = .{ .rmsnorm = .scale } },
2090 .rows = 3,
2091 .cols = 4,
2092 } } }) == null);
2093 try std.testing.expect(select(.{ .fused = .{ .row_normalization = .{
2094 .dtype = .f32,
2095 .kind = .{ .residual = .{ .rmsnorm = .scale } },
2096 .rows = 2,
2097 .cols = 5,
2098 } } }) == null);
2099 try std.testing.expect(select(.{ .fused = .{ .row_normalization = .{
2100 .dtype = .f32,
2101 .kind = .{ .residual = .{ .rmsnorm = .scale } },
2102 .rows = 2,
2103 .cols = 4,
2104 .schedule = .{ .thread_blocks = .{ .x = 8, .y = 2 } },
2105 } } }) == null);
2106 }
2107
2108 test "kernel library catalog rejects unavailable activation variants" {
2109 try std.testing.expect(select(.{ .activation = .{
2110 .dtype = .i32,
2111 .kind = .gelu,
2112 .extent = 8,
2113 } }) == null);
2114 try std.testing.expect(select(.{ .activation = .{
2115 .dtype = .f32,
2116 .kind = .gelu,
2117 .extent = 0,
2118 } }) == null);
2119 try std.testing.expect(select(.{ .activation = .{
2120 .dtype = .f32,
2121 .kind = .gelu,
2122 .extent = 16,
2123 } }) == null);
2124 try std.testing.expect(select(.{ .activation = .{
2125 .dtype = .i32,
2126 .kind = .relu,
2127 .extent = 8,
2128 } }) == null);
2129 try std.testing.expect(select(.{ .activation = .{
2130 .dtype = .f32,
2131 .kind = .relu,
2132 .extent = 0,
2133 } }) == null);
2134 try std.testing.expect(select(.{ .activation = .{
2135 .dtype = .f32,
2136 .kind = .relu,
2137 .extent = 16,
2138 } }) == null);
2139 try std.testing.expect(select(.{ .activation = .{
2140 .dtype = .i32,
2141 .kind = .silu,
2142 .extent = 8,
2143 } }) == null);
2144 try std.testing.expect(select(.{ .activation = .{
2145 .dtype = .f32,
2146 .kind = .silu,
2147 .extent = 0,
2148 } }) == null);
2149 try std.testing.expect(select(.{ .activation = .{
2150 .dtype = .f32,
2151 .kind = .silu,
2152 .extent = 16,
2153 } }) == null);
2154 }
2155
2156 test "kernel library catalog rejects unavailable transpose variants" {
2157 const input_dims = [_]i64{ 8, 16 };
2158 const output_dims = [_]i64{ 16, 8 };
2159 const wrong_output_dims = [_]i64{ 8, 16 };
2160 const unavailable_dims = [_]i64{ 16, 8 };
2161 const unavailable_output_dims = [_]i64{ 8, 16 };
2162
2163 try std.testing.expect(select(.{ .layout = .{
2164 .dtype = .i32,
2165 .kind = .transpose,
2166 .input_indices = "ij",
2167 .output_indices = "ji",
2168 .input_dims = &input_dims,
2169 .output_dims = &output_dims,
2170 } }) == null);
2171 try std.testing.expect(select(.{ .layout = .{
2172 .dtype = .f32,
2173 .kind = .transpose,
2174 .input_indices = "ij",
2175 .output_indices = "ij",
2176 .input_dims = &input_dims,
2177 .output_dims = &output_dims,
2178 } }) == null);
2179 try std.testing.expect(select(.{ .layout = .{
2180 .dtype = .f32,
2181 .kind = .transpose,
2182 .input_indices = "ij",
2183 .output_indices = "ji",
2184 .input_dims = &input_dims,
2185 .output_dims = &wrong_output_dims,
2186 } }) == null);
2187 try std.testing.expect(select(.{ .layout = .{
2188 .dtype = .f32,
2189 .kind = .transpose,
2190 .input_indices = "ij",
2191 .output_indices = "ji",
2192 .input_dims = &unavailable_dims,
2193 .output_dims = &unavailable_output_dims,
2194 } }) == null);
2195 }
2196 };
2197
2198 const normalization_tests = struct {
2199 test "kernel library catalog selects row normalization by aggregate contract" {
2200 const softmax = select(.{ .row_normalization = .{
2201 .dtype = .f32,
2202 .kind = .softmax,
2203 .rows = 2,
2204 .cols = 4,
2205 } }) orelse return error.TestExpectedRowSoftmaxSelection;
2206 const log_softmax = select(.{ .row_normalization = .{
2207 .dtype = .f32,
2208 .kind = .log_softmax,
2209 .rows = 2,
2210 .cols = 4,
2211 } }) orelse return error.TestExpectedRowLogSoftmaxSelection;
2212 const rmsnorm = select(.{ .row_normalization = .{
2213 .dtype = .f32,
2214 .kind = .{ .rmsnorm = .scale },
2215 .rows = 2,
2216 .cols = 4,
2217 } }) orelse return error.TestExpectedRowRmsNormSelection;
2218 const layernorm = select(.{ .row_normalization = .{
2219 .dtype = .f32,
2220 .kind = .{ .layernorm = .none },
2221 .rows = 2,
2222 .cols = 4,
2223 } }) orelse return error.TestExpectedRowLayerNormSelection;
2224 const affine_layernorm = select(.{ .row_normalization = .{
2225 .dtype = .f32,
2226 .kind = .{ .layernorm = .scale_bias },
2227 .rows = 2,
2228 .cols = 4,
2229 } }) orelse return error.TestExpectedRowAffineLayerNormSelection;
2230
2231 try std.testing.expectEqualStrings(normalization.RowSoftmax2x4F32.target, softmax.metadata.target);
2232 try std.testing.expectEqualStrings(normalization.RowLogSoftmax2x4F32.target, log_softmax.metadata.target);
2233 try std.testing.expectEqualStrings(normalization.RowRmsNorm2x4F32.target, rmsnorm.metadata.target);
2234 try std.testing.expectEqualStrings(normalization.RowLayerNorm2x4F32.target, layernorm.metadata.target);
2235 try std.testing.expectEqualStrings(normalization.RowAffineLayerNorm2x4F32.target, affine_layernorm.metadata.target);
2236 }
2237
2238 test "kernel library catalog selects row normalization schedule specialization" {
2239 const default_softmax = select(.{ .row_normalization = .{
2240 .dtype = .f32,
2241 .kind = .softmax,
2242 .rows = 2,
2243 .cols = 4,
2244 } }) orelse return error.TestExpectedRowSoftmaxSelection;
2245 const scheduled_softmax = select(.{ .row_normalization = .{
2246 .dtype = .f32,
2247 .kind = .softmax,
2248 .rows = 2,
2249 .cols = 4,
2250 .schedule = .{ .thread_blocks = .{ .x = 2, .y = 2 } },
2251 } }) orelse return error.TestExpectedRowSoftmaxScheduleSelection;
2252
2253 try std.testing.expectEqualStrings(normalization.RowSoftmax2x4F32.target, default_softmax.metadata.target);
2254 try std.testing.expectEqualStrings(normalization.RowSoftmax2x4ThreadBlocks2x2F32.target, scheduled_softmax.metadata.target);
2255 try std.testing.expectEqualStrings(normalization.RowSoftmax2x4ThreadBlocks2x2F32.name, scheduled_softmax.name);
2256 try std.testing.expectEqual(@as(u32, 2), scheduled_softmax.metadata.specialization.launch.?.threadgroup[0]);
2257 try std.testing.expectEqual(@as(u32, 2), scheduled_softmax.metadata.specialization.launch.?.threadgroup[1]);
2258 }
2259
2260 test "kernel library catalog rejects unavailable row normalization variants" {
2261 try std.testing.expect(select(.{ .row_normalization = .{
2262 .dtype = .f32,
2263 .kind = .softmax,
2264 .rows = 3,
2265 .cols = 4,
2266 } }) == null);
2267 try std.testing.expect(select(.{ .row_normalization = .{
2268 .dtype = .f32,
2269 .kind = .log_softmax,
2270 .rows = 2,
2271 .cols = 5,
2272 } }) == null);
2273 try std.testing.expect(select(.{ .row_normalization = .{
2274 .dtype = .i32,
2275 .kind = .{ .rmsnorm = .scale },
2276 .rows = 2,
2277 .cols = 4,
2278 } }) == null);
2279 try std.testing.expect(select(.{ .row_normalization = .{
2280 .dtype = .f32,
2281 .kind = .{ .rmsnorm = .none },
2282 .rows = 2,
2283 .cols = 4,
2284 } }) == null);
2285 try std.testing.expect(select(.{ .row_normalization = .{
2286 .dtype = .f32,
2287 .kind = .{ .layernorm = .none },
2288 .rows = 2,
2289 .cols = 5,
2290 } }) == null);
2291 try std.testing.expect(select(.{ .row_normalization = .{
2292 .dtype = .f32,
2293 .kind = .{ .layernorm = .scale },
2294 .rows = 2,
2295 .cols = 4,
2296 } }) == null);
2297 try std.testing.expect(select(.{ .row_normalization = .{
2298 .dtype = .f32,
2299 .kind = .softmax,
2300 .rows = 2,
2301 .cols = 4,
2302 .schedule = .{ .thread_blocks = .{ .x = 8, .y = 2 } },
2303 } }) == null);
2304 }
2305
2306 test "kernel library catalog rejects transposed matrix product layouts" {
2307 const lhs_dims = [_]i64{ 8, 4 };
2308 const rhs_dims = [_]i64{ 8, 16 };
2309 const output_dims = [_]i64{ 4, 16 };
2310 try std.testing.expect(select(.{ .matrix_product = .{
2311 .dtype = .f32,
2312 .lhs_indices = "ki",
2313 .rhs_indices = "kj",
2314 .output_indices = "ij",
2315 .lhs_dims = &lhs_dims,
2316 .rhs_dims = &rhs_dims,
2317 .output_dims = &output_dims,
2318 } }) == null);
2319 }
2320
2321 test "kernel library catalog rejects unavailable attention variants" {
2322 const query_dims = [_]i64{ 2, 2, 2 };
2323 const key_dims = [_]i64{ 2, 3, 2 };
2324 const value_dims = [_]i64{ 2, 3, 2 };
2325 const output_dims = [_]i64{ 2, 2, 2 };
2326 const wrong_key_dims = [_]i64{ 2, 4, 2 };
2327 const unavailable_query_dims = [_]i64{ 3, 2, 2 };
2328 const unavailable_key_dims = [_]i64{ 3, 3, 2 };
2329 const unavailable_value_dims = [_]i64{ 3, 3, 2 };
2330 const unavailable_output_dims = [_]i64{ 3, 2, 2 };
2331
2332 try std.testing.expect(select(.{ .attention = .{
2333 .dtype = .i32,
2334 .kind = .scaled_dot_product,
2335 .query_indices = "bqh",
2336 .key_indices = "bkh",
2337 .value_indices = "bkv",
2338 .output_indices = "bqv",
2339 .query_dims = &query_dims,
2340 .key_dims = &key_dims,
2341 .value_dims = &value_dims,
2342 .output_dims = &output_dims,
2343 } }) == null);
2344 try std.testing.expect(select(.{ .attention = .{
2345 .dtype = .f32,
2346 .kind = .scaled_dot_product,
2347 .query_indices = "bqh",
2348 .key_indices = "ckh",
2349 .value_indices = "bkv",
2350 .output_indices = "bqv",
2351 .query_dims = &query_dims,
2352 .key_dims = &key_dims,
2353 .value_dims = &value_dims,
2354 .output_dims = &output_dims,
2355 } }) == null);
2356 try std.testing.expect(select(.{ .attention = .{
2357 .dtype = .f32,
2358 .kind = .scaled_dot_product,
2359 .query_indices = "bqh",
2360 .key_indices = "bkh",
2361 .value_indices = "bsv",
2362 .output_indices = "bqv",
2363 .query_dims = &query_dims,
2364 .key_dims = &key_dims,
2365 .value_dims = &value_dims,
2366 .output_dims = &output_dims,
2367 } }) == null);
2368 try std.testing.expect(select(.{ .attention = .{
2369 .dtype = .f32,
2370 .kind = .scaled_dot_product,
2371 .query_indices = "bqh",
2372 .key_indices = "bkh",
2373 .value_indices = "bkv",
2374 .output_indices = "bqv",
2375 .query_dims = &query_dims,
2376 .key_dims = &wrong_key_dims,
2377 .value_dims = &value_dims,
2378 .output_dims = &output_dims,
2379 } }) == null);
2380 try std.testing.expect(select(.{ .attention = .{
2381 .dtype = .f32,
2382 .kind = .scaled_dot_product,
2383 .query_indices = "bqh",
2384 .key_indices = "bkh",
2385 .value_indices = "bkv",
2386 .output_indices = "bqv",
2387 .query_dims = &query_dims,
2388 .key_dims = &key_dims,
2389 .value_dims = &value_dims,
2390 .output_dims = &output_dims,
2391 .schedule = .{ .thread_blocks = .{ .x = 4, .y = 2, .z = 1 } },
2392 } }) == null);
2393 try std.testing.expect(select(.{ .attention = .{
2394 .dtype = .f32,
2395 .kind = .scaled_dot_product,
2396 .query_indices = "bqh",
2397 .key_indices = "bkh",
2398 .value_indices = "bkv",
2399 .output_indices = "bqv",
2400 .query_dims = &unavailable_query_dims,
2401 .key_dims = &unavailable_key_dims,
2402 .value_dims = &unavailable_value_dims,
2403 .output_dims = &unavailable_output_dims,
2404 } }) == null);
2405 }
2406 };
2407
2408 const invalid_tests = struct {
2409 test "kernel library catalog rejects unavailable linalg schedules" {
2410 const lhs_dims = [_]i64{ 4, 8 };
2411 const rhs_dims = [_]i64{ 8, 16 };
2412 const matrix_output_dims = [_]i64{ 4, 16 };
2413 const batched_lhs_dims = [_]i64{ 2, 2, 4 };
2414 const batched_rhs_dims = [_]i64{ 2, 4, 3 };
2415 const batched_output_dims = [_]i64{ 2, 2, 3 };
2416 const vector_dims = [_]i64{8};
2417 const vector_output_dims = [_]i64{4};
2418 const outer_lhs_dims = [_]i64{4};
2419 const outer_rhs_dims = [_]i64{3};
2420 const outer_output_dims = [_]i64{ 4, 3 };
2421
2422 try std.testing.expect(select(.{ .matrix_product = .{
2423 .dtype = .f32,
2424 .lhs_indices = "ik",
2425 .rhs_indices = "kj",
2426 .output_indices = "ij",
2427 .lhs_dims = &lhs_dims,
2428 .rhs_dims = &rhs_dims,
2429 .output_dims = &matrix_output_dims,
2430 .schedule = .{ .thread_blocks = .{ .x = 16, .y = 8 } },
2431 } }) == null);
2432
2433 try std.testing.expect(select(.{ .batched_matrix_product = .{
2434 .dtype = .f32,
2435 .lhs_indices = "bmk",
2436 .rhs_indices = "bkn",
2437 .output_indices = "bmn",
2438 .lhs_dims = &batched_lhs_dims,
2439 .rhs_dims = &batched_rhs_dims,
2440 .output_dims = &batched_output_dims,
2441 .schedule = .{ .thread_blocks = .{ .x = 4, .y = 2, .z = 2 } },
2442 } }) == null);
2443
2444 try std.testing.expect(select(.{ .matrix_vector_product = .{
2445 .dtype = .f32,
2446 .matrix_indices = "mk",
2447 .vector_indices = "k",
2448 .output_indices = "m",
2449 .matrix_dims = &lhs_dims,
2450 .vector_dims = &vector_dims,
2451 .output_dims = &vector_output_dims,
2452 .schedule = .{ .thread_blocks = 8 },
2453 } }) == null);
2454
2455 try std.testing.expect(select(.{ .outer_product = .{
2456 .dtype = .f32,
2457 .lhs_indices = "m",
2458 .rhs_indices = "n",
2459 .output_indices = "mn",
2460 .lhs_dims = &outer_lhs_dims,
2461 .rhs_dims = &outer_rhs_dims,
2462 .output_dims = &outer_output_dims,
2463 .schedule = .{ .thread_blocks = .{ .x = 4, .y = 2 } },
2464 } }) == null);
2465 }
2466
2467 test "kernel library catalog rejects unavailable batched matrix product variants" {
2468 const lhs_dims = [_]i64{ 2, 2, 4 };
2469 const rhs_dims = [_]i64{ 2, 4, 3 };
2470 const output_dims = [_]i64{ 2, 2, 3 };
2471 const wrong_rhs_dims = [_]i64{ 3, 4, 3 };
2472 const unavailable_lhs_dims = [_]i64{ 3, 2, 4 };
2473 const unavailable_output_dims = [_]i64{ 3, 2, 3 };
2474
2475 try std.testing.expect(select(.{ .batched_matrix_product = .{
2476 .dtype = .i32,
2477 .lhs_indices = "bmk",
2478 .rhs_indices = "bkn",
2479 .output_indices = "bmn",
2480 .lhs_dims = &lhs_dims,
2481 .rhs_dims = &rhs_dims,
2482 .output_dims = &output_dims,
2483 } }) == null);
2484 try std.testing.expect(select(.{ .batched_matrix_product = .{
2485 .dtype = .f32,
2486 .lhs_indices = "mbk",
2487 .rhs_indices = "bkn",
2488 .output_indices = "bmn",
2489 .lhs_dims = &lhs_dims,
2490 .rhs_dims = &rhs_dims,
2491 .output_dims = &output_dims,
2492 } }) == null);
2493 try std.testing.expect(select(.{ .batched_matrix_product = .{
2494 .dtype = .f32,
2495 .lhs_indices = "bmk",
2496 .rhs_indices = "ckn",
2497 .output_indices = "bmn",
2498 .lhs_dims = &lhs_dims,
2499 .rhs_dims = &rhs_dims,
2500 .output_dims = &output_dims,
2501 } }) == null);
2502 try std.testing.expect(select(.{ .batched_matrix_product = .{
2503 .dtype = .f32,
2504 .lhs_indices = "bmk",
2505 .rhs_indices = "bkn",
2506 .output_indices = "bmn",
2507 .lhs_dims = &lhs_dims,
2508 .rhs_dims = &wrong_rhs_dims,
2509 .output_dims = &output_dims,
2510 } }) == null);
2511 try std.testing.expect(select(.{ .batched_matrix_product = .{
2512 .dtype = .f32,
2513 .lhs_indices = "bmk",
2514 .rhs_indices = "bkn",
2515 .output_indices = "bmn",
2516 .lhs_dims = &unavailable_lhs_dims,
2517 .rhs_dims = &wrong_rhs_dims,
2518 .output_dims = &unavailable_output_dims,
2519 } }) == null);
2520 }
2521
2522 test "kernel library catalog rejects unavailable matrix vector product variants" {
2523 const matrix_dims = [_]i64{ 4, 8 };
2524 const vector_dims = [_]i64{8};
2525 const output_dims = [_]i64{4};
2526 const wrong_vector_dims = [_]i64{7};
2527 const unavailable_matrix_dims = [_]i64{ 5, 8 };
2528 const unavailable_output_dims = [_]i64{5};
2529
2530 try std.testing.expect(select(.{ .matrix_vector_product = .{
2531 .dtype = .i32,
2532 .matrix_indices = "mk",
2533 .vector_indices = "k",
2534 .output_indices = "m",
2535 .matrix_dims = &matrix_dims,
2536 .vector_dims = &vector_dims,
2537 .output_dims = &output_dims,
2538 } }) == null);
2539 try std.testing.expect(select(.{ .matrix_vector_product = .{
2540 .dtype = .f32,
2541 .matrix_indices = "km",
2542 .vector_indices = "k",
2543 .output_indices = "m",
2544 .matrix_dims = &matrix_dims,
2545 .vector_dims = &vector_dims,
2546 .output_dims = &output_dims,
2547 } }) == null);
2548 try std.testing.expect(select(.{ .matrix_vector_product = .{
2549 .dtype = .f32,
2550 .matrix_indices = "mk",
2551 .vector_indices = "n",
2552 .output_indices = "m",
2553 .matrix_dims = &matrix_dims,
2554 .vector_dims = &vector_dims,
2555 .output_dims = &output_dims,
2556 } }) == null);
2557 try std.testing.expect(select(.{ .matrix_vector_product = .{
2558 .dtype = .f32,
2559 .matrix_indices = "mk",
2560 .vector_indices = "k",
2561 .output_indices = "m",
2562 .matrix_dims = &matrix_dims,
2563 .vector_dims = &wrong_vector_dims,
2564 .output_dims = &output_dims,
2565 } }) == null);
2566 try std.testing.expect(select(.{ .matrix_vector_product = .{
2567 .dtype = .f32,
2568 .matrix_indices = "mk",
2569 .vector_indices = "k",
2570 .output_indices = "m",
2571 .matrix_dims = &unavailable_matrix_dims,
2572 .vector_dims = &vector_dims,
2573 .output_dims = &unavailable_output_dims,
2574 } }) == null);
2575 }
2576
2577 test "kernel library catalog rejects unavailable outer product variants" {
2578 const lhs_dims = [_]i64{4};
2579 const rhs_dims = [_]i64{3};
2580 const output_dims = [_]i64{ 4, 3 };
2581 const wrong_rhs_dims = [_]i64{4};
2582 const unavailable_lhs_dims = [_]i64{5};
2583 const unavailable_output_dims = [_]i64{ 5, 3 };
2584
2585 try std.testing.expect(select(.{ .outer_product = .{
2586 .dtype = .i32,
2587 .lhs_indices = "m",
2588 .rhs_indices = "n",
2589 .output_indices = "mn",
2590 .lhs_dims = &lhs_dims,
2591 .rhs_dims = &rhs_dims,
2592 .output_dims = &output_dims,
2593 } }) == null);
2594 try std.testing.expect(select(.{ .outer_product = .{
2595 .dtype = .f32,
2596 .lhs_indices = "m",
2597 .rhs_indices = "m",
2598 .output_indices = "mm",
2599 .lhs_dims = &lhs_dims,
2600 .rhs_dims = &rhs_dims,
2601 .output_dims = &output_dims,
2602 } }) == null);
2603 try std.testing.expect(select(.{ .outer_product = .{
2604 .dtype = .f32,
2605 .lhs_indices = "m",
2606 .rhs_indices = "n",
2607 .output_indices = "nm",
2608 .lhs_dims = &lhs_dims,
2609 .rhs_dims = &rhs_dims,
2610 .output_dims = &output_dims,
2611 } }) == null);
2612 try std.testing.expect(select(.{ .outer_product = .{
2613 .dtype = .f32,
2614 .lhs_indices = "m",
2615 .rhs_indices = "n",
2616 .output_indices = "mn",
2617 .lhs_dims = &lhs_dims,
2618 .rhs_dims = &wrong_rhs_dims,
2619 .output_dims = &output_dims,
2620 } }) == null);
2621 try std.testing.expect(select(.{ .outer_product = .{
2622 .dtype = .f32,
2623 .lhs_indices = "m",
2624 .rhs_indices = "n",
2625 .output_indices = "mn",
2626 .lhs_dims = &unavailable_lhs_dims,
2627 .rhs_dims = &rhs_dims,
2628 .output_dims = &unavailable_output_dims,
2629 } }) == null);
2630 }
2631 };
2632
2633 const family_tests = struct {
2634 test "kernel library catalog creates registry-ready artifact by target" {
2635 const allocator = std.testing.allocator;
2636 var state = gpu.recording.BackendState{
2637 .allocator = allocator,
2638 .kind = .cuda,
2639 .format = .cuda_ptx,
2640 };
2641
2642 var call_artifact = try createKernelCallArtifact(allocator, state.handle(), .{
2643 .target = linalg.MatrixProduct4x16x8F32.target,
2644 .options = .{ .limits = linalg.MatrixProduct4x16x8F32.Limits.testing },
2645 });
2646 defer call_artifact.deinit();
2647
2648 const artifact = call_artifact.registry().find(linalg.MatrixProduct4x16x8F32.target, linalg.MatrixProduct4x16x8F32.version, .cuda_ptx) orelse {
2649 return error.TestExpectedKernelCallArtifact;
2650 };
2651 try std.testing.expectEqualStrings(linalg.MatrixProduct4x16x8F32.name, artifact.entry_name);
2652 }
2653
2654 test "kernel library catalog selects stencil window descriptors" {
2655 const static_descriptor = select(.{ .stencil = .{
2656 .dtype = .f32,
2657 .kind = .window,
2658 .rows = 2,
2659 .cols = 3,
2660 .radius = 1,
2661 } }) orelse return error.TestExpectedCatalogDescriptor;
2662 try std.testing.expectEqualStrings(stencil.Window2x3R1F32.target, static_descriptor.metadata.target);
2663
2664 var owned = (try selectOwned(std.testing.allocator, .{ .stencil = .{
2665 .dtype = .f32,
2666 .kind = .window,
2667 .rows = 17,
2668 .cols = 17,
2669 .radius = 1,
2670 } })) orelse return error.TestExpectedCatalogDescriptor;
2671 defer owned.deinit();
2672 try std.testing.expectEqualStrings("accy.kernel.stencil.window_family_r1_17x9_f32", owned.descriptor.metadata.target);
2673 try std.testing.expect(owned.descriptor.metadata.specialization.shape_family != null);
2674
2675 var scheduled = (try selectOwned(std.testing.allocator, .{ .stencil = .{
2676 .dtype = .f32,
2677 .kind = .window,
2678 .rows = 17,
2679 .cols = 17,
2680 .radius = 1,
2681 .schedule = .{ .thread_blocks = .{ .x = 4, .y = 2 } },
2682 } })) orelse return error.TestExpectedCatalogDescriptor;
2683 defer scheduled.deinit();
2684 try std.testing.expectEqualStrings("accy.kernel.stencil.window_family_r1_4x2_f32", scheduled.descriptor.metadata.target);
2685
2686 try std.testing.expectEqual(
2687 @as(?Descriptor, null),
2688 select(.{ .stencil = .{ .dtype = .f32, .kind = .window, .rows = 0, .cols = 3, .radius = 1 } }),
2689 );
2690 try std.testing.expectEqual(
2691 @as(?OwnedDescriptor, null),
2692 try selectOwned(std.testing.allocator, .{ .stencil = .{
2693 .dtype = .f32,
2694 .kind = .window,
2695 .rows = 4,
2696 .cols = 4,
2697 .radius = stencil.window_radius_max + 1,
2698 } }),
2699 );
2700 }
2701
2702 test "kernel library catalog selects owned stencil window family candidates" {
2703 var candidates = try selectOwnedStencilWindowCandidates(std.testing.allocator, .{
2704 .dtype = .f32,
2705 .kind = .window,
2706 .rows = 17,
2707 .cols = 17,
2708 .radius = 1,
2709 });
2710 defer candidates.deinit();
2711
2712 try std.testing.expect(candidates.count > 2);
2713 try std.testing.expectEqualStrings("accy.kernel.stencil.window_family_r1_17x9_f32", candidates.items[0].descriptor.metadata.target);
2714 for (candidates.slice(), 0..) |candidate, index| {
2715 try std.testing.expect(candidate.specialization != null);
2716 try std.testing.expect(candidate.descriptor.metadata.specialization.shape_family != null);
2717 try std.testing.expect(stencilWindowDescriptorMatches(candidate.descriptor, .{
2718 .dtype = .f32,
2719 .kind = .window,
2720 .rows = 17,
2721 .cols = 17,
2722 .radius = 1,
2723 }));
2724 for (candidates.slice()[0..index]) |previous| {
2725 try std.testing.expect(!std.mem.eql(u8, previous.descriptor.metadata.target, candidate.descriptor.metadata.target));
2726 }
2727 }
2728 }
2729
2730 test "kernel library catalog selects image descriptors" {
2731 try std.testing.expectEqual(
2732 @as(?Descriptor, null),
2733 select(.{ .image = .{
2734 .dtype = .u32,
2735 .kind = .{ .blur_pass = .{ .radius = 2, .axis = .horizontal } },
2736 .width = 24,
2737 .height = 10,
2738 } }),
2739 );
2740
2741 var blur = (try selectOwned(std.testing.allocator, .{ .image = .{
2742 .dtype = .u32,
2743 .kind = .{ .blur_pass = .{ .radius = 2, .axis = .horizontal } },
2744 .width = 24,
2745 .height = 10,
2746 } })) orelse return error.TestExpectedImageDescriptor;
2747 defer blur.deinit();
2748 try std.testing.expectEqualStrings("image_blur_pass_family_r2h_24x10_rgba8", blur.descriptor.metadata.target);
2749 try std.testing.expectEqual(entry.Category.image, blur.descriptor.metadata.category);
2750 try std.testing.expect(blur.descriptor.metadata.specialization.operationIs(.{ .image = .blur_pass }));
2751 try std.testing.expect(blur.descriptor.metadata.specialization.shape_family != null);
2752 try std.testing.expect(blur.descriptor.metadata.specialization.staticParameterMatches("radius", 2));
2753 try std.testing.expect(blur.descriptor.metadata.specialization.staticParameterMatches("axis", image.blurPassAxisParameter(.horizontal)));
2754
2755 var scheduled_blur = (try selectOwned(std.testing.allocator, .{ .image = .{
2756 .dtype = .u32,
2757 .kind = .{ .blur_pass = .{ .radius = 2, .axis = .horizontal } },
2758 .width = 24,
2759 .height = 10,
2760 .schedule = .{ .thread_blocks = .{ .x = 8, .y = 4 } },
2761 } })) orelse return error.TestExpectedImageDescriptor;
2762 defer scheduled_blur.deinit();
2763 try std.testing.expectEqualStrings("image_blur_pass_family_r2h_8x4_rgba8", scheduled_blur.descriptor.metadata.target);
2764
2765 var resize = (try selectOwned(std.testing.allocator, .{ .image = .{
2766 .dtype = .u32,
2767 .kind = .{ .resize_bilinear = .{ .src_width = 20, .src_height = 12 } },
2768 .width = 11,
2769 .height = 7,
2770 } })) orelse return error.TestExpectedImageDescriptor;
2771 defer resize.deinit();
2772 try std.testing.expectEqualStrings("image_resize_bilinear_family_11x7_rgba8", resize.descriptor.metadata.target);
2773 try std.testing.expectEqual(entry.Category.image, resize.descriptor.metadata.category);
2774 try std.testing.expect(resize.descriptor.metadata.specialization.operationIs(.{ .image = .resize_bilinear }));
2775 try std.testing.expect(resize.descriptor.metadata.specialization.shape_family != null);
2776
2777 try std.testing.expectEqual(
2778 @as(?OwnedDescriptor, null),
2779 try selectOwned(std.testing.allocator, .{ .image = .{
2780 .dtype = .f32,
2781 .kind = .{ .blur_pass = .{ .radius = 2, .axis = .horizontal } },
2782 .width = 24,
2783 .height = 10,
2784 } }),
2785 );
2786 try std.testing.expectEqual(
2787 @as(?OwnedDescriptor, null),
2788 try selectOwned(std.testing.allocator, .{ .image = .{
2789 .dtype = .u32,
2790 .kind = .{ .blur_pass = .{ .radius = 0, .axis = .horizontal } },
2791 .width = 24,
2792 .height = 10,
2793 } }),
2794 );
2795 try std.testing.expectEqual(
2796 @as(?OwnedDescriptor, null),
2797 try selectOwned(std.testing.allocator, .{ .image = .{
2798 .dtype = .u32,
2799 .kind = .{ .blur_pass = .{ .radius = 2, .axis = .horizontal } },
2800 .width = 0,
2801 .height = 10,
2802 } }),
2803 );
2804 try std.testing.expectEqual(
2805 @as(?OwnedDescriptor, null),
2806 try selectOwned(std.testing.allocator, .{ .image = .{
2807 .dtype = .u32,
2808 .kind = .{ .resize_bilinear = .{ .src_width = 20, .src_height = 0 } },
2809 .width = 11,
2810 .height = 7,
2811 } }),
2812 );
2813 }
2814
2815 test "kernel library catalog selects owned image family candidates" {
2816 var candidates = try selectOwnedImageCandidates(std.testing.allocator, .{
2817 .dtype = .u32,
2818 .kind = .{ .blur_pass = .{ .radius = 2, .axis = .horizontal } },
2819 .width = 24,
2820 .height = 10,
2821 });
2822 defer candidates.deinit();
2823
2824 try std.testing.expect(candidates.count > 1);
2825 try std.testing.expectEqualStrings("image_blur_pass_family_r2h_24x10_rgba8", candidates.items[0].descriptor.metadata.target);
2826 for (candidates.slice(), 0..) |candidate, index| {
2827 try std.testing.expect(candidate.specialization != null);
2828 try std.testing.expectEqual(entry.Category.image, candidate.descriptor.metadata.category);
2829 try std.testing.expect(candidate.descriptor.metadata.specialization.shape_family != null);
2830 try std.testing.expect(candidate.descriptor.metadata.specialization.operationIs(.{ .image = .blur_pass }));
2831 for (candidates.slice()[0..index]) |previous| {
2832 try std.testing.expect(!std.mem.eql(u8, previous.descriptor.metadata.target, candidate.descriptor.metadata.target));
2833 }
2834 }
2835 }
2836
2837 test "kernel library catalog builds stencil window scheduled artifact registry" {
2838 const allocator = std.testing.allocator;
2839 var state = gpu.recording.BackendState{
2840 .allocator = allocator,
2841 .kind = .cuda,
2842 .format = .cuda_ptx,
2843 };
2844 var candidates = try selectOwnedStencilWindowCandidates(allocator, .{
2845 .dtype = .f32,
2846 .kind = .window,
2847 .rows = 17,
2848 .cols = 17,
2849 .radius = 1,
2850 .schedule = .{ .thread_blocks = .{ .x = 17, .y = 9 } },
2851 });
2852 defer candidates.deinit();
2853
2854 var owned_registry = try createOwnedKernelCallArtifactRegistry(allocator, state.handle(), candidates.slice(), .{ .limits = .testing });
2855 defer owned_registry.deinit();
2856 const registry = owned_registry.registry();
2857 const selected = registry.find("accy.kernel.stencil.window_family_r1_17x9_f32", stencil.window_family_version, .cuda_ptx) orelse {
2858 return error.TestExpectedKernelCallArtifact;
2859 };
2860 const launch = switch (selected.launch) {
2861 .derived => |derived| derived,
2862 .fixed => return error.TestExpectedDerivedLaunch,
2863 };
2864 const profile = selected.shape_profile orelse return error.TestExpectedShapeProfile;
2865
2866 try std.testing.expectEqual(candidates.count, registry.entries.len);
2867 try std.testing.expectEqual(@as(u32, 5), selected.argument_count);
2868 try std.testing.expectEqual(@as(u32, 2), selected.runtime_scalar_argument_count);
2869 try std.testing.expectEqual(@as(u32, 17), launch.threadgroup[0]);
2870 try std.testing.expectEqual(@as(u32, 9), launch.threadgroup[1]);
2871 try std.testing.expectEqual(@as(usize, 2), profile.dimensions.len);
2872 try std.testing.expect(selected.required_dtypes.contains(.f32));
2873 try std.testing.expect(selected.required_dtypes.contains(.i32));
2874 }
2875
2876 test "kernel library catalog selects gather descriptors" {
2877 const static_descriptor = select(.{ .gather = .{
2878 .dtype = .f32,
2879 .axis_size = 8,
2880 .gathered = 8,
2881 } }) orelse return error.TestExpectedCatalogDescriptor;
2882 try std.testing.expectEqualStrings(indexing.Gather8F32.target, static_descriptor.metadata.target);
2883
2884 var owned = (try selectOwned(std.testing.allocator, .{ .gather = .{
2885 .dtype = .f32,
2886 .outer = 2,
2887 .axis_size = 1024,
2888 .gathered = 500,
2889 .inner = 1,
2890 } })) orelse return error.TestExpectedCatalogDescriptor;
2891 defer owned.deinit();
2892 try std.testing.expectEqualStrings("accy.kernel.indexing.gather_family_256_f32", owned.descriptor.metadata.target);
2893 try std.testing.expect(owned.descriptor.metadata.specialization.shape_family != null);
2894
2895 var scheduled = (try selectOwned(std.testing.allocator, .{ .gather = .{
2896 .dtype = .f32,
2897 .outer = 2,
2898 .axis_size = 1024,
2899 .gathered = 500,
2900 .inner = 1,
2901 .schedule = .{ .thread_blocks = 64 },
2902 } })) orelse return error.TestExpectedCatalogDescriptor;
2903 defer scheduled.deinit();
2904 try std.testing.expectEqualStrings("accy.kernel.indexing.gather_family_64_f32", scheduled.descriptor.metadata.target);
2905
2906 try std.testing.expectEqual(
2907 @as(?OwnedDescriptor, null),
2908 try selectOwned(std.testing.allocator, .{ .gather = .{ .dtype = .f32, .axis_size = 0, .gathered = 4 } }),
2909 );
2910 }
2911
2912 test "kernel library catalog builds gather scheduled artifact registry" {
2913 const allocator = std.testing.allocator;
2914 var state = gpu.recording.BackendState{
2915 .allocator = allocator,
2916 .kind = .cuda,
2917 .format = .cuda_ptx,
2918 };
2919 var candidates = try selectOwnedGatherCandidates(allocator, .{
2920 .dtype = .f32,
2921 .outer = 2,
2922 .axis_size = 1024,
2923 .gathered = 500,
2924 .inner = 1,
2925 .schedule = .{ .thread_blocks = 256 },
2926 });
2927 defer candidates.deinit();
2928
2929 try std.testing.expectEqual(@as(usize, 1), candidates.count);
2930
2931 var owned_registry = try createOwnedKernelCallArtifactRegistry(allocator, state.handle(), candidates.slice(), .{ .limits = .testing });
2932 defer owned_registry.deinit();
2933 const registry = owned_registry.registry();
2934 const selected = registry.find("accy.kernel.indexing.gather_family_256_f32", indexing.gather_family_version, .cuda_ptx) orelse {
2935 return error.TestExpectedKernelCallArtifact;
2936 };
2937
2938 try std.testing.expectEqual(candidates.count, registry.entries.len);
2939 try std.testing.expectEqual(@as(u32, 8), selected.argument_count);
2940 try std.testing.expectEqual(@as(u32, 5), selected.runtime_scalar_argument_count);
2941 try std.testing.expect(selected.required_dtypes.contains(.i32));
2942 const launch = switch (selected.launch) {
2943 .derived => |derived| derived,
2944 .fixed => return error.TestExpectedDerivedLaunch,
2945 };
2946 try std.testing.expectEqual(@as(u32, 256), launch.threadgroup[0]);
2947 }
2948 };
2949
2950 const random_tests = struct {
2951 test "kernel library catalog builds random scheduled artifact registry" {
2952 const allocator = std.testing.allocator;
2953 var state = gpu.recording.BackendState{
2954 .allocator = allocator,
2955 .kind = .cuda,
2956 .format = .cuda_ptx,
2957 };
2958 var candidates = try selectOwnedRandomCandidates(allocator, .{
2959 .dtype = .f32,
2960 .algorithm = .philox,
2961 .count = 1 << 20,
2962 .schedule = .{ .thread_blocks = 256 },
2963 });
2964 defer candidates.deinit();
2965
2966 try std.testing.expectEqual(@as(usize, 1), candidates.count);
2967
2968 var owned_registry = try createOwnedKernelCallArtifactRegistry(allocator, state.handle(), candidates.slice(), .{ .limits = .testing });
2969 defer owned_registry.deinit();
2970 const registry = owned_registry.registry();
2971 const selected = registry.find("accy.kernel.random.philox_family_10r_256_f32", random.philox_family_version, .cuda_ptx) orelse {
2972 return error.TestExpectedKernelCallArtifact;
2973 };
2974
2975 try std.testing.expectEqual(candidates.count, registry.entries.len);
2976 try std.testing.expectEqual(@as(u32, 4), selected.argument_count);
2977 try std.testing.expectEqual(@as(u32, 3), selected.runtime_scalar_argument_count);
2978 const launch = switch (selected.launch) {
2979 .derived => |derived| derived,
2980 .fixed => return error.TestExpectedDerivedLaunch,
2981 };
2982 try std.testing.expectEqual(@as(u32, 256), launch.threadgroup[0]);
2983 }
2984
2985 test "kernel library catalog selects random descriptors across rounds" {
2986 const static_descriptor = select(.{ .random = .{
2987 .dtype = .f32,
2988 .algorithm = .philox,
2989 .count = 8,
2990 .schedule = .{ .thread_blocks = 2 },
2991 } }) orelse return error.TestExpectedCatalogDescriptor;
2992 try std.testing.expectEqualStrings(random.Philox8F32.target, static_descriptor.metadata.target);
2993
2994 var default_owned = (try selectOwned(std.testing.allocator, .{ .random = .{
2995 .dtype = .f32,
2996 .algorithm = .threefry,
2997 .count = 4096,
2998 } })) orelse return error.TestExpectedCatalogDescriptor;
2999 defer default_owned.deinit();
3000 try std.testing.expect(std.mem.indexOf(u8, default_owned.descriptor.metadata.target, "threefry_family_20r_") != null);
3001
3002 var reduced_owned = (try selectOwned(std.testing.allocator, .{ .random = .{
3003 .dtype = .i32,
3004 .algorithm = .philox,
3005 .count = 4096,
3006 .rounds = 7,
3007 } })) orelse return error.TestExpectedCatalogDescriptor;
3008 defer reduced_owned.deinit();
3009 try std.testing.expect(std.mem.indexOf(u8, reduced_owned.descriptor.metadata.target, "philox_family_7r_") != null);
3010
3011 try std.testing.expectEqual(
3012 @as(?OwnedDescriptor, null),
3013 try selectOwned(std.testing.allocator, .{ .random = .{ .dtype = .f32, .count = 0 } }),
3014 );
3015 try std.testing.expectEqual(
3016 @as(?OwnedDescriptor, null),
3017 try selectOwned(std.testing.allocator, .{ .random = .{ .dtype = .f16, .count = 64 } }),
3018 );
3019 try std.testing.expectEqual(
3020 @as(?OwnedDescriptor, null),
3021 try selectOwned(std.testing.allocator, .{ .random = .{ .dtype = .f32, .count = 64, .rounds = 99 } }),
3022 );
3023 }
3024
3025 test "kernel library catalog builds filter scheduled artifact registry" {
3026 const allocator = std.testing.allocator;
3027 var state = gpu.recording.BackendState{
3028 .allocator = allocator,
3029 .kind = .cuda,
3030 .format = .cuda_ptx,
3031 };
3032 var candidates = try selectOwnedFilterCandidates(allocator, .{
3033 .dtype = .f32,
3034 .extent = 1 << 20,
3035 .schedule = .{ .thread_blocks = 256 },
3036 });
3037 defer candidates.deinit();
3038
3039 try std.testing.expectEqual(@as(usize, 1), candidates.count);
3040
3041 var owned_registry = try createOwnedKernelCallArtifactRegistry(allocator, state.handle(), candidates.slice(), .{ .limits = .testing });
3042 defer owned_registry.deinit();
3043 const registry = owned_registry.registry();
3044 const selected = registry.find("accy.kernel.compaction.filter_family_nonzero_256_f32", compaction.filter_family_version, .cuda_ptx) orelse {
3045 return error.TestExpectedKernelCallArtifact;
3046 };
3047
3048 try std.testing.expectEqual(candidates.count, registry.entries.len);
3049 try std.testing.expectEqual(@as(u32, 3), selected.argument_count);
3050 try std.testing.expectEqual(@as(u32, 1), selected.runtime_scalar_argument_count);
3051 }
3052
3053 test "kernel library catalog builds greater filter scheduled artifact registry" {
3054 const allocator = std.testing.allocator;
3055 var state = gpu.recording.BackendState{
3056 .allocator = allocator,
3057 .kind = .cuda,
3058 .format = .cuda_ptx,
3059 };
3060 var candidates = try selectOwnedFilterCandidates(allocator, .{
3061 .dtype = .f32,
3062 .predicate = .greater_than,
3063 .extent = 1 << 20,
3064 .schedule = .{ .thread_blocks = 256 },
3065 });
3066 defer candidates.deinit();
3067
3068 try std.testing.expectEqual(@as(usize, 1), candidates.count);
3069 for (candidates.slice()) |candidate| {
3070 const instance = compaction.filterInstanceFromSpecialization(
3071 candidate.specialization.?.value,
3072 ) orelse return error.TestExpectedFilterInstance;
3073 try std.testing.expectEqual(entry.CompactionPredicate.greater_than, instance.predicate);
3074 }
3075
3076 var owned_registry = try createOwnedKernelCallArtifactRegistry(allocator, state.handle(), candidates.slice(), .{ .limits = .testing });
3077 defer owned_registry.deinit();
3078 const registry = owned_registry.registry();
3079 const selected = registry.find("accy.kernel.compaction.filter_family_greater_256_f32", compaction.filter_family_version, .cuda_ptx) orelse {
3080 return error.TestExpectedKernelCallArtifact;
3081 };
3082
3083 try std.testing.expectEqual(@as(u32, 4), selected.argument_count);
3084 try std.testing.expectEqual(@as(u32, 2), selected.runtime_scalar_argument_count);
3085 }
3086 };
3087
3088 const segment_tests = struct {
3089 test "kernel library catalog selects segment sum descriptors" {
3090 const static_descriptor = select(.{ .segmented = .{
3091 .dtype = .f32,
3092 .kind = .segment_sum,
3093 .segments = 4,
3094 .total = 16,
3095 } }) orelse return error.TestExpectedCatalogDescriptor;
3096 try std.testing.expectEqualStrings(segmented.SegmentSum4F32.target, static_descriptor.metadata.target);
3097
3098 var owned = (try selectOwned(std.testing.allocator, .{ .segmented = .{
3099 .dtype = .f32,
3100 .kind = .segment_sum,
3101 .segments = 1000,
3102 .total = 65536,
3103 } })) orelse return error.TestExpectedCatalogDescriptor;
3104 defer owned.deinit();
3105 try std.testing.expectEqualStrings("accy.kernel.segmented.segment_sum_family_thread_256_f32", owned.descriptor.metadata.target);
3106 try std.testing.expect(owned.descriptor.metadata.specialization.shape_family != null);
3107
3108 var scheduled = (try selectOwned(std.testing.allocator, .{ .segmented = .{
3109 .dtype = .f32,
3110 .kind = .segment_sum,
3111 .segments = 1000,
3112 .total = 65536,
3113 .schedule = .{ .thread_blocks = 64 },
3114 } })) orelse return error.TestExpectedCatalogDescriptor;
3115 defer scheduled.deinit();
3116 try std.testing.expectEqualStrings("accy.kernel.segmented.segment_sum_family_thread_64_f32", scheduled.descriptor.metadata.target);
3117
3118 var warp_scheduled = (try selectOwned(std.testing.allocator, .{ .segmented = .{
3119 .dtype = .f32,
3120 .kind = .segment_sum,
3121 .segments = 1000,
3122 .total = 65536,
3123 .schedule = .{ .warp_blocks = 64 },
3124 } })) orelse return error.TestExpectedCatalogDescriptor;
3125 defer warp_scheduled.deinit();
3126 try std.testing.expectEqualStrings("accy.kernel.segmented.segment_sum_family_warp_64_f32", warp_scheduled.descriptor.metadata.target);
3127
3128 try std.testing.expectEqual(
3129 @as(?OwnedDescriptor, null),
3130 try selectOwned(std.testing.allocator, .{ .segmented = .{
3131 .dtype = .f32,
3132 .kind = .segment_sum,
3133 .segments = 0,
3134 .total = 16,
3135 } }),
3136 );
3137
3138 try std.testing.expectEqual(
3139 @as(?OwnedDescriptor, null),
3140 try selectOwned(std.testing.allocator, .{ .segmented = .{
3141 .dtype = .f32,
3142 .kind = .segment_sum,
3143 .segments = 8,
3144 .total = 64,
3145 .schedule = .{ .thread_blocks = 64 },
3146 } }),
3147 );
3148
3149 try std.testing.expectEqual(
3150 @as(?OwnedDescriptor, null),
3151 try selectOwned(std.testing.allocator, .{ .segmented = .{
3152 .dtype = .f32,
3153 .kind = .segment_sum,
3154 .segments = 8,
3155 .total = 64,
3156 .schedule = .{ .warp_blocks = 48 },
3157 } }),
3158 );
3159 }
3160
3161 test "kernel library catalog builds segment sum scheduled artifact registry" {
3162 const allocator = std.testing.allocator;
3163 var state = gpu.recording.BackendState{
3164 .allocator = allocator,
3165 .kind = .cuda,
3166 .format = .cuda_ptx,
3167 };
3168 var registry_descriptors: [2]OwnedDescriptor = undefined;
3169 var descriptor_count: usize = 0;
3170 defer for (registry_descriptors[0..descriptor_count]) |*descriptor| descriptor.deinit();
3171 registry_descriptors[0] = (try selectOwned(allocator, .{ .segmented = .{
3172 .dtype = .f32,
3173 .kind = .segment_sum,
3174 .segments = 1000,
3175 .total = 65536,
3176 .schedule = .{ .thread_blocks = 256 },
3177 } })) orelse return error.TestExpectedCatalogDescriptor;
3178 descriptor_count = 1;
3179 registry_descriptors[1] = (try selectOwned(allocator, .{ .segmented = .{
3180 .dtype = .f32,
3181 .kind = .segment_sum,
3182 .segments = 1000,
3183 .total = 65536,
3184 .schedule = .{ .warp_blocks = 256 },
3185 } })) orelse return error.TestExpectedCatalogDescriptor;
3186 descriptor_count = 2;
3187
3188 var owned_registry = try createOwnedKernelCallArtifactRegistry(allocator, state.handle(), registry_descriptors[0..], .{ .limits = .testing });
3189 defer owned_registry.deinit();
3190 const registry = owned_registry.registry();
3191 const selected = registry.find("accy.kernel.segmented.segment_sum_family_thread_256_f32", segmented.segment_sum_family_version, .cuda_ptx) orelse {
3192 return error.TestExpectedKernelCallArtifact;
3193 };
3194
3195 try std.testing.expectEqual(@as(usize, 2), registry.entries.len);
3196 try std.testing.expectEqual(@as(u32, 5), selected.argument_count);
3197 try std.testing.expectEqual(@as(u32, 2), selected.runtime_scalar_argument_count);
3198 try std.testing.expect(selected.required_dtypes.contains(.i32));
3199 const launch = switch (selected.launch) {
3200 .derived => |derived| derived,
3201 .fixed => return error.TestExpectedDerivedLaunch,
3202 };
3203 try std.testing.expectEqual(@as(u32, 256), launch.threadgroup[0]);
3204
3205 const warp_selected = registry.find("accy.kernel.segmented.segment_sum_family_warp_256_f32", segmented.segment_sum_family_version, .cuda_ptx) orelse {
3206 return error.TestExpectedKernelCallArtifact;
3207 };
3208 const warp_launch = switch (warp_selected.launch) {
3209 .derived => |derived| derived,
3210 .fixed => return error.TestExpectedDerivedLaunch,
3211 };
3212 try std.testing.expectEqual(@as(u32, 256), warp_launch.threadgroup[0]);
3213 switch (warp_launch.grid[0]) {
3214 .runtime_u32_ceil_div => |term| try std.testing.expectEqual(@as(u32, 8), term.divisor),
3215 else => return error.TestExpectedDerivedLaunch,
3216 }
3217 }
3218
3219 test "kernel library catalog selects scatter descriptors" {
3220 const static_descriptor = select(.{ .scatter = .{
3221 .dtype = .f32,
3222 .axis_size = 8,
3223 .updates = 4,
3224 } }) orelse return error.TestExpectedCatalogDescriptor;
3225 try std.testing.expectEqualStrings(indexing.Scatter8F32.target, static_descriptor.metadata.target);
3226
3227 var owned = (try selectOwned(std.testing.allocator, .{ .scatter = .{
3228 .dtype = .f32,
3229 .outer = 2,
3230 .axis_size = 1024,
3231 .updates = 500,
3232 .inner = 1,
3233 } })) orelse return error.TestExpectedCatalogDescriptor;
3234 defer owned.deinit();
3235 try std.testing.expectEqualStrings("accy.kernel.indexing.scatter_family_256_f32", owned.descriptor.metadata.target);
3236
3237 var scheduled = (try selectOwned(std.testing.allocator, .{ .scatter = .{
3238 .dtype = .f32,
3239 .outer = 2,
3240 .axis_size = 1024,
3241 .updates = 500,
3242 .inner = 1,
3243 .schedule = .{ .thread_blocks = 64 },
3244 } })) orelse return error.TestExpectedCatalogDescriptor;
3245 defer scheduled.deinit();
3246 try std.testing.expectEqualStrings("accy.kernel.indexing.scatter_family_64_f32", scheduled.descriptor.metadata.target);
3247
3248 try std.testing.expectEqual(
3249 @as(?OwnedDescriptor, null),
3250 try selectOwned(std.testing.allocator, .{ .scatter = .{ .dtype = .f32, .axis_size = 0, .updates = 4 } }),
3251 );
3252 }
3253
3254 test "kernel library catalog builds scatter scheduled artifact registry" {
3255 const allocator = std.testing.allocator;
3256 var state = gpu.recording.BackendState{
3257 .allocator = allocator,
3258 .kind = .cuda,
3259 .format = .cuda_ptx,
3260 };
3261 var candidates = try selectOwnedScatterCandidates(allocator, .{
3262 .dtype = .f32,
3263 .outer = 2,
3264 .axis_size = 1024,
3265 .updates = 500,
3266 .inner = 1,
3267 .schedule = .{ .thread_blocks = 256 },
3268 });
3269 defer candidates.deinit();
3270
3271 try std.testing.expectEqual(@as(usize, 1), candidates.count);
3272 var owned_registry = try createOwnedKernelCallArtifactRegistry(allocator, state.handle(), candidates.slice(), .{ .limits = .testing });
3273 defer owned_registry.deinit();
3274 const registry = owned_registry.registry();
3275 const selected = registry.find("accy.kernel.indexing.scatter_family_256_f32", indexing.scatter_family_version, .cuda_ptx) orelse {
3276 return error.TestExpectedKernelCallArtifact;
3277 };
3278
3279 try std.testing.expectEqual(candidates.count, registry.entries.len);
3280 try std.testing.expectEqual(@as(u32, 9), selected.argument_count);
3281 try std.testing.expectEqual(@as(u32, 5), selected.runtime_scalar_argument_count);
3282 const launch = switch (selected.launch) {
3283 .derived => |derived| derived,
3284 .fixed => return error.TestExpectedDerivedLaunch,
3285 };
3286 try std.testing.expectEqual(@as(u32, 256), launch.threadgroup[0]);
3287 }
3288
3289 test "kernel library catalog builds scatter add scheduled artifact registry" {
3290 const allocator = std.testing.allocator;
3291 var state = gpu.recording.BackendState{
3292 .allocator = allocator,
3293 .kind = .cuda,
3294 .format = .cuda_ptx,
3295 };
3296 var candidates = try selectOwnedScatterAddCandidates(allocator, .{
3297 .dtype = .i32,
3298 .axis_size = 1024,
3299 .updates = 4096,
3300 .schedule = .{ .thread_blocks = 128 },
3301 });
3302 defer candidates.deinit();
3303
3304 try std.testing.expectEqual(@as(usize, 1), candidates.count);
3305 var owned_registry = try createOwnedKernelCallArtifactRegistry(allocator, state.handle(), candidates.slice(), .{ .limits = .testing });
3306 defer owned_registry.deinit();
3307 const registry = owned_registry.registry();
3308 const selected = registry.find("accy.kernel.indexing.scatter_add_family_128_i32", indexing.scatter_add_family_version, .cuda_ptx) orelse {
3309 return error.TestExpectedKernelCallArtifact;
3310 };
3311 const profile = selected.shape_profile orelse return error.TestExpectedShapeProfile;
3312
3313 try std.testing.expectEqual(candidates.count, registry.entries.len);
3314 try std.testing.expectEqual(@as(u32, 9), selected.argument_count);
3315 try std.testing.expectEqual(@as(u32, 5), selected.runtime_scalar_argument_count);
3316 try std.testing.expectEqualStrings("scatter_add", profile.name);
3317 try std.testing.expect(selected.required_dtypes.contains(.i32));
3318 const launch = switch (selected.launch) {
3319 .derived => |derived| derived,
3320 .fixed => return error.TestExpectedDerivedLaunch,
3321 };
3322 try std.testing.expectEqual(@as(u32, 128), launch.threadgroup[0]);
3323 }
3324 };
3325
3326 const scan_tests = struct {
3327 test "kernel library catalog selects prefix sum descriptors" {
3328 const static_descriptor = select(.{ .scan = .{
3329 .dtype = .f32,
3330 .kind = .prefix_sum,
3331 .extent = 8,
3332 .schedule = .{ .thread_blocks = 32 },
3333 } }) orelse return error.TestExpectedCatalogDescriptor;
3334 try std.testing.expectEqualStrings(scan_mod.PrefixSum8F32.target, static_descriptor.metadata.target);
3335
3336 var owned = (try selectOwned(std.testing.allocator, .{ .scan = .{
3337 .dtype = .f32,
3338 .kind = .prefix_sum,
3339 .extent = 100,
3340 } })) orelse return error.TestExpectedCatalogDescriptor;
3341 defer owned.deinit();
3342 try std.testing.expectEqualStrings("accy.kernel.scan.prefix_sum_family_128_f32", owned.descriptor.metadata.target);
3343
3344 var scheduled = (try selectOwned(std.testing.allocator, .{ .scan = .{
3345 .dtype = .f32,
3346 .kind = .prefix_sum,
3347 .extent = 100,
3348 .schedule = .{ .thread_blocks = 256 },
3349 } })) orelse return error.TestExpectedCatalogDescriptor;
3350 defer scheduled.deinit();
3351 try std.testing.expectEqualStrings("accy.kernel.scan.prefix_sum_family_256_f32", scheduled.descriptor.metadata.target);
3352
3353 var u32_owned = (try selectOwned(std.testing.allocator, .{ .scan = .{
3354 .dtype = .u32,
3355 .kind = .prefix_sum,
3356 .extent = 100,
3357 } })) orelse return error.TestExpectedCatalogDescriptor;
3358 defer u32_owned.deinit();
3359 try std.testing.expectEqualStrings("accy.kernel.scan.prefix_sum_family_128_u32", u32_owned.descriptor.metadata.target);
3360
3361 try std.testing.expectEqual(
3362 @as(?OwnedDescriptor, null),
3363 try selectOwned(std.testing.allocator, .{ .scan = .{
3364 .dtype = .f32,
3365 .kind = .prefix_sum,
3366 .extent = 100,
3367 .schedule = .{ .thread_blocks = 64 },
3368 } }),
3369 );
3370 }
3371
3372 test "kernel library catalog selects device scan pipeline descriptors past one block" {
3373 var owned = (try selectOwned(std.testing.allocator, .{ .scan = .{
3374 .dtype = .f32,
3375 .kind = .prefix_sum,
3376 .extent = 2000,
3377 } })) orelse return error.TestExpectedCatalogDescriptor;
3378 defer owned.deinit();
3379 try std.testing.expectEqualStrings("accy.kernel.scan.device_prefix_sum_family_32_f32", owned.descriptor.metadata.target);
3380 try std.testing.expectEqual(scan_mod.device_scan_family_version, owned.descriptor.metadata.version);
3381
3382 var capped = (try selectOwned(std.testing.allocator, .{ .scan = .{
3383 .dtype = .f32,
3384 .kind = .prefix_sum,
3385 .extent = 1 << 20,
3386 } })) orelse return error.TestExpectedCatalogDescriptor;
3387 defer capped.deinit();
3388 try std.testing.expectEqualStrings("accy.kernel.scan.device_prefix_sum_family_1024_f32", capped.descriptor.metadata.target);
3389
3390 try std.testing.expectEqual(
3391 @as(?OwnedDescriptor, null),
3392 try selectOwned(std.testing.allocator, .{ .scan = .{
3393 .dtype = .f32,
3394 .kind = .prefix_sum,
3395 .extent = (1 << 20) + 1,
3396 } }),
3397 );
3398
3399 var exclusive = (try selectOwned(std.testing.allocator, .{ .scan = .{
3400 .dtype = .f32,
3401 .kind = .prefix_sum_exclusive,
3402 .extent = 2000,
3403 } })) orelse return error.TestExpectedCatalogDescriptor;
3404 defer exclusive.deinit();
3405 try std.testing.expectEqualStrings(
3406 "accy.kernel.scan.device_prefix_sum_exclusive_family_32_f32",
3407 exclusive.descriptor.metadata.target,
3408 );
3409
3410 var scheduled = (try selectOwned(std.testing.allocator, .{ .scan = .{
3411 .dtype = .f32,
3412 .kind = .prefix_sum,
3413 .extent = 2000,
3414 .schedule = .{ .thread_blocks = 128 },
3415 } })) orelse return error.TestExpectedCatalogDescriptor;
3416 defer scheduled.deinit();
3417 try std.testing.expectEqualStrings("accy.kernel.scan.device_prefix_sum_family_128_f32", scheduled.descriptor.metadata.target);
3418
3419 var f16_owned = (try selectOwned(std.testing.allocator, .{ .scan = .{
3420 .dtype = .f16,
3421 .kind = .prefix_sum,
3422 .extent = 2000,
3423 } })) orelse return error.TestExpectedCatalogDescriptor;
3424 defer f16_owned.deinit();
3425 try std.testing.expectEqualStrings("accy.kernel.scan.device_prefix_sum_family_32_f16", f16_owned.descriptor.metadata.target);
3426
3427 var u32_owned = (try selectOwned(std.testing.allocator, .{ .scan = .{
3428 .dtype = .u32,
3429 .kind = .prefix_sum,
3430 .extent = 2000,
3431 } })) orelse return error.TestExpectedCatalogDescriptor;
3432 defer u32_owned.deinit();
3433 try std.testing.expectEqualStrings("accy.kernel.scan.device_prefix_sum_family_32_u32", u32_owned.descriptor.metadata.target);
3434 }
3435
3436 test "kernel library catalog device scan pipeline descriptor declines artifact creation" {
3437 const allocator = std.testing.allocator;
3438 var state = gpu.recording.BackendState{
3439 .allocator = allocator,
3440 .kind = .cuda,
3441 .format = .cuda_ptx,
3442 };
3443 var owned = (try selectOwned(allocator, .{ .scan = .{
3444 .dtype = .f32,
3445 .kind = .prefix_sum,
3446 .extent = 2000,
3447 } })) orelse return error.TestExpectedCatalogDescriptor;
3448 defer owned.deinit();
3449 try std.testing.expectError(
3450 error.UnknownKernelLibraryEntry,
3451 createOwnedKernelCallArtifact(allocator, state.handle(), owned, .{ .limits = .testing }),
3452 );
3453 }
3454
3455 test "kernel library catalog selects exclusive prefix sum descriptors" {
3456 var owned = (try selectOwned(std.testing.allocator, .{ .scan = .{
3457 .dtype = .f32,
3458 .kind = .prefix_sum_exclusive,
3459 .extent = 100,
3460 } })) orelse return error.TestExpectedCatalogDescriptor;
3461 defer owned.deinit();
3462 try std.testing.expectEqualStrings("accy.kernel.scan.prefix_sum_exclusive_family_128_f32", owned.descriptor.metadata.target);
3463 }
3464
3465 test "kernel library catalog builds prefix sum scheduled artifact registry" {
3466 const allocator = std.testing.allocator;
3467 var state = gpu.recording.BackendState{
3468 .allocator = allocator,
3469 .kind = .cuda,
3470 .format = .cuda_ptx,
3471 };
3472 var candidates = try selectOwnedPrefixSumCandidates(allocator, .{
3473 .dtype = .f32,
3474 .kind = .prefix_sum,
3475 .extent = 100,
3476 .schedule = .{ .thread_blocks = 128 },
3477 });
3478 defer candidates.deinit();
3479
3480 try std.testing.expectEqual(@as(usize, 1), candidates.count);
3481 var owned_registry = try createOwnedKernelCallArtifactRegistry(allocator, state.handle(), candidates.slice(), .{ .limits = .testing });
3482 defer owned_registry.deinit();
3483 const registry = owned_registry.registry();
3484 const selected = registry.find("accy.kernel.scan.prefix_sum_family_128_f32", scan_mod.prefix_sum_family_version, .cuda_ptx) orelse {
3485 return error.TestExpectedKernelCallArtifact;
3486 };
3487
3488 try std.testing.expectEqual(candidates.count, registry.entries.len);
3489 try std.testing.expectEqual(@as(u32, 3), selected.argument_count);
3490 try std.testing.expectEqual(@as(u32, 1), selected.runtime_scalar_argument_count);
3491 }
3492
3493 test "kernel library catalog dispatches selected device scan pipelines" {
3494 const allocator = std.testing.allocator;
3495 var state = gpu.recording.BackendState{
3496 .allocator = allocator,
3497 .kind = .cuda,
3498 .format = .cuda_ptx,
3499 };
3500
3501 var owned = (try selectOwned(allocator, .{ .scan = .{
3502 .dtype = .f32,
3503 .kind = .prefix_sum,
3504 .extent = 2000,
3505 } })) orelse return error.TestExpectedCatalogDescriptor;
3506 defer owned.deinit();
3507
3508 var package = (try createOwnedKernelCallPipelinePackage(allocator, state.handle(), owned, .{ .limits = .testing })) orelse {
3509 return error.TestExpectedPipelinePackage;
3510 };
3511 defer package.deinit();
3512
3513 try std.testing.expectEqual(@as(usize, 3), package.artifacts.len);
3514 try std.testing.expectEqualStrings(
3515 "accy.kernel.scan.device_prefix_sum_family_32_f32",
3516 package.pipeline.value.target,
3517 );
3518 try std.testing.expect(package.registry().find(
3519 "accy.kernel.scan.device_prefix_sum_block_scan_family_32_f32",
3520 scan_mod.device_scan_family_version,
3521 .cuda_ptx,
3522 ) != null);
3523
3524 const encoded = try accy_artifact.wire.encode(allocator, package.entries, &.{package.pipeline.value});
3525 defer allocator.free(encoded);
3526 var decoded = try accy_artifact.wire.decode(allocator, encoded);
3527 defer decoded.deinit();
3528 const found = accy_artifact.findPipeline(
3529 decoded.pipelines,
3530 package.pipeline.value.target,
3531 scan_mod.device_scan_family_version,
3532 ) orelse return error.TestExpectedPipeline;
3533 try found.validate(decoded.registry(), .cuda_ptx);
3534
3535 var f16_owned = (try selectOwned(allocator, .{ .scan = .{
3536 .dtype = .f16,
3537 .kind = .prefix_sum,
3538 .extent = 2000,
3539 } })) orelse return error.TestExpectedCatalogDescriptor;
3540 defer f16_owned.deinit();
3541
3542 var f16_package = (try createOwnedKernelCallPipelinePackage(allocator, state.handle(), f16_owned, .{ .limits = .testing })) orelse {
3543 return error.TestExpectedPipelinePackage;
3544 };
3545 defer f16_package.deinit();
3546
3547 try std.testing.expectEqualStrings(
3548 "accy.kernel.scan.device_prefix_sum_family_32_f16",
3549 f16_package.pipeline.value.target,
3550 );
3551 try std.testing.expectEqual(@as(usize, 3), f16_package.pipeline.value.intermediates.len);
3552 try std.testing.expect(f16_package.registry().find(
3553 "accy.kernel.scan.device_prefix_sum_block_scan_family_32_f16",
3554 scan_mod.device_scan_family_version,
3555 .cuda_ptx,
3556 ) != null);
3557 try std.testing.expect(f16_package.registry().find(
3558 "accy.kernel.scan.device_add_base_family_32_f16",
3559 scan_mod.device_scan_family_version,
3560 .cuda_ptx,
3561 ) != null);
3562
3563 const f16_encoded = try accy_artifact.wire.encode(allocator, f16_package.entries, &.{f16_package.pipeline.value});
3564 defer allocator.free(f16_encoded);
3565 var f16_decoded = try accy_artifact.wire.decode(allocator, f16_encoded);
3566 defer f16_decoded.deinit();
3567 const f16_found = accy_artifact.findPipeline(
3568 f16_decoded.pipelines,
3569 f16_package.pipeline.value.target,
3570 scan_mod.device_scan_family_version,
3571 ) orelse return error.TestExpectedPipeline;
3572 try f16_found.validate(f16_decoded.registry(), .cuda_ptx);
3573
3574 var u32_owned = (try selectOwned(allocator, .{ .scan = .{
3575 .dtype = .u32,
3576 .kind = .prefix_sum,
3577 .extent = 2000,
3578 } })) orelse return error.TestExpectedCatalogDescriptor;
3579 defer u32_owned.deinit();
3580
3581 var u32_package = (try createOwnedKernelCallPipelinePackage(allocator, state.handle(), u32_owned, .{ .limits = .testing })) orelse {
3582 return error.TestExpectedPipelinePackage;
3583 };
3584 defer u32_package.deinit();
3585
3586 try std.testing.expectEqualStrings(
3587 "accy.kernel.scan.device_prefix_sum_family_32_u32",
3588 u32_package.pipeline.value.target,
3589 );
3590 try std.testing.expectEqual(@as(usize, 2), u32_package.pipeline.value.intermediates.len);
3591 try std.testing.expectEqual(choir_abi.DType.u32, u32_package.pipeline.value.intermediates[0].dtype);
3592 try std.testing.expectEqual(choir_abi.DType.u32, u32_package.pipeline.value.intermediates[1].dtype);
3593 try std.testing.expect(u32_package.registry().find(
3594 "accy.kernel.scan.device_prefix_sum_block_scan_family_32_u32",
3595 scan_mod.device_scan_family_version,
3596 .cuda_ptx,
3597 ) != null);
3598 try std.testing.expect(u32_package.registry().find(
3599 "accy.kernel.scan.device_add_base_family_32_u32",
3600 scan_mod.device_scan_family_version,
3601 .cuda_ptx,
3602 ) != null);
3603
3604 const u32_encoded = try accy_artifact.wire.encode(allocator, u32_package.entries, &.{u32_package.pipeline.value});
3605 defer allocator.free(u32_encoded);
3606 var u32_decoded = try accy_artifact.wire.decode(allocator, u32_encoded);
3607 defer u32_decoded.deinit();
3608 const u32_found = accy_artifact.findPipeline(
3609 u32_decoded.pipelines,
3610 u32_package.pipeline.value.target,
3611 scan_mod.device_scan_family_version,
3612 ) orelse return error.TestExpectedPipeline;
3613 try u32_found.validate(u32_decoded.registry(), .cuda_ptx);
3614
3615 var single_block = (try selectOwned(allocator, .{ .scan = .{
3616 .dtype = .f32,
3617 .kind = .prefix_sum,
3618 .extent = 100,
3619 } })) orelse return error.TestExpectedCatalogDescriptor;
3620 defer single_block.deinit();
3621 try std.testing.expectEqual(
3622 @as(?OwnedKernelCallPipelinePackage, null),
3623 try createOwnedKernelCallPipelinePackage(allocator, state.handle(), single_block, .{ .limits = .testing }),
3624 );
3625 }
3626 };
3627
3628 const sparse_tests = struct {
3629 const tuning_mod = tuning;
3630 const SparseQuery = catalog.SparseQuery;
3631 const selectSparse = selectOwnedSparse;
3632 const selectSparseCandidates = selectOwnedSparseCandidates;
3633
3634 test "catalog sparse selection resolves structures from density and overrides" {
3635 const allocator = std.testing.allocator;
3636
3637 var dense = (try selectSparse(allocator, .{
3638 .dtype = .f32,
3639 .kind = .{ .csr_spmv = .{ .rows = 70, .nnz = 560, .x_extent = 40 } },
3640 })) orelse return error.TestExpectedSparseDescriptor;
3641 defer dense.deinit();
3642 try std.testing.expectEqualStrings(
3643 "accy.kernel.sparse.spmv_csr_row_warp_family_256_f32",
3644 dense.descriptor.metadata.target,
3645 );
3646
3647 var sparse_rows = (try selectSparse(allocator, .{
3648 .dtype = .f32,
3649 .kind = .{ .csr_spmv = .{ .rows = 70, .nnz = 90, .x_extent = 40 } },
3650 })) orelse return error.TestExpectedSparseDescriptor;
3651 defer sparse_rows.deinit();
3652 try std.testing.expectEqualStrings(
3653 "accy.kernel.sparse.spmv_csr_row_thread_family_70_f32",
3654 sparse_rows.descriptor.metadata.target,
3655 );
3656
3657 var overridden = (try selectSparse(allocator, .{
3658 .dtype = .f32,
3659 .kind = .{ .csr_spmv = .{ .rows = 70, .nnz = 90, .x_extent = 40 } },
3660 .structure = .row_warp,
3661 .schedule = .{ .thread_blocks = 64 },
3662 })) orelse return error.TestExpectedSparseDescriptor;
3663 defer overridden.deinit();
3664 try std.testing.expectEqualStrings(
3665 "accy.kernel.sparse.spmv_csr_row_warp_family_64_f32",
3666 overridden.descriptor.metadata.target,
3667 );
3668 try std.testing.expect(overridden.descriptor.metadata.specialization.structureIs("row_warp"));
3669
3670 var half = (try selectSparse(allocator, .{
3671 .dtype = .f16,
3672 .kind = .{ .csr_spmv = .{ .rows = 70, .nnz = 560, .x_extent = 40 } },
3673 .schedule = .{ .thread_blocks = 64 },
3674 })) orelse return error.TestExpectedSparseDescriptor;
3675 defer half.deinit();
3676 try std.testing.expectEqualStrings(
3677 "accy.kernel.sparse.spmv_csr_row_warp_family_64_f16",
3678 half.descriptor.metadata.target,
3679 );
3680 try std.testing.expectEqual(@as(?choir_abi.DType, .f16), half.descriptor.metadata.specialization.dtype);
3681 try std.testing.expectEqual(@as(?choir_abi.DType, .f32), half.descriptor.metadata.specialization.accumulation_dtype);
3682
3683 var double = (try selectSparse(allocator, .{
3684 .dtype = .f64,
3685 .kind = .{ .csr_spmv = .{ .rows = 70, .nnz = 560, .x_extent = 40 } },
3686 .schedule = .{ .thread_blocks = 64 },
3687 })) orelse return error.TestExpectedSparseDescriptor;
3688 defer double.deinit();
3689 try std.testing.expectEqualStrings(
3690 "accy.kernel.sparse.spmv_csr_row_warp_family_64_f64",
3691 double.descriptor.metadata.target,
3692 );
3693 try std.testing.expectEqual(@as(?choir_abi.DType, .f64), double.descriptor.metadata.specialization.dtype);
3694 try std.testing.expectEqual(@as(?choir_abi.DType, .f64), double.descriptor.metadata.specialization.accumulation_dtype);
3695
3696 var coo = (try selectSparse(allocator, .{
3697 .dtype = .f32,
3698 .kind = .{ .coo_spmv = .{ .rows = 70, .nnz = 560, .x_extent = 40 } },
3699 .schedule = .{ .thread_blocks = 64 },
3700 })) orelse return error.TestExpectedSparseDescriptor;
3701 defer coo.deinit();
3702 try std.testing.expectEqualStrings(
3703 "accy.kernel.sparse.spmv_coo_element_thread_family_64_f32",
3704 coo.descriptor.metadata.target,
3705 );
3706 try std.testing.expect(coo.descriptor.metadata.specialization.operationIs(.{ .sparse = .coo_spmv }));
3707 try std.testing.expect(coo.descriptor.metadata.specialization.structureIs("element_thread"));
3708 try std.testing.expectEqual(@as(?choir_abi.DType, .f32), coo.descriptor.metadata.specialization.accumulation_dtype);
3709
3710 var coo_half = (try selectSparse(allocator, .{
3711 .dtype = .f16,
3712 .kind = .{ .coo_spmv = .{ .rows = 70, .nnz = 560, .x_extent = 40 } },
3713 .schedule = .{ .thread_blocks = 64 },
3714 })) orelse return error.TestExpectedSparseDescriptor;
3715 defer coo_half.deinit();
3716 try std.testing.expectEqualStrings(
3717 "accy.kernel.sparse.spmv_coo_row_thread_family_64_f16",
3718 coo_half.descriptor.metadata.target,
3719 );
3720 try std.testing.expect(coo_half.descriptor.metadata.specialization.structureIs("row_thread"));
3721 try std.testing.expectEqual(@as(?choir_abi.DType, .f16), coo_half.descriptor.metadata.specialization.dtype);
3722 try std.testing.expectEqual(@as(?choir_abi.DType, .f32), coo_half.descriptor.metadata.specialization.accumulation_dtype);
3723
3724 var coo_double = (try selectSparse(allocator, .{
3725 .dtype = .f64,
3726 .kind = .{ .coo_spmv = .{ .rows = 70, .nnz = 560, .x_extent = 40 } },
3727 .schedule = .{ .thread_blocks = 64 },
3728 })) orelse return error.TestExpectedSparseDescriptor;
3729 defer coo_double.deinit();
3730 try std.testing.expectEqualStrings(
3731 "accy.kernel.sparse.spmv_coo_row_thread_family_64_f64",
3732 coo_double.descriptor.metadata.target,
3733 );
3734 try std.testing.expect(coo_double.descriptor.metadata.specialization.structureIs("row_thread"));
3735 try std.testing.expectEqual(@as(?choir_abi.DType, .f64), coo_double.descriptor.metadata.specialization.dtype);
3736 try std.testing.expectEqual(@as(?choir_abi.DType, .f64), coo_double.descriptor.metadata.specialization.accumulation_dtype);
3737
3738 var coo_row_thread = (try selectSparse(allocator, .{
3739 .dtype = .f32,
3740 .kind = .{ .coo_spmv = .{ .rows = 70, .nnz = 560, .x_extent = 40 } },
3741 .structure = .row_thread,
3742 .schedule = .{ .thread_blocks = 64 },
3743 })) orelse return error.TestExpectedSparseDescriptor;
3744 defer coo_row_thread.deinit();
3745 try std.testing.expectEqualStrings(
3746 "accy.kernel.sparse.spmv_coo_row_thread_family_64_f32",
3747 coo_row_thread.descriptor.metadata.target,
3748 );
3749
3750 var spmm = (try selectSparse(allocator, .{
3751 .dtype = .f32,
3752 .kind = .{ .csr_spmm = .{ .rows = 70, .columns = 45, .nnz = 512, .x_extent = 40 } },
3753 .schedule = .{ .thread_blocks_2d = .{ .x = 8, .y = 4 } },
3754 })) orelse return error.TestExpectedSparseDescriptor;
3755 defer spmm.deinit();
3756 try std.testing.expectEqualStrings(
3757 "accy.kernel.sparse.spmm_csr_row_column_thread_family_8x4_f32",
3758 spmm.descriptor.metadata.target,
3759 );
3760 try std.testing.expect(spmm.descriptor.metadata.specialization.operationIs(.{ .sparse = .csr_spmm }));
3761 try std.testing.expect(spmm.descriptor.metadata.specialization.structureIs("row_column_thread"));
3762
3763 var ell = (try selectSparse(allocator, .{
3764 .dtype = .f32,
3765 .kind = .{ .ell_spmv = .{ .rows = 70, .slots = 8, .x_extent = 40 } },
3766 .schedule = .{ .thread_blocks = 64 },
3767 })) orelse return error.TestExpectedSparseDescriptor;
3768 defer ell.deinit();
3769 try std.testing.expectEqualStrings(
3770 "accy.kernel.sparse.spmv_ell_row_thread_family_64_f32",
3771 ell.descriptor.metadata.target,
3772 );
3773 try std.testing.expect(ell.descriptor.metadata.specialization.operationIs(.{ .sparse = .ell_spmv }));
3774 try std.testing.expect(ell.descriptor.metadata.specialization.structureIs("row_thread"));
3775
3776 var sell = (try selectSparse(allocator, .{
3777 .dtype = .f32,
3778 .kind = .{ .sell_spmv = .{ .rows = 70, .slice_size = 8, .values_size = 400, .x_extent = 40 } },
3779 .schedule = .{ .thread_blocks = 64 },
3780 })) orelse return error.TestExpectedSparseDescriptor;
3781 defer sell.deinit();
3782 try std.testing.expectEqualStrings(
3783 "accy.kernel.sparse.spmv_sell_row_thread_slice8_family_64_f32",
3784 sell.descriptor.metadata.target,
3785 );
3786 try std.testing.expect(sell.descriptor.metadata.specialization.operationIs(.{ .sparse = .sell_spmv }));
3787 try std.testing.expect(sell.descriptor.metadata.specialization.structureIs("row_thread"));
3788 try std.testing.expect(sell.descriptor.metadata.specialization.staticParameterMatches(sparse_mod.spmv_sell_slice_size_parameter, 8));
3789
3790 var spmm_double = (try selectSparse(allocator, .{
3791 .dtype = .f64,
3792 .kind = .{ .csr_spmm = .{ .rows = 70, .columns = 45, .nnz = 512, .x_extent = 40 } },
3793 .schedule = .{ .thread_blocks_2d = .{ .x = 8, .y = 4 } },
3794 })) orelse return error.TestExpectedSparseDescriptor;
3795 defer spmm_double.deinit();
3796 try std.testing.expectEqualStrings(
3797 "accy.kernel.sparse.spmm_csr_row_column_thread_family_8x4_f64",
3798 spmm_double.descriptor.metadata.target,
3799 );
3800 try std.testing.expectEqual(@as(?choir_abi.DType, .f64), spmm_double.descriptor.metadata.specialization.dtype);
3801 try std.testing.expectEqual(@as(?choir_abi.DType, .f64), spmm_double.descriptor.metadata.specialization.accumulation_dtype);
3802
3803 try std.testing.expectEqual(@as(?OwnedDescriptor, null), try selectSparse(allocator, .{
3804 .dtype = .i32,
3805 .kind = .{ .csr_spmv = .{ .rows = 70, .nnz = 560, .x_extent = 40 } },
3806 }));
3807 try std.testing.expectEqual(@as(?OwnedDescriptor, null), try selectSparse(allocator, .{
3808 .dtype = .f32,
3809 .kind = .{ .csr_spmv = .{ .rows = 70, .nnz = 0, .x_extent = 40 } },
3810 }));
3811 try std.testing.expectEqual(@as(?OwnedDescriptor, null), try selectSparse(allocator, .{
3812 .dtype = .f64,
3813 .kind = .{ .coo_spmv = .{ .rows = 70, .nnz = 560, .x_extent = 40 } },
3814 .structure = .element_thread,
3815 }));
3816 try std.testing.expectEqual(@as(?OwnedDescriptor, null), try selectSparse(allocator, .{
3817 .dtype = .f16,
3818 .kind = .{ .coo_spmv = .{ .rows = 70, .nnz = 560, .x_extent = 40 } },
3819 .structure = .element_thread,
3820 }));
3821 try std.testing.expectEqual(@as(?OwnedDescriptor, null), try selectSparse(allocator, .{
3822 .dtype = .f32,
3823 .kind = .{ .coo_spmv = .{ .rows = 70, .nnz = 0, .x_extent = 40 } },
3824 }));
3825 try std.testing.expectEqual(@as(?OwnedDescriptor, null), try selectSparse(allocator, .{
3826 .dtype = .f32,
3827 .kind = .{ .coo_spmv = .{ .rows = 70, .nnz = 560, .x_extent = 40 } },
3828 .schedule = .{ .thread_blocks_2d = .{ .x = 8, .y = 4 } },
3829 }));
3830 try std.testing.expectEqual(@as(?OwnedDescriptor, null), try selectSparse(allocator, .{
3831 .dtype = .f32,
3832 .kind = .{ .csr_spmv = .{ .rows = 70, .nnz = 90, .x_extent = 40 } },
3833 .structure = .row_warp,
3834 .schedule = .{ .thread_blocks = 48 },
3835 }));
3836 try std.testing.expectEqual(@as(?OwnedDescriptor, null), try selectSparse(allocator, .{
3837 .dtype = .f32,
3838 .kind = .{ .csr_spmm = .{ .rows = 70, .columns = 45, .nnz = 512, .x_extent = 40 } },
3839 .structure = .row_warp,
3840 .schedule = .{ .thread_blocks_2d = .{ .x = 8, .y = 4 } },
3841 }));
3842 try std.testing.expectEqual(@as(?OwnedDescriptor, null), try selectSparse(allocator, .{
3843 .dtype = .f32,
3844 .kind = .{ .csr_spmm = .{ .rows = 70, .columns = 45, .nnz = 512, .x_extent = 40 } },
3845 .schedule = .{ .thread_blocks = 64 },
3846 }));
3847 try std.testing.expectEqual(@as(?OwnedDescriptor, null), try selectSparse(allocator, .{
3848 .dtype = .f32,
3849 .kind = .{ .ell_spmv = .{ .rows = 70, .slots = 0, .x_extent = 40 } },
3850 }));
3851 try std.testing.expectEqual(@as(?OwnedDescriptor, null), try selectSparse(allocator, .{
3852 .dtype = .f32,
3853 .kind = .{ .ell_spmv = .{ .rows = 70, .slots = 8, .x_extent = 40 } },
3854 .structure = .row_warp,
3855 .schedule = .{ .thread_blocks = 64 },
3856 }));
3857 try std.testing.expectEqual(@as(?OwnedDescriptor, null), try selectSparse(allocator, .{
3858 .dtype = .f32,
3859 .kind = .{ .ell_spmv = .{ .rows = 70, .slots = 8, .x_extent = 40 } },
3860 .schedule = .{ .thread_blocks_2d = .{ .x = 8, .y = 4 } },
3861 }));
3862 try std.testing.expectEqual(@as(?OwnedDescriptor, null), try selectSparse(allocator, .{
3863 .dtype = .f32,
3864 .kind = .{ .sell_spmv = .{ .rows = 70, .slice_size = 0, .values_size = 400, .x_extent = 40 } },
3865 }));
3866 try std.testing.expectEqual(@as(?OwnedDescriptor, null), try selectSparse(allocator, .{
3867 .dtype = .f32,
3868 .kind = .{ .sell_spmv = .{ .rows = 70, .slice_size = 8, .values_size = 400, .x_extent = 40 } },
3869 .structure = .row_warp,
3870 .schedule = .{ .thread_blocks = 64 },
3871 }));
3872 try std.testing.expectEqual(@as(?OwnedDescriptor, null), try selectSparse(allocator, .{
3873 .dtype = .f32,
3874 .kind = .{ .sell_spmv = .{ .rows = 70, .slice_size = 8, .values_size = 400, .x_extent = 40 } },
3875 .schedule = .{ .thread_blocks_2d = .{ .x = 8, .y = 4 } },
3876 }));
3877 }
3878
3879 test "catalog sparse selection consults csr family tuning for density override" {
3880 const allocator = std.testing.allocator;
3881 const device: u64 = 0xacc70001;
3882 const query = SparseQuery{
3883 .dtype = .f32,
3884 .kind = .{ .csr_spmv = .{ .rows = 70, .nnz = 560, .x_extent = 40 } },
3885 };
3886 const probe = sparse_mod.SpmvCsr{
3887 .rows = 70,
3888 .nnz = 560,
3889 .x_extent = 40,
3890 .threads = 256,
3891 .structure = .row_warp,
3892 };
3893 var winner = probe;
3894 winner.structure = .row_thread;
3895 winner.threads = sparse_mod.spmvCsrRepresentableThreads(winner) orelse
3896 return error.TestExpectedSparseDescriptor;
3897 const winner_target = try sparse_mod.spmvCsrFamilyTarget(allocator, winner);
3898 defer allocator.free(winner_target);
3899 const records = [_]tuning_mod.FamilyTuningRecord{.{
3900 .key = try sparse_mod.spmvCsrFamilyTuningKey(allocator, device, probe),
3901 .target = winner_target,
3902 .winner_median_ns = 800,
3903 .runner_up_median_ns = 1100,
3904 .sample_count = 30,
3905 }};
3906 const reader = tuning_mod.FamilyTuningReader{
3907 .device_fingerprint = device,
3908 .table = .{ .records = records[0..] },
3909 };
3910
3911 var tuned = (try selectSparseWithTuning(allocator, query, reader)) orelse
3912 return error.TestExpectedSparseDescriptor;
3913 defer tuned.deinit();
3914 try std.testing.expectEqualStrings(
3915 "accy.kernel.sparse.spmv_csr_row_thread_family_70_f32",
3916 tuned.descriptor.metadata.target,
3917 );
3918 try std.testing.expect(tuned.descriptor.metadata.specialization.structureIs("row_thread"));
3919
3920 var miss = (try selectSparseWithTuning(allocator, .{
3921 .dtype = .f32,
3922 .kind = .{ .csr_spmv = .{ .rows = 70, .nnz = 561, .x_extent = 40 } },
3923 }, reader)) orelse return error.TestExpectedSparseDescriptor;
3924 defer miss.deinit();
3925 try std.testing.expectEqualStrings(
3926 "accy.kernel.sparse.spmv_csr_row_warp_family_256_f32",
3927 miss.descriptor.metadata.target,
3928 );
3929
3930 var explicit = (try selectSparseWithTuning(allocator, .{
3931 .dtype = .f32,
3932 .kind = .{ .csr_spmv = .{ .rows = 70, .nnz = 560, .x_extent = 40 } },
3933 .structure = .row_warp,
3934 }, reader)) orelse return error.TestExpectedSparseDescriptor;
3935 defer explicit.deinit();
3936 try std.testing.expectEqualStrings(
3937 "accy.kernel.sparse.spmv_csr_row_warp_family_256_f32",
3938 explicit.descriptor.metadata.target,
3939 );
3940 }
3941
3942 test "catalog sparse selection consults coo family tuning for structure override" {
3943 const allocator = std.testing.allocator;
3944 const device: u64 = 0xacc70002;
3945 const query = SparseQuery{
3946 .dtype = .f32,
3947 .kind = .{ .coo_spmv = .{ .rows = 70, .nnz = 560, .x_extent = 40 } },
3948 };
3949 const probe = sparse_mod.SpmvCoo{
3950 .rows = 70,
3951 .nnz = 560,
3952 .x_extent = 40,
3953 .threads = 256,
3954 .structure = .element_thread,
3955 };
3956 var winner = probe;
3957 winner.structure = .row_thread;
3958 winner.threads = sparse_mod.spmvCooRepresentableThreads(winner) orelse
3959 return error.TestExpectedSparseDescriptor;
3960 const winner_target = try sparse_mod.spmvCooFamilyTarget(allocator, winner);
3961 defer allocator.free(winner_target);
3962 const records = [_]tuning_mod.FamilyTuningRecord{.{
3963 .key = try sparse_mod.spmvCooFamilyTuningKey(allocator, device, probe),
3964 .target = winner_target,
3965 .winner_median_ns = 900,
3966 .runner_up_median_ns = 1200,
3967 .sample_count = 30,
3968 }};
3969 const reader = tuning_mod.FamilyTuningReader{
3970 .device_fingerprint = device,
3971 .table = .{ .records = records[0..] },
3972 };
3973
3974 var tuned = (try selectSparseWithTuning(allocator, query, reader)) orelse
3975 return error.TestExpectedSparseDescriptor;
3976 defer tuned.deinit();
3977 try std.testing.expectEqualStrings(
3978 "accy.kernel.sparse.spmv_coo_row_thread_family_70_f32",
3979 tuned.descriptor.metadata.target,
3980 );
3981 try std.testing.expect(tuned.descriptor.metadata.specialization.structureIs("row_thread"));
3982
3983 var miss = (try selectSparseWithTuning(allocator, .{
3984 .dtype = .f32,
3985 .kind = .{ .coo_spmv = .{ .rows = 70, .nnz = 561, .x_extent = 40 } },
3986 }, reader)) orelse return error.TestExpectedSparseDescriptor;
3987 defer miss.deinit();
3988 try std.testing.expectEqualStrings(
3989 "accy.kernel.sparse.spmv_coo_element_thread_family_256_f32",
3990 miss.descriptor.metadata.target,
3991 );
3992
3993 var explicit = (try selectSparseWithTuning(allocator, .{
3994 .dtype = .f32,
3995 .kind = .{ .coo_spmv = .{ .rows = 70, .nnz = 560, .x_extent = 40 } },
3996 .structure = .element_thread,
3997 }, reader)) orelse return error.TestExpectedSparseDescriptor;
3998 defer explicit.deinit();
3999 try std.testing.expectEqualStrings(
4000 "accy.kernel.sparse.spmv_coo_element_thread_family_256_f32",
4001 explicit.descriptor.metadata.target,
4002 );
4003 }
4004
4005 test "catalog sparse selection consults ell family tuning for thread blocks" {
4006 const allocator = std.testing.allocator;
4007 const device: u64 = 0xacc70003;
4008 const query = SparseQuery{
4009 .dtype = .f32,
4010 .kind = .{ .ell_spmv = .{ .rows = 70, .slots = 8, .x_extent = 40 } },
4011 };
4012 const probe = sparse_mod.SpmvEll{
4013 .rows = 70,
4014 .slots = 8,
4015 .x_extent = 40,
4016 .threads = 70,
4017 };
4018 var winner = probe;
4019 winner.threads = 64;
4020 const winner_target = try sparse_mod.spmvEllFamilyTarget(allocator, winner);
4021 defer allocator.free(winner_target);
4022 const records = [_]tuning_mod.FamilyTuningRecord{.{
4023 .key = try sparse_mod.spmvEllFamilyTuningKey(allocator, device, probe),
4024 .target = winner_target,
4025 .winner_median_ns = 900,
4026 .runner_up_median_ns = 1200,
4027 .sample_count = 30,
4028 }};
4029 const reader = tuning_mod.FamilyTuningReader{
4030 .device_fingerprint = device,
4031 .table = .{ .records = records[0..] },
4032 };
4033
4034 var tuned = (try selectSparseWithTuning(allocator, query, reader)) orelse
4035 return error.TestExpectedSparseDescriptor;
4036 defer tuned.deinit();
4037 try std.testing.expectEqualStrings(
4038 "accy.kernel.sparse.spmv_ell_row_thread_family_64_f32",
4039 tuned.descriptor.metadata.target,
4040 );
4041 try std.testing.expectEqual(@as(u32, 64), tuned.descriptor.metadata.specialization.launch.?.threadgroup[0]);
4042
4043 var miss = (try selectSparseWithTuning(allocator, .{
4044 .dtype = .f32,
4045 .kind = .{ .ell_spmv = .{ .rows = 70, .slots = 16, .x_extent = 40 } },
4046 }, reader)) orelse return error.TestExpectedSparseDescriptor;
4047 defer miss.deinit();
4048 try std.testing.expectEqualStrings(
4049 "accy.kernel.sparse.spmv_ell_row_thread_family_70_f32",
4050 miss.descriptor.metadata.target,
4051 );
4052
4053 var explicit = (try selectSparseWithTuning(allocator, .{
4054 .dtype = .f32,
4055 .kind = .{ .ell_spmv = .{ .rows = 70, .slots = 8, .x_extent = 40 } },
4056 .schedule = .{ .thread_blocks = 32 },
4057 }, reader)) orelse return error.TestExpectedSparseDescriptor;
4058 defer explicit.deinit();
4059 try std.testing.expectEqualStrings(
4060 "accy.kernel.sparse.spmv_ell_row_thread_family_32_f32",
4061 explicit.descriptor.metadata.target,
4062 );
4063 }
4064
4065 test "catalog sparse selection consults sell family tuning for thread blocks" {
4066 const allocator = std.testing.allocator;
4067 const device: u64 = 0xacc70004;
4068 const query = SparseQuery{
4069 .dtype = .f32,
4070 .kind = .{ .sell_spmv = .{ .rows = 70, .slice_size = 8, .values_size = 400, .x_extent = 40 } },
4071 };
4072 const probe = sparse_mod.SpmvSell{
4073 .rows = 70,
4074 .slice_size = 8,
4075 .values_size = 400,
4076 .x_extent = 40,
4077 .threads = 70,
4078 };
4079 var winner = probe;
4080 winner.threads = 64;
4081 const winner_target = try sparse_mod.spmvSellFamilyTarget(allocator, winner);
4082 defer allocator.free(winner_target);
4083 const records = [_]tuning_mod.FamilyTuningRecord{.{
4084 .key = try sparse_mod.spmvSellFamilyTuningKey(allocator, device, probe),
4085 .target = winner_target,
4086 .winner_median_ns = 900,
4087 .runner_up_median_ns = 1200,
4088 .sample_count = 30,
4089 }};
4090 const reader = tuning_mod.FamilyTuningReader{
4091 .device_fingerprint = device,
4092 .table = .{ .records = records[0..] },
4093 };
4094
4095 var tuned = (try selectSparseWithTuning(allocator, query, reader)) orelse
4096 return error.TestExpectedSparseDescriptor;
4097 defer tuned.deinit();
4098 try std.testing.expectEqualStrings(
4099 "accy.kernel.sparse.spmv_sell_row_thread_slice8_family_64_f32",
4100 tuned.descriptor.metadata.target,
4101 );
4102 try std.testing.expectEqual(@as(u32, 64), tuned.descriptor.metadata.specialization.launch.?.threadgroup[0]);
4103
4104 var miss = (try selectSparseWithTuning(allocator, .{
4105 .dtype = .f32,
4106 .kind = .{ .sell_spmv = .{ .rows = 70, .slice_size = 16, .values_size = 400, .x_extent = 40 } },
4107 }, reader)) orelse return error.TestExpectedSparseDescriptor;
4108 defer miss.deinit();
4109 try std.testing.expectEqualStrings(
4110 "accy.kernel.sparse.spmv_sell_row_thread_slice16_family_70_f32",
4111 miss.descriptor.metadata.target,
4112 );
4113
4114 var explicit = (try selectSparseWithTuning(allocator, .{
4115 .dtype = .f32,
4116 .kind = .{ .sell_spmv = .{ .rows = 70, .slice_size = 8, .values_size = 400, .x_extent = 40 } },
4117 .schedule = .{ .thread_blocks = 32 },
4118 }, reader)) orelse return error.TestExpectedSparseDescriptor;
4119 defer explicit.deinit();
4120 try std.testing.expectEqualStrings(
4121 "accy.kernel.sparse.spmv_sell_row_thread_slice8_family_32_f32",
4122 explicit.descriptor.metadata.target,
4123 );
4124 }
4125
4126 test "catalog sparse selection consults csr spmm family tuning for thread blocks" {
4127 const allocator = std.testing.allocator;
4128 const device: u64 = 0xacc70005;
4129 const query = SparseQuery{
4130 .dtype = .f32,
4131 .kind = .{ .csr_spmm = .{ .rows = 70, .columns = 45, .nnz = 512, .x_extent = 40 } },
4132 };
4133 const probe = sparse_mod.SpmmCsr{
4134 .rows = 70,
4135 .columns = 45,
4136 .nnz = 512,
4137 .x_extent = 40,
4138 .threads = sparse_mod.spmmCsrThreadsForExtents(70, 45),
4139 };
4140 const candidates = sparse_mod.spmmCsrThreadCandidatesForExtents(probe.rows, probe.columns);
4141 try std.testing.expect(candidates.count > 1);
4142 var winner = probe;
4143 winner.threads = candidates.items[1];
4144 winner.threads = sparse_mod.spmmCsrRepresentableThreads(winner) orelse
4145 return error.TestExpectedSparseDescriptor;
4146 const winner_target = try sparse_mod.spmmCsrFamilyTarget(allocator, winner);
4147 defer allocator.free(winner_target);
4148 const records = [_]tuning_mod.FamilyTuningRecord{.{
4149 .key = try sparse_mod.spmmCsrFamilyTuningKey(allocator, device, probe),
4150 .target = winner_target,
4151 .winner_median_ns = 900,
4152 .runner_up_median_ns = 1200,
4153 .sample_count = 30,
4154 }};
4155 const reader = tuning_mod.FamilyTuningReader{
4156 .device_fingerprint = device,
4157 .table = .{ .records = records[0..] },
4158 };
4159
4160 var tuned = (try selectSparseWithTuning(allocator, query, reader)) orelse
4161 return error.TestExpectedSparseDescriptor;
4162 defer tuned.deinit();
4163 try std.testing.expectEqualStrings(winner_target, tuned.descriptor.metadata.target);
4164 try std.testing.expectEqual(winner.threads.x, tuned.descriptor.metadata.specialization.launch.?.threadgroup[0]);
4165 try std.testing.expectEqual(winner.threads.y, tuned.descriptor.metadata.specialization.launch.?.threadgroup[1]);
4166
4167 var miss = (try selectSparseWithTuning(allocator, .{
4168 .dtype = .f32,
4169 .kind = .{ .csr_spmm = .{ .rows = 70, .columns = 46, .nnz = 512, .x_extent = 40 } },
4170 }, reader)) orelse return error.TestExpectedSparseDescriptor;
4171 defer miss.deinit();
4172 var fallback_instance = sparse_mod.SpmmCsr{
4173 .rows = 70,
4174 .columns = 46,
4175 .nnz = 512,
4176 .x_extent = 40,
4177 .threads = sparse_mod.spmmCsrThreadsForExtents(70, 46),
4178 };
4179 fallback_instance.threads = sparse_mod.spmmCsrRepresentableThreads(fallback_instance) orelse
4180 return error.TestExpectedSparseDescriptor;
4181 const fallback_target = try sparse_mod.spmmCsrFamilyTarget(allocator, fallback_instance);
4182 defer allocator.free(fallback_target);
4183 try std.testing.expectEqualStrings(fallback_target, miss.descriptor.metadata.target);
4184
4185 var explicit = (try selectSparseWithTuning(allocator, .{
4186 .dtype = .f32,
4187 .kind = .{ .csr_spmm = .{ .rows = 70, .columns = 45, .nnz = 512, .x_extent = 40 } },
4188 .schedule = .{ .thread_blocks_2d = .{ .x = 8, .y = 4 } },
4189 }, reader)) orelse return error.TestExpectedSparseDescriptor;
4190 defer explicit.deinit();
4191 try std.testing.expectEqualStrings(
4192 "accy.kernel.sparse.spmm_csr_row_column_thread_family_8x4_f32",
4193 explicit.descriptor.metadata.target,
4194 );
4195 }
4196
4197 test "catalog sparse candidates enumerate sparse structures and schedules" {
4198 const allocator = std.testing.allocator;
4199
4200 var csr_candidates = try selectSparseCandidates(allocator, .{
4201 .dtype = .f32,
4202 .kind = .{ .csr_spmv = .{ .rows = 70, .nnz = 560, .x_extent = 40 } },
4203 });
4204 defer csr_candidates.deinit();
4205 try std.testing.expectEqual(@as(usize, 2), csr_candidates.count);
4206 try std.testing.expectEqualStrings(
4207 "accy.kernel.sparse.spmv_csr_row_thread_family_70_f32",
4208 csr_candidates.items[0].descriptor.metadata.target,
4209 );
4210 try std.testing.expectEqualStrings(
4211 "accy.kernel.sparse.spmv_csr_row_warp_family_256_f32",
4212 csr_candidates.items[1].descriptor.metadata.target,
4213 );
4214
4215 var coo_candidates = try selectSparseCandidates(allocator, .{
4216 .dtype = .f32,
4217 .kind = .{ .coo_spmv = .{ .rows = 70, .nnz = 560, .x_extent = 40 } },
4218 });
4219 defer coo_candidates.deinit();
4220 try std.testing.expectEqual(@as(usize, 2), coo_candidates.count);
4221 try std.testing.expectEqualStrings(
4222 "accy.kernel.sparse.spmv_coo_element_thread_family_256_f32",
4223 coo_candidates.items[0].descriptor.metadata.target,
4224 );
4225 try std.testing.expectEqualStrings(
4226 "accy.kernel.sparse.spmv_coo_row_thread_family_70_f32",
4227 coo_candidates.items[1].descriptor.metadata.target,
4228 );
4229
4230 var coo_half_candidates = try selectSparseCandidates(allocator, .{
4231 .dtype = .f16,
4232 .kind = .{ .coo_spmv = .{ .rows = 70, .nnz = 560, .x_extent = 40 } },
4233 });
4234 defer coo_half_candidates.deinit();
4235 try std.testing.expectEqual(@as(usize, 1), coo_half_candidates.count);
4236 try std.testing.expectEqualStrings(
4237 "accy.kernel.sparse.spmv_coo_row_thread_family_70_f16",
4238 coo_half_candidates.items[0].descriptor.metadata.target,
4239 );
4240
4241 var ell_candidates = try selectSparseCandidates(allocator, .{
4242 .dtype = .f32,
4243 .kind = .{ .ell_spmv = .{ .rows = 70, .slots = 8, .x_extent = 40 } },
4244 });
4245 defer ell_candidates.deinit();
4246 try std.testing.expectEqual(@as(usize, 3), ell_candidates.count);
4247 try std.testing.expectEqualStrings(
4248 "accy.kernel.sparse.spmv_ell_row_thread_family_70_f32",
4249 ell_candidates.items[0].descriptor.metadata.target,
4250 );
4251 try std.testing.expectEqualStrings(
4252 "accy.kernel.sparse.spmv_ell_row_thread_family_32_f32",
4253 ell_candidates.items[1].descriptor.metadata.target,
4254 );
4255 try std.testing.expectEqualStrings(
4256 "accy.kernel.sparse.spmv_ell_row_thread_family_64_f32",
4257 ell_candidates.items[2].descriptor.metadata.target,
4258 );
4259
4260 var sell_candidates = try selectSparseCandidates(allocator, .{
4261 .dtype = .f32,
4262 .kind = .{ .sell_spmv = .{ .rows = 70, .slice_size = 8, .values_size = 400, .x_extent = 40 } },
4263 });
4264 defer sell_candidates.deinit();
4265 try std.testing.expectEqual(@as(usize, 3), sell_candidates.count);
4266 try std.testing.expectEqualStrings(
4267 "accy.kernel.sparse.spmv_sell_row_thread_slice8_family_70_f32",
4268 sell_candidates.items[0].descriptor.metadata.target,
4269 );
4270 try std.testing.expectEqualStrings(
4271 "accy.kernel.sparse.spmv_sell_row_thread_slice8_family_32_f32",
4272 sell_candidates.items[1].descriptor.metadata.target,
4273 );
4274 try std.testing.expectEqualStrings(
4275 "accy.kernel.sparse.spmv_sell_row_thread_slice8_family_64_f32",
4276 sell_candidates.items[2].descriptor.metadata.target,
4277 );
4278
4279 var spmm_candidates = try selectSparseCandidates(allocator, .{
4280 .dtype = .f32,
4281 .kind = .{ .csr_spmm = .{ .rows = 70, .columns = 45, .nnz = 560, .x_extent = 40 } },
4282 });
4283 defer spmm_candidates.deinit();
4284 try std.testing.expect(spmm_candidates.count > 1);
4285 var spmm_default = sparse_mod.SpmmCsr{
4286 .rows = 70,
4287 .columns = 45,
4288 .nnz = 560,
4289 .x_extent = 40,
4290 .threads = sparse_mod.spmmCsrThreadsForExtents(70, 45),
4291 };
4292 spmm_default.threads = sparse_mod.spmmCsrRepresentableThreads(spmm_default) orelse
4293 return error.TestExpectedSparseDescriptor;
4294 const spmm_default_target = try sparse_mod.spmmCsrFamilyTarget(allocator, spmm_default);
4295 defer allocator.free(spmm_default_target);
4296 try std.testing.expectEqualStrings(
4297 spmm_default_target,
4298 spmm_candidates.items[0].descriptor.metadata.target,
4299 );
4300
4301 var explicit = try selectSparseCandidates(allocator, .{
4302 .dtype = .f32,
4303 .kind = .{ .csr_spmv = .{ .rows = 70, .nnz = 560, .x_extent = 40 } },
4304 .structure = .row_warp,
4305 .schedule = .{ .thread_blocks = 64 },
4306 });
4307 defer explicit.deinit();
4308 try std.testing.expectEqual(@as(usize, 1), explicit.count);
4309 try std.testing.expectEqualStrings(
4310 "accy.kernel.sparse.spmv_csr_row_warp_family_64_f32",
4311 explicit.items[0].descriptor.metadata.target,
4312 );
4313 }
4314 };
4315
4316 const sort_tests = struct {
4317 test "kernel library catalog selects and dispatches radix sort pipelines" {
4318 const allocator = std.testing.allocator;
4319 var state = gpu.recording.BackendState{
4320 .allocator = allocator,
4321 .kind = .cuda,
4322 .format = .cuda_ptx,
4323 };
4324
4325 var owned = (try selectOwned(allocator, .{ .sort = .{
4326 .dtype = .i32,
4327 .kind = .radix_ascending,
4328 .extent = 2000,
4329 } })) orelse return error.TestExpectedCatalogDescriptor;
4330 defer owned.deinit();
4331 try std.testing.expectEqualStrings("accy.kernel.sort.radix_digit_family_32_i32", owned.descriptor.metadata.target);
4332 try std.testing.expectEqual(entry.Category.sort, owned.descriptor.metadata.category);
4333 try std.testing.expect(owned.descriptor.metadata.specialization.structureIs("radix_digit"));
4334
4335 var split = (try selectOwned(allocator, .{ .sort = .{
4336 .dtype = .i32,
4337 .kind = .radix_ascending,
4338 .extent = 2000,
4339 .structure = .radix_split,
4340 } })) orelse return error.TestExpectedCatalogDescriptor;
4341 defer split.deinit();
4342 try std.testing.expectEqualStrings("accy.kernel.sort.radix_split_family_32_i32", split.descriptor.metadata.target);
4343 try std.testing.expect(split.descriptor.metadata.specialization.structureIs("radix_split"));
4344
4345 var scheduled = (try selectOwned(allocator, .{ .sort = .{
4346 .dtype = .i32,
4347 .kind = .radix_ascending,
4348 .extent = 2000,
4349 .structure = .radix_split,
4350 .schedule = .{ .thread_blocks = 128 },
4351 } })) orelse return error.TestExpectedCatalogDescriptor;
4352 defer scheduled.deinit();
4353 try std.testing.expectEqualStrings("accy.kernel.sort.radix_split_family_128_i32", scheduled.descriptor.metadata.target);
4354
4355 var bitonic = (try selectOwned(allocator, .{ .sort = .{
4356 .dtype = .i32,
4357 .kind = .radix_ascending,
4358 .extent = 45,
4359 .structure = .bitonic_block,
4360 } })) orelse return error.TestExpectedCatalogDescriptor;
4361 defer bitonic.deinit();
4362 try std.testing.expectEqualStrings("accy.kernel.sort.bitonic_block_family_64_i32", bitonic.descriptor.metadata.target);
4363 try std.testing.expectEqual(sort_mod.bitonic_block_family_version, bitonic.descriptor.metadata.version);
4364 try std.testing.expect(bitonic.descriptor.metadata.specialization.structureIs("bitonic_block"));
4365
4366 try std.testing.expectEqual(
4367 @as(?OwnedDescriptor, null),
4368 try selectOwned(allocator, .{ .sort = .{
4369 .dtype = .i32,
4370 .kind = .radix_ascending,
4371 .extent = 1024 * 1024 + 1,
4372 } }),
4373 );
4374 try std.testing.expectEqual(
4375 @as(?OwnedDescriptor, null),
4376 try selectOwned(allocator, .{ .sort = .{ .dtype = .f32, .kind = .radix_ascending, .extent = 2000 } }),
4377 );
4378 try std.testing.expectEqual(
4379 @as(?OwnedDescriptor, null),
4380 try selectOwned(allocator, .{ .sort = .{
4381 .dtype = .i32,
4382 .kind = .radix_ascending,
4383 .extent = 1025,
4384 .structure = .bitonic_block,
4385 } }),
4386 );
4387
4388 var package = (try createOwnedKernelCallPipelinePackage(allocator, state.handle(), owned, .{ .limits = .testing })) orelse {
4389 return error.TestExpectedPipelinePackage;
4390 };
4391 defer package.deinit();
4392 try std.testing.expectEqual(@as(usize, 5), package.artifacts.len);
4393 try std.testing.expectEqualStrings("accy.kernel.sort.radix_digit_family_32_i32", package.pipeline.value.target);
4394 try std.testing.expect(package.registry().find(
4395 "accy.kernel.sort.radix_digit_histogram_family_32_i32",
4396 sort_mod.radix_split_family_version,
4397 .cuda_ptx,
4398 ) != null);
4399
4400 var split_package = (try createOwnedKernelCallPipelinePackage(allocator, state.handle(), split, .{ .limits = .testing })) orelse {
4401 return error.TestExpectedPipelinePackage;
4402 };
4403 defer split_package.deinit();
4404 try std.testing.expectEqual(@as(usize, 5), split_package.artifacts.len);
4405 try std.testing.expectEqualStrings("accy.kernel.sort.radix_split_family_32_i32", split_package.pipeline.value.target);
4406 try std.testing.expect(split_package.registry().find(
4407 "accy.kernel.sort.radix_split_flags_family_32_i32",
4408 sort_mod.radix_split_family_version,
4409 .cuda_ptx,
4410 ) != null);
4411
4412 var bitonic_artifact = try createOwnedKernelCallArtifact(allocator, state.handle(), bitonic, .{ .limits = .testing });
4413 defer bitonic_artifact.deinit();
4414 const bitonic_entry = bitonic_artifact.entry();
4415 try std.testing.expectEqualStrings("accy.kernel.sort.bitonic_block_family_64_i32", bitonic_entry.target);
4416 try std.testing.expectEqual(@as(u32, 1), bitonic_entry.runtime_scalar_argument_count);
4417 try std.testing.expect(try createOwnedKernelCallPipelinePackage(allocator, state.handle(), bitonic, .{ .limits = .testing }) == null);
4418 }
4419
4420 test "kernel library catalog consults tuning for sort structure" {
4421 const allocator = std.testing.allocator;
4422 const sort_lib = sort_mod;
4423 const device: u64 = 0xfeed_dead_beef_0002;
4424 const instance = sort_lib.RadixSplit{ .extent = 5000, .threads = 32 };
4425
4426 var accumulator = tuning.FamilyMeasurementAccumulator.init(allocator);
4427 defer accumulator.deinit();
4428 const key = try sort_lib.radixSplitFamilyTuningKey(allocator, device, instance);
4429 const split_target = try sort_lib.radixSplitPipelineTarget(allocator, instance);
4430 defer allocator.free(split_target);
4431 const digit_target = try sort_lib.radixDigitPipelineTarget(allocator, instance);
4432 defer allocator.free(digit_target);
4433 try accumulator.append(key, digit_target, 4_953_684, 20);
4434 try accumulator.append(key, split_target, 897_993, 20);
4435
4436 var winners = try accumulator.selectWinners(allocator, tuning.family_tuning_default_margin_percent);
4437 defer winners.deinit();
4438 const encoded = try tuning.encodeFamilyTuningArtifact(allocator, winners.records);
4439 defer allocator.free(encoded);
4440 var decoded = try tuning.decodeFamilyTuningArtifact(allocator, encoded);
4441 defer decoded.deinit();
4442 const reader = tuning.FamilyTuningReader{
4443 .device_fingerprint = device,
4444 .table = decoded.table(),
4445 };
4446
4447 var tuned = (try selectRadixSortWithTuning(allocator, .{
4448 .dtype = .i32,
4449 .kind = .radix_ascending,
4450 .extent = 5000,
4451 }, reader)) orelse return error.TestExpectedCatalogDescriptor;
4452 defer tuned.deinit();
4453 try std.testing.expectEqualStrings("accy.kernel.sort.radix_split_family_32_i32", tuned.descriptor.metadata.target);
4454 try std.testing.expect(tuned.descriptor.metadata.specialization.structureIs("radix_split"));
4455
4456 var missed = (try selectRadixSortWithTuning(allocator, .{
4457 .dtype = .i32,
4458 .kind = .radix_ascending,
4459 .extent = 9000,
4460 }, reader)) orelse return error.TestExpectedCatalogDescriptor;
4461 defer missed.deinit();
4462 try std.testing.expectEqualStrings("accy.kernel.sort.radix_digit_family_32_i32", missed.descriptor.metadata.target);
4463
4464 var pinned = (try selectRadixSortWithTuning(allocator, .{
4465 .dtype = .i32,
4466 .kind = .radix_ascending,
4467 .extent = 5000,
4468 .structure = .radix_digit,
4469 }, reader)) orelse return error.TestExpectedCatalogDescriptor;
4470 defer pinned.deinit();
4471 try std.testing.expectEqualStrings("accy.kernel.sort.radix_digit_family_32_i32", pinned.descriptor.metadata.target);
4472 }
4473
4474 test "kernel library catalog selects and dispatches block top-k families" {
4475 const allocator = std.testing.allocator;
4476 var state = gpu.recording.BackendState{
4477 .allocator = allocator,
4478 .kind = .cuda,
4479 .format = .cuda_ptx,
4480 };
4481
4482 var keys = (try selectOwned(allocator, .{ .sort = .{
4483 .dtype = .i32,
4484 .kind = .top_k_smallest,
4485 .extent = 45,
4486 .k = 8,
4487 } })) orelse return error.TestExpectedCatalogDescriptor;
4488 defer keys.deinit();
4489 try std.testing.expectEqualStrings("accy.kernel.sort.top_k_block_family_64x8_i32", keys.descriptor.metadata.target);
4490 try std.testing.expectEqual(entry.Category.sort, keys.descriptor.metadata.category);
4491 try std.testing.expect(keys.descriptor.metadata.specialization.operationIs(.{ .sort = .top_k_smallest }));
4492 try std.testing.expect(keys.descriptor.metadata.specialization.structureIs(sort_mod.top_k_block_structure_name));
4493 try std.testing.expect(keys.descriptor.metadata.specialization.inputHasExtents(0, &.{45}));
4494 try std.testing.expect(keys.descriptor.metadata.specialization.outputHasExtents(0, &.{8}));
4495
4496 var scheduled = (try selectOwned(allocator, .{ .sort = .{
4497 .dtype = .i32,
4498 .kind = .top_k_smallest,
4499 .extent = 45,
4500 .k = 8,
4501 .schedule = .{ .thread_blocks = 128 },
4502 } })) orelse return error.TestExpectedCatalogDescriptor;
4503 defer scheduled.deinit();
4504 try std.testing.expectEqualStrings("accy.kernel.sort.top_k_block_family_128x8_i32", scheduled.descriptor.metadata.target);
4505
4506 var pairs = (try selectOwned(allocator, .{ .sort = .{
4507 .dtype = .i32,
4508 .kind = .top_k_smallest,
4509 .extent = 45,
4510 .k = 8,
4511 .structure = .top_k_block_pairs,
4512 } })) orelse return error.TestExpectedCatalogDescriptor;
4513 defer pairs.deinit();
4514 try std.testing.expectEqualStrings("accy.kernel.sort.top_k_block_pairs_family_64x8_i32", pairs.descriptor.metadata.target);
4515 try std.testing.expectEqual(sort_mod.top_k_block_pairs_family_version, pairs.descriptor.metadata.version);
4516 try std.testing.expect(pairs.descriptor.metadata.specialization.structureIs(sort_mod.top_k_block_pairs_structure_name));
4517 try std.testing.expect(pairs.descriptor.metadata.specialization.inputHasExtents(0, &.{45}));
4518 try std.testing.expect(pairs.descriptor.metadata.specialization.inputHasExtents(1, &.{45}));
4519 try std.testing.expect(pairs.descriptor.metadata.specialization.outputHasExtents(0, &.{8}));
4520 try std.testing.expect(pairs.descriptor.metadata.specialization.outputHasExtents(1, &.{8}));
4521
4522 try std.testing.expectEqual(
4523 @as(?OwnedDescriptor, null),
4524 try selectOwned(allocator, .{ .sort = .{
4525 .dtype = .i32,
4526 .kind = .top_k_smallest,
4527 .extent = 45,
4528 .k = 0,
4529 } }),
4530 );
4531 try std.testing.expectEqual(
4532 @as(?OwnedDescriptor, null),
4533 try selectOwned(allocator, .{ .sort = .{
4534 .dtype = .i32,
4535 .kind = .top_k_smallest,
4536 .extent = 45,
4537 .k = 46,
4538 } }),
4539 );
4540 try std.testing.expectEqual(
4541 @as(?OwnedDescriptor, null),
4542 try selectOwned(allocator, .{ .sort = .{
4543 .dtype = .f32,
4544 .kind = .top_k_smallest,
4545 .extent = 45,
4546 .k = 8,
4547 } }),
4548 );
4549 try std.testing.expectEqual(
4550 @as(?OwnedDescriptor, null),
4551 try selectOwned(allocator, .{ .sort = .{
4552 .dtype = .i32,
4553 .kind = .top_k_smallest,
4554 .extent = 45,
4555 .k = 8,
4556 .structure = .radix_digit,
4557 } }),
4558 );
4559
4560 var key_artifact = try createOwnedKernelCallArtifact(allocator, state.handle(), keys, .{ .limits = .testing });
4561 defer key_artifact.deinit();
4562 const key_entry = key_artifact.entry();
4563 try std.testing.expectEqualStrings("accy.kernel.sort.top_k_block_family_64x8_i32", key_entry.target);
4564 try std.testing.expectEqual(sort_mod.top_k_block_family_version, key_entry.version);
4565 try std.testing.expectEqual(@as(u32, 1), key_entry.runtime_scalar_argument_count);
4566
4567 var pair_artifact = try createOwnedKernelCallArtifact(allocator, state.handle(), pairs, .{ .limits = .testing });
4568 defer pair_artifact.deinit();
4569 const pair_entry = pair_artifact.entry();
4570 try std.testing.expectEqualStrings("accy.kernel.sort.top_k_block_pairs_family_64x8_i32", pair_entry.target);
4571 try std.testing.expectEqual(sort_mod.top_k_block_pairs_family_version, pair_entry.version);
4572 try std.testing.expectEqual(@as(u32, 1), pair_entry.runtime_scalar_argument_count);
4573 try std.testing.expect(try createOwnedKernelCallPipelinePackage(allocator, state.handle(), pairs, .{ .limits = .testing }) == null);
4574 }
4575 };
4576
4577 const selection_batched_matrix_product_tests = struct {
4578 test "kernel library catalog selects batched matrix product family by canonical einsum roles" {
4579 const lhs_dims = [_]i64{ 2, 2, 4 };
4580 const rhs_dims = [_]i64{ 2, 4, 3 };
4581 const output_dims = [_]i64{ 2, 2, 3 };
4582 const query = Query{ .batched_matrix_product = .{
4583 .dtype = .f32,
4584 .lhs_indices = "bmk",
4585 .rhs_indices = "bkn",
4586 .output_indices = "bmn",
4587 .lhs_dims = &lhs_dims,
4588 .rhs_dims = &rhs_dims,
4589 .output_dims = &output_dims,
4590 } };
4591 try std.testing.expect(select(query) == null);
4592 var selected = (try selectOwned(std.testing.allocator, query)) orelse return error.TestExpectedBatchedMatrixProductFamilySelection;
4593 defer selected.deinit();
4594 const specialization = selected.descriptor.metadata.specialization;
4595
4596 try std.testing.expect(selected.specialization != null);
4597 try std.testing.expectEqualStrings("accy.kernel.linalg.batched_matmul_family_3x2x2_f32", selected.descriptor.metadata.target);
4598 try std.testing.expectEqualStrings("accy_kernel_linalg_batched_matmul_family_3x2x2_f32", selected.descriptor.name);
4599 try std.testing.expect(specialization.operationIs(.{ .linalg = .batched_matrix_product }));
4600 try std.testing.expect(specialization.shape_family != null);
4601 try std.testing.expectEqual(@as(u32, 3), specialization.launch.?.threadgroup[0]);
4602 try std.testing.expectEqual(@as(u32, 2), specialization.launch.?.threadgroup[1]);
4603 try std.testing.expectEqual(@as(u32, 2), specialization.launch.?.threadgroup[2]);
4604 }
4605
4606 test "kernel library catalog selects batched matrix product family schedule specialization" {
4607 const lhs_dims = [_]i64{ 2, 2, 4 };
4608 const rhs_dims = [_]i64{ 2, 4, 3 };
4609 const output_dims = [_]i64{ 2, 2, 3 };
4610 const query = Query{ .batched_matrix_product = .{
4611 .dtype = .f32,
4612 .lhs_indices = "bmk",
4613 .rhs_indices = "bkn",
4614 .output_indices = "bmn",
4615 .lhs_dims = &lhs_dims,
4616 .rhs_dims = &rhs_dims,
4617 .output_dims = &output_dims,
4618 .schedule = .{ .thread_blocks = .{ .x = 3, .y = 2, .z = 2 } },
4619 } };
4620 try std.testing.expect(select(query) == null);
4621 var selected = (try selectOwned(std.testing.allocator, query)) orelse return error.TestExpectedBatchedMatrixProductScheduleSelection;
4622 defer selected.deinit();
4623 const specialization = selected.descriptor.metadata.specialization;
4624
4625 try std.testing.expect(selected.specialization != null);
4626 try std.testing.expectEqualStrings("accy.kernel.linalg.batched_matmul_family_3x2x2_f32", selected.descriptor.metadata.target);
4627 try std.testing.expectEqual(@as(u32, 3), specialization.launch.?.threadgroup[0]);
4628 try std.testing.expectEqual(@as(u32, 2), specialization.launch.?.threadgroup[1]);
4629 try std.testing.expectEqual(@as(u32, 2), specialization.launch.?.threadgroup[2]);
4630 }
4631 };
4632
4633 const selection_matrix_product_tests = struct {
4634 test "kernel library catalog selects matrix product by canonical einsum roles" {
4635 const lhs_dims = [_]i64{ 4, 8 };
4636 const rhs_dims = [_]i64{ 8, 16 };
4637 const output_dims = [_]i64{ 4, 16 };
4638 const selected = select(.{ .matrix_product = .{
4639 .dtype = .f32,
4640 .lhs_indices = "ik",
4641 .rhs_indices = "kj",
4642 .output_indices = "ij",
4643 .lhs_dims = &lhs_dims,
4644 .rhs_dims = &rhs_dims,
4645 .output_dims = &output_dims,
4646 } }) orelse return error.TestExpectedMatrixProductSelection;
4647
4648 try std.testing.expectEqualStrings(linalg.MatrixProduct4x16x8F32.target, selected.metadata.target);
4649 try std.testing.expectEqualStrings(linalg.MatrixProduct4x16x8F32.name, selected.name);
4650 }
4651
4652 test "kernel library catalog selects matrix product schedule specialization" {
4653 const lhs_dims = [_]i64{ 4, 8 };
4654 const rhs_dims = [_]i64{ 8, 16 };
4655 const output_dims = [_]i64{ 4, 16 };
4656 const selected = select(.{ .matrix_product = .{
4657 .dtype = .f32,
4658 .lhs_indices = "ik",
4659 .rhs_indices = "kj",
4660 .output_indices = "ij",
4661 .lhs_dims = &lhs_dims,
4662 .rhs_dims = &rhs_dims,
4663 .output_dims = &output_dims,
4664 .schedule = .{ .thread_blocks = .{ .x = 4, .y = 2 } },
4665 } }) orelse return error.TestExpectedMatrixProductScheduleSelection;
4666
4667 try std.testing.expectEqualStrings(linalg.MatrixProduct4x16x8ThreadBlocks4x2F32.target, selected.metadata.target);
4668 try std.testing.expectEqualStrings(linalg.MatrixProduct4x16x8ThreadBlocks4x2F32.name, selected.name);
4669 try std.testing.expectEqual(@as(u32, 4), selected.metadata.specialization.launch.?.threadgroup[0]);
4670 try std.testing.expectEqual(@as(u32, 2), selected.metadata.specialization.launch.?.threadgroup[1]);
4671 }
4672
4673 test "kernel library catalog owned selection preserves fixed matrix product descriptors" {
4674 const lhs_dims = [_]i64{ 4, 8 };
4675 const rhs_dims = [_]i64{ 8, 16 };
4676 const output_dims = [_]i64{ 4, 16 };
4677 var selected = (try selectOwned(std.testing.allocator, .{ .matrix_product = .{
4678 .dtype = .f32,
4679 .lhs_indices = "ik",
4680 .rhs_indices = "kj",
4681 .output_indices = "ij",
4682 .lhs_dims = &lhs_dims,
4683 .rhs_dims = &rhs_dims,
4684 .output_dims = &output_dims,
4685 } })) orelse return error.TestExpectedMatrixProductSelection;
4686 defer selected.deinit();
4687
4688 try std.testing.expect(selected.specialization == null);
4689 try std.testing.expectEqualStrings(linalg.MatrixProduct4x16x8F32.target, selected.descriptor.metadata.target);
4690 try std.testing.expectEqualStrings(linalg.MatrixProduct4x16x8F32.name, selected.descriptor.name);
4691 }
4692
4693 test "kernel library catalog selects matrix product family for runtime extents" {
4694 const lhs_dims = [_]i64{ 5, 3 };
4695 const rhs_dims = [_]i64{ 3, 7 };
4696 const output_dims = [_]i64{ 5, 7 };
4697 const query = Query{ .matrix_product = .{
4698 .dtype = .f32,
4699 .lhs_indices = "mk",
4700 .rhs_indices = "kn",
4701 .output_indices = "mn",
4702 .lhs_dims = &lhs_dims,
4703 .rhs_dims = &rhs_dims,
4704 .output_dims = &output_dims,
4705 } };
4706
4707 try std.testing.expect(select(query) == null);
4708 var selected = (try selectOwned(std.testing.allocator, query)) orelse return error.TestExpectedMatrixProductFamilySelection;
4709 defer selected.deinit();
4710 const specialization = selected.descriptor.metadata.specialization;
4711
4712 try std.testing.expect(selected.specialization != null);
4713 try std.testing.expectEqualStrings("accy.kernel.linalg.matmul_family_7x5_f32", selected.descriptor.metadata.target);
4714 try std.testing.expectEqualStrings("accy_kernel_linalg_matmul_family_7x5_f32", selected.descriptor.name);
4715 try std.testing.expectEqual(linalg.matrix_product_family_version, selected.descriptor.metadata.version);
4716 try std.testing.expect(specialization.operationIs(.{ .linalg = .matrix_product }));
4717 try std.testing.expect(specialization.shape_family != null);
4718 try std.testing.expect(specialization.scheduleMatchesLaunch());
4719 try std.testing.expect(matrixProductDescriptorMatches(selected.descriptor, .{ .m = 5, .n = 7, .k = 3 }, null, .f32));
4720 try std.testing.expectEqual(@as(u32, 1), specialization.launch.?.grid[0]);
4721 try std.testing.expectEqual(@as(u32, 1), specialization.launch.?.grid[1]);
4722 try std.testing.expectEqual(@as(u32, 7), specialization.launch.?.threadgroup[0]);
4723 try std.testing.expectEqual(@as(u32, 5), specialization.launch.?.threadgroup[1]);
4724 }
4725
4726 test "kernel library catalog selects f16 matrix product family for runtime extents" {
4727 const lhs_dims = [_]i64{ 5, 3 };
4728 const rhs_dims = [_]i64{ 3, 7 };
4729 const output_dims = [_]i64{ 5, 7 };
4730 const query = Query{ .matrix_product = .{
4731 .dtype = .f16,
4732 .lhs_indices = "mk",
4733 .rhs_indices = "kn",
4734 .output_indices = "mn",
4735 .lhs_dims = &lhs_dims,
4736 .rhs_dims = &rhs_dims,
4737 .output_dims = &output_dims,
4738 } };
4739
4740 try std.testing.expect(select(query) == null);
4741 var selected = (try selectOwned(std.testing.allocator, query)) orelse return error.TestExpectedMatrixProductFamilySelection;
4742 defer selected.deinit();
4743 const specialization = selected.descriptor.metadata.specialization;
4744
4745 try std.testing.expect(selected.specialization != null);
4746 try std.testing.expectEqualStrings("accy.kernel.linalg.matmul_family_7x5_f16", selected.descriptor.metadata.target);
4747 try std.testing.expectEqualStrings("accy_kernel_linalg_matmul_family_7x5_f16", selected.descriptor.name);
4748 try std.testing.expectEqual(@as(?choir_abi.DType, .f16), specialization.dtype);
4749 try std.testing.expectEqual(@as(?choir_abi.DType, .f32), specialization.accumulation_dtype);
4750 try std.testing.expect(matrixProductDescriptorMatches(selected.descriptor, .{ .m = 5, .n = 7, .k = 3 }, null, .f16));
4751 }
4752
4753 fn expectMatrixProductCandidateTargetsContain(
4754 candidates: MatrixProductCandidateDescriptors,
4755 target: []const u8,
4756 ) !void {
4757 for (candidates.slice()) |candidate| {
4758 if (std.mem.eql(u8, candidate.descriptor.metadata.target, target)) return;
4759 }
4760 return error.TestExpectedMatrixProductCandidate;
4761 }
4762
4763 test "kernel library catalog selects owned matrix product family candidates" {
4764 const lhs_dims = [_]i64{ 17, 13 };
4765 const rhs_dims = [_]i64{ 13, 17 };
4766 const output_dims = [_]i64{ 17, 17 };
4767 var candidates = try selectOwnedMatrixProductCandidates(std.testing.allocator, .{
4768 .dtype = .f32,
4769 .lhs_indices = "mk",
4770 .rhs_indices = "kn",
4771 .output_indices = "mn",
4772 .lhs_dims = &lhs_dims,
4773 .rhs_dims = &rhs_dims,
4774 .output_dims = &output_dims,
4775 });
4776 defer candidates.deinit();
4777
4778 try std.testing.expect(candidates.count > 2);
4779 try std.testing.expectEqualStrings("accy.kernel.linalg.matmul_family_17x9_f32", candidates.items[0].descriptor.metadata.target);
4780 try expectMatrixProductCandidateTargetsContain(candidates, "accy.kernel.linalg.matmul_family_16x16_f32");
4781 for (candidates.slice(), 0..) |candidate, index| {
4782 try std.testing.expect(candidate.specialization != null);
4783 try std.testing.expect(candidate.descriptor.metadata.specialization.shape_family != null);
4784 try std.testing.expect(matrixProductDescriptorMatches(candidate.descriptor, .{ .m = 17, .n = 17, .k = 13 }, null, .f32));
4785 for (candidates.slice()[0..index]) |previous| {
4786 try std.testing.expect(!std.mem.eql(u8, previous.descriptor.metadata.target, candidate.descriptor.metadata.target));
4787 }
4788 }
4789 }
4790
4791 test "kernel library catalog builds matrix product scheduled artifact registry" {
4792 const allocator = std.testing.allocator;
4793 var state = gpu.recording.BackendState{
4794 .allocator = allocator,
4795 .kind = .cuda,
4796 .format = .cuda_ptx,
4797 };
4798 const lhs_dims = [_]i64{ 17, 13 };
4799 const rhs_dims = [_]i64{ 13, 17 };
4800 const output_dims = [_]i64{ 17, 17 };
4801 var registry_descriptors: [2]OwnedDescriptor = undefined;
4802 var descriptor_count: usize = 0;
4803 defer for (registry_descriptors[0..descriptor_count]) |*descriptor| descriptor.deinit();
4804 registry_descriptors[0] = (try selectOwned(allocator, .{ .matrix_product = .{
4805 .dtype = .f32,
4806 .lhs_indices = "mk",
4807 .rhs_indices = "kn",
4808 .output_indices = "mn",
4809 .lhs_dims = &lhs_dims,
4810 .rhs_dims = &rhs_dims,
4811 .output_dims = &output_dims,
4812 .schedule = .{ .thread_blocks = .{ .x = 17, .y = 9 } },
4813 } })) orelse return error.TestExpectedMatrixProductCandidate;
4814 descriptor_count = 1;
4815 registry_descriptors[1] = (try selectOwned(allocator, .{ .matrix_product = .{
4816 .dtype = .f32,
4817 .lhs_indices = "mk",
4818 .rhs_indices = "kn",
4819 .output_indices = "mn",
4820 .lhs_dims = &lhs_dims,
4821 .rhs_dims = &rhs_dims,
4822 .output_dims = &output_dims,
4823 .schedule = .{ .thread_blocks = .{ .x = 16, .y = 16 } },
4824 } })) orelse return error.TestExpectedMatrixProductCandidate;
4825 descriptor_count = 2;
4826
4827 var owned_registry = try createOwnedKernelCallArtifactRegistry(allocator, state.handle(), registry_descriptors[0..], .{ .limits = .testing });
4828 defer owned_registry.deinit();
4829 const registry = owned_registry.registry();
4830 const selected = registry.find("accy.kernel.linalg.matmul_family_17x9_f32", linalg.matrix_product_family_version, .cuda_ptx) orelse {
4831 return error.TestExpectedKernelCallArtifact;
4832 };
4833 const baseline = registry.find("accy.kernel.linalg.matmul_family_16x16_f32", linalg.matrix_product_family_version, .cuda_ptx) orelse {
4834 return error.TestExpectedKernelCallArtifact;
4835 };
4836 const launch = switch (selected.launch) {
4837 .derived => |derived| derived,
4838 .fixed => return error.TestExpectedDerivedLaunch,
4839 };
4840 const profile = selected.shape_profile orelse return error.TestExpectedShapeProfile;
4841
4842 try std.testing.expectEqual(@as(usize, 2), registry.entries.len);
4843 try std.testing.expectEqual(@as(u32, 6), selected.argument_count);
4844 try std.testing.expectEqual(@as(u32, 3), selected.runtime_scalar_argument_count);
4845 try std.testing.expectEqual(@as(u32, 17), launch.threadgroup[0]);
4846 try std.testing.expectEqual(@as(u32, 9), launch.threadgroup[1]);
4847 try std.testing.expectEqual(@as(usize, 3), profile.dimensions.len);
4848 try std.testing.expect(selected.required_dtypes.contains(.f32));
4849 try std.testing.expect(selected.required_dtypes.contains(.i32));
4850 try std.testing.expectEqual(@as(u32, 16), baseline.launch.derived.threadgroup[0]);
4851 try std.testing.expectEqual(@as(u32, 16), baseline.launch.derived.threadgroup[1]);
4852 }
4853 };
4854
4855 const selection_matrix_vector_tests = struct {
4856 test "kernel library catalog selects matrix vector product family for runtime extents" {
4857 const matrix_dims = [_]i64{ 5, 3 };
4858 const vector_dims = [_]i64{3};
4859 const output_dims = [_]i64{5};
4860 const query = Query{ .matrix_vector_product = .{
4861 .dtype = .f32,
4862 .matrix_indices = "mk",
4863 .vector_indices = "k",
4864 .output_indices = "m",
4865 .matrix_dims = &matrix_dims,
4866 .vector_dims = &vector_dims,
4867 .output_dims = &output_dims,
4868 } };
4869
4870 try std.testing.expect(select(query) == null);
4871 var selected = (try selectOwned(std.testing.allocator, query)) orelse return error.TestExpectedMatrixVectorProductFamilySelection;
4872 defer selected.deinit();
4873 const specialization = selected.descriptor.metadata.specialization;
4874
4875 try std.testing.expect(selected.specialization != null);
4876 try std.testing.expectEqualStrings("accy.kernel.linalg.matvec_family_5x_f32", selected.descriptor.metadata.target);
4877 try std.testing.expectEqualStrings("accy_kernel_linalg_matvec_family_5x_f32", selected.descriptor.name);
4878 try std.testing.expectEqual(linalg.matrix_vector_product_family_version, selected.descriptor.metadata.version);
4879 try std.testing.expect(specialization.operationIs(.{ .linalg = .matrix_vector_product }));
4880 try std.testing.expect(specialization.shape_family != null);
4881 try std.testing.expect(specialization.scheduleMatchesLaunch());
4882 try std.testing.expectEqual(@as(u32, 1), specialization.launch.?.grid[0]);
4883 try std.testing.expectEqual(@as(u32, 5), specialization.launch.?.threadgroup[0]);
4884 try std.testing.expect(specialization.inputHasExtents(0, &.{ 5, 3 }));
4885 try std.testing.expect(specialization.inputHasExtents(1, &.{3}));
4886 try std.testing.expect(specialization.outputHasExtents(0, &.{5}));
4887 try std.testing.expect(specialization.reductionMatches(0, .{ .name = "dot", .operator = .dot_product, .extents = &.{3} }));
4888 }
4889
4890 test "kernel library catalog builds matrix vector product scheduled artifact" {
4891 const allocator = std.testing.allocator;
4892 var state = gpu.recording.BackendState{
4893 .allocator = allocator,
4894 .kind = .cuda,
4895 .format = .cuda_ptx,
4896 };
4897 const matrix_dims = [_]i64{ 5, 3 };
4898 const vector_dims = [_]i64{3};
4899 const output_dims = [_]i64{5};
4900 var selected = (try selectOwned(allocator, .{ .matrix_vector_product = .{
4901 .dtype = .f32,
4902 .matrix_indices = "mk",
4903 .vector_indices = "k",
4904 .output_indices = "m",
4905 .matrix_dims = &matrix_dims,
4906 .vector_dims = &vector_dims,
4907 .output_dims = &output_dims,
4908 .schedule = .{ .thread_blocks = 4 },
4909 } })) orelse return error.TestExpectedMatrixVectorProductFamilySelection;
4910 defer selected.deinit();
4911
4912 var call_artifact = try createOwnedKernelCallArtifact(allocator, state.handle(), selected, .{ .limits = .testing });
4913 defer call_artifact.deinit();
4914
4915 const artifact = call_artifact.registry().find(
4916 "accy.kernel.linalg.matvec_family_4x_f32",
4917 linalg.matrix_vector_product_family_version,
4918 .cuda_ptx,
4919 ) orelse return error.TestExpectedKernelCallArtifact;
4920 const launch = switch (artifact.launch) {
4921 .derived => |derived| derived,
4922 .fixed => return error.TestExpectedDerivedLaunch,
4923 };
4924 const profile = artifact.shape_profile orelse return error.TestExpectedShapeProfile;
4925
4926 try std.testing.expectEqual(@as(u32, 5), artifact.argument_count);
4927 try std.testing.expectEqual(@as(u32, 2), artifact.runtime_scalar_argument_count);
4928 try std.testing.expectEqual(@as(u32, 4), launch.threadgroup[0]);
4929 try std.testing.expectEqual(@as(usize, 2), profile.dimensions.len);
4930 try std.testing.expect(artifact.required_dtypes.contains(.f32));
4931 try std.testing.expect(artifact.required_dtypes.contains(.i32));
4932 }
4933 };
4934
4935 test {
4936 _ = @import("artifact/test.zig");
4937 _ = @import("family/test.zig");
4938 _ = @import("match/test.zig");
4939 _ = @import("select/test.zig");
4940 @import("test_discovery").discover(registry_tests);
4941 @import("test_discovery").discover(artifact_tests);
4942 @import("test_discovery").discover(variants_tests);
4943 @import("test_discovery").discover(fused_tests);
4944 @import("test_discovery").discover(reject_tests);
4945 @import("test_discovery").discover(normalization_tests);
4946 @import("test_discovery").discover(invalid_tests);
4947 @import("test_discovery").discover(family_tests);
4948 @import("test_discovery").discover(random_tests);
4949 @import("test_discovery").discover(segment_tests);
4950 @import("test_discovery").discover(scan_tests);
4951 @import("test_discovery").discover(sparse_tests);
4952 @import("test_discovery").discover(sort_tests);
4953 @import("test_discovery").discover(selection_batched_matrix_product_tests);
4954 @import("test_discovery").discover(selection_matrix_product_tests);
4955 @import("test_discovery").discover(selection_matrix_vector_tests);
4956 @import("test_discovery").discover(catalog);
4957 }
4958
4959 test "accy kernel library catalog declaration coverage" {
4960 std.testing.refAllDecls(registry_tests);
4961 std.testing.refAllDecls(artifact_tests);
4962 std.testing.refAllDecls(variants_tests);
4963 std.testing.refAllDecls(fused_tests);
4964 std.testing.refAllDecls(reject_tests);
4965 std.testing.refAllDecls(normalization_tests);
4966 std.testing.refAllDecls(invalid_tests);
4967 std.testing.refAllDecls(family_tests);
4968 std.testing.refAllDecls(random_tests);
4969 std.testing.refAllDecls(segment_tests);
4970 std.testing.refAllDecls(scan_tests);
4971 std.testing.refAllDecls(sparse_tests);
4972 std.testing.refAllDecls(sort_tests);
4973 std.testing.refAllDecls(selection_batched_matrix_product_tests);
4974 std.testing.refAllDecls(selection_matrix_product_tests);
4975 std.testing.refAllDecls(selection_matrix_vector_tests);
4976 std.testing.refAllDecls(catalog);
4977 }