diff --git a/Test/Passes/InstructionSelection/RISCV64/binop_constant_decoding.mlir b/Test/Passes/InstructionSelection/RISCV64/binop_constant_decoding.mlir new file mode 100644 index 0000000000..e2f0a75780 --- /dev/null +++ b/Test/Passes/InstructionSelection/RISCV64/binop_constant_decoding.mlir @@ -0,0 +1,21 @@ +// RUN: veir-opt %s -p=isel-sdag-riscv64 | filecheck %s + +"builtin.module"() ({ + "func.func"() <{sym_name = "binop_constants", function_type = (i64, i32) -> (i64, i64, i64, i32)}> ({ + ^bb0(%x: i64, %y: i32): + %minusOne = "llvm.mlir.constant"() <{value = 255 : i8}> : () -> i64 + %one = "llvm.mlir.constant"() <{value = -1 : i1}> : () -> i64 + %sum = "llvm.add"(%x, %minusOne) : (i64, i64) -> i64 + %masked = "llvm.and"(%x, %minusOne) : (i64, i64) -> i64 + %shifted = "llvm.shl"(%x, %one) : (i64, i64) -> i64 + %truncatedOne = "llvm.mlir.constant"() <{value = 4294967297 : i64}> : () -> i32 + %sum32 = "llvm.add"(%y, %truncatedOne) : (i32, i32) -> i32 + "func.return"(%sum, %masked, %shifted, %sum32) : (i64, i64, i64, i32) -> () + }) : () -> () +}) : () -> () + +// CHECK-LABEL: func.func @binop_constants +// CHECK: "riscv.addi"({{.*}}) <{"value" = -1 : i64}> +// CHECK: "riscv.andi"({{.*}}) <{"value" = -1 : i64}> +// CHECK: "riscv.slli"({{.*}}) <{"value" = 1 : i64}> +// CHECK: "riscv.addiw"({{.*}}) <{"value" = 1 : i64}> diff --git a/Test/Passes/RISCVCombines/mul_neg_one_flags.mlir b/Test/Passes/RISCVCombines/mul_neg_one_flags.mlir new file mode 100644 index 0000000000..4535ed13c6 --- /dev/null +++ b/Test/Passes/RISCVCombines/mul_neg_one_flags.mlir @@ -0,0 +1,29 @@ +// RUN: veir-opt %s -p=riscv-combine | filecheck %s + +// 1 * 255 is defined unsigned, but 0 - 1 underflows. Multiplication by all-ones +// can become negation only if nuw is dropped. Signed negation preserves nsw. +// overflowFlags: 1 = nsw (no signed wrap), 2 = nuw (no unsigned wrap), 3 = both. +"builtin.module"() ({ + "func.func"() <{sym_name = "mul_neg_one", function_type = (i8) -> (i8, i8, i8, i8)}> ({ + ^bb0(%x: i8): + %unsigned = "llvm.mlir.constant"() <{value = 255 : i8}> : () -> i8 + %signed = "llvm.mlir.constant"() <{value = -1 : i8}> : () -> i8 + %a = "llvm.mul"(%x, %unsigned) <{overflowFlags = 2 : i32}> : (i8, i8) -> i8 + %b = "llvm.mul"(%x, %signed) <{overflowFlags = 2 : i32}> : (i8, i8) -> i8 + %c = "llvm.mul"(%x, %signed) <{overflowFlags = 1 : i32}> : (i8, i8) -> i8 + %d = "llvm.mul"(%x, %signed) <{overflowFlags = 3 : i32}> : (i8, i8) -> i8 + "func.return"(%a, %b, %c, %d) : (i8, i8, i8, i8) -> () + }) : () -> () +}) : () -> () + +// CHECK-LABEL: func.func @mul_neg_one +// CHECK-SAME: (%[[X:.*]]: i8) +// CHECK-NEXT: %[[Z0:.*]] = "llvm.mlir.constant"() <{"value" = 0 : i8}> +// CHECK-NEXT: %[[A:.*]] = "llvm.sub"(%[[Z0]], %[[X]]) : (i8, i8) -> i8 +// CHECK-NEXT: %[[Z1:.*]] = "llvm.mlir.constant"() <{"value" = 0 : i8}> +// CHECK-NEXT: %[[B:.*]] = "llvm.sub"(%[[Z1]], %[[X]]) : (i8, i8) -> i8 +// CHECK-NEXT: %[[Z2:.*]] = "llvm.mlir.constant"() <{"value" = 0 : i8}> +// CHECK-NEXT: %[[C:.*]] = "llvm.sub"(%[[Z2]], %[[X]]) <{"overflowFlags" = 1 : i32}> : (i8, i8) -> i8 +// CHECK-NEXT: %[[Z3:.*]] = "llvm.mlir.constant"() <{"value" = 0 : i8}> +// CHECK-NEXT: %[[D:.*]] = "llvm.sub"(%[[Z3]], %[[X]]) <{"overflowFlags" = 1 : i32}> : (i8, i8) -> i8 +// CHECK-NEXT: "func.return"(%[[A]], %[[B]], %[[C]], %[[D]]) : (i8, i8, i8, i8) -> () diff --git a/Test/Passes/RISCVCombines/select_i1.mlir b/Test/Passes/RISCVCombines/select_i1.mlir new file mode 100644 index 0000000000..30c5f2a260 --- /dev/null +++ b/Test/Passes/RISCVCombines/select_i1.mlir @@ -0,0 +1,20 @@ +// RUN: veir-opt %s -p=riscv-combine | filecheck %s + +// A boolean result needs the condition or its inverse, without an extension. +"builtin.module"() ({ + "func.func"() <{sym_name = "select_i1", function_type = (i1) -> (i1, i1)}> ({ + ^bb0(%cond: i1): + %one = "llvm.mlir.constant"() <{value = 1 : i1}> : () -> i1 + %minusOne = "llvm.mlir.constant"() <{value = -1 : i1}> : () -> i1 + %zero = "llvm.mlir.constant"() <{value = 0 : i1}> : () -> i1 + %a = "llvm.select"(%cond, %one, %zero) : (i1, i1, i1) -> i1 + %b = "llvm.select"(%cond, %zero, %minusOne) : (i1, i1, i1) -> i1 + "func.return"(%a, %b) : (i1, i1) -> () + }) : () -> () +}) : () -> () + +// CHECK-LABEL: func.func @select_i1 +// CHECK-SAME: (%[[COND:.*]]: i1) +// CHECK-NEXT: %[[ONE:.*]] = "llvm.mlir.constant"() <{"value" = -1 : i1}> : () -> i1 +// CHECK-NEXT: %[[NOT:.*]] = "llvm.xor"(%[[COND]], %[[ONE]]) : (i1, i1) -> i1 +// CHECK-NEXT: "func.return"(%[[COND]], %[[NOT]]) : (i1, i1) -> () diff --git a/Veir/Passes/InstCombine.lean b/Veir/Passes/InstCombine.lean index 0603c9f630..3e911ab53f 100644 --- a/Veir/Passes/InstCombine.lean +++ b/Veir/Passes/InstCombine.lean @@ -57,9 +57,7 @@ def mulIOneToX_local (ctx : WfIRContext OpCode) (op : OperationPtr) : Option (WfIRContext OpCode × Option (Array OperationPtr × Array ValuePtr)) := do let some (lhs, rhs, _) := matchMuli op ctx.raw | return (ctx, none) - let some cst := matchConstantIntVal rhs ctx.raw - | return (ctx, none) - if cst ≠ 1 then + if !isConstantOne rhs ctx.raw then return (ctx, none) some (ctx, some (#[], #[lhs])) diff --git a/Veir/Passes/Matching/LLVM/Basic.lean b/Veir/Passes/Matching/LLVM/Basic.lean index d5da98b40a..f52b890dea 100644 --- a/Veir/Passes/Matching/LLVM/Basic.lean +++ b/Veir/Passes/Matching/LLVM/Basic.lean @@ -51,6 +51,7 @@ def matchXori (op : OperationPtr) (ctx : IRContext OpCode) : let (op, _) ← matchOp op ctx (Llvm.xor) 2 return (op[0]!, op[1]!) +/-- Match the raw integer attribute; use `matchConstantIntVal` for the result's value. -/ def matchConstantIntOp (op : OperationPtr) (ctx : IRContext OpCode) : Option IntegerAttr := do let Llvm.mlir__constant := toDialect? Llvm (op.getOpType! ctx) | none @@ -58,14 +59,20 @@ def matchConstantIntOp (op : OperationPtr) (ctx : IRContext OpCode) : let .integer intAttr := properties.value | none return intAttr -/-- Match the raw integer attribute value of an LLVM constant, without adjusting -it to the attribute or result width. -/ def matchConstantIntVal (val : ValuePtr) (ctx : IRContext OpCode) : Option Int := do let .opResult opResultPtr := val | none let op := opResultPtr.op let attr ← matchConstantIntOp op ctx - return attr.value + let .integerType type := (val.getType! ctx).val | none + return (BitVec.ofInt type.bitwidth (decodeLLVMIntegerConstant attr)).toInt + +/-- Recognize the one bit pattern, including i1 true whose signed value is -1. -/ +def isConstantOne (val : ValuePtr) (ctx : IRContext OpCode) : Bool := + match matchConstantIntVal val ctx, (val.getType! ctx).val with + | some value, .integerType type => + (BitVec.ofInt type.bitwidth value).toNat == 1 + | _, _ => false /-- Match a constant integer with value zero, returning `val` itself. -/ def matchConstantZero (val : ValuePtr) (ctx : IRContext OpCode) : Option ValuePtr := do diff --git a/Veir/Passes/Matching/LLVM/Lemmas.lean b/Veir/Passes/Matching/LLVM/Lemmas.lean index 0a59020499..7818512d21 100644 --- a/Veir/Passes/Matching/LLVM/Lemmas.lean +++ b/Veir/Passes/Matching/LLVM/Lemmas.lean @@ -110,9 +110,10 @@ theorem matchConstantIntOp_implies {op : OperationPtr} {ctx : IRContext OpCode} /-- What matching a constant integer value (via `matchConstantIntVal`) syntactically guarantees. -/ theorem matchConstantIntVal_implies {val : ValuePtr} {ctx : IRContext OpCode} {value} : matchConstantIntVal val ctx = some value → - ∃ opResultPtr intAttr, val = .opResult opResultPtr ∧ - matchConstantIntOp opResultPtr.op ctx = some intAttr ∧ - intAttr.value = value := by + ∃ opResultPtr attr type, val = .opResult opResultPtr ∧ + matchConstantIntOp opResultPtr.op ctx = some attr ∧ + (val.getType! ctx).val = .integerType type ∧ + value = (BitVec.ofInt type.bitwidth (decodeLLVMIntegerConstant attr)).toInt := by intro hmatch simp only [matchConstantIntVal, bind, Option.bind, pure] at hmatch grind diff --git a/Veir/Passes/RISCVCombines/Combine.lean b/Veir/Passes/RISCVCombines/Combine.lean index 09b9076f9e..88b541b692 100644 --- a/Veir/Passes/RISCVCombines/Combine.lean +++ b/Veir/Passes/RISCVCombines/Combine.lean @@ -63,8 +63,7 @@ def select_same_val_self (rewriter : PatternRewriter OpCode) (op : OperationPtr) def select_constant_cmp_true_local (ctx : WfIRContext OpCode) (op : OperationPtr) : Option (WfIRContext OpCode × Option (Array OperationPtr × Array ValuePtr)) := do let some (cond, tval, _fval) := matchSelect op ctx.raw | return (ctx, none) - let some cst := matchConstantIntVal cond ctx.raw | return (ctx, none) - if cst ≠ 1 then return (ctx, none) + if !isConstantOne cond ctx.raw then return (ctx, none) some (ctx, some (#[], #[tval])) def select_constant_cmp_true (rewriter : PatternRewriter OpCode) (op : OperationPtr) @@ -958,6 +957,9 @@ def select_neg1_0_local (ctx : WfIRContext OpCode) (op : OperationPtr) : if ct ≠ -1 then return (ctx, none) let some cf := matchConstantIntVal fv ctx.raw | return (ctx, none) if cf ≠ 0 then return (ctx, none) + -- At i1, -1 is true and no extension is needed. + if (op.getResult 0 : ValuePtr).getType! ctx.raw = IntegerType.mk 1 then + return (ctx, some (#[], #[cond])) let (ctx, newOp) ← WfRewriter.createOp! ctx Llvm.sext #[(op.getResult 0 : ValuePtr).getType! ctx.raw] #[cond] #[] #[] () none some (ctx, some (#[newOp], #[newOp.getResult 0])) @@ -1002,6 +1004,8 @@ def select_0_neg1_local (ctx : WfIRContext OpCode) (op : OperationPtr) : #[] #[] m1 none let (ctx, ncond) ← WfRewriter.createOp! ctx Llvm.xor #[cond.getType! ctx.raw] #[cond, (c1.getResult 0)] #[] #[] () none + if (op.getResult 0 : ValuePtr).getType! ctx.raw = IntegerType.mk 1 then + return (ctx, some (#[c1, ncond], #[ncond.getResult 0])) let (ctx, newOp) ← WfRewriter.createOp! ctx Llvm.sext #[(op.getResult 0 : ValuePtr).getType! ctx.raw] #[(ncond.getResult 0)] #[] #[] () none some (ctx, some (#[c1, ncond, newOp], #[newOp.getResult 0])) diff --git a/Veir/Passes/RISCVCombines/MIRCombinesVeir.lean b/Veir/Passes/RISCVCombines/MIRCombinesVeir.lean index f566a33a2b..5f622678b9 100644 --- a/Veir/Passes/RISCVCombines/MIRCombinesVeir.lean +++ b/Veir/Passes/RISCVCombines/MIRCombinesVeir.lean @@ -107,8 +107,7 @@ def right_identity_zero_6 (rewriter : PatternRewriter OpCode) (op : OperationPtr def right_identity_one_int_local (ctx : WfIRContext OpCode) (op : OperationPtr) : Option (WfIRContext OpCode × Option (Array OperationPtr × Array ValuePtr)) := do let some (x, rhs, _props) := matchMul op ctx.raw | return (ctx, none) - let some cst := matchConstantIntVal rhs ctx.raw | return (ctx, none) - if cst ≠ 1 then return (ctx, none) + if !isConstantOne rhs ctx.raw then return (ctx, none) some (ctx, some (#[], #[x])) def right_identity_one_int (rewriter : PatternRewriter OpCode) (op : OperationPtr) @@ -182,15 +181,16 @@ def binop_right_to_zero (rewriter : PatternRewriter OpCode) (op : OperationPtr) def mul_by_neg_one_local (ctx : WfIRContext OpCode) (op : OperationPtr) : Option (WfIRContext OpCode × Option (Array OperationPtr × Array ValuePtr)) := do - let some (x, rhs, _props) := matchMul op ctx.raw | return (ctx, none) + let some (x, rhs, props) := matchMul op ctx.raw | return (ctx, none) let some cst := matchConstantIntVal rhs ctx.raw | return (ctx, none) if cst ≠ -1 then return (ctx, none) let .integerType ctype := (x.getType! ctx.raw).val | return (ctx, none) let cstOpProp := LLVMConstantProperties.mk (.integer (IntegerAttr.mk (0) ctype)) let (ctx, cstOp) ← WfRewriter.createOp! ctx Llvm.mlir__constant #[x.getType! ctx.raw] #[] #[] #[] cstOpProp none + -- Multiplying 1 by all-ones does not overflow unsigned, but 0 - 1 does. let (ctx, newOp) ← WfRewriter.createOp! ctx Llvm.sub #[x.getType! ctx.raw] #[(cstOp.getResult 0), x] - #[] #[] _props none + #[] #[] { props with nuw := false } none some (ctx, some (#[cstOp, newOp], #[newOp.getResult 0])) def mul_by_neg_one (rewriter : PatternRewriter OpCode) (op : OperationPtr)