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 }