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 }