30#include "llvm/ADT/TypeSwitch.h"
36#define GEN_PASS_DEF_CONVERTVECTORTOXEGPU
37#include "mlir/Conversion/Passes.h.inc"
45static bool isZeroConstant(
Value val) {
51 .Case([](FloatAttr floatAttr) {
return floatAttr.getValue().isZero(); })
52 .Case([](IntegerAttr intAttr) {
return intAttr.getValue().isZero(); })
61static bool isZeroOrPoisonPadding(
Value val) {
62 return isZeroConstant(val) || val.
getDefiningOp<ub::PoisonOp>();
72static bool isInnermostTwoDimsTransposed(
AffineMap map) {
79 for (
unsigned i = 0; i + 2 < numResults; ++i)
102 bool hasTranspose =
false) {
103 if (
shape.size() < 2)
108 const auto *blockInst = dyn_cast<xegpu::uArch::BlockIOInstructionInterface>(
113 int width =
static_cast<int>(
shape.back());
114 int height =
static_cast<int>(
shape[
shape.size() - 2]);
120 auto fitsBlockShapes = [&](
bool hasTransform) {
121 std::optional<xegpu::uArch::BlockIOInstructionInterface::BlockShapes>
122 blockShapes = blockInst->getBlockWidthHeightCount(elemTy, hasTransform,
126 auto [widths, heights, counts] = *blockShapes;
135 return fitsBlockShapes(
false) ||
136 fitsBlockShapes(
true);
140 VectorTransferOpInterface xferOp) {
141 if (xferOp.getMask())
143 "Masked transfer is not supported");
145 auto srcTy = dyn_cast<MemRefType>(xferOp.getShapedType());
152 if (
failed(srcTy.getStridesAndOffset(strides, offset)))
154 "The memref strides cannot be inferred");
157 if (strides.back() != 1)
159 xferOp,
"Buffer must be contiguous in the innermost dimension");
161 VectorType vecTy = xferOp.getVectorType();
162 unsigned vecRank = vecTy.getRank();
171 auto dim = dyn_cast<AffineDimExpr>(expr);
172 if (dim.getPosition() < (numInputDims - vecRank))
174 xferOp,
"Only the innermost dimensions can be accessed");
197static void adjustStridesForPermutation(
AffineMap permMap,
212 typename = std::enable_if_t<llvm::is_one_of<
213 std::decay_t<OpType>, vector::TransferReadOp, vector::TransferWriteOp,
214 vector::GatherOp, vector::ScatterOp>::value>>
215static std::pair<SmallVector<Value>,
Value>
218 Value baseMemref = xferOp.getBase();
219 MemRefType memrefType = dyn_cast<MemRefType>(baseMemref.
getType());
222 Value offsetVal =
nullptr;
223 if (memrefType.hasStaticShape()) {
226 if (
failed(memrefType.getStridesAndOffset(intStrides, offset)))
227 return {{}, offsetVal};
228 bool hasDynamicStrides = llvm::any_of(intStrides, [](
int64_t strideVal) {
229 return ShapedType::isDynamic(strideVal);
232 if (!hasDynamicStrides)
236 if (!ShapedType::isDynamic(offset))
240 if (strides.empty() || !offsetVal) {
243 unsigned rank = memrefType.getRank();
249 resultTypes.push_back(MemRefType::get(
250 {}, memrefType.getElementType()));
251 resultTypes.push_back(indexType);
253 for (
unsigned i = 0; i < rank; ++i)
254 resultTypes.push_back(indexType);
256 for (
unsigned i = 0; i < rank; ++i)
257 resultTypes.push_back(indexType);
259 auto meta = memref::ExtractStridedMetadataOp::create(
260 rewriter, loc, resultTypes, baseMemref);
263 strides.append(meta.getStrides().begin(), meta.getStrides().end());
266 offsetVal = meta.getOffset();
271 return {strides, offsetVal};
276static Value computeBaseOffset(VectorTransferOpInterface xferOp,
280 for (
auto [
index, stride] : llvm::zip_equal(xferOp.getIndices(), strides)) {
281 Value contrib = arith::MulIOp::create(rewriter, loc,
index, stride);
282 baseOffset = arith::AddIOp::create(rewriter, loc, baseOffset, contrib);
315static Value computeOffsets(VectorTransferOpInterface xferOp,
319 VectorType vectorType = xferOp.getVectorType();
325 auto stepType = VectorType::get({dim}, rewriter.
getIndexType());
326 auto stepOp = vector::StepOp::create(rewriter, loc, stepType);
327 stepVectors.push_back(stepOp);
334 adjustStridesForPermutation(xferOp.getPermutationMap(), permutedStrides);
337 size_t memrefRank = permutedStrides.size();
340 for (
size_t i = 0; i < vectorRank; ++i) {
341 size_t memrefDim = memrefRank - vectorRank + i;
342 Value strideValue = permutedStrides[memrefDim];
343 auto mulType = dyn_cast<VectorType>(stepVectors[i].
getType());
345 vector::BroadcastOp::create(rewriter, loc, mulType, strideValue);
346 auto mulOp = arith::MulIOp::create(rewriter, loc, stepVectors[i], bcastOp);
347 strideMultiplied.push_back(mulOp);
352 for (
size_t i = 0; i < vectorRank; ++i) {
355 auto newType = VectorType::get(newShape, rewriter.
getIndexType());
356 auto castOp = vector::ShapeCastOp::create(rewriter, loc, newType,
357 strideMultiplied[i]);
358 shapeCasted.push_back(castOp);
363 auto fullIndexVectorType =
365 for (
Value shapeCastVal : shapeCasted) {
366 auto broadcastOp = vector::BroadcastOp::create(
367 rewriter, loc, fullIndexVectorType, shapeCastVal);
368 broadcasted.push_back(broadcastOp);
372 Value localOffsets = broadcasted[0];
373 for (
size_t i = 1; i < broadcasted.size(); ++i)
375 arith::AddIOp::create(rewriter, loc, localOffsets, broadcasted[i]);
378 baseOffset = computeBaseOffset(xferOp, rewriter, strides, baseOffset);
379 Value bcastBase = vector::BroadcastOp::create(
380 rewriter, loc, fullIndexVectorType, baseOffset);
381 localOffsets = arith::AddIOp::create(rewriter, loc, bcastBase, localOffsets);
386static Value getMemrefDimSize(VectorTransferOpInterface xferOp,
unsigned dim,
389 auto memrefTy = cast<MemRefType>(xferOp.getShapedType());
390 if (memrefTy.isDynamicDim(dim))
391 return memref::DimOp::create(rewriter, loc, xferOp.getBase(), dim)
412static Value computeInBoundsMask(VectorTransferOpInterface xferOp,
421 for (
unsigned v = 0, e =
vectorShape.size(); v < e; ++v) {
422 if (xferOp.isDimInBounds(v))
424 unsigned d = cast<AffineDimExpr>(map.
getResult(v)).getPosition();
425 Value bound = getMemrefDimSize(xferOp, d, rewriter);
428 Value limit = arith::SubIOp::create(rewriter, loc, bound,
indices[d]);
430 Value step = vector::StepOp::create(rewriter, loc, stepType);
432 vector::BroadcastOp::create(rewriter, loc, stepType, limit);
433 Value dimMask = arith::CmpIOp::create(
434 rewriter, loc, arith::CmpIPredicate::slt, step, limitVec);
438 dimMask = vector::ShapeCastOp::create(
439 rewriter, loc, VectorType::get(expandedShape, rewriter.
getI1Type()),
441 dimMask = vector::BroadcastOp::create(rewriter, loc, maskType, dimMask);
444 ? arith::AndIOp::create(rewriter, loc, mask, dimMask).getResult()
449 return vector::ConstantMaskOp::create(rewriter, loc, maskType,
vectorShape);
460static Value computeUnitInBoundsMask(VectorTransferOpInterface xferOp,
467 for (
unsigned v = 0, e = xferOp.getVectorType().getRank(); v < e; ++v) {
468 if (xferOp.isDimInBounds(v))
470 unsigned d = cast<AffineDimExpr>(map.
getResult(v)).getPosition();
471 Value bound = getMemrefDimSize(xferOp, d, rewriter);
472 Value dimMask = arith::CmpIOp::create(
473 rewriter, loc, arith::CmpIPredicate::slt,
indices[d], bound);
475 ? arith::AndIOp::create(rewriter, loc, mask, dimMask).getResult()
480 return arith::ConstantOp::create(rewriter, loc, rewriter.
getBoolAttr(
true));
490 typename = std::enable_if_t<llvm::is_one_of<
491 std::decay_t<OpType>, vector::GatherOp, vector::ScatterOp>::value>>
496 for (
size_t i = 0; i < offsets.size(); ++i) {
497 Value offsetContrib =
498 arith::MulIOp::create(rewriter, loc, offsets[i], strides[i]);
500 arith::AddIOp::create(rewriter, loc, baseOffset, offsetContrib);
503 VectorType vecType = cast<VectorType>(
indices.getType());
506 vector::BroadcastOp::create(rewriter, loc, vecType, strides.back())
508 Value stridedIndices =
509 arith::MulIOp::create(rewriter, loc, strideVector,
indices).getResult();
512 vector::BroadcastOp::create(
514 VectorType::get(vecType.getShape(), rewriter.
getIndexType()),
517 return arith::AddIOp::create(rewriter, loc, baseVector, stridedIndices)
526static std::pair<Value, SmallVector<OpFoldResult>>
531 auto memrefType = cast<MemRefType>(
memref.getType());
532 unsigned rank = memrefType.getRank();
534 if (rank <= targetRank)
537 int64_t numCombinedDims = rank - targetRank;
543 for (
unsigned i = 0; i < numCombinedDims; ++i) {
544 subviewOffsets.push_back(offsets[i]);
551 auto originalShape = memrefType.getShape();
552 auto meta = memref::ExtractStridedMetadataOp::create(rewriter, loc,
memref);
553 for (
unsigned i = numCombinedDims; i < rank; ++i) {
555 if (ShapedType::isDynamic(originalShape[i])) {
556 subviewSizes.push_back(meta.getSizes()[i]);
557 resultShape.push_back(ShapedType::kDynamic);
560 resultShape.push_back(originalShape[i]);
565 auto resultType = memref::SubViewOp::inferRankReducedResultType(
566 resultShape, memrefType, subviewOffsets, subviewSizes, subviewStrides);
568 memref::SubViewOp::create(rewriter, loc, resultType,
memref,
569 subviewOffsets, subviewSizes, subviewStrides);
574 return {subviewOp.getResult(), newOffsets};
579 typename = std::enable_if_t<llvm::is_one_of<
580 std::decay_t<OpType>, vector::TransferReadOp, vector::TransferWriteOp,
581 vector::GatherOp, vector::ScatterOp>::value>>
585 auto indexPtr = memref::ExtractAlignedPointerAsIndexOp::create(
586 rewriter, loc, xferOp.getBase())
588 return arith::IndexCastOp::create(rewriter, loc, rewriter.
getI64Type(),
595static bool isUsedAsScalar(
Value vec) {
599 auto extractOp = dyn_cast<vector::ExtractOp>(user);
600 return extractOp && !isa<VectorType>(extractOp.getResult().getType());
614static LogicalResult lowerToScalarLoadOp(vector::TransferReadOp readOp,
617 VectorType vectorType = readOp.getVectorType();
618 if (!isa<MemRefType>(readOp.getShapedType()))
621 auto meta = computeMemrefMeta(readOp, rewriter);
622 if (meta.first.empty())
625 Value offset = computeBaseOffset(readOp, rewriter, meta.first, meta.second);
626 Value flatMemref = memrefToIndexPtr(readOp, rewriter);
628 Value mask = computeUnitInBoundsMask(readOp, rewriter);
629 auto loadOp = xegpu::LoadGatherOp::create(
630 rewriter, loc, vectorType.getElementType(), flatMemref, offset, mask,
631 xegpu::CachePolicyAttr{},
632 xegpu::CachePolicyAttr{},
633 xegpu::CachePolicyAttr{},
638 Value scalar = loadOp.getResult();
639 if (readOp.hasOutOfBoundsDim() &&
640 !readOp.getPadding().getDefiningOp<ub::PoisonOp>())
641 scalar = arith::SelectOp::create(rewriter, loc, mask, scalar,
642 readOp.getPadding());
648static LogicalResult lowerToScatteredLoadOp(vector::TransferReadOp readOp,
652 VectorType vectorType = readOp.getVectorType();
653 auto memrefType = dyn_cast<MemRefType>(readOp.getShapedType());
657 auto meta = computeMemrefMeta(readOp, rewriter);
658 if (meta.first.empty())
662 computeOffsets(readOp, rewriter, meta.first, meta.second);
664 Value flatMemref = memrefToIndexPtr(readOp, rewriter);
666 Value mask = computeInBoundsMask(readOp, rewriter);
667 auto gatherOp = xegpu::LoadGatherOp::create(
668 rewriter, loc, vectorType, flatMemref, localOffsets, mask,
669 xegpu::CachePolicyAttr{},
670 xegpu::CachePolicyAttr{},
671 xegpu::CachePolicyAttr{},
678 if (readOp.hasOutOfBoundsDim() &&
679 !readOp.getPadding().getDefiningOp<ub::PoisonOp>()) {
680 Value padding = vector::BroadcastOp::create(rewriter, loc, vectorType,
681 readOp.getPadding());
682 result = arith::SelectOp::create(rewriter, loc, mask,
result, padding);
689static LogicalResult lowerToScatteredStoreOp(vector::TransferWriteOp writeOp,
693 auto memrefType = dyn_cast<MemRefType>(writeOp.getShapedType());
697 auto meta = computeMemrefMeta(writeOp, rewriter);
698 if (meta.first.empty())
702 computeOffsets(writeOp, rewriter, meta.first, meta.second);
704 Value flatMemref = memrefToIndexPtr(writeOp, rewriter);
708 Value mask = computeInBoundsMask(writeOp, rewriter);
709 xegpu::StoreScatterOp::create(rewriter, loc, writeOp.getVector(), flatMemref,
711 xegpu::CachePolicyAttr{},
712 xegpu::CachePolicyAttr{},
713 xegpu::CachePolicyAttr{},
719struct TransferReadLowering :
public OpRewritePattern<vector::TransferReadOp> {
722 LogicalResult matchAndRewrite(vector::TransferReadOp readOp,
723 PatternRewriter &rewriter)
const override {
724 Location loc = readOp.getLoc();
726 if (
failed(transferPreconditions(rewriter, readOp)))
728 auto readMemTy = cast<MemRefType>(readOp.getShapedType());
729 VectorType loadedVecTy = readOp.getVectorType();
730 bool isOutOfBounds = readOp.hasOutOfBoundsDim();
732 bool isSharedMemory = xegpu::XeGPUDialect::isSharedMemory(readMemTy);
736 if (loadedVecTy.getRank() != 1 && loadedVecTy.getRank() != 2)
738 readOp,
"Only 1D and 2D vector loads are supported for SLM");
739 AffineMap readMap = readOp.getPermutationMap();
743 "Non identity transposition is not supported for SLM loads.");
747 readOp,
"Out-of-bounds access is not supported for SLM loads");
751 xegpu::MemDescType::get(rewriter.
getContext(), readMemTy.getShape(),
752 readMemTy.getElementType(),
754 auto createMemDescOp = xegpu::CreateMemDescOp::create(
755 rewriter, loc, memDescType, readOp.getBase());
757 SmallVector<OpFoldResult>
indices =
759 auto loadMatrixOp = xegpu::LoadMatrixOp::create(
760 rewriter, loc, loadedVecTy, createMemDescOp.getResult(),
indices,
763 rewriter.
replaceOp(readOp, loadMatrixOp.getResult());
771 if (loadedVecTy.getNumElements() == 1 && isUsedAsScalar(readOp.getResult()))
772 return lowerToScalarLoadOp(readOp, rewriter);
774 const xegpu::uArch::uArch *uArch =
782 bool isTransposeLoad = isInnermostTwoDimsTransposed(readMap);
786 Type elementType = loadedVecTy.getElementType();
787 SmallVector<int64_t> descShape(loadedVecTy.getShape());
788 if (isTransposeLoad) {
789 size_t rank = descShape.size();
790 assert(rank >= 2 &&
"Transpose requires at least 2 dimensions");
791 std::swap(descShape[rank - 1], descShape[rank - 2]);
800 bool canLowerToLoadNd =
801 loadedVecTy.getRank() > 1 &&
803 readMemTy.getElementType().isIntOrFloat() &&
804 (!isOutOfBounds || isZeroOrPoisonPadding(readOp.getPadding())) &&
805 isSupportedBlockShape(
806 uArch, xegpu::uArch::InstructionKind::Subgroup2DBlockLoad,
807 descShape, elementType, isTransposeLoad);
809 if (canLowerToLoadNd) {
813 loadedVecTy = VectorType::get(descShape, elementType);
814 auto descType = xegpu::TensorDescType::get(
815 descShape, elementType, 1,
816 isOutOfBounds, xegpu::MemorySpace::Global);
817 auto [src,
indices] = convertMemrefAndOffsetsToTargetRank(
818 rewriter, loc, readOp.getBase(),
821 xegpu::CachePolicyAttr hint =
nullptr;
822 xegpu::CreateNdDescOp ndDesc = xegpu::CreateNdDescOp::create(
825 Operation *loadedOp =
826 xegpu::LoadNdOp::create(rewriter, loc, loadedVecTy, ndDesc,
indices,
831 if (isTransposeLoad) {
835 int64_t rank = loadedVecTy.getRank();
836 SmallVector<int64_t> perm(llvm::to_vector(llvm::seq<int64_t>(0, rank)));
837 std::swap(perm[rank - 1], perm[rank - 2]);
838 loadedOp = vector::TransposeOp::create(rewriter, loc,
847 return lowerToScatteredLoadOp(readOp, rewriter);
851struct TransferWriteLowering
855 LogicalResult matchAndRewrite(vector::TransferWriteOp writeOp,
856 PatternRewriter &rewriter)
const override {
857 Location loc = writeOp.getLoc();
859 if (
failed(transferPreconditions(rewriter, writeOp)))
862 VectorType vecTy = writeOp.getVectorType();
863 auto writeMemTy = cast<MemRefType>(writeOp.getShapedType());
865 bool isSharedMemory = xegpu::XeGPUDialect::isSharedMemory(writeMemTy);
871 if (vecTy.getRank() != 1 && vecTy.getRank() != 2)
873 writeOp,
"Only 1D and 2D vector stores are supported for SLM");
875 if (writeOp.hasOutOfBoundsDim())
877 writeOp,
"Out-of-bounds access is not supported for SLM stores");
880 xegpu::MemDescType::get(rewriter.
getContext(), writeMemTy.getShape(),
881 writeMemTy.getElementType(),
884 auto createMemDescOp = xegpu::CreateMemDescOp::create(
885 rewriter, loc, memDescType, writeOp.getBase());
888 SmallVector<OpFoldResult>
indices =
891 xegpu::StoreMatrixOp::create(rewriter, loc, writeOp.getVector(),
892 createMemDescOp.getResult(),
indices,
899 const xegpu::uArch::uArch *uArch =
909 bool canLowerToStoreNd =
911 writeMemTy.getElementType().isIntOrFloat() &&
912 isSupportedBlockShape(
913 uArch, xegpu::uArch::InstructionKind::Subgroup2DBlockStore,
914 vecTy.getShape(), vecTy.getElementType());
916 if (canLowerToStoreNd) {
917 auto [src,
indices] = convertMemrefAndOffsetsToTargetRank(
918 rewriter, loc, writeOp.getBase(),
921 auto descType = xegpu::TensorDescType::get(
922 vecTy.getShape(), vecTy.getElementType(),
923 1, writeOp.hasOutOfBoundsDim(),
924 xegpu::MemorySpace::Global);
926 xegpu::CachePolicyAttr hint =
nullptr;
927 xegpu::CreateNdDescOp ndDesc = xegpu::CreateNdDescOp::create(
930 auto storeOp = xegpu::StoreNdOp::create(
931 rewriter, loc, writeOp.getVector(), ndDesc,
indices,
941 return lowerToScatteredStoreOp(writeOp, rewriter);
948 LogicalResult matchAndRewrite(vector::GatherOp gatherOp,
949 PatternRewriter &rewriter)
const override {
950 auto srcTy = dyn_cast<MemRefType>(gatherOp.getBase().getType());
954 Location loc = gatherOp.getLoc();
955 VectorType vectorType = gatherOp.getVectorType();
957 auto meta = computeMemrefMeta(gatherOp, rewriter);
958 if (meta.first.empty())
962 computeOffsets(rewriter, gatherOp, meta.first, meta.second);
963 Value flatMemref = memrefToIndexPtr(gatherOp, rewriter);
965 auto xeGatherOp = xegpu::LoadGatherOp::create(
966 rewriter, loc, vectorType, flatMemref, localOffsets, gatherOp.getMask(),
967 xegpu::CachePolicyAttr{},
968 xegpu::CachePolicyAttr{},
969 xegpu::CachePolicyAttr{},
973 arith::SelectOp::create(rewriter, loc, gatherOp.getMask(),
974 xeGatherOp.getResult(), gatherOp.getPassThru());
975 rewriter.
replaceOp(gatherOp, selectOp.getResult());
983 LogicalResult matchAndRewrite(vector::ScatterOp scatterOp,
984 PatternRewriter &rewriter)
const override {
985 auto srcTy = dyn_cast<MemRefType>(scatterOp.getBase().getType());
989 Location loc = scatterOp.getLoc();
990 auto meta = computeMemrefMeta(scatterOp, rewriter);
991 if (meta.first.empty())
993 "Failed to compute strides");
996 computeOffsets(rewriter, scatterOp, meta.first, meta.second);
997 Value flatMemref = memrefToIndexPtr(scatterOp, rewriter);
999 xegpu::StoreScatterOp::create(rewriter, loc, scatterOp.getValueToStore(),
1000 flatMemref, localOffsets, scatterOp.getMask(),
1001 xegpu::CachePolicyAttr{},
1002 xegpu::CachePolicyAttr{},
1003 xegpu::CachePolicyAttr{},
1014 LogicalResult matchAndRewrite(vector::LoadOp loadOp,
1015 PatternRewriter &rewriter)
const override {
1016 Location loc = loadOp.getLoc();
1018 VectorType vecTy = loadOp.getResult().getType();
1019 MemRefType memTy = loadOp.getBase().getType();
1021 if (vecTy.getRank() != 1 && vecTy.getRank() != 2)
1023 if (!memTy.getElementType().isIntOrFloat())
1025 loadOp,
"Unsupported memref element type: expected integer or float");
1028 bool boundaryCheck = vecTy.getRank() > 1;
1030 xegpu::CachePolicyAttr hint =
nullptr;
1032 auto [src,
indices] = convertMemrefAndOffsetsToTargetRank(
1036 auto descType = xegpu::TensorDescType::get(
1037 vecTy.getShape(), vecTy.getElementType(), 1,
1038 boundaryCheck, xegpu::MemorySpace::Global);
1040 xegpu::CreateNdDescOp ndDesc = xegpu::CreateNdDescOp::create(
1043 xegpu::LoadNdOp::create(rewriter, loc, vecTy, ndDesc,
indices,
1057 LogicalResult matchAndRewrite(vector::StoreOp storeOp,
1058 PatternRewriter &rewriter)
const override {
1059 Location loc = storeOp.getLoc();
1062 VectorType vecTy = vector.getType();
1063 MemRefType memTy = storeOp.getBase().getType();
1065 if (vecTy.getRank() != 1 && vecTy.getRank() != 2)
1067 if (!memTy.getElementType().isIntOrFloat())
1070 "Unsupported memref element type: expected integer or float");
1073 bool boundaryCheck = vecTy.getRank() > 1;
1075 auto [src,
indices] = convertMemrefAndOffsetsToTargetRank(
1076 rewriter, loc, storeOp.getBase(),
1079 auto descType = xegpu::TensorDescType::get(
1080 vecTy.getShape(), vecTy.getElementType(),
1081 1, boundaryCheck, xegpu::MemorySpace::Global);
1084 xegpu::CachePolicyAttr hint =
nullptr;
1085 xegpu::CreateNdDescOp ndDesc = xegpu::CreateNdDescOp::create(
1089 xegpu::StoreNdOp::create(rewriter, loc, vector, ndDesc,
indices,
1104static std::optional<int64_t>
1105getRowMajorMatmulBatchRank(
ArrayAttr indexingMaps) {
1106 if (indexingMaps.size() != 3)
1107 return std::nullopt;
1109 AffineMap mapA = cast<AffineMapAttr>(indexingMaps[0]).getValue();
1110 AffineMap mapB = cast<AffineMapAttr>(indexingMaps[1]).getValue();
1111 AffineMap mapC = cast<AffineMapAttr>(indexingMaps[2]).getValue();
1115 return std::nullopt;
1120 unsigned numDims =
static_cast<unsigned>(batchRank) + 3;
1121 unsigned numOperandResults =
static_cast<unsigned>(batchRank) + 2;
1124 return std::nullopt;
1127 return std::nullopt;
1147 auto expected = ArrayAttr::get(
1151 AffineMapAttr::get(
AffineMap::get(numDims, 0, cDims, context))});
1152 if (indexingMaps != expected)
1153 return std::nullopt;
1157struct ContractionLowering :
public OpRewritePattern<vector::ContractionOp> {
1160 LogicalResult matchAndRewrite(vector::ContractionOp contractOp,
1161 PatternRewriter &rewriter)
const override {
1162 Location loc = contractOp.getLoc();
1164 if (contractOp.getKind() != vector::CombiningKind::ADD)
1166 "Expects add combining kind");
1171 VectorType accType = dyn_cast<VectorType>(acc.getType());
1175 std::optional<int64_t> batchRank =
1176 getRowMajorMatmulBatchRank(contractOp.getIndexingMapsAttr());
1180 "Expects a (batched) row-major matmul: leading dims must "
1181 "be batch dims shared by lhs, rhs, and acc; innermost two "
1182 "dims must be (M, K), (K, N), and (M, N)");
1187 "Expects operands of rank 4 or less");
1189 auto dpasOp = xegpu::DpasOp::create(
1190 rewriter, loc, contractOp.getResultType(),
lhs,
rhs, acc,
1191 nullptr,
nullptr,
nullptr);
1200static vector::ShapeCastOp getFlattenCast(
Value flat, VectorType ndType) {
1202 if (shapeCast && shapeCast.getSourceVectorType() == ndType)
1214static vector::BroadcastOp getSplatBroadcast(
Value flat) {
1223static bool canUnflatten(
Value flat, VectorType ndType) {
1224 return getFlattenCast(flat, ndType) || getDenseConstant(flat) ||
1225 getSplatBroadcast(flat);
1230 VectorType ndType) {
1231 assert(canUnflatten(flat, ndType) &&
"expected the cast to fold away");
1232 return vector::ShapeCastOp::create(rewriter, flat.
getLoc(), ndType, flat);
1256template <
typename OpTy>
1258 using OpRewritePattern<OpTy>::OpRewritePattern;
1260 LogicalResult matchAndRewrite(OpTy op,
1261 PatternRewriter &rewriter)
const override {
1262 constexpr bool isGather = std::is_same_v<OpTy, vector::GatherOp>;
1264 if (!isa<MemRefType>(op.getBase().getType()))
1267 if (op.getIndexVectorType().getRank() != 1)
1273 op.getIndices().template getDefiningOp<vector::ShapeCastOp>();
1274 if (!indexCast || indexCast.getSourceVectorType().getRank() < 2)
1276 op,
"index vector is not a shape_cast of an N-D vector");
1277 VectorType ndIndexType = indexCast.getSourceVectorType();
1278 VectorType ndMaskType =
1279 ndIndexType.cloneWith(std::nullopt, rewriter.
getI1Type());
1280 VectorType ndType = ndIndexType.cloneWith(
1281 std::nullopt, op.getVectorType().getElementType());
1285 if (!canUnflatten(op.getMask(), ndMaskType))
1288 if constexpr (isGather) {
1289 if (!canUnflatten(op.getPassThru(), ndType))
1291 "cannot un-flatten the pass-thru");
1293 Value mask = unflatten(rewriter, op.getMask(), ndMaskType);
1294 Value passThru = unflatten(rewriter, op.getPassThru(), ndType);
1295 auto ndGather = vector::GatherOp::create(
1296 rewriter, op.getLoc(), ndType, op.getBase(), op.getOffsets(),
1297 indexCast.getSource(), mask, passThru, op.getAlignmentAttr());
1298 ndGather->setDiscardableAttrs(op->getDiscardableAttrDictionary());
1302 if (!canUnflatten(op.getValueToStore(), ndType))
1304 op,
"cannot un-flatten the stored value");
1306 Value mask = unflatten(rewriter, op.getMask(), ndMaskType);
1307 Value valueToStore = unflatten(rewriter, op.getValueToStore(), ndType);
1311 op.getIndicesMutable().assign(indexCast.getSource());
1312 op.getMaskMutable().assign(mask);
1313 op.getValueToStoreMutable().assign(valueToStore);
1323static LogicalResult unflattenGatherScatter(
Operation *root) {
1326 patterns.add<UnflattenGatherScatter<vector::GatherOp>,
1327 UnflattenGatherScatter<vector::ScatterOp>>(ctx);
1328 vector::ShapeCastOp::getCanonicalizationPatterns(patterns, ctx);
1333static MemRefType withMemorySpace(MemRefType memrefTy,
Attribute newMemSpace) {
1334 return MemRefType::get(memrefTy.getShape(), memrefTy.getElementType(),
1335 memrefTy.getLayout(), newMemSpace);
1348static void promoteAllocasToSLM(
Operation *root) {
1350 Attribute slmAttr = IntegerAttr::get(IntegerType::get(ctx, 64), 3);
1356 auto isMemrefResultOp = [](
Operation *op) {
1359 return llvm::any_of(op->getResultTypes(),
1360 [](
Type t) { return isa<MemRefType>(t); });
1366 auto memrefTy = dyn_cast<MemRefType>(v.getType());
1367 if (!memrefTy || xegpu::XeGPUDialect::isSharedMemory(memrefTy))
1369 v.setType(withMemorySpace(memrefTy, slmAttr));
1371 if (!isMemrefResultOp(user))
1379 root->
walk([&](memref::AllocaOp op) {
1380 auto memrefTy = dyn_cast<MemRefType>(op.getResult().getType());
1381 if (!memrefTy || xegpu::XeGPUDialect::isSharedMemory(memrefTy))
1383 allocas.push_back(op);
1386 for (memref::AllocaOp alloca : allocas) {
1388 auto memrefTy = cast<MemRefType>(alloca.getResult().getType());
1389 auto newTy = withMemorySpace(memrefTy, slmAttr);
1390 auto newOp = memref::AllocaOp::create(
1391 builder, alloca.getLoc(), newTy, alloca.getDynamicSizes(),
1392 alloca.getSymbolOperands(), alloca.getAlignmentAttr());
1393 alloca.getResult().replaceAllUsesWith(newOp.getResult());
1397 if (!isMemrefResultOp(user))
1405struct ConvertVectorToXeGPUPass
1407 void runOnOperation()
override {
1410 promoteAllocasToSLM(getOperation());
1414 if (
failed(unflattenGatherScatter(getOperation())))
1415 return signalPassFailure();
1421 return signalPassFailure();
1430 .
add<TransferReadLowering, TransferWriteLowering, LoadLowering,
1431 ScatterLowering, GatherLowering, StoreLowering, ContractionLowering>(
static std::optional< VectorShape > vectorShape(Type type)
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 bool isSharedMemory(MemRefType type)
Return true if this is a shared memory memref type.
Base type for affine expression.
A multi-dimensional affine map Affine map's are immutable like Type's, and they are uniqued.
MLIRContext * getContext() const
bool isMinorIdentity() const
Returns true if this affine map is a minor identity, i.e.
static AffineMap get(MLIRContext *context)
Returns a zero result affine map with no dimensions or symbols: () -> ().
bool isProjectedPermutation(bool allowZeroInResults=false) const
Returns true if the AffineMap represents a subset (i.e.
ArrayRef< AffineExpr > getResults() const
bool isPermutationOfMinorIdentityWithBroadcasting(SmallVectorImpl< unsigned > &permutedDims) const
Return true if this affine map can be converted to a minor identity with broadcast by doing a permute...
unsigned getNumResults() const
unsigned getNumInputs() const
AffineExpr getResult(unsigned idx) const
static AffineMap getPermutationMap(ArrayRef< unsigned > permutation, MLIRContext *context)
Returns an AffineMap representing a permutation.
Attributes are known-constant values of operations.
IntegerAttr getI64IntegerAttr(int64_t value)
BoolAttr getBoolAttr(bool value)
MLIRContext * getContext() const
An attribute that represents a reference to a dense vector or tensor object.
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.
This class helps build Operations.
This class implements the operand iterators for the Operation class.
Operation is the basic unit of execution within MLIR.
OpResult getResult(unsigned idx)
Get the 'idx'th result of this operation.
std::enable_if_t< llvm::function_traits< std::decay_t< FnT > >::num_args==1, RetT > walk(FnT &&callback)
Walk the operation by calling the callback for each nested operation (including this one),...
user_range getUsers()
Returns a range of all users.
MLIRContext * getContext()
Return the context this operation is associated with.
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...
virtual void eraseOp(Operation *op)
This method erases an operation that is known to have no uses.
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,...
void modifyOpInPlace(Operation *root, CallableT &&callable)
This method is a utility wrapper around an in-place modification of an operation.
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...
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
bool use_empty() const
Returns true if this value has no uses.
Type getType() const
Return the type of this value.
user_range getUsers() const
Location getLoc() const
Return the location of this value.
Operation * getDefiningOp() const
If this value is the result of an operation, return the operation that defines it.
static ConstantIndexOp create(OpBuilder &builder, Location location, int64_t value)
const uArch * getUArch(llvm::StringRef archName)
int getLargestDivisor(T dim, ArrayRef< T > candidates, ArrayRef< T > candidateMultiples={})
Helper Function to find a proper instruction multiple for the user-supplied sg-level data shape (dive...
std::optional< std::string > getChipStr(Operation *op)
Retrieves the chip string from the XeVM target attribute of the parent GPU module operation.
Include the generated interface declarations.
bool matchPattern(Value value, const Pattern &pattern)
Entry point for matching a pattern over a Value.
void populatePrepareVectorToMMAPatterns(RewritePatternSet &patterns, bool useNvGpu=false)
Patterns to transform vector ops into a canonical form to convert to MMA matrix operations.
Type getType(OpFoldResult ofr)
Returns the int type of the integer in ofr.
AffineMap inverseAndBroadcastProjectedPermutation(AffineMap map)
Return the reverse map of a projected permutation where the projected dimensions are transformed into...
SmallVector< T > applyPermutation(ArrayRef< T > input, ArrayRef< int64_t > permutation)
LogicalResult applyPatternsGreedily(Region ®ion, const FrozenRewritePatternSet &patterns, GreedyRewriteConfig config=GreedyRewriteConfig(), bool *changed=nullptr)
Rewrite ops in the given region, which must be isolated from above, by repeatedly applying the highes...
bool isMemoryEffectFree(Operation *op)
Returns true if the given operation is free of memory effects.
std::conditional_t< std::is_same_v< Ty, mlir::Type >, mlir::Value, detail::TypedValue< Ty > > TypedValue
If Ty is mlir::Type this will select Value instead of having a wrapper around it.
void populateVectorToXeGPUConversionPatterns(RewritePatternSet &patterns)
Collect a set of patterns to convert from the vector to XeGPU ops.
llvm::TypeSwitch< T, ResultT > TypeSwitch
OpFoldResult getAsOpFoldResult(Value val)
Given a value, try to extract a constant Attribute.
detail::constant_op_matcher m_Constant()
Matches a constant foldable operation.
AffineExpr getAffineDimExpr(unsigned position, MLIRContext *context)
These free functions allow clients of the API to not use classes in detail.
OpRewritePattern is a wrapper around RewritePattern that allows for matching and rewriting against an...
bool isSupportedInstruction(InstructionKind instr) const
const Instruction * getInstruction(InstructionKind instKind) const