lib/choir/src/dialects/arith/patterns.zig
daab053ee43316e1809a84551d573ddd1e5bf3d2
1 const std = @import("std");
2 const ir = @import("../../core/root.zig");
3 const rewrite = ir.rewrite;
4 const fold_mod = @import("folds.zig");
5
6 pub fn Patterns(comptime Dialect: type) type {
7 return struct {
8 const folds = fold_mod.Folds(Dialect);
9
10 pub const canonicalization_patterns = [_]rewrite.RewritePattern{
11 rewrite.RewritePattern.init(
12 .{
13 .name = "arith-select-canonicalization",
14 .root_op_name = Dialect.SelectOp.operation_name,
15 .benefit = 50,
16 .products = .{ .operations = &.{Dialect.NotOp.operation_name} },
17 },
18 rewriteSelectCanonicalization,
19 ),
20 rewrite.RewritePattern.init(
21 .{
22 .name = "arith-not-canonicalization",
23 .root_op_name = Dialect.NotOp.operation_name,
24 .benefit = 50,
25 .products = .{ .operations = &.{Dialect.ConstantOp.operation_name} },
26 },
27 rewriteNotCanonicalization,
28 ),
29 };
30
31 fn rewriteSelectCanonicalization(op: *ir.Operation, rewriter: *rewrite.PatternRewriter) rewrite.PatternResult {
32 var declaration = ir.interfaces.effects.inspect(op.allocator, op) catch return .failure;
33 defer declaration.deinit(op.allocator);
34 if (!ir.interfaces.effects.repeatableExpression(declaration.facts)) return .failure;
35 if (op.getNumResults() != 1) return .failure;
36 const select = Dialect.SelectOp{ .op = op };
37 const condition_value = select.getCondition();
38 const true_value = select.getTrueValue();
39 const false_value = select.getFalseValue();
40 if (true_value == false_value) return rewriteSameTypeForwarding(op, true_value, rewriter);
41 if (true_value == condition_value) {
42 if (folds.constantBoolEquals(false_value, false)) return rewriteSameTypeForwarding(op, condition_value, rewriter);
43 if (folds.constantBoolEquals(false_value, true)) return rewriteSameTypeForwarding(op, false_value, rewriter);
44 }
45 if (false_value == condition_value) {
46 if (folds.constantBoolEquals(true_value, true)) return rewriteSameTypeForwarding(op, condition_value, rewriter);
47 if (folds.constantBoolEquals(true_value, false)) return rewriteSameTypeForwarding(op, true_value, rewriter);
48 }
49 if (folds.constantBoolEquals(true_value, true) and folds.constantBoolEquals(false_value, false)) {
50 return rewriteSameTypeForwarding(op, condition_value, rewriter);
51 }
52 if (folds.constantBoolEquals(true_value, false) and folds.constantBoolEquals(false_value, true)) {
53 return rewriteBoolNot(op, condition_value, rewriter);
54 }
55 const condition = folds.constantBoolFromValue(condition_value) orelse return .failure;
56 return rewriteSameTypeForwarding(op, if (condition) true_value else false_value, rewriter);
57 }
58
59 fn rewriteNotCanonicalization(op: *ir.Operation, rewriter: *rewrite.PatternRewriter) rewrite.PatternResult {
60 var declaration = ir.interfaces.effects.inspect(op.allocator, op) catch return .failure;
61 defer declaration.deinit(op.allocator);
62 if (!ir.interfaces.effects.repeatableExpression(declaration.facts)) return .failure;
63 if (op.getNumResults() != 1) return .failure;
64 const not = Dialect.NotOp{ .op = op };
65 const input = not.getInput();
66 if (folds.constantBoolFromValue(input)) |value| {
67 return rewriteBoolConstant(op, !value, rewriter);
68 }
69
70 const def_any = input.getDefiningOp() orelse return .failure;
71 const def_op: *ir.Operation = @ptrCast(@alignCast(def_any));
72 if (!std.mem.eql(u8, def_op.name.name, Dialect.NotOp.operation_name)) return .failure;
73 const inner_not = Dialect.NotOp{ .op = def_op };
74 return rewriteSameTypeForwarding(op, inner_not.getInput(), rewriter);
75 }
76
77 fn rewriteSameTypeForwarding(
78 op: *ir.Operation,
79 input: *ir.Value,
80 rewriter: *rewrite.PatternRewriter,
81 ) rewrite.PatternResult {
82 if (op.getNumResults() != 1) return .failure;
83 const result = op.getResult(0) orelse return .failure;
84 if (!result.type.eql(input.type)) return .failure;
85 rewriter.replaceOpWithValue(op, input) catch return .failure;
86 return .success;
87 }
88
89 fn rewriteBoolConstant(
90 op: *ir.Operation,
91 value: bool,
92 rewriter: *rewrite.PatternRewriter,
93 ) rewrite.PatternResult {
94 if (op.getNumResults() != 1) return .failure;
95 const result = op.getResult(0) orelse return .failure;
96 if (!folds.isBoolType(result.type)) return .failure;
97
98 rewriter.setInsertionPointBefore(op);
99 var state = ir.Operation.State.init(Dialect.ConstantOp.operation_name, op.location);
100 state.addTypes(&.{result.type});
101 const attr = Dialect.getBoolAttr(rewriter.ir_ctx, value) catch return .failure;
102 const uses_properties = state.setPropertiesAttrIfRegistered(rewriter.ir_ctx, attr) catch return .failure;
103 if (!uses_properties) state.addAttributes(&.{.{ .name = "value", .value = attr }});
104 _ = rewriter.replaceOpWithNewOp(op, state) catch return .failure;
105 return .success;
106 }
107
108 fn rewriteBoolNot(
109 op: *ir.Operation,
110 input: *ir.Value,
111 rewriter: *rewrite.PatternRewriter,
112 ) rewrite.PatternResult {
113 if (op.getNumResults() != 1) return .failure;
114 const result = op.getResult(0) orelse return .failure;
115 if (!result.type.eql(input.type) or !folds.isBoolType(result.type)) return .failure;
116
117 rewriter.setInsertionPointBefore(op);
118 var state = ir.Operation.State.init(Dialect.NotOp.operation_name, op.location);
119 state.addOperands(&.{input});
120 state.addTypes(&.{result.type});
121 _ = rewriter.replaceOpWithNewOp(op, state) catch return .failure;
122 return .success;
123 }
124 };
125 }