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 }