lib/choir/src/backends/regalloc/loops.zig
daab053ee43316e1809a84551d573ddd1e5bf3d2
1 const std = @import("std");
2 const ir = @import("../../core/root.zig");
3 const interval = @import("interval.zig");
4 const position = @import("position.zig");
5
6 pub const Interval = struct {
7 entry: u32,
8 trailing: u32,
9
10 pub fn containsPosition(self: Interval, pos: u32) bool {
11 return pos >= self.entry and pos <= self.trailing;
12 }
13 };
14
15 pub fn extendAcrossLoops(
16 comptime Register: type,
17 comptime Mask: type,
18 candidates: []interval.Candidate(Register, Mask),
19 loops: []const Interval,
20 ) void {
21 for (candidates) |*candidate| {
22 for (candidate.use_positions.items) |use| {
23 for (loops) |loop| {
24 if (!loop.containsPosition(use.point.position)) continue;
25 if (candidate.start() >= loop.entry) continue;
26 const end_point = position.Point.source(loop.trailing);
27 if (end_point.greaterThan(candidate.endPoint())) {
28 candidate.range.end = end_point.position;
29 candidate.range.end_phase = end_point.phase;
30 }
31 }
32 }
33 }
34 }
35
36 pub fn rangeCrossesLoopEntry(start: u32, end: u32, loops: []const Interval) bool {
37 for (loops) |loop| {
38 if (start <= loop.entry and end >= loop.entry) return true;
39 }
40 return false;
41 }
42
43 pub fn evictionUnsafeAt(victim_start: u32, at: u32, loops: []const Interval) bool {
44 for (loops) |loop| {
45 if (victim_start < loop.entry and loop.containsPosition(at)) return true;
46 }
47 return false;
48 }
49
50 pub fn innermostTrailing(use_position: u32, loops: []const Interval) ?u32 {
51 var clamp: ?u32 = null;
52 for (loops) |loop| {
53 if (!loop.containsPosition(use_position)) continue;
54 if (clamp) |current| {
55 if (loop.trailing < current) clamp = loop.trailing;
56 } else {
57 clamp = loop.trailing;
58 }
59 }
60 return clamp;
61 }
62
63 test "loop intervals extend crossing candidates to the trailing position" {
64 var owner: u8 = 0;
65 var outside = ir.Value{
66 .kind = .{ .op_result = .{ .owner = &owner, .result_number = 0 } },
67 .type = undefined,
68 .id = 0,
69 };
70 var inside = ir.Value{
71 .kind = .{ .op_result = .{ .owner = &owner, .result_number = 1 } },
72 .type = undefined,
73 .id = 1,
74 };
75
76 const CandidateType = interval.Candidate(u8, u8);
77 var outside_uses: std.ArrayListUnmanaged(interval.UsePosition(u8, u8)) = .empty;
78 defer outside_uses.deinit(std.testing.allocator);
79 try outside_uses.append(std.testing.allocator, .{
80 .point = position.Point.source(4),
81 .requirement = .any,
82 .source_blockers = 0,
83 });
84 var inside_uses: std.ArrayListUnmanaged(interval.UsePosition(u8, u8)) = .empty;
85 defer inside_uses.deinit(std.testing.allocator);
86 try inside_uses.append(std.testing.allocator, .{
87 .point = position.Point.source(6),
88 .requirement = .any,
89 .source_blockers = 0,
90 });
91
92 var candidates = [_]CandidateType{
93 .{
94 .value = &outside,
95 .range = .{ .start = 0, .end = 4, .end_phase = .source },
96 .use_positions = outside_uses,
97 .definition = .{
98 .point = position.Point.definition(0),
99 .requirement = .any,
100 .source = .any,
101 },
102 .order = 0,
103 .is_constant = false,
104 },
105 .{
106 .value = &inside,
107 .range = .{ .start = 5, .end = 6, .end_phase = .source },
108 .use_positions = inside_uses,
109 .definition = .{
110 .point = position.Point.definition(5),
111 .requirement = .any,
112 .source = .any,
113 },
114 .order = 1,
115 .is_constant = false,
116 },
117 };
118
119 const loops = [_]Interval{.{ .entry = 3, .trailing = 8 }};
120 extendAcrossLoops(u8, u8, &candidates, &loops);
121
122 try std.testing.expectEqual(@as(u32, 8), candidates[0].range.end);
123 try std.testing.expectEqual(position.Phase.source, candidates[0].range.end_phase);
124 try std.testing.expectEqual(@as(u32, 6), candidates[1].range.end);
125 }
126
127 test "loop crossing and eviction predicates track entry positions" {
128 const loops = [_]Interval{
129 .{ .entry = 3, .trailing = 8 },
130 .{ .entry = 5, .trailing = 7 },
131 };
132
133 try std.testing.expect(rangeCrossesLoopEntry(0, 4, &loops));
134 try std.testing.expect(rangeCrossesLoopEntry(3, 8, &loops));
135 try std.testing.expect(rangeCrossesLoopEntry(5, 7, &loops));
136 try std.testing.expect(!rangeCrossesLoopEntry(0, 2, &loops));
137
138 try std.testing.expect(evictionUnsafeAt(0, 6, &loops));
139 try std.testing.expect(!evictionUnsafeAt(0, 2, &loops));
140 try std.testing.expect(!evictionUnsafeAt(3, 4, &loops));
141 try std.testing.expect(evictionUnsafeAt(3, 6, &loops));
142
143 try std.testing.expectEqual(@as(?u32, 7), innermostTrailing(6, &loops));
144 try std.testing.expectEqual(@as(?u32, 8), innermostTrailing(4, &loops));
145 try std.testing.expectEqual(@as(?u32, null), innermostTrailing(9, &loops));
146 }