lib/choir/src/core/rewrite/set.zig
daab053ee43316e1809a84551d573ddd1e5bf3d2
1 const std = @import("std");
2 const ir = @import("../root.zig");
3 const alloc_phase = @import("alloc_phase");
4 const PatternBenefit = ir.rewrite.PatternBenefit;
5 const RewritePatternSpec = ir.rewrite.RewritePatternSpec;
6 const RewritePattern = ir.rewrite.RewritePattern;
7 const PatternResult = ir.rewrite.PatternResult;
8 const tryApplyRewritePattern = ir.rewrite.tryApplyRewritePattern;
9 const PatternRewriter = ir.rewrite.PatternRewriter;
10 const PatternIndex = ir.rewrite.PatternIndex;
11
12 pub const RewritePatternSet = struct {
13 allocator: std.mem.Allocator,
14 patterns: std.ArrayListUnmanaged(RewritePattern),
15 index: ?PatternIndex,
16
17 pub fn init(allocator: std.mem.Allocator) RewritePatternSet {
18 return .{
19 .allocator = allocator,
20 .patterns = .empty,
21 .index = null,
22 };
23 }
24
25 pub fn deinit(self: *RewritePatternSet) void {
26 if (self.index) |*index| index.deinit(self.allocator);
27 self.patterns.deinit(self.allocator);
28 self.* = undefined;
29 }
30
31 pub fn add(self: *RewritePatternSet, pattern: RewritePattern) !void {
32 if (self.index != null) return error.PatternSetSealed;
33 var ordered = pattern;
34 ordered.order = self.patterns.items.len;
35 try self.patterns.append(self.allocator, ordered);
36 }
37
38 pub fn ensureUnusedCapacity(self: *RewritePatternSet, additional_count: usize) !void {
39 if (self.index != null) return error.PatternSetSealed;
40 const required_count = std.math.add(
41 usize,
42 self.patterns.items.len,
43 additional_count,
44 ) catch return error.CapacityOverflow;
45 try self.patterns.ensureTotalCapacity(self.allocator, required_count);
46 }
47
48 pub fn addPattern(
49 self: *RewritePatternSet,
50 op_name: []const u8,
51 rewrite_fn: *const fn (*ir.Operation, *PatternRewriter) PatternResult,
52 ) !void {
53 try self.add(RewritePattern.init(.{
54 .name = op_name,
55 .root_op_name = op_name,
56 }, rewrite_fn));
57 }
58
59 pub fn count(self: *const RewritePatternSet) usize {
60 if (self.index) |index| return index.capacity.facts.pattern_count;
61 return self.patterns.items.len;
62 }
63
64 pub fn seal(self: *RewritePatternSet) !void {
65 if (self.index != null) return;
66 std.mem.sort(RewritePattern, self.patterns.items, {}, RewritePattern.lessThan);
67 const limits = try PatternIndex.Limits.inspect(self.patterns.items);
68 var index = try PatternIndex.init(self.allocator, limits);
69 errdefer index.deinit(self.allocator);
70 try index.activate();
71 self.patterns.deinit(self.allocator);
72 self.patterns = .empty;
73 self.index = index;
74 }
75
76 pub fn getMatchingPatterns(self: *const RewritePatternSet, op: *ir.Operation) []const RewritePattern {
77 if (self.index) |*index| {
78 return index.matching(op.name.name);
79 }
80 @panic("rewrite pattern set is not sealed");
81 }
82
83 pub fn applyFirstMatchingPattern(
84 self: *const RewritePatternSet,
85 op: *ir.Operation,
86 rewriter: *PatternRewriter,
87 ) bool {
88 for (self.getMatchingPatterns(op)) |*pattern| {
89 if (tryApplyRewritePattern(pattern, op, rewriter)) return true;
90 }
91 return false;
92 }
93 };
94
95 test "rewrite pattern builder reuses reserved capacity" {
96 const testing = std.testing;
97
98 var phase_allocator = try alloc_phase.SealedPhaseAllocator.init(testing.allocator);
99 var maybe_patterns: ?RewritePatternSet = RewritePatternSet.init(
100 phase_allocator.initializationAllocator(),
101 );
102 errdefer {
103 if (phase_allocator.phase() == .initialization) {
104 phase_allocator.abortInitialization();
105 }
106 if (phase_allocator.phase() == .steady) phase_allocator.beginTeardown();
107 if (maybe_patterns) |*patterns| patterns.deinit();
108 if (phase_allocator.phase() == .teardown) phase_allocator.deinit();
109 }
110 const patterns = &maybe_patterns.?;
111
112 try patterns.ensureUnusedCapacity(3);
113 const base_pointer = patterns.patterns.items.ptr;
114 phase_allocator.seal();
115
116 try patterns.add(RewritePattern.init(
117 testRewriteSpec("test.pattern_a", 1),
118 rewriteNoopForPatternOrder,
119 ));
120 try patterns.add(RewritePattern.init(
121 testRewriteSpec("test.pattern_b", 2),
122 rewriteNoopForPatternOrder,
123 ));
124 try patterns.add(RewritePattern.init(
125 testRewriteSpec("test.pattern_c", 3),
126 rewriteNoopForPatternOrder,
127 ));
128 try testing.expectEqual(base_pointer, patterns.patterns.items.ptr);
129 try testing.expectError(
130 error.CapacityOverflow,
131 patterns.ensureUnusedCapacity(std.math.maxInt(usize)),
132 );
133 try testing.expectEqual(
134 alloc_phase.PhaseViolations{},
135 phase_allocator.violations(),
136 );
137
138 phase_allocator.beginTeardown();
139 patterns.deinit();
140 maybe_patterns = null;
141 try testing.expectEqual(
142 alloc_phase.PhaseViolations{},
143 phase_allocator.violations(),
144 );
145 phase_allocator.deinit();
146 }
147
148 test "RewritePatternSet seal preserves equal-benefit insertion order per root" {
149 const testing = std.testing;
150 const allocator = testing.allocator;
151
152 var patterns = RewritePatternSet.init(allocator);
153 defer patterns.deinit();
154
155 try patterns.add(RewritePattern.init(testRewriteSpec("test.a", 1), rewriteNoopForPatternOrder));
156 try patterns.add(RewritePattern.init(testRewriteSpec("test.b", 1), rewriteNoopForPatternOrder));
157 try patterns.add(RewritePattern.init(testRewriteSpec("test.a", 1), rewriteNoopForPatternOrder));
158
159 try patterns.seal();
160
161 try testing.expectEqual(@as(usize, 0), patterns.patterns.items.len);
162 try testing.expectEqual(@as(usize, 3), patterns.count());
163 const matching_a = patterns.index.?.matching("test.a");
164 try testing.expectEqual(@as(usize, 2), matching_a.len);
165 try testing.expectEqual(@as(usize, 0), matching_a[0].order);
166 try testing.expectEqual(@as(usize, 2), matching_a[1].order);
167 const matching_b = patterns.index.?.matching("test.b");
168 try testing.expectEqual(@as(usize, 1), matching_b.len);
169 try testing.expectEqual(@as(usize, 1), matching_b[0].order);
170 }
171
172 test "RewritePatternSet indexes matching patterns by root operation" {
173 const testing = std.testing;
174 const allocator = testing.allocator;
175
176 var ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
177 defer ctx.deinit(allocator);
178 try ctx.allowUnregistered();
179
180 const a = try ctx.createOperation(ir.Operation.State.init("test.a", .unknown));
181 const b = try ctx.createOperation(ir.Operation.State.init("test.b", .unknown));
182 const c = try ctx.createOperation(ir.Operation.State.init("test.c", .unknown));
183
184 var patterns = RewritePatternSet.init(allocator);
185 defer patterns.deinit();
186
187 try patterns.add(RewritePattern.init(testRewriteSpec("test.b", 1), rewriteNoopForPatternOrder));
188 try patterns.add(RewritePattern.init(testRewriteSpec("test.a", 1), rewriteNoopForPatternOrder));
189 try patterns.add(RewritePattern.init(testRewriteSpec("test.a", 2), rewriteNoopForPatternOrder));
190 try patterns.seal();
191
192 const a_patterns = patterns.getMatchingPatterns(a);
193 try testing.expectEqual(@as(usize, 2), a_patterns.len);
194 const first_a = a_patterns[0];
195 const second_a = a_patterns[1];
196 try testing.expectEqual(@as(PatternBenefit, 2), first_a.spec.benefit);
197 try testing.expectEqual(@as(PatternBenefit, 1), second_a.spec.benefit);
198 try testing.expectEqualStrings("test.a", first_a.spec.root_op_name);
199
200 const b_patterns = patterns.getMatchingPatterns(b);
201 try testing.expectEqual(@as(usize, 1), b_patterns.len);
202 try testing.expectEqualStrings("test.b", b_patterns[0].spec.root_op_name);
203
204 try testing.expectEqual(@as(usize, 0), patterns.getMatchingPatterns(c).len);
205 }
206
207 test "rewrite pattern set" {
208 const testing = std.testing;
209 const allocator = testing.allocator;
210
211 var patterns = RewritePatternSet.init(allocator);
212 defer patterns.deinit();
213
214 try patterns.add(RewritePattern.init(testRewriteSpec("arith.addi", 1), dummyRewrite));
215 try patterns.add(RewritePattern.init(testRewriteSpec("arith.muli", 2), dummyRewrite));
216 try patterns.add(RewritePattern.init(testRewriteSpec("arith.subi", 3), dummyRewrite));
217
218 try testing.expectEqual(@as(usize, 3), patterns.patterns.items.len);
219
220 try patterns.seal();
221
222 try testing.expectEqual(alloc_phase.capacity.Phase.steady, patterns.index.?.phase);
223 }
224
225 fn dummyRewrite(op: *ir.Operation, rewriter: *PatternRewriter) PatternResult {
226 rewriter.eraseOp(op) catch return .failure;
227 return .success;
228 }
229
230 fn rewriteNoopForPatternOrder(_: *ir.Operation, _: *PatternRewriter) PatternResult {
231 return .failure;
232 }
233
234 fn testRewriteSpec(root_op_name: []const u8, benefit: PatternBenefit) RewritePatternSpec {
235 return .{
236 .name = root_op_name,
237 .root_op_name = root_op_name,
238 .benefit = benefit,
239 };
240 }