lib/accy/src/kernel/library/catalog/match/image.zig
daab053ee43316e1809a84551d573ddd1e5bf3d2
1 const std = @import("std");
2
3 const catalog = @import("../root.zig");
4 const library = @import("../../root.zig");
5
6 const common = @import("common.zig");
7
8 const entry = library.entry;
9 const image = library.image;
10
11 const Descriptor = catalog.Descriptor;
12 const ImageQuery = catalog.ImageQuery;
13 const ImageSchedule = catalog.ImageSchedule;
14
15 const selectableSpecialization = common.selectableSpecialization;
16 const specializationThreadgroup2DMatches = common.specializationThreadgroup2DMatches;
17
18 pub fn imageDescriptorMatches(descriptor: Descriptor, query: ImageQuery) bool {
19 const metadata = descriptor.metadata;
20 if (metadata.category != .image) return false;
21 const specialization = metadata.specialization;
22 if (!selectableSpecialization(specialization)) return false;
23 if (!specialization.scheduleMatchesLaunch()) return false;
24 if (specialization.dtype != query.dtype) return false;
25 if (specialization.accumulation_dtype != .f32) return false;
26 if (!imageScheduleMatches(specialization, query.schedule)) return false;
27 return switch (query.kind) {
28 .blur_pass => |facts| blurPassDescriptorMatches(specialization, query, facts.radius, facts.axis),
29 .resize_bilinear => |facts| resizeDescriptorMatches(specialization, query, facts.src_width, facts.src_height),
30 };
31 }
32
33 fn blurPassDescriptorMatches(
34 specialization: entry.Specialization,
35 query: ImageQuery,
36 radius: u32,
37 axis: image.Axis,
38 ) bool {
39 if (!specialization.operationIs(.{ .image = .blur_pass })) return false;
40 if (specialization.inputs.len != 2 or specialization.outputs.len != 1) return false;
41 if (specialization.reductions.len != 1 or specialization.static_parameters.len != 2) return false;
42 if (!image.blurPassInstanceValid(.{
43 .radius = radius,
44 .axis = axis,
45 .width = query.width,
46 .height = query.height,
47 })) return false;
48 const taps = 2 * @as(u64, radius) + 1;
49 return specialization.inputHasExtents(0, &.{ query.height, query.width }) and
50 specialization.inputHasExtents(1, &.{taps}) and
51 specialization.outputHasExtents(0, &.{ query.height, query.width }) and
52 specialization.reductionMatches(0, .{
53 .name = "blur_tap",
54 .operator = .weighted_sum,
55 .extents = &.{taps},
56 }) and
57 specialization.staticParameterMatches("radius", radius) and
58 specialization.staticParameterMatches("axis", image.blurPassAxisParameter(axis));
59 }
60
61 fn resizeDescriptorMatches(
62 specialization: entry.Specialization,
63 query: ImageQuery,
64 src_width: u64,
65 src_height: u64,
66 ) bool {
67 if (!specialization.operationIs(.{ .image = .resize_bilinear })) return false;
68 if (specialization.inputs.len != 1 or specialization.outputs.len != 1) return false;
69 if (specialization.reductions.len != 0 or specialization.static_parameters.len != 0) return false;
70 if (!image.resizeInstanceValid(.{
71 .dst_width = query.width,
72 .dst_height = query.height,
73 .src_width = src_width,
74 .src_height = src_height,
75 })) return false;
76 return specialization.inputHasExtents(0, &.{ src_height, src_width }) and
77 specialization.outputHasExtents(0, &.{ query.height, query.width });
78 }
79
80 fn imageScheduleMatches(specialization: entry.Specialization, schedule: ?ImageSchedule) bool {
81 const requested = schedule orelse return true;
82 return switch (requested) {
83 .thread_blocks => |threads| specializationThreadgroup2DMatches(specialization, threads),
84 };
85 }
86
87 test {
88 std.testing.refAllDecls(@This());
89 }