MLIR 24.0.0git
IntRangeOptimizations.cpp
Go to the documentation of this file.
1//===- IntRangeOptimizations.cpp - Optimizations based on integer ranges --===//
2//
3// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.
4// See https://llvm.org/LICENSE.txt for license information.
5// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
6//
7//===----------------------------------------------------------------------===//
8
9#include <utility>
10
11#include "llvm/ADT/TypeSwitch.h"
12
17
21#include "mlir/IR/IRMapping.h"
22#include "mlir/IR/Matchers.h"
29
30namespace mlir::arith {
31#define GEN_PASS_DEF_ARITHINTRANGEOPTS
32#include "mlir/Dialect/Arith/Transforms/Passes.h.inc"
33
34#define GEN_PASS_DEF_ARITHINTRANGENARROWING
35#include "mlir/Dialect/Arith/Transforms/Passes.h.inc"
36} // namespace mlir::arith
37
38using namespace mlir;
39using namespace mlir::arith;
40using namespace mlir::dataflow;
41
42static std::optional<APInt> getMaybeConstantValue(DataFlowSolver &solver,
43 Value value) {
44 auto *maybeInferredRange =
46 if (!maybeInferredRange || maybeInferredRange->getValue().isUninitialized())
47 return std::nullopt;
48 const ConstantIntRanges &inferredRange =
49 maybeInferredRange->getValue().getValue();
50 return inferredRange.getConstantValue();
51}
52
53static void copyIntegerRange(DataFlowSolver &solver, Value oldVal,
54 Value newVal) {
55 auto *oldState = solver.lookupState<IntegerValueRangeLattice>(oldVal);
56 if (!oldState)
57 return;
59 *oldState);
60}
61
62namespace mlir::dataflow {
63/// Patterned after SCCP
65 RewriterBase &rewriter, Value value) {
66 if (value.use_empty())
67 return failure();
68 std::optional<APInt> maybeConstValue = getMaybeConstantValue(solver, value);
69 if (!maybeConstValue.has_value())
70 return failure();
71
72 Type type = value.getType();
73 // If the type or element type is non-integral, the attribute constructor
74 // will crash, so eagerly check for an integer type to avoid this.
75 if (!getElementTypeOrSelf(type).isIntOrIndex())
76 return failure();
77
78 // Bail out if the inferred APInt bitwidth does not match the storage width
79 // of the IR type; IntegerAttr::get would assert otherwise.
80 unsigned storageWidth = ConstantIntRanges::getStorageBitwidth(type);
81 if (storageWidth != 0 && maybeConstValue->getBitWidth() != storageWidth)
82 return failure();
83
84 Location loc = value.getLoc();
85 Operation *maybeDefiningOp = value.getDefiningOp();
86 Dialect *valueDialect =
87 maybeDefiningOp ? maybeDefiningOp->getDialect()
89
90 Attribute constAttr;
91 if (auto shaped = dyn_cast<ShapedType>(type)) {
92 constAttr = mlir::DenseIntElementsAttr::get(shaped, *maybeConstValue);
93 } else {
94 constAttr = rewriter.getIntegerAttr(type, *maybeConstValue);
95 }
96 Operation *constOp =
97 valueDialect->materializeConstant(rewriter, constAttr, type, loc);
98 // Fall back to arith.constant if the dialect materializer doesn't know what
99 // to do with an integer constant.
100 if (!constOp)
101 constOp = rewriter.getContext()
102 ->getLoadedDialect<ArithDialect>()
103 ->materializeConstant(rewriter, constAttr, type, loc);
104 if (!constOp)
105 return failure();
106
107 OpResult res = constOp->getResult(0);
109 solver.eraseState(res);
110 copyIntegerRange(solver, value, res);
111 rewriter.replaceAllUsesWith(value, res);
112 return success();
113}
114} // namespace mlir::dataflow
115
116namespace {
117class DataFlowListener : public RewriterBase::Listener {
118public:
119 DataFlowListener(DataFlowSolver &s) : s(s) {}
120
121protected:
122 void notifyOperationErased(Operation *op) override {
123 s.eraseState(s.getProgramPointAfter(op));
124 for (Value res : op->getResults())
125 s.eraseState(res);
126 }
127
128 DataFlowSolver &s;
129};
130
131/// Rewrite any results of `op` that were inferred to be constant integers to
132/// and replace their uses with that constant. Return success() if all results
133/// where thus replaced and the operation is erased. Also replace any block
134/// arguments with their constant values.
135struct MaterializeKnownConstantValues : public RewritePattern {
136 MaterializeKnownConstantValues(MLIRContext *context, DataFlowSolver &s)
137 : RewritePattern::RewritePattern(Pattern::MatchAnyOpTypeTag(),
138 /*benefit=*/1, context),
139 solver(s) {}
140
141 LogicalResult matchAndRewrite(Operation *op,
142 PatternRewriter &rewriter) const override {
143 if (matchPattern(op, m_Constant()))
144 return failure();
145
146 // We need to check isIntOrIndex() and APInt bitwidth compatibility here
147 // as well to avoid infinite loops in the greedy pattern rewriter. If we
148 // only check in maybeReplaceWithConstant, this lambda might still return
149 // true for values that cannot be materialized, causing the pattern to
150 // match and claim success without making any changes, leading to
151 // non-convergence.
152 auto needsReplacing = [&](Value v) {
153 if (!getElementTypeOrSelf(v.getType()).isIntOrIndex())
154 return false;
155 std::optional<APInt> maybeConstValue = getMaybeConstantValue(solver, v);
156 if (!maybeConstValue.has_value() || v.use_empty())
157 return false;
158 unsigned storageWidth =
160 return storageWidth == 0 ||
161 maybeConstValue->getBitWidth() == storageWidth;
162 };
163 bool hasConstantResults = llvm::any_of(op->getResults(), needsReplacing);
164 if (op->getNumRegions() == 0)
165 if (!hasConstantResults)
166 return failure();
167 bool hasConstantRegionArgs = false;
168 for (Region &region : op->getRegions()) {
169 for (Block &block : region.getBlocks()) {
170 hasConstantRegionArgs |=
171 llvm::any_of(block.getArguments(), needsReplacing);
172 }
173 }
174 if (!hasConstantResults && !hasConstantRegionArgs)
175 return failure();
176
177 bool replacedAll = (op->getNumResults() != 0);
178 for (Value v : op->getResults())
179 replacedAll &=
180 (succeeded(maybeReplaceWithConstant(solver, rewriter, v)) ||
181 v.use_empty());
182 if (replacedAll && isOpTriviallyDead(op)) {
183 rewriter.eraseOp(op);
184 return success();
185 }
186
187 PatternRewriter::InsertionGuard guard(rewriter);
188 for (Region &region : op->getRegions()) {
189 for (Block &block : region.getBlocks()) {
190 rewriter.setInsertionPointToStart(&block);
191 for (BlockArgument &arg : block.getArguments()) {
192 (void)maybeReplaceWithConstant(solver, rewriter, arg);
193 }
194 }
195 }
196
197 return success();
198 }
199
200private:
201 DataFlowSolver &solver;
202};
203
204template <typename RemOp>
205struct DeleteTrivialRem : public OpRewritePattern<RemOp> {
206 DeleteTrivialRem(MLIRContext *context, DataFlowSolver &s)
207 : OpRewritePattern<RemOp>(context), solver(s) {}
208
209 LogicalResult matchAndRewrite(RemOp op,
210 PatternRewriter &rewriter) const override {
211 Value lhs = op.getOperand(0);
212 Value rhs = op.getOperand(1);
213 // TODO: Support index types once integer range inference can use the target
214 // index bitwidth, e.g. from DLTI.
215 if (isa<IndexType>(getElementTypeOrSelf(lhs.getType())))
216 return failure();
217 APInt modulus;
218 bool isUnsigned = isa<RemUIOp>(op);
219 // Any nonzero bit pattern is a valid unsigned modulus. Keep the existing
220 // positive-modulus restriction for signed remainder.
221 if (!matchPattern(rhs, m_ConstantInt(&modulus)) || modulus.isZero() ||
222 (!isUnsigned && modulus.isNegative()))
223 return failure();
224 auto *maybeLhsRange = solver.lookupState<IntegerValueRangeLattice>(lhs);
225 if (!maybeLhsRange || maybeLhsRange->getValue().isUninitialized())
226 return failure();
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())
232 return failure();
233 // The minima and maxima here are given as closed ranges, we must be
234 // non-negative for signed remainder and strictly less than the modulus.
235 if ((!isUnsigned && min.isNegative()) || min.uge(modulus))
236 return failure();
237 if ((!isUnsigned && max.isNegative()) || max.uge(modulus))
238 return failure();
239 if (!min.ule(max))
240 return failure();
241
242 // With all those conditions out of the way, we know thas this invocation of
243 // a remainder is a noop because the input is strictly within the range
244 // [0, modulus), so get rid of it.
245 rewriter.replaceOp(op, ValueRange{lhs});
246 return success();
247 }
248
249private:
250 DataFlowSolver &solver;
251};
252
253/// Gather ranges for all the values in `values`. Appends to the existing
254/// vector.
255static LogicalResult collectRanges(DataFlowSolver &solver, ValueRange values,
257 for (Value val : values) {
258 auto *maybeInferredRange =
260 if (!maybeInferredRange || maybeInferredRange->getValue().isUninitialized())
261 return failure();
262
263 const ConstantIntRanges &inferredRange =
264 maybeInferredRange->getValue().getValue();
265 ranges.push_back(inferredRange);
266 }
267 return success();
268}
269
270/// Return int type truncated to `targetBitwidth`. If `srcType` is shaped,
271/// return shaped type as well.
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);
276
277 assert(srcType.isIntOrIndex() && "Invalid src type");
278 return dstType;
279}
280
281namespace {
282// Enum for tracking which type of truncation should be performed
283// to narrow an operation, if any.
284enum class CastKind : uint8_t { None, Signed, Unsigned, Both };
285} // namespace
286
287/// If the values within `range` can be represented using only `width` bits,
288/// return the kind of truncation needed to preserve that property.
289///
290/// This check relies on the fact that the signed and unsigned ranges are both
291/// always correct, but that one might be an approximation of the other,
292/// so we want to use the correct truncation operation.
293static CastKind checkTruncatability(const ConstantIntRanges &range,
294 unsigned targetWidth) {
295 unsigned srcWidth = range.smin().getBitWidth();
296 if (srcWidth <= targetWidth)
297 return CastKind::None;
298 unsigned removedWidth = srcWidth - targetWidth;
299 // The sign bits need to extend into the sign bit of the target width. For
300 // example, if we're truncating 64 bits to 32, we need 64 - 32 + 1 = 33 sign
301 // bits.
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;
314}
315
316static CastKind mergeCastKinds(CastKind lhs, CastKind rhs) {
317 if (lhs == CastKind::None || rhs == CastKind::None)
318 return CastKind::None;
319 if (lhs == CastKind::Both)
320 return rhs;
321 if (rhs == CastKind::Both)
322 return lhs;
323 if (lhs == rhs)
324 return lhs;
325 return CastKind::None;
326}
327
328static Value doCast(OpBuilder &builder, Location loc, Value src, Type dstType,
329 CastKind castKind) {
330 Type srcType = src.getType();
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");
334 Type srcElemType = getElementTypeOrSelf(srcType);
335 Type dstElemType = getElementTypeOrSelf(dstType);
336 assert(srcElemType.isIntOrIndex() && "Invalid src type");
337 assert(dstElemType.isIntOrIndex() && "Invalid dst type");
338 if (srcType == dstType)
339 return src;
340
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);
345 }
346
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);
351
352 if (castKind == CastKind::Signed)
353 return arith::ExtSIOp::create(builder, loc, dstType, src);
354 return arith::ExtUIOp::create(builder, loc, dstType, src);
355}
356
357struct NarrowElementwise final : OpTraitRewritePattern<OpTrait::Elementwise> {
358 NarrowElementwise(MLIRContext *context, DataFlowSolver &s,
359 ArrayRef<unsigned> target)
360 : OpTraitRewritePattern(context), solver(s), targetBitwidths(target) {}
361
363 LogicalResult matchAndRewrite(Operation *op,
364 PatternRewriter &rewriter) const override {
365 if (op->getNumResults() == 0)
366 return rewriter.notifyMatchFailure(op, "can't narrow resultless op");
367
368 // Inline size chosen empirically based on compilation profiling.
369 // Profiled: 2.6M calls, avg=1.7+-1.3. N=4 covers >95% of cases inline.
370 SmallVector<ConstantIntRanges, 4> ranges;
371 if (failed(collectRanges(solver, op->getOperands(), ranges)))
372 return rewriter.notifyMatchFailure(op, "input without specified range");
373 if (failed(collectRanges(solver, op->getResults(), ranges)))
374 return rewriter.notifyMatchFailure(op, "output without specified range");
375
376 Type srcType = op->getResult(0).getType();
377 if (!llvm::all_equal(op->getResultTypes()))
378 return rewriter.notifyMatchFailure(op, "mismatched result types");
379 if (op->getNumOperands() == 0 ||
380 !llvm::all_of(op->getOperandTypes(),
381 [=](Type t) { return t == srcType; }))
382 return rewriter.notifyMatchFailure(
383 op, "no operands or operand types don't match result type");
384
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)
391 break;
392 }
393 // For operations that explicitly treat the values as signed, we should
394 // only do signed casts, if those are deemed possible as such based on the
395 // value range.
396 auto castKindForOp =
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)
404 continue;
405 // A shift by an amount >= the bitwidth is poison, so only narrow shifts
406 // when the shift amount (second operand) stays below the target width.
407 if (isa<arith::ShLIOp, arith::ShRSIOp, arith::ShRUIOp>(op) &&
408 !ranges[1].umax().ult(targetBitwidth))
409 continue;
410 Type targetType = getTargetType(srcType, targetBitwidth);
411 if (targetType == srcType)
412 continue;
413
414 Location loc = op->getLoc();
415 IRMapping mapping;
416 for (auto [arg, argRange] : llvm::zip_first(op->getOperands(), ranges)) {
417 CastKind argCastKind = castKind;
418 // When dealing with `index` values, preserve non-negativity in the
419 // index_casts since we can't recover this in unsigned when equivalent.
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);
424 }
425
426 Operation *newOp = rewriter.clone(*op, mapping);
427 rewriter.modifyOpInPlace(newOp, [&]() {
428 for (OpResult res : newOp->getResults()) {
429 res.setType(targetType);
430 }
431 });
432 SmallVector<Value> newResults;
433 for (auto [newRes, oldRes] :
434 llvm::zip_equal(newOp->getResults(), op->getResults())) {
435 Value castBack = doCast(rewriter, loc, newRes, srcType, castKind);
436 copyIntegerRange(solver, oldRes, castBack);
437 newResults.push_back(castBack);
438 }
439
440 rewriter.replaceOp(op, newResults);
441 return success();
442 }
443 return failure();
444 }
445
446private:
447 DataFlowSolver &solver;
448 SmallVector<unsigned, 4> targetBitwidths;
449};
450
451struct NarrowCmpI final : OpRewritePattern<arith::CmpIOp> {
452 NarrowCmpI(MLIRContext *context, DataFlowSolver &s, ArrayRef<unsigned> target)
453 : OpRewritePattern(context), solver(s), targetBitwidths(target) {}
454
455 LogicalResult matchAndRewrite(arith::CmpIOp op,
456 PatternRewriter &rewriter) const override {
457 Value lhs = op.getLhs();
458 Value rhs = op.getRhs();
459
460 SmallVector<ConstantIntRanges> ranges;
461 if (failed(collectRanges(solver, op.getOperands(), ranges)))
462 return failure();
463 const ConstantIntRanges &lhsRange = ranges[0];
464 const ConstantIntRanges &rhsRange = ranges[1];
465
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;
471 };
472 // If we're to narrow the input values via a cast, we should preserve the
473 // sign.
474 CastKind predicateBasedCastRestriction =
475 isSignedCmpPredicate(op.getPredicate()) ? CastKind::Signed
476 : CastKind::Both;
477
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);
484 // Note: this includes target width > src width, as well as the unsigned
485 // truncatability & signed predicate scenario.
486 if (castKind == CastKind::None)
487 continue;
488
489 Type targetType = getTargetType(srcType, targetBitwidth);
490 if (targetType == srcType)
491 continue;
492
493 Location loc = op->getLoc();
494 IRMapping mapping;
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);
499
500 Operation *newOp = rewriter.clone(*op, mapping);
501 copyIntegerRange(solver, op.getResult(), newOp->getResult(0));
502 rewriter.replaceOp(op, newOp->getResults());
503 return success();
504 }
505 return failure();
506 }
507
508private:
509 DataFlowSolver &solver;
510 SmallVector<unsigned, 4> targetBitwidths;
511};
512
513/// Fold index_cast(index_cast(%arg: i8, index), i8) -> %arg
514/// This pattern assumes all passed `targetBitwidths` are not wider than index
515/// type.
516template <typename CastOp>
517struct FoldIndexCastChain final : OpRewritePattern<CastOp> {
518 FoldIndexCastChain(MLIRContext *context, ArrayRef<unsigned> target)
519 : OpRewritePattern<CastOp>(context), targetBitwidths(target) {}
520
521 LogicalResult matchAndRewrite(CastOp op,
522 PatternRewriter &rewriter) const override {
523 auto srcOp = op.getIn().template getDefiningOp<CastOp>();
524 if (!srcOp)
525 return rewriter.notifyMatchFailure(op, "doesn't come from an index cast");
526
527 Value src = srcOp.getIn();
528 if (src.getType() != op.getType())
529 return rewriter.notifyMatchFailure(op, "outer types don't match");
530
531 if (!srcOp.getType().isIndex())
532 return rewriter.notifyMatchFailure(op, "intermediate type isn't index");
533
534 auto intType = dyn_cast<IntegerType>(op.getType());
535 if (!intType || !llvm::is_contained(targetBitwidths, intType.getWidth()))
536 return failure();
537
538 rewriter.replaceOp(op, src);
539 return success();
540 }
541
542private:
543 SmallVector<unsigned, 4> targetBitwidths;
544};
545
546struct NarrowLoopBounds final : OpInterfaceRewritePattern<LoopLikeOpInterface> {
547 NarrowLoopBounds(MLIRContext *context, DataFlowSolver &s,
548 ArrayRef<unsigned> target)
549 : OpInterfaceRewritePattern<LoopLikeOpInterface>(context), solver(s),
550 targetBitwidths(target),
551 boundsNarrowingFailedAttr(
552 StringAttr::get(context, "arith.bounds_narrowing_failed")) {}
553
554 LogicalResult matchAndRewrite(LoopLikeOpInterface loopLike,
555 PatternRewriter &rewriter) const override {
556 // Skip ops where bounds narrowing previously failed.
557 if (loopLike->hasDiscardableAttr(boundsNarrowingFailedAttr))
558 return rewriter.notifyMatchFailure(loopLike,
559 "bounds narrowing previously failed");
560
561 std::optional<SmallVector<Value>> inductionVars =
562 loopLike.getLoopInductionVars();
563 if (!inductionVars.has_value() || inductionVars->empty())
564 return rewriter.notifyMatchFailure(loopLike, "no induction variables");
565
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();
571
572 if (!lowerBounds.has_value() || !upperBounds.has_value() ||
573 !steps.has_value())
574 return rewriter.notifyMatchFailure(loopLike, "no loop bounds or steps");
575
576 if (lowerBounds->size() != inductionVars->size() ||
577 upperBounds->size() != inductionVars->size() ||
578 steps->size() != inductionVars->size())
579 return rewriter.notifyMatchFailure(loopLike,
580 "mismatched bounds/steps count");
581
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;
587
588 // Check each (indVar, lb, ub, step) tuple.
589 for (auto [idx, indVar, lbOFR, ubOFR, stepOFR] :
590 llvm::enumerate(*inductionVars, *lowerBounds, *upperBounds, *steps)) {
591
592 // Only process value operands, skip attributes.
593 auto maybeLb = dyn_cast<Value>(lbOFR);
594 auto maybeUb = dyn_cast<Value>(ubOFR);
595 auto maybeStep = dyn_cast<Value>(stepOFR);
596
597 if (!maybeLb || !maybeUb || !maybeStep)
598 continue;
599
600 // Collect ranges for (lb, ub, step, indVar).
601 SmallVector<ConstantIntRanges> ranges;
602 if (failed(collectRanges(
603 solver, ValueRange{maybeLb, maybeUb, maybeStep, indVar}, ranges)))
604 continue;
605
606 const ConstantIntRanges &stepRange = ranges[2];
607 const ConstantIntRanges &indVarRange = ranges[3];
608
609 Type srcType = maybeLb.getType();
610
611 // Try each target bitwidth.
612 for (unsigned targetBitwidth : targetBitwidths) {
613 Type targetType = getTargetType(srcType, targetBitwidth);
614 if (targetType == srcType)
615 continue;
616
617 // Check if the target type is valid for this loop's induction
618 // variables.
619 if (!loopLike.isValidInductionVarType(targetType))
620 continue;
621
622 // Check if all values in this tuple can be truncated.
623 CastKind castKind = CastKind::Both;
624 for (const ConstantIntRanges &range : ranges) {
625 castKind = mergeCastKinds(castKind,
626 checkTruncatability(range, targetBitwidth));
627 if (castKind == CastKind::None)
628 break;
629 }
630
631 if (castKind == CastKind::None)
632 continue;
633
634 // Check if indVar + step fits in the narrowed type.
635 // This is critical for loop correctness: the loop computes
636 // iv_next = iv_current + step in the narrowed type, then compares
637 // iv_next < ub. If iv_current + step overflows, the comparison may
638 // produce incorrect results and break loop termination.
639 // Both signed and unsigned interpretations must fit because loop
640 // semantics are unknown (integer types are signless).
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()));
646
647 if (checkTruncatability(indVarPlusStepRange, targetBitwidth) !=
648 CastKind::Both)
649 continue;
650
651 // Narrow the bounds and step values.
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);
655
656 newLowerBounds[idx] = newLb;
657 newUpperBounds[idx] = newUb;
658 newSteps[idx] = newStep;
659 narrowings.push_back({idx, targetType, castKind});
660 break;
661 }
662 }
663
664 if (narrowings.empty())
665 return rewriter.notifyMatchFailure(loopLike, "no narrowings found");
666
667 // Save original types before modifying.
668 SmallVector<Type> origTypes;
669 for (auto [idx, targetType, castKind] : narrowings) {
670 Value indVar = (*inductionVars)[idx];
671 origTypes.push_back(indVar.getType());
672 }
673
674 // Attempt to update bounds and induction variable types.
675 // If this fails, mark the op so we don't try again.
676 bool updateFailed = false;
677 rewriter.modifyOpInPlace(loopLike, [&]() {
678 // Update the loop bounds and steps.
679 if (failed(loopLike.setLoopLowerBounds(newLowerBounds)) ||
680 failed(loopLike.setLoopUpperBounds(newUpperBounds)) ||
681 failed(loopLike.setLoopSteps(newSteps))) {
682 // Mark op to prevent future attempts. IR was modified (attribute
683 // added), so we must return success() from the pattern.
684 loopLike->setDiscardableAttr(boundsNarrowingFailedAttr,
685 rewriter.getUnitAttr());
686 updateFailed = true;
687 return;
688 }
689
690 // Update induction variable types.
691 for (auto [idx, targetType, castKind] : narrowings) {
692 Value indVar = (*inductionVars)[idx];
693 auto blockArg = cast<BlockArgument>(indVar);
694
695 // Change the block argument type.
696 blockArg.setType(targetType);
697 }
698 });
699
700 if (updateFailed)
701 return success();
702
703 // Insert casts back to original type for uses.
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];
709
710 OpBuilder::InsertionGuard guard(rewriter);
711 rewriter.setInsertionPointToStart(blockArg.getOwner());
712 Value casted = doCast(rewriter, loc, blockArg, origType, castKind);
713 copyIntegerRange(solver, blockArg, casted);
714
715 // Replace all uses of the narrowed indVar with the casted value.
716 rewriter.replaceAllUsesExcept(blockArg, casted, casted.getDefiningOp());
717 }
718
719 return success();
720 }
721
722private:
723 DataFlowSolver &solver;
724 SmallVector<unsigned, 4> targetBitwidths;
725 StringAttr boundsNarrowingFailedAttr;
726};
727
728struct IntRangeOptimizationsPass final
729 : arith::impl::ArithIntRangeOptsBase<IntRangeOptimizationsPass> {
730
731 void runOnOperation() override {
732 Operation *op = getOperation();
733 MLIRContext *ctx = op->getContext();
734 DataFlowSolver solver;
735 loadBaselineAnalyses(solver);
736 solver.load<IntegerRangeAnalysis>();
737 if (failed(solver.initializeAndRun(op)))
738 return signalPassFailure();
739
740 DataFlowListener listener(solver);
741
742 RewritePatternSet patterns(ctx);
744
745 // Disable folding and region simplification to avoid breaking the solver
746 // state. Both can remove block arguments (folding via control-flow
747 // simplification, region simplification via dead-arg elimination), which
748 // frees their underlying storage. A subsequent allocation may reuse the
749 // same address for a different block argument, causing stale solver state
750 // to be associated with the new argument and producing incorrect constants.
751 if (failed(
752 applyPatternsGreedily(op, std::move(patterns),
753 GreedyRewriteConfig()
754 .enableFolding(false)
755 .setRegionSimplificationLevel(
756 GreedySimplifyRegionLevel::Disabled)
757 .setListener(&listener))))
758 signalPassFailure();
759 }
760};
761
762struct IntRangeNarrowingPass final
763 : arith::impl::ArithIntRangeNarrowingBase<IntRangeNarrowingPass> {
764 using ArithIntRangeNarrowingBase::ArithIntRangeNarrowingBase;
765
766 void runOnOperation() override {
767 Operation *op = getOperation();
768 MLIRContext *ctx = op->getContext();
769 DataFlowSolver solver;
770 loadBaselineAnalyses(solver);
771 solver.load<IntegerRangeAnalysis>();
772 if (failed(solver.initializeAndRun(op)))
773 return signalPassFailure();
774
775 DataFlowListener listener(solver);
776
777 RewritePatternSet patterns(ctx);
778 populateIntRangeNarrowingPatterns(patterns, solver, bitwidthsSupported);
780 bitwidthsSupported);
781
782 // We specifically need bottom-up traversal as cmpi pattern needs range
783 // data, attached to its original argument values.
785 op, std::move(patterns),
786 GreedyRewriteConfig().setUseTopDownTraversal(false).setListener(
787 &listener))))
788 signalPassFailure();
789 }
790};
791} // namespace
792
794 RewritePatternSet &patterns, DataFlowSolver &solver) {
795 patterns.add<MaterializeKnownConstantValues, DeleteTrivialRem<RemSIOp>,
796 DeleteTrivialRem<RemUIOp>>(patterns.getContext(), solver);
797}
798
800 RewritePatternSet &patterns, DataFlowSolver &solver,
801 ArrayRef<unsigned> bitwidthsSupported) {
802 patterns.add<NarrowElementwise, NarrowCmpI>(patterns.getContext(), solver,
803 bitwidthsSupported);
804 patterns.add<FoldIndexCastChain<arith::IndexCastUIOp>,
805 FoldIndexCastChain<arith::IndexCastOp>>(patterns.getContext(),
806 bitwidthsSupported);
807}
808
810 RewritePatternSet &patterns, DataFlowSolver &solver,
811 ArrayRef<unsigned> bitwidthsSupported) {
812 patterns.add<NarrowLoopBounds>(patterns.getContext(), solver,
813 bitwidthsSupported);
814}
815
817 return std::make_unique<IntRangeOptimizationsPass>();
818}
return success()
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.
Definition FoldUtils.cpp:51
lhs
static void copyIntegerRange(DataFlowSolver &solver, Value oldVal, Value newVal)
static std::optional< APInt > getMaybeConstantValue(DataFlowSolver &solver, Value value)
@ None
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.
Definition Attributes.h:25
UnitAttr getUnitAttr()
Definition Builders.cpp:106
IntegerAttr getIntegerAttr(Type type, int64_t value)
Definition Builders.cpp:237
MLIRContext * getContext() const
Definition Builders.h:56
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...
Definition Dialect.h:38
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...
Definition Dialect.h:83
void map(Value from, Value to)
Inserts a new mapping for 'from' to 'to'.
Definition IRMapping.h:30
This class defines the main interface for locations in MLIR and acts as a non-nullable wrapper around...
Definition Location.h:76
Dialect * getLoadedDialect(StringRef name)
Get a registered IR dialect with the given namespace.
This class helps build Operations.
Definition Builders.h:210
Operation * clone(Operation &op, IRMapping &mapper)
Creates a deep copy of the specified operation, remapping any operands that use values outside of the...
Definition Builders.cpp:581
void setInsertionPointToStart(Block *block)
Sets the insertion point to the start of the specified block.
Definition Builders.h:434
This is a value defined by a result of an operation.
Definition Value.h:454
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.
Definition Operation.h:87
Dialect * getDialect()
Return the dialect this operation is associated with, or nullptr if the associated dialect is not loa...
Definition Operation.h:237
OpResult getResult(unsigned idx)
Get the 'idx'th result of this operation.
Definition Operation.h:432
unsigned getNumRegions()
Returns the number of regions held by this operation.
Definition Operation.h:726
Location getLoc()
The source location the operation was defined or derived from.
Definition Operation.h:240
unsigned getNumOperands()
Definition Operation.h:371
operand_type_range getOperandTypes()
Definition Operation.h:422
MutableArrayRef< Region > getRegions()
Returns the regions held by this operation.
Definition Operation.h:729
result_type_range getResultTypes()
Definition Operation.h:453
operand_range getOperands()
Returns an iterator on the underlying Value's.
Definition Operation.h:403
result_range getResults()
Definition Operation.h:440
MLIRContext * getContext()
Return the context this operation is associated with.
Definition Operation.h:233
unsigned getNumResults()
Return the number of results held by this operation.
Definition Operation.h:429
Operation * getParentOp()
Return the parent operation this region is attached to.
Definition Region.h:198
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...
Definition Types.h:74
MLIRContext * getContext() const
Return the MLIRContext in which this type was uniqued.
Definition Types.cpp:35
bool isIntOrIndex() const
Return true if this is an integer (of any signedness) or an index type.
Definition Types.cpp:114
This class provides an abstraction over the different types of ranges over Values.
Definition ValueRange.h:389
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
Definition Value.h:96
bool use_empty() const
Returns true if this value has no uses.
Definition Value.h:208
void setType(Type newType)
Mutate the type of this Value to be of the specified type.
Definition Value.h:116
Type getType() const
Return the type of this value.
Definition Value.h:105
Location getLoc() const
Return the location of this value.
Definition Value.cpp:24
Operation * getDefiningOp() const
If this value is the result of an operation, return the operation that defines it.
Definition Value.cpp:18
Region * getParentRegion()
Return the Region in which this Value is defined.
Definition Value.cpp:39
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...
Definition Utils.h:29
detail::InFlightRemark failed(Location loc, RemarkOpts opts)
Report an optimization remark that failed.
Definition Remarks.h:734
Include the generated interface declarations.
bool matchPattern(Value value, const Pattern &pattern)
Entry point for matching a pattern over a Value.
Definition Matchers.h:490
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...
Definition Matchers.h:527
LogicalResult applyPatternsGreedily(Region &region, 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.
Definition Matchers.h:369
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...