lib/accy/src/kernel/logical/selection/fused.zig

daab053ee43316e1809a84551d573ddd1e5bf3d2

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