MLIR 24.0.0git
ArithOps.cpp
Go to the documentation of this file.
1//===- ArithOps.cpp - MLIR Arith dialect ops implementation -----===//
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 <cassert>
10#include <cstdint>
11#include <functional>
12#include <type_traits>
13#include <utility>
14
18#include "mlir/IR/Builders.h"
21#include "mlir/IR/Matchers.h"
26
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"
34
35using namespace mlir;
36using namespace mlir::arith;
37
38/// Default rounding mode according to default LLVM floating-point environment.
39static constexpr llvm::RoundingMode kDefaultRoundingMode =
40 llvm::RoundingMode::NearestTiesToEven;
41
42//===----------------------------------------------------------------------===//
43// Pattern helpers
44//===----------------------------------------------------------------------===//
45
46static IntegerAttr
48 Attribute rhs,
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);
54}
55
56static IntegerAttr addIntegerAttrs(PatternRewriter &builder, Value res,
57 Attribute lhs, Attribute rhs) {
58 return applyToIntegerAttrs(builder, res, lhs, rhs, std::plus<APInt>());
59}
60
61static IntegerAttr subIntegerAttrs(PatternRewriter &builder, Value res,
62 Attribute lhs, Attribute rhs) {
63 return applyToIntegerAttrs(builder, res, lhs, rhs, std::minus<APInt>());
64}
65
66static IntegerAttr mulIntegerAttrs(PatternRewriter &builder, Value res,
67 Attribute lhs, Attribute rhs) {
68 return applyToIntegerAttrs(builder, res, lhs, rhs, std::multiplies<APInt>());
69}
70
71static IntegerAttr andIntegerAttrs(PatternRewriter &builder, Value res,
72 Attribute lhs, Attribute rhs) {
73 return applyToIntegerAttrs(builder, res, lhs, rhs, std::bit_and<APInt>());
74}
75
76static IntegerAttr orIntegerAttrs(PatternRewriter &builder, Value res,
77 Attribute lhs, Attribute rhs) {
78 return applyToIntegerAttrs(builder, res, lhs, rhs, std::bit_or<APInt>());
79}
80
81static IntegerAttr xorIntegerAttrs(PatternRewriter &builder, Value res,
82 Attribute lhs, Attribute rhs) {
83 return applyToIntegerAttrs(builder, res, lhs, rhs, std::bit_xor<APInt>());
84}
85
86// Merge overflow flags from 2 ops, selecting the most conservative combination.
87static IntegerOverflowFlagsAttr
88mergeOverflowFlags(IntegerOverflowFlagsAttr val1,
89 IntegerOverflowFlagsAttr val2) {
90 return IntegerOverflowFlagsAttr::get(val1.getContext(),
91 val1.getValue() & val2.getValue());
92}
93
94/// Invert an integer comparison predicate.
95arith::CmpIPredicate arith::invertPredicate(arith::CmpIPredicate pred) {
96 switch (pred) {
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;
117 }
118 llvm_unreachable("unknown cmpi predicate kind");
119}
120
121/// Equivalent to
122/// convertRoundingModeToLLVM(convertArithRoundingModeToLLVM(roundingMode)).
123///
124/// Not possible to implement as chain of calls as this would introduce a
125/// circular dependency with MLIRArithAttrToLLVMConversion and make arith depend
126/// on the LLVM dialect and on translation to LLVM.
127static llvm::RoundingMode
128convertArithRoundingModeToLLVMIR(std::optional<RoundingMode> roundingMode) {
129 if (!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;
142 }
143 llvm_unreachable("Unhandled rounding mode");
144}
145
146static arith::CmpIPredicateAttr invertPredicate(arith::CmpIPredicateAttr pred) {
147 return arith::CmpIPredicateAttr::get(pred.getContext(),
148 invertPredicate(pred.getValue()));
149}
150
152 Type elemTy = getElementTypeOrSelf(type);
153 if (elemTy.isIntOrFloat())
154 return elemTy.getIntOrFloatBitWidth();
155
156 return -1;
157}
158
160 return getScalarOrElementWidth(value.getType());
161}
162
163static FailureOr<APInt> getIntOrSplatIntValue(Attribute attr) {
164 APInt value;
165 if (matchPattern(attr, m_ConstantInt(&value)))
166 return value;
167
168 return failure();
169}
170
171static Attribute getBoolAttribute(Type type, bool value) {
172 auto boolAttr = BoolAttr::get(type.getContext(), value);
173 ShapedType shapedType = dyn_cast_or_null<ShapedType>(type);
174 if (!shapedType)
175 return boolAttr;
176 // DenseElementsAttr requires a static shape.
177 if (!shapedType.hasStaticShape())
178 return {};
179 return DenseElementsAttr::get(shapedType, boolAttr);
180}
181
182/// Return a scalar or splat integer attribute of `type` (an integer/index type
183/// or a shaped type thereof) holding `value`. Returns a null attribute for
184/// shaped types with a dynamic shape, so callers can bail out of folding.
186 auto scalarAttr = IntegerAttr::get(getElementTypeOrSelf(type), value);
187 ShapedType shapedType = dyn_cast<ShapedType>(type);
188 if (!shapedType)
189 return scalarAttr;
190 if (!shapedType.hasStaticShape())
191 return {};
192 return DenseElementsAttr::get(shapedType, scalarAttr);
193}
194
195//===----------------------------------------------------------------------===//
196// TableGen'd canonicalization patterns
197//===----------------------------------------------------------------------===//
198
199namespace {
200#include "ArithCanonicalization.inc"
201} // namespace
202
203//===----------------------------------------------------------------------===//
204// Common helpers
205//===----------------------------------------------------------------------===//
206
207/// Return the type of the same shape (scalar, vector or tensor) containing i1.
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);
214 return i1Type;
215}
216
217//===----------------------------------------------------------------------===//
218// ConstantOp
219//===----------------------------------------------------------------------===//
220
221void arith::ConstantOp::getAsmResultNames(
222 function_ref<void(Value, StringRef)> setNameFn) {
223 auto type = getType();
224 if (auto intCst = dyn_cast<IntegerAttr>(getValue())) {
225 auto intType = dyn_cast<IntegerType>(type);
226
227 // Sugar i1 constants with 'true' and 'false'.
228 if (intType && intType.getWidth() == 1)
229 return setNameFn(getResult(), (intCst.getInt() ? "true" : "false"));
230
231 // Otherwise, build a complex name with the value and type.
232 SmallString<32> specialNameBuffer;
233 llvm::raw_svector_ostream specialName(specialNameBuffer);
234 specialName << 'c' << intCst.getValue();
235 if (intType)
236 specialName << '_' << type;
237 setNameFn(getResult(), specialName.str());
238 } else {
239 setNameFn(getResult(), "cst");
240 }
241}
242
243/// TODO: disallow arith.constant to return anything other than signless integer
244/// or float like.
245LogicalResult arith::ConstantOp::verify() {
246 auto type = getType();
247 // Integer values must be signless.
248 if (auto intType = dyn_cast<IntegerType>(getElementTypeOrSelf(type));
249 intType && !intType.isSignless())
250 return emitOpError("integer return type must be signless");
251 // Any float or elements attribute are acceptable.
252 if (!llvm::isa<IntegerAttr, FloatAttr, ElementsAttr>(getValue())) {
253 return emitOpError(
254 "value must be an integer, float, or elements attribute");
255 }
256
257 // Note, we could relax this for vectors with 1 scalable dim, e.g.:
258 // * arith.constant dense<[[3, 3], [1, 1]]> : vector<2 x [2] x i32>
259 // However, this would most likely require updating the lowerings to LLVM.
260 if (isa<ScalableVectorType>(type) && !isa<SplatElementsAttr>(getValue()))
261 return emitOpError(
262 "initializing scalable vectors with elements attribute is not supported"
263 " unless it's a vector splat");
264 return success();
265}
266
267bool arith::ConstantOp::isBuildableWith(Attribute value, Type type) {
268 // The value's type must be the same as the provided type.
269 auto typedAttr = dyn_cast<TypedAttr>(value);
270 if (!typedAttr || typedAttr.getType() != type)
271 return false;
272 // Integer values must be signless.
273 if (auto intType = dyn_cast<IntegerType>(getElementTypeOrSelf(type))) {
274 if (!intType.isSignless())
275 return false;
276 }
277 // Integer, float, and element attributes are buildable.
278 return llvm::isa<IntegerAttr, FloatAttr, ElementsAttr>(value);
279}
280
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));
285 return nullptr;
286}
287
288OpFoldResult arith::ConstantOp::fold(FoldAdaptor adaptor) { return getValue(); }
289
291 int64_t value, unsigned width) {
292 auto type = builder.getIntegerType(width);
293 arith::ConstantOp::build(builder, result, type,
294 builder.getIntegerAttr(type, value));
295}
296
298 Location location,
300 unsigned width) {
301 mlir::OperationState state(location, getOperationName());
302 build(builder, state, value, width);
303 auto result = dyn_cast<ConstantIntOp>(builder.create(state));
304 assert(result && "builder didn't return the right type");
305 return result;
306}
307
310 unsigned width) {
311 return create(builder, builder.getLoc(), value, width);
312}
313
315 Type type, int64_t value) {
316 arith::ConstantOp::build(builder, result, type,
317 builder.getIntegerAttr(type, value));
318}
319
321 Location location, Type type,
322 int64_t value) {
323 mlir::OperationState state(location, getOperationName());
324 build(builder, state, type, value);
325 auto result = dyn_cast<ConstantIntOp>(builder.create(state));
326 assert(result && "builder didn't return the right type");
327 return result;
328}
329
331 Type type, int64_t value) {
332 return create(builder, builder.getLoc(), type, value);
333}
334
336 Type type, const APInt &value) {
337 arith::ConstantOp::build(builder, result, type,
338 builder.getIntegerAttr(type, value));
339}
340
342 Location location, Type type,
343 const APInt &value) {
344 mlir::OperationState state(location, getOperationName());
345 build(builder, state, type, value);
346 auto result = dyn_cast<ConstantIntOp>(builder.create(state));
347 assert(result && "builder didn't return the right type");
348 return result;
349}
350
352 Type type,
353 const APInt &value) {
354 return create(builder, builder.getLoc(), type, value);
355}
356
358 if (auto constOp = dyn_cast_or_null<arith::ConstantOp>(op))
359 return constOp.getType().isSignlessInteger();
360 return false;
361}
362
364 FloatType type, const APFloat &value) {
365 arith::ConstantOp::build(builder, result, type,
366 builder.getFloatAttr(type, value));
367}
368
370 Location location,
371 FloatType type,
372 const APFloat &value) {
373 mlir::OperationState state(location, getOperationName());
374 build(builder, state, type, value);
375 auto result = dyn_cast<ConstantFloatOp>(builder.create(state));
376 assert(result && "builder didn't return the right type");
377 return result;
378}
379
382 const APFloat &value) {
383 return create(builder, builder.getLoc(), type, value);
384}
385
387 if (auto constOp = dyn_cast_or_null<arith::ConstantOp>(op))
388 return llvm::isa<FloatType>(constOp.getType());
389 return false;
390}
391
393 int64_t value) {
394 arith::ConstantOp::build(builder, result, builder.getIndexType(),
395 builder.getIndexAttr(value));
396}
397
399 Location location,
400 int64_t value) {
401 mlir::OperationState state(location, getOperationName());
402 build(builder, state, value);
403 auto result = dyn_cast<ConstantIndexOp>(builder.create(state));
404 assert(result && "builder didn't return the right type");
405 return result;
406}
407
410 return create(builder, builder.getLoc(), value);
411}
412
414 if (auto constOp = dyn_cast_or_null<arith::ConstantOp>(op))
415 return constOp.getType().isIndex();
416 return false;
417}
418
420 Type type) {
421 // TODO: Incorporate this check to `FloatAttr::get*`.
422 assert(!isa<Float8E8M0FNUType>(getElementTypeOrSelf(type)) &&
423 "type doesn't have a zero representation");
424 TypedAttr zeroAttr = builder.getZeroAttr(type);
425 assert(zeroAttr && "unsupported type for zero attribute");
426 return arith::ConstantOp::create(builder, loc, zeroAttr);
427}
428
429//===----------------------------------------------------------------------===//
430// AddIOp
431//===----------------------------------------------------------------------===//
432
433OpFoldResult arith::AddIOp::fold(FoldAdaptor adaptor) {
434 // addi(x, 0) -> x
435 if (matchPattern(adaptor.getRhs(), m_Zero()))
436 return getLhs();
437
438 // addi(subi(a, b), b) -> a
439 if (auto sub = getLhs().getDefiningOp<SubIOp>())
440 if (getRhs() == sub.getRhs())
441 return sub.getLhs();
442
443 // addi(b, subi(a, b)) -> a
444 if (auto sub = getRhs().getDefiningOp<SubIOp>())
445 if (getLhs() == sub.getRhs())
446 return sub.getLhs();
447
449 adaptor.getOperands(),
450 [](APInt a, const APInt &b) { return std::move(a) + b; });
451}
452
453void arith::AddIOp::getCanonicalizationPatterns(RewritePatternSet &patterns,
454 MLIRContext *context) {
455 patterns.add<AddIAddConstant, AddISubConstantRHS, AddISubConstantLHS,
456 AddIMulNegativeOneRhs, AddIMulNegativeOneLhs>(context);
457}
458
459//===----------------------------------------------------------------------===//
460// AddUIExtendedOp
461//===----------------------------------------------------------------------===//
462
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());
467 return std::nullopt;
468}
469
470// Returns the overflow bit, assuming that `sum` is the result of unsigned
471// addition of `operand` and another number.
472static APInt calculateUnsignedOverflow(const APInt &sum, const APInt &operand) {
473 return sum.ult(operand) ? APInt::getAllOnes(1) : APInt::getZero(1);
474}
475
476LogicalResult
477arith::AddUIExtendedOp::fold(FoldAdaptor adaptor,
478 SmallVectorImpl<OpFoldResult> &results) {
479 Type overflowTy = getOverflow().getType();
480 // addui_extended(x, 0) -> x, false
481 if (matchPattern(getRhs(), m_Zero())) {
482 Builder builder(getContext());
483 auto falseValue = builder.getZeroAttr(overflowTy);
484
485 results.push_back(getLhs());
486 results.push_back(falseValue);
487 return success();
488 }
489
490 // addui_extended(constant_a, constant_b) -> constant_sum, constant_carry
491 // Let the `constFoldBinaryOp` utility attempt to fold the sum of both
492 // operands. If that succeeds, calculate the overflow bit based on the sum
493 // and the first (constant) operand, `lhs`.
494 if (Attribute sumAttr = constFoldBinaryOp<IntegerAttr>(
495 adaptor.getOperands(),
496 [](APInt a, const APInt &b) { return std::move(a) + b; })) {
497 // If any operand is poison, propagate poison to both results.
498 if (matchPattern(sumAttr, ub::m_Poison())) {
499 results.push_back(sumAttr);
500 results.push_back(sumAttr);
501 return success();
502 }
503 Attribute overflowAttr = constFoldBinaryOp<IntegerAttr>(
504 ArrayRef({sumAttr, adaptor.getLhs()}),
505 getI1SameShape(llvm::cast<TypedAttr>(sumAttr).getType()),
507 if (!overflowAttr)
508 return failure();
509
510 results.push_back(sumAttr);
511 results.push_back(overflowAttr);
512 return success();
513 }
514
515 return failure();
516}
517
518void arith::AddUIExtendedOp::getCanonicalizationPatterns(
519 RewritePatternSet &patterns, MLIRContext *context) {
520 patterns.add<AddUIExtendedToAddI>(context);
521}
522
523//===----------------------------------------------------------------------===//
524// SubUIExtendedOp
525//===----------------------------------------------------------------------===//
526
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());
531 return std::nullopt;
532}
533
534// Returns the borrow bit, assuming `lhs` and `rhs` are operands of an unsigned
535// subtraction whose mathematical result underflows iff `lhs < rhs`.
536static APInt calculateUnsignedBorrow(const APInt &lhs, const APInt &rhs) {
537 return lhs.ult(rhs) ? APInt::getAllOnes(1) : APInt::getZero(1);
538}
539
540LogicalResult
541arith::SubUIExtendedOp::fold(FoldAdaptor adaptor,
542 SmallVectorImpl<OpFoldResult> &results) {
543 Type borrowTy = getBorrow().getType();
544 // subui_extended(x, 0) -> x, false
545 if (matchPattern(getRhs(), m_Zero())) {
546 Builder builder(getContext());
547 auto falseValue = builder.getZeroAttr(borrowTy);
548
549 results.push_back(getLhs());
550 results.push_back(falseValue);
551 return success();
552 }
553
554 // subui_extended(x, x) -> 0, false
555 if (getLhs() == getRhs()) {
556 // A dynamically-shaped result cannot be a constant; bail before
557 // getZeroAttr, which would assert on a non-static shape.
558 auto shapedType = dyn_cast<ShapedType>(getDiff().getType());
559 if (shapedType && !shapedType.hasStaticShape())
560 return failure();
561 Builder builder(getContext());
562 auto zeroDiff = builder.getZeroAttr(getDiff().getType());
563 auto falseValue = builder.getZeroAttr(borrowTy);
564 if (!zeroDiff)
565 return failure();
566
567 results.push_back(zeroDiff);
568 results.push_back(falseValue);
569 return success();
570 }
571
572 // subui_extended(constant_a, constant_b) -> constant_diff, constant_borrow
573 if (Attribute diffAttr = constFoldBinaryOp<IntegerAttr>(
574 adaptor.getOperands(),
575 [](APInt a, const APInt &b) { return std::move(a) - b; })) {
576 // If any operand is poison, propagate poison to both results.
577 if (matchPattern(diffAttr, ub::m_Poison())) {
578 results.push_back(diffAttr);
579 results.push_back(diffAttr);
580 return success();
581 }
582 Attribute borrowAttr = constFoldBinaryOp<IntegerAttr>(
583 adaptor.getOperands(),
584 getI1SameShape(llvm::cast<TypedAttr>(diffAttr).getType()),
586 if (!borrowAttr)
587 return failure();
588
589 results.push_back(diffAttr);
590 results.push_back(borrowAttr);
591 return success();
592 }
593
594 return failure();
595}
596
597void arith::SubUIExtendedOp::getCanonicalizationPatterns(
598 RewritePatternSet &patterns, MLIRContext *context) {
599 patterns.add<SubUIExtendedToSubI>(context);
600}
601
602//===----------------------------------------------------------------------===//
603// SubIOp
604//===----------------------------------------------------------------------===//
605
606OpFoldResult arith::SubIOp::fold(FoldAdaptor adaptor) {
607 // subi(x,x) -> 0
608 if (getOperand(0) == getOperand(1)) {
609 auto shapedType = dyn_cast<ShapedType>(getType());
610 // We can't generate a constant with a dynamic shaped tensor.
611 if (!shapedType || shapedType.hasStaticShape())
612 return Builder(getContext()).getZeroAttr(getType());
613 }
614 // subi(x,0) -> x
615 if (matchPattern(adaptor.getRhs(), m_Zero()))
616 return getLhs();
617
618 if (auto add = getLhs().getDefiningOp<AddIOp>()) {
619 // subi(addi(a, b), b) -> a
620 if (getRhs() == add.getRhs())
621 return add.getLhs();
622 // subi(addi(a, b), a) -> b
623 if (getRhs() == add.getLhs())
624 return add.getRhs();
625 }
626
627 // subi(a, subi(a, b)) -> b
628 if (auto sub = getRhs().getDefiningOp<SubIOp>())
629 if (getLhs() == sub.getLhs())
630 return sub.getRhs();
631
633 adaptor.getOperands(),
634 [](APInt a, const APInt &b) { return std::move(a) - b; });
635}
636
637void arith::SubIOp::getCanonicalizationPatterns(RewritePatternSet &patterns,
638 MLIRContext *context) {
639 patterns.add<SubIRHSAddConstant, SubILHSAddConstant, SubIRHSSubConstantRHS,
640 SubIRHSSubConstantLHS, SubILHSSubConstantRHS,
641 SubILHSSubConstantLHS, SubISubILHSRHSLHS>(context);
642}
643
644//===----------------------------------------------------------------------===//
645// MulIOp
646//===----------------------------------------------------------------------===//
647
648OpFoldResult arith::MulIOp::fold(FoldAdaptor adaptor) {
649 // muli(x, 0) -> 0
650 if (matchPattern(adaptor.getRhs(), m_Zero()))
651 return getRhs();
652 // muli(x, 1) -> x
653 if (matchPattern(adaptor.getRhs(), m_One()))
654 return getLhs();
655 // TODO: Handle the overflow case.
656
657 // default folder
659 adaptor.getOperands(),
660 [](const APInt &a, const APInt &b) { return a * b; });
661}
662
663void arith::MulIOp::getAsmResultNames(
664 function_ref<void(Value, StringRef)> setNameFn) {
665 if (!isa<IndexType>(getType()))
666 return;
667
668 // Match vector.vscale by name to avoid depending on the vector dialect (which
669 // is a circular dependency).
670 auto isVscale = [](Operation *op) {
671 return op && op->getName().getStringRef() == "vector.vscale";
672 };
673
674 IntegerAttr baseValue;
675 auto isVscaleExpr = [&](Value a, Value b) {
676 return matchPattern(a, m_Constant(&baseValue)) &&
677 isVscale(b.getDefiningOp());
678 };
679
680 if (!isVscaleExpr(getLhs(), getRhs()) && !isVscaleExpr(getRhs(), getLhs()))
681 return;
682
683 // Name `base * vscale` or `vscale * base` as `c<base_value>_vscale`.
684 SmallString<32> specialNameBuffer;
685 llvm::raw_svector_ostream specialName(specialNameBuffer);
686 specialName << 'c' << baseValue.getInt() << "_vscale";
687 setNameFn(getResult(), specialName.str());
688}
689
690void arith::MulIOp::getCanonicalizationPatterns(RewritePatternSet &patterns,
691 MLIRContext *context) {
692 patterns.add<MulIMulIConstant>(context);
693}
694
695//===----------------------------------------------------------------------===//
696// MulSIExtendedOp
697//===----------------------------------------------------------------------===//
698
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());
703 return std::nullopt;
704}
705
706LogicalResult
707arith::MulSIExtendedOp::fold(FoldAdaptor adaptor,
708 SmallVectorImpl<OpFoldResult> &results) {
709 // mulsi_extended(x, 0) -> 0, 0
710 if (matchPattern(adaptor.getRhs(), m_Zero())) {
711 Attribute zero = adaptor.getRhs();
712 results.push_back(zero);
713 results.push_back(zero);
714 return success();
715 }
716
717 // mulsi_extended(cst_a, cst_b) -> cst_low, cst_high
718 if (Attribute lowAttr = constFoldBinaryOp<IntegerAttr>(
719 adaptor.getOperands(),
720 [](const APInt &a, const APInt &b) { return a * b; })) {
721 // Invoke the constant fold helper again to calculate the 'high' result.
722 Attribute highAttr = constFoldBinaryOp<IntegerAttr>(adaptor.getOperands(),
723 llvm::APIntOps::mulhs);
724 assert(highAttr && "Unexpected constant-folding failure");
725
726 results.push_back(lowAttr);
727 results.push_back(highAttr);
728 return success();
729 }
730
731 return failure();
732}
733
734void arith::MulSIExtendedOp::getCanonicalizationPatterns(
735 RewritePatternSet &patterns, MLIRContext *context) {
736 patterns.add<MulSIExtendedToMulI, MulSIExtendedRHSOne>(context);
737}
738
739//===----------------------------------------------------------------------===//
740// MulUIExtendedOp
741//===----------------------------------------------------------------------===//
742
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());
747 return std::nullopt;
748}
749
750LogicalResult
751arith::MulUIExtendedOp::fold(FoldAdaptor adaptor,
752 SmallVectorImpl<OpFoldResult> &results) {
753 // mului_extended(x, 0) -> 0, 0
754 if (matchPattern(adaptor.getRhs(), m_Zero())) {
755 Attribute zero = adaptor.getRhs();
756 results.push_back(zero);
757 results.push_back(zero);
758 return success();
759 }
760
761 // mului_extended(x, 1) -> x, 0
762 if (matchPattern(adaptor.getRhs(), m_One())) {
763 Builder builder(getContext());
764 Attribute zero = builder.getZeroAttr(getLhs().getType());
765 results.push_back(getLhs());
766 results.push_back(zero);
767 return success();
768 }
769
770 // mului_extended(cst_a, cst_b) -> cst_low, cst_high
771 if (Attribute lowAttr = constFoldBinaryOp<IntegerAttr>(
772 adaptor.getOperands(),
773 [](const APInt &a, const APInt &b) { return a * b; })) {
774 // Invoke the constant fold helper again to calculate the 'high' result.
775 Attribute highAttr = constFoldBinaryOp<IntegerAttr>(adaptor.getOperands(),
776 llvm::APIntOps::mulhu);
777 assert(highAttr && "Unexpected constant-folding failure");
778
779 results.push_back(lowAttr);
780 results.push_back(highAttr);
781 return success();
782 }
783
784 return failure();
785}
786
787void arith::MulUIExtendedOp::getCanonicalizationPatterns(
788 RewritePatternSet &patterns, MLIRContext *context) {
789 patterns.add<MulUIExtendedToMulI>(context);
790}
791
792//===----------------------------------------------------------------------===//
793// DivUIOp
794//===----------------------------------------------------------------------===//
795
796/// Fold `(a * b) / b -> a`
797static Value foldDivMul(Value lhs, Value rhs,
798 arith::IntegerOverflowFlags ovfFlags) {
799 auto mul = lhs.getDefiningOp<mlir::arith::MulIOp>();
800 if (!mul || !bitEnumContainsAll(mul.getOverflowFlags(), ovfFlags))
801 return {};
802
803 if (mul.getLhs() == rhs)
804 return mul.getRhs();
805
806 if (mul.getRhs() == rhs)
807 return mul.getLhs();
808
809 return {};
810}
811
812OpFoldResult arith::DivUIOp::fold(FoldAdaptor adaptor) {
813 // TODO: divui (x, 0) -> poison. Division by zero is undefined behaviour and
814 // could fold to poison, but that would make the arith dialect depend on the
815 // ub dialect to materialize ub.poison; left out for now.
816
817 // divui (x, 1) -> x.
818 if (matchPattern(adaptor.getRhs(), m_One()))
819 return getLhs();
820
821 // divui (0, x) -> 0. Division by zero is UB, so refining to 0 is valid.
822 if (matchPattern(adaptor.getLhs(), m_Zero()))
823 return getLhs();
824
825 // divui (x, x) -> 1.
826 if (getLhs() == getRhs())
827 return getIntegerAttrOfType(getType(), 1);
828
829 // (a * b) / b -> a
830 if (Value val = foldDivMul(getLhs(), getRhs(), IntegerOverflowFlags::nuw))
831 return val;
832
833 // Don't fold if it would require a division by zero.
834 bool div0 = false;
835 auto result = constFoldBinaryOp<IntegerAttr>(adaptor.getOperands(),
836 [&](APInt a, const APInt &b) {
837 if (div0 || !b) {
838 div0 = true;
839 return a;
840 }
841 return a.udiv(b);
842 });
843
844 return div0 ? Attribute() : result;
845}
846
847/// Returns whether an unsigned division by `divisor` is speculatable.
849 // X / 0 => UB
850 if (matchPattern(divisor, m_IntRangeWithoutZeroU()))
852
854}
855
856Speculation::Speculatability arith::DivUIOp::getSpeculatability() {
857 return getDivUISpeculatability(getRhs());
858}
859
860//===----------------------------------------------------------------------===//
861// DivSIOp
862//===----------------------------------------------------------------------===//
863
864OpFoldResult arith::DivSIOp::fold(FoldAdaptor adaptor) {
865 // TODO: divsi (x, 0) -> poison. Division by zero is undefined behaviour and
866 // could fold to poison, but that would make the arith dialect depend on the
867 // ub dialect to materialize ub.poison; left out for now.
868
869 // divsi (x, 1) -> x.
870 if (matchPattern(adaptor.getRhs(), m_One()))
871 return getLhs();
872
873 // divsi (0, x) -> 0. Division by zero is UB, so refining to 0 is valid.
874 if (matchPattern(adaptor.getLhs(), m_Zero()))
875 return getLhs();
876
877 // divsi (x, x) -> 1.
878 if (getLhs() == getRhs())
879 return getIntegerAttrOfType(getType(), 1);
880
881 // (a * b) / b -> a
882 if (Value val = foldDivMul(getLhs(), getRhs(), IntegerOverflowFlags::nsw))
883 return val;
884
885 // Don't fold if it would overflow or if it requires a division by zero.
886 bool overflowOrDiv0 = false;
888 adaptor.getOperands(), [&](APInt a, const APInt &b) {
889 if (overflowOrDiv0 || !b) {
890 overflowOrDiv0 = true;
891 return a;
892 }
893 return a.sdiv_ov(b, overflowOrDiv0);
894 });
895
896 return overflowOrDiv0 ? Attribute() : result;
897}
898
899/// Returns whether a signed division by `divisor` is speculatable. This
900/// function conservatively assumes that all signed division by -1 are not
901/// speculatable.
903 // X / 0 => UB
904 // INT_MIN / -1 => UB
905 if (matchPattern(divisor, m_IntRangeWithoutZeroS()) &&
908
910}
911
912Speculation::Speculatability arith::DivSIOp::getSpeculatability() {
913 return getDivSISpeculatability(getRhs());
914}
915
916//===----------------------------------------------------------------------===//
917// CeilDivUIOp
918//===----------------------------------------------------------------------===//
919
920OpFoldResult arith::CeilDivUIOp::fold(FoldAdaptor adaptor) {
921 // TODO: ceildivui (x, 0) -> poison. Division by zero is undefined behaviour
922 // and could fold to poison, but that would make the arith dialect depend on
923 // the ub dialect to materialize ub.poison; left out for now.
924
925 // ceildivui (x, 1) -> x.
926 if (matchPattern(adaptor.getRhs(), m_One()))
927 return getLhs();
928
929 // ceildivui (0, x) -> 0. Division by zero is UB, so refining to 0 is valid.
930 if (matchPattern(adaptor.getLhs(), m_Zero()))
931 return getLhs();
932
933 // ceildivui (x, x) -> 1.
934 if (getLhs() == getRhs())
935 return getIntegerAttrOfType(getType(), 1);
936
937 bool overflowOrDiv0 = false;
939 adaptor.getOperands(), [&](APInt a, const APInt &b) {
940 if (overflowOrDiv0 || !b) {
941 overflowOrDiv0 = true;
942 return a;
943 }
944 APInt quotient = a.udiv(b);
945 if (!a.urem(b))
946 return quotient;
947 APInt one(a.getBitWidth(), 1, true);
948 return quotient.uadd_ov(one, overflowOrDiv0);
949 });
950
951 return overflowOrDiv0 ? Attribute() : result;
952}
953
954Speculation::Speculatability arith::CeilDivUIOp::getSpeculatability() {
955 return getDivUISpeculatability(getRhs());
956}
957
958//===----------------------------------------------------------------------===//
959// CeilDivSIOp
960//===----------------------------------------------------------------------===//
961
962OpFoldResult arith::CeilDivSIOp::fold(FoldAdaptor adaptor) {
963 // TODO: ceildivsi (x, 0) -> poison. Division by zero is undefined behaviour
964 // and could fold to poison, but that would make the arith dialect depend on
965 // the ub dialect to materialize ub.poison; left out for now.
966
967 // ceildivsi (x, 1) -> x.
968 if (matchPattern(adaptor.getRhs(), m_One()))
969 return getLhs();
970
971 // ceildivsi (0, x) -> 0. Division by zero is UB, so refining to 0 is valid.
972 if (matchPattern(adaptor.getLhs(), m_Zero()))
973 return getLhs();
974
975 // ceildivsi (x, x) -> 1.
976 if (getLhs() == getRhs())
977 return getIntegerAttrOfType(getType(), 1);
978
979 // Don't fold if it would overflow or if it requires a division by zero.
980 bool overflowOrDiv0 = false;
982 adaptor.getOperands(), [&](const APInt &a, const APInt &b) {
983 if (overflowOrDiv0 || !b) {
984 overflowOrDiv0 = true;
985 return a;
986 }
987 // Compute the ceiling without negating either operand, so that MININT
988 // operands still fold whenever the result is representable.
989 //
990 // sdiv truncates towards zero, so it already rounds up whenever the
991 // exact quotient is negative. When the exact quotient is positive, i.e.
992 // when the operands have the same sign, an inexact division has to be
993 // corrected by one. This mirrors the expansion in ExpandOps.cpp.
994 bool overflowDiv = false;
995 APInt quotient = a.sdiv_ov(b, overflowDiv);
996 if (overflowDiv) {
997 // MININT / -1. The exact result is -MININT, which is not
998 // representable.
999 overflowOrDiv0 = true;
1000 return a;
1001 }
1002 if (a.isNegative() != b.isNegative() || quotient * b == a)
1003 return quotient;
1004
1005 // The correction cannot overflow: it only applies when the exact
1006 // quotient is positive and the division is inexact, which bounds the
1007 // quotient well below the maximum. Check anyway, at no cost.
1008 APInt one(a.getBitWidth(), 1, /*isSigned=*/true);
1009 return quotient.sadd_ov(one, overflowOrDiv0);
1010 });
1011
1012 return overflowOrDiv0 ? Attribute() : result;
1013}
1014
1015Speculation::Speculatability arith::CeilDivSIOp::getSpeculatability() {
1016 return getDivSISpeculatability(getRhs());
1017}
1018
1019//===----------------------------------------------------------------------===//
1020// FloorDivSIOp
1021//===----------------------------------------------------------------------===//
1022
1023OpFoldResult arith::FloorDivSIOp::fold(FoldAdaptor adaptor) {
1024 // TODO: floordivsi (x, 0) -> poison. Division by zero is undefined behaviour
1025 // and could fold to poison, but that would make the arith dialect depend on
1026 // the ub dialect to materialize ub.poison; left out for now.
1027
1028 // floordivsi (x, 1) -> x.
1029 if (matchPattern(adaptor.getRhs(), m_One()))
1030 return getLhs();
1031
1032 // floordivsi (0, x) -> 0. Division by zero is UB, so refining to 0 is valid.
1033 if (matchPattern(adaptor.getLhs(), m_Zero()))
1034 return getLhs();
1035
1036 // floordivsi (x, x) -> 1.
1037 if (getLhs() == getRhs())
1038 return getIntegerAttrOfType(getType(), 1);
1039
1040 // Don't fold if it would overflow or if it requires a division by zero.
1041 bool overflowOrDiv = false;
1043 adaptor.getOperands(), [&](APInt a, const APInt &b) {
1044 if (b.isZero()) {
1045 overflowOrDiv = true;
1046 return a;
1047 }
1048 return a.sfloordiv_ov(b, overflowOrDiv);
1049 });
1050
1051 return overflowOrDiv ? Attribute() : result;
1052}
1053
1054//===----------------------------------------------------------------------===//
1055// RemUIOp
1056//===----------------------------------------------------------------------===//
1057
1058OpFoldResult arith::RemUIOp::fold(FoldAdaptor adaptor) {
1059 // TODO: remui (x, 0) -> poison. Remainder by zero is undefined behaviour and
1060 // could fold to poison, but that would make the arith dialect depend on the
1061 // ub dialect to materialize ub.poison; left out for now.
1062
1063 // remui (x, 1) -> 0.
1064 if (matchPattern(adaptor.getRhs(), m_One()))
1065 return getIntegerAttrOfType(getType(), 0);
1066
1067 // remui (0, x) -> 0 and remui (x, x) -> 0. Division by zero is UB, so
1068 // refining to 0 is valid.
1069 if (matchPattern(adaptor.getLhs(), m_Zero()) || getLhs() == getRhs())
1070 return getIntegerAttrOfType(getType(), 0);
1071
1072 // Don't fold if it would require a division by zero.
1073 bool div0 = false;
1074 auto result = constFoldBinaryOp<IntegerAttr>(adaptor.getOperands(),
1075 [&](APInt a, const APInt &b) {
1076 if (div0 || b.isZero()) {
1077 div0 = true;
1078 return a;
1079 }
1080 return a.urem(b);
1081 });
1082
1083 return div0 ? Attribute() : result;
1084}
1085
1086Speculation::Speculatability arith::RemUIOp::getSpeculatability() {
1087 return getDivUISpeculatability(getRhs());
1088}
1089
1090//===----------------------------------------------------------------------===//
1091// RemSIOp
1092//===----------------------------------------------------------------------===//
1093
1094OpFoldResult arith::RemSIOp::fold(FoldAdaptor adaptor) {
1095 // TODO: remsi (x, 0) -> poison. Remainder by zero is undefined behaviour and
1096 // could fold to poison, but that would make the arith dialect depend on the
1097 // ub dialect to materialize ub.poison; left out for now.
1098
1099 // remsi (x, 1) -> 0.
1100 if (matchPattern(adaptor.getRhs(), m_One()))
1101 return getIntegerAttrOfType(getType(), 0);
1102
1103 // remsi (0, x) -> 0 and remsi (x, x) -> 0. Division by zero is UB, so
1104 // refining to 0 is valid.
1105 if (matchPattern(adaptor.getLhs(), m_Zero()) || getLhs() == getRhs())
1106 return getIntegerAttrOfType(getType(), 0);
1107
1108 // Don't fold if it would require a division by zero.
1109 bool div0 = false;
1110 auto result = constFoldBinaryOp<IntegerAttr>(adaptor.getOperands(),
1111 [&](APInt a, const APInt &b) {
1112 if (div0 || b.isZero()) {
1113 div0 = true;
1114 return a;
1115 }
1116 return a.srem(b);
1117 });
1118
1119 return div0 ? Attribute() : result;
1120}
1121
1122Speculation::Speculatability arith::RemSIOp::getSpeculatability() {
1123 // X % 0 => UB
1124 // X % -1 is well-defined (always 0), unlike X / -1 which can overflow.
1125 if (matchPattern(getRhs(), m_IntRangeWithoutZeroS()))
1127
1129}
1130
1131//===----------------------------------------------------------------------===//
1132// AndIOp
1133//===----------------------------------------------------------------------===//
1134
1135/// Fold `op(a, op(a, b))` to `op(a, b)` for an associative, commutative and
1136/// idempotent `op` (e.g. `and`, `or`).
1137template <typename OpTy>
1139 for (bool reversePrev : {false, true}) {
1140 auto prev = (reversePrev ? op.getRhs() : op.getLhs())
1141 .template getDefiningOp<OpTy>();
1142 if (!prev)
1143 continue;
1144
1145 Value other = (reversePrev ? op.getLhs() : op.getRhs());
1146 if (other != prev.getLhs() && other != prev.getRhs())
1147 continue;
1148
1149 return prev.getResult();
1150 }
1151 return {};
1152}
1153
1154OpFoldResult arith::AndIOp::fold(FoldAdaptor adaptor) {
1155 /// and(x, 0) -> 0
1156 if (matchPattern(adaptor.getRhs(), m_Zero()))
1157 return getRhs();
1158 /// and(x, allOnes) -> x
1159 APInt intValue;
1160 if (matchPattern(adaptor.getRhs(), m_ConstantInt(&intValue)) &&
1161 intValue.isAllOnes())
1162 return getLhs();
1163 /// and(x, not(x)) -> 0
1164 if (matchPattern(getRhs(), m_Op<XOrIOp>(matchers::m_Val(getLhs()),
1165 m_ConstantInt(&intValue))) &&
1166 intValue.isAllOnes())
1167 return Builder(getContext()).getZeroAttr(getType());
1168 /// and(not(x), x) -> 0
1169 if (matchPattern(getLhs(), m_Op<XOrIOp>(matchers::m_Val(getRhs()),
1170 m_ConstantInt(&intValue))) &&
1171 intValue.isAllOnes())
1172 return Builder(getContext()).getZeroAttr(getType());
1173
1174 /// and(a, and(a, b)) -> and(a, b)
1175 if (Value result = foldIdempotentOfSameOp(*this))
1176 return result;
1177
1179 adaptor.getOperands(),
1180 [](APInt a, const APInt &b) { return std::move(a) & b; });
1181}
1182
1183//===----------------------------------------------------------------------===//
1184// OrIOp
1185//===----------------------------------------------------------------------===//
1186
1187OpFoldResult arith::OrIOp::fold(FoldAdaptor adaptor) {
1188 if (APInt rhsVal; matchPattern(adaptor.getRhs(), m_ConstantInt(&rhsVal))) {
1189 /// or(x, 0) -> x
1190 if (rhsVal.isZero())
1191 return getLhs();
1192 /// or(x, <all ones>) -> <all ones>
1193 if (rhsVal.isAllOnes())
1194 return adaptor.getRhs();
1195 }
1196
1197 APInt intValue;
1198 /// or(x, xor(x, 1)) -> 1
1199 if (matchPattern(getRhs(), m_Op<XOrIOp>(matchers::m_Val(getLhs()),
1200 m_ConstantInt(&intValue))) &&
1201 intValue.isAllOnes())
1202 return getRhs().getDefiningOp<XOrIOp>().getRhs();
1203 /// or(xor(x, 1), x) -> 1
1204 if (matchPattern(getLhs(), m_Op<XOrIOp>(matchers::m_Val(getRhs()),
1205 m_ConstantInt(&intValue))) &&
1206 intValue.isAllOnes())
1207 return getLhs().getDefiningOp<XOrIOp>().getRhs();
1208
1209 /// or(a, or(a, b)) -> or(a, b)
1210 if (Value result = foldIdempotentOfSameOp(*this))
1211 return result;
1212
1214 adaptor.getOperands(),
1215 [](APInt a, const APInt &b) { return std::move(a) | b; });
1216}
1217
1218//===----------------------------------------------------------------------===//
1219// XOrIOp
1220//===----------------------------------------------------------------------===//
1221
1222OpFoldResult arith::XOrIOp::fold(FoldAdaptor adaptor) {
1223 /// xor(x, 0) -> x
1224 if (matchPattern(adaptor.getRhs(), m_Zero()))
1225 return getLhs();
1226 /// xor(x, x) -> 0
1227 if (getLhs() == getRhs()) {
1228 // A dynamically-shaped result cannot be a constant; bail before
1229 // getZeroAttr, which would assert on a non-static shape.
1230 auto shapedType = dyn_cast<ShapedType>(getType());
1231 if (!shapedType || shapedType.hasStaticShape())
1232 return Builder(getContext()).getZeroAttr(getType());
1233 }
1234 /// xor(xor(x, a), a) -> x
1235 /// xor(xor(a, x), a) -> x
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();
1241 }
1242 /// xor(a, xor(x, a)) -> x
1243 /// xor(a, xor(a, x)) -> x
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();
1249 }
1250
1252 adaptor.getOperands(),
1253 [](APInt a, const APInt &b) { return std::move(a) ^ b; });
1254}
1255
1256void arith::XOrIOp::getCanonicalizationPatterns(RewritePatternSet &patterns,
1257 MLIRContext *context) {
1258 patterns.add<XOrIXOrIConstant, XOrINotCmpI, XOrIOfExtUI, XOrIOfExtSI>(
1259 context);
1260}
1261
1262//===----------------------------------------------------------------------===//
1263// NegFOp
1264//===----------------------------------------------------------------------===//
1265
1266OpFoldResult arith::NegFOp::fold(FoldAdaptor adaptor) {
1267 /// negf(negf(x)) -> x
1268 if (auto op = this->getOperand().getDefiningOp<arith::NegFOp>())
1269 return op.getOperand();
1270 return constFoldUnaryOp<FloatAttr>(adaptor.getOperands(),
1271 [](const APFloat &a) { return -a; });
1272}
1273
1274//===----------------------------------------------------------------------===//
1275// FlushDenormalsOp
1276//===----------------------------------------------------------------------===//
1277
1278OpFoldResult arith::FlushDenormalsOp::fold(FoldAdaptor adaptor) {
1279 // TODO: Fold flush_denormals if the floating-point type does not support
1280 // denormals. There is currently no API to query this information from
1281 // APFloat.
1282
1283 // flush_denormals(flush_denormals(x)) -> flush_denormals(x)
1284 if (auto op = this->getOperand().getDefiningOp<arith::FlushDenormalsOp>())
1285 return op.getResult();
1286
1287 // Constant-fold flush_denormals if the operand is a constant.
1289 adaptor.getOperands(), [](const APFloat &a) {
1290 if (a.isDenormal())
1291 return APFloat::getZero(a.getSemantics(), a.isNegative());
1292 return a;
1293 });
1294}
1295
1296//===----------------------------------------------------------------------===//
1297// AddFOp
1298//===----------------------------------------------------------------------===//
1299
1300OpFoldResult arith::AddFOp::fold(FoldAdaptor adaptor) {
1301 // addf(x, -0) -> x
1302 if (matchPattern(adaptor.getRhs(), m_NegZeroFloat()))
1303 return getLhs();
1304 if (matchPattern(adaptor.getRhs(), m_PosZeroFloat()) &&
1305 bitEnumContainsAll(adaptor.getFastmath(), FastMathFlags::nsz))
1306 return getLhs();
1307
1308 auto rm = getRoundingmode();
1310 adaptor.getOperands(), [rm](const APFloat &a, const APFloat &b) {
1311 APFloat result(a);
1312 result.add(b, convertArithRoundingModeToLLVMIR(rm));
1313 return result;
1314 });
1315}
1316
1317void arith::AddFOp::getCanonicalizationPatterns(RewritePatternSet &patterns,
1318 MLIRContext *context) {
1319 patterns.add<AddFOfNegFLhs, AddFOfNegFRhs>(context);
1320}
1321
1322//===----------------------------------------------------------------------===//
1323// SubFOp
1324//===----------------------------------------------------------------------===//
1325
1326OpFoldResult arith::SubFOp::fold(FoldAdaptor adaptor) {
1327 // subf(x, +0) -> x
1328 if (matchPattern(adaptor.getRhs(), m_PosZeroFloat()))
1329 return getLhs();
1330 if (matchPattern(adaptor.getRhs(), m_NegZeroFloat()) &&
1331 bitEnumContainsAll(adaptor.getFastmath(), FastMathFlags::nsz))
1332 return getLhs();
1333
1334 auto rm = getRoundingmode();
1336 adaptor.getOperands(), [rm](const APFloat &a, const APFloat &b) {
1337 APFloat result(a);
1338 result.subtract(b, convertArithRoundingModeToLLVMIR(rm));
1339 return result;
1340 });
1341}
1342
1343void arith::SubFOp::getCanonicalizationPatterns(RewritePatternSet &patterns,
1344 MLIRContext *context) {
1345 patterns.add<SubFOfNegZero>(context);
1346}
1347
1348namespace {
1349
1350/// Narrow an extremum whose operands were extended from the result type:
1351///
1352/// trunc(extremum(ext(lhs), ext(rhs))) -> extremum(lhs, rhs)
1353///
1354/// The concrete extension is part of the pattern so each extremum is only
1355/// registered with extensions that preserve its ordering.
1356/// For floating-point types, also require the extension to preserve every
1357/// source value relevant under the extremum's fast-math flags.
1358template <typename TruncOp, typename ExtOp, typename ExtremumOp>
1359struct NarrowExtremum final : OpRewritePattern<TruncOp> {
1360 using OpRewritePattern<TruncOp>::OpRewritePattern;
1361
1362 LogicalResult matchAndRewrite(TruncOp truncOp,
1363 PatternRewriter &rewriter) const override {
1364 auto extremumOp = truncOp.getIn().template getDefiningOp<ExtremumOp>();
1365 if (!extremumOp || !extremumOp->hasOneUse())
1366 return failure();
1367
1368 auto lhsExt = extremumOp.getLhs().template getDefiningOp<ExtOp>();
1369 auto rhsExt = extremumOp.getRhs().template getDefiningOp<ExtOp>();
1370 if (!lhsExt || !rhsExt)
1371 return failure();
1372
1373 Value lhs = lhsExt.getIn();
1374 Value rhs = rhsExt.getIn();
1375 Type narrowType = truncOp.getType();
1376 if (lhs.getType() != narrowType || rhs.getType() != narrowType)
1377 return failure();
1378
1379 // A floating-point extension is not necessarily lossless between arbitrary
1380 // floating-point semantics, even when the destination has a larger bit
1381 // width. In particular, it may lose the sign of zero or quiet a signaling
1382 // NaN, either of which can change an extremum's result. `nnan` lets us
1383 // disregard NaN representation differences, but all other relevant source
1384 // values must be preserved.
1385 if (auto narrowFloatType =
1386 dyn_cast<FloatType>(getElementTypeOrSelf(narrowType))) {
1387 auto wideFloatType =
1388 dyn_cast<FloatType>(getElementTypeOrSelf(extremumOp.getType()));
1389 if (!wideFloatType)
1390 return failure();
1391
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>)
1398 ignoreNaNs =
1399 bitEnumContainsAll(extremumOp.getFastmath(), FastMathFlags::nnan);
1400 if (!llvm::APFloatBase::isLosslesslyConvertibleTo(
1401 narrowSemantics, wideSemantics, ignoreNaNs))
1402 return failure();
1403 }
1404
1405 rewriter.replaceOpWithNewOp<ExtremumOp>(
1406 truncOp, TypeRange{narrowType}, ValueRange{lhs, rhs},
1407 extremumOp.getProperties(),
1408 extremumOp->getDiscardableAttrDictionary().getValue());
1409 return success();
1410 }
1411};
1412
1413} // namespace
1414
1415//===----------------------------------------------------------------------===//
1416// MaximumFOp
1417//===----------------------------------------------------------------------===//
1418
1419OpFoldResult arith::MaximumFOp::fold(FoldAdaptor adaptor) {
1420 // maximumf(x,x) -> x
1421 if (getLhs() == getRhs())
1422 return getRhs();
1423
1424 // maximumf(x, -inf) -> x
1425 if (matchPattern(adaptor.getRhs(), m_NegInfFloat()))
1426 return getLhs();
1427
1428 return constFoldBinaryOp<FloatAttr>(adaptor.getOperands(), llvm::maximum);
1429}
1430
1431//===----------------------------------------------------------------------===//
1432// MaxNumFOp
1433//===----------------------------------------------------------------------===//
1434
1435OpFoldResult arith::MaxNumFOp::fold(FoldAdaptor adaptor) {
1436 // maxnumf(x,x) -> x
1437 if (getLhs() == getRhs())
1438 return getRhs();
1439
1440 // maxnumf(x, NaN) -> x
1441 if (matchPattern(adaptor.getRhs(), m_NaNFloat()))
1442 return getLhs();
1443
1444 return constFoldBinaryOp<FloatAttr>(adaptor.getOperands(), llvm::maxnum);
1445}
1446
1447//===----------------------------------------------------------------------===//
1448// MaximumNumFOp
1449//===----------------------------------------------------------------------===//
1450
1451OpFoldResult arith::MaximumNumFOp::fold(FoldAdaptor adaptor) {
1452 // maximumnumf(x,x) -> x
1453 if (getLhs() == getRhs())
1454 return getRhs();
1455
1456 // maximumnumf(x, NaN) -> x
1457 if (matchPattern(adaptor.getRhs(), m_NaNFloat()))
1458 return getLhs();
1459
1460 return constFoldBinaryOp<FloatAttr>(adaptor.getOperands(), llvm::maximumnum);
1461}
1462
1463//===----------------------------------------------------------------------===//
1464// MaxSIOp
1465//===----------------------------------------------------------------------===//
1466
1467OpFoldResult MaxSIOp::fold(FoldAdaptor adaptor) {
1468 // maxsi(x,x) -> x
1469 if (getLhs() == getRhs())
1470 return getRhs();
1471
1472 if (APInt intValue;
1473 matchPattern(adaptor.getRhs(), m_ConstantInt(&intValue))) {
1474 // maxsi(x,MAX_INT) -> MAX_INT
1475 if (intValue.isMaxSignedValue())
1476 return getRhs();
1477 // maxsi(x, MIN_INT) -> x
1478 if (intValue.isMinSignedValue())
1479 return getLhs();
1480 }
1481
1482 return constFoldBinaryOp<IntegerAttr>(adaptor.getOperands(),
1483 llvm::APIntOps::smax);
1484}
1485
1486//===----------------------------------------------------------------------===//
1487// MaxUIOp
1488//===----------------------------------------------------------------------===//
1489
1490OpFoldResult MaxUIOp::fold(FoldAdaptor adaptor) {
1491 // maxui(x,x) -> x
1492 if (getLhs() == getRhs())
1493 return getRhs();
1494
1495 if (APInt intValue;
1496 matchPattern(adaptor.getRhs(), m_ConstantInt(&intValue))) {
1497 // maxui(x,MAX_INT) -> MAX_INT
1498 if (intValue.isMaxValue())
1499 return getRhs();
1500 // maxui(x, MIN_INT) -> x
1501 if (intValue.isMinValue())
1502 return getLhs();
1503 }
1504
1505 return constFoldBinaryOp<IntegerAttr>(adaptor.getOperands(),
1506 llvm::APIntOps::umax);
1507}
1508
1509//===----------------------------------------------------------------------===//
1510// MinimumFOp
1511//===----------------------------------------------------------------------===//
1512
1513OpFoldResult arith::MinimumFOp::fold(FoldAdaptor adaptor) {
1514 // minimumf(x,x) -> x
1515 if (getLhs() == getRhs())
1516 return getRhs();
1517
1518 // minimumf(x, +inf) -> x
1519 if (matchPattern(adaptor.getRhs(), m_PosInfFloat()))
1520 return getLhs();
1521
1522 return constFoldBinaryOp<FloatAttr>(adaptor.getOperands(), llvm::minimum);
1523}
1524
1525//===----------------------------------------------------------------------===//
1526// MinNumFOp
1527//===----------------------------------------------------------------------===//
1528
1529OpFoldResult arith::MinNumFOp::fold(FoldAdaptor adaptor) {
1530 // minnumf(x,x) -> x
1531 if (getLhs() == getRhs())
1532 return getRhs();
1533
1534 // minnumf(x, NaN) -> x
1535 if (matchPattern(adaptor.getRhs(), m_NaNFloat()))
1536 return getLhs();
1537
1538 return constFoldBinaryOp<FloatAttr>(adaptor.getOperands(), llvm::minnum);
1539}
1540
1541//===----------------------------------------------------------------------===//
1542// MinimumNumFOp
1543//===----------------------------------------------------------------------===//
1544
1545OpFoldResult arith::MinimumNumFOp::fold(FoldAdaptor adaptor) {
1546 // minimumnumf(x,x) -> x
1547 if (getLhs() == getRhs())
1548 return getRhs();
1549
1550 // minimumnumf(x, NaN) -> x
1551 if (matchPattern(adaptor.getRhs(), m_NaNFloat()))
1552 return getLhs();
1553
1554 return constFoldBinaryOp<FloatAttr>(adaptor.getOperands(), llvm::minimumnum);
1555}
1556
1557//===----------------------------------------------------------------------===//
1558// MinSIOp
1559//===----------------------------------------------------------------------===//
1560
1561OpFoldResult MinSIOp::fold(FoldAdaptor adaptor) {
1562 // minsi(x,x) -> x
1563 if (getLhs() == getRhs())
1564 return getRhs();
1565
1566 if (APInt intValue;
1567 matchPattern(adaptor.getRhs(), m_ConstantInt(&intValue))) {
1568 // minsi(x,MIN_INT) -> MIN_INT
1569 if (intValue.isMinSignedValue())
1570 return getRhs();
1571 // minsi(x, MAX_INT) -> x
1572 if (intValue.isMaxSignedValue())
1573 return getLhs();
1574 }
1575
1576 return constFoldBinaryOp<IntegerAttr>(adaptor.getOperands(),
1577 llvm::APIntOps::smin);
1578}
1579
1580//===----------------------------------------------------------------------===//
1581// MinUIOp
1582//===----------------------------------------------------------------------===//
1583
1584OpFoldResult MinUIOp::fold(FoldAdaptor adaptor) {
1585 // minui(x,x) -> x
1586 if (getLhs() == getRhs())
1587 return getRhs();
1588
1589 if (APInt intValue;
1590 matchPattern(adaptor.getRhs(), m_ConstantInt(&intValue))) {
1591 // minui(x,MIN_INT) -> MIN_INT
1592 if (intValue.isMinValue())
1593 return getRhs();
1594 // minui(x, MAX_INT) -> x
1595 if (intValue.isMaxValue())
1596 return getLhs();
1597 }
1598
1599 return constFoldBinaryOp<IntegerAttr>(adaptor.getOperands(),
1600 llvm::APIntOps::umin);
1601}
1602
1603//===----------------------------------------------------------------------===//
1604// MulFOp
1605//===----------------------------------------------------------------------===//
1606
1607OpFoldResult arith::MulFOp::fold(FoldAdaptor adaptor) {
1608 // mulf(x, 1) -> x
1609 if (matchPattern(adaptor.getRhs(), m_OneFloat()))
1610 return getLhs();
1611
1612 if (arith::bitEnumContainsAll(getFastmath(), arith::FastMathFlags::nnan |
1613 arith::FastMathFlags::nsz)) {
1614 // mulf(x, 0) -> 0
1615 if (matchPattern(adaptor.getRhs(), m_AnyZeroFloat()))
1616 return getRhs();
1617 }
1618
1619 auto rm = getRoundingmode();
1621 adaptor.getOperands(), [rm](const APFloat &a, const APFloat &b) {
1622 APFloat result(a);
1623 result.multiply(b, convertArithRoundingModeToLLVMIR(rm));
1624 return result;
1625 });
1626}
1627
1628void arith::MulFOp::getCanonicalizationPatterns(RewritePatternSet &patterns,
1629 MLIRContext *context) {
1630 patterns.add<MulFOfNegF>(context);
1631}
1632
1633//===----------------------------------------------------------------------===//
1634// DivFOp
1635//===----------------------------------------------------------------------===//
1636
1637OpFoldResult arith::DivFOp::fold(FoldAdaptor adaptor) {
1638 // divf(x, 1) -> x
1639 if (matchPattern(adaptor.getRhs(), m_OneFloat()))
1640 return getLhs();
1641
1642 auto rm = getRoundingmode();
1644 adaptor.getOperands(), [rm](const APFloat &a, const APFloat &b) {
1645 APFloat result(a);
1646 result.divide(b, convertArithRoundingModeToLLVMIR(rm));
1647 return result;
1648 });
1649}
1650
1651void arith::DivFOp::getCanonicalizationPatterns(RewritePatternSet &patterns,
1652 MLIRContext *context) {
1653 patterns.add<DivFOfNegF>(context);
1654}
1655
1656//===----------------------------------------------------------------------===//
1657// RemFOp
1658//===----------------------------------------------------------------------===//
1659
1660OpFoldResult arith::RemFOp::fold(FoldAdaptor adaptor) {
1661 return constFoldBinaryOp<FloatAttr>(adaptor.getOperands(),
1662 [](const APFloat &a, const APFloat &b) {
1663 APFloat result(a);
1664 // APFloat::mod() offers the remainder
1665 // behavior we want, i.e. the result has
1666 // the sign of LHS operand.
1667 (void)result.mod(b);
1668 return result;
1669 });
1670}
1671
1672//===----------------------------------------------------------------------===//
1673// Utility functions for verifying cast ops
1674//===----------------------------------------------------------------------===//
1675
1676template <typename... Types>
1677using type_list = std::tuple<Types...> *;
1678
1679/// Returns a non-null type only if the provided type is one of the allowed
1680/// types or one of the allowed shaped types of the allowed types. Returns the
1681/// element type if a valid shaped type is provided.
1682template <typename... ShapedTypes, typename... ElementTypes>
1685 if (llvm::isa<ShapedType>(type) && !llvm::isa<ShapedTypes...>(type))
1686 return {};
1687
1688 auto underlyingType = getElementTypeOrSelf(type);
1689 if (!llvm::isa<ElementTypes...>(underlyingType))
1690 return {};
1691
1692 return underlyingType;
1693}
1694
1695/// Get allowed underlying types for vectors and tensors.
1696template <typename... ElementTypes>
1701
1702/// Get allowed underlying types for vectors, tensors, and memrefs.
1703template <typename... ElementTypes>
1709
1710/// Return false if both types are ranked tensor with mismatching encoding.
1711static bool hasSameEncoding(Type typeA, Type typeB) {
1712 auto rankedTensorA = dyn_cast<RankedTensorType>(typeA);
1713 auto rankedTensorB = dyn_cast<RankedTensorType>(typeB);
1714 if (!rankedTensorA || !rankedTensorB)
1715 return true;
1716 return rankedTensorA.getEncoding() == rankedTensorB.getEncoding();
1717}
1718
1720 if (inputs.size() != 1 || outputs.size() != 1)
1721 return false;
1722 if (!hasSameEncoding(inputs.front(), outputs.front()))
1723 return false;
1724 return succeeded(verifyCompatibleShapes(inputs.front(), outputs.front()));
1725}
1726
1727//===----------------------------------------------------------------------===//
1728// Verifiers for integer and floating point extension/truncation ops
1729//===----------------------------------------------------------------------===//
1730
1731// Extend ops can only extend to a wider type.
1732template <typename ValType, typename Op>
1733static LogicalResult verifyExtOp(Op op) {
1734 Type srcType = getElementTypeOrSelf(op.getIn().getType());
1735 Type dstType = getElementTypeOrSelf(op.getType());
1736
1737 if (llvm::cast<ValType>(srcType).getWidth() >=
1738 llvm::cast<ValType>(dstType).getWidth())
1739 return op.emitError("result type ")
1740 << dstType << " must be wider than operand type " << srcType;
1741
1742 return success();
1743}
1744
1745// Truncate ops can only truncate to a shorter type.
1746template <typename ValType, typename Op>
1747static LogicalResult verifyTruncateOp(Op op) {
1748 Type srcType = getElementTypeOrSelf(op.getIn().getType());
1749 Type dstType = getElementTypeOrSelf(op.getType());
1750
1751 if (llvm::cast<ValType>(srcType).getWidth() <=
1752 llvm::cast<ValType>(dstType).getWidth())
1753 return op.emitError("result type ")
1754 << dstType << " must be shorter than operand type " << srcType;
1755
1756 return success();
1757}
1758
1759/// Validate a cast that changes the width of a type.
1760template <template <typename> class WidthComparator, typename... ElementTypes>
1761static bool checkWidthChangeCast(TypeRange inputs, TypeRange outputs) {
1762 if (!areValidCastInputsAndOutputs(inputs, outputs))
1763 return false;
1764
1765 auto srcType = getTypeIfLike<ElementTypes...>(inputs.front());
1766 auto dstType = getTypeIfLike<ElementTypes...>(outputs.front());
1767 if (!srcType || !dstType)
1768 return false;
1769
1770 return WidthComparator<unsigned>()(dstType.getIntOrFloatBitWidth(),
1771 srcType.getIntOrFloatBitWidth());
1772}
1773
1774/// Attempts to convert `sourceValue` to an APFloat value with
1775/// `targetSemantics` and `roundingMode`, without any information loss.
1776static FailureOr<APFloat>
1777convertFloatValue(APFloat sourceValue,
1778 const llvm::fltSemantics &targetSemantics,
1779 llvm::RoundingMode roundingMode = kDefaultRoundingMode) {
1780 // Reject special values that are not representable in the target type before
1781 // calling APFloat::convert, which would llvm_unreachable on them.
1782 using fltNonfiniteBehavior = llvm::fltNonfiniteBehavior;
1783 if (sourceValue.isInfinity() &&
1784 (targetSemantics.nonFiniteBehavior == fltNonfiniteBehavior::NanOnly ||
1785 targetSemantics.nonFiniteBehavior == fltNonfiniteBehavior::FiniteOnly))
1786 return failure();
1787 if (sourceValue.isNaN() &&
1788 targetSemantics.nonFiniteBehavior == fltNonfiniteBehavior::FiniteOnly)
1789 return failure();
1790
1791 bool losesInfo = false;
1792 auto status = sourceValue.convert(targetSemantics, roundingMode, &losesInfo);
1793 if (losesInfo || status != APFloat::opOK)
1794 return failure();
1795
1796 return sourceValue;
1797}
1798
1799//===----------------------------------------------------------------------===//
1800// ExtUIOp
1801//===----------------------------------------------------------------------===//
1802
1803OpFoldResult arith::ExtUIOp::fold(FoldAdaptor adaptor) {
1804 if (auto lhs = getIn().getDefiningOp<ExtUIOp>()) {
1805 // Only the inner extension's nneg speaks about the surviving source; the
1806 // outer flag described the already-extended value.
1807 setNonNeg(lhs.getNonNeg());
1808 getInMutable().assign(lhs.getIn());
1809 return getResult();
1810 }
1811
1812 Type resType = getElementTypeOrSelf(getType());
1813 unsigned bitWidth = llvm::cast<IntegerType>(resType).getWidth();
1815 adaptor.getOperands(), getType(),
1816 [bitWidth](const APInt &a, bool &castStatus) {
1817 return a.zext(bitWidth);
1818 });
1819}
1820
1821bool arith::ExtUIOp::areCastCompatible(TypeRange inputs, TypeRange outputs) {
1823}
1824
1825LogicalResult arith::ExtUIOp::verify() {
1826 return verifyExtOp<IntegerType>(*this);
1827}
1828
1829//===----------------------------------------------------------------------===//
1830// ExtSIOp
1831//===----------------------------------------------------------------------===//
1832
1833OpFoldResult arith::ExtSIOp::fold(FoldAdaptor adaptor) {
1834 if (auto lhs = getIn().getDefiningOp<ExtSIOp>()) {
1835 getInMutable().assign(lhs.getIn());
1836 return getResult();
1837 }
1838
1839 Type resType = getElementTypeOrSelf(getType());
1840 unsigned bitWidth = llvm::cast<IntegerType>(resType).getWidth();
1842 adaptor.getOperands(), getType(),
1843 [bitWidth](const APInt &a, bool &castStatus) {
1844 return a.sext(bitWidth);
1845 });
1846}
1847
1848bool arith::ExtSIOp::areCastCompatible(TypeRange inputs, TypeRange outputs) {
1850}
1851
1852void arith::ExtSIOp::getCanonicalizationPatterns(RewritePatternSet &patterns,
1853 MLIRContext *context) {
1854 patterns.add<ExtSIOfExtUI>(context);
1855}
1856
1857LogicalResult arith::ExtSIOp::verify() {
1858 return verifyExtOp<IntegerType>(*this);
1859}
1860
1861//===----------------------------------------------------------------------===//
1862// ExtFOp
1863//===----------------------------------------------------------------------===//
1864
1865/// Fold extension of float constants when there is no information loss due the
1866/// difference in fp semantics.
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();
1880 }
1881 }
1882 }
1883
1884 auto resElemType = cast<FloatType>(getElementTypeOrSelf(getType()));
1885 const llvm::fltSemantics &targetSemantics = resElemType.getFloatSemantics();
1887 adaptor.getOperands(), getType(),
1888 [&targetSemantics](const APFloat &a, bool &castStatus) {
1889 FailureOr<APFloat> result = convertFloatValue(a, targetSemantics);
1890 if (failed(result)) {
1891 castStatus = false;
1892 return a;
1893 }
1894 return *result;
1895 });
1896}
1897
1898bool arith::ExtFOp::areCastCompatible(TypeRange inputs, TypeRange outputs) {
1899 return checkWidthChangeCast<std::greater, FloatType>(inputs, outputs);
1900}
1901
1902LogicalResult arith::ExtFOp::verify() { return verifyExtOp<FloatType>(*this); }
1903
1904//===----------------------------------------------------------------------===//
1905// ScalingExtFOp
1906//===----------------------------------------------------------------------===//
1907
1908/// Fold `calculate` element-wise over the operands of a scaling cast op. The
1909/// `constFoldBinaryOp` helpers cannot be used: they bail out unless both
1910/// operands have the same type, and `in` and `scale` never do.
1912 Attribute inAttr, Attribute scaleAttr, Type resultType,
1913 function_ref<std::optional<APFloat>(const APFloat &, const APFloat &)>
1914 calculate) {
1915 // Poison propagates, as it does in the generic constant folders.
1916 if (isa_and_nonnull<ub::PoisonAttr>(inAttr))
1917 return inAttr;
1918 if (isa_and_nonnull<ub::PoisonAttr>(scaleAttr))
1919 return scaleAttr;
1920
1921 if (!inAttr || !scaleAttr || !resultType)
1922 return {};
1923
1924 if (auto inFloat = dyn_cast<FloatAttr>(inAttr)) {
1925 auto scaleFloat = dyn_cast<FloatAttr>(scaleAttr);
1926 if (!scaleFloat)
1927 return {};
1928 std::optional<APFloat> result =
1929 calculate(inFloat.getValue(), scaleFloat.getValue());
1930 if (!result)
1931 return {};
1932 return FloatAttr::get(resultType, *result);
1933 }
1934
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())
1941 return {};
1942
1943 // Both operands are splats, so avoid expanding the elements out.
1944 if (inElements.isSplat() && scaleElements.isSplat()) {
1945 std::optional<APFloat> result =
1946 calculate(inElements.getSplatValue<APFloat>(),
1947 scaleElements.getSplatValue<APFloat>());
1948 if (!result)
1949 return {};
1950 return DenseElementsAttr::get(shapedResultType, *result);
1951 }
1952
1953 SmallVector<APFloat> results;
1954 results.reserve(inElements.getNumElements());
1955 for (const auto &[in, scale] : llvm::zip_equal(inElements, scaleElements)) {
1956 std::optional<APFloat> result = calculate(in, scale);
1957 if (!result)
1958 return {};
1959 results.push_back(*result);
1960 }
1961 return DenseElementsAttr::get(shapedResultType, results);
1962}
1963
1964/// Only scales that already are f8E8M0FNU fold. What a wider scale means is
1965/// unsettled -- the tree does not say whether truncating one to f8E8M0FNU
1966/// rounds or takes its exponent -- so a folder should not settle it, see
1967/// https://github.com/llvm/llvm-project/issues/215295.
1968static bool isFoldableScalingScale(Value scale) {
1969 return isa<Float8E8M0FNUType>(getElementTypeOrSelf(scale.getType()));
1970}
1971
1972OpFoldResult arith::ScalingExtFOp::fold(FoldAdaptor adaptor) {
1973 // scaling_extf(in, scale) -> mulf(extf(in), extf(scale)), matching the
1974 // expansion in ExpandOps.cpp. As in arith.extf, the widening steps only fold
1975 // when they are lossless.
1976 if (!isFoldableScalingScale(getScale()))
1977 return {};
1978
1979 auto resElemType = cast<FloatType>(getElementTypeOrSelf(getType()));
1980 const llvm::fltSemantics &resSemantics = resElemType.getFloatSemantics();
1981 return foldScalingCastOp(
1982 adaptor.getIn(), adaptor.getScale(), getType(),
1983 [&resSemantics](const APFloat &in,
1984 const APFloat &scale) -> std::optional<APFloat> {
1985 FailureOr<APFloat> inExt = convertFloatValue(in, resSemantics);
1986 FailureOr<APFloat> scaleExt = convertFloatValue(scale, resSemantics);
1987 if (failed(inExt) || failed(scaleExt))
1988 return std::nullopt;
1989 APFloat result(*inExt);
1990 result.multiply(*scaleExt, kDefaultRoundingMode);
1991 return result;
1992 });
1993}
1994
1995bool arith::ScalingExtFOp::areCastCompatible(TypeRange inputs,
1996 TypeRange outputs) {
1997 return checkWidthChangeCast<std::greater, FloatType>(inputs.front(), outputs);
1998}
1999
2000LogicalResult arith::ScalingExtFOp::verify() {
2001 return verifyExtOp<FloatType>(*this);
2002}
2003
2004//===----------------------------------------------------------------------===//
2005// TruncIOp
2006//===----------------------------------------------------------------------===//
2007
2008OpFoldResult arith::TruncIOp::fold(FoldAdaptor adaptor) {
2009 if (matchPattern(getOperand(), m_Op<arith::ExtUIOp>()) ||
2010 matchPattern(getOperand(), m_Op<arith::ExtSIOp>())) {
2011 Value src = getOperand().getDefiningOp()->getOperand(0);
2012 Type srcType = getElementTypeOrSelf(src.getType());
2013 Type dstType = getElementTypeOrSelf(getType());
2014 // trunci(zexti(a)) -> trunci(a)
2015 // trunci(sexti(a)) -> trunci(a)
2016 if (llvm::cast<IntegerType>(srcType).getWidth() >
2017 llvm::cast<IntegerType>(dstType).getWidth()) {
2018 setOperand(src);
2019 return getResult();
2020 }
2021
2022 // trunci(zexti(a)) -> a
2023 // trunci(sexti(a)) -> a
2024 if (srcType == dstType)
2025 return src;
2026 }
2027
2028 // trunci(trunci(a)) -> trunci(a))
2029 if (matchPattern(getOperand(), m_Op<arith::TruncIOp>())) {
2030 setOperand(getOperand().getDefiningOp()->getOperand(0));
2031 return getResult();
2032 }
2033
2034 Type resType = getElementTypeOrSelf(getType());
2035 unsigned bitWidth = llvm::cast<IntegerType>(resType).getWidth();
2037 adaptor.getOperands(), getType(),
2038 [bitWidth](const APInt &a, bool &castStatus) {
2039 return a.trunc(bitWidth);
2040 });
2041}
2042
2043bool arith::TruncIOp::areCastCompatible(TypeRange inputs, TypeRange outputs) {
2044 return checkWidthChangeCast<std::less, IntegerType>(inputs, outputs);
2045}
2046
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);
2054}
2055
2056LogicalResult arith::TruncIOp::verify() {
2057 return verifyTruncateOp<IntegerType>(*this);
2058}
2059
2060//===----------------------------------------------------------------------===//
2061// TruncFOp
2062//===----------------------------------------------------------------------===//
2063
2064/// Perform safe const propagation for truncf, i.e., only propagate if FP value
2065/// can be represented without precision loss.
2066OpFoldResult arith::TruncFOp::fold(FoldAdaptor adaptor) {
2067 auto resElemType = cast<FloatType>(getElementTypeOrSelf(getType()));
2068 if (auto extOp = getOperand().getDefiningOp<arith::ExtFOp>()) {
2069 Value src = extOp.getIn();
2070 auto srcType = cast<FloatType>(getElementTypeOrSelf(src.getType()));
2071 auto intermediateType =
2072 cast<FloatType>(getElementTypeOrSelf(extOp.getType()));
2073 // Check whether every source value round-trips through the intermediate
2074 // type, including signaling NaNs and signed zero.
2075 if (llvm::APFloatBase::isLosslesslyConvertibleTo(
2076 srcType.getFloatSemantics(),
2077 intermediateType.getFloatSemantics())) {
2078 // truncf(extf(a)) -> truncf(a)
2079 if (srcType.getWidth() > resElemType.getWidth()) {
2080 setOperand(src);
2081 return getResult();
2082 }
2083
2084 // truncf(extf(a)) -> a
2085 if (srcType == resElemType)
2086 return src;
2087 }
2088 }
2089
2090 const llvm::fltSemantics &targetSemantics = resElemType.getFloatSemantics();
2092 adaptor.getOperands(), getType(),
2093 [this, &targetSemantics](const APFloat &a, bool &castStatus) {
2094 llvm::RoundingMode llvmRoundingMode =
2095 convertArithRoundingModeToLLVMIR(getRoundingmode());
2096 FailureOr<APFloat> result =
2097 convertFloatValue(a, targetSemantics, llvmRoundingMode);
2098 if (failed(result)) {
2099 castStatus = false;
2100 return a;
2101 }
2102 return *result;
2103 });
2104}
2105
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);
2115}
2116
2117bool arith::TruncFOp::areCastCompatible(TypeRange inputs, TypeRange outputs) {
2118 return checkWidthChangeCast<std::less, FloatType>(inputs, outputs);
2119}
2120
2121LogicalResult arith::TruncFOp::verify() {
2122 return verifyTruncateOp<FloatType>(*this);
2123}
2124
2125//===----------------------------------------------------------------------===//
2126// ConvertFOp
2127//===----------------------------------------------------------------------===//
2128
2129OpFoldResult arith::ConvertFOp::fold(FoldAdaptor adaptor) {
2130 auto resElemType = cast<FloatType>(getElementTypeOrSelf(getType()));
2131 const llvm::fltSemantics &targetSemantics = resElemType.getFloatSemantics();
2133 adaptor.getOperands(), getType(),
2134 [this, &targetSemantics](const APFloat &a, bool &castStatus) {
2135 llvm::RoundingMode llvmRoundingMode =
2136 convertArithRoundingModeToLLVMIR(getRoundingmode());
2137 FailureOr<APFloat> result =
2138 convertFloatValue(a, targetSemantics, llvmRoundingMode);
2139 if (failed(result)) {
2140 castStatus = false;
2141 return a;
2142 }
2143 return *result;
2144 });
2145}
2146
2147bool arith::ConvertFOp::areCastCompatible(TypeRange inputs, TypeRange outputs) {
2148 if (!areValidCastInputsAndOutputs(inputs, outputs))
2149 return false;
2150 auto srcType = getTypeIfLike<FloatType>(inputs.front());
2151 auto dstType = getTypeIfLike<FloatType>(outputs.front());
2152 if (!srcType || !dstType)
2153 return false;
2154 return srcType != dstType &&
2155 srcType.getIntOrFloatBitWidth() == dstType.getIntOrFloatBitWidth();
2156}
2157
2158LogicalResult arith::ConvertFOp::verify() {
2159 auto srcType = cast<FloatType>(getElementTypeOrSelf(getIn().getType()));
2160 auto dstType = cast<FloatType>(getElementTypeOrSelf(getType()));
2161 if (srcType == dstType)
2162 return emitError("result element type ")
2163 << dstType << " must be different from operand element type "
2164 << srcType;
2165 if (srcType.getWidth() != dstType.getWidth())
2166 return emitError("result element type ")
2167 << dstType << " must have the same bitwidth as operand element type "
2168 << srcType;
2169 return success();
2170}
2171
2172//===----------------------------------------------------------------------===//
2173// ScalingTruncFOp
2174//===----------------------------------------------------------------------===//
2175
2176OpFoldResult arith::ScalingTruncFOp::fold(FoldAdaptor adaptor) {
2177 // scaling_truncf(in, scale) -> truncf(in / extf(scale)), matching the
2178 // expansion in ExpandOps.cpp. Unlike scaling_extf, the scale is widened to
2179 // the type of `in` rather than to the result type.
2180 if (!isFoldableScalingScale(getScale()))
2181 return {};
2182
2183 auto inElemType = cast<FloatType>(getElementTypeOrSelf(getIn().getType()));
2184 auto resElemType = cast<FloatType>(getElementTypeOrSelf(getType()));
2185 const llvm::fltSemantics &inSemantics = inElemType.getFloatSemantics();
2186 const llvm::fltSemantics &resSemantics = resElemType.getFloatSemantics();
2187 llvm::RoundingMode roundingMode =
2188 convertArithRoundingModeToLLVMIR(getRoundingmode());
2189 return foldScalingCastOp(
2190 adaptor.getIn(), adaptor.getScale(), getType(),
2191 [&](const APFloat &in, const APFloat &scale) -> std::optional<APFloat> {
2192 FailureOr<APFloat> scaleExt = convertFloatValue(scale, inSemantics);
2193 if (failed(scaleExt))
2194 return std::nullopt;
2195 APFloat quotient(in);
2196 quotient.divide(*scaleExt, kDefaultRoundingMode);
2197 FailureOr<APFloat> result =
2198 convertFloatValue(quotient, resSemantics, roundingMode);
2199 if (failed(result))
2200 return std::nullopt;
2201 return *result;
2202 });
2203}
2204
2205bool arith::ScalingTruncFOp::areCastCompatible(TypeRange inputs,
2206 TypeRange outputs) {
2207 return checkWidthChangeCast<std::less, FloatType>(inputs.front(), outputs);
2208}
2209
2210LogicalResult arith::ScalingTruncFOp::verify() {
2211 return verifyTruncateOp<FloatType>(*this);
2212}
2213
2214//===----------------------------------------------------------------------===//
2215// AndIOp
2216//===----------------------------------------------------------------------===//
2217
2218void arith::AndIOp::getCanonicalizationPatterns(RewritePatternSet &patterns,
2219 MLIRContext *context) {
2220 patterns.add<AndIAndIConstant, AndOfExtUI, AndOfExtSI>(context);
2221}
2222
2223//===----------------------------------------------------------------------===//
2224// OrIOp
2225//===----------------------------------------------------------------------===//
2226
2227void arith::OrIOp::getCanonicalizationPatterns(RewritePatternSet &patterns,
2228 MLIRContext *context) {
2229 patterns.add<OrIOrIConstant, OrOfExtUI, OrOfExtSI>(context);
2230}
2231
2232//===----------------------------------------------------------------------===//
2233// Verifiers for casts between integers and floats.
2234//===----------------------------------------------------------------------===//
2235
2236template <typename From, typename To>
2237static bool checkIntFloatCast(TypeRange inputs, TypeRange outputs) {
2238 if (!areValidCastInputsAndOutputs(inputs, outputs))
2239 return false;
2240
2241 auto srcType = getTypeIfLike<From>(inputs.front());
2242 auto dstType = getTypeIfLike<To>(outputs.back());
2243
2244 return srcType && dstType;
2245}
2246
2247//===----------------------------------------------------------------------===//
2248// UIToFPOp
2249//===----------------------------------------------------------------------===//
2250
2251bool arith::UIToFPOp::areCastCompatible(TypeRange inputs, TypeRange outputs) {
2252 return checkIntFloatCast<IntegerType, FloatType>(inputs, outputs);
2253}
2254
2255OpFoldResult arith::UIToFPOp::fold(FoldAdaptor adaptor) {
2256 Type resEleType = getElementTypeOrSelf(getType());
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, /*IsSigned=*/false,
2264 APFloat::rmNearestTiesToEven);
2265 return apf;
2266 });
2267}
2268
2269void arith::UIToFPOp::getCanonicalizationPatterns(RewritePatternSet &patterns,
2270 MLIRContext *context) {
2271 patterns.add<UIToFPOfExtUI>(context);
2272}
2273
2274//===----------------------------------------------------------------------===//
2275// SIToFPOp
2276//===----------------------------------------------------------------------===//
2277
2278bool arith::SIToFPOp::areCastCompatible(TypeRange inputs, TypeRange outputs) {
2279 return checkIntFloatCast<IntegerType, FloatType>(inputs, outputs);
2280}
2281
2282OpFoldResult arith::SIToFPOp::fold(FoldAdaptor adaptor) {
2283 Type resEleType = getElementTypeOrSelf(getType());
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, /*IsSigned=*/true,
2291 APFloat::rmNearestTiesToEven);
2292 return apf;
2293 });
2294}
2295
2296void arith::SIToFPOp::getCanonicalizationPatterns(RewritePatternSet &patterns,
2297 MLIRContext *context) {
2298 patterns.add<SIToFPOfExtSI, SIToFPOfExtUI>(context);
2299}
2300
2301//===----------------------------------------------------------------------===//
2302// FPToUIOp
2303//===----------------------------------------------------------------------===//
2304
2305bool arith::FPToUIOp::areCastCompatible(TypeRange inputs, TypeRange outputs) {
2306 return checkIntFloatCast<FloatType, IntegerType>(inputs, outputs);
2307}
2308
2309OpFoldResult arith::FPToUIOp::fold(FoldAdaptor adaptor) {
2310 Type resType = getElementTypeOrSelf(getType());
2311 unsigned bitWidth = llvm::cast<IntegerType>(resType).getWidth();
2313 adaptor.getOperands(), getType(),
2314 [&bitWidth](const APFloat &a, bool &castStatus) {
2315 bool ignored;
2316 APSInt api(bitWidth, /*isUnsigned=*/true);
2317 castStatus = APFloat::opInvalidOp !=
2318 a.convertToInteger(api, APFloat::rmTowardZero, &ignored);
2319 return api;
2320 });
2321}
2322
2323//===----------------------------------------------------------------------===//
2324// FPToSIOp
2325//===----------------------------------------------------------------------===//
2326
2327bool arith::FPToSIOp::areCastCompatible(TypeRange inputs, TypeRange outputs) {
2328 return checkIntFloatCast<FloatType, IntegerType>(inputs, outputs);
2329}
2330
2331OpFoldResult arith::FPToSIOp::fold(FoldAdaptor adaptor) {
2332 Type resType = getElementTypeOrSelf(getType());
2333 unsigned bitWidth = llvm::cast<IntegerType>(resType).getWidth();
2335 adaptor.getOperands(), getType(),
2336 [&bitWidth](const APFloat &a, bool &castStatus) {
2337 bool ignored;
2338 APSInt api(bitWidth, /*isUnsigned=*/false);
2339 castStatus = APFloat::opInvalidOp !=
2340 a.convertToInteger(api, APFloat::rmTowardZero, &ignored);
2341 return api;
2342 });
2343}
2344
2345//===----------------------------------------------------------------------===//
2346// IndexCastOp
2347//===----------------------------------------------------------------------===//
2348
2349/// Return the bit-width of \p t for the purpose of index_cast width checks.
2350/// For vector types use the element type; index maps to its internal storage
2351/// width (64 on all current targets).
2352static unsigned getIndexCastWidth(Type t) {
2353 if (auto intTy = dyn_cast<IntegerType>(getElementTypeOrSelf(t)))
2354 return intTy.getWidth();
2355 return IndexType::kInternalStorageBitWidth;
2356}
2357
2358static bool areIndexCastCompatible(TypeRange inputs, TypeRange outputs) {
2359 if (!areValidCastInputsAndOutputs(inputs, outputs))
2360 return false;
2361
2362 auto srcType = getTypeIfLikeOrMemRef<IntegerType, IndexType>(inputs.front());
2363 auto dstType = getTypeIfLikeOrMemRef<IntegerType, IndexType>(outputs.front());
2364 if (!srcType || !dstType)
2365 return false;
2366
2367 return (srcType.isIndex() && dstType.isSignlessInteger()) ||
2368 (srcType.isSignlessInteger() && dstType.isIndex());
2369}
2370
2371bool arith::IndexCastOp::areCastCompatible(TypeRange inputs,
2372 TypeRange outputs) {
2373 return areIndexCastCompatible(inputs, outputs);
2374}
2375
2376OpFoldResult arith::IndexCastOp::fold(FoldAdaptor adaptor) {
2377 // index_cast(constant) -> constant
2378 unsigned resultBitwidth = 64; // Default for index integer attributes.
2379 if (auto intTy = dyn_cast<IntegerType>(getElementTypeOrSelf(getType())))
2380 resultBitwidth = intTy.getWidth();
2381
2382 if (auto foldResult = constFoldCastOp<IntegerAttr, IntegerAttr>(
2383 adaptor.getOperands(), getType(),
2384 [resultBitwidth](const APInt &a, bool & /*castStatus*/) {
2385 return a.sextOrTrunc(resultBitwidth);
2386 }))
2387 return foldResult;
2388
2389 // index_cast(index_cast(x : A) : B) : A -> x, but only when B is at least
2390 // as wide as A. If B is narrower, the inner cast truncates and the outer
2391 // cast sign-extends, so the round-trip is lossy.
2392 if (auto inner = getOperand().getDefiningOp<arith::IndexCastOp>()) {
2393 Value x = inner.getOperand();
2394 if (x.getType() == getType()) {
2395 if (getIndexCastWidth(inner.getType()) >= getIndexCastWidth(x.getType()))
2396 return x;
2397 }
2398 }
2399 return {};
2400}
2401
2402void arith::IndexCastOp::getCanonicalizationPatterns(
2403 RewritePatternSet &patterns, MLIRContext *context) {
2404 patterns.add<IndexCastOfExtSI>(context);
2405}
2406
2407//===----------------------------------------------------------------------===//
2408// IndexCastUIOp
2409//===----------------------------------------------------------------------===//
2410
2411bool arith::IndexCastUIOp::areCastCompatible(TypeRange inputs,
2412 TypeRange outputs) {
2413 return areIndexCastCompatible(inputs, outputs);
2414}
2415
2416OpFoldResult arith::IndexCastUIOp::fold(FoldAdaptor adaptor) {
2417 // index_castui(constant) -> constant
2418 unsigned resultBitwidth = 64; // Default for index integer attributes.
2419 if (auto intTy = dyn_cast<IntegerType>(getElementTypeOrSelf(getType())))
2420 resultBitwidth = intTy.getWidth();
2421
2422 if (auto foldResult = constFoldCastOp<IntegerAttr, IntegerAttr>(
2423 adaptor.getOperands(), getType(),
2424 [resultBitwidth](const APInt &a, bool & /*castStatus*/) {
2425 return a.zextOrTrunc(resultBitwidth);
2426 }))
2427 return foldResult;
2428
2429 // index_castui(index_castui(x : A) : B) : A -> x, but only when B is at
2430 // least as wide as A. If B is narrower, the inner cast truncates and the
2431 // outer cast zero-extends, so the round-trip is lossy.
2432 if (auto inner = getOperand().getDefiningOp<arith::IndexCastUIOp>()) {
2433 Value x = inner.getOperand();
2434 if (x.getType() == getType()) {
2435 if (getIndexCastWidth(inner.getType()) >= getIndexCastWidth(x.getType()))
2436 return x;
2437 }
2438 }
2439 return {};
2440}
2441
2442void arith::IndexCastUIOp::getCanonicalizationPatterns(
2443 RewritePatternSet &patterns, MLIRContext *context) {
2444 patterns.add<IndexCastUIOfExtUI>(context);
2445}
2446
2447//===----------------------------------------------------------------------===//
2448// BitcastOp
2449//===----------------------------------------------------------------------===//
2450
2451bool arith::BitcastOp::areCastCompatible(TypeRange inputs, TypeRange outputs) {
2452 if (!areValidCastInputsAndOutputs(inputs, outputs))
2453 return false;
2454
2455 auto srcType = getTypeIfLikeOrMemRef<IntegerType, FloatType>(inputs.front());
2456 auto dstType = getTypeIfLikeOrMemRef<IntegerType, FloatType>(outputs.front());
2457 if (!srcType || !dstType)
2458 return false;
2459
2460 return srcType.getIntOrFloatBitWidth() == dstType.getIntOrFloatBitWidth();
2461}
2462
2463OpFoldResult arith::BitcastOp::fold(FoldAdaptor adaptor) {
2464 auto resType = getType();
2465 auto operand = adaptor.getIn();
2466 if (!operand)
2467 return {};
2468
2469 /// Bitcast dense elements.
2470 if (auto denseAttr = dyn_cast_or_null<DenseElementsAttr>(operand))
2471 return denseAttr.bitcast(llvm::cast<ShapedType>(resType).getElementType());
2472 /// Other shaped types unhandled.
2473 if (llvm::isa<ShapedType>(resType))
2474 return {};
2475
2476 /// Bitcast poison.
2477 if (matchPattern(operand, ub::m_Poison()))
2478 return ub::PoisonAttr::get(getContext());
2479
2480 /// Bitcast integer or float to integer or float.
2481 if (!llvm::isa<FloatAttr, IntegerAttr>(operand))
2482 return {};
2483
2484 APInt bits = llvm::isa<FloatAttr>(operand)
2485 ? llvm::cast<FloatAttr>(operand).getValue().bitcastToAPInt()
2486 : llvm::cast<IntegerAttr>(operand).getValue();
2487 assert(resType.getIntOrFloatBitWidth() == bits.getBitWidth() &&
2488 "trying to fold on broken IR: operands have incompatible types");
2489
2490 if (auto resFloatType = dyn_cast<FloatType>(resType))
2491 return FloatAttr::get(resType,
2492 APFloat(resFloatType.getFloatSemantics(), bits));
2493 return IntegerAttr::get(resType, bits);
2494}
2495
2496void arith::BitcastOp::getCanonicalizationPatterns(RewritePatternSet &patterns,
2497 MLIRContext *context) {
2498 patterns.add<BitcastOfBitcast>(context);
2499}
2500
2501//===----------------------------------------------------------------------===//
2502// CmpIOp
2503//===----------------------------------------------------------------------===//
2504
2505/// Compute `lhs` `pred` `rhs`, where `pred` is one of the known integer
2506/// comparison predicates.
2507bool mlir::arith::applyCmpPredicate(arith::CmpIPredicate predicate,
2508 const APInt &lhs, const APInt &rhs) {
2509 switch (predicate) {
2510 case arith::CmpIPredicate::eq:
2511 return lhs.eq(rhs);
2512 case arith::CmpIPredicate::ne:
2513 return lhs.ne(rhs);
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);
2530 }
2531 llvm_unreachable("unknown cmpi predicate kind");
2532}
2533
2534/// Returns true if the predicate is true for two equal operands.
2535static bool applyCmpPredicateToEqualOperands(arith::CmpIPredicate predicate) {
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:
2542 return true;
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:
2548 return false;
2549 }
2550 llvm_unreachable("unknown cmpi predicate kind");
2551}
2552
2553static std::optional<int64_t> getIntegerWidth(Type t) {
2554 if (auto intType = dyn_cast<IntegerType>(t)) {
2555 return intType.getWidth();
2556 }
2557 if (auto vectorIntType = dyn_cast<VectorType>(t)) {
2558 return llvm::cast<IntegerType>(vectorIntType.getElementType()).getWidth();
2559 }
2560 return std::nullopt;
2561}
2562
2563OpFoldResult arith::CmpIOp::fold(FoldAdaptor adaptor) {
2564 // cmpi(pred, x, x)
2565 if (getLhs() == getRhs()) {
2566 auto val = applyCmpPredicateToEqualOperands(getPredicate());
2567 return getBoolAttribute(getType(), val);
2568 }
2569
2570 if (matchPattern(adaptor.getRhs(), m_Zero())) {
2571 if (auto extOp = getLhs().getDefiningOp<ExtSIOp>()) {
2572 // extsi(%x : i1 -> iN) != 0 -> %x
2573 std::optional<int64_t> integerWidth =
2574 getIntegerWidth(extOp.getOperand().getType());
2575 if (integerWidth && integerWidth.value() == 1 &&
2576 getPredicate() == arith::CmpIPredicate::ne)
2577 return extOp.getOperand();
2578 }
2579 if (auto extOp = getLhs().getDefiningOp<ExtUIOp>()) {
2580 // extui(%x : i1 -> iN) != 0 -> %x
2581 std::optional<int64_t> integerWidth =
2582 getIntegerWidth(extOp.getOperand().getType());
2583 if (integerWidth && integerWidth.value() == 1 &&
2584 getPredicate() == arith::CmpIPredicate::ne)
2585 return extOp.getOperand();
2586 }
2587
2588 // arith.cmpi ne, %val, %zero : i1 -> %val
2589 if (getElementTypeOrSelf(getLhs().getType()).isInteger(1) &&
2590 getPredicate() == arith::CmpIPredicate::ne)
2591 return getLhs();
2592 }
2593
2594 if (matchPattern(adaptor.getRhs(), m_One())) {
2595 // arith.cmpi eq, %val, %one : i1 -> %val
2596 if (getElementTypeOrSelf(getLhs().getType()).isInteger(1) &&
2597 getPredicate() == arith::CmpIPredicate::eq)
2598 return getLhs();
2599 }
2600
2601 // Move constant to the right side.
2602 if (adaptor.getLhs() && !adaptor.getRhs()) {
2603 // Do not use invertPredicate, as it will change eq to ne and vice versa.
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},
2610 };
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);
2619 return getResult();
2620 }
2621 }
2622 llvm_unreachable("unknown cmpi predicate kind");
2623 }
2624
2625 // We are moving constants to the right side; So if lhs is constant rhs is
2626 // guaranteed to be a constant.
2627 if (auto lhs = dyn_cast_if_present<TypedAttr>(adaptor.getLhs())) {
2629 adaptor.getOperands(), getI1SameShape(lhs.getType()),
2630 [pred = getPredicate()](const APInt &lhs, const APInt &rhs) {
2631 return APInt(1,
2632 static_cast<int64_t>(applyCmpPredicate(pred, lhs, rhs)));
2633 });
2634 }
2635
2636 return {};
2637}
2638
2639void arith::CmpIOp::getCanonicalizationPatterns(RewritePatternSet &patterns,
2640 MLIRContext *context) {
2641 patterns.insert<CmpIExtSI, CmpIExtUI>(context);
2642}
2643
2644//===----------------------------------------------------------------------===//
2645// CmpFOp
2646//===----------------------------------------------------------------------===//
2647
2648/// Compute `lhs` `pred` `rhs`, where `pred` is one of the known floating point
2649/// comparison predicates.
2650bool mlir::arith::applyCmpPredicate(arith::CmpFPredicate predicate,
2651 const APFloat &lhs, const APFloat &rhs) {
2652 auto cmpResult = lhs.compare(rhs);
2653 switch (predicate) {
2654 case arith::CmpFPredicate::AlwaysFalse:
2655 return false;
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:
2691 return true;
2692 }
2693 llvm_unreachable("unknown cmpf predicate kind");
2694}
2695
2696OpFoldResult arith::CmpFOp::fold(FoldAdaptor adaptor) {
2697 auto lhs = dyn_cast_if_present<FloatAttr>(adaptor.getLhs());
2698 auto rhs = dyn_cast_if_present<FloatAttr>(adaptor.getRhs());
2699
2700 // If one operand is NaN, making them both NaN does not change the result.
2701 if (lhs && lhs.getValue().isNaN())
2702 rhs = lhs;
2703 if (rhs && rhs.getValue().isNaN())
2704 lhs = rhs;
2705
2706 if (!lhs || !rhs)
2707 return {};
2708
2709 auto val = applyCmpPredicate(getPredicate(), lhs.getValue(), rhs.getValue());
2710 return BoolAttr::get(getContext(), val);
2711}
2712
2713class CmpFIntToFPConst final : public OpRewritePattern<CmpFOp> {
2714public:
2715 using Base::Base;
2716
2717 static CmpIPredicate convertToIntegerPredicate(CmpFPredicate pred,
2718 bool isUnsigned) {
2719 using namespace arith;
2720 switch (pred) {
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;
2739 default:
2740 llvm_unreachable("Unexpected predicate!");
2741 }
2742 }
2743
2744 LogicalResult matchAndRewrite(CmpFOp op,
2745 PatternRewriter &rewriter) const override {
2746 FloatAttr flt;
2747 if (!matchPattern(op.getRhs(), m_Constant(&flt)))
2748 return failure();
2749
2750 const APFloat &rhs = flt.getValue();
2751
2752 // Don't attempt to fold a nan.
2753 if (rhs.isNaN())
2754 return failure();
2755
2756 // Get the width of the mantissa. We don't want to hack on conversions that
2757 // might lose information from the integer, e.g. "i64 -> float"
2758 FloatType floatTy = llvm::cast<FloatType>(op.getRhs().getType());
2759 int mantissaWidth = floatTy.getFPMantissaWidth();
2760 if (mantissaWidth <= 0)
2761 return failure();
2762
2763 bool isUnsigned;
2764 Value intVal;
2765
2766 if (auto si = op.getLhs().getDefiningOp<SIToFPOp>()) {
2767 isUnsigned = false;
2768 intVal = si.getIn();
2769 } else if (auto ui = op.getLhs().getDefiningOp<UIToFPOp>()) {
2770 isUnsigned = true;
2771 intVal = ui.getIn();
2772 } else {
2773 return failure();
2774 }
2775
2776 // Check to see that the input is converted from an integer type that is
2777 // small enough that preserves all bits.
2778 auto intTy = llvm::cast<IntegerType>(intVal.getType());
2779 auto intWidth = intTy.getWidth();
2780
2781 // Number of bits representing values, as opposed to the sign
2782 auto valueBits = isUnsigned ? intWidth : (intWidth - 1);
2783
2784 // Following test does NOT adjust intWidth downwards for signed inputs,
2785 // because the most negative value still requires all the mantissa bits
2786 // to distinguish it from one less than that value.
2787 if ((int)intWidth > mantissaWidth) {
2788 // Conversion would lose accuracy. Check if loss can impact comparison.
2789 int exponent = ilogb(rhs);
2790 if (exponent == APFloat::IEK_Inf) {
2791 int maxExponent = ilogb(APFloat::getLargest(rhs.getSemantics()));
2792 if (maxExponent < (int)valueBits) {
2793 // Conversion could create infinity.
2794 return failure();
2795 }
2796 } else {
2797 // Note that if rhs is zero or NaN, then Exp is negative
2798 // and first condition is trivially false.
2799 if (mantissaWidth <= exponent && exponent <= (int)valueBits) {
2800 // Conversion could affect comparison.
2801 return failure();
2802 }
2803 }
2804 }
2805
2806 // Convert to equivalent cmpi predicate
2807 CmpIPredicate pred;
2808 switch (op.getPredicate()) {
2809 case CmpFPredicate::ORD:
2810 // Int to fp conversion doesn't create a nan (ord checks neither is a nan)
2811 rewriter.replaceOpWithNewOp<ConstantIntOp>(op, /*value=*/true,
2812 /*width=*/1);
2813 return success();
2814 case CmpFPredicate::UNO:
2815 // Int to fp conversion doesn't create a nan (uno checks either is a nan)
2816 rewriter.replaceOpWithNewOp<ConstantIntOp>(op, /*value=*/false,
2817 /*width=*/1);
2818 return success();
2819 default:
2820 pred = convertToIntegerPredicate(op.getPredicate(), isUnsigned);
2821 break;
2822 }
2823
2824 if (!isUnsigned) {
2825 // If the rhs value is > SignedMax, fold the comparison. This handles
2826 // +INF and large values.
2827 APFloat signedMax(rhs.getSemantics());
2828 signedMax.convertFromAPInt(APInt::getSignedMaxValue(intWidth), true,
2829 APFloat::rmNearestTiesToEven);
2830 if (signedMax < rhs) { // smax < 13123.0
2831 if (pred == CmpIPredicate::ne || pred == CmpIPredicate::slt ||
2832 pred == CmpIPredicate::sle)
2833 rewriter.replaceOpWithNewOp<ConstantIntOp>(op, /*value=*/true,
2834 /*width=*/1);
2835 else
2836 rewriter.replaceOpWithNewOp<ConstantIntOp>(op, /*value=*/false,
2837 /*width=*/1);
2838 return success();
2839 }
2840 } else {
2841 // If the rhs value is > UnsignedMax, fold the comparison. This handles
2842 // +INF and large values.
2843 APFloat unsignedMax(rhs.getSemantics());
2844 unsignedMax.convertFromAPInt(APInt::getMaxValue(intWidth), false,
2845 APFloat::rmNearestTiesToEven);
2846 if (unsignedMax < rhs) { // umax < 13123.0
2847 if (pred == CmpIPredicate::ne || pred == CmpIPredicate::ult ||
2848 pred == CmpIPredicate::ule)
2849 rewriter.replaceOpWithNewOp<ConstantIntOp>(op, /*value=*/true,
2850 /*width=*/1);
2851 else
2852 rewriter.replaceOpWithNewOp<ConstantIntOp>(op, /*value=*/false,
2853 /*width=*/1);
2854 return success();
2855 }
2856 }
2857
2858 if (!isUnsigned) {
2859 // See if the rhs value is < SignedMin.
2860 APFloat signedMin(rhs.getSemantics());
2861 signedMin.convertFromAPInt(APInt::getSignedMinValue(intWidth), true,
2862 APFloat::rmNearestTiesToEven);
2863 if (signedMin > rhs) { // smin > 12312.0
2864 if (pred == CmpIPredicate::ne || pred == CmpIPredicate::sgt ||
2865 pred == CmpIPredicate::sge)
2866 rewriter.replaceOpWithNewOp<ConstantIntOp>(op, /*value=*/true,
2867 /*width=*/1);
2868 else
2869 rewriter.replaceOpWithNewOp<ConstantIntOp>(op, /*value=*/false,
2870 /*width=*/1);
2871 return success();
2872 }
2873 } else {
2874 // See if the rhs value is < UnsignedMin.
2875 APFloat unsignedMin(rhs.getSemantics());
2876 unsignedMin.convertFromAPInt(APInt::getMinValue(intWidth), false,
2877 APFloat::rmNearestTiesToEven);
2878 if (unsignedMin > rhs) { // umin > 12312.0
2879 if (pred == CmpIPredicate::ne || pred == CmpIPredicate::ugt ||
2880 pred == CmpIPredicate::uge)
2881 rewriter.replaceOpWithNewOp<ConstantIntOp>(op, /*value=*/true,
2882 /*width=*/1);
2883 else
2884 rewriter.replaceOpWithNewOp<ConstantIntOp>(op, /*value=*/false,
2885 /*width=*/1);
2886 return success();
2887 }
2888 }
2889
2890 // Okay, now we know that the FP constant fits in the range [SMIN, SMAX] or
2891 // [0, UMAX], but it may still be fractional. See if it is fractional by
2892 // casting the FP value to the integer value and back, checking for
2893 // equality. Don't do this for zero, because -0.0 is not fractional.
2894 bool ignored;
2895 APSInt rhsInt(intWidth, isUnsigned);
2896 if (APFloat::opInvalidOp ==
2897 rhs.convertToInteger(rhsInt, APFloat::rmTowardZero, &ignored)) {
2898 // Undefined behavior invoked - the destination type can't represent
2899 // the input constant.
2900 return failure();
2901 }
2902
2903 if (!rhs.isZero()) {
2904 APFloat apf(floatTy.getFloatSemantics(),
2905 APInt::getZero(floatTy.getWidth()));
2906 apf.convertFromAPInt(rhsInt, !isUnsigned, APFloat::rmNearestTiesToEven);
2907
2908 bool equal = apf == rhs;
2909 if (!equal) {
2910 // If we had a comparison against a fractional value, we have to adjust
2911 // the compare predicate and sometimes the value. rhsInt is rounded
2912 // towards zero at this point.
2913 switch (pred) {
2914 case CmpIPredicate::ne: // (float)int != 4.4 --> true
2915 rewriter.replaceOpWithNewOp<ConstantIntOp>(op, /*value=*/true,
2916 /*width=*/1);
2917 return success();
2918 case CmpIPredicate::eq: // (float)int == 4.4 --> false
2919 rewriter.replaceOpWithNewOp<ConstantIntOp>(op, /*value=*/false,
2920 /*width=*/1);
2921 return success();
2922 case CmpIPredicate::ule:
2923 // (float)int <= 4.4 --> int <= 4
2924 // (float)int <= -4.4 --> false
2925 if (rhs.isNegative()) {
2926 rewriter.replaceOpWithNewOp<ConstantIntOp>(op, /*value=*/false,
2927 /*width=*/1);
2928 return success();
2929 }
2930 break;
2931 case CmpIPredicate::sle:
2932 // (float)int <= 4.4 --> int <= 4
2933 // (float)int <= -4.4 --> int < -4
2934 if (rhs.isNegative())
2935 pred = CmpIPredicate::slt;
2936 break;
2937 case CmpIPredicate::ult:
2938 // (float)int < -4.4 --> false
2939 // (float)int < 4.4 --> int <= 4
2940 if (rhs.isNegative()) {
2941 rewriter.replaceOpWithNewOp<ConstantIntOp>(op, /*value=*/false,
2942 /*width=*/1);
2943 return success();
2944 }
2945 pred = CmpIPredicate::ule;
2946 break;
2947 case CmpIPredicate::slt:
2948 // (float)int < -4.4 --> int < -4
2949 // (float)int < 4.4 --> int <= 4
2950 if (!rhs.isNegative())
2951 pred = CmpIPredicate::sle;
2952 break;
2953 case CmpIPredicate::ugt:
2954 // (float)int > 4.4 --> int > 4
2955 // (float)int > -4.4 --> true
2956 if (rhs.isNegative()) {
2957 rewriter.replaceOpWithNewOp<ConstantIntOp>(op, /*value=*/true,
2958 /*width=*/1);
2959 return success();
2960 }
2961 break;
2962 case CmpIPredicate::sgt:
2963 // (float)int > 4.4 --> int > 4
2964 // (float)int > -4.4 --> int >= -4
2965 if (rhs.isNegative())
2966 pred = CmpIPredicate::sge;
2967 break;
2968 case CmpIPredicate::uge:
2969 // (float)int >= -4.4 --> true
2970 // (float)int >= 4.4 --> int > 4
2971 if (rhs.isNegative()) {
2972 rewriter.replaceOpWithNewOp<ConstantIntOp>(op, /*value=*/true,
2973 /*width=*/1);
2974 return success();
2975 }
2976 pred = CmpIPredicate::ugt;
2977 break;
2978 case CmpIPredicate::sge:
2979 // (float)int >= -4.4 --> int >= -4
2980 // (float)int >= 4.4 --> int > 4
2981 if (!rhs.isNegative())
2982 pred = CmpIPredicate::sgt;
2983 break;
2984 }
2985 }
2986 }
2987
2988 // Lower this FP comparison into an appropriate integer version of the
2989 // comparison.
2990 rewriter.replaceOpWithNewOp<CmpIOp>(
2991 op, pred, intVal,
2992 ConstantOp::create(rewriter, op.getLoc(), intVal.getType(),
2993 rewriter.getIntegerAttr(intVal.getType(), rhsInt)));
2994 return success();
2995 }
2996};
2997
2998void arith::CmpFOp::getCanonicalizationPatterns(RewritePatternSet &patterns,
2999 MLIRContext *context) {
3000 patterns.insert<CmpFIntToFPConst>(context);
3001}
3002
3003//===----------------------------------------------------------------------===//
3004// SelectOp
3005//===----------------------------------------------------------------------===//
3006
3007// select %arg, %c1, %c0 => extui %arg
3008struct SelectToExtUI : public OpRewritePattern<arith::SelectOp> {
3009 using Base::Base;
3010
3011 LogicalResult matchAndRewrite(arith::SelectOp op,
3012 PatternRewriter &rewriter) const override {
3013 // Cannot extui i1 to i1, or i1 to f32
3014 if (!llvm::isa<IntegerType>(op.getType()) || op.getType().isInteger(1))
3015 return failure();
3016
3017 // select %x, c1, %c0 => extui %arg
3018 if (matchPattern(op.getTrueValue(), m_One()) &&
3019 matchPattern(op.getFalseValue(), m_Zero())) {
3020 rewriter.replaceOpWithNewOp<arith::ExtUIOp>(op, op.getType(),
3021 op.getCondition());
3022 return success();
3023 }
3024
3025 // select %x, c0, %c1 => extui (xor %arg, true)
3026 if (matchPattern(op.getTrueValue(), m_Zero()) &&
3027 matchPattern(op.getFalseValue(), m_One())) {
3028 rewriter.replaceOpWithNewOp<arith::ExtUIOp>(
3029 op, op.getType(),
3030 arith::XOrIOp::create(
3031 rewriter, op.getLoc(), op.getCondition(),
3032 arith::ConstantIntOp::create(rewriter, op.getLoc(),
3033 op.getCondition().getType(), 1)));
3034 return success();
3035 }
3036
3037 return failure();
3038 }
3039};
3040
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>(
3050 context);
3051}
3052
3053OpFoldResult arith::SelectOp::fold(FoldAdaptor adaptor) {
3054 Value trueVal = getTrueValue();
3055 Value falseVal = getFalseValue();
3056 if (trueVal == falseVal)
3057 return trueVal;
3058
3059 Value condition = getCondition();
3060
3061 // select true, %0, %1 => %0
3062 if (matchPattern(adaptor.getCondition(), m_One()))
3063 return trueVal;
3064
3065 // select false, %0, %1 => %1
3066 if (matchPattern(adaptor.getCondition(), m_Zero()))
3067 return falseVal;
3068
3069 // If either operand is fully poisoned, return the other.
3070 if (matchPattern(adaptor.getTrueValue(), ub::m_Poison()))
3071 return falseVal;
3072
3073 if (matchPattern(adaptor.getFalseValue(), ub::m_Poison()))
3074 return trueVal;
3075
3076 // select %x, true, false => %x
3077 if (getType().isSignlessInteger(1) &&
3078 matchPattern(adaptor.getTrueValue(), m_One()) &&
3079 matchPattern(adaptor.getFalseValue(), m_Zero()))
3080 return condition;
3081
3082 if (auto cmp = condition.getDefiningOp<arith::CmpIOp>()) {
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();
3087
3088 // %0 = arith.cmpi eq, %arg0, %arg1
3089 // %1 = arith.select %0, %arg0, %arg1 => %arg1
3090
3091 // %0 = arith.cmpi ne, %arg0, %arg1
3092 // %1 = arith.select %0, %arg0, %arg1 => %arg0
3093
3094 if ((cmpLhs == trueVal && cmpRhs == falseVal) ||
3095 (cmpRhs == trueVal && cmpLhs == falseVal))
3096 return pred == arith::CmpIPredicate::ne ? trueVal : falseVal;
3097 }
3098 }
3099
3100 // Constant-fold constant operands over non-splat constant condition.
3101 // select %cst_vec, %cst0, %cst1 => %cst2
3102 if (auto cond =
3103 dyn_cast_if_present<DenseElementsAttr>(adaptor.getCondition())) {
3104 // DenseElementsAttr by construction always has a static shape.
3105 assert(cond.getType().hasStaticShape() &&
3106 "DenseElementsAttr must have static shape");
3107 if (auto lhs =
3108 dyn_cast_if_present<DenseElementsAttr>(adaptor.getTrueValue())) {
3109 if (auto rhs =
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>());
3119
3120 for (auto [condVal, lhsVal, rhsVal] :
3121 llvm::zip_equal(condVals, lhsVals, rhsVals))
3122 results.push_back(condVal.getValue() ? lhsVal : rhsVal);
3123
3124 return DenseElementsAttr::get(lhs.getType(), results);
3125 }
3126 }
3127 }
3128
3129 return nullptr;
3130}
3131
3132ParseResult SelectOp::parse(OpAsmParser &parser, OperationState &result) {
3133 Type conditionType, resultType;
3134 SmallVector<OpAsmParser::UnresolvedOperand, 3> operands;
3135 if (parser.parseOperandList(operands, /*requiredOperandCount=*/3) ||
3136 parser.parseOptionalAttrDict(result.attributes) ||
3137 parser.parseColonType(resultType))
3138 return failure();
3139
3140 // Check for the explicit condition type if this is a masked tensor or vector.
3141 if (succeeded(parser.parseOptionalComma())) {
3142 conditionType = resultType;
3143 if (parser.parseType(resultType))
3144 return failure();
3145 } else {
3146 conditionType = parser.getBuilder().getI1Type();
3147 }
3148
3149 result.addTypes(resultType);
3150 return parser.resolveOperands(operands,
3151 {conditionType, resultType, resultType},
3152 parser.getNameLoc(), result.operands);
3153}
3154
3155void arith::SelectOp::print(OpAsmPrinter &p) {
3156 p << " " << getOperands();
3157 p.printOptionalAttrDict((*this)->getDiscardableAttrDictionary().getValue());
3158 p << " : ";
3159 if (ShapedType condType = dyn_cast<ShapedType>(getCondition().getType()))
3160 p << condType << ", ";
3161 p << getType();
3162}
3163
3164LogicalResult arith::SelectOp::verify() {
3165 Type conditionType = getCondition().getType();
3166 if (conditionType.isSignlessInteger(1))
3167 return success();
3168
3169 // If the result type is a vector or tensor, the type can be a mask with the
3170 // same elements.
3171 Type resultType = getType();
3172 if (!llvm::isa<TensorType, VectorType>(resultType))
3173 return emitOpError() << "expected condition to be a signless i1, but got "
3174 << conditionType;
3175 Type shapedConditionType = getI1SameShape(resultType);
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 "
3180 << conditionType;
3181 }
3182 return success();
3183}
3184//===----------------------------------------------------------------------===//
3185// ShLIOp
3186//===----------------------------------------------------------------------===//
3187
3188OpFoldResult arith::ShLIOp::fold(FoldAdaptor adaptor) {
3189 // TODO: shli(x, c) -> poison when c is out of range (c >= bit width). An
3190 // out-of-range shift amount is undefined behaviour and could fold to poison,
3191 // but that would make the arith dialect depend on the ub dialect to
3192 // materialize ub.poison; left out for now.
3193
3194 // shli(x, 0) -> x
3195 if (matchPattern(adaptor.getRhs(), m_Zero()))
3196 return getLhs();
3197 // shli(0, x) -> 0. An out-of-range shift amount yields poison, so refining
3198 // it to 0 is valid.
3199 if (matchPattern(adaptor.getLhs(), m_Zero()))
3200 return getLhs();
3201 // Don't fold if shifting more or equal than the bit width.
3202 bool bounded = false;
3204 adaptor.getOperands(), [&](const APInt &a, const APInt &b) {
3205 bounded = b.ult(b.getBitWidth());
3206 return a.shl(b);
3207 });
3208 return bounded ? result : Attribute();
3209}
3210
3211//===----------------------------------------------------------------------===//
3212// ShRUIOp
3213//===----------------------------------------------------------------------===//
3214
3215OpFoldResult arith::ShRUIOp::fold(FoldAdaptor adaptor) {
3216 // TODO: shrui(x, c) -> poison when c is out of range (c >= bit width). An
3217 // out-of-range shift amount is undefined behaviour and could fold to poison,
3218 // but that would make the arith dialect depend on the ub dialect to
3219 // materialize ub.poison; left out for now.
3220
3221 // shrui(x, 0) -> x
3222 if (matchPattern(adaptor.getRhs(), m_Zero()))
3223 return getLhs();
3224 // shrui(0, x) -> 0. An out-of-range shift amount yields poison, so refining
3225 // it to 0 is valid.
3226 if (matchPattern(adaptor.getLhs(), m_Zero()))
3227 return getLhs();
3228 // shrui(x, x) -> 0. For any in-range shift amount v < bitwidth, v >> v == 0
3229 // (v < 2^v); out-of-range amounts yield poison, so 0 is a valid refinement.
3230 if (getLhs() == getRhs())
3231 return getIntegerAttrOfType(getType(), 0);
3232 // Don't fold if shifting more or equal than the bit width.
3233 bool bounded = false;
3235 adaptor.getOperands(), [&](const APInt &a, const APInt &b) {
3236 bounded = b.ult(b.getBitWidth());
3237 return a.lshr(b);
3238 });
3239 return bounded ? result : Attribute();
3240}
3241
3242//===----------------------------------------------------------------------===//
3243// ShRSIOp
3244//===----------------------------------------------------------------------===//
3245
3246OpFoldResult arith::ShRSIOp::fold(FoldAdaptor adaptor) {
3247 // TODO: shrsi(x, c) -> poison when c is out of range (c >= bit width). An
3248 // out-of-range shift amount is undefined behaviour and could fold to poison,
3249 // but that would make the arith dialect depend on the ub dialect to
3250 // materialize ub.poison; left out for now.
3251
3252 // shrsi(x, 0) -> x
3253 if (matchPattern(adaptor.getRhs(), m_Zero()))
3254 return getLhs();
3255 // shrsi(0, x) -> 0. An out-of-range shift amount yields poison, so refining
3256 // it to 0 is valid.
3257 if (matchPattern(adaptor.getLhs(), m_Zero()))
3258 return getLhs();
3259 // shrsi(x, x) -> 0. For any in-range shift amount v < bitwidth, v is a small
3260 // non-negative value and v >> v == 0; out-of-range amounts yield poison.
3261 if (getLhs() == getRhs())
3262 return getIntegerAttrOfType(getType(), 0);
3263 // shrsi(-1, x) -> -1. Arithmetic shift of all-ones is all-ones for any
3264 // in-range amount; out-of-range amounts yield poison.
3265 if (APInt val;
3266 matchPattern(adaptor.getLhs(), m_ConstantInt(&val)) && val.isAllOnes())
3267 return getLhs();
3268 // Don't fold if shifting more or equal than the bit width.
3269 bool bounded = false;
3271 adaptor.getOperands(), [&](const APInt &a, const APInt &b) {
3272 bounded = b.ult(b.getBitWidth());
3273 return a.ashr(b);
3274 });
3275 return bounded ? result : Attribute();
3276}
3277
3278//===----------------------------------------------------------------------===//
3279// Atomic Enum
3280//===----------------------------------------------------------------------===//
3281
3282/// Returns the identity value attribute associated with an AtomicRMWKind op.
3283TypedAttr mlir::arith::getIdentityValueAttr(AtomicRMWKind kind, Type resultType,
3284 OpBuilder &builder, Location loc,
3285 bool useOnlyFiniteValue) {
3286 switch (kind) {
3287 case AtomicRMWKind::maximumf: {
3288 const llvm::fltSemantics &semantic =
3289 llvm::cast<FloatType>(resultType).getFloatSemantics();
3290 APFloat identity = useOnlyFiniteValue
3291 ? APFloat::getLargest(semantic, /*Negative=*/true)
3292 : APFloat::getInf(semantic, /*Negative=*/true);
3293 return builder.getFloatAttr(resultType, identity);
3294 }
3295 case AtomicRMWKind::maxnumf: {
3296 const llvm::fltSemantics &semantic =
3297 llvm::cast<FloatType>(resultType).getFloatSemantics();
3298 APFloat identity = APFloat::getNaN(semantic, /*Negative=*/true);
3299 return builder.getFloatAttr(resultType, identity);
3300 }
3301 case AtomicRMWKind::addf:
3302 case AtomicRMWKind::addi:
3303 case AtomicRMWKind::maxu:
3304 case AtomicRMWKind::ori:
3305 case AtomicRMWKind::xori:
3306 return builder.getZeroAttr(resultType);
3307 case AtomicRMWKind::andi:
3308 return builder.getIntegerAttr(
3309 resultType,
3310 APInt::getAllOnes(llvm::cast<IntegerType>(resultType).getWidth()));
3311 case AtomicRMWKind::maxs:
3312 return builder.getIntegerAttr(
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, /*Negative=*/false)
3320 : APFloat::getInf(semantic, /*Negative=*/false);
3321
3322 return builder.getFloatAttr(resultType, identity);
3323 }
3324 case AtomicRMWKind::minnumf: {
3325 const llvm::fltSemantics &semantic =
3326 llvm::cast<FloatType>(resultType).getFloatSemantics();
3327 APFloat identity = APFloat::getNaN(semantic, /*Negative=*/false);
3328 return builder.getFloatAttr(resultType, identity);
3329 }
3330 case AtomicRMWKind::mins:
3331 return builder.getIntegerAttr(
3332 resultType, APInt::getSignedMaxValue(
3333 llvm::cast<IntegerType>(resultType).getWidth()));
3334 case AtomicRMWKind::minu:
3335 return builder.getIntegerAttr(
3336 resultType,
3337 APInt::getMaxValue(llvm::cast<IntegerType>(resultType).getWidth()));
3338 case AtomicRMWKind::muli:
3339 return builder.getIntegerAttr(resultType, 1);
3340 case AtomicRMWKind::mulf:
3341 return builder.getFloatAttr(resultType, 1);
3342 // `assign` is not a reduction and has no identity element.
3343 case AtomicRMWKind::assign:
3344 break;
3345 }
3346 (void)emitOptionalError(loc, "Reduction operation type not supported");
3347 return nullptr;
3348}
3349
3350/// Returns the identity numeric value of the given op.
3351std::optional<TypedAttr> mlir::arith::getNeutralElement(Operation *op) {
3352 std::optional<AtomicRMWKind> maybeKind =
3354 // Floating-point operations.
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; })
3361 // Integer operations.
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);
3372 if (!maybeKind) {
3373 return std::nullopt;
3374 }
3375
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);
3382 }
3383
3384 // Builder only used as helper for attribute creation.
3385 OpBuilder b(op->getContext());
3386 Type resultType = op->getResult(0).getType();
3387
3388 return getIdentityValueAttr(*maybeKind, resultType, b, op->getLoc(),
3389 useOnlyFiniteValue);
3390}
3391
3392/// Returns the identity value associated with an AtomicRMWKind op.
3393Value mlir::arith::getIdentityValue(AtomicRMWKind op, Type resultType,
3394 OpBuilder &builder, Location loc,
3395 bool useOnlyFiniteValue) {
3396 if (auto attr = getIdentityValueAttr(op, resultType, builder, loc,
3397 useOnlyFiniteValue))
3398 return arith::ConstantOp::create(builder, loc, attr);
3399 return {};
3400}
3401
3402/// Return the value obtained by applying the reduction operation kind
3403/// associated with a binary AtomicRMWKind op to `lhs` and `rhs`.
3405 Location loc, Value lhs, Value rhs) {
3406 switch (op) {
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);
3437 // `assign` is not a reduction and has no corresponding binary operation.
3438 case AtomicRMWKind::assign:
3439 break;
3440 }
3441 (void)emitOptionalError(loc, "Reduction operation type not supported");
3442 return nullptr;
3443}
3444
3445//===----------------------------------------------------------------------===//
3446// TableGen'd op method definitions
3447//===----------------------------------------------------------------------===//
3448
3449#define GET_OP_CLASSES
3450#include "mlir/Dialect/Arith/IR/ArithOps.cpp.inc"
3451
3452//===----------------------------------------------------------------------===//
3453// TableGen'd enum attribute definitions
3454//===----------------------------------------------------------------------===//
3455
3456#include "mlir/Dialect/Arith/IR/ArithOpsEnums.cpp.inc"
return success()
if(failed(verifyVectorMemoryOp(getOperation(), memrefType, getVectorType()))) return failure()
static Speculation::Speculatability getDivUISpeculatability(Value divisor)
Returns whether an unsigned division by divisor is speculatable.
Definition ArithOps.cpp:848
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)
Definition ArithOps.cpp:66
static IntegerOverflowFlagsAttr mergeOverflowFlags(IntegerOverflowFlagsAttr val1, IntegerOverflowFlagsAttr val2)
Definition ArithOps.cpp:88
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.
Definition ArithOps.cpp:39
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
Definition ArithOps.cpp:797
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)).
Definition ArithOps.cpp:128
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.
Definition ArithOps.cpp:902
static IntegerAttr orIntegerAttrs(PatternRewriter &builder, Value res, Attribute lhs, Attribute rhs)
Definition ArithOps.cpp:76
static IntegerAttr addIntegerAttrs(PatternRewriter &builder, Value res, Attribute lhs, Attribute rhs)
Definition ArithOps.cpp:56
static Attribute getBoolAttribute(Type type, bool value)
Definition ArithOps.cpp:171
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)
Definition ArithOps.cpp:61
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...
Definition ArithOps.cpp:185
static IntegerAttr andIntegerAttrs(PatternRewriter &builder, Value res, Attribute lhs, Attribute rhs)
Definition ArithOps.cpp:71
static int64_t getScalarOrElementWidth(Type type)
Definition ArithOps.cpp:151
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.
Definition ArithOps.cpp:208
static bool areValidCastInputsAndOutputs(TypeRange inputs, TypeRange outputs)
static IntegerAttr xorIntegerAttrs(PatternRewriter &builder, Value res, Attribute lhs, Attribute rhs)
Definition ArithOps.cpp:81
static APInt calculateUnsignedBorrow(const APInt &lhs, const APInt &rhs)
Definition ArithOps.cpp:536
std::tuple< Types... > * type_list
static IntegerAttr applyToIntegerAttrs(PatternRewriter &builder, Value res, Attribute lhs, Attribute rhs, function_ref< APInt(const APInt &, const APInt &)> binFn)
Definition ArithOps.cpp:47
static APInt calculateUnsignedOverflow(const APInt &sum, const APInt &operand)
Definition ArithOps.cpp:472
static FailureOr< APInt > getIntOrSplatIntValue(Attribute attr)
Definition ArithOps.cpp:163
static unsigned getIndexCastWidth(Type t)
Return the bit-width of t for the purpose of index_cast width checks.
static LogicalResult verifyTruncateOp(Op op)
lhs
b
Return true if permutation is a valid permutation of the outer_dims_perm (case OuterOrInnerPerm::Oute...
b getContext())
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,...
Definition SPIRVOps.cpp:229
#define mul(a, b)
#define add(a, b)
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.
Definition Attributes.h:25
static BoolAttr get(MLIRContext *context, bool value)
IntegerAttr getIndexAttr(int64_t value)
Definition Builders.cpp:116
IntegerAttr getIntegerAttr(Type type, int64_t value)
Definition Builders.cpp:237
FloatAttr getFloatAttr(Type type, double value)
Definition Builders.cpp:263
IntegerType getIntegerType(unsigned width)
Definition Builders.cpp:75
Ty getType(Args &&...args)
Get or construct an instance of the type Ty with provided arguments.
Definition Builders.h:94
TypedAttr getZeroAttr(Type type)
Definition Builders.cpp:333
IntegerType getI1Type()
Definition Builders.cpp:61
IndexType getIndexType()
Definition Builders.cpp:59
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...
Definition Builders.h:632
Location getLoc() const
Accessors for the implied location.
Definition Builders.h:665
This class defines the main interface for locations in MLIR and acts as a non-nullable wrapper around...
Definition Location.h:76
MLIRContext is the top-level object for a collection of MLIR operations.
Definition MLIRContext.h:63
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.
Definition Builders.h:210
Operation * create(const OperationState &state)
Creates an operation given the fields represented as an OperationState.
Definition Builders.cpp:466
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.
Definition Operation.h:87
OpResult getResult(unsigned idx)
Get the 'idx'th result of this operation.
Definition Operation.h:432
Location getLoc()
The source location the operation was defined or derived from.
Definition Operation.h:240
MLIRContext * getContext()
Return the context this operation is associated with.
Definition Operation.h:233
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.
Definition TypeRange.h:40
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 isSignlessInteger() const
Return true if this is a signless integer type (with the specified width).
Definition Types.cpp:66
bool isIndex() const
Definition Types.cpp:56
bool isIntOrFloat() const
Return true if this is an integer (of any signedness) or a float type.
Definition Types.cpp:118
unsigned getIntOrFloatBitWidth() const
Return the bit width of an integer or a float type, assert failure on other types.
Definition Types.cpp:124
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
Definition Value.h:96
Type getType() const
Return the type of this value.
Definition Value.h:105
Operation * getDefiningOp() const
If this value is the result of an operation, return the operation that defines it.
Definition Value.cpp:18
Specialization of arith.constant op that returns a floating point value.
Definition Arith.h:72
static ConstantFloatOp create(OpBuilder &builder, Location location, FloatType type, const APFloat &value)
Definition ArithOps.cpp:369
static bool classof(Operation *op)
Definition ArithOps.cpp:386
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.
Definition ArithOps.cpp:363
Specialization of arith.constant op that returns an integer of index type.
Definition Arith.h:93
static void build(OpBuilder &builder, OperationState &result, int64_t value)
Build a constant int op that produces an index.
Definition ArithOps.cpp:392
static bool classof(Operation *op)
Definition ArithOps.cpp:413
static ConstantIndexOp create(OpBuilder &builder, Location location, int64_t value)
Definition ArithOps.cpp:398
Specialization of arith.constant op that returns an integer value.
Definition Arith.h:34
static ConstantIntOp create(OpBuilder &builder, Location location, int64_t value, unsigned width)
Definition ArithOps.cpp:297
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.
Definition ArithOps.cpp:290
static bool classof(Operation *op)
Definition ArithOps.cpp:357
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.
Definition ArithOps.cpp:95
Value getZeroConstant(OpBuilder &builder, Location loc, Type type)
Creates an arith.constant operation with a zero value of type type.
Definition ArithOps.cpp:419
auto m_Val(Value v)
Definition Matchers.h:539
detail::InFlightRemark failed(Location loc, RemarkOpts opts)
Report an optimization remark that failed.
Definition Remarks.h:734
detail::poison_attr_matcher m_Poison()
Matches a poison constant (any attribute implementing PoisonAttrInterface).
Definition UBMatchers.h:46
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
detail::constant_float_predicate_matcher m_NaNFloat()
Matches a constant scalar / vector splat / tensor splat float ones.
Definition Matchers.h:421
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.
Definition Utils.cpp:311
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...
Definition Matchers.h:471
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.
Definition Matchers.h:404
detail::constant_int_predicate_matcher m_Zero()
Matches a constant scalar / vector splat / tensor splat integer zero.
Definition Matchers.h:442
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.
Definition Matchers.h:399
detail::constant_int_predicate_matcher m_One()
Matches a constant scalar / vector splat / tensor splat integer one.
Definition Matchers.h:478
detail::constant_float_predicate_matcher m_NegInfFloat()
Matches a constant scalar / vector splat / tensor splat float negative infinity.
Definition Matchers.h:435
detail::constant_float_predicate_matcher m_NegZeroFloat()
Matches a constant scalar / vector splat / tensor splat float negative zero.
Definition Matchers.h:409
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...
Definition Matchers.h:462
detail::op_matcher< OpClass > m_Op()
Matches the given OpClass.
Definition Matchers.h:484
detail::constant_op_matcher m_Constant()
Matches a constant foldable operation.
Definition Matchers.h:369
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.
Definition Matchers.h:427
llvm::function_ref< Fn > function_ref
Definition LLVM.h:147
detail::constant_float_predicate_matcher m_OneFloat()
Matches a constant scalar / vector splat / tensor splat float ones.
Definition Matchers.h:414
detail::constant_int_range_predicate_matcher m_IntRangeWithoutZeroU()
Matches a constant scalar / vector splat / tensor splat integer or a unsigned integer range that does...
Definition Matchers.h:455
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.