11#include "llvm/ADT/TypeSwitch.h"
31#define GEN_PASS_DEF_ARITHINTRANGEOPTS
32#include "mlir/Dialect/Arith/Transforms/Passes.h.inc"
34#define GEN_PASS_DEF_ARITHINTRANGENARROWING
35#include "mlir/Dialect/Arith/Transforms/Passes.h.inc"
44 auto *maybeInferredRange =
46 if (!maybeInferredRange || maybeInferredRange->getValue().isUninitialized())
49 maybeInferredRange->getValue().getValue();
69 if (!maybeConstValue.has_value())
81 if (storageWidth != 0 && maybeConstValue->getBitWidth() != storageWidth)
87 maybeDefiningOp ? maybeDefiningOp->
getDialect()
91 if (
auto shaped = dyn_cast<ShapedType>(type)) {
122 void notifyOperationErased(Operation *op)
override {
136 MaterializeKnownConstantValues(MLIRContext *context, DataFlowSolver &s)
137 : RewritePattern::RewritePattern(Pattern::MatchAnyOpTypeTag(),
141 LogicalResult matchAndRewrite(Operation *op,
142 PatternRewriter &rewriter)
const override {
152 auto needsReplacing = [&](Value v) {
156 if (!maybeConstValue.has_value() || v.use_empty())
158 unsigned storageWidth =
160 return storageWidth == 0 ||
161 maybeConstValue->getBitWidth() == storageWidth;
163 bool hasConstantResults = llvm::any_of(op->
getResults(), needsReplacing);
165 if (!hasConstantResults)
167 bool hasConstantRegionArgs =
false;
169 for (
Block &block : region.getBlocks()) {
170 hasConstantRegionArgs |=
171 llvm::any_of(block.getArguments(), needsReplacing);
174 if (!hasConstantResults && !hasConstantRegionArgs)
187 PatternRewriter::InsertionGuard guard(rewriter);
189 for (
Block &block : region.getBlocks()) {
191 for (BlockArgument &arg : block.getArguments()) {
201 DataFlowSolver &solver;
204template <
typename RemOp>
206 DeleteTrivialRem(MLIRContext *context, DataFlowSolver &s)
207 : OpRewritePattern<RemOp>(context), solver(s) {}
209 LogicalResult matchAndRewrite(RemOp op,
210 PatternRewriter &rewriter)
const override {
211 Value
lhs = op.getOperand(0);
212 Value
rhs = op.getOperand(1);
218 bool isUnsigned = isa<RemUIOp>(op);
222 (!isUnsigned && modulus.isNegative()))
224 auto *maybeLhsRange = solver.lookupState<IntegerValueRangeLattice>(
lhs);
225 if (!maybeLhsRange || maybeLhsRange->getValue().isUninitialized())
227 const ConstantIntRanges &lhsRange = maybeLhsRange->getValue().getValue();
228 const APInt &
min = isUnsigned ? lhsRange.
umin() : lhsRange.
smin();
229 const APInt &
max = isUnsigned ? lhsRange.
umax() : lhsRange.
smax();
230 if (
min.getBitWidth() != modulus.getBitWidth() ||
231 max.getBitWidth() != modulus.getBitWidth())
235 if ((!isUnsigned &&
min.isNegative()) ||
min.uge(modulus))
237 if ((!isUnsigned &&
max.isNegative()) ||
max.uge(modulus))
250 DataFlowSolver &solver;
257 for (
Value val : values) {
258 auto *maybeInferredRange =
260 if (!maybeInferredRange || maybeInferredRange->getValue().isUninitialized())
264 maybeInferredRange->getValue().getValue();
265 ranges.push_back(inferredRange);
272static Type getTargetType(
Type srcType,
unsigned targetBitwidth) {
273 auto dstType = IntegerType::get(srcType.
getContext(), targetBitwidth);
274 if (
auto shaped = dyn_cast<ShapedType>(srcType))
275 return shaped.clone(dstType);
294 unsigned targetWidth) {
295 unsigned srcWidth = range.
smin().getBitWidth();
296 if (srcWidth <= targetWidth)
297 return CastKind::None;
298 unsigned removedWidth = srcWidth - targetWidth;
302 bool canTruncateSigned =
303 range.
smin().getNumSignBits() >= (removedWidth + 1) &&
304 range.
smax().getNumSignBits() >= (removedWidth + 1);
305 bool canTruncateUnsigned = range.
umin().countLeadingZeros() >= removedWidth &&
306 range.
umax().countLeadingZeros() >= removedWidth;
307 if (canTruncateSigned && canTruncateUnsigned)
308 return CastKind::Both;
309 if (canTruncateSigned)
310 return CastKind::Signed;
311 if (canTruncateUnsigned)
312 return CastKind::Unsigned;
313 return CastKind::None;
316static CastKind mergeCastKinds(CastKind lhs, CastKind rhs) {
317 if (lhs == CastKind::None || rhs == CastKind::None)
318 return CastKind::None;
319 if (lhs == CastKind::Both)
321 if (rhs == CastKind::Both)
325 return CastKind::None;
331 assert(isa<VectorType>(srcType) == isa<VectorType>(dstType) &&
332 "Mixing vector and non-vector types");
333 assert(castKind != CastKind::None &&
"Can't cast when casting isn't allowed");
336 assert(srcElemType.
isIntOrIndex() &&
"Invalid src type");
337 assert(dstElemType.
isIntOrIndex() &&
"Invalid dst type");
338 if (srcType == dstType)
341 if (isa<IndexType>(srcElemType) || isa<IndexType>(dstElemType)) {
342 if (castKind == CastKind::Signed)
343 return arith::IndexCastOp::create(builder, loc, dstType, src);
344 return arith::IndexCastUIOp::create(builder, loc, dstType, src);
347 auto srcInt = cast<IntegerType>(srcElemType);
348 auto dstInt = cast<IntegerType>(dstElemType);
349 if (dstInt.getWidth() < srcInt.getWidth())
350 return arith::TruncIOp::create(builder, loc, dstType, src);
352 if (castKind == CastKind::Signed)
353 return arith::ExtSIOp::create(builder, loc, dstType, src);
354 return arith::ExtUIOp::create(builder, loc, dstType, src);
358 NarrowElementwise(MLIRContext *context, DataFlowSolver &s,
359 ArrayRef<unsigned>
target)
360 : OpTraitRewritePattern(context), solver(s), targetBitwidths(
target) {}
363 LogicalResult matchAndRewrite(Operation *op,
364 PatternRewriter &rewriter)
const override {
370 SmallVector<ConstantIntRanges, 4> ranges;
381 [=](Type t) { return t == srcType; }))
383 op,
"no operands or operand types don't match result type");
385 for (
unsigned targetBitwidth : targetBitwidths) {
386 CastKind castKind = CastKind::Both;
387 for (
const ConstantIntRanges &range : ranges) {
388 castKind = mergeCastKinds(castKind,
389 checkTruncatability(range, targetBitwidth));
390 if (castKind == CastKind::None)
397 llvm::TypeSwitch<Operation *, CastKind>(op)
398 .Case<arith::DivSIOp, arith::CeilDivSIOp, arith::FloorDivSIOp,
399 arith::RemSIOp, arith::MaxSIOp, arith::MinSIOp,
400 arith::ShRSIOp>([](
auto) {
return CastKind::Signed; })
401 .Default(CastKind::Both);
402 castKind = mergeCastKinds(castKind, castKindForOp);
403 if (castKind == CastKind::None)
407 if (isa<arith::ShLIOp, arith::ShRSIOp, arith::ShRUIOp>(op) &&
408 !ranges[1].umax().
ult(targetBitwidth))
410 Type targetType = getTargetType(srcType, targetBitwidth);
411 if (targetType == srcType)
414 Location loc = op->
getLoc();
416 for (
auto [arg, argRange] : llvm::zip_first(op->
getOperands(), ranges)) {
417 CastKind argCastKind = castKind;
420 if (argCastKind == CastKind::Signed && argRange.smin().isNonNegative())
421 argCastKind = CastKind::Both;
422 Value newArg = doCast(rewriter, loc, arg, targetType, argCastKind);
423 mapping.
map(arg, newArg);
426 Operation *newOp = rewriter.
clone(*op, mapping);
432 SmallVector<Value> newResults;
433 for (
auto [newRes, oldRes] :
435 Value castBack = doCast(rewriter, loc, newRes, srcType, castKind);
437 newResults.push_back(castBack);
447 DataFlowSolver &solver;
448 SmallVector<unsigned, 4> targetBitwidths;
452 NarrowCmpI(MLIRContext *context, DataFlowSolver &s, ArrayRef<unsigned>
target)
453 : OpRewritePattern(context), solver(s), targetBitwidths(
target) {}
455 LogicalResult matchAndRewrite(arith::CmpIOp op,
456 PatternRewriter &rewriter)
const override {
457 Value
lhs = op.getLhs();
458 Value
rhs = op.getRhs();
460 SmallVector<ConstantIntRanges> ranges;
461 if (
failed(collectRanges(solver, op.getOperands(), ranges)))
463 const ConstantIntRanges &lhsRange = ranges[0];
464 const ConstantIntRanges &rhsRange = ranges[1];
466 auto isSignedCmpPredicate = [](arith::CmpIPredicate pred) ->
bool {
467 return pred == arith::CmpIPredicate::sge ||
468 pred == arith::CmpIPredicate::sgt ||
469 pred == arith::CmpIPredicate::sle ||
470 pred == arith::CmpIPredicate::slt;
474 CastKind predicateBasedCastRestriction =
475 isSignedCmpPredicate(op.getPredicate()) ? CastKind::Signed
478 Type srcType =
lhs.getType();
479 for (
unsigned targetBitwidth : targetBitwidths) {
480 CastKind lhsCastKind = checkTruncatability(lhsRange, targetBitwidth);
481 CastKind rhsCastKind = checkTruncatability(rhsRange, targetBitwidth);
482 CastKind castKind = mergeCastKinds(lhsCastKind, rhsCastKind);
483 castKind = mergeCastKinds(castKind, predicateBasedCastRestriction);
486 if (castKind == CastKind::None)
489 Type targetType = getTargetType(srcType, targetBitwidth);
490 if (targetType == srcType)
493 Location loc = op->getLoc();
495 Value lhsCast = doCast(rewriter, loc,
lhs, targetType, lhsCastKind);
496 Value rhsCast = doCast(rewriter, loc,
rhs, targetType, rhsCastKind);
497 mapping.
map(
lhs, lhsCast);
498 mapping.
map(
rhs, rhsCast);
500 Operation *newOp = rewriter.
clone(*op, mapping);
509 DataFlowSolver &solver;
510 SmallVector<unsigned, 4> targetBitwidths;
516template <
typename CastOp>
518 FoldIndexCastChain(MLIRContext *context, ArrayRef<unsigned>
target)
519 : OpRewritePattern<CastOp>(context), targetBitwidths(
target) {}
521 LogicalResult matchAndRewrite(CastOp op,
522 PatternRewriter &rewriter)
const override {
523 auto srcOp = op.getIn().template getDefiningOp<CastOp>();
527 Value src = srcOp.getIn();
528 if (src.
getType() != op.getType())
531 if (!srcOp.getType().isIndex())
534 auto intType = dyn_cast<IntegerType>(op.getType());
535 if (!intType || !llvm::is_contained(targetBitwidths, intType.getWidth()))
543 SmallVector<unsigned, 4> targetBitwidths;
547 NarrowLoopBounds(MLIRContext *context, DataFlowSolver &s,
548 ArrayRef<unsigned>
target)
549 : OpInterfaceRewritePattern<LoopLikeOpInterface>(context), solver(s),
551 boundsNarrowingFailedAttr(
552 StringAttr::
get(context,
"arith.bounds_narrowing_failed")) {}
554 LogicalResult matchAndRewrite(LoopLikeOpInterface loopLike,
555 PatternRewriter &rewriter)
const override {
557 if (loopLike->hasDiscardableAttr(boundsNarrowingFailedAttr))
559 "bounds narrowing previously failed");
561 std::optional<SmallVector<Value>> inductionVars =
562 loopLike.getLoopInductionVars();
563 if (!inductionVars.has_value() || inductionVars->empty())
566 std::optional<SmallVector<OpFoldResult>> lowerBounds =
567 loopLike.getLoopLowerBounds();
568 std::optional<SmallVector<OpFoldResult>> upperBounds =
569 loopLike.getLoopUpperBounds();
570 std::optional<SmallVector<OpFoldResult>> steps = loopLike.getLoopSteps();
572 if (!lowerBounds.has_value() || !upperBounds.has_value() ||
576 if (lowerBounds->size() != inductionVars->size() ||
577 upperBounds->size() != inductionVars->size() ||
578 steps->size() != inductionVars->size())
580 "mismatched bounds/steps count");
582 Location loc = loopLike->getLoc();
583 SmallVector<OpFoldResult> newLowerBounds(*lowerBounds);
584 SmallVector<OpFoldResult> newUpperBounds(*upperBounds);
585 SmallVector<OpFoldResult> newSteps(*steps);
586 SmallVector<std::tuple<size_t, Type, CastKind>> narrowings;
589 for (
auto [idx, indVar, lbOFR, ubOFR, stepOFR] :
590 llvm::enumerate(*inductionVars, *lowerBounds, *upperBounds, *steps)) {
593 auto maybeLb = dyn_cast<Value>(lbOFR);
594 auto maybeUb = dyn_cast<Value>(ubOFR);
595 auto maybeStep = dyn_cast<Value>(stepOFR);
597 if (!maybeLb || !maybeUb || !maybeStep)
601 SmallVector<ConstantIntRanges> ranges;
603 solver,
ValueRange{maybeLb, maybeUb, maybeStep, indVar}, ranges)))
606 const ConstantIntRanges &stepRange = ranges[2];
607 const ConstantIntRanges &indVarRange = ranges[3];
609 Type srcType = maybeLb.getType();
612 for (
unsigned targetBitwidth : targetBitwidths) {
613 Type targetType = getTargetType(srcType, targetBitwidth);
614 if (targetType == srcType)
619 if (!loopLike.isValidInductionVarType(targetType))
623 CastKind castKind = CastKind::Both;
624 for (
const ConstantIntRanges &range : ranges) {
625 castKind = mergeCastKinds(castKind,
626 checkTruncatability(range, targetBitwidth));
627 if (castKind == CastKind::None)
631 if (castKind == CastKind::None)
641 ConstantIntRanges indVarPlusStepRange(
642 indVarRange.
smin().sadd_sat(stepRange.
smin()),
643 indVarRange.
smax().sadd_sat(stepRange.
smax()),
644 indVarRange.
umin().uadd_sat(stepRange.
umin()),
645 indVarRange.
umax().uadd_sat(stepRange.
umax()));
647 if (checkTruncatability(indVarPlusStepRange, targetBitwidth) !=
652 Value newLb = doCast(rewriter, loc, maybeLb, targetType, castKind);
653 Value newUb = doCast(rewriter, loc, maybeUb, targetType, castKind);
654 Value newStep = doCast(rewriter, loc, maybeStep, targetType, castKind);
656 newLowerBounds[idx] = newLb;
657 newUpperBounds[idx] = newUb;
658 newSteps[idx] = newStep;
659 narrowings.push_back({idx, targetType, castKind});
664 if (narrowings.empty())
668 SmallVector<Type> origTypes;
669 for (
auto [idx, targetType, castKind] : narrowings) {
670 Value indVar = (*inductionVars)[idx];
671 origTypes.push_back(indVar.
getType());
676 bool updateFailed =
false;
679 if (
failed(loopLike.setLoopLowerBounds(newLowerBounds)) ||
680 failed(loopLike.setLoopUpperBounds(newUpperBounds)) ||
681 failed(loopLike.setLoopSteps(newSteps))) {
684 loopLike->setDiscardableAttr(boundsNarrowingFailedAttr,
691 for (
auto [idx, targetType, castKind] : narrowings) {
692 Value indVar = (*inductionVars)[idx];
693 auto blockArg = cast<BlockArgument>(indVar);
696 blockArg.setType(targetType);
704 for (
auto [narrowingIdx, narrowingInfo] : llvm::enumerate(narrowings)) {
705 auto [idx, targetType, castKind] = narrowingInfo;
706 Value indVar = (*inductionVars)[idx];
707 auto blockArg = cast<BlockArgument>(indVar);
708 Type origType = origTypes[narrowingIdx];
710 OpBuilder::InsertionGuard guard(rewriter);
712 Value casted = doCast(rewriter, loc, blockArg, origType, castKind);
723 DataFlowSolver &solver;
724 SmallVector<unsigned, 4> targetBitwidths;
725 StringAttr boundsNarrowingFailedAttr;
728struct IntRangeOptimizationsPass final
729 : arith::impl::ArithIntRangeOptsBase<IntRangeOptimizationsPass> {
731 void runOnOperation()
override {
732 Operation *op = getOperation();
734 DataFlowSolver solver;
736 solver.
load<IntegerRangeAnalysis>();
738 return signalPassFailure();
740 DataFlowListener listener(solver);
742 RewritePatternSet patterns(ctx);
753 GreedyRewriteConfig()
754 .enableFolding(
false)
755 .setRegionSimplificationLevel(
756 GreedySimplifyRegionLevel::Disabled)
757 .setListener(&listener))))
762struct IntRangeNarrowingPass final
763 : arith::impl::ArithIntRangeNarrowingBase<IntRangeNarrowingPass> {
764 using ArithIntRangeNarrowingBase::ArithIntRangeNarrowingBase;
766 void runOnOperation()
override {
767 Operation *op = getOperation();
769 DataFlowSolver solver;
771 solver.
load<IntegerRangeAnalysis>();
773 return signalPassFailure();
775 DataFlowListener listener(solver);
777 RewritePatternSet patterns(ctx);
785 op, std::move(patterns),
786 GreedyRewriteConfig().setUseTopDownTraversal(
false).setListener(
795 patterns.
add<MaterializeKnownConstantValues, DeleteTrivialRem<RemSIOp>,
796 DeleteTrivialRem<RemUIOp>>(patterns.
getContext(), solver);
802 patterns.
add<NarrowElementwise, NarrowCmpI>(patterns.
getContext(), solver,
804 patterns.
add<FoldIndexCastChain<arith::IndexCastUIOp>,
805 FoldIndexCastChain<arith::IndexCastOp>>(patterns.
getContext(),
812 patterns.
add<NarrowLoopBounds>(patterns.
getContext(), solver,
817 return std::make_unique<IntRangeOptimizationsPass>();
static Operation * materializeConstant(Dialect *dialect, OpBuilder &builder, Attribute value, Type type, Location loc)
A utility function used to materialize a constant for a given attribute and type.
static void copyIntegerRange(DataFlowSolver &solver, Value oldVal, Value newVal)
static std::optional< APInt > getMaybeConstantValue(DataFlowSolver &solver, Value value)
static Value max(ImplicitLocOpBuilder &builder, Value value, Value bound)
static Value min(ImplicitLocOpBuilder &builder, Value value, Value bound)
Attributes are known-constant values of operations.
IntegerAttr getIntegerAttr(Type type, int64_t value)
MLIRContext * getContext() const
A set of arbitrary-precision integers representing bounds on a given integer value.
const APInt & smax() const
The maximum value of an integer when it is interpreted as signed.
const APInt & smin() const
The minimum value of an integer when it is interpreted as signed.
static unsigned getStorageBitwidth(Type type)
Return the bitwidth that should be used for integer ranges describing type.
std::optional< APInt > getConstantValue() const
If either the signed or unsigned interpretations of the range indicate that the value it bounds is a ...
const APInt & umax() const
The maximum value of an integer when it is interpreted as unsigned.
const APInt & umin() const
The minimum value of an integer when it is interpreted as unsigned.
The general data-flow analysis solver.
LogicalResult initializeAndRun(Operation *top, llvm::function_ref< bool(DataFlowAnalysis &)> analysisFilter=nullptr)
Initialize analyses starting from the provided top-level operation and run the analysis until fixpoin...
void eraseState(AnchorT anchor)
Erase any analysis state associated with the given lattice anchor.
const StateT * lookupState(AnchorT anchor) const
Lookup an analysis state for the given lattice anchor.
StateT * getOrCreateState(AnchorT anchor)
Get the state associated with the given lattice anchor.
AnalysisT * load(Args &&...args)
Load an analysis into the solver. Return the analysis instance.
static DenseIntElementsAttr get(const ShapedType &type, Arg &&arg)
Get an instance of a DenseIntElementsAttr with the given arguments.
Dialects are groups of MLIR operations, types and attributes, as well as behavior associated with the...
virtual Operation * materializeConstant(OpBuilder &builder, Attribute value, Type type, Location loc)
Registered hook to materialize a single constant operation from a given attribute value with the desi...
void map(Value from, Value to)
Inserts a new mapping for 'from' to 'to'.
This class defines the main interface for locations in MLIR and acts as a non-nullable wrapper around...
Dialect * getLoadedDialect(StringRef name)
Get a registered IR dialect with the given namespace.
This class helps build Operations.
Operation * clone(Operation &op, IRMapping &mapper)
Creates a deep copy of the specified operation, remapping any operands that use values outside of the...
void setInsertionPointToStart(Block *block)
Sets the insertion point to the start of the specified block.
This is a value defined by a result of an operation.
OpTraitRewritePattern is a wrapper around RewritePattern that allows for matching and rewriting again...
OpTraitRewritePattern(MLIRContext *context, PatternBenefit benefit=1)
Operation is the basic unit of execution within MLIR.
Dialect * getDialect()
Return the dialect this operation is associated with, or nullptr if the associated dialect is not loa...
OpResult getResult(unsigned idx)
Get the 'idx'th result of this operation.
unsigned getNumRegions()
Returns the number of regions held by this operation.
Location getLoc()
The source location the operation was defined or derived from.
unsigned getNumOperands()
operand_type_range getOperandTypes()
MutableArrayRef< Region > getRegions()
Returns the regions held by this operation.
result_type_range getResultTypes()
operand_range getOperands()
Returns an iterator on the underlying Value's.
result_range getResults()
MLIRContext * getContext()
Return the context this operation is associated with.
unsigned getNumResults()
Return the number of results held by this operation.
Operation * getParentOp()
Return the parent operation this region is attached to.
MLIRContext * getContext() const
RewritePatternSet & add(ConstructorArg &&arg, ConstructorArgs &&...args)
Add an instance of each of the pattern types 'Ts' to the pattern list with the given arguments.
RewritePattern is the common base class for all DAG to DAG replacements.
This class coordinates the application of a rewrite on a set of IR, providing a way for clients to tr...
virtual void replaceOp(Operation *op, ValueRange newValues)
Replace the results of the given (original) operation with the specified list of values (replacements...
virtual void eraseOp(Operation *op)
This method erases an operation that is known to have no uses.
void replaceAllUsesExcept(Value from, Value to, Operation *exceptedUser)
Find uses of from and replace them with to except if the user is exceptedUser.
std::enable_if_t<!std::is_convertible< CallbackT, Twine >::value, LogicalResult > notifyMatchFailure(Location loc, CallbackT &&reasonCallback)
Used to notify the listener that the IR failed to be rewritten because of a match failure,...
void modifyOpInPlace(Operation *root, CallableT &&callable)
This method is a utility wrapper around an in-place modification of an operation.
virtual void replaceAllUsesWith(Value from, Value to)
Find uses of from and replace them with to.
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 isIntOrIndex() const
Return true if this is an integer (of any signedness) or an index type.
This class provides an abstraction over the different types of ranges over Values.
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
bool use_empty() const
Returns true if this value has no uses.
void setType(Type newType)
Mutate the type of this Value to be of the specified type.
Type getType() const
Return the type of this value.
Location getLoc() const
Return the location of this value.
Operation * getDefiningOp() const
If this value is the result of an operation, return the operation that defines it.
Region * getParentRegion()
Return the Region in which this Value is defined.
This lattice element represents the integer value range of an SSA value.
ChangeResult join(const AbstractSparseLattice &rhs) override
Join the information contained in 'rhs' into this lattice.
std::unique_ptr< Pass > createIntRangeOptimizationsPass()
Create a pass which do optimizations based on integer range analysis.
void populateControlFlowValuesNarrowingPatterns(RewritePatternSet &patterns, DataFlowSolver &solver, ArrayRef< unsigned > bitwidthsSupported)
Add patterns for narrowing control flow values (loop bounds, steps, etc.) based on int range analysis...
void populateIntRangeOptimizationsPatterns(RewritePatternSet &patterns, DataFlowSolver &solver)
Add patterns for int range based optimizations.
void populateIntRangeNarrowingPatterns(RewritePatternSet &patterns, DataFlowSolver &solver, ArrayRef< unsigned > bitwidthsSupported)
Add patterns for int range based narrowing.
LogicalResult maybeReplaceWithConstant(DataFlowSolver &solver, RewriterBase &rewriter, Value value)
Patterned after SCCP.
void loadBaselineAnalyses(DataFlowSolver &solver)
Populates a DataFlowSolver with analyses that are required to ensure user-defined analyses are run pr...
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...
LogicalResult applyPatternsGreedily(Region ®ion, const FrozenRewritePatternSet &patterns, GreedyRewriteConfig config=GreedyRewriteConfig(), bool *changed=nullptr)
Rewrite ops in the given region, which must be isolated from above, by repeatedly applying the highes...
Type getElementTypeOrSelf(Type type)
Return the element type or return the type itself.
bool isOpTriviallyDead(Operation *op)
Return true if the given operation is unused, and has no side effects on memory that prevent erasing.
auto get(MLIRContext *context, Ts &&...params)
Helper method that injects context only if needed, this helps unify some of the attribute constructio...
detail::constant_op_matcher m_Constant()
Matches a constant foldable operation.
OpInterfaceRewritePattern is a wrapper around RewritePattern that allows for matching and rewriting a...
OpRewritePattern is a wrapper around RewritePattern that allows for matching and rewriting against an...