23#include "llvm/ADT/STLExtras.h"
24#include "llvm/ADT/SmallVectorExtras.h"
38 if (
auto boolAttr = dyn_cast<BoolAttr>(attr))
39 return boolAttr.getValue();
40 if (
auto splatAttr = dyn_cast<SplatElementsAttr>(attr))
41 if (splatAttr.getElementType().isInteger(1))
42 return splatAttr.getSplatValue<
bool>();
57 if (
auto vector = dyn_cast<ElementsAttr>(composite)) {
58 assert(
indices.size() == 1 &&
"must have exactly one index for a vector");
62 if (
auto array = dyn_cast<ArrayAttr>(composite)) {
63 assert(!
indices.empty() &&
"must have at least one index for an array");
72 bool div0 =
b.isZero();
73 bool overflow = a.isMinSignedValue() &&
b.isAllOnes();
75 return div0 || overflow;
83#include "SPIRVCanonicalization.inc"
93template <
typename AccessChainOp>
95 using OpRewritePattern<AccessChainOp>::OpRewritePattern;
97 LogicalResult matchAndRewrite(AccessChainOp accessChainOp,
98 PatternRewriter &rewriter)
const override {
99 auto parentAccessChainOp =
100 accessChainOp.getBasePtr().template getDefiningOp<AccessChainOp>();
102 if (!parentAccessChainOp) {
107 SmallVector<Value, 4>
indices(parentAccessChainOp.getIndices());
108 llvm::append_range(
indices, accessChainOp.getIndices());
111 accessChainOp, parentAccessChainOp.getBasePtr(),
indices);
118void spirv::AccessChainOp::getCanonicalizationPatterns(
120 results.
add<CombineChainedAccessChain<spirv::AccessChainOp>>(context);
123void spirv::InBoundsAccessChainOp::getCanonicalizationPatterns(
125 results.
add<CombineChainedAccessChain<spirv::InBoundsAccessChainOp>>(context);
132template <
typename Op>
136 static constexpr bool IsSub = std::is_same_v<Op, spirv::ISubBorrowOp>;
140 Value lhs = op.getOperand1();
141 Value rhs = op.getOperand2();
146 std::array<Value, 2> constituents =
147 IsSub ? std::array{lhs, rhs} : std::array{rhs, lhs};
161 [](
const APInt &a,
const APInt &
b) {
return IsSub ? a -
b : a +
b; });
166 {lhsAttr, rhsAttr}, [](
const APInt &a,
const APInt &
b) {
167 bool wrapped =
IsSub ? a.ult(
b) : (a +
b).ult(a);
168 return APInt(a.getBitWidth(), wrapped ? 1 : 0);
174 op, op.getType(), rewriter.
getArrayAttr({lowBits, wrapBit}));
180void spirv::IAddCarryOp::getCanonicalizationPatterns(
186void spirv::ISubBorrowOp::getCanonicalizationPatterns(
195template <
typename MulOp,
bool IsSigned>
202 Value lhs = op.getOperand1();
203 Value rhs = op.getOperand2();
204 Type constituentType = lhs.getType();
208 Value zero = spirv::ConstantOp::getZero(constituentType, loc, rewriter);
209 Value constituents[2] = {zero, zero};
231 [](
const APInt &a,
const APInt &
b) {
return a *
b; });
237 {lhsAttr, rhsAttr}, [](
const APInt &a,
const APInt &
b) {
239 return llvm::APIntOps::mulhs(a,
b);
241 return llvm::APIntOps::mulhu(a,
b);
248 op, op.getType(), rewriter.
getArrayAttr({lowBits, highBits}));
253template <
typename MulOp>
260 Value lhs = op.getOperand1();
261 Value rhs = op.getOperand2();
262 Type constituentType = lhs.getType();
266 Value zero = spirv::ConstantOp::getZero(constituentType, loc, rewriter);
267 Value constituents[2] = {lhs, zero};
279void spirv::SMulExtendedOp::getCanonicalizationPatterns(
286void spirv::UMulExtendedOp::getCanonicalizationPatterns(
309 auto prevUMod = umodOp.getOperand(0).getDefiningOp<spirv::UModOp>();
321 bool isApplicable =
false;
322 if (
auto prevInt = dyn_cast<IntegerAttr>(prevValue)) {
323 auto currInt = cast<IntegerAttr>(currValue);
324 if (currInt.getValue().isZero())
326 isApplicable = prevInt.getValue().urem(currInt.getValue()) == 0;
327 }
else if (
auto prevVec = dyn_cast<DenseElementsAttr>(prevValue)) {
328 auto currVec = cast<DenseElementsAttr>(currValue);
329 if (llvm::any_of(currVec.getValues<APInt>(),
330 [](
const APInt &curr) { return curr.isZero(); }))
332 isApplicable = llvm::all_of(llvm::zip_equal(prevVec.getValues<APInt>(),
333 currVec.getValues<APInt>()),
334 [](
const auto &pair) {
335 auto &[prev, curr] = pair;
336 return prev.urem(curr) == 0;
346 umodOp, umodOp.getType(), prevUMod.getOperand(0), umodOp.getOperand(1));
362 Value curInput = getOperand();
367 if (
auto prevCast = curInput.
getDefiningOp<spirv::BitcastOp>()) {
368 Value prevInput = prevCast.getOperand();
372 getOperandMutable().assign(prevInput);
384OpFoldResult spirv::CompositeExtractOp::fold(FoldAdaptor adaptor) {
385 Value compositeOp = getComposite();
387 while (
auto insertOp =
390 return insertOp.getObject();
391 compositeOp = insertOp.getComposite();
394 if (
auto constructOp =
396 auto type = cast<spirv::CompositeType>(constructOp.getType());
398 constructOp.getConstituents().size() == type.getNumElements()) {
399 auto i = cast<IntegerAttr>(*
getIndices().begin());
400 if (i.getValue().getSExtValue() <
401 static_cast<int64_t>(constructOp.getConstituents().size()))
402 return constructOp.getConstituents()[i.getValue().getSExtValue()];
407 return static_cast<unsigned>(cast<IntegerAttr>(attr).getInt());
427 return getOperand1();
435 adaptor.getOperands(),
436 [](APInt a,
const APInt &
b) { return std::move(a) + b; });
446 return getOperand2();
449 return getOperand1();
457 adaptor.getOperands(),
458 [](
const APInt &a,
const APInt &
b) { return a * b; });
467 if (getOperand1() == getOperand2())
476 adaptor.getOperands(),
477 [](APInt a,
const APInt &
b) { return std::move(a) - b; });
487 return getOperand1();
497 bool div0OrOverflow =
false;
499 adaptor.getOperands(), [&](
const APInt &a,
const APInt &
b) {
500 if (div0OrOverflow || isDivZeroOrOverflow(a, b)) {
501 div0OrOverflow = true;
506 return div0OrOverflow ?
Attribute() : res;
528 bool div0OrOverflow =
false;
530 adaptor.getOperands(), [&](
const APInt &a,
const APInt &
b) {
531 if (div0OrOverflow || isDivZeroOrOverflow(a, b)) {
532 div0OrOverflow = true;
535 APInt c = a.abs().urem(
b.abs());
538 if (
b.isNegative()) {
539 APInt zero = APInt::getZero(c.getBitWidth());
540 return a.isNegative() ? (std::move(zero) - c) : (b + std::move(c));
543 return b - std::move(c);
546 return div0OrOverflow ?
Attribute() : res;
568 bool div0OrOverflow =
false;
570 adaptor.getOperands(), [&](APInt a,
const APInt &
b) {
571 if (div0OrOverflow || isDivZeroOrOverflow(a, b)) {
572 div0OrOverflow = true;
577 return div0OrOverflow ?
Attribute() : res;
587 return getOperand1();
597 adaptor.getOperands(), [&](
const APInt &a,
const APInt &
b) {
598 if (div0 || b.isZero()) {
624 adaptor.getOperands(), [&](
const APInt &a,
const APInt &
b) {
625 if (div0 || b.isZero()) {
638OpFoldResult spirv::SNegateOp::fold(FoldAdaptor adaptor) {
640 auto op = getOperand();
641 if (
auto negateOp = op.getDefiningOp<spirv::SNegateOp>())
642 return negateOp->getOperand(0);
648 adaptor.getOperands(), [](
const APInt &a) {
649 APInt zero = APInt::getZero(a.getBitWidth());
650 return std::move(zero) - a;
658OpFoldResult spirv::NotOp::fold(spirv::NotOp::FoldAdaptor adaptor) {
660 auto op = getOperand();
661 if (
auto notOp = op.getDefiningOp<spirv::NotOp>())
662 return notOp->getOperand(0);
677OpFoldResult spirv::LogicalAndOp::fold(FoldAdaptor adaptor) {
678 if (std::optional<bool> rhs =
682 return getOperand1();
686 return adaptor.getOperand2();
697spirv::LogicalEqualOp::fold(spirv::LogicalEqualOp::FoldAdaptor adaptor) {
699 if (getOperand1() == getOperand2()) {
701 if (isa<IntegerType>(
getType()))
703 if (
auto vecTy = dyn_cast<VectorType>(
getType()))
708 adaptor.getOperands(), [](
const APInt &a,
const APInt &
b) {
709 return a == b ? APInt::getAllOnes(1) : APInt::getZero(1);
717OpFoldResult spirv::LogicalNotEqualOp::fold(FoldAdaptor adaptor) {
718 if (std::optional<bool> rhs =
722 return getOperand1();
726 if (getOperand1() == getOperand2()) {
728 if (isa<IntegerType>(
getType()))
730 if (
auto vecTy = dyn_cast<VectorType>(
getType()))
735 adaptor.getOperands(), [](
const APInt &a,
const APInt &
b) {
736 return a == b ? APInt::getZero(1) : APInt::getAllOnes(1);
744OpFoldResult spirv::LogicalNotOp::fold(FoldAdaptor adaptor) {
746 auto op = getOperand();
747 if (
auto notOp = op.getDefiningOp<spirv::LogicalNotOp>())
748 return notOp->getOperand(0);
754 adaptor.getOperands(), [](
const APInt &a) {
755 return a == 1 ? APInt::getZero(1) : APInt::getAllOnes(1);
759void spirv::LogicalNotOp::getCanonicalizationPatterns(
762 .
add<ConvertLogicalNotOfIEqual, ConvertLogicalNotOfINotEqual,
763 ConvertLogicalNotOfLogicalEqual, ConvertLogicalNotOfLogicalNotEqual>(
771OpFoldResult spirv::LogicalOrOp::fold(FoldAdaptor adaptor) {
775 return adaptor.getOperand2();
780 return getOperand1();
791OpFoldResult spirv::SelectOp::fold(FoldAdaptor adaptor) {
793 Value trueVals = getTrueValue();
794 Value falseVals = getFalseValue();
795 if (trueVals == falseVals)
803 return *boolAttr ? trueVals : falseVals;
806 if (!operands[0] || !operands[1] || !operands[2])
812 auto condAttrs = dyn_cast<DenseElementsAttr>(operands[0]);
813 auto trueAttrs = dyn_cast<DenseElementsAttr>(operands[1]);
814 auto falseAttrs = dyn_cast<DenseElementsAttr>(operands[2]);
815 if (!condAttrs || !trueAttrs || !falseAttrs)
818 auto elementResults = llvm::to_vector<4>(trueAttrs.getValues<
Attribute>());
819 auto iters = llvm::zip_equal(elementResults, condAttrs.getValues<
BoolAttr>(),
821 for (
auto [
result, cond, falseRes] : iters) {
822 if (!cond.getValue())
826 auto resultType = trueAttrs.getType();
834OpFoldResult spirv::IEqualOp::fold(spirv::IEqualOp::FoldAdaptor adaptor) {
836 if (getOperand1() == getOperand2()) {
838 if (isa<IntegerType>(
getType()))
840 if (
auto vecTy = dyn_cast<VectorType>(
getType()))
845 adaptor.getOperands(),
getType(), [](
const APInt &a,
const APInt &
b) {
846 return a ==
b ? APInt::getAllOnes(1) : APInt::
getZero(1);
854OpFoldResult spirv::INotEqualOp::fold(spirv::INotEqualOp::FoldAdaptor adaptor) {
856 if (getOperand1() == getOperand2()) {
858 if (isa<IntegerType>(
getType()))
860 if (
auto vecTy = dyn_cast<VectorType>(
getType()))
865 adaptor.getOperands(),
getType(), [](
const APInt &a,
const APInt &
b) {
866 return a ==
b ? APInt::getZero(1) : APInt::getAllOnes(1);
875spirv::SGreaterThanOp::fold(spirv::SGreaterThanOp::FoldAdaptor adaptor) {
877 if (getOperand1() == getOperand2()) {
879 if (isa<IntegerType>(
getType()))
881 if (
auto vecTy = dyn_cast<VectorType>(
getType()))
886 adaptor.getOperands(),
getType(), [](
const APInt &a,
const APInt &
b) {
887 return a.sgt(
b) ? APInt::getAllOnes(1) : APInt::
getZero(1);
896 spirv::SGreaterThanEqualOp::FoldAdaptor adaptor) {
898 if (getOperand1() == getOperand2()) {
900 if (isa<IntegerType>(
getType()))
902 if (
auto vecTy = dyn_cast<VectorType>(
getType()))
907 adaptor.getOperands(),
getType(), [](
const APInt &a,
const APInt &
b) {
908 return a.sge(
b) ? APInt::getAllOnes(1) : APInt::
getZero(1);
917spirv::UGreaterThanOp::fold(spirv::UGreaterThanOp::FoldAdaptor adaptor) {
919 if (getOperand1() == getOperand2()) {
921 if (isa<IntegerType>(
getType()))
923 if (
auto vecTy = dyn_cast<VectorType>(
getType()))
928 adaptor.getOperands(),
getType(), [](
const APInt &a,
const APInt &
b) {
929 return a.ugt(
b) ? APInt::getAllOnes(1) : APInt::
getZero(1);
938 spirv::UGreaterThanEqualOp::FoldAdaptor adaptor) {
940 if (getOperand1() == getOperand2()) {
942 if (isa<IntegerType>(
getType()))
944 if (
auto vecTy = dyn_cast<VectorType>(
getType()))
949 adaptor.getOperands(),
getType(), [](
const APInt &a,
const APInt &
b) {
950 return a.uge(
b) ? APInt::getAllOnes(1) : APInt::
getZero(1);
958OpFoldResult spirv::SLessThanOp::fold(spirv::SLessThanOp::FoldAdaptor adaptor) {
960 if (getOperand1() == getOperand2()) {
962 if (isa<IntegerType>(
getType()))
964 if (
auto vecTy = dyn_cast<VectorType>(
getType()))
969 adaptor.getOperands(),
getType(), [](
const APInt &a,
const APInt &
b) {
970 return a.slt(
b) ? APInt::getAllOnes(1) : APInt::
getZero(1);
979spirv::SLessThanEqualOp::fold(spirv::SLessThanEqualOp::FoldAdaptor adaptor) {
981 if (getOperand1() == getOperand2()) {
983 if (isa<IntegerType>(
getType()))
985 if (
auto vecTy = dyn_cast<VectorType>(
getType()))
990 adaptor.getOperands(),
getType(), [](
const APInt &a,
const APInt &
b) {
991 return a.sle(
b) ? APInt::getAllOnes(1) : APInt::
getZero(1);
999OpFoldResult spirv::ULessThanOp::fold(spirv::ULessThanOp::FoldAdaptor adaptor) {
1001 if (getOperand1() == getOperand2()) {
1003 if (isa<IntegerType>(
getType()))
1005 if (
auto vecTy = dyn_cast<VectorType>(
getType()))
1010 adaptor.getOperands(),
getType(), [](
const APInt &a,
const APInt &
b) {
1011 return a.ult(
b) ? APInt::getAllOnes(1) : APInt::
getZero(1);
1020spirv::ULessThanEqualOp::fold(spirv::ULessThanEqualOp::FoldAdaptor adaptor) {
1022 if (getOperand1() == getOperand2()) {
1024 if (isa<IntegerType>(
getType()))
1026 if (
auto vecTy = dyn_cast<VectorType>(
getType()))
1031 adaptor.getOperands(),
getType(), [](
const APInt &a,
const APInt &
b) {
1032 return a.ule(
b) ? APInt::getAllOnes(1) : APInt::
getZero(1);
1041 spirv::ShiftLeftLogicalOp::FoldAdaptor adaptor) {
1044 return getOperand1();
1055 bool shiftToLarge =
false;
1057 adaptor.getOperands(), [&](
const APInt &a,
const APInt &
b) {
1058 if (shiftToLarge || b.uge(a.getBitWidth())) {
1059 shiftToLarge = true;
1064 return shiftToLarge ?
Attribute() : res;
1072 spirv::ShiftRightArithmeticOp::FoldAdaptor adaptor) {
1075 return getOperand1();
1086 bool shiftToLarge =
false;
1088 adaptor.getOperands(), [&](
const APInt &a,
const APInt &
b) {
1089 if (shiftToLarge || b.uge(a.getBitWidth())) {
1090 shiftToLarge = true;
1095 return shiftToLarge ?
Attribute() : res;
1103 spirv::ShiftRightLogicalOp::FoldAdaptor adaptor) {
1106 return getOperand1();
1117 bool shiftToLarge =
false;
1119 adaptor.getOperands(), [&](
const APInt &a,
const APInt &
b) {
1120 if (shiftToLarge || b.uge(a.getBitWidth())) {
1121 shiftToLarge = true;
1126 return shiftToLarge ?
Attribute() : res;
1134spirv::BitwiseAndOp::fold(spirv::BitwiseAndOp::FoldAdaptor adaptor) {
1136 if (getOperand1() == getOperand2()) {
1137 return getOperand1();
1143 if (rhsMask.isZero())
1144 return getOperand2();
1147 if (rhsMask.isAllOnes())
1148 return getOperand1();
1151 if (
auto zext = getOperand1().getDefiningOp<spirv::UConvertOp>()) {
1154 if (rhsMask.zextOrTrunc(valueBits).isAllOnes())
1155 return getOperand1();
1165 adaptor.getOperands(),
1166 [](
const APInt &a,
const APInt &
b) { return a & b; });
1173OpFoldResult spirv::BitwiseOrOp::fold(spirv::BitwiseOrOp::FoldAdaptor adaptor) {
1175 if (getOperand1() == getOperand2()) {
1176 return getOperand1();
1182 if (rhsMask.isZero())
1183 return getOperand1();
1186 if (rhsMask.isAllOnes())
1187 return getOperand2();
1196 adaptor.getOperands(),
1197 [](
const APInt &a,
const APInt &
b) { return a | b; });
1205spirv::BitwiseXorOp::fold(spirv::BitwiseXorOp::FoldAdaptor adaptor) {
1208 return getOperand1();
1212 if (getOperand1() == getOperand2())
1221 adaptor.getOperands(),
1222 [](
const APInt &a,
const APInt &
b) { return a ^ b; });
1255struct ConvertSelectionOpToSelect final :
OpRewritePattern<spirv::SelectionOp> {
1258 LogicalResult matchAndRewrite(spirv::SelectionOp selectionOp,
1259 PatternRewriter &rewriter)
const override {
1260 Operation *op = selectionOp.getOperation();
1269 if (llvm::range_size(body) != 4) {
1273 Block *headerBlock = selectionOp.getHeaderBlock();
1274 if (!onlyContainsBranchConditionalOp(headerBlock)) {
1278 auto brConditionalOp =
1279 cast<spirv::BranchConditionalOp>(headerBlock->
front());
1281 Block *trueBlock = brConditionalOp.getSuccessor(0);
1282 Block *falseBlock = brConditionalOp.getSuccessor(1);
1283 Block *mergeBlock = selectionOp.getMergeBlock();
1285 if (
failed(canCanonicalizeSelection(trueBlock, falseBlock, mergeBlock)))
1288 Value trueValue = getSrcValue(trueBlock);
1289 Value falseValue = getSrcValue(falseBlock);
1290 Value ptrValue = getDstPtr(trueBlock);
1291 auto storeOp = cast<spirv::StoreOp>(trueBlock->
front());
1293 auto selectOp = spirv::SelectOp::create(
1294 rewriter, selectionOp.getLoc(), trueValue.
getType(),
1295 brConditionalOp.getCondition(), trueValue, falseValue);
1296 auto newStore = spirv::StoreOp::create(
1297 rewriter, selectOp.getLoc(), ptrValue, selectOp.getResult(),
1298 storeOp.getMemoryAccessAttr(), storeOp.getAlignmentAttr());
1299 newStore->setDiscardableAttrs(storeOp->getDiscardableAttrDictionary());
1313 LogicalResult canCanonicalizeSelection(
Block *trueBlock,
Block *falseBlock,
1314 Block *mergeBlock)
const;
1316 bool onlyContainsBranchConditionalOp(
Block *block)
const {
1317 return llvm::hasSingleElement(*block) &&
1318 isa<spirv::BranchConditionalOp>(block->
front());
1321 bool isSameAttrList(spirv::StoreOp
lhs, spirv::StoreOp
rhs)
const {
1322 return lhs->getDiscardableAttrDictionary() ==
1323 rhs->getDiscardableAttrDictionary() &&
1324 lhs.getProperties() ==
rhs.getProperties();
1328 Value getSrcValue(
Block *block)
const {
1329 auto storeOp = cast<spirv::StoreOp>(block->
front());
1330 return storeOp.getValue();
1334 Value getDstPtr(
Block *block)
const {
1335 auto storeOp = cast<spirv::StoreOp>(block->
front());
1336 return storeOp.getPtr();
1340LogicalResult ConvertSelectionOpToSelect::canCanonicalizeSelection(
1343 if (llvm::range_size(*trueBlock) != 2 || llvm::range_size(*falseBlock) != 2) {
1347 auto trueBrStoreOp = dyn_cast<spirv::StoreOp>(trueBlock->
front());
1348 auto trueBrBranchOp =
1349 dyn_cast<spirv::BranchOp>(*std::next(trueBlock->
begin()));
1350 auto falseBrStoreOp = dyn_cast<spirv::StoreOp>(falseBlock->
front());
1351 auto falseBrBranchOp =
1352 dyn_cast<spirv::BranchOp>(*std::next(falseBlock->
begin()));
1354 if (!trueBrStoreOp || !trueBrBranchOp || !falseBrStoreOp ||
1364 bool isScalarOrVector =
1365 cast<spirv::SPIRVType>(trueBrStoreOp.getValue().getType())
1366 .isScalarOrVector();
1370 if ((trueBrStoreOp.getPtr() != falseBrStoreOp.getPtr()) ||
1371 !isSameAttrList(trueBrStoreOp, falseBrStoreOp) || !isScalarOrVector) {
1375 if ((trueBrBranchOp->getSuccessor(0) != mergeBlock) ||
1376 (falseBrBranchOp->getSuccessor(0) != mergeBlock)) {
1384void spirv::SelectionOp::getCanonicalizationPatterns(RewritePatternSet &results,
1385 MLIRContext *context) {
1386 results.
add<ConvertSelectionOpToSelect>(context);
if(failed(verifyVectorMemoryOp(getOperation(), memrefType, getVectorType()))) return failure()
static Value getZero(OpBuilder &b, Location loc, Type elementType)
Get zero value for an element type.
static uint64_t zext(uint32_t arg)
ArithmeticExtendedBinaryFold< spirv::ISubBorrowOp > ISubBorrowFold
MulExtendedOpXOne< spirv::SMulExtendedOp > SMulExtendedOpXOne
static Attribute extractCompositeElement(Attribute composite, ArrayRef< unsigned > indices)
MulExtendedFold< spirv::UMulExtendedOp, false > UMulExtendedOpFold
MulExtendedOpXOne< spirv::UMulExtendedOp > UMulExtendedOpXOne
MulExtendedFold< spirv::SMulExtendedOp, true > SMulExtendedOpFold
static std::optional< bool > getScalarOrSplatBoolAttr(Attribute attr)
Returns the boolean value under the hood if the given boolAttr is a scalar or splat vector bool const...
static bool isDivZeroOrOverflow(const APInt &a, const APInt &b)
ArithmeticExtendedBinaryFold< spirv::IAddCarryOp > IAddCarryFold
Attributes are known-constant values of operations.
Block represents an ordered list of Operations.
Special case of IntegerAttr to represent boolean integers, i.e., signless i1 integers.
static BoolAttr get(MLIRContext *context, bool value)
This class is a general helper class for creating context-global objects like types,...
TypedAttr getZeroAttr(Type type)
ArrayAttr getArrayAttr(ArrayRef< Attribute > value)
static DenseElementsAttr get(ShapedType type, ArrayRef< Attribute > values)
Constructs a dense elements attribute from an array of element values.
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.
This class represents a single result from folding an operation.
This provides public APIs that all operations should have.
Region & getRegion(unsigned index)
Returns the region held by this operation at position 'index'.
A special type of RewriterBase that coordinates the application of a rewrite pattern on the current I...
RewritePatternSet & add(ConstructorArg &&arg, ConstructorArgs &&...args)
Add an instance of each of the pattern types 'Ts' to the pattern list with the given arguments.
virtual void eraseOp(Operation *op)
This method erases an operation that is known to have no uses.
OpTy replaceOpWithNewOp(Operation *op, Args &&...args)
Replace the results of the given (original) op with a new op that is created without verification (re...
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
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.
Operation::operand_range getIndices(Operation *op)
Get the indices that the given load/store operation is operating on.
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...
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_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_int_predicate_matcher m_One()
Matches a constant scalar / vector splat / tensor splat integer one.
detail::constant_op_matcher m_Constant()
Matches a constant foldable operation.
Attribute constFoldUnaryOp(ArrayRef< Attribute > operands, Type resultType, CalculationT &&calculate)
LogicalResult matchAndRewrite(Op op, PatternRewriter &rewriter) const override
static constexpr bool IsSub
LogicalResult matchAndRewrite(MulOp op, PatternRewriter &rewriter) const override
LogicalResult matchAndRewrite(MulOp op, PatternRewriter &rewriter) const override
LogicalResult matchAndRewrite(spirv::UModOp umodOp, PatternRewriter &rewriter) const override
OpRewritePattern is a wrapper around RewritePattern that allows for matching and rewriting against an...
OpRewritePattern Base
Type alias to allow derived classes to inherit constructors with using Base::Base;.
OpRewritePattern(MLIRContext *context, PatternBenefit benefit=1, ArrayRef< StringRef > generatedNames={})