lib/choir/src/backends/regalloc/locations.zig
daab053ee43316e1809a84551d573ddd1e5bf3d2
1 const std = @import("std");
2 const ir = @import("../../core/root.zig");
3 const position = @import("position.zig");
4 const range = @import("range.zig");
5
6 pub fn ValueLocationIndex(comptime Register: type) type {
7 const Range = range.ValueLocationRange(Register);
8 const Link = struct {
9 range_index: usize,
10 next: ?usize,
11 };
12 const no_link = std.math.maxInt(usize);
13 const Chain = struct {
14 first: usize = no_link,
15 last: usize = no_link,
16 };
17 const no_interval = std.math.maxInt(usize);
18 const IntervalNode = struct {
19 range_index: usize,
20 start: u64,
21 end: u64,
22 max_end: u64,
23 left: usize = no_interval,
24 right: usize = no_interval,
25 };
26
27 return struct {
28 value_keys: std.ArrayListUnmanaged(usize) = .empty,
29 value_heads: std.ArrayListUnmanaged(Chain) = .empty,
30 value_links: std.ArrayListUnmanaged(Link) = .empty,
31 register_roots: std.AutoHashMapUnmanaged(Register, usize) = .empty,
32 register_intervals: std.ArrayListUnmanaged(IntervalNode) = .empty,
33 start_heads: std.ArrayListUnmanaged(Chain) = .empty,
34 start_links: std.ArrayListUnmanaged(Link) = .empty,
35 end_heads: std.ArrayListUnmanaged(Chain) = .empty,
36 end_links: std.ArrayListUnmanaged(Link) = .empty,
37
38 const Self = @This();
39
40 pub const Iterator = struct {
41 links: []const Link,
42 link_index: ?usize,
43
44 pub fn next(self: *Iterator) ?usize {
45 const index = self.link_index orelse return null;
46 const link = self.links[index];
47 self.link_index = link.next;
48 return link.range_index;
49 }
50 };
51
52 pub fn deinit(self: *Self, allocator: std.mem.Allocator) void {
53 self.value_keys.deinit(allocator);
54 self.value_heads.deinit(allocator);
55 self.value_links.deinit(allocator);
56 self.register_roots.deinit(allocator);
57 self.register_intervals.deinit(allocator);
58 self.start_heads.deinit(allocator);
59 self.start_links.deinit(allocator);
60 self.end_heads.deinit(allocator);
61 self.end_links.deinit(allocator);
62 }
63
64 pub fn clearRetainingCapacity(self: *Self) void {
65 self.value_keys.clearRetainingCapacity();
66 self.value_heads.clearRetainingCapacity();
67 self.value_links.clearRetainingCapacity();
68 self.register_roots.clearRetainingCapacity();
69 self.register_intervals.clearRetainingCapacity();
70 self.start_heads.clearRetainingCapacity();
71 self.start_links.clearRetainingCapacity();
72 self.end_heads.clearRetainingCapacity();
73 self.end_links.clearRetainingCapacity();
74 }
75
76 pub fn prepare(
77 self: *Self,
78 allocator: std.mem.Allocator,
79 value_capacity: usize,
80 position_count: usize,
81 range_capacity: usize,
82 register_capacity: usize,
83 ) !void {
84 self.clearRetainingCapacity();
85 try self.value_keys.ensureTotalCapacityPrecise(allocator, value_capacity);
86 try self.value_heads.ensureTotalCapacityPrecise(allocator, value_capacity);
87 try resizeHeads(&self.start_heads, allocator, position_count);
88 try resizeHeads(&self.end_heads, allocator, position_count);
89 try self.value_links.ensureTotalCapacityPrecise(allocator, range_capacity);
90 try self.register_roots.ensureTotalCapacity(allocator, @intCast(register_capacity));
91 try self.register_intervals.ensureTotalCapacityPrecise(allocator, range_capacity);
92 try self.start_links.ensureTotalCapacityPrecise(allocator, range_capacity);
93 try self.end_links.ensureTotalCapacityPrecise(allocator, range_capacity);
94 }
95
96 pub fn includeValue(self: *Self, value: *ir.Value) void {
97 std.debug.assert(self.value_keys.items.len < self.value_keys.capacity);
98 self.value_keys.appendAssumeCapacity(valueKey(value));
99 }
100
101 pub fn sealValues(self: *Self, allocator: std.mem.Allocator) !void {
102 std.mem.sort(usize, self.value_keys.items, {}, std.sort.asc(usize));
103 if (self.value_keys.items.len != 0) {
104 var count: usize = 1;
105 for (self.value_keys.items[1..]) |key| {
106 if (key == self.value_keys.items[count - 1]) continue;
107 self.value_keys.items[count] = key;
108 count += 1;
109 }
110 self.value_keys.items.len = count;
111 }
112 try resizeHeads(&self.value_heads, allocator, self.value_keys.items.len);
113 }
114
115 pub fn rebuild(self: *Self, allocator: std.mem.Allocator, ranges: []const Range) !void {
116 var position_count: usize = 0;
117 for (ranges) |entry| {
118 position_count = @max(position_count, @as(usize, @max(entry.start, entry.end)) + 1);
119 }
120 try self.prepare(allocator, ranges.len, position_count, ranges.len, ranges.len);
121 for (ranges) |entry| self.includeValue(entry.value);
122 try self.sealValues(allocator);
123 for (ranges, 0..) |entry, index| self.append(index, entry);
124 }
125
126 pub fn append(self: *Self, range_index: usize, entry: Range) void {
127 const value_index = self.valueIndex(entry.value) orelse unreachable;
128 appendDenseLink(&self.value_heads, &self.value_links, value_index, range_index);
129 self.appendRegisterInterval(range_index, entry);
130 appendDenseLink(&self.start_heads, &self.start_links, entry.start, range_index);
131 appendDenseLink(&self.end_heads, &self.end_links, entry.end, range_index);
132 }
133
134 pub fn valueLocationAtPoint(self: Self, ranges: []const Range, value: *ir.Value, point: position.Point) ?Register {
135 const location = self.valueRangeAtPoint(ranges, value, point) orelse return null;
136 return location.reg;
137 }
138
139 pub fn valueRangeAtPoint(self: Self, ranges: []const Range, value: *ir.Value, point: position.Point) ?Range {
140 var iterator = self.valueIterator(value);
141 while (iterator.next()) |index| {
142 const entry = ranges[index];
143 if (entry.coversValueAt(value, point)) return entry;
144 }
145 return null;
146 }
147
148 pub fn valueHasRangeAtPoint(self: Self, ranges: []const Range, value: *ir.Value, point: position.Point) bool {
149 return self.valueLocationAtPoint(ranges, value, point) != null;
150 }
151
152 pub fn valueHasLocationRange(self: Self, value: *ir.Value) bool {
153 const value_index = self.valueIndex(value) orelse return false;
154 return headAt(self.value_heads.items, value_index) != null;
155 }
156
157 pub fn valueHasEntry(self: Self, ranges: []const Range, value: *ir.Value, entry: range.Entry) bool {
158 var iterator = self.valueIterator(value);
159 while (iterator.next()) |index| {
160 if (ranges[index].entry == entry) return true;
161 }
162 return false;
163 }
164
165 pub fn valueLocationStart(self: Self, ranges: []const Range, value: *ir.Value) ?u32 {
166 var iterator = self.valueIterator(value);
167 const index = iterator.next() orelse return null;
168 return ranges[index].start;
169 }
170
171 pub fn registerBlocksAt(self: Self, reg: Register, point: position.Point) bool {
172 const root = self.register_roots.get(reg) orelse no_interval;
173 return self.firstRegisterOverlap(root, point.rank(), point.rank()) != null;
174 }
175
176 pub fn registerRangeFirstOverlapStart(
177 self: Self,
178 ranges: []const Range,
179 value: *ir.Value,
180 start: u32,
181 end: u32,
182 reg: Register,
183 entry: range.Entry,
184 end_phase: position.Phase,
185 ) ?position.Point {
186 if (start >= end) return null;
187 const proposed = Range{
188 .value = value,
189 .start = start,
190 .end = end,
191 .reg = reg,
192 .entry = entry,
193 .end_phase = end_phase,
194 };
195 const proposed_start = proposed.startPoint().rank();
196 const proposed_end = proposed.endPoint().rank();
197 if (proposed_start > proposed_end) return null;
198 const range_index = self.firstRegisterOverlap(
199 self.register_roots.get(reg) orelse no_interval,
200 proposed_start,
201 proposed_end,
202 ) orelse return null;
203 const existing = ranges[range_index];
204 return if (existing.startPoint().lessThan(proposed.startPoint())) proposed.startPoint() else existing.startPoint();
205 }
206
207 pub fn rangesStartingAt(self: Self, point: u32) Iterator {
208 return eventIterator(headAt(self.start_heads.items, point), self.start_links.items);
209 }
210
211 pub fn rangesEndingAt(self: Self, point: u32) Iterator {
212 return eventIterator(headAt(self.end_heads.items, point), self.end_links.items);
213 }
214
215 fn valueIterator(self: Self, value: *ir.Value) Iterator {
216 const value_index = self.valueIndex(value) orelse return eventIterator(null, self.value_links.items);
217 return eventIterator(headAt(self.value_heads.items, value_index), self.value_links.items);
218 }
219
220 fn eventIterator(chain: ?Chain, links: []const Link) Iterator {
221 return .{ .links = links, .link_index = if (chain) |entry| entry.first else null };
222 }
223
224 fn appendRegisterInterval(self: *Self, range_index: usize, entry: Range) void {
225 const start = entry.startPoint().rank();
226 const end = entry.endPoint().rank();
227 const node_index = self.register_intervals.items.len;
228 std.debug.assert(node_index < self.register_intervals.capacity);
229 self.register_intervals.appendAssumeCapacity(.{
230 .range_index = range_index,
231 .start = start,
232 .end = end,
233 .max_end = end,
234 });
235 const root = self.insertRegisterInterval(self.register_roots.get(entry.reg) orelse no_interval, node_index);
236 self.register_roots.putAssumeCapacity(entry.reg, root);
237 }
238
239 fn insertRegisterInterval(self: *Self, root_index: usize, node_index: usize) usize {
240 if (root_index == no_interval) return node_index;
241 if (intervalBefore(self.register_intervals.items[node_index], self.register_intervals.items[root_index])) {
242 const child = self.insertRegisterInterval(self.register_intervals.items[root_index].left, node_index);
243 self.register_intervals.items[root_index].left = child;
244 if (intervalPriority(self.register_intervals.items[child].range_index) < intervalPriority(self.register_intervals.items[root_index].range_index)) {
245 return self.rotateRegisterRight(root_index);
246 }
247 } else {
248 const child = self.insertRegisterInterval(self.register_intervals.items[root_index].right, node_index);
249 self.register_intervals.items[root_index].right = child;
250 if (intervalPriority(self.register_intervals.items[child].range_index) < intervalPriority(self.register_intervals.items[root_index].range_index)) {
251 return self.rotateRegisterLeft(root_index);
252 }
253 }
254 self.updateRegisterMaxEnd(root_index);
255 return root_index;
256 }
257
258 fn rotateRegisterRight(self: *Self, root_index: usize) usize {
259 const pivot_index = self.register_intervals.items[root_index].left;
260 self.register_intervals.items[root_index].left = self.register_intervals.items[pivot_index].right;
261 self.register_intervals.items[pivot_index].right = root_index;
262 self.updateRegisterMaxEnd(root_index);
263 self.updateRegisterMaxEnd(pivot_index);
264 return pivot_index;
265 }
266
267 fn rotateRegisterLeft(self: *Self, root_index: usize) usize {
268 const pivot_index = self.register_intervals.items[root_index].right;
269 self.register_intervals.items[root_index].right = self.register_intervals.items[pivot_index].left;
270 self.register_intervals.items[pivot_index].left = root_index;
271 self.updateRegisterMaxEnd(root_index);
272 self.updateRegisterMaxEnd(pivot_index);
273 return pivot_index;
274 }
275
276 fn updateRegisterMaxEnd(self: *Self, node_index: usize) void {
277 const node = &self.register_intervals.items[node_index];
278 node.max_end = node.end;
279 if (node.left != no_interval) node.max_end = @max(node.max_end, self.register_intervals.items[node.left].max_end);
280 if (node.right != no_interval) node.max_end = @max(node.max_end, self.register_intervals.items[node.right].max_end);
281 }
282
283 fn firstRegisterOverlap(self: Self, node_index: usize, start: u64, end: u64) ?usize {
284 if (node_index == no_interval) return null;
285 const node = self.register_intervals.items[node_index];
286 if (node.left != no_interval and self.register_intervals.items[node.left].max_end >= start) {
287 if (self.firstRegisterOverlap(node.left, start, end)) |range_index| return range_index;
288 }
289 if (node.start <= end and node.end >= start) return node.range_index;
290 if (node.start > end) return null;
291 return self.firstRegisterOverlap(node.right, start, end);
292 }
293
294 fn intervalBefore(lhs: IntervalNode, rhs: IntervalNode) bool {
295 if (lhs.start != rhs.start) return lhs.start < rhs.start;
296 return lhs.range_index < rhs.range_index;
297 }
298
299 fn intervalPriority(range_index: usize) u64 {
300 var value = @as(u64, @intCast(range_index)) +% 0x9e3779b97f4a7c15;
301 value = (value ^ (value >> 30)) *% 0xbf58476d1ce4e5b9;
302 value = (value ^ (value >> 27)) *% 0x94d049bb133111eb;
303 return value ^ (value >> 31);
304 }
305
306 fn appendDenseLink(
307 heads: *std.ArrayListUnmanaged(Chain),
308 links: *std.ArrayListUnmanaged(Link),
309 key: usize,
310 range_index: usize,
311 ) void {
312 std.debug.assert(key < heads.items.len);
313 const link_index = links.items.len;
314 std.debug.assert(link_index < links.capacity);
315 links.appendAssumeCapacity(.{ .range_index = range_index, .next = null });
316 const head = &heads.items[key];
317 if (head.first != no_link) {
318 links.items[head.last].next = link_index;
319 head.last = link_index;
320 } else {
321 head.* = .{ .first = link_index, .last = link_index };
322 }
323 }
324
325 fn resizeHeads(
326 heads: *std.ArrayListUnmanaged(Chain),
327 allocator: std.mem.Allocator,
328 count: usize,
329 ) !void {
330 const old_len = heads.items.len;
331 if (count <= old_len) return;
332 try heads.ensureTotalCapacityPrecise(allocator, count);
333 try heads.resize(allocator, count);
334 @memset(heads.items[old_len..], .{});
335 }
336
337 fn headAt(heads: []const Chain, index: usize) ?Chain {
338 if (index >= heads.len) return null;
339 return if (heads[index].first == no_link) null else heads[index];
340 }
341
342 fn valueIndex(self: Self, value: *ir.Value) ?usize {
343 return std.sort.binarySearch(usize, self.value_keys.items, valueKey(value), valueKeyOrder);
344 }
345
346 fn valueKey(value: *ir.Value) usize {
347 return @intFromPtr(value);
348 }
349
350 fn valueKeyOrder(key: usize, candidate: usize) std.math.Order {
351 return std.math.order(key, candidate);
352 }
353 };
354 }
355
356 test "value location index tracks values registers and events" {
357 var owner: u8 = 0;
358 var value = ir.Value{
359 .kind = .{ .block_argument = .{ .owner = &owner, .arg_number = 0 } },
360 .type = undefined,
361 .id = 0,
362 };
363 var other = ir.Value{
364 .kind = .{ .block_argument = .{ .owner = &owner, .arg_number = 1 } },
365 .type = undefined,
366 .id = 0,
367 };
368 const Range = range.ValueLocationRange(u8);
369 var ranges = std.ArrayListUnmanaged(Range).empty;
370 defer ranges.deinit(std.testing.allocator);
371 var index = ValueLocationIndex(u8){};
372 defer index.deinit(std.testing.allocator);
373
374 try index.prepare(std.testing.allocator, 2, 7, 3, 2);
375 index.includeValue(&value);
376 index.includeValue(&other);
377 try index.sealValues(std.testing.allocator);
378 try std.testing.expectEqual(@as(usize, 2), index.value_keys.items.len);
379 try std.testing.expect(index.value_keys.items[0] != index.value_keys.items[1]);
380 try std.testing.expectEqual(@as(usize, 2), index.value_heads.items.len);
381 const capacities = .{
382 index.value_keys.capacity,
383 index.value_heads.capacity,
384 index.value_links.capacity,
385 index.register_roots.capacity(),
386 index.register_intervals.capacity,
387 index.start_heads.capacity,
388 index.start_links.capacity,
389 index.end_heads.capacity,
390 index.end_links.capacity,
391 };
392
393 try ranges.append(std.testing.allocator, .{ .value = &value, .start = 1, .end = 4, .reg = 2 });
394 index.append(0, ranges.items[0]);
395 try ranges.append(std.testing.allocator, .{ .value = &other, .start = 3, .end = 6, .reg = 2 });
396 index.append(1, ranges.items[1]);
397 try ranges.append(std.testing.allocator, .{ .value = &value, .start = 2, .end = 3, .reg = 4 });
398 index.append(2, ranges.items[2]);
399
400 try std.testing.expectEqual(capacities[0], index.value_keys.capacity);
401 try std.testing.expectEqual(capacities[1], index.value_heads.capacity);
402 try std.testing.expectEqual(capacities[2], index.value_links.capacity);
403 try std.testing.expectEqual(capacities[3], index.register_roots.capacity());
404 try std.testing.expectEqual(capacities[4], index.register_intervals.capacity);
405 try std.testing.expectEqual(capacities[5], index.start_heads.capacity);
406 try std.testing.expectEqual(capacities[6], index.start_links.capacity);
407 try std.testing.expectEqual(capacities[7], index.end_heads.capacity);
408 try std.testing.expectEqual(capacities[8], index.end_links.capacity);
409
410 try std.testing.expectEqual(@as(?u8, 2), index.valueLocationAtPoint(ranges.items, &value, position.Point.definition(2)));
411 try std.testing.expectEqual(@as(?u32, 1), index.valueLocationStart(ranges.items, &value));
412 try std.testing.expectEqual(@as(?u32, 3), index.valueLocationStart(ranges.items, &other));
413 try std.testing.expect(index.registerBlocksAt(2, position.Point.source(3)));
414 try std.testing.expectEqual(position.Point.source(3), index.registerRangeFirstOverlapStart(ranges.items, &value, 3, 5, 2, .reload, .definition).?);
415
416 var starts = index.rangesStartingAt(3);
417 try std.testing.expectEqual(@as(?usize, 1), starts.next());
418 try std.testing.expectEqual(@as(?usize, null), starts.next());
419 var ends = index.rangesEndingAt(4);
420 try std.testing.expectEqual(@as(?usize, 0), ends.next());
421 try std.testing.expectEqual(@as(?usize, null), ends.next());
422 }
423
424 test "value location index preserves phase and earliest overlap semantics" {
425 var owner: u8 = 0;
426 var value = ir.Value{
427 .kind = .{ .op_result = .{ .owner = &owner, .result_number = 0 } },
428 .type = undefined,
429 .id = 0,
430 };
431 var other = ir.Value{
432 .kind = .{ .op_result = .{ .owner = &owner, .result_number = 1 } },
433 .type = undefined,
434 .id = 1,
435 };
436 const Range = range.ValueLocationRange(u8);
437 const ranges = [_]Range{
438 .{ .value = &value, .start = 2, .end = 4, .reg = 7, .end_phase = .source },
439 .{ .value = &value, .start = 7, .end = 9, .reg = 2 },
440 .{ .value = &value, .start = 6, .end = 7, .reg = 2, .entry = .reload },
441 .{ .value = &other, .start = 3, .end = 5, .reg = 9 },
442 };
443 var index = ValueLocationIndex(u8){};
444 defer index.deinit(std.testing.allocator);
445 try index.rebuild(std.testing.allocator, &ranges);
446
447 try std.testing.expectEqual(@as(?u8, 7), index.valueLocationAtPoint(&ranges, &value, position.Point.source(3)));
448 try std.testing.expectEqual(@as(?u8, null), index.valueLocationAtPoint(&ranges, &value, position.Point.definition(3)));
449 try std.testing.expectEqual(@as(?u32, 2), index.valueLocationStart(&ranges, &value));
450 try std.testing.expect(index.valueHasLocationRange(&other));
451 try std.testing.expectEqual(position.Point.source(6), index.registerRangeFirstOverlapStart(&ranges, &other, 5, 10, 2, .reload, .definition).?);
452 try std.testing.expectEqual(@as(?position.Point, null), index.registerRangeFirstOverlapStart(&ranges, &other, 5, 10, 4, .reload, .definition));
453 try std.testing.expectEqual(@as(?position.Point, null), index.registerRangeFirstOverlapStart(&ranges, &other, 3, 4, 9, .resident, .source));
454 }