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 }