lib/choir/src/egraph/rules.zig
daab053ee43316e1809a84551d573ddd1e5bf3d2
1 const std = @import("std");
2 const ir = @import("../core/root.zig");
3 const graph_mod = @import("graph.zig");
4 const pattern_mod = @import("pattern.zig");
5 const extract_mod = @import("extract.zig");
6
7 const ClassId = graph_mod.ClassId;
8 const Node = graph_mod.Node;
9 const Graph = graph_mod.Graph;
10
11 pub const RewriteContext = struct {
12 graph: *Graph,
13
14 pub fn merge(self: *RewriteContext, lhs: ClassId, rhs: ClassId) !bool {
15 return self.graph.merge(lhs, rhs);
16 }
17
18 pub fn nodes(self: *RewriteContext, id: ClassId) []const Node {
19 return self.graph.nodes(id);
20 }
21 };
22
23 pub const RewriteFn = *const fn (*RewriteContext, ClassId, *const Node) anyerror!bool;
24
25 pub const RewriteRule = struct {
26 name: []const u8,
27 benefit: u32 = 1,
28 apply: RewriteFn,
29 };
30
31 fn rewrite_rule_benefit_greater_than(_: void, lhs: RewriteRule, rhs: RewriteRule) bool {
32 return lhs.benefit > rhs.benefit;
33 }
34
35 fn pattern_rule_benefit_greater_than(
36 _: void,
37 lhs: pattern_mod.Rule,
38 rhs: pattern_mod.Rule,
39 ) bool {
40 return lhs.benefit > rhs.benefit;
41 }
42
43 pub const RewriteSet = struct {
44 allocator: std.mem.Allocator,
45 rules: std.ArrayListUnmanaged(RewriteRule) = .empty,
46 patterns: std.ArrayListUnmanaged(pattern_mod.Rule) = .empty,
47 constant_model: ?pattern_mod.ConstantModel = null,
48 cost_model: ?extract_mod.CostModel = null,
49
50 pub fn init(allocator: std.mem.Allocator) RewriteSet {
51 return .{ .allocator = allocator };
52 }
53
54 pub fn deinit(self: *RewriteSet) void {
55 self.rules.deinit(self.allocator);
56 self.patterns.deinit(self.allocator);
57 }
58
59 pub fn add(self: *RewriteSet, rule: RewriteRule) !void {
60 try self.rules.append(self.allocator, rule);
61 }
62
63 pub fn addPattern(self: *RewriteSet, rule: pattern_mod.Rule) !void {
64 try self.patterns.append(self.allocator, rule);
65 }
66
67 pub fn addPatterns(self: *RewriteSet, pattern_rules: []const pattern_mod.Rule) !void {
68 try self.patterns.appendSlice(self.allocator, pattern_rules);
69 }
70
71 pub fn setConstantModel(self: *RewriteSet, model: pattern_mod.ConstantModel) void {
72 self.constant_model = model;
73 }
74
75 pub fn setCostModel(self: *RewriteSet, model: extract_mod.CostModel) void {
76 self.cost_model = model;
77 }
78
79 pub fn resolvedCostModel(self: *const RewriteSet) extract_mod.CostModel {
80 var model = self.cost_model orelse extract_mod.CostModel{};
81 if (model.constant_op_name == null) {
82 if (self.constant_model) |constant_model| {
83 model.constant_op_name = constant_model.op_name;
84 }
85 }
86 return model;
87 }
88
89 pub fn sortByBenefit(self: *RewriteSet) void {
90 std.mem.sort(RewriteRule, self.rules.items, {}, rewrite_rule_benefit_greater_than);
91 std.mem.sort(pattern_mod.Rule, self.patterns.items, {}, pattern_rule_benefit_greater_than);
92 }
93 };
94
95 pub const CandidateFn = *const fn (?*anyopaque, *ir.Operation) anyerror!bool;
96 pub const PopulateRulesFn = *const fn (*RewriteSet) anyerror!void;