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 }