lib/accy/src/kernel/library/geometry.zig

daab053ee43316e1809a84551d573ddd1e5bf3d2

  1 const std = @import("std");
  2 
  3 const entry = @import("entry.zig");
  4 
  5 pub const max_thread_candidates: usize = 8;
  6 
  7 pub const Grid2D = struct {
  8     rows: u64,
  9     cols: u64,
 10 };
 11 
 12 pub const ThreadCaps = struct {
 13     budget: u32,
 14     x_max: u32,
 15     y_max: u32,
 16     baseline_x_floor: u32 = 16,
 17     baseline_occupancy_numerator: u64 = 3,
 18     baseline_occupancy_denominator: u64 = 4,
 19 };
 20 
 21 pub const ThreadCandidates = struct {
 22     count: usize = 0,
 23     items: [max_thread_candidates]entry.Threads2D = @as([max_thread_candidates]entry.Threads2D, @splat(.{})),
 24 
 25     pub fn slice(self: *const ThreadCandidates) []const entry.Threads2D {
 26         return self.items[0..self.count];
 27     }
 28 };
 29 
 30 pub fn threadsForGrid(grid: Grid2D, caps: ThreadCaps) entry.Threads2D {
 31     return threadCandidatesForGrid(grid, caps).items[0];
 32 }
 33 
 34 pub fn threadCandidatesForGrid(grid: Grid2D, caps: ThreadCaps) ThreadCandidates {
 35     const baseline = baselineThreadsForGrid(grid, caps);
 36     const ranked = rankedThreadCandidatesForGrid(grid, caps);
 37     const selected = if (threadOccupancyAtLeast(
 38         grid,
 39         baseline,
 40         caps.baseline_occupancy_numerator,
 41         caps.baseline_occupancy_denominator,
 42     ))
 43         baseline
 44     else if (ranked.count != 0)
 45         ranked.items[0]
 46     else
 47         baseline;
 48 
 49     var result = ThreadCandidates{};
 50     appendThreadCandidate(&result, selected);
 51     appendThreadCandidate(&result, baseline);
 52     for (ranked.slice()) |candidate| appendThreadCandidate(&result, candidate);
 53     return result;
 54 }
 55 
 56 pub fn threadCandidatesEqual(lhs: entry.Threads2D, rhs: entry.Threads2D) bool {
 57     return lhs.x == rhs.x and lhs.y == rhs.y;
 58 }
 59 
 60 fn baselineThreadsForGrid(grid: Grid2D, caps: ThreadCaps) entry.Threads2D {
 61     const ty: u32 = @intCast(@max(@as(u64, 1), @min(grid.rows, caps.y_max)));
 62     const budget = caps.budget / ty;
 63     const tx: u32 = @intCast(@max(@as(u64, 1), @min(grid.cols, @min(@max(@as(u64, budget), caps.baseline_x_floor), caps.x_max))));
 64     return .{ .x = tx, .y = ty };
 65 }
 66 
 67 fn rankedThreadCandidatesForGrid(grid: Grid2D, caps: ThreadCaps) ThreadCandidates {
 68     var candidates = ThreadCandidates{};
 69     var y: u32 = 1;
 70     while (y <= @min(@max(grid.rows, 1), caps.y_max)) : (y += 1) {
 71         var x: u32 = 1;
 72         while (x <= @min(@max(grid.cols, 1), caps.x_max)) : (x += 1) {
 73             if (x * y > caps.budget) continue;
 74             insertRankedThreadCandidate(&candidates, grid, .{ .x = x, .y = y });
 75         }
 76     }
 77     return candidates;
 78 }
 79 
 80 fn appendThreadCandidate(candidates: *ThreadCandidates, candidate: entry.Threads2D) void {
 81     for (candidates.slice()) |existing| {
 82         if (threadCandidatesEqual(existing, candidate)) return;
 83     }
 84     if (candidates.count >= max_thread_candidates) return;
 85     candidates.items[candidates.count] = candidate;
 86     candidates.count += 1;
 87 }
 88 
 89 fn insertRankedThreadCandidate(
 90     candidates: *ThreadCandidates,
 91     grid: Grid2D,
 92     candidate: entry.Threads2D,
 93 ) void {
 94     for (candidates.slice()) |existing| {
 95         if (threadCandidatesEqual(existing, candidate)) return;
 96     }
 97 
 98     var insert_index: usize = 0;
 99     while (insert_index < candidates.count and !threadCandidateBeats(grid, candidate, candidates.items[insert_index])) {
100         insert_index += 1;
101     }
102     if (insert_index >= max_thread_candidates) return;
103 
104     if (candidates.count < max_thread_candidates) candidates.count += 1;
105     var index = candidates.count - 1;
106     while (index > insert_index) : (index -= 1) {
107         candidates.items[index] = candidates.items[index - 1];
108     }
109     candidates.items[insert_index] = candidate;
110 }
111 
112 fn threadCandidateBeats(grid: Grid2D, candidate: entry.Threads2D, best: entry.Threads2D) bool {
113     const candidate_score = threadScore(grid, candidate);
114     const best_score = threadScore(grid, best);
115     if (candidate_score.blocks != best_score.blocks) return candidate_score.blocks < best_score.blocks;
116     if (candidate_score.waste != best_score.waste) return candidate_score.waste < best_score.waste;
117     if (candidate_score.area != best_score.area) return candidate_score.area > best_score.area;
118     if (candidate_score.balance != best_score.balance) return candidate_score.balance < best_score.balance;
119     if (candidate.x != best.x) return candidate.x > best.x;
120     return candidate.y > best.y;
121 }
122 
123 const ThreadScore = struct {
124     blocks: u64,
125     waste: u64,
126     area: u32,
127     balance: u32,
128 };
129 
130 fn threadScore(grid: Grid2D, threads: entry.Threads2D) ThreadScore {
131     const grid_x = ceilDivU64(grid.cols, threads.x);
132     const grid_y = ceilDivU64(grid.rows, threads.y);
133     const blocks = grid_x *| grid_y;
134     const launched = blocks *| threads.x *| threads.y;
135     const useful = grid.rows *| grid.cols;
136     return .{
137         .blocks = blocks,
138         .waste = launched -| useful,
139         .area = threads.x * threads.y,
140         .balance = if (threads.x > threads.y) threads.x - threads.y else threads.y - threads.x,
141     };
142 }
143 
144 fn threadOccupancyAtLeast(
145     grid: Grid2D,
146     threads: entry.Threads2D,
147     numerator: u64,
148     denominator: u64,
149 ) bool {
150     const grid_x = ceilDivU64(grid.cols, threads.x);
151     const grid_y = ceilDivU64(grid.rows, threads.y);
152     const launched = grid_x *| grid_y *| threads.x *| threads.y;
153     const useful = grid.rows *| grid.cols;
154     return useful *| denominator >= launched *| numerator;
155 }
156 
157 fn ceilDivU64(numerator: u64, denominator: u32) u64 {
158     return numerator / denominator + @as(u64, @intFromBool(numerator % denominator != 0));
159 }
160 
161 const test_caps = ThreadCaps{ .budget = 256, .x_max = 64, .y_max = 16 };
162 
163 fn testOccupancy(grid: Grid2D, threads: entry.Threads2D) f64 {
164     const grid_x = (grid.cols + threads.x - 1) / threads.x;
165     const grid_y = (grid.rows + threads.y - 1) / threads.y;
166     const launched = grid_x * grid_y * threads.x * threads.y;
167     return @as(f64, @floatFromInt(grid.rows * grid.cols)) / @as(f64, @floatFromInt(launched));
168 }
169 
170 fn expectThreadCandidatesLegal(candidates: ThreadCandidates, grid: Grid2D, caps: ThreadCaps) !void {
171     try std.testing.expect(candidates.count != 0);
172     for (candidates.slice(), 0..) |candidate, index| {
173         try std.testing.expect(candidate.x != 0);
174         try std.testing.expect(candidate.y != 0);
175         try std.testing.expect(candidate.x <= @min(@max(grid.cols, 1), caps.x_max));
176         try std.testing.expect(candidate.y <= @min(@max(grid.rows, 1), caps.y_max));
177         try std.testing.expect(candidate.x * candidate.y <= caps.budget);
178         for (candidates.slice()[0..index]) |previous| {
179             try std.testing.expect(!threadCandidatesEqual(previous, candidate));
180         }
181     }
182 }
183 
184 test "geometry thread selection keeps occupancy high across grid regimes" {
185     const skinny = threadsForGrid(.{ .rows = 1, .cols = 1000 }, test_caps);
186     try std.testing.expectEqual(@as(u32, 1), skinny.y);
187     try std.testing.expectEqual(@as(u32, 64), skinny.x);
188     try std.testing.expect(testOccupancy(.{ .rows = 1, .cols = 1000 }, skinny) >= 0.9);
189 
190     const tall = threadsForGrid(.{ .rows = 1000, .cols = 2 }, test_caps);
191     try std.testing.expectEqual(@as(u32, 16), tall.y);
192     try std.testing.expectEqual(@as(u32, 2), tall.x);
193     try std.testing.expect(testOccupancy(.{ .rows = 1000, .cols = 2 }, tall) >= 0.9);
194 
195     const tiny = threadsForGrid(.{ .rows = 5, .cols = 7 }, test_caps);
196     try std.testing.expectEqual(@as(u32, 5), tiny.y);
197     try std.testing.expectEqual(@as(u32, 7), tiny.x);
198     try std.testing.expect(testOccupancy(.{ .rows = 5, .cols = 7 }, tiny) == 1.0);
199 
200     const dense = threadsForGrid(.{ .rows = 1024, .cols = 1024 }, test_caps);
201     try std.testing.expectEqual(@as(u32, 16), dense.y);
202     try std.testing.expectEqual(@as(u32, 16), dense.x);
203     try std.testing.expect(testOccupancy(.{ .rows = 1024, .cols = 1024 }, dense) == 1.0);
204 
205     const near_tile = threadsForGrid(.{ .rows = 17, .cols = 17 }, test_caps);
206     try std.testing.expectEqual(@as(u32, 9), near_tile.y);
207     try std.testing.expectEqual(@as(u32, 17), near_tile.x);
208     try std.testing.expect(testOccupancy(.{ .rows = 17, .cols = 17 }, near_tile) >= 0.9);
209 }
210 
211 test "geometry thread candidates stay legal, unique, and lead with the selected default" {
212     const near_square = threadCandidatesForGrid(.{ .rows = 17, .cols = 17 }, test_caps);
213     try expectThreadCandidatesLegal(near_square, .{ .rows = 17, .cols = 17 }, test_caps);
214     try std.testing.expect(near_square.count > 2);
215     try std.testing.expect(threadCandidatesEqual(near_square.items[0], threadsForGrid(.{ .rows = 17, .cols = 17 }, test_caps)));
216 
217     const skinny = threadCandidatesForGrid(.{ .rows = 1, .cols = 1000 }, test_caps);
218     try expectThreadCandidatesLegal(skinny, .{ .rows = 1, .cols = 1000 }, test_caps);
219     try std.testing.expect(skinny.count > 1);
220     try std.testing.expect(threadCandidatesEqual(skinny.items[0], .{ .x = 64, .y = 1 }));
221 }
222 
223 test "geometry thread selection honors caller caps" {
224     const caps = ThreadCaps{ .budget = 64, .x_max = 8, .y_max = 8, .baseline_x_floor = 4 };
225 
226     const wide = threadsForGrid(.{ .rows = 3, .cols = 1000 }, caps);
227     try std.testing.expect(wide.x <= 8);
228     try std.testing.expect(wide.y <= 8);
229     try std.testing.expect(wide.x * wide.y <= 64);
230 
231     const candidates = threadCandidatesForGrid(.{ .rows = 33, .cols = 33 }, caps);
232     try expectThreadCandidatesLegal(candidates, .{ .rows = 33, .cols = 33 }, caps);
233     for (candidates.slice()) |candidate| {
234         try std.testing.expect(candidate.x * candidate.y <= caps.budget);
235         try std.testing.expect(candidate.x <= caps.x_max);
236         try std.testing.expect(candidate.y <= caps.y_max);
237     }
238 }
239 
240 pub const ThreadCaps1D = struct {
241     budget: u32 = 1024,
242     preferred: u32 = 256,
243     lane_quantum: u32 = 32,
244 };
245 
246 pub const Thread1DCandidates = struct {
247     count: usize = 0,
248     items: [max_thread_candidates]u32 = @as([max_thread_candidates]u32, @splat(0)),
249 
250     pub fn slice(self: *const Thread1DCandidates) []const u32 {
251         return self.items[0..self.count];
252     }
253 };
254 
255 pub fn threadsForExtent(extent: u64, caps: ThreadCaps1D) u32 {
256     return threadCandidatesForExtent(extent, caps).items[0];
257 }
258 
259 pub fn threadCandidatesForExtent(extent: u64, caps: ThreadCaps1D) Thread1DCandidates {
260     var result = Thread1DCandidates{};
261     const extent_cap: u64 = @max(extent, 1);
262     if (extent_cap < caps.lane_quantum) {
263         insertRanked1DCandidate(&result, extent, caps, @intCast(extent_cap));
264         return result;
265     }
266     var threads: u32 = caps.lane_quantum;
267     while (threads <= caps.budget and @as(u64, threads) <= extent_cap) : (threads *= 2) {
268         insertRanked1DCandidate(&result, extent, caps, threads);
269     }
270     if (result.count == 0) insertRanked1DCandidate(&result, extent, caps, caps.lane_quantum);
271     return result;
272 }
273 
274 fn insertRanked1DCandidate(candidates: *Thread1DCandidates, extent: u64, caps: ThreadCaps1D, candidate: u32) void {
275     for (candidates.slice()) |existing| {
276         if (existing == candidate) return;
277     }
278 
279     var insert_index: usize = 0;
280     while (insert_index < candidates.count and !thread1DCandidateBeats(extent, caps, candidate, candidates.items[insert_index])) {
281         insert_index += 1;
282     }
283     if (insert_index >= max_thread_candidates) return;
284 
285     if (candidates.count < max_thread_candidates) candidates.count += 1;
286     var index = candidates.count - 1;
287     while (index > insert_index) : (index -= 1) {
288         candidates.items[index] = candidates.items[index - 1];
289     }
290     candidates.items[insert_index] = candidate;
291 }
292 
293 fn thread1DCandidateBeats(extent: u64, caps: ThreadCaps1D, candidate: u32, best: u32) bool {
294     const candidate_waste = lane1DWaste(extent, caps, candidate);
295     const best_waste = lane1DWaste(extent, caps, best);
296     if (candidate_waste != best_waste) return candidate_waste < best_waste;
297     const candidate_distance = preferredDistance(candidate, caps.preferred);
298     const best_distance = preferredDistance(best, caps.preferred);
299     if (candidate_distance != best_distance) return candidate_distance < best_distance;
300     return candidate > best;
301 }
302 
303 fn lane1DWaste(extent: u64, caps: ThreadCaps1D, threads: u32) u64 {
304     const lanes = laneRoundUp(@as(u64, threads), caps.lane_quantum);
305     const blocks = ceilDivU64(@max(extent, 1), threads);
306     return blocks *| lanes -| @max(extent, 1);
307 }
308 
309 fn laneRoundUp(value: u64, quantum: u32) u64 {
310     return ceilDivU64(value, quantum) *| quantum;
311 }
312 
313 fn preferredDistance(threads: u32, preferred: u32) u32 {
314     return if (threads > preferred) threads - preferred else preferred - threads;
315 }
316 
317 const test_caps_1d = ThreadCaps1D{};
318 
319 test "geometry 1D thread candidates prefer warp-aligned high-occupancy blocks" {
320     const huge = threadCandidatesForExtent(1 << 20, test_caps_1d);
321     try std.testing.expect(huge.count >= 5);
322     try std.testing.expectEqual(@as(u32, 256), huge.items[0]);
323     for (huge.slice(), 0..) |candidate, index| {
324         try std.testing.expect(candidate >= 32 and candidate <= 1024);
325         for (huge.slice()[0..index]) |previous| try std.testing.expect(previous != candidate);
326     }
327 
328     const tiny = threadCandidatesForExtent(8, test_caps_1d);
329     try std.testing.expectEqual(@as(u32, 8), tiny.items[0]);
330     try std.testing.expectEqual(@as(usize, 1), tiny.count);
331 
332     const uneven = threadCandidatesForExtent(1000, test_caps_1d);
333     try std.testing.expectEqual(@as(u32, 256), uneven.items[0]);
334     try std.testing.expectEqual(@as(u32, 256), threadsForExtent(1000, test_caps_1d));
335 }