lib/choir/src/passes/saturation.zig

daab053ee43316e1809a84551d573ddd1e5bf3d2

  1 const std = @import("std");
  2 const ir = @import("../core/root.zig");
  3 const passes = @import("root.zig");
  4 const pass_mod = passes.pass;
  5 const rewrite = ir.rewrite;
  6 const effects = passes.effects;
  7 const egraph = @import("../egraph/root.zig");
  8 const graph_mod = egraph.graph;
  9 const rules_mod = egraph.rules;
 10 const pattern_mod = egraph.pattern;
 11 const extract_mod = egraph.extract;
 12 
 13 const ClassId = graph_mod.ClassId;
 14 const Graph = graph_mod.Graph;
 15 const RewriteContext = rules_mod.RewriteContext;
 16 const RewriteSet = rules_mod.RewriteSet;
 17 const CandidateFn = rules_mod.CandidateFn;
 18 const PopulateRulesFn = rules_mod.PopulateRulesFn;
 19 const work = pass_mod.work;
 20 
 21 /// Bounds for a fixed rule set that only merges a result with a dominating SSA
 22 /// value. Such rules add no nodes. Their cost model permits only linear failed
 23 /// materialization paths: extraction can explore a cycle before retaining an
 24 /// operation, but cannot create an expression. Arbitrary rules do not qualify.
 25 pub fn eliminationWorkBound(input: work.Input, rule_count: u64, iterations: u32) !work.Bounds {
 26     if (iterations == 0) return error.MissingWorkContract;
 27     const census = try work.Census.inspect(input.operation);
 28     const nodes = try work.add(census.values, census.operands);
 29     const blocks = try work.add(nodes, 1);
 30     const rebuilds = try work.add(nodes, try work.multiply(iterations, blocks));
 31     var storage = Graph.mergeStorageBound(nodes, census.atoms, blocks, rebuilds) catch
 32         return error.WorkOverflow;
 33     storage = try work.add(storage, extract_mod.Extraction.choiceStorageBound(nodes, blocks) catch
 34         return error.WorkOverflow);
 35     const candidates = try work.arrayListGrowth(CandidateResult, census.operations);
 36     const values = try work.hashMapGrowth(*ir.Value, ClassId, nodes);
 37     const operands = try work.arrayListGrowth(ClassId, census.operands);
 38     storage = try work.add(storage, try work.multiply(blocks, try work.add(
 39         candidates,
 40         try work.add(values, operands),
 41     )));
 42     storage = try work.add(storage, try work.arrayListGrowth(rules_mod.RewriteRule, rule_count));
 43     const visiting = try work.hashMapGrowth(u32, void, nodes);
 44     const operand_values = try work.arrayListGrowth(*ir.Value, census.operands);
 45     const materialization = try work.add(visiting, try work.multiply(nodes, operand_values));
 46     storage = try work.add(storage, try work.multiply(census.operations, materialization));
 47     const erasures = try work.arrayListGrowth(*ir.Operation, census.operations);
 48     storage = try work.add(storage, try work.multiply(2, erasures));
 49     const population = try work.add(nodes, 1);
 50     const scans = try work.multiply(try work.add(rebuilds, iterations), try work.multiply(
 51         population,
 52         try work.multiply(population, population),
 53     ));
 54     const units = try work.add(try work.add(census.atoms, census.input_bytes), 1);
 55     const attempts_per_iteration = try work.multiply(rule_count, try work.multiply(population, population));
 56     const attempts = try work.multiply(iterations, attempts_per_iteration);
 57     const visits = try work.multiply(scans, try work.multiply(try work.add(rule_count, 64), units));
 58     return .{
 59         .work = .{
 60             .input_bytes = census.input_bytes,
 61             .output_bytes = input.operation.context.capacity.storage_bytes,
 62             .structural_visits = visits,
 63             .rewrite_attempts = attempts,
 64             .allocation_capacity = try work.add(storage, input.operation.context.capacity.storage_bytes),
 65         },
 66         .workspace = storage,
 67     };
 68 }
 69 
 70 pub const OptimizationOptions = struct {
 71     max_iterations: u32 = 8,
 72     max_rewrites: u32 = 1024,
 73     candidate: CandidateFn = defaultCandidate,
 74     candidate_context: ?*anyopaque = null,
 75 };
 76 
 77 pub const OptimizationStats = struct {
 78     classes_created: usize = 0,
 79     nodes_added: usize = 0,
 80     unions: usize = 0,
 81     rebuilds: usize = 0,
 82     rewrites: usize = 0,
 83     replacements: usize = 0,
 84     iterations: u64 = 0,
 85     termination: OptimizationTermination = .converged,
 86 
 87     pub fn modified(self: OptimizationStats) bool {
 88         return self.replacements != 0;
 89     }
 90 };
 91 
 92 pub const OptimizationTermination = enum { converged, iteration_limit, rewrite_limit };
 93 
 94 /// Physical counters do not replace pre-pass admission. Transient callers may
 95 /// consume partial optimization; an accounted producer cannot publish it as success.
 96 pub fn observeOptimization(ctx: *pass_mod.PassContext, stats: OptimizationStats) bool {
 97     const ledger = ctx.analysis_cache.accounting orelse return true;
 98     if (stats.termination != .converged) ledger.fail(.exhausted);
 99     ledger.observeCounters(.{
100         .successful_rewrites = stats.rewrites,
101         .rewrite_iterations = stats.iterations,
102     }) catch return false;
103     return stats.termination == .converged;
104 }
105 
106 pub const PassConfig = struct {
107     name: []const u8,
108     description: []const u8,
109     populate_rules: PopulateRulesFn,
110     options: OptimizationOptions = .{},
111     mutation_scope: pass_mod.PassMutationScope = .isolated,
112 };
113 
114 pub fn EGraphPass(comptime config: PassConfig) type {
115     return struct {
116         pub fn create() pass_mod.Pass {
117             return .{
118                 .name = config.name,
119                 .description = config.description,
120                 .run_fn = run,
121                 .mutation_scope = config.mutation_scope,
122             };
123         }
124 
125         fn run(ctx: *pass_mod.PassContext) pass_mod.PassResult {
126             var rules = RewriteSet.init(ctx.allocator);
127             defer rules.deinit();
128 
129             config.populate_rules(&rules) catch return .failure;
130             const stats = runOptimization(ctx.allocator, ctx.ir_ctx, ctx.op, &rules, config.options) catch return .failure;
131             if (stats.modified()) {
132                 ctx.markModified();
133             } else {
134                 ctx.preserveAllAnalyses();
135             }
136             if (!observeOptimization(ctx, stats)) return .failure;
137             return .success;
138         }
139     };
140 }
141 
142 pub fn defaultCandidate(_: ?*anyopaque, op: *ir.Operation) anyerror!bool {
143     if (op.getNumResults() != 1) return false;
144     if (op.getNumRegions() != 0) return false;
145     if (op.getNumSuccessors() != 0) return false;
146 
147     const traits = op.getTraits();
148     if (traits.is_terminator) return false;
149     if (op.hasInterface(ir.interfaces.SymbolOpInterface)) return false;
150     return effects.permitsRepeatableExpression(op);
151 }
152 
153 pub fn runOptimization(
154     allocator: std.mem.Allocator,
155     ir_ctx: *ir.Context,
156     root: *ir.Operation,
157     rules: *RewriteSet,
158     options: OptimizationOptions,
159 ) !OptimizationStats {
160     rules.sortByBenefit();
161 
162     var rewriter = rewrite.PatternRewriter.init(allocator, ir_ctx);
163     defer rewriter.deinit();
164 
165     var state = OptimizerState{
166         .allocator = allocator,
167         .ir_ctx = ir_ctx,
168         .rules = rules,
169         .options = options,
170         .rewriter = &rewriter,
171     };
172     try state.optimizeOperation(root);
173     if (state.stats.replacements > 0) {
174         rewriter.finalize(root);
175     }
176     return state.stats;
177 }
178 
179 const CandidateResult = struct {
180     op: *ir.Operation,
181     value: *ir.Value,
182     class: ClassId,
183     order: usize,
184 };
185 
186 const OptimizerState = struct {
187     allocator: std.mem.Allocator,
188     ir_ctx: *ir.Context,
189     rules: *RewriteSet,
190     options: OptimizationOptions,
191     rewriter: *rewrite.PatternRewriter,
192     stats: OptimizationStats = .{},
193     value_order: usize = 0,
194 
195     fn optimizeOperation(self: *OptimizerState, op: *ir.Operation) !void {
196         for (op.regions.items) |*region| {
197             var block_iter = region.getBlocks();
198             while (block_iter.next()) |block| {
199                 try self.optimizeBlock(block);
200 
201                 var current: ?*ir.Operation = @ptrCast(@alignCast(block.operations.head));
202                 while (current) |nested_op| {
203                     try self.optimizeOperation(nested_op);
204                     current = nested_op.next_op;
205                 }
206             }
207         }
208     }
209 
210     fn optimizeBlock(self: *OptimizerState, block: *ir.Block) !void {
211         var graph = Graph.init(self.allocator);
212         defer graph.deinit();
213 
214         var value_classes = std.AutoHashMap(*ir.Value, ClassId).init(self.allocator);
215         defer value_classes.deinit();
216 
217         var candidates: std.ArrayListUnmanaged(CandidateResult) = .empty;
218         defer candidates.deinit(self.allocator);
219 
220         for (block.arguments.items) |argument| {
221             const class = try graph.addValue(argument, 0, self.nextOrder());
222             try value_classes.put(argument, class);
223         }
224 
225         var current: ?*ir.Operation = @ptrCast(@alignCast(block.operations.head));
226         while (current) |op| {
227             const candidate = op.getNumResults() == 1 and
228                 op.getNumRegions() == 0 and
229                 op.getNumSuccessors() == 0 and
230                 try defaultCandidate(null, op) and
231                 try self.options.candidate(self.options.candidate_context, op);
232 
233             if (candidate) {
234                 var operand_classes = std.ArrayListUnmanaged(ClassId).empty;
235                 defer operand_classes.deinit(self.allocator);
236                 try operand_classes.ensureTotalCapacity(self.allocator, op.getOperandValues().len);
237 
238                 for (op.getOperandValues()) |operand| {
239                     operand_classes.appendAssumeCapacity(try self.classForValue(&graph, &value_classes, operand));
240                 }
241 
242                 const order = self.nextOrder();
243                 const class = try graph.addOperation(op, operand_classes.items, 1, order);
244                 const result = op.getResult(0).?;
245                 try value_classes.put(result, class);
246                 try candidates.append(self.allocator, .{ .op = op, .value = result, .class = class, .order = order });
247             } else {
248                 for (op.results.items) |*result| {
249                     const class = try graph.addValue(result, 0, self.nextOrder());
250                     try value_classes.put(result, class);
251                 }
252             }
253 
254             current = op.next_op;
255         }
256 
257         try self.saturate(&graph);
258         try self.replaceCandidates(&graph, &value_classes, candidates.items);
259 
260         self.stats.classes_created += graph.stats.classes_created;
261         self.stats.nodes_added += graph.stats.nodes_added;
262         self.stats.unions += graph.stats.unions;
263         self.stats.rebuilds += graph.stats.rebuilds;
264     }
265 
266     fn nextOrder(self: *OptimizerState) usize {
267         const order = self.value_order;
268         self.value_order += 1;
269         return order;
270     }
271 
272     fn classForValue(
273         self: *OptimizerState,
274         graph: *Graph,
275         value_classes: *std.AutoHashMap(*ir.Value, ClassId),
276         value: *ir.Value,
277     ) !ClassId {
278         if (value_classes.get(value)) |existing| return graph.find(existing);
279         const class = try graph.addValue(value, 0, self.nextOrder());
280         try value_classes.put(value, class);
281         return class;
282     }
283 
284     fn saturate(self: *OptimizerState, graph: *Graph) !void {
285         const iteration_limit = if (self.options.max_iterations == 0) std.math.maxInt(u32) else self.options.max_iterations;
286         const rewrite_limit = if (self.options.max_rewrites == 0) std.math.maxInt(u32) else self.options.max_rewrites;
287 
288         var iteration: u32 = 0;
289         while (iteration < iteration_limit) : (iteration += 1) {
290             self.stats.iterations += 1;
291             const budget = rewrite_limit - self.stats.rewrites;
292             const applied = try self.applyRulesOnce(graph, budget);
293             self.stats.rewrites += applied;
294             if (applied == 0) return;
295             _ = try graph.rebuild();
296             if (self.stats.rewrites >= rewrite_limit) {
297                 self.noteTermination(.rewrite_limit);
298                 return;
299             }
300         }
301         self.noteTermination(.iteration_limit);
302     }
303 
304     fn noteTermination(self: *OptimizerState, cause: OptimizationTermination) void {
305         if (self.stats.termination == .converged) self.stats.termination = cause;
306     }
307 
308     fn applyRulesOnce(self: *OptimizerState, graph: *Graph, budget: usize) !usize {
309         var ctx = RewriteContext{ .graph = graph };
310         var applied: usize = 0;
311 
312         const class_count = graph.classCount();
313         for (0..class_count) |index| {
314             const id = ClassId{ .index = @intCast(index) };
315             const root = graph.find(id);
316             if (!root.eql(id)) continue;
317 
318             var node_index: usize = 0;
319             while (node_index < graph.nodes(root).len) : (node_index += 1) {
320                 const node = graph.nodes(root)[node_index];
321                 for (self.rules.rules.items) |rule| {
322                     if (applied >= budget) return applied;
323                     if (try rule.apply(&ctx, root, &node)) {
324                         applied += 1;
325                     }
326                 }
327                 if (self.rules.constant_model) |*model| {
328                     for (self.rules.patterns.items) |*rule| {
329                         if (applied >= budget) return applied;
330                         if (try pattern_mod.applyRule(graph, self.ir_ctx, model, rule, root, &node)) {
331                             applied += 1;
332                         }
333                     }
334                 }
335             }
336         }
337 
338         return applied;
339     }
340 
341     fn replaceCandidates(
342         self: *OptimizerState,
343         graph: *Graph,
344         value_classes: *std.AutoHashMap(*ir.Value, ClassId),
345         candidates: []const CandidateResult,
346     ) !void {
347         if (candidates.len == 0) return;
348 
349         var extraction = try extract_mod.Extraction.init(self.allocator, graph, self.rules.resolvedCostModel());
350         defer extraction.deinit();
351         extraction.analyze();
352 
353         for (candidates) |candidate| {
354             const root = graph.find(candidate.class);
355 
356             if (extraction.usableValue(root, candidate.order, candidate.value)) |value| {
357                 try self.rewriter.replaceOpWithValue(candidate.op, value);
358                 self.stats.replacements += 1;
359                 continue;
360             }
361 
362             var own_cost: u64 = extraction.model.operationCost(candidate.op.name.name);
363             for (candidate.op.getOperandValues()) |operand| {
364                 const operand_class = value_classes.get(operand) orelse continue;
365                 own_cost +|= extraction.classCost(operand_class);
366             }
367 
368             if (extraction.nodeCost(root) >= own_cost) continue;
369             const replacement = try extraction.materialize(
370                 self.rewriter,
371                 root,
372                 candidate.op,
373                 candidate.order,
374                 candidate.value,
375             ) orelse continue;
376             if (replacement == candidate.value) continue;
377             try self.rewriter.replaceOpWithValue(candidate.op, replacement);
378             self.stats.replacements += 1;
379         }
380     }
381 };
382 
383 test "saturation termination distinguishes convergence and both internal limits" {
384     try checkSaturationTermination(.{ .max_iterations = 1 }, .iteration_limit, 1);
385     try checkSaturationTermination(.{ .max_rewrites = 1 }, .rewrite_limit, 1);
386     try checkSaturationTermination(.{ .max_iterations = 2 }, .converged, 2);
387     try checkSaturationTermination(.{ .max_iterations = 0, .max_rewrites = 0 }, .converged, 2);
388 }
389 
390 fn mergeFirstClass(ctx: *RewriteContext, class: ClassId, _: *const graph_mod.Node) !bool {
391     return ctx.merge(class, .{ .index = 0 });
392 }
393 
394 fn checkSaturationTermination(
395     options: OptimizationOptions,
396     expected: OptimizationTermination,
397     iterations: u64,
398 ) !void {
399     const allocator = std.testing.allocator;
400     var context = try ir.Context.init(allocator, ir.Context.Limits.testing);
401     defer context.deinit(allocator);
402     try context.allowUnregistered();
403     var op_builder = ir.OperationBuilder.init(&context);
404     const root = try op_builder.create(ir.Operation.State.init("test.root", ir.Location.getUnknown()));
405     var rewriter = rewrite.PatternRewriter.init(allocator, &context);
406     defer rewriter.deinit();
407     var graph = Graph.init(allocator);
408     defer graph.deinit();
409     _ = try graph.addNode(&.{ .kind = .operation, .op_name = "test.left" });
410     _ = try graph.addNode(&.{ .kind = .operation, .op_name = "test.right" });
411     var rules = RewriteSet.init(allocator);
412     defer rules.deinit();
413     try rules.add(.{ .name = "merge-first", .apply = mergeFirstClass });
414     var state = OptimizerState{
415         .allocator = allocator,
416         .ir_ctx = &context,
417         .rules = &rules,
418         .options = options,
419         .rewriter = &rewriter,
420     };
421     try state.saturate(&graph);
422     try std.testing.expectEqual(expected, state.stats.termination);
423     try std.testing.expectEqual(iterations, state.stats.iterations);
424     try std.testing.expectEqual(1, state.stats.rewrites);
425     try checkOptimizationReceipt(root, state.stats);
426 }
427 
428 fn checkOptimizationReceipt(root: *ir.Operation, stats: OptimizationStats) !void {
429     const allocator = std.testing.allocator;
430     const revision = @import("../product/root.zig").revision;
431     const identity = revision.record.Version{ .name = "saturation-test", .version = 1 };
432     const ledger = try revision.AccountingV1.create(allocator, .{
433         .allowance = .uniform(std.math.maxInt(u64)),
434         .workspace = 1024 * 1024,
435         .events = 16,
436     }, &.{identity});
437     defer ledger.destroy();
438     var cache = try pass_mod.AnalysisCache.initAccounted(allocator, null, ledger, .{}, 8);
439     defer cache.deinit();
440     var context = pass_mod.PassContext.init(root, root.context, allocator, &cache);
441     defer context.deinit();
442     const token = try ledger.begin(.pass, .{ .identity = identity, .work = .{} });
443     try std.testing.expectEqual(
444         stats.termination == .converged,
445         observeOptimization(&context, stats),
446     );
447     try ledger.finish(token, .success, .{});
448     const receipt = ledger.view();
449     try std.testing.expectEqual(stats.rewrites, receipt.executed.counters.successful_rewrites);
450     try std.testing.expectEqual(stats.iterations, receipt.executed.counters.rewrite_iterations);
451     const exhausted = stats.termination != .converged;
452     try std.testing.expectEqual(exhausted, receipt.outcome == .exhausted);
453     try std.testing.expectEqual(exhausted, receipt.events[token].outcome == .exhausted);
454 }