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 }