19#define GEN_PASS_DEF_ARITHEXPANDOPSPASS
20#include "mlir/Dialect/Arith/Transforms/Passes.h.inc"
30 if (
auto shapedTy = dyn_cast<ShapedType>(type)) {
31 return arith::ConstantOp::create(rewriter, loc,
34 return arith::ConstantOp::create(rewriter, loc, attr);
41 if (
auto shapedTy = dyn_cast<ShapedType>(type)) {
42 return arith::ConstantOp::create(rewriter, loc,
45 return arith::ConstantOp::create(rewriter, loc, attr);
52 if (
auto shapedTy = dyn_cast<ShapedType>(type)) {
53 return arith::ConstantOp::create(rewriter, loc,
57 return arith::ConstantOp::create(rewriter, loc, attr);
62 if (
auto shapedTy = dyn_cast<ShapedType>(cloneFrom)) {
63 return shapedTy.clone(cloneTo);
74 LogicalResult matchAndRewrite(arith::CeilDivUIOp op,
75 PatternRewriter &rewriter)
const final {
76 Location loc = op.getLoc();
77 Value a = op.getLhs();
78 Value
b = op.getRhs();
81 arith::CmpIOp::create(rewriter, loc, arith::CmpIPredicate::eq, a, zero);
83 Value minusOne = arith::SubIOp::create(rewriter, loc, a, one);
84 Value quotient = arith::DivUIOp::create(rewriter, loc, minusOne,
b);
85 Value plusOne = arith::AddIOp::create(rewriter, loc, quotient, one);
86 rewriter.replaceOpWithNewOp<arith::SelectOp>(op,
compare, zero, plusOne);
100 LogicalResult matchAndRewrite(arith::CeilDivSIOp op,
101 PatternRewriter &rewriter)
const final {
102 Location loc = op.getLoc();
103 Type type = op.getType();
104 Value a = op.getLhs();
105 Value
b = op.getRhs();
110 Value quotient = arith::DivSIOp::create(rewriter, loc, a,
b);
111 Value
product = arith::MulIOp::create(rewriter, loc, quotient,
b);
112 Value notEqualDivisor = arith::CmpIOp::create(
113 rewriter, loc, arith::CmpIPredicate::ne, a,
product);
115 Value aNeg = arith::CmpIOp::create(rewriter, loc, arith::CmpIPredicate::slt,
117 Value bNeg = arith::CmpIOp::create(rewriter, loc, arith::CmpIPredicate::slt,
120 Value signEqual = arith::CmpIOp::create(
121 rewriter, loc, arith::CmpIPredicate::eq, aNeg, bNeg);
123 arith::AndIOp::create(rewriter, loc, notEqualDivisor, signEqual);
125 Value quotientPlusOne = arith::AddIOp::create(rewriter, loc, quotient, one);
127 rewriter.replaceOpWithNewOp<arith::SelectOp>(op, cond, quotientPlusOne,
140struct FloorDivSIOpConverter :
public OpRewritePattern<arith::FloorDivSIOp> {
142 LogicalResult matchAndRewrite(arith::FloorDivSIOp op,
143 PatternRewriter &rewriter)
const final {
144 Location loc = op.getLoc();
145 Type type = op.getType();
146 Value a = op.getLhs();
147 Value
b = op.getRhs();
149 Value quotient = arith::DivSIOp::create(rewriter, loc, a,
b);
150 Value
product = arith::MulIOp::create(rewriter, loc, quotient,
b);
151 Value notEqualDivisor = arith::CmpIOp::create(
152 rewriter, loc, arith::CmpIPredicate::ne, a,
product);
155 Value aNeg = arith::CmpIOp::create(rewriter, loc, arith::CmpIPredicate::slt,
157 Value bNeg = arith::CmpIOp::create(rewriter, loc, arith::CmpIPredicate::slt,
160 Value signOpposite = arith::CmpIOp::create(
161 rewriter, loc, arith::CmpIPredicate::ne, aNeg, bNeg);
163 arith::AndIOp::create(rewriter, loc, notEqualDivisor, signOpposite);
165 Value minusOne =
createConst(loc, type, -1, rewriter);
166 Value quotientMinusOne =
167 arith::AddIOp::create(rewriter, loc, quotient, minusOne);
169 rewriter.replaceOpWithNewOp<arith::SelectOp>(op, cond, quotientMinusOne,
175template <
typename OpTy, arith::CmpIPredicate pred>
178 using OpRewritePattern<OpTy>::OpRewritePattern;
180 LogicalResult matchAndRewrite(OpTy op,
181 PatternRewriter &rewriter)
const final {
182 Value
lhs = op.getLhs();
183 Value
rhs = op.getRhs();
185 Value cmp = arith::CmpIOp::create(rewriter, op.getLoc(), pred,
lhs,
rhs);
186 rewriter.replaceOpWithNewOp<arith::SelectOp>(op, cmp,
lhs,
rhs);
191template <
typename OpTy, arith::CmpFPredicate pred>
194 using OpRewritePattern<OpTy>::OpRewritePattern;
196 LogicalResult matchAndRewrite(OpTy op,
197 PatternRewriter &rewriter)
const final {
198 Value
lhs = op.getLhs();
199 Value
rhs = op.getRhs();
201 Location loc = op.getLoc();
203 static_assert(pred == arith::CmpFPredicate::UGT ||
204 pred == arith::CmpFPredicate::ULT,
205 "pred must be either UGT or ULT");
206 Value cmp = arith::CmpFOp::create(rewriter, loc, pred,
lhs,
rhs);
207 Value select = arith::SelectOp::create(rewriter, loc, cmp,
lhs,
rhs);
210 Value isNaN = arith::CmpFOp::create(rewriter, loc,
211 arith::CmpFPredicate::UNO,
rhs,
rhs);
212 rewriter.replaceOpWithNewOp<arith::SelectOp>(op, isNaN,
rhs, select);
217template <
typename OpTy, arith::CmpFPredicate pred>
220 using OpRewritePattern<OpTy>::OpRewritePattern;
222 LogicalResult matchAndRewrite(OpTy op,
223 PatternRewriter &rewriter)
const final {
224 Value
lhs = op.getLhs();
225 Value
rhs = op.getRhs();
227 Location loc = op.getLoc();
229 static_assert(pred == arith::CmpFPredicate::UGT ||
230 pred == arith::CmpFPredicate::ULT,
231 "pred must be either UGT or ULT");
232 Value cmp = arith::CmpFOp::create(rewriter, loc, pred,
lhs,
rhs);
233 Value select = arith::SelectOp::create(rewriter, loc, cmp,
lhs,
rhs);
236 Value isNaN = arith::CmpFOp::create(rewriter, loc,
237 arith::CmpFPredicate::UNO,
lhs,
lhs);
238 rewriter.replaceOpWithNewOp<arith::SelectOp>(op, isNaN,
rhs, select);
245 LogicalResult matchAndRewrite(arith::ExtFOp op,
246 PatternRewriter &rewriter)
const final {
247 ImplicitLocOpBuilder
b(op.getLoc(), rewriter);
248 auto operand = op.getOperand();
249 Type operandTy = operand.getType();
250 Type resultTy = op.getType();
255 return rewriter.notifyMatchFailure(op,
"not a ext of bf16 to f32.");
261 Value bitcast = arith::BitcastOp::create(
b, i16Ty, operand);
262 Value exti = arith::ExtUIOp::create(
b, i32Ty, bitcast);
264 Value c16 =
createConst(op.getLoc(), i32Ty, 16, rewriter);
265 Value shl = arith::ShLIOp::create(
b, exti, c16);
266 Value
result = arith::BitcastOp::create(
b, resultTy, shl);
268 rewriter.replaceOp(op,
result);
273struct BFloat16TruncFOpConverter :
public OpRewritePattern<arith::TruncFOp> {
275 LogicalResult matchAndRewrite(arith::TruncFOp op,
276 PatternRewriter &rewriter)
const final {
277 ImplicitLocOpBuilder
b(op.getLoc(), rewriter);
278 auto operand = op.getOperand();
279 Type operandTy = operand.getType();
280 Type resultTy = op.getType();
285 return rewriter.notifyMatchFailure(op,
"not a trunc of f32 to bf16.");
288 if (op.getRoundingmodeAttr()) {
289 return rewriter.notifyMatchFailure(
290 op,
"only applicable to default rounding mode.");
310 arith::CmpFOp::create(
b, arith::CmpFPredicate::UNE, operand, operand);
312 Value c7FFF =
createConst(op.getLoc(), i32Ty, 0x7fff, rewriter);
314 Value c7FC0I16 =
createConst(op.getLoc(), i16Ty, 0x7fc0, rewriter);
316 Value c16 =
createConst(op.getLoc(), i32Ty, 16, rewriter);
317 Value c1 =
createConst(op.getLoc(), i32Ty, 1, rewriter);
319 Value bitcast = arith::BitcastOp::create(
b, i32Ty, operand);
322 arith::AndIOp::create(
b, arith::ShRUIOp::create(
b, bitcast, c16), c1);
325 Value roundingBias = arith::AddIOp::create(
b, bit16, c7FFF);
332 Value biased = arith::AddIOp::create(
b, bitcast, roundingBias);
335 Value biasedAndShifted = arith::ShRUIOp::create(
b, biased, c16);
336 Value normalCaseResultI16 =
337 arith::TruncIOp::create(
b, i16Ty, biasedAndShifted);
341 arith::SelectOp::create(
b, isNan, c7FC0I16, normalCaseResultI16);
342 Value
result = arith::BitcastOp::create(
b, resultTy, select);
343 rewriter.replaceOp(op,
result);
379 LogicalResult matchAndRewrite(arith::ExtFOp op,
380 PatternRewriter &rewriter)
const final {
381 Location loc = op.getLoc();
382 ImplicitLocOpBuilder
b(loc, rewriter);
383 Value operand = op.getOperand();
384 Type operandTy = operand.
getType();
385 Type resultTy = op.getType();
389 if (!isa<Float4E2M1FNType>(operandETy))
390 return rewriter.notifyMatchFailure(op,
"not a ext of F4E2M1FN");
395 Value i4Bits = arith::BitcastOp::create(
b, i4Ty, operand);
397 Value c0x0 =
createConst(loc, i4Ty, 0x0, rewriter);
398 Value c0x1 =
createConst(loc, i4Ty, 0x1, rewriter);
399 Value c0x2 =
createConst(loc, i4Ty, 0x2, rewriter);
400 Value c0x4 =
createConst(loc, i4Ty, 0x4, rewriter);
401 Value c0x7 =
createConst(loc, i4Ty, 0x7, rewriter);
403 Value i4BitsNoSign = arith::AndIOp::create(
b, i4Bits, c0x7);
406 Value c0x00000014 =
createConst(loc, i32Ty, 0x14, rewriter);
407 Value bits1To24 = arith::ShLIOp::create(
b, i4BitsNoSign, c0x2);
409 arith::CmpIOp::create(
b, arith::CmpIPredicate::eq, i4BitsNoSign, c0x1);
410 bits1To24 = arith::SelectOp::create(
b, isHalf, c0x0, bits1To24);
411 bits1To24 = arith::ExtUIOp::create(
b, i32Ty, bits1To24);
412 bits1To24 = arith::ShLIOp::create(
b, bits1To24, c0x00000014);
415 Value zeroExpBits =
createConst(loc, i32Ty, 0x00000000, rewriter);
416 Value highExpBits =
createConst(loc, i32Ty, 0x40000000, rewriter);
417 Value lowExpBits =
createConst(loc, i32Ty, 0x3f000000, rewriter);
419 arith::CmpIOp::create(
b, arith::CmpIPredicate::uge, i4BitsNoSign, c0x4);
421 arith::SelectOp::create(
b, useLargerExp, highExpBits, lowExpBits);
423 arith::CmpIOp::create(
b, arith::CmpIPredicate::eq, i4BitsNoSign, c0x0);
424 bits25To31 = arith::SelectOp::create(
b, zeroExp, zeroExpBits, bits25To31);
427 Value c0x80000000 =
createConst(loc, i32Ty, 0x80000000, rewriter);
428 Value c0x8 =
createConst(loc, i4Ty, 0x8, rewriter);
430 arith::CmpIOp::create(
b, arith::CmpIPredicate::uge, i4Bits, c0x8);
432 arith::SelectOp::create(
b, negative, c0x80000000, zeroExpBits);
435 Value bits1To31 = arith::AddIOp::create(
b, bits1To24, bits25To31);
436 Value bits1To32 = arith::AddIOp::create(
b, bits1To31, bit32);
437 Value
result = arith::BitcastOp::create(
b, f32Ty, bits1To32);
438 if (!isa<Float32Type>(resultETy))
441 rewriter.replaceOp(op,
result);
448 LogicalResult matchAndRewrite(arith::ExtFOp op,
449 PatternRewriter &rewriter)
const final {
450 ImplicitLocOpBuilder
b(op.getLoc(), rewriter);
451 Value operand = op.getOperand();
452 Type operandTy = operand.
getType();
453 Type resultTy = op.getType();
457 if (!llvm::isa<Float8E8M0FNUType>(operandETy)) {
458 return rewriter.notifyMatchFailure(op,
"not a ext of F8E8M0FNU");
465 Value bitcast = arith::BitcastOp::create(
b, i8Ty, operand);
466 Value cF32MantissaWidth =
createConst(op->getLoc(), i32Ty, 23, rewriter);
467 Value exti = arith::ExtUIOp::create(
b, i32Ty, bitcast);
468 Value f32Bits = arith::ShLIOp::create(
b, exti, cF32MantissaWidth);
471 auto fastMath = op.getFastmathAttr();
472 bool NoNaN = fastMath
473 ? (fastMath.getValue() & arith::FastMathFlags::nnan) ==
474 arith::FastMathFlags::nnan
477 Value cF8NaN =
createConst(op.getLoc(), i8Ty, 0xff, rewriter);
478 Value cF32NaN =
createConst(op.getLoc(), i32Ty, 0xffffffff, rewriter);
480 arith::CmpIOp::create(
b, arith::CmpIPredicate::eq, bitcast, cF8NaN);
482 f32Bits = arith::SelectOp::create(
b, isNan, cF32NaN, f32Bits);
485 Value
result = arith::BitcastOp::create(
b, f32Ty, f32Bits);
487 result = arith::TruncFOp::create(
b, resultTy,
result,
nullptr,
488 op.getFastmathAttr());
490 result = arith::ExtFOp::create(
b, resultTy,
result, op.getFastmathAttr());
492 rewriter.replaceOp(op,
result);
527 LogicalResult matchAndRewrite(arith::TruncFOp op,
528 PatternRewriter &rewriter)
const final {
529 Location loc = op.getLoc();
530 ImplicitLocOpBuilder
b(loc, rewriter);
531 Value operand = op.getOperand();
532 Type operandTy = operand.
getType();
533 Type resultTy = op.getType();
542 if (!isa<Float4E2M1FNType>(resultETy))
543 return rewriter.notifyMatchFailure(op,
"not a trunc of F4E2M1FN");
544 if (!isa<Float32Type>(operandETy))
546 arith::ExtFOp::create(
b, f32Ty, operand, arith::FastMathFlagsAttr{});
550 Value c0x00000016 =
createConst(loc, i32Ty, 22, rewriter);
551 Value c0x00 =
createConst(loc, i8Ty, 0x00, rewriter);
552 Value c0xff =
createConst(loc, i8Ty, 0xff, rewriter);
553 Value zeroExpBits =
createConst(loc, i32Ty, 0, rewriter);
558 Value operandClamped = arith::MinNumFOp::create(
b, cHigherBound, operand);
559 operandClamped = arith::MaxNumFOp::create(
b, cLowerBound, operandClamped);
560 Value f32Bits = arith::BitcastOp::create(
b, i32Ty, operandClamped);
563 Value cF32ExpManWidth =
createConst(loc, i32Ty, 31, rewriter);
564 Value f32Sign = arith::ShRUIOp::create(
b, f32Bits, cF32ExpManWidth);
565 Value f4Sign = arith::TruncIOp::create(
b, i4Ty, f32Sign);
566 Value f4Bits = arith::ShLIOp::create(
b, f4Sign, c0x3);
569 Value biasAdjustment =
createConst(loc, i32Ty, 0x7e, rewriter);
570 Value cF4MantissaWidth = c0x1;
571 Value cF32MantissaWidth =
createConst(loc, i32Ty, 23, rewriter);
572 Value f32SignExp = arith::ShRUIOp::create(
b, f32Bits, cF32MantissaWidth);
573 Value biasAdjustedSignExp =
574 arith::SubIOp::create(
b, f32SignExp, biasAdjustment);
575 Value f4Exp = arith::TruncIOp::create(
b, i4Ty, biasAdjustedSignExp);
576 f4Exp = arith::ShLIOp::create(
b, f4Exp, cF4MantissaWidth);
577 f4Bits = arith::AddIOp::create(
b, f4Bits, f4Exp);
580 Value cF32FirstBitMask =
createConst(loc, i32Ty, 0x400000, rewriter);
581 Value man1Bit = arith::AndIOp::create(
b, f32Bits, cF32FirstBitMask);
582 man1Bit = arith::ShRUIOp::create(
b, man1Bit, c0x00000016);
583 Value f4Man = arith::TruncIOp::create(
b, i4Ty, man1Bit);
584 f4Bits = arith::AddIOp::create(
b, f4Bits, f4Man);
587 Value cF32MantissaMask =
createConst(loc, i32Ty, 0x7fffff, rewriter);
588 Value f8Exp = arith::TruncIOp::create(
b, i8Ty, biasAdjustedSignExp);
590 arith::CmpIOp::create(
b, arith::CmpIPredicate::sle, f8Exp, c0x00);
592 arith::CmpIOp::create(
b, arith::CmpIPredicate::eq, f8Exp, c0xff);
593 Value man23Bits = arith::AndIOp::create(
b, f32Bits, cF32MantissaMask);
594 Value isNonZeroMan = arith::CmpIOp::create(
b, arith::CmpIPredicate::ugt,
595 man23Bits, zeroExpBits);
596 Value roundToHalf = arith::AndIOp::create(
b, isNegOneExp, isNonZeroMan);
598 arith::CmpIOp::create(
b, arith::CmpIPredicate::eq, f8Exp, c0x00);
599 Value subnormalF4Bits =
createConst(loc, i4Ty, 0xf, rewriter);
600 Value halfF4Bits =
createConst(loc, i4Ty, 0x0, rewriter);
602 arith::SelectOp::create(
b, isSubnormal, subnormalF4Bits, f4Bits);
603 subResult = arith::SelectOp::create(
b, roundToHalf, halfF4Bits, subResult);
604 f4Bits = arith::SelectOp::create(
b, isZeroExp, f4Bits, subResult);
607 Value cF32Last22BitMask =
createConst(loc, i32Ty, 0x3fffff, rewriter);
608 Value cRound =
createConst(loc, i32Ty, 0x200000, rewriter);
609 Value man22Bits = arith::AndIOp::create(
b, f32Bits, cF32Last22BitMask);
611 arith::CmpIOp::create(
b, arith::CmpIPredicate::uge, man22Bits, cRound);
612 shouldRound = arith::OrIOp::create(
b, shouldRound, isSubnormal);
613 Value roundedF4Bits = arith::AddIOp::create(
b, f4Bits, c0x1);
614 f4Bits = arith::SelectOp::create(
b, shouldRound, roundedF4Bits, f4Bits);
616 Value
result = arith::BitcastOp::create(
b, resultTy, f4Bits);
617 rewriter.replaceOp(op,
result);
629 LogicalResult matchAndRewrite(arith::TruncFOp op,
630 PatternRewriter &rewriter)
const final {
631 ImplicitLocOpBuilder
b(op.getLoc(), rewriter);
632 Value operand = op.getOperand();
633 Type operandTy = operand.
getType();
635 Type resultTy = op.getType();
637 if (!llvm::isa<Float8E8M0FNUType>(resultETy)) {
638 return rewriter.notifyMatchFailure(op,
"not a truncf to f8E8M0FNU");
641 if (op.getRoundingmodeAttr()) {
642 return rewriter.notifyMatchFailure(
643 op,
"only applicable to default rounding mode.");
651 operand = arith::ExtFOp::create(
b, f32Ty, operand, op.getFastmathAttr());
653 operand = arith::TruncFOp::create(
654 b, f32Ty, operand, op.getRoundingmodeAttr(), op.getFastmathAttr());
656 Value f32Bits = arith::BitcastOp::create(
b, i32Ty, operand);
657 Value cF32MantissaWidth =
createConst(op->getLoc(), i32Ty, 23, rewriter);
658 Value f32SignExp = arith::ShRUIOp::create(
b, f32Bits, cF32MantissaWidth);
659 Value exp8Bits = arith::TruncIOp::create(
b, i8Ty, f32SignExp);
660 Value
result = arith::BitcastOp::create(
b, resultTy, exp8Bits);
661 rewriter.replaceOp(op,
result);
674 LogicalResult matchAndRewrite(arith::ExtFOp op,
675 PatternRewriter &rewriter)
const final {
676 ImplicitLocOpBuilder
b(op.getLoc(), rewriter);
677 Value operand = op.getOperand();
678 Type operandTy = operand.
getType();
679 Type resultTy = op.getType();
683 if (!llvm::isa<Float8E5M2Type>(operandETy))
684 return rewriter.notifyMatchFailure(op,
"not a ext of F8E5M2");
690 Value bitcast = arith::BitcastOp::create(
b, i8Ty, operand);
691 Value exti = arith::ExtUIOp::create(
b, i16Ty, bitcast);
692 Value c8 =
createConst(op.getLoc(), i16Ty, 8, rewriter);
693 Value f16Bits = arith::ShLIOp::create(
b, exti, c8);
694 Value f16 = arith::BitcastOp::create(
b, f16Ty, f16Bits);
697 if (!resultETy.
isF16()) {
699 result = arith::TruncFOp::create(
b, resultTy, f16,
nullptr,
700 op.getFastmathAttr());
702 result = arith::ExtFOp::create(
b, resultTy, f16, op.getFastmathAttr());
704 rewriter.replaceOp(op,
result);
718 LogicalResult matchAndRewrite(arith::TruncFOp op,
719 PatternRewriter &rewriter)
const final {
720 ImplicitLocOpBuilder
b(op.getLoc(), rewriter);
721 Value operand = op.getOperand();
722 Type operandTy = operand.
getType();
723 Type resultTy = op.getType();
727 if (!llvm::isa<Float8E5M2Type>(resultETy))
728 return rewriter.notifyMatchFailure(op,
"not a trunc to F8E5M2");
729 if (op.getRoundingmodeAttr())
730 return rewriter.notifyMatchFailure(
731 op,
"only applicable to default rounding mode.");
738 if (!operandETy.
isF16())
739 h16 = arith::TruncFOp::create(
b, f16Ty, operand,
nullptr,
740 op.getFastmathAttr());
742 Value isNan = arith::CmpFOp::create(
b, arith::CmpFPredicate::UNE, h16, h16);
743 Value h16Bits = arith::BitcastOp::create(
b, i16Ty, h16);
745 Value c7F =
createConst(op.getLoc(), i16Ty, 0x7f, rewriter);
746 Value c8 =
createConst(op.getLoc(), i16Ty, 8, rewriter);
747 Value c1 =
createConst(op.getLoc(), i16Ty, 1, rewriter);
749 arith::AndIOp::create(
b, arith::ShRUIOp::create(
b, h16Bits, c8), c1);
750 Value roundingBias = arith::AddIOp::create(
b, bit8, c7F);
751 Value biased = arith::AddIOp::create(
b, h16Bits, roundingBias);
752 Value biasedAndShifted = arith::ShRUIOp::create(
b, biased, c8);
753 Value normalCaseResult = arith::TruncIOp::create(
b, i8Ty, biasedAndShifted);
755 Value cNan =
createConst(op.getLoc(), i8Ty, 0x7e, rewriter);
756 Value select = arith::SelectOp::create(
b, isNan, cNan, normalCaseResult);
757 Value
result = arith::BitcastOp::create(
b, resultTy, select);
758 rewriter.replaceOp(op,
result);
773 LogicalResult matchAndRewrite(arith::ExtFOp op,
774 PatternRewriter &rewriter)
const final {
775 ImplicitLocOpBuilder
b(op.getLoc(), rewriter);
776 Value operand = op.getOperand();
777 Type operandTy = operand.
getType();
778 Type resultTy = op.getType();
782 if (!llvm::isa<Float8E4M3FNType>(operandETy))
783 return rewriter.notifyMatchFailure(op,
"not a ext of F8E4M3FN");
791 Value bits = arith::BitcastOp::create(
b, i8Ty, operand);
792 Value c7F8 =
createConst(op.getLoc(), i8Ty, 0x7f, rewriter);
793 Value mag8 = arith::AndIOp::create(
b, bits, c7F8);
795 Value mag16 = arith::ExtUIOp::create(
b, i16Ty, mag8);
796 Value c7 =
createConst(op.getLoc(), i16Ty, 7, rewriter);
797 Value g16Bits = arith::ShLIOp::create(
b, mag16, c7);
798 Value g16 = arith::BitcastOp::create(
b, f16Ty, g16Bits);
799 Value gF32 = arith::ExtFOp::create(
b, f32Ty, g16, op.getFastmathAttr());
802 Value magF32 = arith::MulFOp::create(
b, gF32, c256, op.getFastmathAttr());
804 Value magI32 = arith::BitcastOp::create(
b, i32Ty, magF32);
805 Value c80I8 =
createConst(op.getLoc(), i8Ty, 0x80, rewriter);
806 Value sign8 = arith::AndIOp::create(
b, bits, c80I8);
807 Value sign32 = arith::ExtUIOp::create(
b, i32Ty, sign8);
808 Value c24 =
createConst(op.getLoc(), i32Ty, 24, rewriter);
809 Value signBit = arith::ShLIOp::create(
b, sign32, c24);
810 Value signedI32 = arith::OrIOp::create(
b, magI32, signBit);
811 Value signedF32 = arith::BitcastOp::create(
b, f32Ty, signedI32);
814 arith::CmpIOp::create(
b, arith::CmpIPredicate::eq, mag8, c7F8);
815 Value cNan32 =
createConst(op.getLoc(), i32Ty, 0x7fc00000, rewriter);
816 Value nanSigned = arith::OrIOp::create(
b, cNan32, signBit);
817 Value nanF32 = arith::BitcastOp::create(
b, f32Ty, nanSigned);
818 Value resultF32 = arith::SelectOp::create(
b, isNan, nanF32, signedF32);
821 if (!resultETy.
isF32()) {
823 result = arith::TruncFOp::create(
b, resultTy, resultF32,
nullptr,
824 op.getFastmathAttr());
827 arith::ExtFOp::create(
b, resultTy, resultF32, op.getFastmathAttr());
829 rewriter.replaceOp(op,
result);
842struct F8E4M3FNTruncFOpConverter :
public OpRewritePattern<arith::TruncFOp> {
844 LogicalResult matchAndRewrite(arith::TruncFOp op,
845 PatternRewriter &rewriter)
const final {
846 ImplicitLocOpBuilder
b(op.getLoc(), rewriter);
847 Value operand = op.getOperand();
848 Type operandTy = operand.
getType();
849 Type resultTy = op.getType();
853 if (!llvm::isa<Float8E4M3FNType>(resultETy))
854 return rewriter.notifyMatchFailure(op,
"not a trunc to F8E4M3FN");
855 if (op.getRoundingmodeAttr())
856 return rewriter.notifyMatchFailure(
857 op,
"only applicable to default rounding mode.");
866 if (!operandETy.
isF32()) {
868 f32 = arith::ExtFOp::create(
b, f32Ty, operand, op.getFastmathAttr());
870 f32 = arith::TruncFOp::create(
b, f32Ty, operand,
nullptr,
871 op.getFastmathAttr());
874 Value isNan = arith::CmpFOp::create(
b, arith::CmpFPredicate::UNE, f32, f32);
876 Value f32Bits = arith::BitcastOp::create(
b, i32Ty, f32);
877 Value cSignMask =
createConst(op.getLoc(), i32Ty, 0x80000000, rewriter);
878 Value cAbsMask =
createConst(op.getLoc(), i32Ty, 0x7fffffff, rewriter);
879 Value signBits = arith::AndIOp::create(
b, f32Bits, cSignMask);
880 Value absBits = arith::AndIOp::create(
b, f32Bits, cAbsMask);
881 Value absF32 = arith::BitcastOp::create(
b, f32Ty, absBits);
887 arith::CmpFOp::create(
b, arith::CmpFPredicate::OGT, absF32, cOverflow);
892 absF32 = arith::MinNumFOp::create(
b, absF32, cMax);
895 Value scaled = arith::MulFOp::create(
b, absF32, cInv256,
nullptr);
896 Value h16 = arith::TruncFOp::create(
b, f16Ty, scaled,
nullptr,
897 op.getFastmathAttr());
898 Value h16Bits = arith::BitcastOp::create(
b, i16Ty, h16);
900 Value c3F =
createConst(op.getLoc(), i16Ty, 0x3f, rewriter);
901 Value c7 =
createConst(op.getLoc(), i16Ty, 7, rewriter);
902 Value c1 =
createConst(op.getLoc(), i16Ty, 1, rewriter);
904 arith::AndIOp::create(
b, arith::ShRUIOp::create(
b, h16Bits, c7), c1);
905 Value roundingBias = arith::AddIOp::create(
b, bit7, c3F);
906 Value biased = arith::AddIOp::create(
b, h16Bits, roundingBias);
907 Value shifted = arith::ShRUIOp::create(
b, biased, c7);
908 Value mag8 = arith::TruncIOp::create(
b, i8Ty, shifted);
909 Value c7F8 =
createConst(op.getLoc(), i8Ty, 0x7f, rewriter);
910 mag8 = arith::AndIOp::create(
b, mag8, c7F8);
912 Value c24 =
createConst(op.getLoc(), i32Ty, 24, rewriter);
913 Value sign8 = arith::TruncIOp::create(
914 b, i8Ty, arith::ShRUIOp::create(
b, signBits, c24));
915 Value res8 = arith::OrIOp::create(
b, mag8, sign8);
917 Value isNanOrOverflow = arith::OrIOp::create(
b, isNan, isOverflow);
918 Value cNan8 =
createConst(op.getLoc(), i8Ty, 0x7f, rewriter);
919 Value res = arith::SelectOp::create(
b, isNanOrOverflow, cNan8, res8);
920 Value
result = arith::BitcastOp::create(
b, resultTy, res);
921 rewriter.replaceOp(op,
result);
926struct ScalingExtFOpConverter :
public OpRewritePattern<arith::ScalingExtFOp> {
928 LogicalResult matchAndRewrite(arith::ScalingExtFOp op,
929 PatternRewriter &rewriter)
const final {
930 ImplicitLocOpBuilder
b(op.getLoc(), rewriter);
931 Value inputOperand = op.getIn();
932 Value scaleOperand = op.getScale();
933 Type scaleTy = scaleOperand.
getType();
937 scaleETy =
b.getF8E8M0Type();
939 scaleOperand = arith::TruncFOp::create(
b, scaleTy, scaleOperand,
nullptr,
940 op.getFastmathAttr());
943 if (!llvm::isa<Float8E8M0FNUType>(scaleETy)) {
944 return rewriter.notifyMatchFailure(
945 op,
"scaling_extf is using scales of type which can not be converted "
948 Type resultTy = op.getType();
952 arith::ExtFOp::create(
b, resultTy, scaleOperand, op.getFastmathAttr());
954 arith::ExtFOp::create(
b, resultTy, inputOperand, op.getFastmathAttr());
956 arith::MulFOp::create(
b, inputExt, scaleExt, op.getFastmathAttr());
957 rewriter.replaceOp(op,
result);
967struct ScalingTruncFOpConverter
970 LogicalResult matchAndRewrite(arith::ScalingTruncFOp op,
971 PatternRewriter &rewriter)
const final {
972 ImplicitLocOpBuilder
b(op.getLoc(), rewriter);
973 Value inputOperand = op.getIn();
974 Value scaleOperand = op.getScale();
975 Type scaleTy = scaleOperand.
getType();
979 scaleETy =
b.getF8E8M0Type();
981 scaleOperand = arith::TruncFOp::create(
b, scaleTy, scaleOperand,
nullptr,
982 op.getFastmathAttr());
984 if (!llvm::isa<Float8E8M0FNUType>(scaleETy)) {
985 return rewriter.notifyMatchFailure(
986 op,
"scaling_truncf is using scales type which can not be converted "
989 Type resultTy = op.getType();
990 Type inputTy = inputOperand.
getType();
994 arith::ExtFOp::create(
b, inputTy, scaleOperand, op.getFastmathAttr());
995 Value
result = arith::DivFOp::create(
b, inputOperand, scaleOperand,
996 op.getFastmathAttr());
997 Value resultCast = arith::TruncFOp::create(
998 b, resultTy,
result, op.getRoundingmodeAttr(), op.getFastmathAttr());
999 rewriter.replaceOp(op, resultCast);
1020struct FlushDenormalsOpConverter
1023 LogicalResult matchAndRewrite(arith::FlushDenormalsOp op,
1024 PatternRewriter &rewriter)
const final {
1025 Location loc = op.getLoc();
1026 ImplicitLocOpBuilder
b(loc, rewriter);
1027 Value operand = op.getOperand();
1028 Type operandTy = operand.
getType();
1031 return rewriter.notifyMatchFailure(op,
"operand is not a float type");
1033 const llvm::fltSemantics &sem = floatTy.getFloatSemantics();
1036 if (!llvm::APFloatBase::isIEEELikeFP(sem))
1037 return rewriter.notifyMatchFailure(
1038 op,
"only IEEE-like floating-point types are supported");
1040 unsigned totalBits = llvm::APFloatBase::semanticsSizeInBits(sem);
1041 unsigned precision = llvm::APFloatBase::semanticsPrecision(sem);
1044 if (precision < 1 || precision > totalBits)
1045 return rewriter.notifyMatchFailure(op,
"unexpected float semantics");
1046 unsigned mantissaBits = precision - 1;
1047 unsigned expBits = totalBits - 1 - mantissaBits;
1048 if (expBits == 0 || mantissaBits == 0)
1049 return rewriter.notifyMatchFailure(
1050 op,
"degenerate float encoding has no exponent or mantissa");
1054 Value bits = arith::BitcastOp::create(
b, intTy, operand);
1056 APInt::getBitsSet(totalBits, mantissaBits, mantissaBits + expBits);
1057 APInt clearMantissaMaskVal = ~APInt::getLowBitsSet(totalBits, mantissaBits);
1058 APInt zeroVal = APInt::getZero(totalBits);
1060 Value clearMantissaMask =
1065 Value expField = arith::AndIOp::create(
b, bits, expMask);
1067 arith::CmpIOp::create(
b, arith::CmpIPredicate::eq, expField, zero);
1070 Value cleared = arith::AndIOp::create(
b, bits, clearMantissaMask);
1071 Value resultBits = arith::SelectOp::create(
b, expIsZero, cleared, bits);
1072 Value
result = arith::BitcastOp::create(
b, operandTy, resultBits);
1074 rewriter.replaceOp(op,
result);
1079struct ArithExpandOpsPass
1080 :
public arith::impl::ArithExpandOpsPassBase<ArithExpandOpsPass> {
1081 using ArithExpandOpsPassBase::ArithExpandOpsPassBase;
1083 void runOnOperation()
override {
1087 arith::populateCeilFloorDivExpandOpsPatterns(patterns);
1088 arith::populateExpandScalingExtTruncPatterns(patterns);
1090 target.addLegalDialect<arith::ArithDialect>();
1091 target.addLegalDialect<vector::VectorDialect>();
1097 arith::FloorDivSIOp,
1098 arith::ScalingExtFOp,
1099 arith::ScalingTruncFOp
1109 if (includeMinMaxF) {
1110 arith::populateExpandMinMaxFPatterns(patterns);
1120 if (includeMinMaxI) {
1121 arith::populateExpandMinMaxIPatterns(patterns);
1133 arith::populateExpandBFloat16Patterns(patterns);
1135 arith::populateExpandF8E8M0Patterns(patterns);
1137 arith::populateExpandF4E2M1Patterns(patterns);
1139 arith::populateExpandF8E5M2Patterns(patterns);
1140 if (includeF8E4M3FN)
1141 arith::populateExpandF8E4M3FNPatterns(patterns);
1142 if (includeFlushDenormals) {
1143 arith::populateExpandFlushDenormalsPatterns(patterns);
1146 target.addDynamicallyLegalOp<arith::FlushDenormalsOp>(
1147 [](arith::FlushDenormalsOp op) {
1152 return !llvm::APFloatBase::isIEEELikeFP(
1153 floatTy.getFloatSemantics());
1157 target.addDynamicallyLegalOp<arith::ExtFOp>([=](arith::ExtFOp op) {
1160 bool legalTypes =
true;
1162 legalTypes &= !(inETy.
isBF16() && outETy.
isF32());
1164 legalTypes &= !llvm::isa<Float8E8M0FNUType>(inETy);
1166 legalTypes &= !llvm::isa<Float4E2M1FNType>(inETy);
1168 legalTypes &= !llvm::isa<Float8E5M2Type>(inETy);
1169 if (includeF8E4M3FN)
1170 legalTypes &= !llvm::isa<Float8E4M3FNType>(inETy);
1174 target.addDynamicallyLegalOp<arith::TruncFOp>([=](arith::TruncFOp op) {
1177 bool legalTypes =
true;
1179 legalTypes &= !(inETy.
isF32() && outETy.
isBF16());
1181 legalTypes &= !(llvm::isa<Float8E8M0FNUType>(outETy));
1183 legalTypes &= !llvm::isa<Float4E2M1FNType>(outETy);
1185 legalTypes &= !llvm::isa<Float8E5M2Type>(outETy);
1186 if (includeF8E4M3FN)
1187 legalTypes &= !llvm::isa<Float8E4M3FNType>(outETy);
1192 if (
failed(applyPartialConversion(getOperation(),
target,
1193 std::move(patterns))))
1194 signalPassFailure();
1203 .
add<CeilDivSIOpConverter, CeilDivUIOpConverter, FloorDivSIOpConverter>(
1208 patterns.
add<BFloat16ExtFOpConverter, BFloat16TruncFOpConverter>(
1213 patterns.
add<F4E2M1ExtFOpConverter, F4E2M1TruncFOpConverter>(
1218 patterns.
add<F8E5M2ExtFOpConverter, F8E5M2TruncFOpConverter>(
1223 patterns.
add<F8E4M3FNExtFOpConverter, F8E4M3FNTruncFOpConverter>(
1228 patterns.
add<F8E8M0ExtFOpConverter, F8E8M0TruncFOpConverter>(
1234 patterns.
add<ScalingExtFOpConverter, ScalingTruncFOpConverter>(
1240 patterns.
add<FlushDenormalsOpConverter>(patterns.
getContext());
1246 MaximumMinimumFOpConverter<MaximumFOp, arith::CmpFPredicate::UGT>,
1247 MaximumMinimumFOpConverter<MinimumFOp, arith::CmpFPredicate::ULT>,
1248 MaxNumMinNumFOpConverter<MaxNumFOp, arith::CmpFPredicate::UGT>,
1249 MaxNumMinNumFOpConverter<MinNumFOp, arith::CmpFPredicate::ULT>
1257 MaxMinIOpConverter<MaxSIOp, arith::CmpIPredicate::sgt>,
1258 MaxMinIOpConverter<MaxUIOp, arith::CmpIPredicate::ugt>,
1259 MaxMinIOpConverter<MinSIOp, arith::CmpIPredicate::slt>,
1260 MaxMinIOpConverter<MinUIOp, arith::CmpIPredicate::ult>
static int64_t product(ArrayRef< int64_t > vals)
IntegerAttr getIntegerAttr(Type type, int64_t value)
FloatAttr getFloatAttr(Type type, double value)
static DenseElementsAttr get(ShapedType type, ArrayRef< Attribute > values)
Constructs a dense elements attribute from an array of element values.
This class defines the main interface for locations in MLIR and acts as a non-nullable wrapper around...
A special type of RewriterBase that coordinates the application of a rewrite pattern on the current I...
MLIRContext * getContext() const
RewritePatternSet & add(ConstructorArg &&arg, ConstructorArgs &&...args)
Add an instance of each of the pattern types 'Ts' to the pattern list with the given arguments.
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
unsigned getIntOrFloatBitWidth() const
Return the bit width of an integer or a float type, assert failure on other types.
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
Type getType() const
Return the type of this value.
void populateExpandF8E4M3FNPatterns(RewritePatternSet &patterns)
Add patterns to expand Arith f8e4m3fn patterns to lower level bitcasts/shifts.
void populateExpandBFloat16Patterns(RewritePatternSet &patterns)
Add patterns to expand Arith bf16 patterns to lower level bitcasts/shifts.
void populateExpandScalingExtTruncPatterns(RewritePatternSet &patterns)
Add patterns to expand scaling ExtF/TruncF ops to equivalent arith ops.
void populateExpandF8E8M0Patterns(RewritePatternSet &patterns)
Add patterns to expand Arith f8e8m0 patterns to lower level bitcasts/shifts.
void populateCeilFloorDivExpandOpsPatterns(RewritePatternSet &patterns)
Add patterns to expand Arith ceil/floor division ops.
void populateExpandF4E2M1Patterns(RewritePatternSet &patterns)
Add patterns to expand Arith f4e2m1 patterns to lower level bitcasts/shifts.
void populateExpandFlushDenormalsPatterns(RewritePatternSet &patterns)
Add patterns to expand arith.flush_denormals into integer arithmetic (bitcast + bit masks + compare +...
void populateExpandMinMaxFPatterns(RewritePatternSet &patterns)
Add patterns to expand the floating-point min/max ops (arith.maximumf/ minimumf/maxnumf/minnumf) into...
void populateExpandMinMaxIPatterns(RewritePatternSet &patterns)
Add patterns to expand the signed/unsigned integer min/max ops (arith.maxsi/maxui/minsi/minui) into c...
void populateExpandMinMaxPatterns(RewritePatternSet &patterns)
Add patterns to expand both the floating-point and integer min/max ops into cmpf/cmpi + select sequen...
void populateArithExpandOpsPatterns(RewritePatternSet &patterns)
Add patterns to expand Arith ops.
void populateExpandF8E5M2Patterns(RewritePatternSet &patterns)
Add patterns to expand Arith f8e5m2 patterns to lower level bitcasts/shifts.
int compare(const Fraction &x, const Fraction &y)
Three-way comparison between two fractions.
Include the generated interface declarations.
Type getElementTypeOrSelf(Type type)
Return the element type or return the type itself.
OpRewritePattern is a wrapper around RewritePattern that allows for matching and rewriting against an...