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;
730struct VectorToScalarMapper<
LLVM::vector_reduce_fmaximumnum> {
731 using Type = LLVM::MaximumNumOp;
734struct VectorToScalarMapper<
LLVM::vector_reduce_fminimumnum> {
735 using Type = LLVM::MinimumNumOp;
739template <
class LLVMRedIntrinOp>
740static Value createFPReductionComparisonOpLowering(
741 ConversionPatternRewriter &rewriter,
Location loc,
Type llvmType,
742 Value vectorOperand,
Value accumulator, LLVM::FastmathFlagsAttr fmf) {
744 LLVMRedIntrinOp::create(rewriter, loc, llvmType, vectorOperand, fmf);
747 result = VectorToScalarMapper<LLVMRedIntrinOp>::Type::create(
748 rewriter, loc,
result, accumulator);
755class MaskNeutralFMaximumNum {};
756class MaskNeutralFMinimumNum {};
762static llvm::APFloat getMaskNeutralValue(MaskNeutralFMaximumNum,
763 const llvm::fltSemantics &semantics,
764 bool noNaNs,
bool noInfs) {
766 return llvm::APFloat::getQNaN(semantics);
768 return llvm::APFloat::getLargest(semantics,
true);
769 return llvm::APFloat::getInf(semantics,
true);
774static llvm::APFloat getMaskNeutralValue(MaskNeutralFMinimumNum,
775 const llvm::fltSemantics &semantics,
776 bool noNaNs,
bool noInfs) {
778 return llvm::APFloat::getQNaN(semantics);
780 return llvm::APFloat::getLargest(semantics);
781 return llvm::APFloat::getInf(semantics);
789template <
class LLVMRedIntrinOp,
class MaskNeutral>
791lowerMaskedReductionWithRegular(ConversionPatternRewriter &rewriter,
794 Value mask, LLVM::FastmathFlagsAttr fmf) {
795 const auto &floatSemantics = cast<FloatType>(llvmType).getFloatSemantics();
796 auto value = getMaskNeutralValue(
797 MaskNeutral{}, floatSemantics,
798 LLVM::bitEnumContainsAny(fmf.getValue(), LLVM::FastmathFlags::nnan),
799 LLVM::bitEnumContainsAny(fmf.getValue(), LLVM::FastmathFlags::ninf));
802 const Value vectorMaskNeutral =
803 LLVM::ConstantOp::create(rewriter, loc, vectorType, denseValue);
804 const Value selectedVectorByMask = LLVM::SelectOp::create(
805 rewriter, loc, mask, vectorOperand, vectorMaskNeutral);
806 return createFPReductionComparisonOpLowering<LLVMRedIntrinOp>(
807 rewriter, loc, llvmType, selectedVectorByMask, accumulator, fmf);
810template <
class LLVMRedIntrinOp,
class ReductionNeutral>
812lowerReductionWithStartValue(ConversionPatternRewriter &rewriter,
Location loc,
814 Value accumulator, LLVM::FastmathFlagsAttr fmf) {
815 accumulator = getOrCreateAccumulator<ReductionNeutral>(rewriter, loc,
816 llvmType, accumulator);
817 return LLVMRedIntrinOp::create(rewriter, loc, llvmType,
818 accumulator, vectorOperand,
822template <
class LLVMVPRedIntrinOp,
class ReductionNeutral>
823static Value lowerPredicatedReductionWithStartValue(
824 ConversionPatternRewriter &rewriter,
Location loc,
Type llvmType,
826 accumulator = getOrCreateAccumulator<ReductionNeutral>(rewriter, loc,
827 llvmType, accumulator);
829 createVectorLengthValue(rewriter, loc, vectorOperand.
getType());
830 return LLVMVPRedIntrinOp::create(rewriter, loc, llvmType,
831 accumulator, vectorOperand,
835template <
class LLVMIntVPRedIntrinOp,
class IntReductionNeutral,
836 class LLVMFPVPRedIntrinOp,
class FPReductionNeutral>
837static Value lowerPredicatedReductionWithStartValue(
838 ConversionPatternRewriter &rewriter,
Location loc,
Type llvmType,
841 return lowerPredicatedReductionWithStartValue<LLVMIntVPRedIntrinOp,
842 IntReductionNeutral>(
843 rewriter, loc, llvmType, vectorOperand, accumulator, mask);
846 return lowerPredicatedReductionWithStartValue<LLVMFPVPRedIntrinOp,
848 rewriter, loc, llvmType, vectorOperand, accumulator, mask);
852class VectorReductionOpConversion
855 explicit VectorReductionOpConversion(
const LLVMTypeConverter &typeConv,
856 bool reassociateFPRed)
857 : ConvertOpToLLVMPattern<vector::ReductionOp>(typeConv),
858 reassociateFPReductions(reassociateFPRed) {}
861 matchAndRewrite(vector::ReductionOp reductionOp, OpAdaptor adaptor,
862 ConversionPatternRewriter &rewriter)
const override {
863 auto kind = reductionOp.getKind();
864 Type eltType = reductionOp.getDest().getType();
865 Type llvmType = typeConverter->convertType(eltType);
866 Value operand = adaptor.getVector();
867 Value acc = adaptor.getAcc();
868 Location loc = reductionOp.getLoc();
874 case vector::CombiningKind::ADD:
876 createIntegerReductionArithmeticOpLowering<LLVM::vector_reduce_add,
878 rewriter, loc, llvmType, operand, acc);
880 case vector::CombiningKind::MUL:
882 createIntegerReductionArithmeticOpLowering<LLVM::vector_reduce_mul,
884 rewriter, loc, llvmType, operand, acc);
886 case vector::CombiningKind::MINUI:
887 result = createIntegerReductionComparisonOpLowering<
888 LLVM::vector_reduce_umin>(rewriter, loc, llvmType, operand, acc,
889 LLVM::ICmpPredicate::ule);
891 case vector::CombiningKind::MINSI:
892 result = createIntegerReductionComparisonOpLowering<
893 LLVM::vector_reduce_smin>(rewriter, loc, llvmType, operand, acc,
894 LLVM::ICmpPredicate::sle);
896 case vector::CombiningKind::MAXUI:
897 result = createIntegerReductionComparisonOpLowering<
898 LLVM::vector_reduce_umax>(rewriter, loc, llvmType, operand, acc,
899 LLVM::ICmpPredicate::uge);
901 case vector::CombiningKind::MAXSI:
902 result = createIntegerReductionComparisonOpLowering<
903 LLVM::vector_reduce_smax>(rewriter, loc, llvmType, operand, acc,
904 LLVM::ICmpPredicate::sge);
906 case vector::CombiningKind::AND:
908 createIntegerReductionArithmeticOpLowering<LLVM::vector_reduce_and,
910 rewriter, loc, llvmType, operand, acc);
912 case vector::CombiningKind::OR:
914 createIntegerReductionArithmeticOpLowering<LLVM::vector_reduce_or,
916 rewriter, loc, llvmType, operand, acc);
918 case vector::CombiningKind::XOR:
920 createIntegerReductionArithmeticOpLowering<LLVM::vector_reduce_xor,
922 rewriter, loc, llvmType, operand, acc);
927 rewriter.replaceOp(reductionOp,
result);
932 if (!isa<FloatType>(eltType))
935 arith::FastMathFlagsAttr fMFAttr = reductionOp.getFastMathFlagsAttr();
936 LLVM::FastmathFlagsAttr fmf = LLVM::FastmathFlagsAttr::get(
937 reductionOp.getContext(),
939 fmf = LLVM::FastmathFlagsAttr::get(
940 reductionOp.getContext(),
941 fmf.getValue() | (reassociateFPReductions ? LLVM::FastmathFlags::reassoc
942 : LLVM::FastmathFlags::none));
946 if (kind == vector::CombiningKind::ADD) {
947 result = lowerReductionWithStartValue<LLVM::vector_reduce_fadd,
948 ReductionNeutralZero>(
949 rewriter, loc, llvmType, operand, acc, fmf);
950 }
else if (kind == vector::CombiningKind::MUL) {
951 result = lowerReductionWithStartValue<LLVM::vector_reduce_fmul,
952 ReductionNeutralFPOne>(
953 rewriter, loc, llvmType, operand, acc, fmf);
954 }
else if (kind == vector::CombiningKind::MINIMUMF) {
956 createFPReductionComparisonOpLowering<LLVM::vector_reduce_fminimum>(
957 rewriter, loc, llvmType, operand, acc, fmf);
958 }
else if (kind == vector::CombiningKind::MAXIMUMF) {
960 createFPReductionComparisonOpLowering<LLVM::vector_reduce_fmaximum>(
961 rewriter, loc, llvmType, operand, acc, fmf);
962 }
else if (kind == vector::CombiningKind::MINNUMF) {
963 result = createFPReductionComparisonOpLowering<LLVM::vector_reduce_fmin>(
964 rewriter, loc, llvmType, operand, acc, fmf);
965 }
else if (kind == vector::CombiningKind::MAXNUMF) {
966 result = createFPReductionComparisonOpLowering<LLVM::vector_reduce_fmax>(
967 rewriter, loc, llvmType, operand, acc, fmf);
968 }
else if (kind == vector::CombiningKind::MAXIMUMNUMF) {
969 result = createFPReductionComparisonOpLowering<
970 LLVM::vector_reduce_fmaximumnum>(rewriter, loc, llvmType, operand,
972 }
else if (kind == vector::CombiningKind::MINIMUMNUMF) {
973 result = createFPReductionComparisonOpLowering<
974 LLVM::vector_reduce_fminimumnum>(rewriter, loc, llvmType, operand,
980 rewriter.replaceOp(reductionOp,
result);
985 const bool reassociateFPReductions;
996template <
class MaskedOp>
997class VectorMaskOpConversionBase
1000 using ConvertOpToLLVMPattern<vector::MaskOp>::ConvertOpToLLVMPattern;
1003 matchAndRewrite(vector::MaskOp maskOp, OpAdaptor adaptor,
1004 ConversionPatternRewriter &rewriter)
const final {
1006 auto maskedOp = llvm::dyn_cast_or_null<MaskedOp>(maskOp.getMaskableOp());
1009 return matchAndRewriteMaskableOp(maskOp, maskedOp, rewriter);
1013 virtual LogicalResult
1014 matchAndRewriteMaskableOp(vector::MaskOp maskOp,
1015 vector::MaskableOpInterface maskableOp,
1016 ConversionPatternRewriter &rewriter)
const = 0;
1019class MaskedReductionOpConversion
1020 :
public VectorMaskOpConversionBase<vector::ReductionOp> {
1023 using VectorMaskOpConversionBase<
1024 vector::ReductionOp>::VectorMaskOpConversionBase;
1026 LogicalResult matchAndRewriteMaskableOp(
1027 vector::MaskOp maskOp, MaskableOpInterface maskableOp,
1028 ConversionPatternRewriter &rewriter)
const override {
1029 auto reductionOp = cast<ReductionOp>(maskableOp.getOperation());
1030 auto kind = reductionOp.getKind();
1031 Type eltType = reductionOp.getDest().getType();
1032 Type llvmType = typeConverter->convertType(eltType);
1033 Value operand = reductionOp.getVector();
1034 Value acc = reductionOp.getAcc();
1035 Location loc = reductionOp.getLoc();
1037 arith::FastMathFlagsAttr fMFAttr = reductionOp.getFastMathFlagsAttr();
1038 LLVM::FastmathFlagsAttr fmf = LLVM::FastmathFlagsAttr::get(
1039 reductionOp.getContext(),
1042 LLVM::bitEnumContainsAny(fmf.getValue(), LLVM::FastmathFlags::ninf);
1046 case vector::CombiningKind::ADD:
1047 result = lowerPredicatedReductionWithStartValue<
1048 LLVM::VPReduceAddOp, ReductionNeutralZero, LLVM::VPReduceFAddOp,
1049 ReductionNeutralZero>(rewriter, loc, llvmType, operand, acc,
1052 case vector::CombiningKind::MUL:
1053 result = lowerPredicatedReductionWithStartValue<
1054 LLVM::VPReduceMulOp, ReductionNeutralIntOne, LLVM::VPReduceFMulOp,
1055 ReductionNeutralFPOne>(rewriter, loc, llvmType, operand, acc,
1058 case vector::CombiningKind::MINUI:
1059 result = lowerPredicatedReductionWithStartValue<LLVM::VPReduceUMinOp,
1060 ReductionNeutralUIntMax>(
1061 rewriter, loc, llvmType, operand, acc, maskOp.getMask());
1063 case vector::CombiningKind::MINSI:
1064 result = lowerPredicatedReductionWithStartValue<LLVM::VPReduceSMinOp,
1065 ReductionNeutralSIntMax>(
1066 rewriter, loc, llvmType, operand, acc, maskOp.getMask());
1068 case vector::CombiningKind::MAXUI:
1069 result = lowerPredicatedReductionWithStartValue<LLVM::VPReduceUMaxOp,
1070 ReductionNeutralUIntMin>(
1071 rewriter, loc, llvmType, operand, acc, maskOp.getMask());
1073 case vector::CombiningKind::MAXSI:
1074 result = lowerPredicatedReductionWithStartValue<LLVM::VPReduceSMaxOp,
1075 ReductionNeutralSIntMin>(
1076 rewriter, loc, llvmType, operand, acc, maskOp.getMask());
1078 case vector::CombiningKind::AND:
1079 result = lowerPredicatedReductionWithStartValue<LLVM::VPReduceAndOp,
1080 ReductionNeutralAllOnes>(
1081 rewriter, loc, llvmType, operand, acc, maskOp.getMask());
1083 case vector::CombiningKind::OR:
1084 result = lowerPredicatedReductionWithStartValue<LLVM::VPReduceOrOp,
1085 ReductionNeutralZero>(
1086 rewriter, loc, llvmType, operand, acc, maskOp.getMask());
1088 case vector::CombiningKind::XOR:
1089 result = lowerPredicatedReductionWithStartValue<LLVM::VPReduceXorOp,
1090 ReductionNeutralZero>(
1091 rewriter, loc, llvmType, operand, acc, maskOp.getMask());
1093 case vector::CombiningKind::MINNUMF:
1095 lowerPredicatedReductionWithStartValue<LLVM::VPReduceFMinOp,
1096 ReductionNeutralFPNegQNaN>(
1097 rewriter, loc, llvmType, operand, acc, maskOp.getMask());
1099 case vector::CombiningKind::MAXNUMF:
1100 result = lowerPredicatedReductionWithStartValue<LLVM::VPReduceFMaxOp,
1101 ReductionNeutralFPQNaN>(
1102 rewriter, loc, llvmType, operand, acc, maskOp.getMask());
1104 case CombiningKind::MAXIMUMF:
1109 ? lowerPredicatedReductionWithStartValue<
1110 LLVM::VPReduceFMaximumOp, ReductionNeutralFPLowestFinite>(
1111 rewriter, loc, llvmType, operand, acc, maskOp.getMask())
1112 : lowerPredicatedReductionWithStartValue<
1113 LLVM::VPReduceFMaximumOp, ReductionNeutralFPNegInf>(
1114 rewriter, loc, llvmType, operand, acc, maskOp.getMask());
1116 case CombiningKind::MINIMUMF:
1119 ? lowerPredicatedReductionWithStartValue<
1120 LLVM::VPReduceFMinimumOp, ReductionNeutralFPLargestFinite>(
1121 rewriter, loc, llvmType, operand, acc, maskOp.getMask())
1122 : lowerPredicatedReductionWithStartValue<
1123 LLVM::VPReduceFMinimumOp, ReductionNeutralFPPosInf>(
1124 rewriter, loc, llvmType, operand, acc, maskOp.getMask());
1126 case CombiningKind::MAXIMUMNUMF:
1127 result = lowerMaskedReductionWithRegular<LLVM::vector_reduce_fmaximumnum,
1128 MaskNeutralFMaximumNum>(
1129 rewriter, loc, llvmType, operand, acc, maskOp.getMask(), fmf);
1131 case CombiningKind::MINIMUMNUMF:
1132 result = lowerMaskedReductionWithRegular<LLVM::vector_reduce_fminimumnum,
1133 MaskNeutralFMinimumNum>(
1134 rewriter, loc, llvmType, operand, acc, maskOp.getMask(), fmf);
1139 rewriter.replaceOp(maskOp,
result);
1144class VectorShuffleOpConversion
1147 using ConvertOpToLLVMPattern<vector::ShuffleOp>::ConvertOpToLLVMPattern;
1150 matchAndRewrite(vector::ShuffleOp shuffleOp, OpAdaptor adaptor,
1151 ConversionPatternRewriter &rewriter)
const override {
1152 auto loc = shuffleOp->getLoc();
1153 auto v1Type = shuffleOp.getV1VectorType();
1154 auto v2Type = shuffleOp.getV2VectorType();
1155 auto vectorType = shuffleOp.getResultVectorType();
1156 Type llvmType = typeConverter->convertType(vectorType);
1157 ArrayRef<int64_t> mask = shuffleOp.getMask();
1164 int64_t rank = vectorType.getRank();
1166 bool wellFormed0DCase =
1167 v1Type.getRank() == 0 && v2Type.getRank() == 0 && rank == 1;
1168 bool wellFormedNDCase =
1169 v1Type.getRank() == rank && v2Type.getRank() == rank;
1170 assert((wellFormed0DCase || wellFormedNDCase) &&
"op is not well-formed");
1175 if (rank <= 1 && v1Type == v2Type) {
1176 Value llvmShuffleOp = LLVM::ShuffleVectorOp::create(
1177 rewriter, loc, adaptor.getV1(), adaptor.getV2(),
1178 llvm::to_vector_of<int32_t>(mask));
1179 rewriter.replaceOp(shuffleOp, llvmShuffleOp);
1184 int64_t v1Dim = v1Type.getDimSize(0);
1186 if (
auto arrayType = dyn_cast<LLVM::LLVMArrayType>(llvmType))
1187 eltType = arrayType.getElementType();
1189 eltType = cast<VectorType>(llvmType).getElementType();
1190 Value insert = LLVM::PoisonOp::create(rewriter, loc, llvmType);
1192 for (int64_t extPos : mask) {
1193 Value value = adaptor.getV1();
1194 if (extPos >= v1Dim) {
1196 value = adaptor.getV2();
1198 Value extract =
extractOne(rewriter, *getTypeConverter(), loc, value,
1199 eltType, rank, extPos);
1200 insert =
insertOne(rewriter, *getTypeConverter(), loc, insert, extract,
1201 llvmType, rank, insPos++);
1203 rewriter.replaceOp(shuffleOp, insert);
1208class VectorExtractOpConversion
1211 using ConvertOpToLLVMPattern<vector::ExtractOp>::ConvertOpToLLVMPattern;
1214 matchAndRewrite(vector::ExtractOp extractOp, OpAdaptor adaptor,
1215 ConversionPatternRewriter &rewriter)
const override {
1216 auto loc = extractOp->getLoc();
1217 auto resultType = extractOp.getResult().getType();
1218 auto llvmResultType = typeConverter->convertType(resultType);
1220 if (!llvmResultType)
1224 adaptor.getStaticPosition(), adaptor.getDynamicPosition(), rewriter);
1238 bool extractsAggregate = extractOp.getSourceVectorType().getRank() >= 2;
1242 bool extractsScalar =
static_cast<int64_t
>(positionVec.size()) ==
1243 extractOp.getSourceVectorType().getRank();
1247 if (extractOp.getSourceVectorType().getRank() == 0) {
1248 Type idxType = typeConverter->convertType(rewriter.getIndexType());
1249 positionVec.push_back(rewriter.getZeroAttr(idxType));
1252 Value extracted = adaptor.getSource();
1253 if (extractsAggregate) {
1254 ArrayRef<OpFoldResult> position(positionVec);
1255 if (extractsScalar) {
1259 position = position.drop_back();
1262 if (!llvm::all_of(position, llvm::IsaPred<Attribute>)) {
1265 extracted = LLVM::ExtractValueOp::create(rewriter, loc, extracted,
1269 if (extractsScalar) {
1270 extracted = LLVM::ExtractElementOp::create(
1271 rewriter, loc, extracted,
1275 rewriter.replaceOp(extractOp, extracted);
1296 using ConvertOpToLLVMPattern<vector::FMAOp>::ConvertOpToLLVMPattern;
1299 matchAndRewrite(vector::FMAOp fmaOp, OpAdaptor adaptor,
1300 ConversionPatternRewriter &rewriter)
const override {
1301 VectorType vType = fmaOp.getVectorType();
1302 if (vType.getRank() > 1)
1305 rewriter.replaceOpWithNewOp<LLVM::FMulAddOp>(
1306 fmaOp, adaptor.getLhs(), adaptor.getRhs(), adaptor.getAcc());
1311class VectorInsertOpConversion
1314 using ConvertOpToLLVMPattern<vector::InsertOp>::ConvertOpToLLVMPattern;
1317 matchAndRewrite(vector::InsertOp insertOp, OpAdaptor adaptor,
1318 ConversionPatternRewriter &rewriter)
const override {
1319 auto loc = insertOp->getLoc();
1320 auto destVectorType = insertOp.getDestVectorType();
1321 auto llvmResultType = typeConverter->convertType(destVectorType);
1323 if (!llvmResultType)
1327 adaptor.getStaticPosition(), adaptor.getDynamicPosition(), rewriter);
1349 bool isNestedAggregate = isa<LLVM::LLVMArrayType>(llvmResultType);
1351 bool insertIntoInnermostDim =
1352 static_cast<int64_t
>(positionVec.size()) == destVectorType.getRank();
1354 ArrayRef<OpFoldResult> positionOf1DVectorWithinAggregate(
1355 positionVec.begin(),
1356 insertIntoInnermostDim ? positionVec.size() - 1 : positionVec.size());
1357 OpFoldResult positionOfScalarWithin1DVector;
1358 if (destVectorType.getRank() == 0) {
1361 Type idxType = typeConverter->convertType(rewriter.getIndexType());
1362 positionOfScalarWithin1DVector = rewriter.getZeroAttr(idxType);
1363 }
else if (insertIntoInnermostDim) {
1364 positionOfScalarWithin1DVector = positionVec.back();
1370 Value sourceAggregate = adaptor.getValueToStore();
1371 if (insertIntoInnermostDim) {
1374 if (isNestedAggregate) {
1377 if (!llvm::all_of(positionOf1DVectorWithinAggregate,
1378 llvm::IsaPred<Attribute>)) {
1382 sourceAggregate = LLVM::ExtractValueOp::create(
1383 rewriter, loc, adaptor.getDest(),
1388 sourceAggregate = adaptor.getDest();
1391 sourceAggregate = LLVM::InsertElementOp::create(
1392 rewriter, loc, sourceAggregate.
getType(), sourceAggregate,
1393 adaptor.getValueToStore(),
1397 Value
result = sourceAggregate;
1398 if (isNestedAggregate) {
1399 if (!llvm::all_of(positionOf1DVectorWithinAggregate,
1400 llvm::IsaPred<Attribute>)) {
1404 result = LLVM::InsertValueOp::create(
1405 rewriter, loc, adaptor.getDest(), sourceAggregate,
1409 rewriter.replaceOp(insertOp,
result);
1415struct VectorScalableInsertOpLowering
1417 using ConvertOpToLLVMPattern<
1418 vector::ScalableInsertOp>::ConvertOpToLLVMPattern;
1421 matchAndRewrite(vector::ScalableInsertOp insOp, OpAdaptor adaptor,
1422 ConversionPatternRewriter &rewriter)
const override {
1423 rewriter.replaceOpWithNewOp<LLVM::vector_insert>(
1424 insOp, adaptor.getDest(), adaptor.getValueToStore(), adaptor.getPos());
1430struct VectorScalableExtractOpLowering
1432 using ConvertOpToLLVMPattern<
1433 vector::ScalableExtractOp>::ConvertOpToLLVMPattern;
1436 matchAndRewrite(vector::ScalableExtractOp extOp, OpAdaptor adaptor,
1437 ConversionPatternRewriter &rewriter)
const override {
1438 rewriter.replaceOpWithNewOp<LLVM::vector_extract>(
1439 extOp, typeConverter->convertType(extOp.getResultVectorType()),
1440 adaptor.getSource(), adaptor.getPos());
1473 setHasBoundedRewriteRecursion();
1476 LogicalResult matchAndRewrite(FMAOp op,
1477 PatternRewriter &rewriter)
const override {
1478 auto vType = op.getVectorType();
1479 if (vType.getRank() < 2)
1482 auto loc = op.getLoc();
1483 auto elemType = vType.getElementType();
1484 Value zero = arith::ConstantOp::create(rewriter, loc, elemType,
1486 Value desc = vector::BroadcastOp::create(rewriter, loc, vType, zero);
1487 for (int64_t i = 0, e = vType.getShape().front(); i != e; ++i) {
1488 Value extrLHS = ExtractOp::create(rewriter, loc, op.getLhs(), i);
1489 Value extrRHS = ExtractOp::create(rewriter, loc, op.getRhs(), i);
1490 Value extrACC = ExtractOp::create(rewriter, loc, op.getAcc(), i);
1491 Value fma = FMAOp::create(rewriter, loc, extrLHS, extrRHS, extrACC);
1492 desc = InsertOp::create(rewriter, loc, fma, desc, i);
1501static std::optional<SmallVector<int64_t, 4>>
1502computeContiguousStrides(MemRefType memRefType) {
1505 if (
failed(memRefType.getStridesAndOffset(strides, offset)))
1506 return std::nullopt;
1507 if (!strides.empty() && strides.back() != 1)
1508 return std::nullopt;
1510 if (memRefType.getLayout().isIdentity())
1517 auto sizes = memRefType.getShape();
1519 if (ShapedType::isDynamic(sizes[
index + 1]) ||
1520 ShapedType::isDynamic(strides[
index]) ||
1521 ShapedType::isDynamic(strides[
index + 1]))
1522 return std::nullopt;
1524 return std::nullopt;
1529class VectorTypeCastOpConversion
1532 using ConvertOpToLLVMPattern<vector::TypeCastOp>::ConvertOpToLLVMPattern;
1535 matchAndRewrite(vector::TypeCastOp castOp, OpAdaptor adaptor,
1536 ConversionPatternRewriter &rewriter)
const override {
1537 auto loc = castOp->getLoc();
1538 MemRefType sourceMemRefType =
1539 cast<MemRefType>(castOp.getOperand().getType());
1540 MemRefType targetMemRefType = castOp.getType();
1543 if (!sourceMemRefType.hasStaticShape() ||
1544 !targetMemRefType.hasStaticShape())
1547 auto llvmSourceDescriptorTy =
1548 dyn_cast<LLVM::LLVMStructType>(adaptor.getOperands()[0].getType());
1549 if (!llvmSourceDescriptorTy)
1551 MemRefDescriptor sourceMemRef(adaptor.getOperands()[0]);
1553 auto llvmTargetDescriptorTy = dyn_cast_or_null<LLVM::LLVMStructType>(
1554 typeConverter->convertType(targetMemRefType));
1555 if (!llvmTargetDescriptorTy)
1559 auto sourceStrides = computeContiguousStrides(sourceMemRefType);
1562 auto targetStrides = computeContiguousStrides(targetMemRefType);
1566 if (llvm::any_of(*targetStrides, ShapedType::isDynamic))
1571 Type indexTy = getTypeConverter()->getIndexType();
1574 auto desc = MemRefDescriptor::poison(rewriter, loc, llvmTargetDescriptorTy);
1576 Value allocated = sourceMemRef.allocatedPtr(rewriter, loc);
1577 desc.setAllocatedPtr(rewriter, loc, allocated);
1580 Value ptr = sourceMemRef.alignedPtr(rewriter, loc);
1581 desc.setAlignedPtr(rewriter, loc, ptr);
1583 desc.setOffset(rewriter, loc,
1587 for (
const auto &indexedSize :
1588 llvm::enumerate(targetMemRefType.getShape())) {
1589 int64_t index = indexedSize.index();
1590 desc.setSize(rewriter, loc, index,
1592 indexedSize.value()));
1593 desc.setStride(rewriter, loc, index,
1595 (*targetStrides)[index]));
1598 rewriter.replaceOp(castOp, {desc});
1605class VectorCreateMaskOpConversion
1606 :
public OpConversionPattern<vector::CreateMaskOp> {
1608 explicit VectorCreateMaskOpConversion(MLIRContext *context,
1609 bool enableIndexOpt)
1610 : OpConversionPattern<vector::CreateMaskOp>(context),
1611 force32BitVectorIndices(enableIndexOpt) {}
1614 matchAndRewrite(vector::CreateMaskOp op, OpAdaptor adaptor,
1615 ConversionPatternRewriter &rewriter)
const override {
1616 auto dstType = op.getType();
1617 if (dstType.getRank() != 1 || !cast<VectorType>(dstType).isScalable())
1619 IntegerType idxType =
1620 force32BitVectorIndices ? rewriter.getI32Type() : rewriter.getI64Type();
1621 auto loc = op->getLoc();
1622 Value
indices = LLVM::StepVectorOp::create(
1626 Value maskBound = adaptor.getOperands()[0];
1633 if (force32BitVectorIndices) {
1636 maskBound = arith::MinSIOp::create(rewriter, loc, maskBound, maxBound);
1640 Value bounds = BroadcastOp::create(rewriter, loc,
indices.getType(), bound);
1641 Value comp = arith::CmpIOp::create(rewriter, loc, arith::CmpIPredicate::slt,
1643 rewriter.replaceOp(op, comp);
1648 const bool force32BitVectorIndices;
1652 SymbolTableCollection *symbolTables =
nullptr;
1655 explicit VectorPrintOpConversion(
1656 const LLVMTypeConverter &typeConverter,
1657 SymbolTableCollection *symbolTables =
nullptr)
1658 : ConvertOpToLLVMPattern<vector::PrintOp>(typeConverter),
1659 symbolTables(symbolTables) {}
1675 matchAndRewrite(vector::PrintOp
printOp, OpAdaptor adaptor,
1676 ConversionPatternRewriter &rewriter)
const override {
1677 auto parent =
printOp->getParentOfType<ModuleOp>();
1683 if (
auto value = adaptor.getSource()) {
1685 if (isa<VectorType>(printType)) {
1689 if (
failed(emitScalarPrint(rewriter, parent, loc, printType, value)))
1693 auto punct =
printOp.getPunctuation();
1694 if (
auto stringLiteral =
printOp.getStringLiteral()) {
1697 *stringLiteral, *getTypeConverter(),
1699 if (createResult.failed())
1702 }
else if (punct != PrintPunctuation::NoPunctuation) {
1703 FailureOr<LLVM::LLVMFuncOp> op = [&]() {
1705 case PrintPunctuation::Close:
1708 case PrintPunctuation::Open:
1711 case PrintPunctuation::Comma:
1714 case PrintPunctuation::NewLine:
1718 llvm_unreachable(
"unexpected punctuation");
1723 emitCall(rewriter,
printOp->getLoc(), op.value());
1731 enum class PrintConversion {
1740 LogicalResult emitScalarPrint(ConversionPatternRewriter &rewriter,
1741 ModuleOp parent, Location loc, Type printType,
1742 Value value)
const {
1743 if (typeConverter->convertType(printType) ==
nullptr)
1747 PrintConversion conversion = PrintConversion::None;
1748 FailureOr<Operation *> printer;
1754 conversion = PrintConversion::Bitcast16;
1757 conversion = PrintConversion::Bitcast16;
1761 }
else if (
auto intTy = dyn_cast<IntegerType>(printType)) {
1765 unsigned width = intTy.getWidth();
1766 if (intTy.isUnsigned()) {
1769 conversion = PrintConversion::ZeroExt64;
1776 assert(intTy.isSignless() || intTy.isSigned());
1781 conversion = PrintConversion::ZeroExt64;
1782 else if (width < 64)
1783 conversion = PrintConversion::SignExt64;
1790 }
else if (
auto floatTy = dyn_cast<FloatType>(printType)) {
1793 llvm::APFloatBase::SemanticsToEnum(floatTy.getFloatSemantics());
1794 Value semValue = LLVM::ConstantOp::create(
1795 rewriter, loc, rewriter.getI32Type(),
1796 rewriter.getIntegerAttr(rewriter.getI32Type(), sem));
1798 LLVM::ZExtOp::create(rewriter, loc, rewriter.getI64Type(), value);
1801 emitCall(rewriter, loc, printer.value(),
1810 switch (conversion) {
1811 case PrintConversion::ZeroExt64:
1812 value = arith::ExtUIOp::create(
1813 rewriter, loc, IntegerType::get(rewriter.getContext(), 64), value);
1815 case PrintConversion::SignExt64:
1816 value = arith::ExtSIOp::create(
1817 rewriter, loc, IntegerType::get(rewriter.getContext(), 64), value);
1819 case PrintConversion::Bitcast16:
1820 value = LLVM::BitcastOp::create(
1821 rewriter, loc, IntegerType::get(rewriter.getContext(), 16), value);
1823 case PrintConversion::None:
1826 emitCall(rewriter, loc, printer.value(), value);
1831 static void emitCall(ConversionPatternRewriter &rewriter, Location loc,
1833 LLVM::CallOp::create(rewriter, loc,
TypeRange(), SymbolRefAttr::get(ref),
1841struct VectorBroadcastScalarToLowRankLowering
1843 using ConvertOpToLLVMPattern<vector::BroadcastOp>::ConvertOpToLLVMPattern;
1846 matchAndRewrite(vector::BroadcastOp
broadcast, OpAdaptor adaptor,
1847 ConversionPatternRewriter &rewriter)
const override {
1848 if (isa<VectorType>(
broadcast.getSourceType()))
1849 return rewriter.notifyMatchFailure(
1850 broadcast,
"broadcast from vector type not handled");
1853 if (resultType.getRank() > 1)
1854 return rewriter.notifyMatchFailure(
broadcast,
1855 "broadcast to 2+-d handled elsewhere");
1861 auto zero = LLVM::ConstantOp::create(
1863 typeConverter->convertType(rewriter.getIntegerType(32)),
1864 rewriter.getZeroAttr(rewriter.getIntegerType(32)));
1867 if (resultType.getRank() == 0) {
1868 rewriter.replaceOpWithNewOp<LLVM::InsertElementOp>(
1869 broadcast, vectorType, poison, adaptor.getSource(), zero);
1874 LLVM::InsertElementOp::create(rewriter,
broadcast.
getLoc(), vectorType,
1875 poison, adaptor.getSource(), zero);
1879 SmallVector<int32_t> zeroValues(width, 0);
1882 auto shuffle = rewriter.createOrFold<LLVM::ShuffleVectorOp>(
1893struct VectorBroadcastScalarToNdLowering
1895 using ConvertOpToLLVMPattern<BroadcastOp>::ConvertOpToLLVMPattern;
1898 matchAndRewrite(BroadcastOp
broadcast, OpAdaptor adaptor,
1899 ConversionPatternRewriter &rewriter)
const override {
1900 if (isa<VectorType>(
broadcast.getSourceType()))
1901 return rewriter.notifyMatchFailure(
1902 broadcast,
"broadcast from vector type not handled");
1905 if (resultType.getRank() <= 1)
1906 return rewriter.notifyMatchFailure(
1907 broadcast,
"broadcast to 1-d or 0-d handled elsewhere");
1911 auto vectorTypeInfo =
1913 auto llvmNDVectorTy = vectorTypeInfo.llvmNDVectorTy;
1914 auto llvm1DVectorTy = vectorTypeInfo.llvm1DVectorTy;
1915 if (!llvmNDVectorTy || !llvm1DVectorTy)
1919 Value desc = LLVM::PoisonOp::create(rewriter, loc, llvmNDVectorTy);
1923 Value vdesc = LLVM::PoisonOp::create(rewriter, loc, llvm1DVectorTy);
1924 auto zero = LLVM::ConstantOp::create(
1925 rewriter, loc, typeConverter->convertType(rewriter.getIntegerType(32)),
1926 rewriter.getZeroAttr(rewriter.getIntegerType(32)));
1927 Value v = LLVM::InsertElementOp::create(rewriter, loc, llvm1DVectorTy,
1928 vdesc, adaptor.getSource(), zero);
1931 int64_t width = resultType.getDimSize(resultType.getRank() - 1);
1932 SmallVector<int32_t> zeroValues(width, 0);
1933 v = LLVM::ShuffleVectorOp::create(rewriter, loc, v, v, zeroValues);
1937 nDVectorIterate(vectorTypeInfo, rewriter, [&](ArrayRef<int64_t> position) {
1938 desc = LLVM::InsertValueOp::create(rewriter, loc, desc, v, position);
1947struct VectorInterleaveOpLowering
1952 matchAndRewrite(vector::InterleaveOp interleaveOp, OpAdaptor adaptor,
1953 ConversionPatternRewriter &rewriter)
const override {
1954 VectorType resultType = interleaveOp.getResultVectorType();
1956 if (resultType.getRank() != 1)
1957 return rewriter.notifyMatchFailure(interleaveOp,
1958 "InterleaveOp not rank 1");
1960 if (resultType.isScalable()) {
1961 rewriter.replaceOpWithNewOp<LLVM::vector_interleave2>(
1962 interleaveOp, typeConverter->convertType(resultType),
1963 adaptor.getLhs(), adaptor.getRhs());
1970 int64_t resultVectorSize = resultType.getNumElements();
1971 SmallVector<int32_t> interleaveShuffleMask;
1972 interleaveShuffleMask.reserve(resultVectorSize);
1973 for (
int i = 0, end = resultVectorSize / 2; i < end; ++i) {
1974 interleaveShuffleMask.push_back(i);
1975 interleaveShuffleMask.push_back((resultVectorSize / 2) + i);
1977 rewriter.replaceOpWithNewOp<LLVM::ShuffleVectorOp>(
1978 interleaveOp, adaptor.getLhs(), adaptor.getRhs(),
1979 interleaveShuffleMask);
1986struct VectorDeinterleaveOpLowering
1991 matchAndRewrite(vector::DeinterleaveOp deinterleaveOp, OpAdaptor adaptor,
1992 ConversionPatternRewriter &rewriter)
const override {
1993 VectorType resultType = deinterleaveOp.getResultVectorType();
1994 VectorType sourceType = deinterleaveOp.getSourceVectorType();
1995 auto loc = deinterleaveOp.getLoc();
1999 if (resultType.getRank() != 1)
2000 return rewriter.notifyMatchFailure(deinterleaveOp,
2001 "DeinterleaveOp not rank 1");
2003 if (resultType.isScalable()) {
2004 const auto *llvmTypeConverter = this->getTypeConverter();
2005 auto deinterleaveResults = deinterleaveOp.getResultTypes();
2006 auto packedOpResults =
2007 llvmTypeConverter->packOperationResults(deinterleaveResults);
2008 auto intrinsic = LLVM::vector_deinterleave2::create(
2009 rewriter, loc, packedOpResults, adaptor.getSource());
2011 auto evenResult = LLVM::ExtractValueOp::create(
2012 rewriter, loc, intrinsic->getResult(0), 0);
2013 auto oddResult = LLVM::ExtractValueOp::create(rewriter, loc,
2014 intrinsic->getResult(0), 1);
2016 rewriter.replaceOp(deinterleaveOp,
ValueRange{evenResult, oddResult});
2023 int64_t resultVectorSize = resultType.getNumElements();
2024 SmallVector<int32_t> evenShuffleMask;
2025 SmallVector<int32_t> oddShuffleMask;
2027 evenShuffleMask.reserve(resultVectorSize);
2028 oddShuffleMask.reserve(resultVectorSize);
2030 for (
int i = 0; i < sourceType.getNumElements(); ++i) {
2032 evenShuffleMask.push_back(i);
2034 oddShuffleMask.push_back(i);
2037 auto poison = LLVM::PoisonOp::create(rewriter, loc, sourceType);
2038 auto evenShuffle = LLVM::ShuffleVectorOp::create(
2039 rewriter, loc, adaptor.getSource(), poison, evenShuffleMask);
2040 auto oddShuffle = LLVM::ShuffleVectorOp::create(
2041 rewriter, loc, adaptor.getSource(), poison, oddShuffleMask);
2043 rewriter.replaceOp(deinterleaveOp,
ValueRange{evenShuffle, oddShuffle});
2049struct VectorFromElementsLowering
2054 matchAndRewrite(vector::FromElementsOp fromElementsOp, OpAdaptor adaptor,
2055 ConversionPatternRewriter &rewriter)
const override {
2056 Location loc = fromElementsOp.getLoc();
2057 VectorType vectorType = fromElementsOp.getType();
2061 if (vectorType.getRank() > 1)
2062 return rewriter.notifyMatchFailure(fromElementsOp,
2063 "rank > 1 vectors are not supported");
2064 Type llvmType = typeConverter->convertType(vectorType);
2065 Type llvmIndexType = typeConverter->convertType(rewriter.getIndexType());
2066 Value
result = LLVM::PoisonOp::create(rewriter, loc, llvmType);
2067 for (
auto [idx, val] : llvm::enumerate(adaptor.getElements())) {
2069 LLVM::ConstantOp::create(rewriter, loc, llvmIndexType, idx);
2070 result = LLVM::InsertElementOp::create(rewriter, loc, llvmType,
result,
2073 rewriter.replaceOp(fromElementsOp,
result);
2079struct VectorToElementsLowering
2084 matchAndRewrite(vector::ToElementsOp toElementsOp, OpAdaptor adaptor,
2085 ConversionPatternRewriter &rewriter)
const override {
2086 Location loc = toElementsOp.getLoc();
2087 auto idxType = typeConverter->convertType(rewriter.getIndexType());
2088 Value source = adaptor.getSource();
2090 SmallVector<Value> results(toElementsOp->getNumResults());
2091 for (
auto [idx, element] : llvm::enumerate(toElementsOp.getElements())) {
2093 if (element.use_empty())
2096 auto constIdx = LLVM::ConstantOp::create(
2097 rewriter, loc, idxType, rewriter.getIntegerAttr(idxType, idx));
2098 auto llvmType = typeConverter->convertType(element.getType());
2100 Value
result = LLVM::ExtractElementOp::create(rewriter, loc, llvmType,
2105 rewriter.replaceOp(toElementsOp, results);
2115 matchAndRewrite(vector::StepOp stepOp, OpAdaptor adaptor,
2116 ConversionPatternRewriter &rewriter)
const override {
2117 Type llvmType = typeConverter->convertType(stepOp.getType());
2118 rewriter.replaceOpWithNewOp<LLVM::StepVectorOp>(stepOp, llvmType);
2133class ContractionOpToMatmulOpLowering
2136 using MaskableOpRewritePattern::MaskableOpRewritePattern;
2138 ContractionOpToMatmulOpLowering(MLIRContext *context,
2139 PatternBenefit benefit = 100)
2140 : MaskableOpRewritePattern<vector::ContractionOp>(context, benefit) {}
2143 matchAndRewriteMaskableOp(vector::ContractionOp op, MaskingOpInterface maskOp,
2144 PatternRewriter &rewriter)
const override;
2164FailureOr<Value> ContractionOpToMatmulOpLowering::matchAndRewriteMaskableOp(
2165 vector::ContractionOp op, MaskingOpInterface maskOp,
2171 auto iteratorTypes = op.getIteratorTypes().getValue();
2177 Type opResType = op.getType();
2178 VectorType vecType = dyn_cast<VectorType>(opResType);
2179 if (vecType && vecType.isScalable()) {
2184 Type elementType = op.getLhsType().getElementType();
2188 Type dstElementType = vecType ? vecType.getElementType() : opResType;
2189 if (elementType != dstElementType)
2194 MLIRContext *ctx = op.getContext();
2195 Location loc = op.getLoc();
2199 Value
lhs = op.getLhs();
2200 auto lhsMap = op.getIndexingMapsArray()[0];
2202 lhs = vector::TransposeOp::create(rew, loc,
lhs, ArrayRef<int64_t>{1, 0});
2207 Value
rhs = op.getRhs();
2208 auto rhsMap = op.getIndexingMapsArray()[1];
2210 rhs = vector::TransposeOp::create(rew, loc,
rhs, ArrayRef<int64_t>{1, 0});
2215 VectorType lhsType = cast<VectorType>(
lhs.getType());
2216 VectorType rhsType = cast<VectorType>(
rhs.getType());
2217 int64_t lhsRows = lhsType.getDimSize(0);
2218 int64_t lhsColumns = lhsType.getDimSize(1);
2219 int64_t rhsColumns = rhsType.getDimSize(1);
2221 Type flattenedLHSType =
2222 VectorType::get(lhsType.getNumElements(), lhsType.getElementType());
2223 lhs = vector::ShapeCastOp::create(rew, loc, flattenedLHSType,
lhs);
2225 Type flattenedRHSType =
2226 VectorType::get(rhsType.getNumElements(), rhsType.getElementType());
2227 rhs = vector::ShapeCastOp::create(rew, loc, flattenedRHSType,
rhs);
2229 Value
mul = LLVM::MatrixMultiplyOp::create(
2231 VectorType::get(lhsRows * rhsColumns,
2232 cast<VectorType>(
lhs.getType()).getElementType()),
2233 lhs,
rhs, lhsRows, lhsColumns, rhsColumns);
2235 mul = vector::ShapeCastOp::create(
2237 VectorType::get({lhsRows, rhsColumns},
2242 auto accMap = op.getIndexingMapsArray()[2];
2244 mul = vector::TransposeOp::create(rew, loc,
mul, ArrayRef<int64_t>{1, 0});
2246 llvm_unreachable(
"invalid contraction semantics");
2248 Value res = isa<IntegerType>(elementType)
2249 ?
static_cast<Value
>(
2250 arith::AddIOp::create(rew, loc, op.getAcc(),
mul))
2251 : static_cast<Value>(
2252 arith::AddFOp::create(rew, loc, op.getAcc(),
mul));
2270class TransposeOpToMatrixTransposeOpLowering
2271 :
public OpRewritePattern<vector::TransposeOp> {
2275 LogicalResult matchAndRewrite(vector::TransposeOp op,
2276 PatternRewriter &rewriter)
const override {
2277 auto loc = op.getLoc();
2279 Value input = op.getVector();
2280 VectorType inputType = op.getSourceVectorType();
2281 VectorType resType = op.getResultVectorType();
2283 if (inputType.isScalable())
2285 op,
"This lowering does not support scalable vectors");
2288 ArrayRef<int64_t> transp = op.getPermutation();
2290 if (resType.getRank() != 2 || transp[0] != 1 || transp[1] != 0) {
2294 Type flattenedType =
2295 VectorType::get(resType.getNumElements(), resType.getElementType());
2297 vector::ShapeCastOp::create(rewriter, loc, flattenedType, input);
2300 Value trans = LLVM::MatrixTransposeOp::create(rewriter, loc, flattenedType,
2301 matrix, rows, columns);
2311 patterns.
add<VectorFMAOpNDRewritePattern>(patterns.
getContext());
2316 patterns.
add<ContractionOpToMatmulOpLowering>(patterns.
getContext(), benefit);
2321 patterns.
add<TransposeOpToMatrixTransposeOpLowering>(patterns.
getContext(),
2328 bool reassociateFPReductions,
bool force32BitVectorIndices,
2329 bool useVectorAlignment,
bool enableGEPInboundsNuw) {
2332 patterns.
add<VectorReductionOpConversion>(converter, reassociateFPReductions);
2333 patterns.
add<VectorCreateMaskOpConversion>(ctx, force32BitVectorIndices);
2334 patterns.
add<VectorLoadStoreConversion<vector::LoadOp>,
2335 VectorLoadStoreConversion<vector::MaskedLoadOp>,
2336 VectorLoadStoreConversion<vector::StoreOp>,
2337 VectorLoadStoreConversion<vector::MaskedStoreOp>>(
2338 converter, useVectorAlignment, enableGEPInboundsNuw);
2339 patterns.
add<VectorGatherOpConversion, VectorScatterOpConversion>(
2340 converter, useVectorAlignment);
2341 patterns.
add<VectorBitCastOpConversion, VectorShuffleOpConversion,
2342 VectorExtractOpConversion, VectorFMAOp1DConversion,
2343 VectorInsertOpConversion, VectorPrintOpConversion,
2344 VectorTypeCastOpConversion, VectorScaleOpConversion,
2345 VectorExpandLoadOpConversion, VectorCompressStoreOpConversion,
2346 VectorBroadcastScalarToLowRankLowering,
2347 VectorBroadcastScalarToNdLowering,
2348 VectorScalableInsertOpLowering, VectorScalableExtractOpLowering,
2349 MaskedReductionOpConversion, VectorInterleaveOpLowering,
2350 VectorDeinterleaveOpLowering, VectorFromElementsLowering,
2351 VectorToElementsLowering, VectorStepOpLowering>(converter);
2355struct VectorToLLVMDialectInterface :
public ConvertToLLVMPatternInterface {
2356 VectorToLLVMDialectInterface(
Dialect *dialect)
2357 : ConvertToLLVMPatternInterface(dialect) {}
2359 using ConvertToLLVMPatternInterface::ConvertToLLVMPatternInterface;
2360 void loadDependentDialects(MLIRContext *context)
const final {
2361 context->loadDialect<LLVM::LLVMDialect>();
2366 void populateConvertToLLVMConversionPatterns(
2367 ConversionTarget &
target, LLVMTypeConverter &typeConverter,
2368 RewritePatternSet &patterns)
const final {
2377 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)
static DenseElementsAttr get(ShapedType type, ArrayRef< Attribute > values)
Constructs a dense elements attribute from an array of element values.
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)
FailureOr< LLVM::LLVMFuncOp > lookupOrCreatePrintBF16Fn(OpBuilder &b, Operation *moduleOp, SymbolTableCollection *symbolTables=nullptr)
FailureOr< LLVM::LLVMFuncOp > lookupOrCreatePrintOpenFn(OpBuilder &b, Operation *moduleOp, SymbolTableCollection *symbolTables=nullptr)
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.
FailureOr< LLVM::LLVMFuncOp > lookupOrCreatePrintCommaFn(OpBuilder &b, Operation *moduleOp, SymbolTableCollection *symbolTables=nullptr)
FailureOr< LLVM::LLVMFuncOp > lookupOrCreatePrintI64Fn(OpBuilder &b, Operation *moduleOp, SymbolTableCollection *symbolTables=nullptr)
Helper functions to look up or create the declaration for commonly used external C function calls.
Value createIndexAttrConstant(OpBuilder &builder, Location loc, Type resultType, int64_t value)
Creates an llvm.mlir.constant producing value as resultType, which is expected to be the converted in...
FailureOr< LLVM::LLVMFuncOp > lookupOrCreatePrintNewlineFn(OpBuilder &b, Operation *moduleOp, SymbolTableCollection *symbolTables=nullptr)
FailureOr< LLVM::LLVMFuncOp > lookupOrCreatePrintCloseFn(OpBuilder &b, Operation *moduleOp, SymbolTableCollection *symbolTables=nullptr)
FailureOr< LLVM::LLVMFuncOp > lookupOrCreatePrintU64Fn(OpBuilder &b, Operation *moduleOp, SymbolTableCollection *symbolTables=nullptr)
FailureOr< LLVM::LLVMFuncOp > lookupOrCreatePrintF32Fn(OpBuilder &b, Operation *moduleOp, SymbolTableCollection *symbolTables=nullptr)
FailureOr< LLVM::LLVMFuncOp > lookupOrCreateApFloatPrintFn(OpBuilder &b, Operation *moduleOp, SymbolTableCollection *symbolTables=nullptr)
LogicalResult createPrintStrCall(OpBuilder &builder, Location loc, ModuleOp moduleOp, StringRef symbolName, StringRef string, const LLVMTypeConverter &typeConverter, bool addNewline=true, std::optional< StringRef > runtimeFunctionName={}, SymbolTableCollection *symbolTables=nullptr)
Generate IR that prints the given string to stdout.
FailureOr< LLVM::LLVMFuncOp > lookupOrCreatePrintF16Fn(OpBuilder &b, Operation *moduleOp, SymbolTableCollection *symbolTables=nullptr)
FailureOr< LLVM::LLVMFuncOp > lookupOrCreatePrintF64Fn(OpBuilder &b, Operation *moduleOp, SymbolTableCollection *symbolTables=nullptr)
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.