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 }