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;