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 }