Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
@@ -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}>
29 changes: 29 additions & 0 deletions Test/Passes/RISCVCombines/mul_neg_one_flags.mlir
Original file line number Diff line number Diff line change
@@ -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) -> ()
20 changes: 20 additions & 0 deletions Test/Passes/RISCVCombines/select_i1.mlir
Original file line number Diff line number Diff line change
@@ -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) -> ()
4 changes: 1 addition & 3 deletions Veir/Passes/InstCombine.lean
Original file line number Diff line number Diff line change
Expand Up @@ -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]))

Expand Down
13 changes: 10 additions & 3 deletions Veir/Passes/Matching/LLVM/Basic.lean
Original file line number Diff line number Diff line change
Expand Up @@ -51,21 +51,28 @@ 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
let properties := op.getProperties! ctx Llvm.mlir__constant
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
Expand Down
7 changes: 4 additions & 3 deletions Veir/Passes/Matching/LLVM/Lemmas.lean
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
8 changes: 6 additions & 2 deletions Veir/Passes/RISCVCombines/Combine.lean
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down Expand Up @@ -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]))
Expand Down Expand Up @@ -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]))
Expand Down
8 changes: 4 additions & 4 deletions Veir/Passes/RISCVCombines/MIRCombinesVeir.lean
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down Expand Up @@ -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)
Expand Down
Loading