32#include "llvm/ADT/APFloat.h"
33#include "llvm/IR/LLVMContext.h"
34#include "llvm/Support/Casting.h"
46 assert(rank > 0 &&
"0-D vector corner case should have been handled already");
48 Type idxType = typeConverter.convertType(rewriter.getIndexType());
49 auto constant = LLVM::ConstantOp::create(
50 rewriter, loc, idxType, rewriter.getIntegerAttr(idxType, pos));
51 return LLVM::InsertElementOp::create(rewriter, loc, llvmType, val1, val2,
54 return LLVM::InsertValueOp::create(rewriter, loc, val1, val2, pos);
62 Type idxType = typeConverter.convertType(rewriter.getIndexType());
63 auto constant = LLVM::ConstantOp::create(
64 rewriter, loc, idxType, rewriter.getIntegerAttr(idxType, pos));
65 return LLVM::ExtractElementOp::create(rewriter, loc, llvmType, val,
68 return LLVM::ExtractValueOp::create(rewriter, loc, val, pos);
73 VectorType vectorType,
unsigned &align) {
74 Type convertedVectorTy = typeConverter.convertType(vectorType);
75 if (!convertedVectorTy)
78 llvm::LLVMContext llvmContext;
88 MemRefType memrefType,
unsigned &align) {
89 Type elementTy = typeConverter.convertType(memrefType.getElementType());
95 llvm::LLVMContext llvmContext;
108 VectorType vectorType,
109 MemRefType memrefType,
unsigned &align,
110 bool useVectorAlignment) {
111 if (useVectorAlignment) {
126 if (!memRefType.isLastDimUnitStride())
136 MemRefType memRefType,
Value llvmMemref,
Value base,
139 "unsupported memref type");
140 assert(vectorType.getRank() == 1 &&
"expected a 1-d vector type");
144 vectorType.getScalableDims()[0]);
145 return LLVM::GEPOp::create(
146 rewriter, loc, ptrsType,
147 typeConverter.convertType(memRefType.getElementType()), base,
index);
154 if (
auto attr = dyn_cast<Attribute>(foldResult)) {
155 auto intAttr = cast<IntegerAttr>(attr);
156 return LLVM::ConstantOp::create(builder, loc, intAttr).getResult();
159 return cast<Value>(foldResult);
165using VectorScaleOpConversion =
169class VectorBitCastOpConversion
172 using ConvertOpToLLVMPattern<vector::BitCastOp>::ConvertOpToLLVMPattern;
175 matchAndRewrite(vector::BitCastOp bitCastOp, OpAdaptor adaptor,
176 ConversionPatternRewriter &rewriter)
const override {
178 VectorType resultTy = bitCastOp.getResultVectorType();
179 if (resultTy.getRank() > 1)
181 Type newResultTy = typeConverter->convertType(resultTy);
182 rewriter.replaceOpWithNewOp<LLVM::BitcastOp>(bitCastOp, newResultTy,
183 adaptor.getOperands()[0]);
191static void replaceLoadOrStoreOp(vector::LoadOp loadOp,
192 vector::LoadOpAdaptor adaptor,
193 VectorType vectorTy,
Value ptr,
unsigned align,
194 ConversionPatternRewriter &rewriter) {
195 rewriter.replaceOpWithNewOp<LLVM::LoadOp>(loadOp, vectorTy,
ptr, align,
197 loadOp.getNontemporal());
200static void replaceLoadOrStoreOp(vector::MaskedLoadOp loadOp,
201 vector::MaskedLoadOpAdaptor adaptor,
202 VectorType vectorTy,
Value ptr,
unsigned align,
203 ConversionPatternRewriter &rewriter) {
204 rewriter.replaceOpWithNewOp<LLVM::MaskedLoadOp>(
205 loadOp, vectorTy,
ptr, adaptor.getMask(), adaptor.getPassThru(), align);
208static void replaceLoadOrStoreOp(vector::StoreOp storeOp,
209 vector::StoreOpAdaptor adaptor,
210 VectorType vectorTy,
Value ptr,
unsigned align,
211 ConversionPatternRewriter &rewriter) {
212 rewriter.replaceOpWithNewOp<LLVM::StoreOp>(storeOp, adaptor.getValueToStore(),
214 storeOp.getNontemporal());
217static void replaceLoadOrStoreOp(vector::MaskedStoreOp storeOp,
218 vector::MaskedStoreOpAdaptor adaptor,
219 VectorType vectorTy,
Value ptr,
unsigned align,
220 ConversionPatternRewriter &rewriter) {
221 rewriter.replaceOpWithNewOp<LLVM::MaskedStoreOp>(
222 storeOp, adaptor.getValueToStore(),
ptr, adaptor.getMask(), align);
227template <
class LoadOrStoreOp>
230 explicit VectorLoadStoreConversion(
const LLVMTypeConverter &typeConv,
232 bool enableGEPInboundsNuw)
233 : ConvertOpToLLVMPattern<LoadOrStoreOp>(typeConv),
234 useVectorAlignment(useVectorAlign),
235 enableGEPInboundsNuw(enableGEPInboundsNuw) {}
238 matchAndRewrite(LoadOrStoreOp loadOrStoreOp,
239 typename LoadOrStoreOp::Adaptor adaptor,
240 ConversionPatternRewriter &rewriter)
const override {
242 VectorType vectorTy = loadOrStoreOp.getVectorType();
243 if (vectorTy.getRank() > 1)
246 auto loc = loadOrStoreOp->getLoc();
247 MemRefType memRefTy = loadOrStoreOp.getMemRefType();
251 unsigned align = loadOrStoreOp.getAlignment().value_or(0);
254 memRefTy, align, useVectorAlignment)))
255 return rewriter.notifyMatchFailure(loadOrStoreOp,
256 "could not resolve alignment");
264 LLVM::GEPNoWrapFlags noWrapFlags = LLVM::GEPNoWrapFlags::none;
265 if constexpr (std::is_same_v<LoadOrStoreOp, vector::LoadOp> ||
266 std::is_same_v<LoadOrStoreOp, vector::StoreOp>) {
270 auto [strides, offset] = memRefTy.getStridesAndOffset();
271 assert((strides.empty() || strides.back() == 1) &&
272 "vector.load/store requires unit trailing memref stride");
273 if (enableGEPInboundsNuw) {
274 noWrapFlags = noWrapFlags | LLVM::GEPNoWrapFlags::inbounds;
279 "Invalid MemRef type - should have been rejected by Op verifier.");
280 noWrapFlags = noWrapFlags | LLVM::GEPNoWrapFlags::nuw;
283 auto vtype = cast<VectorType>(
284 this->typeConverter->convertType(loadOrStoreOp.getVectorType()));
287 adaptor.getIndices(), noWrapFlags);
288 replaceLoadOrStoreOp(loadOrStoreOp, adaptor, vtype, dataPtr, align,
298 const bool useVectorAlignment;
299 const bool enableGEPInboundsNuw;
303class VectorGatherOpConversion
306 explicit VectorGatherOpConversion(
const LLVMTypeConverter &typeConv,
308 : ConvertOpToLLVMPattern<vector::GatherOp>(typeConv),
309 useVectorAlignment(useVectorAlign) {}
310 using ConvertOpToLLVMPattern<vector::GatherOp>::ConvertOpToLLVMPattern;
313 matchAndRewrite(vector::GatherOp gather, OpAdaptor adaptor,
314 ConversionPatternRewriter &rewriter)
const override {
315 Location loc = gather->getLoc();
316 MemRefType memRefType = dyn_cast<MemRefType>(gather.getBaseType());
317 assert(memRefType &&
"The base should be bufferized");
321 return rewriter.notifyMatchFailure(gather,
"memref type not supported");
323 VectorType vType = gather.getVectorType();
324 if (vType.getRank() > 1) {
325 return rewriter.notifyMatchFailure(
326 gather,
"only 1-D vectors can be lowered to LLVM");
331 unsigned align = gather.getAlignment().value_or(0);
334 memRefType, align, useVectorAlignment)))
335 return rewriter.notifyMatchFailure(gather,
"could not resolve alignment");
339 adaptor.getBase(), adaptor.getOffsets());
340 Value base = adaptor.getBase();
342 getIndexedPtrs(rewriter, loc, *this->getTypeConverter(), memRefType,
343 base, ptr, adaptor.getIndices(), vType);
346 rewriter.replaceOpWithNewOp<LLVM::masked_gather>(
347 gather, typeConverter->convertType(vType), ptrs, adaptor.getMask(),
348 adaptor.getPassThru(), align);
357 const bool useVectorAlignment;
361class VectorScatterOpConversion
364 explicit VectorScatterOpConversion(
const LLVMTypeConverter &typeConv,
366 : ConvertOpToLLVMPattern<vector::ScatterOp>(typeConv),
367 useVectorAlignment(useVectorAlign) {}
369 using ConvertOpToLLVMPattern<vector::ScatterOp>::ConvertOpToLLVMPattern;
372 matchAndRewrite(vector::ScatterOp scatter, OpAdaptor adaptor,
373 ConversionPatternRewriter &rewriter)
const override {
374 auto loc = scatter->getLoc();
375 auto memRefType = dyn_cast<MemRefType>(scatter.getBaseType());
376 assert(memRefType &&
"The base should be bufferized");
380 return rewriter.notifyMatchFailure(scatter,
"memref type not supported");
382 VectorType vType = scatter.getVectorType();
383 if (vType.getRank() > 1) {
384 return rewriter.notifyMatchFailure(
385 scatter,
"only 1-D vectors can be lowered to LLVM");
390 unsigned align = scatter.getAlignment().value_or(0);
393 memRefType, align, useVectorAlignment)))
394 return rewriter.notifyMatchFailure(scatter,
395 "could not resolve alignment");
399 adaptor.getBase(), adaptor.getOffsets());
401 getIndexedPtrs(rewriter, loc, *this->getTypeConverter(), memRefType,
402 adaptor.getBase(), ptr, adaptor.getIndices(), vType);
405 rewriter.replaceOpWithNewOp<LLVM::masked_scatter>(
406 scatter, adaptor.getValueToStore(), ptrs, adaptor.getMask(), align);
415 const bool useVectorAlignment;
419class VectorExpandLoadOpConversion
422 using ConvertOpToLLVMPattern<vector::ExpandLoadOp>::ConvertOpToLLVMPattern;
425 matchAndRewrite(vector::ExpandLoadOp expand, OpAdaptor adaptor,
426 ConversionPatternRewriter &rewriter)
const override {
427 auto loc = expand->getLoc();
428 MemRefType memRefType = expand.getMemRefType();
431 auto vtype = typeConverter->convertType(expand.getVectorType());
433 adaptor.getBase(), adaptor.getIndices());
438 uint64_t alignment = expand.getAlignment().value_or(1);
440 rewriter.replaceOpWithNewOp<LLVM::masked_expandload>(
441 expand, vtype, ptr, adaptor.getMask(), adaptor.getPassThru(),
448class VectorCompressStoreOpConversion
451 using ConvertOpToLLVMPattern<vector::CompressStoreOp>::ConvertOpToLLVMPattern;
454 matchAndRewrite(vector::CompressStoreOp compress, OpAdaptor adaptor,
455 ConversionPatternRewriter &rewriter)
const override {
456 auto loc = compress->getLoc();
457 MemRefType memRefType = compress.getMemRefType();
461 adaptor.getBase(), adaptor.getIndices());
466 uint64_t alignment = compress.getAlignment().value_or(1);
468 rewriter.replaceOpWithNewOp<LLVM::masked_compressstore>(
469 compress, adaptor.getValueToStore(), ptr, adaptor.getMask(), alignment);
475class ReductionNeutralZero {};
476class ReductionNeutralIntOne {};
477class ReductionNeutralFPOne {};
478class ReductionNeutralAllOnes {};
479class ReductionNeutralSIntMin {};
480class ReductionNeutralUIntMin {};
481class ReductionNeutralSIntMax {};
482class ReductionNeutralUIntMax {};
483class ReductionNeutralFPQNaN {};
484class ReductionNeutralFPNegQNaN {};
485class ReductionNeutralFPNegInf {};
486class ReductionNeutralFPPosInf {};
487class ReductionNeutralFPLowestFinite {};
488class ReductionNeutralFPLargestFinite {};
492 ConversionPatternRewriter &rewriter,
494 return LLVM::ConstantOp::create(rewriter, loc, llvmType,
495 rewriter.getZeroAttr(llvmType));
500 ConversionPatternRewriter &rewriter,
502 return LLVM::ConstantOp::create(rewriter, loc, llvmType,
503 rewriter.getIntegerAttr(llvmType, 1));
508 ConversionPatternRewriter &rewriter,
510 return LLVM::ConstantOp::create(rewriter, loc, llvmType,
511 rewriter.getFloatAttr(llvmType, 1.0));
516 ConversionPatternRewriter &rewriter,
518 return LLVM::ConstantOp::create(
519 rewriter, loc, llvmType,
520 rewriter.getIntegerAttr(
526 ConversionPatternRewriter &rewriter,
528 return LLVM::ConstantOp::create(
529 rewriter, loc, llvmType,
530 rewriter.getIntegerAttr(llvmType, llvm::APInt::getSignedMinValue(
536 ConversionPatternRewriter &rewriter,
538 return LLVM::ConstantOp::create(
539 rewriter, loc, llvmType,
540 rewriter.getIntegerAttr(llvmType, llvm::APInt::getMinValue(
546 ConversionPatternRewriter &rewriter,
548 return LLVM::ConstantOp::create(
549 rewriter, loc, llvmType,
550 rewriter.getIntegerAttr(llvmType, llvm::APInt::getSignedMaxValue(
556 ConversionPatternRewriter &rewriter,
558 return LLVM::ConstantOp::create(
559 rewriter, loc, llvmType,
560 rewriter.getIntegerAttr(llvmType, llvm::APInt::getMaxValue(
566 ConversionPatternRewriter &rewriter,
568 auto floatType = cast<FloatType>(llvmType);
569 return LLVM::ConstantOp::create(
570 rewriter, loc, llvmType,
571 rewriter.getFloatAttr(
572 llvmType, llvm::APFloat::getQNaN(floatType.getFloatSemantics(),
578 ConversionPatternRewriter &rewriter,
580 auto floatType = cast<FloatType>(llvmType);
581 return LLVM::ConstantOp::create(
582 rewriter, loc, llvmType,
583 rewriter.getFloatAttr(
584 llvmType, llvm::APFloat::getQNaN(floatType.getFloatSemantics(),
590 ConversionPatternRewriter &rewriter,
592 auto floatType = cast<FloatType>(llvmType);
593 return LLVM::ConstantOp::create(
594 rewriter, loc, llvmType,
595 rewriter.getFloatAttr(llvmType,
596 llvm::APFloat::getInf(floatType.getFloatSemantics(),
602 ConversionPatternRewriter &rewriter,
604 auto floatType = cast<FloatType>(llvmType);
605 return LLVM::ConstantOp::create(
606 rewriter, loc, llvmType,
607 rewriter.getFloatAttr(llvmType,
608 llvm::APFloat::getInf(floatType.getFloatSemantics(),
614 ConversionPatternRewriter &rewriter,
616 auto floatType = cast<FloatType>(llvmType);
617 return LLVM::ConstantOp::create(
618 rewriter, loc, llvmType,
619 rewriter.getFloatAttr(
620 llvmType, llvm::APFloat::getLargest(floatType.getFloatSemantics(),
627 ConversionPatternRewriter &rewriter,
Location loc,
629 auto floatType = cast<FloatType>(llvmType);
630 return LLVM::ConstantOp::create(
631 rewriter, loc, llvmType,
632 rewriter.getFloatAttr(
633 llvmType, llvm::APFloat::getLargest(floatType.getFloatSemantics(),
639template <
class ReductionNeutral>
640static Value getOrCreateAccumulator(ConversionPatternRewriter &rewriter,
653static Value createVectorLengthValue(ConversionPatternRewriter &rewriter,
655 VectorType vType = cast<VectorType>(llvmType);
656 auto vShape = vType.getShape();
657 assert(vShape.size() == 1 &&
"Unexpected multi-dim vector type");
659 Value baseVecLength = LLVM::ConstantOp::create(
660 rewriter, loc, rewriter.getI32Type(),
661 rewriter.getIntegerAttr(rewriter.getI32Type(), vShape[0]));
663 if (!vType.getScalableDims()[0])
664 return baseVecLength;
667 Value vScale = vector::VectorScaleOp::create(rewriter, loc);
669 arith::IndexCastOp::create(rewriter, loc, rewriter.getI32Type(), vScale);
670 Value scalableVecLength =
671 arith::MulIOp::create(rewriter, loc, baseVecLength, vScale);
672 return scalableVecLength;
679template <
class LLVMRedIntrinOp,
class ScalarOp>
680static Value createIntegerReductionArithmeticOpLowering(
681 ConversionPatternRewriter &rewriter,
Location loc,
Type llvmType,
685 LLVMRedIntrinOp::create(rewriter, loc, llvmType, vectorOperand);
688 result = ScalarOp::create(rewriter, loc, accumulator,
result);
696template <
class LLVMRedIntrinOp>
697static Value createIntegerReductionComparisonOpLowering(
698 ConversionPatternRewriter &rewriter,
Location loc,
Type llvmType,
699 Value vectorOperand,
Value accumulator, LLVM::ICmpPredicate predicate) {
701 LLVMRedIntrinOp::create(rewriter, loc, llvmType, vectorOperand);
704 LLVM::ICmpOp::create(rewriter, loc, predicate, accumulator,
result);
705 result = LLVM::SelectOp::create(rewriter, loc, cmp, accumulator,
result);
711template <
typename Source>
712struct VectorToScalarMapper;
714struct VectorToScalarMapper<
LLVM::vector_reduce_fmaximum> {
715 using Type = LLVM::MaximumOp;
718struct VectorToScalarMapper<
LLVM::vector_reduce_fminimum> {
719 using Type = LLVM::MinimumOp;
722struct VectorToScalarMapper<
LLVM::vector_reduce_fmax> {
723 using Type = LLVM::MaxNumOp;
726struct VectorToScalarMapper<
LLVM::vector_reduce_fmin> {
727 using Type = LLVM::MinNumOp;
731template <
class LLVMRedIntrinOp>
732static Value createFPReductionComparisonOpLowering(
733 ConversionPatternRewriter &rewriter,
Location loc,
Type llvmType,
734 Value vectorOperand,
Value accumulator, LLVM::FastmathFlagsAttr fmf) {
736 LLVMRedIntrinOp::create(rewriter, loc, llvmType, vectorOperand, fmf);
739 result = VectorToScalarMapper<LLVMRedIntrinOp>::Type::create(
740 rewriter, loc,
result, accumulator);
746template <
class LLVMRedIntrinOp,
class ReductionNeutral>
748lowerReductionWithStartValue(ConversionPatternRewriter &rewriter,
Location loc,
750 Value accumulator, LLVM::FastmathFlagsAttr fmf) {
751 accumulator = getOrCreateAccumulator<ReductionNeutral>(rewriter, loc,
752 llvmType, accumulator);
753 return LLVMRedIntrinOp::create(rewriter, loc, llvmType,
754 accumulator, vectorOperand,
758template <
class LLVMVPRedIntrinOp,
class ReductionNeutral>
759static Value lowerPredicatedReductionWithStartValue(
760 ConversionPatternRewriter &rewriter,
Location loc,
Type llvmType,
762 accumulator = getOrCreateAccumulator<ReductionNeutral>(rewriter, loc,
763 llvmType, accumulator);
765 createVectorLengthValue(rewriter, loc, vectorOperand.
getType());
766 return LLVMVPRedIntrinOp::create(rewriter, loc, llvmType,
767 accumulator, vectorOperand,
771template <
class LLVMIntVPRedIntrinOp,
class IntReductionNeutral,
772 class LLVMFPVPRedIntrinOp,
class FPReductionNeutral>
773static Value lowerPredicatedReductionWithStartValue(
774 ConversionPatternRewriter &rewriter,
Location loc,
Type llvmType,
777 return lowerPredicatedReductionWithStartValue<LLVMIntVPRedIntrinOp,
778 IntReductionNeutral>(
779 rewriter, loc, llvmType, vectorOperand, accumulator, mask);
782 return lowerPredicatedReductionWithStartValue<LLVMFPVPRedIntrinOp,
784 rewriter, loc, llvmType, vectorOperand, accumulator, mask);
788class VectorReductionOpConversion
791 explicit VectorReductionOpConversion(
const LLVMTypeConverter &typeConv,
792 bool reassociateFPRed)
793 : ConvertOpToLLVMPattern<vector::ReductionOp>(typeConv),
794 reassociateFPReductions(reassociateFPRed) {}
797 matchAndRewrite(vector::ReductionOp reductionOp, OpAdaptor adaptor,
798 ConversionPatternRewriter &rewriter)
const override {
799 auto kind = reductionOp.getKind();
800 Type eltType = reductionOp.getDest().getType();
801 Type llvmType = typeConverter->convertType(eltType);
802 Value operand = adaptor.getVector();
803 Value acc = adaptor.getAcc();
804 Location loc = reductionOp.getLoc();
810 case vector::CombiningKind::ADD:
812 createIntegerReductionArithmeticOpLowering<LLVM::vector_reduce_add,
814 rewriter, loc, llvmType, operand, acc);
816 case vector::CombiningKind::MUL:
818 createIntegerReductionArithmeticOpLowering<LLVM::vector_reduce_mul,
820 rewriter, loc, llvmType, operand, acc);
822 case vector::CombiningKind::MINUI:
823 result = createIntegerReductionComparisonOpLowering<
824 LLVM::vector_reduce_umin>(rewriter, loc, llvmType, operand, acc,
825 LLVM::ICmpPredicate::ule);
827 case vector::CombiningKind::MINSI:
828 result = createIntegerReductionComparisonOpLowering<
829 LLVM::vector_reduce_smin>(rewriter, loc, llvmType, operand, acc,
830 LLVM::ICmpPredicate::sle);
832 case vector::CombiningKind::MAXUI:
833 result = createIntegerReductionComparisonOpLowering<
834 LLVM::vector_reduce_umax>(rewriter, loc, llvmType, operand, acc,
835 LLVM::ICmpPredicate::uge);
837 case vector::CombiningKind::MAXSI:
838 result = createIntegerReductionComparisonOpLowering<
839 LLVM::vector_reduce_smax>(rewriter, loc, llvmType, operand, acc,
840 LLVM::ICmpPredicate::sge);
842 case vector::CombiningKind::AND:
844 createIntegerReductionArithmeticOpLowering<LLVM::vector_reduce_and,
846 rewriter, loc, llvmType, operand, acc);
848 case vector::CombiningKind::OR:
850 createIntegerReductionArithmeticOpLowering<LLVM::vector_reduce_or,
852 rewriter, loc, llvmType, operand, acc);
854 case vector::CombiningKind::XOR:
856 createIntegerReductionArithmeticOpLowering<LLVM::vector_reduce_xor,
858 rewriter, loc, llvmType, operand, acc);
863 rewriter.replaceOp(reductionOp,
result);
868 if (!isa<FloatType>(eltType))
871 arith::FastMathFlagsAttr fMFAttr = reductionOp.getFastMathFlagsAttr();
872 LLVM::FastmathFlagsAttr fmf = LLVM::FastmathFlagsAttr::get(
873 reductionOp.getContext(),
875 fmf = LLVM::FastmathFlagsAttr::get(
876 reductionOp.getContext(),
877 fmf.getValue() | (reassociateFPReductions ? LLVM::FastmathFlags::reassoc
878 : LLVM::FastmathFlags::none));
882 if (kind == vector::CombiningKind::ADD) {
883 result = lowerReductionWithStartValue<LLVM::vector_reduce_fadd,
884 ReductionNeutralZero>(
885 rewriter, loc, llvmType, operand, acc, fmf);
886 }
else if (kind == vector::CombiningKind::MUL) {
887 result = lowerReductionWithStartValue<LLVM::vector_reduce_fmul,
888 ReductionNeutralFPOne>(
889 rewriter, loc, llvmType, operand, acc, fmf);
890 }
else if (kind == vector::CombiningKind::MINIMUMF) {
892 createFPReductionComparisonOpLowering<LLVM::vector_reduce_fminimum>(
893 rewriter, loc, llvmType, operand, acc, fmf);
894 }
else if (kind == vector::CombiningKind::MAXIMUMF) {
896 createFPReductionComparisonOpLowering<LLVM::vector_reduce_fmaximum>(
897 rewriter, loc, llvmType, operand, acc, fmf);
898 }
else if (kind == vector::CombiningKind::MINNUMF) {
899 result = createFPReductionComparisonOpLowering<LLVM::vector_reduce_fmin>(
900 rewriter, loc, llvmType, operand, acc, fmf);
901 }
else if (kind == vector::CombiningKind::MAXNUMF) {
902 result = createFPReductionComparisonOpLowering<LLVM::vector_reduce_fmax>(
903 rewriter, loc, llvmType, operand, acc, fmf);
908 rewriter.replaceOp(reductionOp,
result);
913 const bool reassociateFPReductions;
924template <
class MaskedOp>
925class VectorMaskOpConversionBase
928 using ConvertOpToLLVMPattern<vector::MaskOp>::ConvertOpToLLVMPattern;
931 matchAndRewrite(vector::MaskOp maskOp, OpAdaptor adaptor,
932 ConversionPatternRewriter &rewriter)
const final {
934 auto maskedOp = llvm::dyn_cast_or_null<MaskedOp>(maskOp.getMaskableOp());
937 return matchAndRewriteMaskableOp(maskOp, maskedOp, rewriter);
941 virtual LogicalResult
942 matchAndRewriteMaskableOp(vector::MaskOp maskOp,
943 vector::MaskableOpInterface maskableOp,
944 ConversionPatternRewriter &rewriter)
const = 0;
947class MaskedReductionOpConversion
948 :
public VectorMaskOpConversionBase<vector::ReductionOp> {
951 using VectorMaskOpConversionBase<
952 vector::ReductionOp>::VectorMaskOpConversionBase;
954 LogicalResult matchAndRewriteMaskableOp(
955 vector::MaskOp maskOp, MaskableOpInterface maskableOp,
956 ConversionPatternRewriter &rewriter)
const override {
957 auto reductionOp = cast<ReductionOp>(maskableOp.getOperation());
958 auto kind = reductionOp.getKind();
959 Type eltType = reductionOp.getDest().getType();
960 Type llvmType = typeConverter->convertType(eltType);
961 Value operand = reductionOp.getVector();
962 Value acc = reductionOp.getAcc();
963 Location loc = reductionOp.getLoc();
965 arith::FastMathFlagsAttr fMFAttr = reductionOp.getFastMathFlagsAttr();
966 LLVM::FastmathFlagsAttr fmf = LLVM::FastmathFlagsAttr::get(
967 reductionOp.getContext(),
970 LLVM::bitEnumContainsAny(fmf.getValue(), LLVM::FastmathFlags::ninf);
974 case vector::CombiningKind::ADD:
975 result = lowerPredicatedReductionWithStartValue<
976 LLVM::VPReduceAddOp, ReductionNeutralZero, LLVM::VPReduceFAddOp,
977 ReductionNeutralZero>(rewriter, loc, llvmType, operand, acc,
980 case vector::CombiningKind::MUL:
981 result = lowerPredicatedReductionWithStartValue<
982 LLVM::VPReduceMulOp, ReductionNeutralIntOne, LLVM::VPReduceFMulOp,
983 ReductionNeutralFPOne>(rewriter, loc, llvmType, operand, acc,
986 case vector::CombiningKind::MINUI:
987 result = lowerPredicatedReductionWithStartValue<LLVM::VPReduceUMinOp,
988 ReductionNeutralUIntMax>(
989 rewriter, loc, llvmType, operand, acc, maskOp.getMask());
991 case vector::CombiningKind::MINSI:
992 result = lowerPredicatedReductionWithStartValue<LLVM::VPReduceSMinOp,
993 ReductionNeutralSIntMax>(
994 rewriter, loc, llvmType, operand, acc, maskOp.getMask());
996 case vector::CombiningKind::MAXUI:
997 result = lowerPredicatedReductionWithStartValue<LLVM::VPReduceUMaxOp,
998 ReductionNeutralUIntMin>(
999 rewriter, loc, llvmType, operand, acc, maskOp.getMask());
1001 case vector::CombiningKind::MAXSI:
1002 result = lowerPredicatedReductionWithStartValue<LLVM::VPReduceSMaxOp,
1003 ReductionNeutralSIntMin>(
1004 rewriter, loc, llvmType, operand, acc, maskOp.getMask());
1006 case vector::CombiningKind::AND:
1007 result = lowerPredicatedReductionWithStartValue<LLVM::VPReduceAndOp,
1008 ReductionNeutralAllOnes>(
1009 rewriter, loc, llvmType, operand, acc, maskOp.getMask());
1011 case vector::CombiningKind::OR:
1012 result = lowerPredicatedReductionWithStartValue<LLVM::VPReduceOrOp,
1013 ReductionNeutralZero>(
1014 rewriter, loc, llvmType, operand, acc, maskOp.getMask());
1016 case vector::CombiningKind::XOR:
1017 result = lowerPredicatedReductionWithStartValue<LLVM::VPReduceXorOp,
1018 ReductionNeutralZero>(
1019 rewriter, loc, llvmType, operand, acc, maskOp.getMask());
1021 case vector::CombiningKind::MINNUMF:
1023 lowerPredicatedReductionWithStartValue<LLVM::VPReduceFMinOp,
1024 ReductionNeutralFPNegQNaN>(
1025 rewriter, loc, llvmType, operand, acc, maskOp.getMask());
1027 case vector::CombiningKind::MAXNUMF:
1028 result = lowerPredicatedReductionWithStartValue<LLVM::VPReduceFMaxOp,
1029 ReductionNeutralFPQNaN>(
1030 rewriter, loc, llvmType, operand, acc, maskOp.getMask());
1032 case CombiningKind::MAXIMUMF:
1037 ? lowerPredicatedReductionWithStartValue<
1038 LLVM::VPReduceFMaximumOp, ReductionNeutralFPLowestFinite>(
1039 rewriter, loc, llvmType, operand, acc, maskOp.getMask())
1040 : lowerPredicatedReductionWithStartValue<
1041 LLVM::VPReduceFMaximumOp, ReductionNeutralFPNegInf>(
1042 rewriter, loc, llvmType, operand, acc, maskOp.getMask());
1044 case CombiningKind::MINIMUMF:
1047 ? lowerPredicatedReductionWithStartValue<
1048 LLVM::VPReduceFMinimumOp, ReductionNeutralFPLargestFinite>(
1049 rewriter, loc, llvmType, operand, acc, maskOp.getMask())
1050 : lowerPredicatedReductionWithStartValue<
1051 LLVM::VPReduceFMinimumOp, ReductionNeutralFPPosInf>(
1052 rewriter, loc, llvmType, operand, acc, maskOp.getMask());
1057 rewriter.replaceOp(maskOp,
result);
1062class VectorShuffleOpConversion
1065 using ConvertOpToLLVMPattern<vector::ShuffleOp>::ConvertOpToLLVMPattern;
1068 matchAndRewrite(vector::ShuffleOp shuffleOp, OpAdaptor adaptor,
1069 ConversionPatternRewriter &rewriter)
const override {
1070 auto loc = shuffleOp->getLoc();
1071 auto v1Type = shuffleOp.getV1VectorType();
1072 auto v2Type = shuffleOp.getV2VectorType();
1073 auto vectorType = shuffleOp.getResultVectorType();
1074 Type llvmType = typeConverter->convertType(vectorType);
1075 ArrayRef<int64_t> mask = shuffleOp.getMask();
1082 int64_t rank = vectorType.getRank();
1084 bool wellFormed0DCase =
1085 v1Type.getRank() == 0 && v2Type.getRank() == 0 && rank == 1;
1086 bool wellFormedNDCase =
1087 v1Type.getRank() == rank && v2Type.getRank() == rank;
1088 assert((wellFormed0DCase || wellFormedNDCase) &&
"op is not well-formed");
1093 if (rank <= 1 && v1Type == v2Type) {
1094 Value llvmShuffleOp = LLVM::ShuffleVectorOp::create(
1095 rewriter, loc, adaptor.getV1(), adaptor.getV2(),
1096 llvm::to_vector_of<int32_t>(mask));
1097 rewriter.replaceOp(shuffleOp, llvmShuffleOp);
1102 int64_t v1Dim = v1Type.getDimSize(0);
1104 if (
auto arrayType = dyn_cast<LLVM::LLVMArrayType>(llvmType))
1105 eltType = arrayType.getElementType();
1107 eltType = cast<VectorType>(llvmType).getElementType();
1108 Value insert = LLVM::PoisonOp::create(rewriter, loc, llvmType);
1110 for (int64_t extPos : mask) {
1111 Value value = adaptor.getV1();
1112 if (extPos >= v1Dim) {
1114 value = adaptor.getV2();
1116 Value extract =
extractOne(rewriter, *getTypeConverter(), loc, value,
1117 eltType, rank, extPos);
1118 insert =
insertOne(rewriter, *getTypeConverter(), loc, insert, extract,
1119 llvmType, rank, insPos++);
1121 rewriter.replaceOp(shuffleOp, insert);
1126class VectorExtractOpConversion
1129 using ConvertOpToLLVMPattern<vector::ExtractOp>::ConvertOpToLLVMPattern;
1132 matchAndRewrite(vector::ExtractOp extractOp, OpAdaptor adaptor,
1133 ConversionPatternRewriter &rewriter)
const override {
1134 auto loc = extractOp->getLoc();
1135 auto resultType = extractOp.getResult().getType();
1136 auto llvmResultType = typeConverter->convertType(resultType);
1138 if (!llvmResultType)
1142 adaptor.getStaticPosition(), adaptor.getDynamicPosition(), rewriter);
1156 bool extractsAggregate = extractOp.getSourceVectorType().getRank() >= 2;
1160 bool extractsScalar =
static_cast<int64_t
>(positionVec.size()) ==
1161 extractOp.getSourceVectorType().getRank();
1165 if (extractOp.getSourceVectorType().getRank() == 0) {
1166 Type idxType = typeConverter->convertType(rewriter.getIndexType());
1167 positionVec.push_back(rewriter.getZeroAttr(idxType));
1170 Value extracted = adaptor.getSource();
1171 if (extractsAggregate) {
1172 ArrayRef<OpFoldResult> position(positionVec);
1173 if (extractsScalar) {
1177 position = position.drop_back();
1180 if (!llvm::all_of(position, llvm::IsaPred<Attribute>)) {
1183 extracted = LLVM::ExtractValueOp::create(rewriter, loc, extracted,
1187 if (extractsScalar) {
1188 extracted = LLVM::ExtractElementOp::create(
1189 rewriter, loc, extracted,
1193 rewriter.replaceOp(extractOp, extracted);
1214 using ConvertOpToLLVMPattern<vector::FMAOp>::ConvertOpToLLVMPattern;
1217 matchAndRewrite(vector::FMAOp fmaOp, OpAdaptor adaptor,
1218 ConversionPatternRewriter &rewriter)
const override {
1219 VectorType vType = fmaOp.getVectorType();
1220 if (vType.getRank() > 1)
1223 rewriter.replaceOpWithNewOp<LLVM::FMulAddOp>(
1224 fmaOp, adaptor.getLhs(), adaptor.getRhs(), adaptor.getAcc());
1229class VectorInsertOpConversion
1232 using ConvertOpToLLVMPattern<vector::InsertOp>::ConvertOpToLLVMPattern;
1235 matchAndRewrite(vector::InsertOp insertOp, OpAdaptor adaptor,
1236 ConversionPatternRewriter &rewriter)
const override {
1237 auto loc = insertOp->getLoc();
1238 auto destVectorType = insertOp.getDestVectorType();
1239 auto llvmResultType = typeConverter->convertType(destVectorType);
1241 if (!llvmResultType)
1245 adaptor.getStaticPosition(), adaptor.getDynamicPosition(), rewriter);
1267 bool isNestedAggregate = isa<LLVM::LLVMArrayType>(llvmResultType);
1269 bool insertIntoInnermostDim =
1270 static_cast<int64_t
>(positionVec.size()) == destVectorType.getRank();
1272 ArrayRef<OpFoldResult> positionOf1DVectorWithinAggregate(
1273 positionVec.begin(),
1274 insertIntoInnermostDim ? positionVec.size() - 1 : positionVec.size());
1275 OpFoldResult positionOfScalarWithin1DVector;
1276 if (destVectorType.getRank() == 0) {
1279 Type idxType = typeConverter->convertType(rewriter.getIndexType());
1280 positionOfScalarWithin1DVector = rewriter.getZeroAttr(idxType);
1281 }
else if (insertIntoInnermostDim) {
1282 positionOfScalarWithin1DVector = positionVec.back();
1288 Value sourceAggregate = adaptor.getValueToStore();
1289 if (insertIntoInnermostDim) {
1292 if (isNestedAggregate) {
1295 if (!llvm::all_of(positionOf1DVectorWithinAggregate,
1296 llvm::IsaPred<Attribute>)) {
1300 sourceAggregate = LLVM::ExtractValueOp::create(
1301 rewriter, loc, adaptor.getDest(),
1306 sourceAggregate = adaptor.getDest();
1309 sourceAggregate = LLVM::InsertElementOp::create(
1310 rewriter, loc, sourceAggregate.
getType(), sourceAggregate,
1311 adaptor.getValueToStore(),
1315 Value
result = sourceAggregate;
1316 if (isNestedAggregate) {
1317 if (!llvm::all_of(positionOf1DVectorWithinAggregate,
1318 llvm::IsaPred<Attribute>)) {
1322 result = LLVM::InsertValueOp::create(
1323 rewriter, loc, adaptor.getDest(), sourceAggregate,
1327 rewriter.replaceOp(insertOp,
result);
1333struct VectorScalableInsertOpLowering
1335 using ConvertOpToLLVMPattern<
1336 vector::ScalableInsertOp>::ConvertOpToLLVMPattern;
1339 matchAndRewrite(vector::ScalableInsertOp insOp, OpAdaptor adaptor,
1340 ConversionPatternRewriter &rewriter)
const override {
1341 rewriter.replaceOpWithNewOp<LLVM::vector_insert>(
1342 insOp, adaptor.getDest(), adaptor.getValueToStore(), adaptor.getPos());
1348struct VectorScalableExtractOpLowering
1350 using ConvertOpToLLVMPattern<
1351 vector::ScalableExtractOp>::ConvertOpToLLVMPattern;
1354 matchAndRewrite(vector::ScalableExtractOp extOp, OpAdaptor adaptor,
1355 ConversionPatternRewriter &rewriter)
const override {
1356 rewriter.replaceOpWithNewOp<LLVM::vector_extract>(
1357 extOp, typeConverter->convertType(extOp.getResultVectorType()),
1358 adaptor.getSource(), adaptor.getPos());
1391 setHasBoundedRewriteRecursion();
1394 LogicalResult matchAndRewrite(FMAOp op,
1395 PatternRewriter &rewriter)
const override {
1396 auto vType = op.getVectorType();
1397 if (vType.getRank() < 2)
1400 auto loc = op.getLoc();
1401 auto elemType = vType.getElementType();
1402 Value zero = arith::ConstantOp::create(rewriter, loc, elemType,
1404 Value desc = vector::BroadcastOp::create(rewriter, loc, vType, zero);
1405 for (int64_t i = 0, e = vType.getShape().front(); i != e; ++i) {
1406 Value extrLHS = ExtractOp::create(rewriter, loc, op.getLhs(), i);
1407 Value extrRHS = ExtractOp::create(rewriter, loc, op.getRhs(), i);
1408 Value extrACC = ExtractOp::create(rewriter, loc, op.getAcc(), i);
1409 Value fma = FMAOp::create(rewriter, loc, extrLHS, extrRHS, extrACC);
1410 desc = InsertOp::create(rewriter, loc, fma, desc, i);
1419static std::optional<SmallVector<int64_t, 4>>
1420computeContiguousStrides(MemRefType memRefType) {
1423 if (
failed(memRefType.getStridesAndOffset(strides, offset)))
1424 return std::nullopt;
1425 if (!strides.empty() && strides.back() != 1)
1426 return std::nullopt;
1428 if (memRefType.getLayout().isIdentity())
1435 auto sizes = memRefType.getShape();
1437 if (ShapedType::isDynamic(sizes[
index + 1]) ||
1438 ShapedType::isDynamic(strides[
index]) ||
1439 ShapedType::isDynamic(strides[
index + 1]))
1440 return std::nullopt;
1442 return std::nullopt;
1447class VectorTypeCastOpConversion
1450 using ConvertOpToLLVMPattern<vector::TypeCastOp>::ConvertOpToLLVMPattern;
1453 matchAndRewrite(vector::TypeCastOp castOp, OpAdaptor adaptor,
1454 ConversionPatternRewriter &rewriter)
const override {
1455 auto loc = castOp->getLoc();
1456 MemRefType sourceMemRefType =
1457 cast<MemRefType>(castOp.getOperand().getType());
1458 MemRefType targetMemRefType = castOp.getType();
1461 if (!sourceMemRefType.hasStaticShape() ||
1462 !targetMemRefType.hasStaticShape())
1465 auto llvmSourceDescriptorTy =
1466 dyn_cast<LLVM::LLVMStructType>(adaptor.getOperands()[0].getType());
1467 if (!llvmSourceDescriptorTy)
1469 MemRefDescriptor sourceMemRef(adaptor.getOperands()[0]);
1471 auto llvmTargetDescriptorTy = dyn_cast_or_null<LLVM::LLVMStructType>(
1472 typeConverter->convertType(targetMemRefType));
1473 if (!llvmTargetDescriptorTy)
1477 auto sourceStrides = computeContiguousStrides(sourceMemRefType);
1480 auto targetStrides = computeContiguousStrides(targetMemRefType);
1484 if (llvm::any_of(*targetStrides, ShapedType::isDynamic))
1489 Type indexTy = getTypeConverter()->getIndexType();
1492 auto desc = MemRefDescriptor::poison(rewriter, loc, llvmTargetDescriptorTy);
1494 Value allocated = sourceMemRef.allocatedPtr(rewriter, loc);
1495 desc.setAllocatedPtr(rewriter, loc, allocated);
1498 Value ptr = sourceMemRef.alignedPtr(rewriter, loc);
1499 desc.setAlignedPtr(rewriter, loc, ptr);
1501 desc.setOffset(rewriter, loc,
1502 LLVM::createIndexAttrConstant(rewriter, loc, indexTy, 0));
1505 for (
const auto &indexedSize :
1506 llvm::enumerate(targetMemRefType.getShape())) {
1507 int64_t index = indexedSize.index();
1508 desc.setSize(rewriter, loc, index,
1509 LLVM::createIndexAttrConstant(rewriter, loc, indexTy,
1510 indexedSize.value()));
1511 desc.setStride(rewriter, loc, index,
1512 LLVM::createIndexAttrConstant(rewriter, loc, indexTy,
1513 (*targetStrides)[index]));
1516 rewriter.replaceOp(castOp, {desc});
1523class VectorCreateMaskOpConversion
1524 :
public OpConversionPattern<vector::CreateMaskOp> {
1526 explicit VectorCreateMaskOpConversion(MLIRContext *context,
1527 bool enableIndexOpt)
1528 : OpConversionPattern<vector::CreateMaskOp>(context),
1529 force32BitVectorIndices(enableIndexOpt) {}
1532 matchAndRewrite(vector::CreateMaskOp op, OpAdaptor adaptor,
1533 ConversionPatternRewriter &rewriter)
const override {
1534 auto dstType = op.getType();
1535 if (dstType.getRank() != 1 || !cast<VectorType>(dstType).isScalable())
1537 IntegerType idxType =
1538 force32BitVectorIndices ? rewriter.getI32Type() : rewriter.getI64Type();
1539 auto loc = op->getLoc();
1540 Value
indices = LLVM::StepVectorOp::create(
1542 LLVM::getVectorType(idxType, dstType.getShape()[0],
1544 Value maskBound = adaptor.getOperands()[0];
1551 if (force32BitVectorIndices) {
1554 maskBound = arith::MinSIOp::create(rewriter, loc, maskBound, maxBound);
1558 Value bounds = BroadcastOp::create(rewriter, loc,
indices.getType(), bound);
1559 Value comp = arith::CmpIOp::create(rewriter, loc, arith::CmpIPredicate::slt,
1561 rewriter.replaceOp(op, comp);
1566 const bool force32BitVectorIndices;
1570 SymbolTableCollection *symbolTables =
nullptr;
1573 explicit VectorPrintOpConversion(
1574 const LLVMTypeConverter &typeConverter,
1575 SymbolTableCollection *symbolTables =
nullptr)
1576 : ConvertOpToLLVMPattern<vector::PrintOp>(typeConverter),
1577 symbolTables(symbolTables) {}
1593 matchAndRewrite(vector::PrintOp
printOp, OpAdaptor adaptor,
1594 ConversionPatternRewriter &rewriter)
const override {
1595 auto parent =
printOp->getParentOfType<ModuleOp>();
1601 if (
auto value = adaptor.getSource()) {
1603 if (isa<VectorType>(printType)) {
1607 if (
failed(emitScalarPrint(rewriter, parent, loc, printType, value)))
1611 auto punct =
printOp.getPunctuation();
1612 if (
auto stringLiteral =
printOp.getStringLiteral()) {
1614 LLVM::createPrintStrCall(rewriter, loc, parent,
"vector_print_str",
1615 *stringLiteral, *getTypeConverter(),
1617 if (createResult.failed())
1620 }
else if (punct != PrintPunctuation::NoPunctuation) {
1621 FailureOr<LLVM::LLVMFuncOp> op = [&]() {
1623 case PrintPunctuation::Close:
1624 return LLVM::lookupOrCreatePrintCloseFn(rewriter, parent,
1626 case PrintPunctuation::Open:
1627 return LLVM::lookupOrCreatePrintOpenFn(rewriter, parent,
1629 case PrintPunctuation::Comma:
1630 return LLVM::lookupOrCreatePrintCommaFn(rewriter, parent,
1632 case PrintPunctuation::NewLine:
1633 return LLVM::lookupOrCreatePrintNewlineFn(rewriter, parent,
1636 llvm_unreachable(
"unexpected punctuation");
1641 emitCall(rewriter,
printOp->getLoc(), op.value());
1649 enum class PrintConversion {
1658 LogicalResult emitScalarPrint(ConversionPatternRewriter &rewriter,
1659 ModuleOp parent, Location loc, Type printType,
1660 Value value)
const {
1661 if (typeConverter->convertType(printType) ==
nullptr)
1665 PrintConversion conversion = PrintConversion::None;
1666 FailureOr<Operation *> printer;
1668 printer = LLVM::lookupOrCreatePrintF32Fn(rewriter, parent, symbolTables);
1670 printer = LLVM::lookupOrCreatePrintF64Fn(rewriter, parent, symbolTables);
1672 conversion = PrintConversion::Bitcast16;
1673 printer = LLVM::lookupOrCreatePrintF16Fn(rewriter, parent, symbolTables);
1675 conversion = PrintConversion::Bitcast16;
1676 printer = LLVM::lookupOrCreatePrintBF16Fn(rewriter, parent, symbolTables);
1678 printer = LLVM::lookupOrCreatePrintU64Fn(rewriter, parent, symbolTables);
1679 }
else if (
auto intTy = dyn_cast<IntegerType>(printType)) {
1683 unsigned width = intTy.getWidth();
1684 if (intTy.isUnsigned()) {
1687 conversion = PrintConversion::ZeroExt64;
1689 LLVM::lookupOrCreatePrintU64Fn(rewriter, parent, symbolTables);
1694 assert(intTy.isSignless() || intTy.isSigned());
1699 conversion = PrintConversion::ZeroExt64;
1700 else if (width < 64)
1701 conversion = PrintConversion::SignExt64;
1703 LLVM::lookupOrCreatePrintI64Fn(rewriter, parent, symbolTables);
1708 }
else if (
auto floatTy = dyn_cast<FloatType>(printType)) {
1711 llvm::APFloatBase::SemanticsToEnum(floatTy.getFloatSemantics());
1712 Value semValue = LLVM::ConstantOp::create(
1713 rewriter, loc, rewriter.getI32Type(),
1714 rewriter.getIntegerAttr(rewriter.getI32Type(), sem));
1716 LLVM::ZExtOp::create(rewriter, loc, rewriter.getI64Type(), value);
1718 LLVM::lookupOrCreateApFloatPrintFn(rewriter, parent, symbolTables);
1719 emitCall(rewriter, loc, printer.value(),
1728 switch (conversion) {
1729 case PrintConversion::ZeroExt64:
1730 value = arith::ExtUIOp::create(
1731 rewriter, loc, IntegerType::get(rewriter.getContext(), 64), value);
1733 case PrintConversion::SignExt64:
1734 value = arith::ExtSIOp::create(
1735 rewriter, loc, IntegerType::get(rewriter.getContext(), 64), value);
1737 case PrintConversion::Bitcast16:
1738 value = LLVM::BitcastOp::create(
1739 rewriter, loc, IntegerType::get(rewriter.getContext(), 16), value);
1741 case PrintConversion::None:
1744 emitCall(rewriter, loc, printer.value(), value);
1749 static void emitCall(ConversionPatternRewriter &rewriter, Location loc,
1751 LLVM::CallOp::create(rewriter, loc,
TypeRange(), SymbolRefAttr::get(ref),
1759struct VectorBroadcastScalarToLowRankLowering
1761 using ConvertOpToLLVMPattern<vector::BroadcastOp>::ConvertOpToLLVMPattern;
1764 matchAndRewrite(vector::BroadcastOp
broadcast, OpAdaptor adaptor,
1765 ConversionPatternRewriter &rewriter)
const override {
1766 if (isa<VectorType>(
broadcast.getSourceType()))
1767 return rewriter.notifyMatchFailure(
1768 broadcast,
"broadcast from vector type not handled");
1771 if (resultType.getRank() > 1)
1772 return rewriter.notifyMatchFailure(
broadcast,
1773 "broadcast to 2+-d handled elsewhere");
1779 auto zero = LLVM::ConstantOp::create(
1781 typeConverter->convertType(rewriter.getIntegerType(32)),
1782 rewriter.getZeroAttr(rewriter.getIntegerType(32)));
1785 if (resultType.getRank() == 0) {
1786 rewriter.replaceOpWithNewOp<LLVM::InsertElementOp>(
1787 broadcast, vectorType, poison, adaptor.getSource(), zero);
1792 LLVM::InsertElementOp::create(rewriter,
broadcast.
getLoc(), vectorType,
1793 poison, adaptor.getSource(), zero);
1797 SmallVector<int32_t> zeroValues(width, 0);
1800 auto shuffle = rewriter.createOrFold<LLVM::ShuffleVectorOp>(
1811struct VectorBroadcastScalarToNdLowering
1813 using ConvertOpToLLVMPattern<BroadcastOp>::ConvertOpToLLVMPattern;
1816 matchAndRewrite(BroadcastOp
broadcast, OpAdaptor adaptor,
1817 ConversionPatternRewriter &rewriter)
const override {
1818 if (isa<VectorType>(
broadcast.getSourceType()))
1819 return rewriter.notifyMatchFailure(
1820 broadcast,
"broadcast from vector type not handled");
1823 if (resultType.getRank() <= 1)
1824 return rewriter.notifyMatchFailure(
1825 broadcast,
"broadcast to 1-d or 0-d handled elsewhere");
1829 auto vectorTypeInfo =
1831 auto llvmNDVectorTy = vectorTypeInfo.llvmNDVectorTy;
1832 auto llvm1DVectorTy = vectorTypeInfo.llvm1DVectorTy;
1833 if (!llvmNDVectorTy || !llvm1DVectorTy)
1837 Value desc = LLVM::PoisonOp::create(rewriter, loc, llvmNDVectorTy);
1841 Value vdesc = LLVM::PoisonOp::create(rewriter, loc, llvm1DVectorTy);
1842 auto zero = LLVM::ConstantOp::create(
1843 rewriter, loc, typeConverter->convertType(rewriter.getIntegerType(32)),
1844 rewriter.getZeroAttr(rewriter.getIntegerType(32)));
1845 Value v = LLVM::InsertElementOp::create(rewriter, loc, llvm1DVectorTy,
1846 vdesc, adaptor.getSource(), zero);
1849 int64_t width = resultType.getDimSize(resultType.getRank() - 1);
1850 SmallVector<int32_t> zeroValues(width, 0);
1851 v = LLVM::ShuffleVectorOp::create(rewriter, loc, v, v, zeroValues);
1855 nDVectorIterate(vectorTypeInfo, rewriter, [&](ArrayRef<int64_t> position) {
1856 desc = LLVM::InsertValueOp::create(rewriter, loc, desc, v, position);
1865struct VectorInterleaveOpLowering
1870 matchAndRewrite(vector::InterleaveOp interleaveOp, OpAdaptor adaptor,
1871 ConversionPatternRewriter &rewriter)
const override {
1872 VectorType resultType = interleaveOp.getResultVectorType();
1874 if (resultType.getRank() != 1)
1875 return rewriter.notifyMatchFailure(interleaveOp,
1876 "InterleaveOp not rank 1");
1878 if (resultType.isScalable()) {
1879 rewriter.replaceOpWithNewOp<LLVM::vector_interleave2>(
1880 interleaveOp, typeConverter->convertType(resultType),
1881 adaptor.getLhs(), adaptor.getRhs());
1888 int64_t resultVectorSize = resultType.getNumElements();
1889 SmallVector<int32_t> interleaveShuffleMask;
1890 interleaveShuffleMask.reserve(resultVectorSize);
1891 for (
int i = 0, end = resultVectorSize / 2; i < end; ++i) {
1892 interleaveShuffleMask.push_back(i);
1893 interleaveShuffleMask.push_back((resultVectorSize / 2) + i);
1895 rewriter.replaceOpWithNewOp<LLVM::ShuffleVectorOp>(
1896 interleaveOp, adaptor.getLhs(), adaptor.getRhs(),
1897 interleaveShuffleMask);
1904struct VectorDeinterleaveOpLowering
1909 matchAndRewrite(vector::DeinterleaveOp deinterleaveOp, OpAdaptor adaptor,
1910 ConversionPatternRewriter &rewriter)
const override {
1911 VectorType resultType = deinterleaveOp.getResultVectorType();
1912 VectorType sourceType = deinterleaveOp.getSourceVectorType();
1913 auto loc = deinterleaveOp.getLoc();
1917 if (resultType.getRank() != 1)
1918 return rewriter.notifyMatchFailure(deinterleaveOp,
1919 "DeinterleaveOp not rank 1");
1921 if (resultType.isScalable()) {
1922 const auto *llvmTypeConverter = this->getTypeConverter();
1923 auto deinterleaveResults = deinterleaveOp.getResultTypes();
1924 auto packedOpResults =
1925 llvmTypeConverter->packOperationResults(deinterleaveResults);
1926 auto intrinsic = LLVM::vector_deinterleave2::create(
1927 rewriter, loc, packedOpResults, adaptor.getSource());
1929 auto evenResult = LLVM::ExtractValueOp::create(
1930 rewriter, loc, intrinsic->getResult(0), 0);
1931 auto oddResult = LLVM::ExtractValueOp::create(rewriter, loc,
1932 intrinsic->getResult(0), 1);
1934 rewriter.replaceOp(deinterleaveOp,
ValueRange{evenResult, oddResult});
1941 int64_t resultVectorSize = resultType.getNumElements();
1942 SmallVector<int32_t> evenShuffleMask;
1943 SmallVector<int32_t> oddShuffleMask;
1945 evenShuffleMask.reserve(resultVectorSize);
1946 oddShuffleMask.reserve(resultVectorSize);
1948 for (
int i = 0; i < sourceType.getNumElements(); ++i) {
1950 evenShuffleMask.push_back(i);
1952 oddShuffleMask.push_back(i);
1955 auto poison = LLVM::PoisonOp::create(rewriter, loc, sourceType);
1956 auto evenShuffle = LLVM::ShuffleVectorOp::create(
1957 rewriter, loc, adaptor.getSource(), poison, evenShuffleMask);
1958 auto oddShuffle = LLVM::ShuffleVectorOp::create(
1959 rewriter, loc, adaptor.getSource(), poison, oddShuffleMask);
1961 rewriter.replaceOp(deinterleaveOp,
ValueRange{evenShuffle, oddShuffle});
1967struct VectorFromElementsLowering
1972 matchAndRewrite(vector::FromElementsOp fromElementsOp, OpAdaptor adaptor,
1973 ConversionPatternRewriter &rewriter)
const override {
1974 Location loc = fromElementsOp.getLoc();
1975 VectorType vectorType = fromElementsOp.getType();
1979 if (vectorType.getRank() > 1)
1980 return rewriter.notifyMatchFailure(fromElementsOp,
1981 "rank > 1 vectors are not supported");
1982 Type llvmType = typeConverter->convertType(vectorType);
1983 Type llvmIndexType = typeConverter->convertType(rewriter.getIndexType());
1984 Value
result = LLVM::PoisonOp::create(rewriter, loc, llvmType);
1985 for (
auto [idx, val] : llvm::enumerate(adaptor.getElements())) {
1987 LLVM::ConstantOp::create(rewriter, loc, llvmIndexType, idx);
1988 result = LLVM::InsertElementOp::create(rewriter, loc, llvmType,
result,
1991 rewriter.replaceOp(fromElementsOp,
result);
1997struct VectorToElementsLowering
2002 matchAndRewrite(vector::ToElementsOp toElementsOp, OpAdaptor adaptor,
2003 ConversionPatternRewriter &rewriter)
const override {
2004 Location loc = toElementsOp.getLoc();
2005 auto idxType = typeConverter->convertType(rewriter.getIndexType());
2006 Value source = adaptor.getSource();
2008 SmallVector<Value> results(toElementsOp->getNumResults());
2009 for (
auto [idx, element] : llvm::enumerate(toElementsOp.getElements())) {
2011 if (element.use_empty())
2014 auto constIdx = LLVM::ConstantOp::create(
2015 rewriter, loc, idxType, rewriter.getIntegerAttr(idxType, idx));
2016 auto llvmType = typeConverter->convertType(element.getType());
2018 Value
result = LLVM::ExtractElementOp::create(rewriter, loc, llvmType,
2023 rewriter.replaceOp(toElementsOp, results);
2033 matchAndRewrite(vector::StepOp stepOp, OpAdaptor adaptor,
2034 ConversionPatternRewriter &rewriter)
const override {
2035 Type llvmType = typeConverter->convertType(stepOp.getType());
2036 rewriter.replaceOpWithNewOp<LLVM::StepVectorOp>(stepOp, llvmType);
2051class ContractionOpToMatmulOpLowering
2054 using MaskableOpRewritePattern::MaskableOpRewritePattern;
2056 ContractionOpToMatmulOpLowering(MLIRContext *context,
2057 PatternBenefit benefit = 100)
2058 : MaskableOpRewritePattern<vector::ContractionOp>(context, benefit) {}
2061 matchAndRewriteMaskableOp(vector::ContractionOp op, MaskingOpInterface maskOp,
2062 PatternRewriter &rewriter)
const override;
2082FailureOr<Value> ContractionOpToMatmulOpLowering::matchAndRewriteMaskableOp(
2083 vector::ContractionOp op, MaskingOpInterface maskOp,
2089 auto iteratorTypes = op.getIteratorTypes().getValue();
2095 Type opResType = op.getType();
2096 VectorType vecType = dyn_cast<VectorType>(opResType);
2097 if (vecType && vecType.isScalable()) {
2102 Type elementType = op.getLhsType().getElementType();
2106 Type dstElementType = vecType ? vecType.getElementType() : opResType;
2107 if (elementType != dstElementType)
2112 MLIRContext *ctx = op.getContext();
2113 Location loc = op.getLoc();
2117 Value
lhs = op.getLhs();
2118 auto lhsMap = op.getIndexingMapsArray()[0];
2120 lhs = vector::TransposeOp::create(rew, loc,
lhs, ArrayRef<int64_t>{1, 0});
2125 Value
rhs = op.getRhs();
2126 auto rhsMap = op.getIndexingMapsArray()[1];
2128 rhs = vector::TransposeOp::create(rew, loc,
rhs, ArrayRef<int64_t>{1, 0});
2133 VectorType lhsType = cast<VectorType>(
lhs.getType());
2134 VectorType rhsType = cast<VectorType>(
rhs.getType());
2135 int64_t lhsRows = lhsType.getDimSize(0);
2136 int64_t lhsColumns = lhsType.getDimSize(1);
2137 int64_t rhsColumns = rhsType.getDimSize(1);
2139 Type flattenedLHSType =
2140 VectorType::get(lhsType.getNumElements(), lhsType.getElementType());
2141 lhs = vector::ShapeCastOp::create(rew, loc, flattenedLHSType,
lhs);
2143 Type flattenedRHSType =
2144 VectorType::get(rhsType.getNumElements(), rhsType.getElementType());
2145 rhs = vector::ShapeCastOp::create(rew, loc, flattenedRHSType,
rhs);
2147 Value
mul = LLVM::MatrixMultiplyOp::create(
2149 VectorType::get(lhsRows * rhsColumns,
2150 cast<VectorType>(
lhs.getType()).getElementType()),
2151 lhs,
rhs, lhsRows, lhsColumns, rhsColumns);
2153 mul = vector::ShapeCastOp::create(
2155 VectorType::get({lhsRows, rhsColumns},
2160 auto accMap = op.getIndexingMapsArray()[2];
2162 mul = vector::TransposeOp::create(rew, loc,
mul, ArrayRef<int64_t>{1, 0});
2164 llvm_unreachable(
"invalid contraction semantics");
2166 Value res = isa<IntegerType>(elementType)
2167 ?
static_cast<Value
>(
2168 arith::AddIOp::create(rew, loc, op.getAcc(),
mul))
2169 : static_cast<Value>(
2170 arith::AddFOp::create(rew, loc, op.getAcc(),
mul));
2188class TransposeOpToMatrixTransposeOpLowering
2189 :
public OpRewritePattern<vector::TransposeOp> {
2193 LogicalResult matchAndRewrite(vector::TransposeOp op,
2194 PatternRewriter &rewriter)
const override {
2195 auto loc = op.getLoc();
2197 Value input = op.getVector();
2198 VectorType inputType = op.getSourceVectorType();
2199 VectorType resType = op.getResultVectorType();
2201 if (inputType.isScalable())
2203 op,
"This lowering does not support scalable vectors");
2206 ArrayRef<int64_t> transp = op.getPermutation();
2208 if (resType.getRank() != 2 || transp[0] != 1 || transp[1] != 0) {
2212 Type flattenedType =
2213 VectorType::get(resType.getNumElements(), resType.getElementType());
2215 vector::ShapeCastOp::create(rewriter, loc, flattenedType, input);
2218 Value trans = LLVM::MatrixTransposeOp::create(rewriter, loc, flattenedType,
2219 matrix, rows, columns);
2229 patterns.
add<VectorFMAOpNDRewritePattern>(patterns.
getContext());
2234 patterns.
add<ContractionOpToMatmulOpLowering>(patterns.
getContext(), benefit);
2239 patterns.
add<TransposeOpToMatrixTransposeOpLowering>(patterns.
getContext(),
2246 bool reassociateFPReductions,
bool force32BitVectorIndices,
2247 bool useVectorAlignment,
bool enableGEPInboundsNuw) {
2250 patterns.
add<VectorReductionOpConversion>(converter, reassociateFPReductions);
2251 patterns.
add<VectorCreateMaskOpConversion>(ctx, force32BitVectorIndices);
2252 patterns.
add<VectorLoadStoreConversion<vector::LoadOp>,
2253 VectorLoadStoreConversion<vector::MaskedLoadOp>,
2254 VectorLoadStoreConversion<vector::StoreOp>,
2255 VectorLoadStoreConversion<vector::MaskedStoreOp>>(
2256 converter, useVectorAlignment, enableGEPInboundsNuw);
2257 patterns.
add<VectorGatherOpConversion, VectorScatterOpConversion>(
2258 converter, useVectorAlignment);
2259 patterns.
add<VectorBitCastOpConversion, VectorShuffleOpConversion,
2260 VectorExtractOpConversion, VectorFMAOp1DConversion,
2261 VectorInsertOpConversion, VectorPrintOpConversion,
2262 VectorTypeCastOpConversion, VectorScaleOpConversion,
2263 VectorExpandLoadOpConversion, VectorCompressStoreOpConversion,
2264 VectorBroadcastScalarToLowRankLowering,
2265 VectorBroadcastScalarToNdLowering,
2266 VectorScalableInsertOpLowering, VectorScalableExtractOpLowering,
2267 MaskedReductionOpConversion, VectorInterleaveOpLowering,
2268 VectorDeinterleaveOpLowering, VectorFromElementsLowering,
2269 VectorToElementsLowering, VectorStepOpLowering>(converter);
2273struct VectorToLLVMDialectInterface :
public ConvertToLLVMPatternInterface {
2274 VectorToLLVMDialectInterface(
Dialect *dialect)
2275 : ConvertToLLVMPatternInterface(dialect) {}
2277 using ConvertToLLVMPatternInterface::ConvertToLLVMPatternInterface;
2278 void loadDependentDialects(MLIRContext *context)
const final {
2279 context->loadDialect<LLVM::LLVMDialect>();
2284 void populateConvertToLLVMConversionPatterns(
2285 ConversionTarget &
target, LLVMTypeConverter &typeConverter,
2286 RewritePatternSet &patterns)
const final {
2295 dialect->addInterfaces<VectorToLLVMDialectInterface>();
static Value getIndexedPtrs(ConversionPatternRewriter &rewriter, Location loc, const LLVMTypeConverter &typeConverter, MemRefType memRefType, Value llvmMemref, Value base, Value index, VectorType vectorType)
LogicalResult getVectorToLLVMAlignment(const LLVMTypeConverter &typeConverter, VectorType vectorType, MemRefType memrefType, unsigned &align, bool useVectorAlignment)
LogicalResult getVectorAlignment(const LLVMTypeConverter &typeConverter, VectorType vectorType, unsigned &align)
LogicalResult getMemRefAlignment(const LLVMTypeConverter &typeConverter, MemRefType memrefType, unsigned &align)
static Value extractOne(ConversionPatternRewriter &rewriter, const LLVMTypeConverter &typeConverter, Location loc, Value val, Type llvmType, int64_t rank, int64_t pos)
static Value insertOne(ConversionPatternRewriter &rewriter, const LLVMTypeConverter &typeConverter, Location loc, Value val1, Value val2, Type llvmType, int64_t rank, int64_t pos)
static Value getAsLLVMValue(OpBuilder &builder, Location loc, OpFoldResult foldResult)
Convert foldResult into a Value.
static LogicalResult isMemRefTypeSupported(MemRefType memRefType, const LLVMTypeConverter &converter)
LogicalResult initialize(unsigned origNumLoops, ArrayRef< ReassociationIndices > foldedIterationDims)
static Value broadcast(Location loc, Value toBroadcast, unsigned numElements, const TypeConverter &typeConverter, ConversionPatternRewriter &rewriter)
Broadcasts the value to vector with numElements number of elements.
static void printOp(llvm::raw_ostream &os, Operation *op, OpPrintingFlags &flags)
static AffineMap get(MLIRContext *context)
Returns a zero result affine map with no dimensions or symbols: () -> ().
IntegerAttr getI32IntegerAttr(int32_t value)
TypedAttr getZeroAttr(Type type)
MLIRContext * getContext() const
Utility class for operation conversions targeting the LLVM dialect that match exactly one source oper...
ConvertOpToLLVMPattern(const LLVMTypeConverter &typeConverter, PatternBenefit benefit=1)
The DialectRegistry maps a dialect namespace to a constructor for the matching dialect.
bool addExtension(TypeID extensionID, std::unique_ptr< DialectExtensionBase > extension)
Add the given extension to the registry.
Dialects are groups of MLIR operations, types and attributes, as well as behavior associated with the...
Conversion from types to the LLVM IR dialect.
const llvm::DataLayout & getDataLayout() const
Returns the data layout to use during and after conversion.
FailureOr< unsigned > getMemRefAddressSpace(BaseMemRefType type) const
Return the LLVM address space corresponding to the memory space of the memref type type or failure if...
LLVM::LLVMDialect * getDialect() const
Returns the LLVM dialect.
Utility class to translate MLIR LLVM dialect types to LLVM IR.
unsigned getPreferredAlignment(Type type, const llvm::DataLayout &layout)
Returns the preferred alignment for the type given the data layout.
This class defines the main interface for locations in MLIR and acts as a non-nullable wrapper around...
MLIRContext is the top-level object for a collection of MLIR operations.
Helper class to produce LLVM dialect operations extracting or inserting elements of a MemRef descript...
LLVM::LLVMPointerType getElementPtrType()
Returns the (LLVM) pointer type this descriptor contains.
Generic implementation of one-to-one conversion from "SourceOp" to "TargetOp" where the latter belong...
This class helps build Operations.
This class represents a single result from folding an operation.
This class represents the benefit of a pattern match in a unitless scheme that ranges from 0 (very li...
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.
virtual void replaceOp(Operation *op, ValueRange newValues)
Replace the results of the given (original) operation with the specified list of values (replacements...
std::enable_if_t<!std::is_convertible< CallbackT, Twine >::value, LogicalResult > notifyMatchFailure(Location loc, CallbackT &&reasonCallback)
Used to notify the listener that the IR failed to be rewritten because of a match failure,...
OpTy replaceOpWithNewOp(Operation *op, Args &&...args)
Replace the results of the given (original) op with a new op that is created without verification (re...
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
bool isIntOrIndex() const
Return true if this is an integer (of any signedness) or an index type.
bool isIntOrFloat() const
Return true if this is an integer (of any signedness) or a float type.
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.
Location getLoc() const
Return the location of this value.
static ConstantIndexOp create(OpBuilder &builder, Location location, int64_t value)
void printType(Type type, AsmPrinter &printer)
Prints an LLVM Dialect type.
void nDVectorIterate(const NDVectorTypeInfo &info, OpBuilder &builder, function_ref< void(ArrayRef< int64_t >)> fun)
NDVectorTypeInfo extractNDVectorTypeInfo(VectorType vectorType, const LLVMTypeConverter &converter)
Value getStridedElementPtr(OpBuilder &builder, Location loc, const LLVMTypeConverter &converter, MemRefType type, Value memRefDesc, ValueRange indices, LLVM::GEPNoWrapFlags noWrapFlags=LLVM::GEPNoWrapFlags::none)
Performs the index computation to get to the element at indices of the memory pointed to by memRefDes...
Type getVectorType(Type elementType, unsigned numElements, bool isScalable=false)
Creates an LLVM dialect-compatible vector type with the given element type and length.
LLVM::FastmathFlags convertArithFastMathFlagsToLLVM(arith::FastMathFlags arithFMF)
Maps arithmetic fastmath enum values to LLVM enum values.
bool hasNegativeStaticStride(MemRefType memRefTy)
Returns true if any stride of memRefTy is statically known to be negative.
bool isReductionIterator(Attribute attr)
Returns true if attr has "reduction" iterator type semantics.
void populateVectorContractToMatrixMultiply(RewritePatternSet &patterns, PatternBenefit benefit=100)
Populate the pattern set with the following patterns:
void populateVectorRankReducingFMAPattern(RewritePatternSet &patterns)
Populates a pattern that rank-reduces n-D FMAs into (n-1)-D FMAs where n > 1.
bool isParallelIterator(Attribute attr)
Returns true if attr has "parallel" iterator type semantics.
void registerConvertVectorToLLVMInterface(DialectRegistry ®istry)
SmallVector< int64_t > getAsIntegers(ArrayRef< Value > values)
Returns the integer numbers in values.
void populateVectorTransposeToFlatTranspose(RewritePatternSet &patterns, PatternBenefit benefit=100)
Populate the pattern set with the following patterns:
Value createReductionNeutralValue(OpBuilder &builder, Location loc, Type type, vector::CombiningKind kind)
Creates a constant filled with the neutral (identity) value for the given reduction kind.
Include the generated interface declarations.
SmallVector< OpFoldResult > getMixedValues(ArrayRef< int64_t > staticValues, ValueRange dynamicValues, MLIRContext *context)
Return a vector of OpFoldResults with the same size a staticValues, but all elements for which Shaped...
void populateVectorToLLVMConversionPatterns(const LLVMTypeConverter &converter, RewritePatternSet &patterns, bool reassociateFPReductions=false, bool force32BitVectorIndices=false, bool useVectorAlignment=false, bool enableGEPInboundsNuw=false)
Collect a set of patterns to convert from the Vector dialect to LLVM.
void bindDims(MLIRContext *ctx, AffineExprTy &...exprs)
Bind a list of AffineExpr references to DimExpr at positions: [0 .
Value getValueOrCreateCastToIndexLike(OpBuilder &b, Location loc, Type targetType, Value value)
Create a cast from an index-like value (index or integer) to another index-like value.
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...
A pattern for ops that implement MaskableOpInterface and that might be masked (i.e.