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 }