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,
761template <
class LLVMVPRedIntrinOp,
class ReductionNeutral>
763lowerPredicatedReductionWithStartValue(ConversionPatternRewriter &rewriter,
766 accumulator = getOrCreateAccumulator<ReductionNeutral>(rewriter, loc,
767 llvmType, accumulator);
768 return LLVMVPRedIntrinOp::create(rewriter, loc, llvmType,
769 accumulator, vectorOperand);
772template <
class LLVMVPRedIntrinOp,
class ReductionNeutral>
773static Value lowerPredicatedReductionWithStartValue(
774 ConversionPatternRewriter &rewriter,
Location loc,
Type llvmType,
776 accumulator = getOrCreateAccumulator<ReductionNeutral>(rewriter, loc,
777 llvmType, accumulator);
779 createVectorLengthValue(rewriter, loc, vectorOperand.
getType());
780 return LLVMVPRedIntrinOp::create(rewriter, loc, llvmType,
781 accumulator, vectorOperand,
785template <
class LLVMIntVPRedIntrinOp,
class IntReductionNeutral,
786 class LLVMFPVPRedIntrinOp,
class FPReductionNeutral>
787static Value lowerPredicatedReductionWithStartValue(
788 ConversionPatternRewriter &rewriter,
Location loc,
Type llvmType,
791 return lowerPredicatedReductionWithStartValue<LLVMIntVPRedIntrinOp,
792 IntReductionNeutral>(
793 rewriter, loc, llvmType, vectorOperand, accumulator, mask);
796 return lowerPredicatedReductionWithStartValue<LLVMFPVPRedIntrinOp,
798 rewriter, loc, llvmType, vectorOperand, accumulator, mask);
802class VectorReductionOpConversion
805 explicit VectorReductionOpConversion(
const LLVMTypeConverter &typeConv,
806 bool reassociateFPRed)
807 : ConvertOpToLLVMPattern<vector::ReductionOp>(typeConv),
808 reassociateFPReductions(reassociateFPRed) {}
811 matchAndRewrite(vector::ReductionOp reductionOp, OpAdaptor adaptor,
812 ConversionPatternRewriter &rewriter)
const override {
813 auto kind = reductionOp.getKind();
814 Type eltType = reductionOp.getDest().getType();
815 Type llvmType = typeConverter->convertType(eltType);
816 Value operand = adaptor.getVector();
817 Value acc = adaptor.getAcc();
818 Location loc = reductionOp.getLoc();
824 case vector::CombiningKind::ADD:
826 createIntegerReductionArithmeticOpLowering<LLVM::vector_reduce_add,
828 rewriter, loc, llvmType, operand, acc);
830 case vector::CombiningKind::MUL:
832 createIntegerReductionArithmeticOpLowering<LLVM::vector_reduce_mul,
834 rewriter, loc, llvmType, operand, acc);
836 case vector::CombiningKind::MINUI:
837 result = createIntegerReductionComparisonOpLowering<
838 LLVM::vector_reduce_umin>(rewriter, loc, llvmType, operand, acc,
839 LLVM::ICmpPredicate::ule);
841 case vector::CombiningKind::MINSI:
842 result = createIntegerReductionComparisonOpLowering<
843 LLVM::vector_reduce_smin>(rewriter, loc, llvmType, operand, acc,
844 LLVM::ICmpPredicate::sle);
846 case vector::CombiningKind::MAXUI:
847 result = createIntegerReductionComparisonOpLowering<
848 LLVM::vector_reduce_umax>(rewriter, loc, llvmType, operand, acc,
849 LLVM::ICmpPredicate::uge);
851 case vector::CombiningKind::MAXSI:
852 result = createIntegerReductionComparisonOpLowering<
853 LLVM::vector_reduce_smax>(rewriter, loc, llvmType, operand, acc,
854 LLVM::ICmpPredicate::sge);
856 case vector::CombiningKind::AND:
858 createIntegerReductionArithmeticOpLowering<LLVM::vector_reduce_and,
860 rewriter, loc, llvmType, operand, acc);
862 case vector::CombiningKind::OR:
864 createIntegerReductionArithmeticOpLowering<LLVM::vector_reduce_or,
866 rewriter, loc, llvmType, operand, acc);
868 case vector::CombiningKind::XOR:
870 createIntegerReductionArithmeticOpLowering<LLVM::vector_reduce_xor,
872 rewriter, loc, llvmType, operand, acc);
877 rewriter.replaceOp(reductionOp,
result);
882 if (!isa<FloatType>(eltType))
885 arith::FastMathFlagsAttr fMFAttr = reductionOp.getFastMathFlagsAttr();
886 LLVM::FastmathFlagsAttr fmf = LLVM::FastmathFlagsAttr::get(
887 reductionOp.getContext(),
889 fmf = LLVM::FastmathFlagsAttr::get(
890 reductionOp.getContext(),
891 fmf.getValue() | (reassociateFPReductions ? LLVM::FastmathFlags::reassoc
892 : LLVM::FastmathFlags::none));
896 if (kind == vector::CombiningKind::ADD) {
897 result = lowerReductionWithStartValue<LLVM::vector_reduce_fadd,
898 ReductionNeutralZero>(
899 rewriter, loc, llvmType, operand, acc, fmf);
900 }
else if (kind == vector::CombiningKind::MUL) {
901 result = lowerReductionWithStartValue<LLVM::vector_reduce_fmul,
902 ReductionNeutralFPOne>(
903 rewriter, loc, llvmType, operand, acc, fmf);
904 }
else if (kind == vector::CombiningKind::MINIMUMF) {
906 createFPReductionComparisonOpLowering<LLVM::vector_reduce_fminimum>(
907 rewriter, loc, llvmType, operand, acc, fmf);
908 }
else if (kind == vector::CombiningKind::MAXIMUMF) {
910 createFPReductionComparisonOpLowering<LLVM::vector_reduce_fmaximum>(
911 rewriter, loc, llvmType, operand, acc, fmf);
912 }
else if (kind == vector::CombiningKind::MINNUMF) {
913 result = createFPReductionComparisonOpLowering<LLVM::vector_reduce_fmin>(
914 rewriter, loc, llvmType, operand, acc, fmf);
915 }
else if (kind == vector::CombiningKind::MAXNUMF) {
916 result = createFPReductionComparisonOpLowering<LLVM::vector_reduce_fmax>(
917 rewriter, loc, llvmType, operand, acc, fmf);
922 rewriter.replaceOp(reductionOp,
result);
927 const bool reassociateFPReductions;
938template <
class MaskedOp>
939class VectorMaskOpConversionBase
942 using ConvertOpToLLVMPattern<vector::MaskOp>::ConvertOpToLLVMPattern;
945 matchAndRewrite(vector::MaskOp maskOp, OpAdaptor adaptor,
946 ConversionPatternRewriter &rewriter)
const final {
948 auto maskedOp = llvm::dyn_cast_or_null<MaskedOp>(maskOp.getMaskableOp());
951 return matchAndRewriteMaskableOp(maskOp, maskedOp, rewriter);
955 virtual LogicalResult
956 matchAndRewriteMaskableOp(vector::MaskOp maskOp,
957 vector::MaskableOpInterface maskableOp,
958 ConversionPatternRewriter &rewriter)
const = 0;
961class MaskedReductionOpConversion
962 :
public VectorMaskOpConversionBase<vector::ReductionOp> {
965 using VectorMaskOpConversionBase<
966 vector::ReductionOp>::VectorMaskOpConversionBase;
968 LogicalResult matchAndRewriteMaskableOp(
969 vector::MaskOp maskOp, MaskableOpInterface maskableOp,
970 ConversionPatternRewriter &rewriter)
const override {
971 auto reductionOp = cast<ReductionOp>(maskableOp.getOperation());
972 auto kind = reductionOp.getKind();
973 Type eltType = reductionOp.getDest().getType();
974 Type llvmType = typeConverter->convertType(eltType);
975 Value operand = reductionOp.getVector();
976 Value acc = reductionOp.getAcc();
977 Location loc = reductionOp.getLoc();
979 arith::FastMathFlagsAttr fMFAttr = reductionOp.getFastMathFlagsAttr();
980 LLVM::FastmathFlagsAttr fmf = LLVM::FastmathFlagsAttr::get(
981 reductionOp.getContext(),
984 LLVM::bitEnumContainsAny(fmf.getValue(), LLVM::FastmathFlags::ninf);
988 case vector::CombiningKind::ADD:
989 result = lowerPredicatedReductionWithStartValue<
990 LLVM::VPReduceAddOp, ReductionNeutralZero, LLVM::VPReduceFAddOp,
991 ReductionNeutralZero>(rewriter, loc, llvmType, operand, acc,
994 case vector::CombiningKind::MUL:
995 result = lowerPredicatedReductionWithStartValue<
996 LLVM::VPReduceMulOp, ReductionNeutralIntOne, LLVM::VPReduceFMulOp,
997 ReductionNeutralFPOne>(rewriter, loc, llvmType, operand, acc,
1000 case vector::CombiningKind::MINUI:
1001 result = lowerPredicatedReductionWithStartValue<LLVM::VPReduceUMinOp,
1002 ReductionNeutralUIntMax>(
1003 rewriter, loc, llvmType, operand, acc, maskOp.getMask());
1005 case vector::CombiningKind::MINSI:
1006 result = lowerPredicatedReductionWithStartValue<LLVM::VPReduceSMinOp,
1007 ReductionNeutralSIntMax>(
1008 rewriter, loc, llvmType, operand, acc, maskOp.getMask());
1010 case vector::CombiningKind::MAXUI:
1011 result = lowerPredicatedReductionWithStartValue<LLVM::VPReduceUMaxOp,
1012 ReductionNeutralUIntMin>(
1013 rewriter, loc, llvmType, operand, acc, maskOp.getMask());
1015 case vector::CombiningKind::MAXSI:
1016 result = lowerPredicatedReductionWithStartValue<LLVM::VPReduceSMaxOp,
1017 ReductionNeutralSIntMin>(
1018 rewriter, loc, llvmType, operand, acc, maskOp.getMask());
1020 case vector::CombiningKind::AND:
1021 result = lowerPredicatedReductionWithStartValue<LLVM::VPReduceAndOp,
1022 ReductionNeutralAllOnes>(
1023 rewriter, loc, llvmType, operand, acc, maskOp.getMask());
1025 case vector::CombiningKind::OR:
1026 result = lowerPredicatedReductionWithStartValue<LLVM::VPReduceOrOp,
1027 ReductionNeutralZero>(
1028 rewriter, loc, llvmType, operand, acc, maskOp.getMask());
1030 case vector::CombiningKind::XOR:
1031 result = lowerPredicatedReductionWithStartValue<LLVM::VPReduceXorOp,
1032 ReductionNeutralZero>(
1033 rewriter, loc, llvmType, operand, acc, maskOp.getMask());
1035 case vector::CombiningKind::MINNUMF:
1037 lowerPredicatedReductionWithStartValue<LLVM::VPReduceFMinOp,
1038 ReductionNeutralFPNegQNaN>(
1039 rewriter, loc, llvmType, operand, acc, maskOp.getMask());
1041 case vector::CombiningKind::MAXNUMF:
1042 result = lowerPredicatedReductionWithStartValue<LLVM::VPReduceFMaxOp,
1043 ReductionNeutralFPQNaN>(
1044 rewriter, loc, llvmType, operand, acc, maskOp.getMask());
1046 case CombiningKind::MAXIMUMF:
1051 ? lowerPredicatedReductionWithStartValue<
1052 LLVM::VPReduceFMaximumOp, ReductionNeutralFPLowestFinite>(
1053 rewriter, loc, llvmType, operand, acc, maskOp.getMask())
1054 : lowerPredicatedReductionWithStartValue<
1055 LLVM::VPReduceFMaximumOp, ReductionNeutralFPNegInf>(
1056 rewriter, loc, llvmType, operand, acc, maskOp.getMask());
1058 case CombiningKind::MINIMUMF:
1061 ? lowerPredicatedReductionWithStartValue<
1062 LLVM::VPReduceFMinimumOp, ReductionNeutralFPLargestFinite>(
1063 rewriter, loc, llvmType, operand, acc, maskOp.getMask())
1064 : lowerPredicatedReductionWithStartValue<
1065 LLVM::VPReduceFMinimumOp, ReductionNeutralFPPosInf>(
1066 rewriter, loc, llvmType, operand, acc, maskOp.getMask());
1071 rewriter.replaceOp(maskOp,
result);
1076class VectorShuffleOpConversion
1079 using ConvertOpToLLVMPattern<vector::ShuffleOp>::ConvertOpToLLVMPattern;
1082 matchAndRewrite(vector::ShuffleOp shuffleOp, OpAdaptor adaptor,
1083 ConversionPatternRewriter &rewriter)
const override {
1084 auto loc = shuffleOp->getLoc();
1085 auto v1Type = shuffleOp.getV1VectorType();
1086 auto v2Type = shuffleOp.getV2VectorType();
1087 auto vectorType = shuffleOp.getResultVectorType();
1088 Type llvmType = typeConverter->convertType(vectorType);
1089 ArrayRef<int64_t> mask = shuffleOp.getMask();
1096 int64_t rank = vectorType.getRank();
1098 bool wellFormed0DCase =
1099 v1Type.getRank() == 0 && v2Type.getRank() == 0 && rank == 1;
1100 bool wellFormedNDCase =
1101 v1Type.getRank() == rank && v2Type.getRank() == rank;
1102 assert((wellFormed0DCase || wellFormedNDCase) &&
"op is not well-formed");
1107 if (rank <= 1 && v1Type == v2Type) {
1108 Value llvmShuffleOp = LLVM::ShuffleVectorOp::create(
1109 rewriter, loc, adaptor.getV1(), adaptor.getV2(),
1110 llvm::to_vector_of<int32_t>(mask));
1111 rewriter.replaceOp(shuffleOp, llvmShuffleOp);
1116 int64_t v1Dim = v1Type.getDimSize(0);
1118 if (
auto arrayType = dyn_cast<LLVM::LLVMArrayType>(llvmType))
1119 eltType = arrayType.getElementType();
1121 eltType = cast<VectorType>(llvmType).getElementType();
1122 Value insert = LLVM::PoisonOp::create(rewriter, loc, llvmType);
1124 for (int64_t extPos : mask) {
1125 Value value = adaptor.getV1();
1126 if (extPos >= v1Dim) {
1128 value = adaptor.getV2();
1130 Value extract =
extractOne(rewriter, *getTypeConverter(), loc, value,
1131 eltType, rank, extPos);
1132 insert =
insertOne(rewriter, *getTypeConverter(), loc, insert, extract,
1133 llvmType, rank, insPos++);
1135 rewriter.replaceOp(shuffleOp, insert);
1140class VectorExtractOpConversion
1143 using ConvertOpToLLVMPattern<vector::ExtractOp>::ConvertOpToLLVMPattern;
1146 matchAndRewrite(vector::ExtractOp extractOp, OpAdaptor adaptor,
1147 ConversionPatternRewriter &rewriter)
const override {
1148 auto loc = extractOp->getLoc();
1149 auto resultType = extractOp.getResult().getType();
1150 auto llvmResultType = typeConverter->convertType(resultType);
1152 if (!llvmResultType)
1156 adaptor.getStaticPosition(), adaptor.getDynamicPosition(), rewriter);
1170 bool extractsAggregate = extractOp.getSourceVectorType().getRank() >= 2;
1174 bool extractsScalar =
static_cast<int64_t
>(positionVec.size()) ==
1175 extractOp.getSourceVectorType().getRank();
1179 if (extractOp.getSourceVectorType().getRank() == 0) {
1180 Type idxType = typeConverter->convertType(rewriter.getIndexType());
1181 positionVec.push_back(rewriter.getZeroAttr(idxType));
1184 Value extracted = adaptor.getSource();
1185 if (extractsAggregate) {
1186 ArrayRef<OpFoldResult> position(positionVec);
1187 if (extractsScalar) {
1191 position = position.drop_back();
1194 if (!llvm::all_of(position, llvm::IsaPred<Attribute>)) {
1197 extracted = LLVM::ExtractValueOp::create(rewriter, loc, extracted,
1201 if (extractsScalar) {
1202 extracted = LLVM::ExtractElementOp::create(
1203 rewriter, loc, extracted,
1207 rewriter.replaceOp(extractOp, extracted);
1228 using ConvertOpToLLVMPattern<vector::FMAOp>::ConvertOpToLLVMPattern;
1231 matchAndRewrite(vector::FMAOp fmaOp, OpAdaptor adaptor,
1232 ConversionPatternRewriter &rewriter)
const override {
1233 VectorType vType = fmaOp.getVectorType();
1234 if (vType.getRank() > 1)
1237 rewriter.replaceOpWithNewOp<LLVM::FMulAddOp>(
1238 fmaOp, adaptor.getLhs(), adaptor.getRhs(), adaptor.getAcc());
1243class VectorInsertOpConversion
1246 using ConvertOpToLLVMPattern<vector::InsertOp>::ConvertOpToLLVMPattern;
1249 matchAndRewrite(vector::InsertOp insertOp, OpAdaptor adaptor,
1250 ConversionPatternRewriter &rewriter)
const override {
1251 auto loc = insertOp->getLoc();
1252 auto destVectorType = insertOp.getDestVectorType();
1253 auto llvmResultType = typeConverter->convertType(destVectorType);
1255 if (!llvmResultType)
1259 adaptor.getStaticPosition(), adaptor.getDynamicPosition(), rewriter);
1281 bool isNestedAggregate = isa<LLVM::LLVMArrayType>(llvmResultType);
1283 bool insertIntoInnermostDim =
1284 static_cast<int64_t
>(positionVec.size()) == destVectorType.getRank();
1286 ArrayRef<OpFoldResult> positionOf1DVectorWithinAggregate(
1287 positionVec.begin(),
1288 insertIntoInnermostDim ? positionVec.size() - 1 : positionVec.size());
1289 OpFoldResult positionOfScalarWithin1DVector;
1290 if (destVectorType.getRank() == 0) {
1293 Type idxType = typeConverter->convertType(rewriter.getIndexType());
1294 positionOfScalarWithin1DVector = rewriter.getZeroAttr(idxType);
1295 }
else if (insertIntoInnermostDim) {
1296 positionOfScalarWithin1DVector = positionVec.back();
1302 Value sourceAggregate = adaptor.getValueToStore();
1303 if (insertIntoInnermostDim) {
1306 if (isNestedAggregate) {
1309 if (!llvm::all_of(positionOf1DVectorWithinAggregate,
1310 llvm::IsaPred<Attribute>)) {
1314 sourceAggregate = LLVM::ExtractValueOp::create(
1315 rewriter, loc, adaptor.getDest(),
1320 sourceAggregate = adaptor.getDest();
1323 sourceAggregate = LLVM::InsertElementOp::create(
1324 rewriter, loc, sourceAggregate.
getType(), sourceAggregate,
1325 adaptor.getValueToStore(),
1329 Value
result = sourceAggregate;
1330 if (isNestedAggregate) {
1331 if (!llvm::all_of(positionOf1DVectorWithinAggregate,
1332 llvm::IsaPred<Attribute>)) {
1336 result = LLVM::InsertValueOp::create(
1337 rewriter, loc, adaptor.getDest(), sourceAggregate,
1341 rewriter.replaceOp(insertOp,
result);
1347struct VectorScalableInsertOpLowering
1349 using ConvertOpToLLVMPattern<
1350 vector::ScalableInsertOp>::ConvertOpToLLVMPattern;
1353 matchAndRewrite(vector::ScalableInsertOp insOp, OpAdaptor adaptor,
1354 ConversionPatternRewriter &rewriter)
const override {
1355 rewriter.replaceOpWithNewOp<LLVM::vector_insert>(
1356 insOp, adaptor.getDest(), adaptor.getValueToStore(), adaptor.getPos());
1362struct VectorScalableExtractOpLowering
1364 using ConvertOpToLLVMPattern<
1365 vector::ScalableExtractOp>::ConvertOpToLLVMPattern;
1368 matchAndRewrite(vector::ScalableExtractOp extOp, OpAdaptor adaptor,
1369 ConversionPatternRewriter &rewriter)
const override {
1370 rewriter.replaceOpWithNewOp<LLVM::vector_extract>(
1371 extOp, typeConverter->convertType(extOp.getResultVectorType()),
1372 adaptor.getSource(), adaptor.getPos());
1405 setHasBoundedRewriteRecursion();
1408 LogicalResult matchAndRewrite(FMAOp op,
1409 PatternRewriter &rewriter)
const override {
1410 auto vType = op.getVectorType();
1411 if (vType.getRank() < 2)
1414 auto loc = op.getLoc();
1415 auto elemType = vType.getElementType();
1416 Value zero = arith::ConstantOp::create(rewriter, loc, elemType,
1418 Value desc = vector::BroadcastOp::create(rewriter, loc, vType, zero);
1419 for (int64_t i = 0, e = vType.getShape().front(); i != e; ++i) {
1420 Value extrLHS = ExtractOp::create(rewriter, loc, op.getLhs(), i);
1421 Value extrRHS = ExtractOp::create(rewriter, loc, op.getRhs(), i);
1422 Value extrACC = ExtractOp::create(rewriter, loc, op.getAcc(), i);
1423 Value fma = FMAOp::create(rewriter, loc, extrLHS, extrRHS, extrACC);
1424 desc = InsertOp::create(rewriter, loc, fma, desc, i);
1433static std::optional<SmallVector<int64_t, 4>>
1434computeContiguousStrides(MemRefType memRefType) {
1437 if (
failed(memRefType.getStridesAndOffset(strides, offset)))
1438 return std::nullopt;
1439 if (!strides.empty() && strides.back() != 1)
1440 return std::nullopt;
1442 if (memRefType.getLayout().isIdentity())
1449 auto sizes = memRefType.getShape();
1451 if (ShapedType::isDynamic(sizes[
index + 1]) ||
1452 ShapedType::isDynamic(strides[
index]) ||
1453 ShapedType::isDynamic(strides[
index + 1]))
1454 return std::nullopt;
1456 return std::nullopt;
1461class VectorTypeCastOpConversion
1464 using ConvertOpToLLVMPattern<vector::TypeCastOp>::ConvertOpToLLVMPattern;
1467 matchAndRewrite(vector::TypeCastOp castOp, OpAdaptor adaptor,
1468 ConversionPatternRewriter &rewriter)
const override {
1469 auto loc = castOp->getLoc();
1470 MemRefType sourceMemRefType =
1471 cast<MemRefType>(castOp.getOperand().getType());
1472 MemRefType targetMemRefType = castOp.getType();
1475 if (!sourceMemRefType.hasStaticShape() ||
1476 !targetMemRefType.hasStaticShape())
1479 auto llvmSourceDescriptorTy =
1480 dyn_cast<LLVM::LLVMStructType>(adaptor.getOperands()[0].getType());
1481 if (!llvmSourceDescriptorTy)
1483 MemRefDescriptor sourceMemRef(adaptor.getOperands()[0]);
1485 auto llvmTargetDescriptorTy = dyn_cast_or_null<LLVM::LLVMStructType>(
1486 typeConverter->convertType(targetMemRefType));
1487 if (!llvmTargetDescriptorTy)
1491 auto sourceStrides = computeContiguousStrides(sourceMemRefType);
1494 auto targetStrides = computeContiguousStrides(targetMemRefType);
1498 if (llvm::any_of(*targetStrides, ShapedType::isDynamic))
1503 Type indexTy = getTypeConverter()->getIndexType();
1506 auto desc = MemRefDescriptor::poison(rewriter, loc, llvmTargetDescriptorTy);
1508 Value allocated = sourceMemRef.allocatedPtr(rewriter, loc);
1509 desc.setAllocatedPtr(rewriter, loc, allocated);
1512 Value ptr = sourceMemRef.alignedPtr(rewriter, loc);
1513 desc.setAlignedPtr(rewriter, loc, ptr);
1515 desc.setOffset(rewriter, loc,
1516 LLVM::createIndexAttrConstant(rewriter, loc, indexTy, 0));
1519 for (
const auto &indexedSize :
1520 llvm::enumerate(targetMemRefType.getShape())) {
1521 int64_t index = indexedSize.index();
1522 desc.setSize(rewriter, loc, index,
1523 LLVM::createIndexAttrConstant(rewriter, loc, indexTy,
1524 indexedSize.value()));
1525 desc.setStride(rewriter, loc, index,
1526 LLVM::createIndexAttrConstant(rewriter, loc, indexTy,
1527 (*targetStrides)[index]));
1530 rewriter.replaceOp(castOp, {desc});
1537class VectorCreateMaskOpConversion
1538 :
public OpConversionPattern<vector::CreateMaskOp> {
1540 explicit VectorCreateMaskOpConversion(MLIRContext *context,
1541 bool enableIndexOpt)
1542 : OpConversionPattern<vector::CreateMaskOp>(context),
1543 force32BitVectorIndices(enableIndexOpt) {}
1546 matchAndRewrite(vector::CreateMaskOp op, OpAdaptor adaptor,
1547 ConversionPatternRewriter &rewriter)
const override {
1548 auto dstType = op.getType();
1549 if (dstType.getRank() != 1 || !cast<VectorType>(dstType).isScalable())
1551 IntegerType idxType =
1552 force32BitVectorIndices ? rewriter.getI32Type() : rewriter.getI64Type();
1553 auto loc = op->getLoc();
1554 Value
indices = LLVM::StepVectorOp::create(
1556 LLVM::getVectorType(idxType, dstType.getShape()[0],
1558 Value maskBound = adaptor.getOperands()[0];
1565 if (force32BitVectorIndices) {
1568 maskBound = arith::MinSIOp::create(rewriter, loc, maskBound, maxBound);
1572 Value bounds = BroadcastOp::create(rewriter, loc,
indices.getType(), bound);
1573 Value comp = arith::CmpIOp::create(rewriter, loc, arith::CmpIPredicate::slt,
1575 rewriter.replaceOp(op, comp);
1580 const bool force32BitVectorIndices;
1584 SymbolTableCollection *symbolTables =
nullptr;
1587 explicit VectorPrintOpConversion(
1588 const LLVMTypeConverter &typeConverter,
1589 SymbolTableCollection *symbolTables =
nullptr)
1590 : ConvertOpToLLVMPattern<vector::PrintOp>(typeConverter),
1591 symbolTables(symbolTables) {}
1607 matchAndRewrite(vector::PrintOp
printOp, OpAdaptor adaptor,
1608 ConversionPatternRewriter &rewriter)
const override {
1609 auto parent =
printOp->getParentOfType<ModuleOp>();
1615 if (
auto value = adaptor.getSource()) {
1617 if (isa<VectorType>(printType)) {
1621 if (
failed(emitScalarPrint(rewriter, parent, loc, printType, value)))
1625 auto punct =
printOp.getPunctuation();
1626 if (
auto stringLiteral =
printOp.getStringLiteral()) {
1628 LLVM::createPrintStrCall(rewriter, loc, parent,
"vector_print_str",
1629 *stringLiteral, *getTypeConverter(),
1631 if (createResult.failed())
1634 }
else if (punct != PrintPunctuation::NoPunctuation) {
1635 FailureOr<LLVM::LLVMFuncOp> op = [&]() {
1637 case PrintPunctuation::Close:
1638 return LLVM::lookupOrCreatePrintCloseFn(rewriter, parent,
1640 case PrintPunctuation::Open:
1641 return LLVM::lookupOrCreatePrintOpenFn(rewriter, parent,
1643 case PrintPunctuation::Comma:
1644 return LLVM::lookupOrCreatePrintCommaFn(rewriter, parent,
1646 case PrintPunctuation::NewLine:
1647 return LLVM::lookupOrCreatePrintNewlineFn(rewriter, parent,
1650 llvm_unreachable(
"unexpected punctuation");
1655 emitCall(rewriter,
printOp->getLoc(), op.value());
1663 enum class PrintConversion {
1672 LogicalResult emitScalarPrint(ConversionPatternRewriter &rewriter,
1673 ModuleOp parent, Location loc, Type printType,
1674 Value value)
const {
1675 if (typeConverter->convertType(printType) ==
nullptr)
1679 PrintConversion conversion = PrintConversion::None;
1680 FailureOr<Operation *> printer;
1682 printer = LLVM::lookupOrCreatePrintF32Fn(rewriter, parent, symbolTables);
1684 printer = LLVM::lookupOrCreatePrintF64Fn(rewriter, parent, symbolTables);
1686 conversion = PrintConversion::Bitcast16;
1687 printer = LLVM::lookupOrCreatePrintF16Fn(rewriter, parent, symbolTables);
1689 conversion = PrintConversion::Bitcast16;
1690 printer = LLVM::lookupOrCreatePrintBF16Fn(rewriter, parent, symbolTables);
1692 printer = LLVM::lookupOrCreatePrintU64Fn(rewriter, parent, symbolTables);
1693 }
else if (
auto intTy = dyn_cast<IntegerType>(printType)) {
1697 unsigned width = intTy.getWidth();
1698 if (intTy.isUnsigned()) {
1701 conversion = PrintConversion::ZeroExt64;
1703 LLVM::lookupOrCreatePrintU64Fn(rewriter, parent, symbolTables);
1708 assert(intTy.isSignless() || intTy.isSigned());
1713 conversion = PrintConversion::ZeroExt64;
1714 else if (width < 64)
1715 conversion = PrintConversion::SignExt64;
1717 LLVM::lookupOrCreatePrintI64Fn(rewriter, parent, symbolTables);
1722 }
else if (
auto floatTy = dyn_cast<FloatType>(printType)) {
1725 llvm::APFloatBase::SemanticsToEnum(floatTy.getFloatSemantics());
1726 Value semValue = LLVM::ConstantOp::create(
1727 rewriter, loc, rewriter.getI32Type(),
1728 rewriter.getIntegerAttr(rewriter.getI32Type(), sem));
1730 LLVM::ZExtOp::create(rewriter, loc, rewriter.getI64Type(), value);
1732 LLVM::lookupOrCreateApFloatPrintFn(rewriter, parent, symbolTables);
1733 emitCall(rewriter, loc, printer.value(),
1742 switch (conversion) {
1743 case PrintConversion::ZeroExt64:
1744 value = arith::ExtUIOp::create(
1745 rewriter, loc, IntegerType::get(rewriter.getContext(), 64), value);
1747 case PrintConversion::SignExt64:
1748 value = arith::ExtSIOp::create(
1749 rewriter, loc, IntegerType::get(rewriter.getContext(), 64), value);
1751 case PrintConversion::Bitcast16:
1752 value = LLVM::BitcastOp::create(
1753 rewriter, loc, IntegerType::get(rewriter.getContext(), 16), value);
1755 case PrintConversion::None:
1758 emitCall(rewriter, loc, printer.value(), value);
1763 static void emitCall(ConversionPatternRewriter &rewriter, Location loc,
1765 LLVM::CallOp::create(rewriter, loc,
TypeRange(), SymbolRefAttr::get(ref),
1773struct VectorBroadcastScalarToLowRankLowering
1775 using ConvertOpToLLVMPattern<vector::BroadcastOp>::ConvertOpToLLVMPattern;
1778 matchAndRewrite(vector::BroadcastOp
broadcast, OpAdaptor adaptor,
1779 ConversionPatternRewriter &rewriter)
const override {
1780 if (isa<VectorType>(
broadcast.getSourceType()))
1781 return rewriter.notifyMatchFailure(
1782 broadcast,
"broadcast from vector type not handled");
1785 if (resultType.getRank() > 1)
1786 return rewriter.notifyMatchFailure(
broadcast,
1787 "broadcast to 2+-d handled elsewhere");
1793 auto zero = LLVM::ConstantOp::create(
1795 typeConverter->convertType(rewriter.getIntegerType(32)),
1796 rewriter.getZeroAttr(rewriter.getIntegerType(32)));
1799 if (resultType.getRank() == 0) {
1800 rewriter.replaceOpWithNewOp<LLVM::InsertElementOp>(
1801 broadcast, vectorType, poison, adaptor.getSource(), zero);
1806 LLVM::InsertElementOp::create(rewriter,
broadcast.
getLoc(), vectorType,
1807 poison, adaptor.getSource(), zero);
1811 SmallVector<int32_t> zeroValues(width, 0);
1814 auto shuffle = rewriter.createOrFold<LLVM::ShuffleVectorOp>(
1825struct VectorBroadcastScalarToNdLowering
1827 using ConvertOpToLLVMPattern<BroadcastOp>::ConvertOpToLLVMPattern;
1830 matchAndRewrite(BroadcastOp
broadcast, OpAdaptor adaptor,
1831 ConversionPatternRewriter &rewriter)
const override {
1832 if (isa<VectorType>(
broadcast.getSourceType()))
1833 return rewriter.notifyMatchFailure(
1834 broadcast,
"broadcast from vector type not handled");
1837 if (resultType.getRank() <= 1)
1838 return rewriter.notifyMatchFailure(
1839 broadcast,
"broadcast to 1-d or 0-d handled elsewhere");
1843 auto vectorTypeInfo =
1845 auto llvmNDVectorTy = vectorTypeInfo.llvmNDVectorTy;
1846 auto llvm1DVectorTy = vectorTypeInfo.llvm1DVectorTy;
1847 if (!llvmNDVectorTy || !llvm1DVectorTy)
1851 Value desc = LLVM::PoisonOp::create(rewriter, loc, llvmNDVectorTy);
1855 Value vdesc = LLVM::PoisonOp::create(rewriter, loc, llvm1DVectorTy);
1856 auto zero = LLVM::ConstantOp::create(
1857 rewriter, loc, typeConverter->convertType(rewriter.getIntegerType(32)),
1858 rewriter.getZeroAttr(rewriter.getIntegerType(32)));
1859 Value v = LLVM::InsertElementOp::create(rewriter, loc, llvm1DVectorTy,
1860 vdesc, adaptor.getSource(), zero);
1863 int64_t width = resultType.getDimSize(resultType.getRank() - 1);
1864 SmallVector<int32_t> zeroValues(width, 0);
1865 v = LLVM::ShuffleVectorOp::create(rewriter, loc, v, v, zeroValues);
1869 nDVectorIterate(vectorTypeInfo, rewriter, [&](ArrayRef<int64_t> position) {
1870 desc = LLVM::InsertValueOp::create(rewriter, loc, desc, v, position);
1879struct VectorInterleaveOpLowering
1884 matchAndRewrite(vector::InterleaveOp interleaveOp, OpAdaptor adaptor,
1885 ConversionPatternRewriter &rewriter)
const override {
1886 VectorType resultType = interleaveOp.getResultVectorType();
1888 if (resultType.getRank() != 1)
1889 return rewriter.notifyMatchFailure(interleaveOp,
1890 "InterleaveOp not rank 1");
1892 if (resultType.isScalable()) {
1893 rewriter.replaceOpWithNewOp<LLVM::vector_interleave2>(
1894 interleaveOp, typeConverter->convertType(resultType),
1895 adaptor.getLhs(), adaptor.getRhs());
1902 int64_t resultVectorSize = resultType.getNumElements();
1903 SmallVector<int32_t> interleaveShuffleMask;
1904 interleaveShuffleMask.reserve(resultVectorSize);
1905 for (
int i = 0, end = resultVectorSize / 2; i < end; ++i) {
1906 interleaveShuffleMask.push_back(i);
1907 interleaveShuffleMask.push_back((resultVectorSize / 2) + i);
1909 rewriter.replaceOpWithNewOp<LLVM::ShuffleVectorOp>(
1910 interleaveOp, adaptor.getLhs(), adaptor.getRhs(),
1911 interleaveShuffleMask);
1918struct VectorDeinterleaveOpLowering
1923 matchAndRewrite(vector::DeinterleaveOp deinterleaveOp, OpAdaptor adaptor,
1924 ConversionPatternRewriter &rewriter)
const override {
1925 VectorType resultType = deinterleaveOp.getResultVectorType();
1926 VectorType sourceType = deinterleaveOp.getSourceVectorType();
1927 auto loc = deinterleaveOp.getLoc();
1931 if (resultType.getRank() != 1)
1932 return rewriter.notifyMatchFailure(deinterleaveOp,
1933 "DeinterleaveOp not rank 1");
1935 if (resultType.isScalable()) {
1936 const auto *llvmTypeConverter = this->getTypeConverter();
1937 auto deinterleaveResults = deinterleaveOp.getResultTypes();
1938 auto packedOpResults =
1939 llvmTypeConverter->packOperationResults(deinterleaveResults);
1940 auto intrinsic = LLVM::vector_deinterleave2::create(
1941 rewriter, loc, packedOpResults, adaptor.getSource());
1943 auto evenResult = LLVM::ExtractValueOp::create(
1944 rewriter, loc, intrinsic->getResult(0), 0);
1945 auto oddResult = LLVM::ExtractValueOp::create(rewriter, loc,
1946 intrinsic->getResult(0), 1);
1948 rewriter.replaceOp(deinterleaveOp,
ValueRange{evenResult, oddResult});
1955 int64_t resultVectorSize = resultType.getNumElements();
1956 SmallVector<int32_t> evenShuffleMask;
1957 SmallVector<int32_t> oddShuffleMask;
1959 evenShuffleMask.reserve(resultVectorSize);
1960 oddShuffleMask.reserve(resultVectorSize);
1962 for (
int i = 0; i < sourceType.getNumElements(); ++i) {
1964 evenShuffleMask.push_back(i);
1966 oddShuffleMask.push_back(i);
1969 auto poison = LLVM::PoisonOp::create(rewriter, loc, sourceType);
1970 auto evenShuffle = LLVM::ShuffleVectorOp::create(
1971 rewriter, loc, adaptor.getSource(), poison, evenShuffleMask);
1972 auto oddShuffle = LLVM::ShuffleVectorOp::create(
1973 rewriter, loc, adaptor.getSource(), poison, oddShuffleMask);
1975 rewriter.replaceOp(deinterleaveOp,
ValueRange{evenShuffle, oddShuffle});
1981struct VectorFromElementsLowering
1986 matchAndRewrite(vector::FromElementsOp fromElementsOp, OpAdaptor adaptor,
1987 ConversionPatternRewriter &rewriter)
const override {
1988 Location loc = fromElementsOp.getLoc();
1989 VectorType vectorType = fromElementsOp.getType();
1993 if (vectorType.getRank() > 1)
1994 return rewriter.notifyMatchFailure(fromElementsOp,
1995 "rank > 1 vectors are not supported");
1996 Type llvmType = typeConverter->convertType(vectorType);
1997 Type llvmIndexType = typeConverter->convertType(rewriter.getIndexType());
1998 Value
result = LLVM::PoisonOp::create(rewriter, loc, llvmType);
1999 for (
auto [idx, val] : llvm::enumerate(adaptor.getElements())) {
2001 LLVM::ConstantOp::create(rewriter, loc, llvmIndexType, idx);
2002 result = LLVM::InsertElementOp::create(rewriter, loc, llvmType,
result,
2005 rewriter.replaceOp(fromElementsOp,
result);
2011struct VectorToElementsLowering
2016 matchAndRewrite(vector::ToElementsOp toElementsOp, OpAdaptor adaptor,
2017 ConversionPatternRewriter &rewriter)
const override {
2018 Location loc = toElementsOp.getLoc();
2019 auto idxType = typeConverter->convertType(rewriter.getIndexType());
2020 Value source = adaptor.getSource();
2022 SmallVector<Value> results(toElementsOp->getNumResults());
2023 for (
auto [idx, element] : llvm::enumerate(toElementsOp.getElements())) {
2025 if (element.use_empty())
2028 auto constIdx = LLVM::ConstantOp::create(
2029 rewriter, loc, idxType, rewriter.getIntegerAttr(idxType, idx));
2030 auto llvmType = typeConverter->convertType(element.getType());
2032 Value
result = LLVM::ExtractElementOp::create(rewriter, loc, llvmType,
2037 rewriter.replaceOp(toElementsOp, results);
2047 matchAndRewrite(vector::StepOp stepOp, OpAdaptor adaptor,
2048 ConversionPatternRewriter &rewriter)
const override {
2049 Type llvmType = typeConverter->convertType(stepOp.getType());
2050 rewriter.replaceOpWithNewOp<LLVM::StepVectorOp>(stepOp, llvmType);
2065class ContractionOpToMatmulOpLowering
2068 using MaskableOpRewritePattern::MaskableOpRewritePattern;
2070 ContractionOpToMatmulOpLowering(MLIRContext *context,
2071 PatternBenefit benefit = 100)
2072 : MaskableOpRewritePattern<vector::ContractionOp>(context, benefit) {}
2075 matchAndRewriteMaskableOp(vector::ContractionOp op, MaskingOpInterface maskOp,
2076 PatternRewriter &rewriter)
const override;
2096FailureOr<Value> ContractionOpToMatmulOpLowering::matchAndRewriteMaskableOp(
2097 vector::ContractionOp op, MaskingOpInterface maskOp,
2103 auto iteratorTypes = op.getIteratorTypes().getValue();
2109 Type opResType = op.getType();
2110 VectorType vecType = dyn_cast<VectorType>(opResType);
2111 if (vecType && vecType.isScalable()) {
2116 Type elementType = op.getLhsType().getElementType();
2120 Type dstElementType = vecType ? vecType.getElementType() : opResType;
2121 if (elementType != dstElementType)
2126 MLIRContext *ctx = op.getContext();
2127 Location loc = op.getLoc();
2131 Value
lhs = op.getLhs();
2132 auto lhsMap = op.getIndexingMapsArray()[0];
2134 lhs = vector::TransposeOp::create(rew, loc,
lhs, ArrayRef<int64_t>{1, 0});
2139 Value
rhs = op.getRhs();
2140 auto rhsMap = op.getIndexingMapsArray()[1];
2142 rhs = vector::TransposeOp::create(rew, loc,
rhs, ArrayRef<int64_t>{1, 0});
2147 VectorType lhsType = cast<VectorType>(
lhs.getType());
2148 VectorType rhsType = cast<VectorType>(
rhs.getType());
2149 int64_t lhsRows = lhsType.getDimSize(0);
2150 int64_t lhsColumns = lhsType.getDimSize(1);
2151 int64_t rhsColumns = rhsType.getDimSize(1);
2153 Type flattenedLHSType =
2154 VectorType::get(lhsType.getNumElements(), lhsType.getElementType());
2155 lhs = vector::ShapeCastOp::create(rew, loc, flattenedLHSType,
lhs);
2157 Type flattenedRHSType =
2158 VectorType::get(rhsType.getNumElements(), rhsType.getElementType());
2159 rhs = vector::ShapeCastOp::create(rew, loc, flattenedRHSType,
rhs);
2161 Value
mul = LLVM::MatrixMultiplyOp::create(
2163 VectorType::get(lhsRows * rhsColumns,
2164 cast<VectorType>(
lhs.getType()).getElementType()),
2165 lhs,
rhs, lhsRows, lhsColumns, rhsColumns);
2167 mul = vector::ShapeCastOp::create(
2169 VectorType::get({lhsRows, rhsColumns},
2174 auto accMap = op.getIndexingMapsArray()[2];
2176 mul = vector::TransposeOp::create(rew, loc,
mul, ArrayRef<int64_t>{1, 0});
2178 llvm_unreachable(
"invalid contraction semantics");
2180 Value res = isa<IntegerType>(elementType)
2181 ?
static_cast<Value
>(
2182 arith::AddIOp::create(rew, loc, op.getAcc(),
mul))
2183 : static_cast<Value>(
2184 arith::AddFOp::create(rew, loc, op.getAcc(),
mul));
2202class TransposeOpToMatrixTransposeOpLowering
2203 :
public OpRewritePattern<vector::TransposeOp> {
2207 LogicalResult matchAndRewrite(vector::TransposeOp op,
2208 PatternRewriter &rewriter)
const override {
2209 auto loc = op.getLoc();
2211 Value input = op.getVector();
2212 VectorType inputType = op.getSourceVectorType();
2213 VectorType resType = op.getResultVectorType();
2215 if (inputType.isScalable())
2217 op,
"This lowering does not support scalable vectors");
2220 ArrayRef<int64_t> transp = op.getPermutation();
2222 if (resType.getRank() != 2 || transp[0] != 1 || transp[1] != 0) {
2226 Type flattenedType =
2227 VectorType::get(resType.getNumElements(), resType.getElementType());
2229 vector::ShapeCastOp::create(rewriter, loc, flattenedType, input);
2232 Value trans = LLVM::MatrixTransposeOp::create(rewriter, loc, flattenedType,
2233 matrix, rows, columns);
2243 patterns.
add<VectorFMAOpNDRewritePattern>(patterns.
getContext());
2248 patterns.
add<ContractionOpToMatmulOpLowering>(patterns.
getContext(), benefit);
2253 patterns.
add<TransposeOpToMatrixTransposeOpLowering>(patterns.
getContext(),
2260 bool reassociateFPReductions,
bool force32BitVectorIndices,
2261 bool useVectorAlignment,
bool enableGEPInboundsNuw) {
2264 patterns.
add<VectorReductionOpConversion>(converter, reassociateFPReductions);
2265 patterns.
add<VectorCreateMaskOpConversion>(ctx, force32BitVectorIndices);
2266 patterns.
add<VectorLoadStoreConversion<vector::LoadOp>,
2267 VectorLoadStoreConversion<vector::MaskedLoadOp>,
2268 VectorLoadStoreConversion<vector::StoreOp>,
2269 VectorLoadStoreConversion<vector::MaskedStoreOp>>(
2270 converter, useVectorAlignment, enableGEPInboundsNuw);
2271 patterns.
add<VectorGatherOpConversion, VectorScatterOpConversion>(
2272 converter, useVectorAlignment);
2273 patterns.
add<VectorBitCastOpConversion, VectorShuffleOpConversion,
2274 VectorExtractOpConversion, VectorFMAOp1DConversion,
2275 VectorInsertOpConversion, VectorPrintOpConversion,
2276 VectorTypeCastOpConversion, VectorScaleOpConversion,
2277 VectorExpandLoadOpConversion, VectorCompressStoreOpConversion,
2278 VectorBroadcastScalarToLowRankLowering,
2279 VectorBroadcastScalarToNdLowering,
2280 VectorScalableInsertOpLowering, VectorScalableExtractOpLowering,
2281 MaskedReductionOpConversion, VectorInterleaveOpLowering,
2282 VectorDeinterleaveOpLowering, VectorFromElementsLowering,
2283 VectorToElementsLowering, VectorStepOpLowering>(converter);
2287struct VectorToLLVMDialectInterface :
public ConvertToLLVMPatternInterface {
2288 VectorToLLVMDialectInterface(
Dialect *dialect)
2289 : ConvertToLLVMPatternInterface(dialect) {}
2291 using ConvertToLLVMPatternInterface::ConvertToLLVMPatternInterface;
2292 void loadDependentDialects(MLIRContext *context)
const final {
2293 context->loadDialect<LLVM::LLVMDialect>();
2298 void populateConvertToLLVMConversionPatterns(
2299 ConversionTarget &
target, LLVMTypeConverter &typeConverter,
2300 RewritePatternSet &patterns)
const final {
2309 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.