lib/accy/src/kernel/library/catalog/match/stencil.zig

daab053ee43316e1809a84551d573ddd1e5bf3d2

 1 const catalog = @import("../root.zig");
 2 const library = @import("../../root.zig");
 3 
 4 const common = @import("common.zig");
 5 
 6 const entry = library.entry;
 7 const stencil = library.stencil;
 8 
 9 const Descriptor = catalog.Descriptor;
10 const StencilQuery = catalog.StencilQuery;
11 const StencilSchedule = catalog.StencilSchedule;
12 
13 const selectableSpecialization = common.selectableSpecialization;
14 const specializationThreadgroup2DMatches = common.specializationThreadgroup2DMatches;
15 
16 pub fn stencilWindowDescriptorMatches(descriptor: Descriptor, query: StencilQuery) bool {
17     const metadata = descriptor.metadata;
18     if (metadata.category != .stencil) return false;
19     const specialization = metadata.specialization;
20     if (!selectableSpecialization(specialization)) return false;
21     if (!specialization.scheduleMatchesLaunch()) return false;
22     if (!specialization.operationIs(.{ .stencil = .window })) return false;
23     if (specialization.dtype != query.dtype) return false;
24     if (specialization.accumulation_dtype != stencil.windowAccumulationDType(query.dtype)) return false;
25     if (specialization.inputs.len != 2 or specialization.outputs.len != 1) return false;
26     if (specialization.reductions.len != 1) return false;
27     if (!stencil.windowRadiusValid(query.radius)) return false;
28     const side = 2 * @as(u64, query.radius) + 1;
29     const taps = side * side;
30     const padded_rows = query.rows + 2 * @as(u64, query.radius);
31     const padded_cols = query.cols + 2 * @as(u64, query.radius);
32     return specialization.inputHasExtents(0, &.{ padded_rows, padded_cols }) and
33         specialization.inputHasExtents(1, &.{taps}) and
34         specialization.outputHasExtents(0, &.{ query.rows, query.cols }) and
35         specialization.reductionMatches(0, .{
36             .name = "window",
37             .operator = .weighted_sum,
38             .extents = &.{taps},
39         }) and stencilWindowScheduleMatches(specialization, query.schedule);
40 }
41 fn stencilWindowScheduleMatches(specialization: entry.Specialization, schedule: ?StencilSchedule) bool {
42     const requested = schedule orelse return true;
43     return switch (requested) {
44         .thread_blocks => |threads| specializationThreadgroup2DMatches(specialization, threads),
45     };
46 }