27#include "llvm/ADT/APFloat.h"
28#include "llvm/ADT/APInt.h"
29#include "llvm/ADT/APSInt.h"
30#include "llvm/ADT/FloatingPointMode.h"
31#include "llvm/ADT/STLExtras.h"
32#include "llvm/ADT/SmallVector.h"
33#include "llvm/ADT/TypeSwitch.h"
40 llvm::RoundingMode::NearestTiesToEven;
49 function_ref<APInt(
const APInt &,
const APInt &)> binFn) {
50 const APInt &lhsVal = llvm::cast<IntegerAttr>(lhs).getValue();
51 const APInt &rhsVal = llvm::cast<IntegerAttr>(rhs).getValue();
52 APInt value = binFn(lhsVal, rhsVal);
53 return IntegerAttr::get(res.getType(), value);
87static IntegerOverflowFlagsAttr
89 IntegerOverflowFlagsAttr val2) {
90 return IntegerOverflowFlagsAttr::get(val1.getContext(),
91 val1.getValue() & val2.getValue());
97 case arith::CmpIPredicate::eq:
98 return arith::CmpIPredicate::ne;
99 case arith::CmpIPredicate::ne:
100 return arith::CmpIPredicate::eq;
101 case arith::CmpIPredicate::slt:
102 return arith::CmpIPredicate::sge;
103 case arith::CmpIPredicate::sle:
104 return arith::CmpIPredicate::sgt;
105 case arith::CmpIPredicate::sgt:
106 return arith::CmpIPredicate::sle;
107 case arith::CmpIPredicate::sge:
108 return arith::CmpIPredicate::slt;
109 case arith::CmpIPredicate::ult:
110 return arith::CmpIPredicate::uge;
111 case arith::CmpIPredicate::ule:
112 return arith::CmpIPredicate::ugt;
113 case arith::CmpIPredicate::ugt:
114 return arith::CmpIPredicate::ule;
115 case arith::CmpIPredicate::uge:
116 return arith::CmpIPredicate::ult;
118 llvm_unreachable(
"unknown cmpi predicate kind");
127static llvm::RoundingMode
131 switch (*roundingMode) {
132 case RoundingMode::downward:
133 return llvm::RoundingMode::TowardNegative;
134 case RoundingMode::to_nearest_away:
135 return llvm::RoundingMode::NearestTiesToAway;
136 case RoundingMode::to_nearest_even:
137 return llvm::RoundingMode::NearestTiesToEven;
138 case RoundingMode::toward_zero:
139 return llvm::RoundingMode::TowardZero;
140 case RoundingMode::upward:
141 return llvm::RoundingMode::TowardPositive;
143 llvm_unreachable(
"Unhandled rounding mode");
147 return arith::CmpIPredicateAttr::get(pred.getContext(),
173 ShapedType shapedType = dyn_cast_or_null<ShapedType>(type);
177 if (!shapedType.hasStaticShape())
187 ShapedType shapedType = dyn_cast<ShapedType>(type);
190 if (!shapedType.hasStaticShape())
200#include "ArithCanonicalization.inc"
209 auto i1Type = IntegerType::get(type.
getContext(), 1);
210 if (
auto shapedType = dyn_cast<ShapedType>(type))
211 return shapedType.cloneWith(std::nullopt, i1Type);
212 if (llvm::isa<UnrankedTensorType>(type))
213 return UnrankedTensorType::get(i1Type);
221void arith::ConstantOp::getAsmResultNames(
224 if (
auto intCst = dyn_cast<IntegerAttr>(getValue())) {
225 auto intType = dyn_cast<IntegerType>(type);
228 if (intType && intType.getWidth() == 1)
229 return setNameFn(getResult(), (intCst.getInt() ?
"true" :
"false"));
232 SmallString<32> specialNameBuffer;
233 llvm::raw_svector_ostream specialName(specialNameBuffer);
234 specialName <<
'c' << intCst.getValue();
236 specialName <<
'_' << type;
237 setNameFn(getResult(), specialName.str());
239 setNameFn(getResult(),
"cst");
245LogicalResult arith::ConstantOp::verify() {
249 intType && !intType.isSignless())
250 return emitOpError(
"integer return type must be signless");
252 if (!llvm::isa<IntegerAttr, FloatAttr, ElementsAttr>(getValue())) {
254 "value must be an integer, float, or elements attribute");
260 if (isa<ScalableVectorType>(type) && !isa<SplatElementsAttr>(getValue()))
262 "initializing scalable vectors with elements attribute is not supported"
263 " unless it's a vector splat");
267bool arith::ConstantOp::isBuildableWith(Attribute value, Type type) {
269 auto typedAttr = dyn_cast<TypedAttr>(value);
270 if (!typedAttr || typedAttr.getType() != type)
274 if (!intType.isSignless())
278 return llvm::isa<IntegerAttr, FloatAttr, ElementsAttr>(value);
281ConstantOp arith::ConstantOp::materialize(OpBuilder &builder, Attribute value,
282 Type type, Location loc) {
283 if (isBuildableWith(value, type))
284 return arith::ConstantOp::create(builder, loc, cast<TypedAttr>(value));
288OpFoldResult arith::ConstantOp::fold(FoldAdaptor adaptor) {
return getValue(); }
293 arith::ConstantOp::build(builder,
result, type,
303 auto result = dyn_cast<ConstantIntOp>(builder.
create(state));
304 assert(
result &&
"builder didn't return the right type");
316 arith::ConstantOp::build(builder,
result, type,
325 auto result = dyn_cast<ConstantIntOp>(builder.
create(state));
326 assert(
result &&
"builder didn't return the right type");
337 arith::ConstantOp::build(builder,
result, type,
343 const APInt &
value) {
346 auto result = dyn_cast<ConstantIntOp>(builder.
create(state));
347 assert(
result &&
"builder didn't return the right type");
353 const APInt &
value) {
358 if (
auto constOp = dyn_cast_or_null<arith::ConstantOp>(op))
359 return constOp.getType().isSignlessInteger();
364 FloatType type,
const APFloat &
value) {
365 arith::ConstantOp::build(builder,
result, type,
372 const APFloat &
value) {
375 auto result = dyn_cast<ConstantFloatOp>(builder.
create(state));
376 assert(
result &&
"builder didn't return the right type");
382 const APFloat &
value) {
387 if (
auto constOp = dyn_cast_or_null<arith::ConstantOp>(op))
388 return llvm::isa<FloatType>(constOp.getType());
403 auto result = dyn_cast<ConstantIndexOp>(builder.
create(state));
404 assert(
result &&
"builder didn't return the right type");
414 if (
auto constOp = dyn_cast_or_null<arith::ConstantOp>(op))
415 return constOp.getType().isIndex();
423 "type doesn't have a zero representation");
425 assert(zeroAttr &&
"unsupported type for zero attribute");
426 return arith::ConstantOp::create(builder, loc, zeroAttr);
439 if (
auto sub = getLhs().getDefiningOp<SubIOp>())
440 if (getRhs() == sub.getRhs())
444 if (
auto sub = getRhs().getDefiningOp<SubIOp>())
445 if (getLhs() == sub.getRhs())
449 adaptor.getOperands(),
450 [](APInt a,
const APInt &
b) { return std::move(a) + b; });
455 patterns.
add<AddIAddConstant, AddISubConstantRHS, AddISubConstantLHS,
456 AddIMulNegativeOneRhs, AddIMulNegativeOneLhs>(context);
463std::optional<SmallVector<int64_t, 4>>
464arith::AddUIExtendedOp::getShapeForUnroll() {
465 if (
auto vt = dyn_cast<VectorType>(
getType(0)))
466 return llvm::to_vector<4>(vt.getShape());
473 return sum.ult(operand) ? APInt::getAllOnes(1) : APInt::getZero(1);
477arith::AddUIExtendedOp::fold(FoldAdaptor adaptor,
478 SmallVectorImpl<OpFoldResult> &results) {
479 Type overflowTy = getOverflow().getType();
485 results.push_back(getLhs());
486 results.push_back(falseValue);
495 adaptor.getOperands(),
496 [](APInt a,
const APInt &
b) { return std::move(a) + b; })) {
499 results.push_back(sumAttr);
500 results.push_back(sumAttr);
504 ArrayRef({sumAttr, adaptor.getLhs()}),
510 results.push_back(sumAttr);
511 results.push_back(overflowAttr);
518void arith::AddUIExtendedOp::getCanonicalizationPatterns(
519 RewritePatternSet &patterns, MLIRContext *context) {
520 patterns.
add<AddUIExtendedToAddI>(context);
527std::optional<SmallVector<int64_t, 4>>
528arith::SubUIExtendedOp::getShapeForUnroll() {
529 if (
auto vt = dyn_cast<VectorType>(
getType(0)))
530 return llvm::to_vector<4>(vt.getShape());
537 return lhs.ult(rhs) ? APInt::getAllOnes(1) : APInt::getZero(1);
541arith::SubUIExtendedOp::fold(FoldAdaptor adaptor,
542 SmallVectorImpl<OpFoldResult> &results) {
543 Type borrowTy = getBorrow().getType();
549 results.push_back(getLhs());
550 results.push_back(falseValue);
555 if (getLhs() == getRhs()) {
558 auto shapedType = dyn_cast<ShapedType>(getDiff().
getType());
559 if (shapedType && !shapedType.hasStaticShape())
567 results.push_back(zeroDiff);
568 results.push_back(falseValue);
574 adaptor.getOperands(),
575 [](APInt a,
const APInt &
b) { return std::move(a) - b; })) {
578 results.push_back(diffAttr);
579 results.push_back(diffAttr);
583 adaptor.getOperands(),
589 results.push_back(diffAttr);
590 results.push_back(borrowAttr);
597void arith::SubUIExtendedOp::getCanonicalizationPatterns(
598 RewritePatternSet &patterns, MLIRContext *context) {
599 patterns.
add<SubUIExtendedToSubI>(context);
606OpFoldResult arith::SubIOp::fold(FoldAdaptor adaptor) {
608 if (getOperand(0) == getOperand(1)) {
609 auto shapedType = dyn_cast<ShapedType>(
getType());
611 if (!shapedType || shapedType.hasStaticShape())
618 if (
auto add = getLhs().getDefiningOp<AddIOp>()) {
620 if (getRhs() ==
add.getRhs())
623 if (getRhs() ==
add.getLhs())
628 if (
auto sub = getRhs().getDefiningOp<SubIOp>())
629 if (getLhs() == sub.getLhs())
633 adaptor.getOperands(),
634 [](APInt a,
const APInt &
b) { return std::move(a) - b; });
637void arith::SubIOp::getCanonicalizationPatterns(RewritePatternSet &patterns,
638 MLIRContext *context) {
639 patterns.
add<SubIRHSAddConstant, SubILHSAddConstant, SubIRHSSubConstantRHS,
640 SubIRHSSubConstantLHS, SubILHSSubConstantRHS,
641 SubILHSSubConstantLHS, SubISubILHSRHSLHS>(context);
648OpFoldResult arith::MulIOp::fold(FoldAdaptor adaptor) {
659 adaptor.getOperands(),
660 [](
const APInt &a,
const APInt &
b) { return a * b; });
663void arith::MulIOp::getAsmResultNames(
665 if (!isa<IndexType>(
getType()))
670 auto isVscale = [](Operation *op) {
671 return op && op->getName().getStringRef() ==
"vector.vscale";
674 IntegerAttr baseValue;
675 auto isVscaleExpr = [&](Value a, Value
b) {
677 isVscale(
b.getDefiningOp());
680 if (!isVscaleExpr(getLhs(), getRhs()) && !isVscaleExpr(getRhs(), getLhs()))
684 SmallString<32> specialNameBuffer;
685 llvm::raw_svector_ostream specialName(specialNameBuffer);
686 specialName <<
'c' << baseValue.getInt() <<
"_vscale";
687 setNameFn(getResult(), specialName.str());
690void arith::MulIOp::getCanonicalizationPatterns(RewritePatternSet &patterns,
691 MLIRContext *context) {
692 patterns.
add<MulIMulIConstant>(context);
699std::optional<SmallVector<int64_t, 4>>
700arith::MulSIExtendedOp::getShapeForUnroll() {
701 if (
auto vt = dyn_cast<VectorType>(
getType(0)))
702 return llvm::to_vector<4>(vt.getShape());
707arith::MulSIExtendedOp::fold(FoldAdaptor adaptor,
708 SmallVectorImpl<OpFoldResult> &results) {
711 Attribute zero = adaptor.getRhs();
712 results.push_back(zero);
713 results.push_back(zero);
719 adaptor.getOperands(),
720 [](
const APInt &a,
const APInt &
b) { return a * b; })) {
723 llvm::APIntOps::mulhs);
724 assert(highAttr &&
"Unexpected constant-folding failure");
726 results.push_back(lowAttr);
727 results.push_back(highAttr);
734void arith::MulSIExtendedOp::getCanonicalizationPatterns(
735 RewritePatternSet &patterns, MLIRContext *context) {
736 patterns.
add<MulSIExtendedToMulI, MulSIExtendedRHSOne>(context);
743std::optional<SmallVector<int64_t, 4>>
744arith::MulUIExtendedOp::getShapeForUnroll() {
745 if (
auto vt = dyn_cast<VectorType>(
getType(0)))
746 return llvm::to_vector<4>(vt.getShape());
751arith::MulUIExtendedOp::fold(FoldAdaptor adaptor,
752 SmallVectorImpl<OpFoldResult> &results) {
755 Attribute zero = adaptor.getRhs();
756 results.push_back(zero);
757 results.push_back(zero);
765 results.push_back(getLhs());
766 results.push_back(zero);
772 adaptor.getOperands(),
773 [](
const APInt &a,
const APInt &
b) { return a * b; })) {
776 llvm::APIntOps::mulhu);
777 assert(highAttr &&
"Unexpected constant-folding failure");
779 results.push_back(lowAttr);
780 results.push_back(highAttr);
787void arith::MulUIExtendedOp::getCanonicalizationPatterns(
788 RewritePatternSet &patterns, MLIRContext *context) {
789 patterns.
add<MulUIExtendedToMulI>(context);
798 arith::IntegerOverflowFlags ovfFlags) {
799 auto mul = lhs.getDefiningOp<mlir::arith::MulIOp>();
800 if (!
mul || !bitEnumContainsAll(
mul.getOverflowFlags(), ovfFlags))
803 if (
mul.getLhs() == rhs)
806 if (
mul.getRhs() == rhs)
812OpFoldResult arith::DivUIOp::fold(FoldAdaptor adaptor) {
826 if (getLhs() == getRhs())
830 if (Value val =
foldDivMul(getLhs(), getRhs(), IntegerOverflowFlags::nuw))
836 [&](APInt a,
const APInt &
b) {
844 return div0 ? Attribute() :
result;
864OpFoldResult arith::DivSIOp::fold(FoldAdaptor adaptor) {
878 if (getLhs() == getRhs())
882 if (Value val =
foldDivMul(getLhs(), getRhs(), IntegerOverflowFlags::nsw))
886 bool overflowOrDiv0 =
false;
888 adaptor.getOperands(), [&](APInt a,
const APInt &
b) {
889 if (overflowOrDiv0 || !b) {
890 overflowOrDiv0 = true;
893 return a.sdiv_ov(
b, overflowOrDiv0);
896 return overflowOrDiv0 ? Attribute() :
result;
920OpFoldResult arith::CeilDivUIOp::fold(FoldAdaptor adaptor) {
934 if (getLhs() == getRhs())
937 bool overflowOrDiv0 =
false;
939 adaptor.getOperands(), [&](APInt a,
const APInt &
b) {
940 if (overflowOrDiv0 || !b) {
941 overflowOrDiv0 = true;
944 APInt quotient = a.udiv(
b);
947 APInt one(a.getBitWidth(), 1,
true);
948 return quotient.uadd_ov(one, overflowOrDiv0);
951 return overflowOrDiv0 ? Attribute() :
result;
962OpFoldResult arith::CeilDivSIOp::fold(FoldAdaptor adaptor) {
976 if (getLhs() == getRhs())
980 bool overflowOrDiv0 =
false;
982 adaptor.getOperands(), [&](
const APInt &a,
const APInt &
b) {
983 if (overflowOrDiv0 || !b) {
984 overflowOrDiv0 = true;
994 bool overflowDiv =
false;
995 APInt quotient = a.sdiv_ov(
b, overflowDiv);
999 overflowOrDiv0 =
true;
1002 if (a.isNegative() !=
b.isNegative() || quotient *
b == a)
1008 APInt one(a.getBitWidth(), 1,
true);
1009 return quotient.sadd_ov(one, overflowOrDiv0);
1012 return overflowOrDiv0 ? Attribute() :
result;
1023OpFoldResult arith::FloorDivSIOp::fold(FoldAdaptor adaptor) {
1037 if (getLhs() == getRhs())
1041 bool overflowOrDiv =
false;
1043 adaptor.getOperands(), [&](APInt a,
const APInt &
b) {
1045 overflowOrDiv = true;
1048 return a.sfloordiv_ov(
b, overflowOrDiv);
1051 return overflowOrDiv ? Attribute() :
result;
1058OpFoldResult arith::RemUIOp::fold(FoldAdaptor adaptor) {
1075 [&](APInt a,
const APInt &
b) {
1076 if (div0 || b.isZero()) {
1083 return div0 ? Attribute() :
result;
1094OpFoldResult arith::RemSIOp::fold(FoldAdaptor adaptor) {
1111 [&](APInt a,
const APInt &
b) {
1112 if (div0 || b.isZero()) {
1119 return div0 ? Attribute() :
result;
1137template <
typename OpTy>
1139 for (
bool reversePrev : {
false,
true}) {
1140 auto prev = (reversePrev ? op.getRhs() : op.getLhs())
1141 .template getDefiningOp<OpTy>();
1145 Value other = (reversePrev ? op.getLhs() : op.getRhs());
1146 if (other != prev.getLhs() && other != prev.getRhs())
1149 return prev.getResult();
1154OpFoldResult arith::AndIOp::fold(FoldAdaptor adaptor) {
1161 intValue.isAllOnes())
1166 intValue.isAllOnes())
1171 intValue.isAllOnes())
1179 adaptor.getOperands(),
1180 [](APInt a,
const APInt &
b) { return std::move(a) & b; });
1187OpFoldResult arith::OrIOp::fold(FoldAdaptor adaptor) {
1190 if (rhsVal.isZero())
1193 if (rhsVal.isAllOnes())
1194 return adaptor.getRhs();
1201 intValue.isAllOnes())
1202 return getRhs().getDefiningOp<XOrIOp>().getRhs();
1206 intValue.isAllOnes())
1207 return getLhs().getDefiningOp<XOrIOp>().getRhs();
1214 adaptor.getOperands(),
1215 [](APInt a,
const APInt &
b) { return std::move(a) | b; });
1222OpFoldResult arith::XOrIOp::fold(FoldAdaptor adaptor) {
1227 if (getLhs() == getRhs()) {
1230 auto shapedType = dyn_cast<ShapedType>(
getType());
1231 if (!shapedType || shapedType.hasStaticShape())
1236 if (arith::XOrIOp prev = getLhs().getDefiningOp<arith::XOrIOp>()) {
1237 if (prev.getRhs() == getRhs())
1238 return prev.getLhs();
1239 if (prev.getLhs() == getRhs())
1240 return prev.getRhs();
1244 if (arith::XOrIOp prev = getRhs().getDefiningOp<arith::XOrIOp>()) {
1245 if (prev.getRhs() == getLhs())
1246 return prev.getLhs();
1247 if (prev.getLhs() == getLhs())
1248 return prev.getRhs();
1252 adaptor.getOperands(),
1253 [](APInt a,
const APInt &
b) { return std::move(a) ^ b; });
1256void arith::XOrIOp::getCanonicalizationPatterns(RewritePatternSet &patterns,
1257 MLIRContext *context) {
1258 patterns.
add<XOrIXOrIConstant, XOrINotCmpI, XOrIOfExtUI, XOrIOfExtSI>(
1266OpFoldResult arith::NegFOp::fold(FoldAdaptor adaptor) {
1268 if (
auto op = this->getOperand().getDefiningOp<arith::NegFOp>())
1269 return op.getOperand();
1271 [](
const APFloat &a) { return -a; });
1278OpFoldResult arith::FlushDenormalsOp::fold(FoldAdaptor adaptor) {
1284 if (
auto op = this->getOperand().getDefiningOp<arith::FlushDenormalsOp>())
1285 return op.getResult();
1289 adaptor.getOperands(), [](
const APFloat &a) {
1291 return APFloat::getZero(a.getSemantics(), a.isNegative());
1300OpFoldResult arith::AddFOp::fold(FoldAdaptor adaptor) {
1305 bitEnumContainsAll(adaptor.getFastmath(), FastMathFlags::nsz))
1308 auto rm = getRoundingmode();
1310 adaptor.getOperands(), [rm](
const APFloat &a,
const APFloat &
b) {
1312 result.add(b, convertArithRoundingModeToLLVMIR(rm));
1317void arith::AddFOp::getCanonicalizationPatterns(RewritePatternSet &patterns,
1318 MLIRContext *context) {
1319 patterns.
add<AddFOfNegFLhs, AddFOfNegFRhs>(context);
1326OpFoldResult arith::SubFOp::fold(FoldAdaptor adaptor) {
1331 bitEnumContainsAll(adaptor.getFastmath(), FastMathFlags::nsz))
1334 auto rm = getRoundingmode();
1336 adaptor.getOperands(), [rm](
const APFloat &a,
const APFloat &
b) {
1338 result.subtract(b, convertArithRoundingModeToLLVMIR(rm));
1343void arith::SubFOp::getCanonicalizationPatterns(RewritePatternSet &patterns,
1344 MLIRContext *context) {
1345 patterns.
add<SubFOfNegZero>(context);
1358template <
typename TruncOp,
typename ExtOp,
typename ExtremumOp>
1359struct NarrowExtremum final : OpRewritePattern<TruncOp> {
1360 using OpRewritePattern<TruncOp>::OpRewritePattern;
1362 LogicalResult matchAndRewrite(TruncOp truncOp,
1363 PatternRewriter &rewriter)
const override {
1364 auto extremumOp = truncOp.getIn().template getDefiningOp<ExtremumOp>();
1365 if (!extremumOp || !extremumOp->hasOneUse())
1368 auto lhsExt = extremumOp.getLhs().template getDefiningOp<ExtOp>();
1369 auto rhsExt = extremumOp.getRhs().template getDefiningOp<ExtOp>();
1370 if (!lhsExt || !rhsExt)
1373 Value
lhs = lhsExt.getIn();
1374 Value
rhs = rhsExt.getIn();
1375 Type narrowType = truncOp.getType();
1376 if (
lhs.getType() != narrowType ||
rhs.getType() != narrowType)
1385 if (
auto narrowFloatType =
1387 auto wideFloatType =
1392 const llvm::fltSemantics &narrowSemantics =
1393 narrowFloatType.getFloatSemantics();
1394 const llvm::fltSemantics &wideSemantics =
1395 wideFloatType.getFloatSemantics();
1396 bool ignoreNaNs =
false;
1397 if constexpr (std::is_same_v<TruncOp, TruncFOp>)
1399 bitEnumContainsAll(extremumOp.getFastmath(), FastMathFlags::nnan);
1400 if (!llvm::APFloatBase::isLosslesslyConvertibleTo(
1401 narrowSemantics, wideSemantics, ignoreNaNs))
1407 extremumOp.getProperties(),
1408 extremumOp->getDiscardableAttrDictionary().getValue());
1419OpFoldResult arith::MaximumFOp::fold(FoldAdaptor adaptor) {
1421 if (getLhs() == getRhs())
1435OpFoldResult arith::MaxNumFOp::fold(FoldAdaptor adaptor) {
1437 if (getLhs() == getRhs())
1451OpFoldResult arith::MaximumNumFOp::fold(FoldAdaptor adaptor) {
1453 if (getLhs() == getRhs())
1467OpFoldResult MaxSIOp::fold(FoldAdaptor adaptor) {
1469 if (getLhs() == getRhs())
1475 if (intValue.isMaxSignedValue())
1478 if (intValue.isMinSignedValue())
1483 llvm::APIntOps::smax);
1490OpFoldResult MaxUIOp::fold(FoldAdaptor adaptor) {
1492 if (getLhs() == getRhs())
1498 if (intValue.isMaxValue())
1501 if (intValue.isMinValue())
1506 llvm::APIntOps::umax);
1513OpFoldResult arith::MinimumFOp::fold(FoldAdaptor adaptor) {
1515 if (getLhs() == getRhs())
1529OpFoldResult arith::MinNumFOp::fold(FoldAdaptor adaptor) {
1531 if (getLhs() == getRhs())
1545OpFoldResult arith::MinimumNumFOp::fold(FoldAdaptor adaptor) {
1547 if (getLhs() == getRhs())
1561OpFoldResult MinSIOp::fold(FoldAdaptor adaptor) {
1563 if (getLhs() == getRhs())
1569 if (intValue.isMinSignedValue())
1572 if (intValue.isMaxSignedValue())
1577 llvm::APIntOps::smin);
1584OpFoldResult MinUIOp::fold(FoldAdaptor adaptor) {
1586 if (getLhs() == getRhs())
1592 if (intValue.isMinValue())
1595 if (intValue.isMaxValue())
1600 llvm::APIntOps::umin);
1607OpFoldResult arith::MulFOp::fold(FoldAdaptor adaptor) {
1612 if (arith::bitEnumContainsAll(getFastmath(), arith::FastMathFlags::nnan |
1613 arith::FastMathFlags::nsz)) {
1619 auto rm = getRoundingmode();
1621 adaptor.getOperands(), [rm](
const APFloat &a,
const APFloat &
b) {
1623 result.multiply(b, convertArithRoundingModeToLLVMIR(rm));
1628void arith::MulFOp::getCanonicalizationPatterns(RewritePatternSet &patterns,
1629 MLIRContext *context) {
1630 patterns.
add<MulFOfNegF>(context);
1637OpFoldResult arith::DivFOp::fold(FoldAdaptor adaptor) {
1642 auto rm = getRoundingmode();
1644 adaptor.getOperands(), [rm](
const APFloat &a,
const APFloat &
b) {
1646 result.divide(b, convertArithRoundingModeToLLVMIR(rm));
1651void arith::DivFOp::getCanonicalizationPatterns(RewritePatternSet &patterns,
1652 MLIRContext *context) {
1653 patterns.
add<DivFOfNegF>(context);
1660OpFoldResult arith::RemFOp::fold(FoldAdaptor adaptor) {
1662 [](
const APFloat &a,
const APFloat &
b) {
1667 (void)result.mod(b);
1676template <
typename... Types>
1682template <
typename... ShapedTypes,
typename... ElementTypes>
1685 if (llvm::isa<ShapedType>(type) && !llvm::isa<ShapedTypes...>(type))
1689 if (!llvm::isa<ElementTypes...>(underlyingType))
1692 return underlyingType;
1696template <
typename... ElementTypes>
1703template <
typename... ElementTypes>
1712 auto rankedTensorA = dyn_cast<RankedTensorType>(typeA);
1713 auto rankedTensorB = dyn_cast<RankedTensorType>(typeB);
1714 if (!rankedTensorA || !rankedTensorB)
1716 return rankedTensorA.getEncoding() == rankedTensorB.getEncoding();
1720 if (inputs.size() != 1 || outputs.size() != 1)
1732template <
typename ValType,
typename Op>
1737 if (llvm::cast<ValType>(srcType).getWidth() >=
1738 llvm::cast<ValType>(dstType).getWidth())
1740 << dstType <<
" must be wider than operand type " << srcType;
1746template <
typename ValType,
typename Op>
1751 if (llvm::cast<ValType>(srcType).getWidth() <=
1752 llvm::cast<ValType>(dstType).getWidth())
1754 << dstType <<
" must be shorter than operand type " << srcType;
1760template <
template <
typename>
class WidthComparator,
typename... ElementTypes>
1765 auto srcType =
getTypeIfLike<ElementTypes...>(inputs.front());
1766 auto dstType =
getTypeIfLike<ElementTypes...>(outputs.front());
1767 if (!srcType || !dstType)
1770 return WidthComparator<unsigned>()(dstType.getIntOrFloatBitWidth(),
1771 srcType.getIntOrFloatBitWidth());
1776static FailureOr<APFloat>
1778 const llvm::fltSemantics &targetSemantics,
1782 using fltNonfiniteBehavior = llvm::fltNonfiniteBehavior;
1783 if (sourceValue.isInfinity() &&
1784 (targetSemantics.nonFiniteBehavior == fltNonfiniteBehavior::NanOnly ||
1785 targetSemantics.nonFiniteBehavior == fltNonfiniteBehavior::FiniteOnly))
1787 if (sourceValue.isNaN() &&
1788 targetSemantics.nonFiniteBehavior == fltNonfiniteBehavior::FiniteOnly)
1791 bool losesInfo =
false;
1792 auto status = sourceValue.convert(targetSemantics, roundingMode, &losesInfo);
1793 if (losesInfo || status != APFloat::opOK)
1803OpFoldResult arith::ExtUIOp::fold(FoldAdaptor adaptor) {
1804 if (
auto lhs = getIn().getDefiningOp<ExtUIOp>()) {
1807 setNonNeg(
lhs.getNonNeg());
1808 getInMutable().assign(
lhs.getIn());
1813 unsigned bitWidth = llvm::cast<IntegerType>(resType).getWidth();
1815 adaptor.getOperands(),
getType(),
1816 [bitWidth](
const APInt &a,
bool &castStatus) {
1817 return a.zext(bitWidth);
1825LogicalResult arith::ExtUIOp::verify() {
1833OpFoldResult arith::ExtSIOp::fold(FoldAdaptor adaptor) {
1834 if (
auto lhs = getIn().getDefiningOp<ExtSIOp>()) {
1835 getInMutable().assign(
lhs.getIn());
1840 unsigned bitWidth = llvm::cast<IntegerType>(resType).getWidth();
1842 adaptor.getOperands(),
getType(),
1843 [bitWidth](
const APInt &a,
bool &castStatus) {
1844 return a.sext(bitWidth);
1852void arith::ExtSIOp::getCanonicalizationPatterns(RewritePatternSet &patterns,
1853 MLIRContext *context) {
1854 patterns.
add<ExtSIOfExtUI>(context);
1857LogicalResult arith::ExtSIOp::verify() {
1867OpFoldResult arith::ExtFOp::fold(FoldAdaptor adaptor) {
1868 if (
auto truncFOp = getOperand().getDefiningOp<TruncFOp>()) {
1869 if (truncFOp.getOperand().getType() ==
getType()) {
1870 arith::FastMathFlags truncFMF =
1871 truncFOp.getFastmath().value_or(arith::FastMathFlags::none);
1872 bool isTruncContract =
1873 bitEnumContainsAll(truncFMF, arith::FastMathFlags::contract);
1874 arith::FastMathFlags extFMF =
1875 getFastmath().value_or(arith::FastMathFlags::none);
1876 bool isExtContract =
1877 bitEnumContainsAll(extFMF, arith::FastMathFlags::contract);
1878 if (isTruncContract && isExtContract) {
1879 return truncFOp.getOperand();
1885 const llvm::fltSemantics &targetSemantics = resElemType.getFloatSemantics();
1887 adaptor.getOperands(),
getType(),
1888 [&targetSemantics](
const APFloat &a,
bool &castStatus) {
1913 function_ref<std::optional<APFloat>(
const APFloat &,
const APFloat &)>
1916 if (isa_and_nonnull<ub::PoisonAttr>(inAttr))
1918 if (isa_and_nonnull<ub::PoisonAttr>(scaleAttr))
1921 if (!inAttr || !scaleAttr || !resultType)
1924 if (
auto inFloat = dyn_cast<FloatAttr>(inAttr)) {
1925 auto scaleFloat = dyn_cast<FloatAttr>(scaleAttr);
1928 std::optional<APFloat>
result =
1929 calculate(inFloat.getValue(), scaleFloat.getValue());
1932 return FloatAttr::get(resultType, *
result);
1935 auto inElements = dyn_cast<DenseFPElementsAttr>(inAttr);
1936 auto scaleElements = dyn_cast<DenseFPElementsAttr>(scaleAttr);
1937 auto shapedResultType = dyn_cast<ShapedType>(resultType);
1938 if (!inElements || !scaleElements || !shapedResultType ||
1939 !shapedResultType.hasStaticShape() ||
1940 inElements.getNumElements() != scaleElements.getNumElements())
1944 if (inElements.isSplat() && scaleElements.isSplat()) {
1945 std::optional<APFloat>
result =
1946 calculate(inElements.getSplatValue<APFloat>(),
1947 scaleElements.getSplatValue<APFloat>());
1954 results.reserve(inElements.getNumElements());
1955 for (
const auto &[in, scale] : llvm::zip_equal(inElements, scaleElements)) {
1956 std::optional<APFloat>
result = calculate(in, scale);
1959 results.push_back(*
result);
1972OpFoldResult arith::ScalingExtFOp::fold(FoldAdaptor adaptor) {
1980 const llvm::fltSemantics &resSemantics = resElemType.getFloatSemantics();
1982 adaptor.getIn(), adaptor.getScale(),
getType(),
1983 [&resSemantics](
const APFloat &in,
1984 const APFloat &scale) -> std::optional<APFloat> {
1988 return std::nullopt;
1995bool arith::ScalingExtFOp::areCastCompatible(
TypeRange inputs,
2000LogicalResult arith::ScalingExtFOp::verify() {
2008OpFoldResult arith::TruncIOp::fold(FoldAdaptor adaptor) {
2011 Value src = getOperand().getDefiningOp()->getOperand(0);
2016 if (llvm::cast<IntegerType>(srcType).getWidth() >
2017 llvm::cast<IntegerType>(dstType).getWidth()) {
2024 if (srcType == dstType)
2030 setOperand(getOperand().getDefiningOp()->getOperand(0));
2035 unsigned bitWidth = llvm::cast<IntegerType>(resType).getWidth();
2037 adaptor.getOperands(),
getType(),
2038 [bitWidth](
const APInt &a,
bool &castStatus) {
2039 return a.trunc(bitWidth);
2047void arith::TruncIOp::getCanonicalizationPatterns(RewritePatternSet &patterns,
2048 MLIRContext *context) {
2049 patterns.
add<NarrowExtremum<TruncIOp, ExtSIOp, MaxSIOp>,
2050 NarrowExtremum<TruncIOp, ExtSIOp, MinSIOp>,
2051 NarrowExtremum<TruncIOp, ExtUIOp, MaxUIOp>,
2052 NarrowExtremum<TruncIOp, ExtUIOp, MinUIOp>, TruncIExtSIToExtSI,
2053 TruncIExtUIToExtUI, TruncIShrSIToTrunciShrUI>(context);
2056LogicalResult arith::TruncIOp::verify() {
2066OpFoldResult arith::TruncFOp::fold(FoldAdaptor adaptor) {
2068 if (
auto extOp = getOperand().getDefiningOp<arith::ExtFOp>()) {
2069 Value src = extOp.getIn();
2071 auto intermediateType =
2075 if (llvm::APFloatBase::isLosslesslyConvertibleTo(
2076 srcType.getFloatSemantics(),
2077 intermediateType.getFloatSemantics())) {
2079 if (srcType.getWidth() > resElemType.getWidth()) {
2085 if (srcType == resElemType)
2090 const llvm::fltSemantics &targetSemantics = resElemType.getFloatSemantics();
2092 adaptor.getOperands(),
getType(),
2093 [
this, &targetSemantics](
const APFloat &a,
bool &castStatus) {
2094 llvm::RoundingMode llvmRoundingMode =
2096 FailureOr<APFloat>
result =
2106void arith::TruncFOp::getCanonicalizationPatterns(RewritePatternSet &patterns,
2107 MLIRContext *context) {
2108 patterns.
add<NarrowExtremum<TruncFOp, ExtFOp, MaximumFOp>,
2109 NarrowExtremum<TruncFOp, ExtFOp, MaxNumFOp>,
2110 NarrowExtremum<TruncFOp, ExtFOp, MaximumNumFOp>,
2111 NarrowExtremum<TruncFOp, ExtFOp, MinimumFOp>,
2112 NarrowExtremum<TruncFOp, ExtFOp, MinNumFOp>,
2113 NarrowExtremum<TruncFOp, ExtFOp, MinimumNumFOp>,
2114 TruncFSIToFPToSIToFP, TruncFUIToFPToUIToFP>(context);
2121LogicalResult arith::TruncFOp::verify() {
2129OpFoldResult arith::ConvertFOp::fold(FoldAdaptor adaptor) {
2131 const llvm::fltSemantics &targetSemantics = resElemType.getFloatSemantics();
2133 adaptor.getOperands(),
getType(),
2134 [
this, &targetSemantics](
const APFloat &a,
bool &castStatus) {
2135 llvm::RoundingMode llvmRoundingMode =
2137 FailureOr<APFloat>
result =
2152 if (!srcType || !dstType)
2154 return srcType != dstType &&
2158LogicalResult arith::ConvertFOp::verify() {
2161 if (srcType == dstType)
2162 return emitError(
"result element type ")
2163 << dstType <<
" must be different from operand element type "
2165 if (srcType.getWidth() != dstType.getWidth())
2166 return emitError(
"result element type ")
2167 << dstType <<
" must have the same bitwidth as operand element type "
2176OpFoldResult arith::ScalingTruncFOp::fold(FoldAdaptor adaptor) {
2185 const llvm::fltSemantics &inSemantics = inElemType.getFloatSemantics();
2186 const llvm::fltSemantics &resSemantics = resElemType.getFloatSemantics();
2187 llvm::RoundingMode roundingMode =
2190 adaptor.getIn(), adaptor.getScale(),
getType(),
2191 [&](
const APFloat &in,
const APFloat &scale) -> std::optional<APFloat> {
2194 return std::nullopt;
2195 APFloat quotient(in);
2197 FailureOr<APFloat>
result =
2200 return std::nullopt;
2205bool arith::ScalingTruncFOp::areCastCompatible(
TypeRange inputs,
2210LogicalResult arith::ScalingTruncFOp::verify() {
2218void arith::AndIOp::getCanonicalizationPatterns(RewritePatternSet &patterns,
2219 MLIRContext *context) {
2220 patterns.
add<AndIAndIConstant, AndOfExtUI, AndOfExtSI>(context);
2227void arith::OrIOp::getCanonicalizationPatterns(RewritePatternSet &patterns,
2228 MLIRContext *context) {
2229 patterns.
add<OrIOrIConstant, OrOfExtUI, OrOfExtSI>(context);
2236template <
typename From,
typename To>
2244 return srcType && dstType;
2255OpFoldResult arith::UIToFPOp::fold(FoldAdaptor adaptor) {
2258 adaptor.getOperands(),
getType(),
2259 [&resEleType](
const APInt &a,
bool &castStatus) {
2260 FloatType floatTy = llvm::cast<FloatType>(resEleType);
2261 APFloat apf(floatTy.getFloatSemantics(),
2262 APInt::getZero(floatTy.getWidth()));
2263 apf.convertFromAPInt(a,
false,
2264 APFloat::rmNearestTiesToEven);
2269void arith::UIToFPOp::getCanonicalizationPatterns(RewritePatternSet &patterns,
2270 MLIRContext *context) {
2271 patterns.
add<UIToFPOfExtUI>(context);
2282OpFoldResult arith::SIToFPOp::fold(FoldAdaptor adaptor) {
2285 adaptor.getOperands(),
getType(),
2286 [&resEleType](
const APInt &a,
bool &castStatus) {
2287 FloatType floatTy = llvm::cast<FloatType>(resEleType);
2288 APFloat apf(floatTy.getFloatSemantics(),
2289 APInt::getZero(floatTy.getWidth()));
2290 apf.convertFromAPInt(a,
true,
2291 APFloat::rmNearestTiesToEven);
2296void arith::SIToFPOp::getCanonicalizationPatterns(RewritePatternSet &patterns,
2297 MLIRContext *context) {
2298 patterns.
add<SIToFPOfExtSI, SIToFPOfExtUI>(context);
2309OpFoldResult arith::FPToUIOp::fold(FoldAdaptor adaptor) {
2311 unsigned bitWidth = llvm::cast<IntegerType>(resType).getWidth();
2313 adaptor.getOperands(),
getType(),
2314 [&bitWidth](
const APFloat &a,
bool &castStatus) {
2316 APSInt api(bitWidth,
true);
2317 castStatus = APFloat::opInvalidOp !=
2318 a.convertToInteger(api, APFloat::rmTowardZero, &ignored);
2331OpFoldResult arith::FPToSIOp::fold(FoldAdaptor adaptor) {
2333 unsigned bitWidth = llvm::cast<IntegerType>(resType).getWidth();
2335 adaptor.getOperands(),
getType(),
2336 [&bitWidth](
const APFloat &a,
bool &castStatus) {
2338 APSInt api(bitWidth,
false);
2339 castStatus = APFloat::opInvalidOp !=
2340 a.convertToInteger(api, APFloat::rmTowardZero, &ignored);
2354 return intTy.getWidth();
2355 return IndexType::kInternalStorageBitWidth;
2364 if (!srcType || !dstType)
2368 (srcType.isSignlessInteger() && dstType.
isIndex());
2371bool arith::IndexCastOp::areCastCompatible(
TypeRange inputs,
2376OpFoldResult arith::IndexCastOp::fold(FoldAdaptor adaptor) {
2378 unsigned resultBitwidth = 64;
2380 resultBitwidth = intTy.getWidth();
2383 adaptor.getOperands(),
getType(),
2384 [resultBitwidth](
const APInt &a,
bool & ) {
2385 return a.sextOrTrunc(resultBitwidth);
2392 if (
auto inner = getOperand().getDefiningOp<arith::IndexCastOp>()) {
2393 Value x = inner.getOperand();
2402void arith::IndexCastOp::getCanonicalizationPatterns(
2403 RewritePatternSet &patterns, MLIRContext *context) {
2404 patterns.
add<IndexCastOfExtSI>(context);
2411bool arith::IndexCastUIOp::areCastCompatible(
TypeRange inputs,
2416OpFoldResult arith::IndexCastUIOp::fold(FoldAdaptor adaptor) {
2418 unsigned resultBitwidth = 64;
2420 resultBitwidth = intTy.getWidth();
2423 adaptor.getOperands(),
getType(),
2424 [resultBitwidth](
const APInt &a,
bool & ) {
2425 return a.zextOrTrunc(resultBitwidth);
2432 if (
auto inner = getOperand().getDefiningOp<arith::IndexCastUIOp>()) {
2433 Value x = inner.getOperand();
2442void arith::IndexCastUIOp::getCanonicalizationPatterns(
2443 RewritePatternSet &patterns, MLIRContext *context) {
2444 patterns.
add<IndexCastUIOfExtUI>(context);
2457 if (!srcType || !dstType)
2463OpFoldResult arith::BitcastOp::fold(FoldAdaptor adaptor) {
2465 auto operand = adaptor.getIn();
2470 if (
auto denseAttr = dyn_cast_or_null<DenseElementsAttr>(operand))
2471 return denseAttr.bitcast(llvm::cast<ShapedType>(resType).
getElementType());
2473 if (llvm::isa<ShapedType>(resType))
2481 if (!llvm::isa<FloatAttr, IntegerAttr>(operand))
2484 APInt bits = llvm::isa<FloatAttr>(operand)
2485 ? llvm::cast<FloatAttr>(operand).getValue().bitcastToAPInt()
2486 : llvm::cast<IntegerAttr>(operand).getValue();
2488 "trying to fold on broken IR: operands have incompatible types");
2490 if (
auto resFloatType = dyn_cast<FloatType>(resType))
2491 return FloatAttr::get(resType,
2492 APFloat(resFloatType.getFloatSemantics(), bits));
2493 return IntegerAttr::get(resType, bits);
2496void arith::BitcastOp::getCanonicalizationPatterns(RewritePatternSet &patterns,
2497 MLIRContext *context) {
2498 patterns.
add<BitcastOfBitcast>(context);
2508 const APInt &lhs,
const APInt &rhs) {
2509 switch (predicate) {
2510 case arith::CmpIPredicate::eq:
2512 case arith::CmpIPredicate::ne:
2514 case arith::CmpIPredicate::slt:
2515 return lhs.slt(rhs);
2516 case arith::CmpIPredicate::sle:
2517 return lhs.sle(rhs);
2518 case arith::CmpIPredicate::sgt:
2519 return lhs.sgt(rhs);
2520 case arith::CmpIPredicate::sge:
2521 return lhs.sge(rhs);
2522 case arith::CmpIPredicate::ult:
2523 return lhs.ult(rhs);
2524 case arith::CmpIPredicate::ule:
2525 return lhs.ule(rhs);
2526 case arith::CmpIPredicate::ugt:
2527 return lhs.ugt(rhs);
2528 case arith::CmpIPredicate::uge:
2529 return lhs.uge(rhs);
2531 llvm_unreachable(
"unknown cmpi predicate kind");
2536 switch (predicate) {
2537 case arith::CmpIPredicate::eq:
2538 case arith::CmpIPredicate::sle:
2539 case arith::CmpIPredicate::sge:
2540 case arith::CmpIPredicate::ule:
2541 case arith::CmpIPredicate::uge:
2543 case arith::CmpIPredicate::ne:
2544 case arith::CmpIPredicate::slt:
2545 case arith::CmpIPredicate::sgt:
2546 case arith::CmpIPredicate::ult:
2547 case arith::CmpIPredicate::ugt:
2550 llvm_unreachable(
"unknown cmpi predicate kind");
2554 if (
auto intType = dyn_cast<IntegerType>(t)) {
2555 return intType.getWidth();
2557 if (
auto vectorIntType = dyn_cast<VectorType>(t)) {
2558 return llvm::cast<IntegerType>(vectorIntType.getElementType()).getWidth();
2560 return std::nullopt;
2563OpFoldResult arith::CmpIOp::fold(FoldAdaptor adaptor) {
2565 if (getLhs() == getRhs()) {
2571 if (
auto extOp = getLhs().getDefiningOp<ExtSIOp>()) {
2573 std::optional<int64_t> integerWidth =
2575 if (integerWidth && integerWidth.value() == 1 &&
2576 getPredicate() == arith::CmpIPredicate::ne)
2577 return extOp.getOperand();
2579 if (
auto extOp = getLhs().getDefiningOp<ExtUIOp>()) {
2581 std::optional<int64_t> integerWidth =
2583 if (integerWidth && integerWidth.value() == 1 &&
2584 getPredicate() == arith::CmpIPredicate::ne)
2585 return extOp.getOperand();
2590 getPredicate() == arith::CmpIPredicate::ne)
2597 getPredicate() == arith::CmpIPredicate::eq)
2602 if (adaptor.getLhs() && !adaptor.getRhs()) {
2604 using Pred = CmpIPredicate;
2605 const std::pair<Pred, Pred> invPreds[] = {
2606 {Pred::slt, Pred::sgt}, {Pred::sgt, Pred::slt}, {Pred::sle, Pred::sge},
2607 {Pred::sge, Pred::sle}, {Pred::ult, Pred::ugt}, {Pred::ugt, Pred::ult},
2608 {Pred::ule, Pred::uge}, {Pred::uge, Pred::ule}, {Pred::eq, Pred::eq},
2609 {Pred::ne, Pred::ne},
2611 Pred origPred = getPredicate();
2612 for (
auto pred : invPreds) {
2613 if (origPred == pred.first) {
2614 setPredicate(pred.second);
2615 Value
lhs = getLhs();
2616 Value
rhs = getRhs();
2617 getLhsMutable().assign(
rhs);
2618 getRhsMutable().assign(
lhs);
2622 llvm_unreachable(
"unknown cmpi predicate kind");
2627 if (
auto lhs = dyn_cast_if_present<TypedAttr>(adaptor.getLhs())) {
2630 [pred = getPredicate()](
const APInt &
lhs,
const APInt &
rhs) {
2639void arith::CmpIOp::getCanonicalizationPatterns(RewritePatternSet &patterns,
2640 MLIRContext *context) {
2641 patterns.
insert<CmpIExtSI, CmpIExtUI>(context);
2651 const APFloat &lhs,
const APFloat &rhs) {
2652 auto cmpResult = lhs.compare(rhs);
2653 switch (predicate) {
2654 case arith::CmpFPredicate::AlwaysFalse:
2656 case arith::CmpFPredicate::OEQ:
2657 return cmpResult == APFloat::cmpEqual;
2658 case arith::CmpFPredicate::OGT:
2659 return cmpResult == APFloat::cmpGreaterThan;
2660 case arith::CmpFPredicate::OGE:
2661 return cmpResult == APFloat::cmpGreaterThan ||
2662 cmpResult == APFloat::cmpEqual;
2663 case arith::CmpFPredicate::OLT:
2664 return cmpResult == APFloat::cmpLessThan;
2665 case arith::CmpFPredicate::OLE:
2666 return cmpResult == APFloat::cmpLessThan || cmpResult == APFloat::cmpEqual;
2667 case arith::CmpFPredicate::ONE:
2668 return cmpResult != APFloat::cmpUnordered && cmpResult != APFloat::cmpEqual;
2669 case arith::CmpFPredicate::ORD:
2670 return cmpResult != APFloat::cmpUnordered;
2671 case arith::CmpFPredicate::UEQ:
2672 return cmpResult == APFloat::cmpUnordered || cmpResult == APFloat::cmpEqual;
2673 case arith::CmpFPredicate::UGT:
2674 return cmpResult == APFloat::cmpUnordered ||
2675 cmpResult == APFloat::cmpGreaterThan;
2676 case arith::CmpFPredicate::UGE:
2677 return cmpResult == APFloat::cmpUnordered ||
2678 cmpResult == APFloat::cmpGreaterThan ||
2679 cmpResult == APFloat::cmpEqual;
2680 case arith::CmpFPredicate::ULT:
2681 return cmpResult == APFloat::cmpUnordered ||
2682 cmpResult == APFloat::cmpLessThan;
2683 case arith::CmpFPredicate::ULE:
2684 return cmpResult == APFloat::cmpUnordered ||
2685 cmpResult == APFloat::cmpLessThan || cmpResult == APFloat::cmpEqual;
2686 case arith::CmpFPredicate::UNE:
2687 return cmpResult != APFloat::cmpEqual;
2688 case arith::CmpFPredicate::UNO:
2689 return cmpResult == APFloat::cmpUnordered;
2690 case arith::CmpFPredicate::AlwaysTrue:
2693 llvm_unreachable(
"unknown cmpf predicate kind");
2697 auto lhs = dyn_cast_if_present<FloatAttr>(adaptor.getLhs());
2698 auto rhs = dyn_cast_if_present<FloatAttr>(adaptor.getRhs());
2701 if (lhs && lhs.getValue().isNaN())
2703 if (rhs && rhs.getValue().isNaN())
2719 using namespace arith;
2721 case CmpFPredicate::UEQ:
2722 case CmpFPredicate::OEQ:
2723 return CmpIPredicate::eq;
2724 case CmpFPredicate::UGT:
2725 case CmpFPredicate::OGT:
2726 return isUnsigned ? CmpIPredicate::ugt : CmpIPredicate::sgt;
2727 case CmpFPredicate::UGE:
2728 case CmpFPredicate::OGE:
2729 return isUnsigned ? CmpIPredicate::uge : CmpIPredicate::sge;
2730 case CmpFPredicate::ULT:
2731 case CmpFPredicate::OLT:
2732 return isUnsigned ? CmpIPredicate::ult : CmpIPredicate::slt;
2733 case CmpFPredicate::ULE:
2734 case CmpFPredicate::OLE:
2735 return isUnsigned ? CmpIPredicate::ule : CmpIPredicate::sle;
2736 case CmpFPredicate::UNE:
2737 case CmpFPredicate::ONE:
2738 return CmpIPredicate::ne;
2740 llvm_unreachable(
"Unexpected predicate!");
2750 const APFloat &rhs = flt.getValue();
2758 FloatType floatTy = llvm::cast<FloatType>(op.getRhs().getType());
2759 int mantissaWidth = floatTy.getFPMantissaWidth();
2760 if (mantissaWidth <= 0)
2766 if (
auto si = op.getLhs().getDefiningOp<SIToFPOp>()) {
2768 intVal = si.getIn();
2769 }
else if (
auto ui = op.getLhs().getDefiningOp<UIToFPOp>()) {
2771 intVal = ui.getIn();
2778 auto intTy = llvm::cast<IntegerType>(intVal.
getType());
2779 auto intWidth = intTy.getWidth();
2782 auto valueBits = isUnsigned ? intWidth : (intWidth - 1);
2787 if ((
int)intWidth > mantissaWidth) {
2789 int exponent = ilogb(rhs);
2790 if (exponent == APFloat::IEK_Inf) {
2791 int maxExponent = ilogb(APFloat::getLargest(rhs.getSemantics()));
2792 if (maxExponent < (
int)valueBits) {
2799 if (mantissaWidth <= exponent && exponent <= (
int)valueBits) {
2808 switch (op.getPredicate()) {
2809 case CmpFPredicate::ORD:
2814 case CmpFPredicate::UNO:
2827 APFloat signedMax(rhs.getSemantics());
2828 signedMax.convertFromAPInt(APInt::getSignedMaxValue(intWidth),
true,
2829 APFloat::rmNearestTiesToEven);
2830 if (signedMax < rhs) {
2831 if (pred == CmpIPredicate::ne || pred == CmpIPredicate::slt ||
2832 pred == CmpIPredicate::sle)
2843 APFloat unsignedMax(rhs.getSemantics());
2844 unsignedMax.convertFromAPInt(APInt::getMaxValue(intWidth),
false,
2845 APFloat::rmNearestTiesToEven);
2846 if (unsignedMax < rhs) {
2847 if (pred == CmpIPredicate::ne || pred == CmpIPredicate::ult ||
2848 pred == CmpIPredicate::ule)
2860 APFloat signedMin(rhs.getSemantics());
2861 signedMin.convertFromAPInt(APInt::getSignedMinValue(intWidth),
true,
2862 APFloat::rmNearestTiesToEven);
2863 if (signedMin > rhs) {
2864 if (pred == CmpIPredicate::ne || pred == CmpIPredicate::sgt ||
2865 pred == CmpIPredicate::sge)
2875 APFloat unsignedMin(rhs.getSemantics());
2876 unsignedMin.convertFromAPInt(APInt::getMinValue(intWidth),
false,
2877 APFloat::rmNearestTiesToEven);
2878 if (unsignedMin > rhs) {
2879 if (pred == CmpIPredicate::ne || pred == CmpIPredicate::ugt ||
2880 pred == CmpIPredicate::uge)
2895 APSInt rhsInt(intWidth, isUnsigned);
2896 if (APFloat::opInvalidOp ==
2897 rhs.convertToInteger(rhsInt, APFloat::rmTowardZero, &ignored)) {
2903 if (!rhs.isZero()) {
2904 APFloat apf(floatTy.getFloatSemantics(),
2905 APInt::getZero(floatTy.getWidth()));
2906 apf.convertFromAPInt(rhsInt, !isUnsigned, APFloat::rmNearestTiesToEven);
2908 bool equal = apf == rhs;
2914 case CmpIPredicate::ne:
2918 case CmpIPredicate::eq:
2922 case CmpIPredicate::ule:
2925 if (rhs.isNegative()) {
2931 case CmpIPredicate::sle:
2934 if (rhs.isNegative())
2935 pred = CmpIPredicate::slt;
2937 case CmpIPredicate::ult:
2940 if (rhs.isNegative()) {
2945 pred = CmpIPredicate::ule;
2947 case CmpIPredicate::slt:
2950 if (!rhs.isNegative())
2951 pred = CmpIPredicate::sle;
2953 case CmpIPredicate::ugt:
2956 if (rhs.isNegative()) {
2962 case CmpIPredicate::sgt:
2965 if (rhs.isNegative())
2966 pred = CmpIPredicate::sge;
2968 case CmpIPredicate::uge:
2971 if (rhs.isNegative()) {
2976 pred = CmpIPredicate::ugt;
2978 case CmpIPredicate::sge:
2981 if (!rhs.isNegative())
2982 pred = CmpIPredicate::sgt;
2992 ConstantOp::create(rewriter, op.getLoc(), intVal.
getType(),
2998void arith::CmpFOp::getCanonicalizationPatterns(RewritePatternSet &patterns,
2999 MLIRContext *context) {
3000 patterns.
insert<CmpFIntToFPConst>(context);
3014 if (!llvm::isa<IntegerType>(op.getType()) || op.getType().isInteger(1))
3030 arith::XOrIOp::create(
3031 rewriter, op.getLoc(), op.getCondition(),
3033 op.getCondition().
getType(), 1)));
3041void arith::SelectOp::getCanonicalizationPatterns(RewritePatternSet &results,
3042 MLIRContext *context) {
3043 results.
add<RedundantSelectFalse, RedundantSelectTrue, SelectNotCond,
3044 SelectI1ToNot, SelectCmpISgeToMaxSI, SelectCmpISgeToMinSI,
3045 SelectCmpISgtToMaxSI, SelectCmpISgtToMinSI, SelectCmpISleToMaxSI,
3046 SelectCmpISleToMinSI, SelectCmpISltToMaxSI, SelectCmpISltToMinSI,
3047 SelectCmpIUgeToMaxUI, SelectCmpIUgeToMinUI, SelectCmpIUgtToMaxUI,
3048 SelectCmpIUgtToMinUI, SelectCmpIUleToMaxUI, SelectCmpIUleToMinUI,
3049 SelectCmpIUltToMaxUI, SelectCmpIUltToMinUI, SelectToExtUI>(
3053OpFoldResult arith::SelectOp::fold(FoldAdaptor adaptor) {
3054 Value trueVal = getTrueValue();
3055 Value falseVal = getFalseValue();
3056 if (trueVal == falseVal)
3059 Value condition = getCondition();
3077 if (
getType().isSignlessInteger(1) &&
3083 auto pred = cmp.getPredicate();
3084 if (pred == arith::CmpIPredicate::eq || pred == arith::CmpIPredicate::ne) {
3085 auto cmpLhs = cmp.getLhs();
3086 auto cmpRhs = cmp.getRhs();
3094 if ((cmpLhs == trueVal && cmpRhs == falseVal) ||
3095 (cmpRhs == trueVal && cmpLhs == falseVal))
3096 return pred == arith::CmpIPredicate::ne ? trueVal : falseVal;
3103 dyn_cast_if_present<DenseElementsAttr>(adaptor.getCondition())) {
3105 assert(cond.getType().hasStaticShape() &&
3106 "DenseElementsAttr must have static shape");
3108 dyn_cast_if_present<DenseElementsAttr>(adaptor.getTrueValue())) {
3110 dyn_cast_if_present<DenseElementsAttr>(adaptor.getFalseValue())) {
3111 SmallVector<Attribute> results;
3112 results.reserve(
static_cast<size_t>(cond.getNumElements()));
3113 auto condVals = llvm::make_range(cond.value_begin<BoolAttr>(),
3114 cond.value_end<BoolAttr>());
3115 auto lhsVals = llvm::make_range(
lhs.value_begin<Attribute>(),
3116 lhs.value_end<Attribute>());
3117 auto rhsVals = llvm::make_range(
rhs.value_begin<Attribute>(),
3118 rhs.value_end<Attribute>());
3120 for (
auto [condVal, lhsVal, rhsVal] :
3121 llvm::zip_equal(condVals, lhsVals, rhsVals))
3122 results.push_back(condVal.getValue() ? lhsVal : rhsVal);
3132ParseResult SelectOp::parse(OpAsmParser &parser, OperationState &
result) {
3133 Type conditionType, resultType;
3134 SmallVector<OpAsmParser::UnresolvedOperand, 3> operands;
3142 conditionType = resultType;
3149 result.addTypes(resultType);
3151 {conditionType, resultType, resultType},
3155void arith::SelectOp::print(OpAsmPrinter &p) {
3156 p <<
" " << getOperands();
3159 if (ShapedType condType = dyn_cast<ShapedType>(getCondition().
getType()))
3160 p << condType <<
", ";
3164LogicalResult arith::SelectOp::verify() {
3165 Type conditionType = getCondition().getType();
3172 if (!llvm::isa<TensorType, VectorType>(resultType))
3173 return emitOpError() <<
"expected condition to be a signless i1, but got "
3176 if (conditionType != shapedConditionType) {
3177 return emitOpError() <<
"expected condition type to have the same shape "
3178 "as the result type, expected "
3179 << shapedConditionType <<
", but got "
3188OpFoldResult arith::ShLIOp::fold(FoldAdaptor adaptor) {
3202 bool bounded =
false;
3204 adaptor.getOperands(), [&](
const APInt &a,
const APInt &
b) {
3205 bounded = b.ult(b.getBitWidth());
3208 return bounded ?
result : Attribute();
3215OpFoldResult arith::ShRUIOp::fold(FoldAdaptor adaptor) {
3230 if (getLhs() == getRhs())
3233 bool bounded =
false;
3235 adaptor.getOperands(), [&](
const APInt &a,
const APInt &
b) {
3236 bounded = b.ult(b.getBitWidth());
3239 return bounded ?
result : Attribute();
3246OpFoldResult arith::ShRSIOp::fold(FoldAdaptor adaptor) {
3261 if (getLhs() == getRhs())
3269 bool bounded =
false;
3271 adaptor.getOperands(), [&](
const APInt &a,
const APInt &
b) {
3272 bounded = b.ult(b.getBitWidth());
3275 return bounded ?
result : Attribute();
3285 bool useOnlyFiniteValue) {
3287 case AtomicRMWKind::maximumf: {
3288 const llvm::fltSemantics &semantic =
3289 llvm::cast<FloatType>(resultType).getFloatSemantics();
3290 APFloat identity = useOnlyFiniteValue
3291 ? APFloat::getLargest(semantic,
true)
3292 : APFloat::getInf(semantic,
true);
3295 case AtomicRMWKind::maxnumf: {
3296 const llvm::fltSemantics &semantic =
3297 llvm::cast<FloatType>(resultType).getFloatSemantics();
3298 APFloat identity = APFloat::getNaN(semantic,
true);
3301 case AtomicRMWKind::addf:
3302 case AtomicRMWKind::addi:
3303 case AtomicRMWKind::maxu:
3304 case AtomicRMWKind::ori:
3305 case AtomicRMWKind::xori:
3307 case AtomicRMWKind::andi:
3310 APInt::getAllOnes(llvm::cast<IntegerType>(resultType).getWidth()));
3311 case AtomicRMWKind::maxs:
3313 resultType, APInt::getSignedMinValue(
3314 llvm::cast<IntegerType>(resultType).getWidth()));
3315 case AtomicRMWKind::minimumf: {
3316 const llvm::fltSemantics &semantic =
3317 llvm::cast<FloatType>(resultType).getFloatSemantics();
3318 APFloat identity = useOnlyFiniteValue
3319 ? APFloat::getLargest(semantic,
false)
3320 : APFloat::getInf(semantic,
false);
3324 case AtomicRMWKind::minnumf: {
3325 const llvm::fltSemantics &semantic =
3326 llvm::cast<FloatType>(resultType).getFloatSemantics();
3327 APFloat identity = APFloat::getNaN(semantic,
false);
3330 case AtomicRMWKind::mins:
3332 resultType, APInt::getSignedMaxValue(
3333 llvm::cast<IntegerType>(resultType).getWidth()));
3334 case AtomicRMWKind::minu:
3337 APInt::getMaxValue(llvm::cast<IntegerType>(resultType).getWidth()));
3338 case AtomicRMWKind::muli:
3340 case AtomicRMWKind::mulf:
3343 case AtomicRMWKind::assign:
3352 std::optional<AtomicRMWKind> maybeKind =
3355 .Case([](arith::AddFOp op) {
return AtomicRMWKind::addf; })
3356 .Case([](arith::MulFOp op) {
return AtomicRMWKind::mulf; })
3357 .Case([](arith::MaximumFOp op) {
return AtomicRMWKind::maximumf; })
3358 .Case([](arith::MinimumFOp op) {
return AtomicRMWKind::minimumf; })
3359 .Case([](arith::MaxNumFOp op) {
return AtomicRMWKind::maxnumf; })
3360 .Case([](arith::MinNumFOp op) {
return AtomicRMWKind::minnumf; })
3362 .Case([](arith::AddIOp op) {
return AtomicRMWKind::addi; })
3363 .Case([](arith::OrIOp op) {
return AtomicRMWKind::ori; })
3364 .Case([](arith::XOrIOp op) {
return AtomicRMWKind::xori; })
3365 .Case([](arith::AndIOp op) {
return AtomicRMWKind::andi; })
3366 .Case([](arith::MaxUIOp op) {
return AtomicRMWKind::maxu; })
3367 .Case([](arith::MinUIOp op) {
return AtomicRMWKind::minu; })
3368 .Case([](arith::MaxSIOp op) {
return AtomicRMWKind::maxs; })
3369 .Case([](arith::MinSIOp op) {
return AtomicRMWKind::mins; })
3370 .Case([](arith::MulIOp op) {
return AtomicRMWKind::muli; })
3371 .Default(std::nullopt);
3373 return std::nullopt;
3376 bool useOnlyFiniteValue =
false;
3377 auto fmfOpInterface = dyn_cast<ArithFastMathInterface>(op);
3378 if (fmfOpInterface) {
3379 arith::FastMathFlagsAttr fmfAttr = fmfOpInterface.getFastMathFlagsAttr();
3380 useOnlyFiniteValue =
3381 bitEnumContainsAny(fmfAttr.getValue(), arith::FastMathFlags::ninf);
3389 useOnlyFiniteValue);
3395 bool useOnlyFiniteValue) {
3397 useOnlyFiniteValue))
3398 return arith::ConstantOp::create(builder, loc, attr);
3407 case AtomicRMWKind::addf:
3408 return arith::AddFOp::create(builder, loc, lhs, rhs);
3409 case AtomicRMWKind::addi:
3410 return arith::AddIOp::create(builder, loc, lhs, rhs);
3411 case AtomicRMWKind::mulf:
3412 return arith::MulFOp::create(builder, loc, lhs, rhs);
3413 case AtomicRMWKind::muli:
3414 return arith::MulIOp::create(builder, loc, lhs, rhs);
3415 case AtomicRMWKind::maximumf:
3416 return arith::MaximumFOp::create(builder, loc, lhs, rhs);
3417 case AtomicRMWKind::minimumf:
3418 return arith::MinimumFOp::create(builder, loc, lhs, rhs);
3419 case AtomicRMWKind::maxnumf:
3420 return arith::MaxNumFOp::create(builder, loc, lhs, rhs);
3421 case AtomicRMWKind::minnumf:
3422 return arith::MinNumFOp::create(builder, loc, lhs, rhs);
3423 case AtomicRMWKind::maxs:
3424 return arith::MaxSIOp::create(builder, loc, lhs, rhs);
3425 case AtomicRMWKind::mins:
3426 return arith::MinSIOp::create(builder, loc, lhs, rhs);
3427 case AtomicRMWKind::maxu:
3428 return arith::MaxUIOp::create(builder, loc, lhs, rhs);
3429 case AtomicRMWKind::minu:
3430 return arith::MinUIOp::create(builder, loc, lhs, rhs);
3431 case AtomicRMWKind::ori:
3432 return arith::OrIOp::create(builder, loc, lhs, rhs);
3433 case AtomicRMWKind::andi:
3434 return arith::AndIOp::create(builder, loc, lhs, rhs);
3435 case AtomicRMWKind::xori:
3436 return arith::XOrIOp::create(builder, loc, lhs, rhs);
3438 case AtomicRMWKind::assign:
3449#define GET_OP_CLASSES
3450#include "mlir/Dialect/Arith/IR/ArithOps.cpp.inc"
3456#include "mlir/Dialect/Arith/IR/ArithOpsEnums.cpp.inc"
if(failed(verifyVectorMemoryOp(getOperation(), memrefType, getVectorType()))) return failure()
static Speculation::Speculatability getDivUISpeculatability(Value divisor)
Returns whether an unsigned division by divisor is speculatable.
static bool checkWidthChangeCast(TypeRange inputs, TypeRange outputs)
Validate a cast that changes the width of a type.
static IntegerAttr mulIntegerAttrs(PatternRewriter &builder, Value res, Attribute lhs, Attribute rhs)
static IntegerOverflowFlagsAttr mergeOverflowFlags(IntegerOverflowFlagsAttr val1, IntegerOverflowFlagsAttr val2)
static Value foldIdempotentOfSameOp(OpTy op)
Fold op(a, op(a, b)) to op(a, b) for an associative, commutative and idempotent op (e....
static constexpr llvm::RoundingMode kDefaultRoundingMode
Default rounding mode according to default LLVM floating-point environment.
static Type getTypeIfLike(Type type)
Get allowed underlying types for vectors and tensors.
static bool applyCmpPredicateToEqualOperands(arith::CmpIPredicate predicate)
Returns true if the predicate is true for two equal operands.
static FailureOr< APFloat > convertFloatValue(APFloat sourceValue, const llvm::fltSemantics &targetSemantics, llvm::RoundingMode roundingMode=kDefaultRoundingMode)
Attempts to convert sourceValue to an APFloat value with targetSemantics and roundingMode,...
static Attribute foldScalingCastOp(Attribute inAttr, Attribute scaleAttr, Type resultType, function_ref< std::optional< APFloat >(const APFloat &, const APFloat &)> calculate)
Fold calculate element-wise over the operands of a scaling cast op.
static Value foldDivMul(Value lhs, Value rhs, arith::IntegerOverflowFlags ovfFlags)
Fold (a * b) / b -> a
static bool hasSameEncoding(Type typeA, Type typeB)
Return false if both types are ranked tensor with mismatching encoding.
static llvm::RoundingMode convertArithRoundingModeToLLVMIR(std::optional< RoundingMode > roundingMode)
Equivalent to convertRoundingModeToLLVM(convertArithRoundingModeToLLVM(roundingMode)).
static Type getUnderlyingType(Type type, type_list< ShapedTypes... >, type_list< ElementTypes... >)
Returns a non-null type only if the provided type is one of the allowed types or one of the allowed s...
static std::optional< int64_t > getIntegerWidth(Type t)
static Speculation::Speculatability getDivSISpeculatability(Value divisor)
Returns whether a signed division by divisor is speculatable.
static IntegerAttr orIntegerAttrs(PatternRewriter &builder, Value res, Attribute lhs, Attribute rhs)
static IntegerAttr addIntegerAttrs(PatternRewriter &builder, Value res, Attribute lhs, Attribute rhs)
static Attribute getBoolAttribute(Type type, bool value)
static bool areIndexCastCompatible(TypeRange inputs, TypeRange outputs)
static bool isFoldableScalingScale(Value scale)
Only scales that already are f8E8M0FNU fold.
static bool checkIntFloatCast(TypeRange inputs, TypeRange outputs)
static LogicalResult verifyExtOp(Op op)
static IntegerAttr subIntegerAttrs(PatternRewriter &builder, Value res, Attribute lhs, Attribute rhs)
static Attribute getIntegerAttrOfType(Type type, int64_t value)
Return a scalar or splat integer attribute of type (an integer/index type or a shaped type thereof) h...
static IntegerAttr andIntegerAttrs(PatternRewriter &builder, Value res, Attribute lhs, Attribute rhs)
static int64_t getScalarOrElementWidth(Type type)
static Type getTypeIfLikeOrMemRef(Type type)
Get allowed underlying types for vectors, tensors, and memrefs.
static Type getI1SameShape(Type type)
Return the type of the same shape (scalar, vector or tensor) containing i1.
static bool areValidCastInputsAndOutputs(TypeRange inputs, TypeRange outputs)
static IntegerAttr xorIntegerAttrs(PatternRewriter &builder, Value res, Attribute lhs, Attribute rhs)
static APInt calculateUnsignedBorrow(const APInt &lhs, const APInt &rhs)
std::tuple< Types... > * type_list
static IntegerAttr applyToIntegerAttrs(PatternRewriter &builder, Value res, Attribute lhs, Attribute rhs, function_ref< APInt(const APInt &, const APInt &)> binFn)
static APInt calculateUnsignedOverflow(const APInt &sum, const APInt &operand)
static FailureOr< APInt > getIntOrSplatIntValue(Attribute attr)
static unsigned getIndexCastWidth(Type t)
Return the bit-width of t for the purpose of index_cast width checks.
static LogicalResult verifyTruncateOp(Op op)
static Type getElementType(Type type, ArrayRef< int32_t > indices, function_ref< InFlightDiagnostic(StringRef)> emitErrorFn)
Walks the given type hierarchy with the given indices, potentially down to component granularity,...
LogicalResult matchAndRewrite(CmpFOp op, PatternRewriter &rewriter) const override
static CmpIPredicate convertToIntegerPredicate(CmpFPredicate pred, bool isUnsigned)
virtual Builder & getBuilder() const =0
Return a builder which provides useful access to MLIRContext, global objects like types and attribute...
virtual ParseResult parseOptionalAttrDict(NamedAttrList &result)=0
Parse a named dictionary into 'result' if it is present.
virtual ParseResult parseColonType(Type &result)=0
Parse a colon followed by a type.
virtual ParseResult parseOptionalComma()=0
Parse a , token if present.
virtual SMLoc getNameLoc() const =0
Return the location of the original name token.
virtual ParseResult parseType(Type &result)=0
Parse a type.
Attributes are known-constant values of operations.
static BoolAttr get(MLIRContext *context, bool value)
IntegerAttr getIndexAttr(int64_t value)
IntegerAttr getIntegerAttr(Type type, int64_t value)
FloatAttr getFloatAttr(Type type, double value)
IntegerType getIntegerType(unsigned width)
Ty getType(Args &&...args)
Get or construct an instance of the type Ty with provided arguments.
TypedAttr getZeroAttr(Type type)
static DenseElementsAttr get(ShapedType type, ArrayRef< Attribute > values)
Constructs a dense elements attribute from an array of element values.
ImplicitLocOpBuilder maintains a 'current location', allowing use of the create<> method without spec...
Location getLoc() const
Accessors for the implied location.
This class defines the main interface for locations in MLIR and acts as a non-nullable wrapper around...
MLIRContext is the top-level object for a collection of MLIR operations.
ParseResult resolveOperands(Operands &&operands, Type type, SmallVectorImpl< Value > &result)
Resolve a list of operands to SSA values, emitting an error on failure, or appending the results to t...
virtual ParseResult parseOperandList(SmallVectorImpl< UnresolvedOperand > &result, Delimiter delimiter=Delimiter::None, bool allowResultNumber=true, int requiredOperandCount=-1)=0
Parse zero or more SSA comma-separated operand references with a specified surrounding delimiter,...
virtual void printOptionalAttrDict(ArrayRef< NamedAttribute > attrs, ArrayRef< StringRef > elidedAttrs={})=0
If the specified operation has attributes, print out an attribute dictionary with their values.
This class helps build Operations.
Operation * create(const OperationState &state)
Creates an operation given the fields represented as an OperationState.
This class represents a single result from folding an operation.
InFlightDiagnostic emitError(const Twine &message={})
Emit an error about fatal conditions with this operation, reporting up to any diagnostic handlers tha...
This provides public APIs that all operations should have.
Operation is the basic unit of execution within MLIR.
OpResult getResult(unsigned idx)
Get the 'idx'th result of this operation.
Location getLoc()
The source location the operation was defined or derived from.
MLIRContext * getContext()
Return the context this operation is associated with.
A special type of RewriterBase that coordinates the application of a rewrite pattern on the current I...
RewritePatternSet & insert(ConstructorArg &&arg, ConstructorArgs &&...args)
Add an instance of each of the pattern types 'Ts' to the pattern list with the given arguments.
RewritePatternSet & add(ConstructorArg &&arg, ConstructorArgs &&...args)
Add an instance of each of the pattern types 'Ts' to the pattern list with the given arguments.
OpTy replaceOpWithNewOp(Operation *op, Args &&...args)
Replace the results of the given (original) op with a new op that is created without verification (re...
This class provides an abstraction over the various different ranges of value types.
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
MLIRContext * getContext() const
Return the MLIRContext in which this type was uniqued.
bool isSignlessInteger() const
Return true if this is a signless integer type (with the specified width).
bool isIntOrFloat() const
Return true if this is an integer (of any signedness) or a float type.
unsigned getIntOrFloatBitWidth() const
Return the bit width of an integer or a float type, assert failure on other types.
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
Type getType() const
Return the type of this value.
Operation * getDefiningOp() const
If this value is the result of an operation, return the operation that defines it.
Specialization of arith.constant op that returns a floating point value.
static ConstantFloatOp create(OpBuilder &builder, Location location, FloatType type, const APFloat &value)
static bool classof(Operation *op)
static void build(OpBuilder &builder, OperationState &result, FloatType type, const APFloat &value)
Build a constant float op that produces a float of the specified type.
Specialization of arith.constant op that returns an integer of index type.
static void build(OpBuilder &builder, OperationState &result, int64_t value)
Build a constant int op that produces an index.
static bool classof(Operation *op)
static ConstantIndexOp create(OpBuilder &builder, Location location, int64_t value)
Specialization of arith.constant op that returns an integer value.
static ConstantIntOp create(OpBuilder &builder, Location location, int64_t value, unsigned width)
static void build(OpBuilder &builder, OperationState &result, int64_t value, unsigned width)
Build a constant int op that produces an integer of the specified width.
static bool classof(Operation *op)
Speculatability
This enum is returned from the getSpeculatability method in the ConditionallySpeculatable op interfac...
constexpr auto Speculatable
constexpr auto NotSpeculatable
std::optional< TypedAttr > getNeutralElement(Operation *op)
Return the identity numeric value associated to the give op.
bool applyCmpPredicate(arith::CmpIPredicate predicate, const APInt &lhs, const APInt &rhs)
Compute lhs pred rhs, where pred is one of the known integer comparison predicates.
TypedAttr getIdentityValueAttr(AtomicRMWKind kind, Type resultType, OpBuilder &builder, Location loc, bool useOnlyFiniteValue=false)
Returns the identity value attribute associated with an AtomicRMWKind op.
Value getReductionOp(AtomicRMWKind op, OpBuilder &builder, Location loc, Value lhs, Value rhs)
Returns the value obtained by applying the reduction operation kind associated with a binary AtomicRM...
Value getIdentityValue(AtomicRMWKind op, Type resultType, OpBuilder &builder, Location loc, bool useOnlyFiniteValue=false)
Returns the identity value associated with an AtomicRMWKind op.
arith::CmpIPredicate invertPredicate(arith::CmpIPredicate pred)
Invert an integer comparison predicate.
Value getZeroConstant(OpBuilder &builder, Location loc, Type type)
Creates an arith.constant operation with a zero value of type type.
detail::poison_attr_matcher m_Poison()
Matches a poison constant (any attribute implementing PoisonAttrInterface).
Include the generated interface declarations.
bool matchPattern(Value value, const Pattern &pattern)
Entry point for matching a pattern over a Value.
detail::constant_int_value_binder m_ConstantInt(IntegerAttr::ValueType *bind_value)
Matches a constant holding a scalar/vector/tensor integer (splat) and writes the integer value to bin...
detail::constant_float_predicate_matcher m_NaNFloat()
Matches a constant scalar / vector splat / tensor splat float ones.
LogicalResult verifyCompatibleShapes(TypeRange types1, TypeRange types2)
Returns success if the given two arrays have the same number of elements and each pair wise entries h...
Attribute constFoldCastOp(ArrayRef< Attribute > operands, Type resType, CalculationT &&calculate)
Attribute constFoldBinaryOp(ArrayRef< Attribute > operands, Type resultType, CalculationT &&calculate)
Type getType(OpFoldResult ofr)
Returns the int type of the integer in ofr.
detail::constant_int_range_predicate_matcher m_IntRangeWithoutNegOneS()
Matches a constant scalar / vector splat / tensor splat integer or a signed integer range that does n...
LogicalResult emitOptionalError(std::optional< Location > loc, Args &&...args)
Overloads of the above emission functions that take an optionally null location.
InFlightDiagnostic emitError(Location loc)
Utility method to emit an error message using this location.
detail::constant_float_predicate_matcher m_PosZeroFloat()
Matches a constant scalar / vector splat / tensor splat float positive zero.
detail::constant_int_predicate_matcher m_Zero()
Matches a constant scalar / vector splat / tensor splat integer zero.
Type getElementTypeOrSelf(Type type)
Return the element type or return the type itself.
detail::constant_float_predicate_matcher m_AnyZeroFloat()
Matches a constant scalar / vector splat / tensor splat float (both positive and negative) zero.
detail::constant_int_predicate_matcher m_One()
Matches a constant scalar / vector splat / tensor splat integer one.
detail::constant_float_predicate_matcher m_NegInfFloat()
Matches a constant scalar / vector splat / tensor splat float negative infinity.
detail::constant_float_predicate_matcher m_NegZeroFloat()
Matches a constant scalar / vector splat / tensor splat float negative zero.
detail::constant_int_range_predicate_matcher m_IntRangeWithoutZeroS()
Matches a constant scalar / vector splat / tensor splat integer or a signed integer range that does n...
detail::op_matcher< OpClass > m_Op()
Matches the given OpClass.
detail::constant_op_matcher m_Constant()
Matches a constant foldable operation.
Attribute constFoldUnaryOp(ArrayRef< Attribute > operands, Type resultType, CalculationT &&calculate)
detail::constant_float_predicate_matcher m_PosInfFloat()
Matches a constant scalar / vector splat / tensor splat float positive infinity.
llvm::function_ref< Fn > function_ref
detail::constant_float_predicate_matcher m_OneFloat()
Matches a constant scalar / vector splat / tensor splat float ones.
detail::constant_int_range_predicate_matcher m_IntRangeWithoutZeroU()
Matches a constant scalar / vector splat / tensor splat integer or a unsigned integer range that does...
LogicalResult matchAndRewrite(arith::SelectOp op, PatternRewriter &rewriter) const override
OpRewritePattern Base
Type alias to allow derived classes to inherit constructors with using Base::Base;.
OpRewritePattern(MLIRContext *context, PatternBenefit benefit=1, ArrayRef< StringRef > generatedNames={})
This represents an operation in an abstracted form, suitable for use with the builder APIs.