lib/accy/src/kernel/library/catalog/family/scatter/add.zig
daab053ee43316e1809a84551d573ddd1e5bf3d2
1 const std = @import("std");
2
3 const descriptor_mod = @import("../../root.zig");
4 const geometry = @import("../../../root.zig").geometry;
5 const indexing = @import("../../../root.zig").indexing;
6 const match_mod = @import("../../match/root.zig");
7 const query_mod = descriptor_mod;
8
9 const Descriptor = descriptor_mod.Descriptor;
10 const OwnedDescriptor = descriptor_mod.OwnedDescriptor;
11 const ScatterAddQuery = query_mod.ScatterAddQuery;
12
13 pub const ScatterAddCandidateDescriptors = struct {
14 count: usize = 0,
15 items: [2 * geometry.max_thread_candidates]OwnedDescriptor = undefined,
16
17 pub fn slice(self: *const ScatterAddCandidateDescriptors) []const OwnedDescriptor {
18 return self.items[0..self.count];
19 }
20
21 pub fn deinit(self: *ScatterAddCandidateDescriptors) void {
22 for (self.items[0..self.count]) |*descriptor| descriptor.deinit();
23 self.* = undefined;
24 }
25 };
26
27 pub fn selectScatterAdd(backing_allocator: std.mem.Allocator, query: ScatterAddQuery) !?OwnedDescriptor {
28 if (!canonicalScatterAdd(query)) return null;
29 const instance = scatterAddFamilyInstance(query) orelse return null;
30 return try scatterAddDescriptorForInstance(backing_allocator, instance, query);
31 }
32
33 pub fn selectScatterAddCandidates(
34 backing_allocator: std.mem.Allocator,
35 query: ScatterAddQuery,
36 ) !ScatterAddCandidateDescriptors {
37 var result = ScatterAddCandidateDescriptors{};
38 errdefer result.deinit();
39
40 if (!canonicalScatterAdd(query)) return result;
41 if (query.schedule != null) {
42 if (try selectScatterAdd(backing_allocator, query)) |descriptor| {
43 result.items[result.count] = descriptor;
44 result.count += 1;
45 }
46 return result;
47 }
48
49 const query_total = scatterAddQueryTotal(query) orelse return result;
50 const thread_candidates = indexing.scatterAddThreadCandidatesForTotal(query_total);
51 const variants = [_]indexing.ScatterAddVariant{ .direct, .shared_bins };
52 for (variants) |variant| {
53 for (thread_candidates.slice()) |threads| {
54 const instance = indexing.ScatterAdd{
55 .outer = query.outer,
56 .axis_size = query.axis_size,
57 .updates = query.updates,
58 .inner = query.inner,
59 .dtype = query.dtype,
60 .variant = variant,
61 .threads = threads,
62 };
63 if (!indexing.scatterAddInstanceValid(instance)) continue;
64 if (result.count >= result.items.len) break;
65 if (try scatterAddDescriptorForInstance(backing_allocator, instance, query)) |descriptor| {
66 result.items[result.count] = descriptor;
67 result.count += 1;
68 }
69 }
70 }
71 return result;
72 }
73
74 fn canonicalScatterAdd(query: ScatterAddQuery) bool {
75 if (!indexing.scatterAddDTypeSupported(query.dtype)) return false;
76 return query.outer != 0 and query.axis_size != 0 and query.updates != 0 and query.inner != 0;
77 }
78
79 fn scatterAddDescriptorForInstance(
80 backing_allocator: std.mem.Allocator,
81 instance: indexing.ScatterAdd,
82 query: ScatterAddQuery,
83 ) !?OwnedDescriptor {
84 var specialization = try indexing.scatterAddFamilySpecialization(backing_allocator, instance);
85 errdefer specialization.deinit();
86 const lifetime_allocator = specialization.allocator();
87 const descriptor = Descriptor{
88 .name = try indexing.scatterAddFamilyEntryName(lifetime_allocator, instance),
89 .metadata = .{
90 .target = try indexing.scatterAddFamilyTarget(lifetime_allocator, instance),
91 .version = indexing.scatter_add_family_version,
92 .layer = .logical,
93 .category = .indexing,
94 .specialization = specialization.value,
95 },
96 };
97 if (!match_mod.scatterAddDescriptorMatches(descriptor, query)) {
98 specialization.deinit();
99 return null;
100 }
101 return .{ .descriptor = descriptor, .specialization = specialization };
102 }
103
104 fn scatterAddFamilyInstance(query: ScatterAddQuery) ?indexing.ScatterAdd {
105 var instance = indexing.ScatterAdd{
106 .outer = query.outer,
107 .axis_size = query.axis_size,
108 .updates = query.updates,
109 .inner = query.inner,
110 .dtype = query.dtype,
111 .threads = indexing.scatterAddThreadsForTotal(scatterAddQueryTotal(query) orelse return null),
112 };
113 if (query.schedule) |requested| {
114 switch (requested) {
115 .thread_blocks => |threads| {
116 if (threads == 0) return null;
117 if (@as(u64, threads) > instance.total()) return null;
118 instance.threads = threads;
119 },
120 .shared_bins => |threads| {
121 if (threads == 0) return null;
122 if (@as(u64, threads) > instance.total()) return null;
123 instance.variant = .shared_bins;
124 instance.threads = threads;
125 },
126 }
127 }
128 if (!indexing.scatterAddInstanceValid(instance)) return null;
129 return instance;
130 }
131
132 fn scatterAddQueryTotal(query: ScatterAddQuery) ?u64 {
133 const outer_updates = std.math.mul(u64, query.outer, query.updates) catch return null;
134 return std.math.mul(u64, outer_updates, query.inner) catch return null;
135 }
136
137 test "catalog scatter add family selection answers canonical queries" {
138 const allocator = std.testing.allocator;
139
140 var selected = (try selectScatterAdd(allocator, .{
141 .dtype = .i32,
142 .axis_size = 256,
143 .updates = 1024,
144 })) orelse return error.TestExpectedScatterAddDescriptor;
145 defer selected.deinit();
146 try std.testing.expect(std.mem.startsWith(u8, selected.descriptor.metadata.target, "accy.kernel.indexing.scatter_add_family_"));
147 try std.testing.expect(selected.descriptor.metadata.specialization.operationIs(.{ .indexing = .scatter_add }));
148
149 var scheduled = (try selectScatterAdd(allocator, .{
150 .dtype = .i32,
151 .axis_size = 256,
152 .updates = 1024,
153 .schedule = .{ .thread_blocks = 64 },
154 })) orelse return error.TestExpectedScatterAddDescriptor;
155 defer scheduled.deinit();
156 try std.testing.expectEqualStrings("accy.kernel.indexing.scatter_add_family_64_i32", scheduled.descriptor.metadata.target);
157
158 var single = (try selectScatterAdd(allocator, .{
159 .dtype = .f32,
160 .axis_size = 256,
161 .updates = 1024,
162 })) orelse return error.TestExpectedScatterAddDescriptor;
163 defer single.deinit();
164 try std.testing.expect(std.mem.endsWith(u8, single.descriptor.metadata.target, "_f32"));
165
166 try std.testing.expectEqual(@as(?OwnedDescriptor, null), try selectScatterAdd(allocator, .{
167 .dtype = .f16,
168 .axis_size = 256,
169 .updates = 1024,
170 }));
171 try std.testing.expectEqual(@as(?OwnedDescriptor, null), try selectScatterAdd(allocator, .{
172 .dtype = .i32,
173 .axis_size = 0,
174 .updates = 1024,
175 }));
176 }
177
178 test "catalog scatter add family candidates enumerate thread geometries" {
179 const allocator = std.testing.allocator;
180
181 var candidates = try selectScatterAddCandidates(allocator, .{
182 .dtype = .i32,
183 .axis_size = 128,
184 .updates = 4096,
185 });
186 defer candidates.deinit();
187 try std.testing.expect(candidates.count >= 2);
188 for (candidates.slice()) |candidate| {
189 try std.testing.expect(candidate.descriptor.metadata.specialization.operationIs(.{ .indexing = .scatter_add }));
190 }
191
192 var pinned = try selectScatterAddCandidates(allocator, .{
193 .dtype = .i32,
194 .axis_size = 128,
195 .updates = 4096,
196 .schedule = .{ .thread_blocks = 128 },
197 });
198 defer pinned.deinit();
199 try std.testing.expectEqual(@as(usize, 1), pinned.count);
200 try std.testing.expectEqualStrings("accy.kernel.indexing.scatter_add_family_128_i32", pinned.slice()[0].descriptor.metadata.target);
201 }
202
203 test "catalog scatter add candidates enumerate both schedule structures" {
204 const allocator = std.testing.allocator;
205
206 var candidates = try selectScatterAddCandidates(allocator, .{
207 .dtype = .i32,
208 .axis_size = 128,
209 .updates = 4096,
210 });
211 defer candidates.deinit();
212
213 var direct_count: usize = 0;
214 var shared_count: usize = 0;
215 for (candidates.slice()) |candidate| {
216 if (candidate.descriptor.metadata.specialization.structureIs("shared_bins")) {
217 shared_count += 1;
218 } else {
219 direct_count += 1;
220 }
221 }
222 try std.testing.expect(direct_count >= 2);
223 try std.testing.expectEqual(direct_count, shared_count);
224
225 var capped = try selectScatterAddCandidates(allocator, .{
226 .dtype = .i32,
227 .axis_size = indexing.scatter_add_shared_bins_cap + 1,
228 .updates = 4096,
229 });
230 defer capped.deinit();
231 for (capped.slice()) |candidate| {
232 try std.testing.expect(!candidate.descriptor.metadata.specialization.structureIs("shared_bins"));
233 }
234
235 var pinned = (try selectScatterAdd(allocator, .{
236 .dtype = .i32,
237 .axis_size = 128,
238 .updates = 4096,
239 .schedule = .{ .shared_bins = 128 },
240 })) orelse return error.TestExpectedScatterAddDescriptor;
241 defer pinned.deinit();
242 try std.testing.expectEqualStrings(
243 "accy.kernel.indexing.scatter_add_family_shared128_128_i32",
244 pinned.descriptor.metadata.target,
245 );
246 }