30#include "llvm/ADT/SetVector.h"
31#include "llvm/Support/LogicalResult.h"
32#include "llvm/Support/raw_ostream.h"
37#define GEN_PASS_DEF_XEGPUSGTOLANEDISTRIBUTE
38#include "mlir/Dialect/XeGPU/Transforms/Passes.h.inc"
44#define DEBUG_TYPE "xegpu-sg-to-lane-distribute"
45#define DBGS() (llvm::dbgs() << "[" DEBUG_TYPE "]: ")
50static Value castValueTo(ConversionPatternRewriter &rewriter,
53 if (v.getType() == expectedTy)
56 if (isa<VectorType>(v.getType()) &&
57 v.getType().getNumElements() == expectedTy.getNumElements())
58 return vector::ShapeCastOp::create(rewriter, v.getLoc(), expectedTy, v);
61 auto newOp = UnrealizedConversionCastOp::create(rewriter, v.getLoc(),
63 return newOp.getResult(0);
69static bool isValidSubgroupMultiReductionOp(vector::MultiDimReductionOp op) {
72 if (!resLayout || !resLayout.isForSubgroup())
75 if (op.getType().isIntOrFloat())
76 return op.getReductionDims().size() == 1;
77 VectorType resTy = dyn_cast<VectorType>(op.getType());
81 FailureOr<VectorType> resDistTypeOrFailure =
82 getDistVecTypeBasedOnLaneLayout(resLayout, resTy);
83 if (failed(resDistTypeOrFailure))
85 return op.getReductionDims().size() == 1;
91static bool isReductionLaneLocal(vector::MultiDimReductionOp op) {
93 assert(isValidSubgroupMultiReductionOp(op) &&
"Expecting a valid subgroup "
94 "MultiDimReductionOp");
97 assert(reductionDims.size() == 1 &&
98 "Expecting single reduction dimension for subgroup multi "
100 int64_t reductionDim = reductionDims[0];
102 assert(reductionDim <
static_cast<int64_t>(srcLaneLayout.size()) &&
103 "Expecting a source lane_layout covering the reduction dimension");
104 return srcLaneLayout[reductionDim] == 1;
110 VectorType distributedType) {
111 assert(originalType.getRank() == distributedType.getRank() &&
112 "original and distributed vector types must have the same rank");
114 for (
int64_t i = 0; i < originalType.getRank(); ++i) {
115 if (distributedType.getDimSize(i) != originalType.getDimSize(i))
116 distributedDims.push_back(i);
118 return distributedDims;
123struct SgToLaneCreateNdDesc
124 :
public OpConversionPattern<xegpu::CreateNdDescOp> {
125 using OpConversionPattern<xegpu::CreateNdDescOp>::OpConversionPattern;
128 matchAndRewrite(xegpu::CreateNdDescOp op, OpAdaptor adaptor,
129 ConversionPatternRewriter &rewriter)
const override {
130 xegpu::TensorDescType resultType = op.getType();
132 if (!resultType.getLayout())
135 auto newOp = xegpu::CreateNdDescOp::create(
136 rewriter, op.getLoc(),
TypeRange{resultType.dropLayouts()},
137 op.getOperands(), op.getProperties(),
138 op->getDiscardableAttrDictionary().getValue());
139 rewriter.replaceOp(op, newOp.getResult());
147struct SgToLaneLoadNd :
public OpConversionPattern<xegpu::LoadNdOp> {
148 using OpConversionPattern<xegpu::LoadNdOp>::OpConversionPattern;
151 matchAndRewrite(xegpu::LoadNdOp op, OpAdaptor adaptor,
152 ConversionPatternRewriter &rewriter)
const override {
153 xegpu::DistributeLayoutAttr layout = op.getAnchorLayout();
159 if (op.getTensorDescType().getLayout() != layout)
160 return rewriter.notifyMatchFailure(
161 op,
"conflicting layout attributes on tensor descriptor and anchor");
165 return rewriter.notifyMatchFailure(
166 op,
"xegpu::LoadNdOp require target attribute attached to "
167 "determine transpose "
169 auto supportedLaneResultTyOrFailure =
171 auto expectedLaneResultTyOrFailure =
173 if (failed(supportedLaneResultTyOrFailure))
174 return rewriter.notifyMatchFailure(
175 op,
"unable to compute the lane vector type for LoadNdOp");
176 if (failed(expectedLaneResultTyOrFailure))
177 return rewriter.notifyMatchFailure(
178 op,
"unable to compute expected lane vector type from lane layout");
179 auto newOp = xegpu::LoadNdOp::create(
180 rewriter, op.getLoc(), supportedLaneResultTyOrFailure.value(),
181 adaptor.getTensorDesc(), op.getMixedOffsets(), op.getPackedAttr(),
182 op.getTransposeAttr(), op.getL1HintAttr(), op.getL2HintAttr(),
183 op.getL3HintAttr(),
nullptr);
189 rewriter.replaceOp(op, castValueTo(rewriter, newOp.getResult(),
190 expectedLaneResultTyOrFailure.value()));
198struct SgToLaneStoreNd :
public OpConversionPattern<xegpu::StoreNdOp> {
199 using OpConversionPattern<xegpu::StoreNdOp>::OpConversionPattern;
202 matchAndRewrite(xegpu::StoreNdOp op, OpAdaptor adaptor,
203 ConversionPatternRewriter &rewriter)
const override {
204 xegpu::DistributeLayoutAttr layout = op.getAnchorLayout();
210 if (op.getTensorDescType().getLayout() != layout)
211 return rewriter.notifyMatchFailure(
212 op,
"conflicting layout attributes on tensor descriptor and anchor");
214 if (valueLayout != layout)
215 return rewriter.notifyMatchFailure(
216 op,
"conflicting layout attributes on value and anchor");
217 auto supportedLaneValueTyOrFailure =
219 if (failed(supportedLaneValueTyOrFailure))
220 return rewriter.notifyMatchFailure(
222 "unable to compute lane vector type for StoreNdOp value from tensor "
225 xegpu::StoreNdOp::create(
226 rewriter, op.getLoc(),
228 supportedLaneValueTyOrFailure.value()),
229 adaptor.getTensorDesc(), op.getMixedOffsets(), op.getL1HintAttr(),
230 op.getL2HintAttr(), op.getL3HintAttr(),
nullptr);
231 rewriter.eraseOp(op);
239struct SgToLaneDpas :
public OpConversionPattern<xegpu::DpasOp> {
240 using OpConversionPattern<xegpu::DpasOp>::OpConversionPattern;
243 matchAndRewrite(xegpu::DpasOp op, OpAdaptor adaptor,
244 ConversionPatternRewriter &rewriter)
const override {
246 auto layoutA = cast<xegpu::LayoutAttr>(op.getLayoutAAttr());
247 auto layoutB = cast<xegpu::LayoutAttr>(op.getLayoutBAttr());
248 auto layoutCd = cast<xegpu::LayoutAttr>(op.getLayoutCdAttr());
249 if (!layoutA || !layoutB || !layoutCd)
251 auto laneResultTyOrFailure =
253 auto laneATypeOrFailure =
255 auto laneBTypeOrFailure =
257 auto expectedLaneResultTyOrFailure =
259 if (failed(laneResultTyOrFailure) || failed(laneATypeOrFailure) ||
260 failed(laneBTypeOrFailure))
261 return rewriter.notifyMatchFailure(
262 op,
"failed to calculate supported lane vector types for DpasOp "
264 if (failed(expectedLaneResultTyOrFailure))
265 return rewriter.notifyMatchFailure(
266 op,
"unable to compute expected lane vector type for DpasOp from "
273 const auto *uArchInstruction =
274 dyn_cast<xegpu::uArch::SubgroupMatrixMultiplyAcc>(
275 uArch->getInstruction(
277 if (uArchInstruction) {
278 auto laneAType = laneATypeOrFailure.value();
279 auto laneBType = laneBTypeOrFailure.value();
281 unsigned aPackedBitWidth =
282 laneAType.getElementTypeBitWidth() * laneAType.getNumElements();
283 unsigned bPackedBitWidth =
284 laneBType.getElementTypeBitWidth() * laneBType.getNumElements();
285 unsigned expectedABitSize = uArchInstruction->getPackedFormatBitSizeA();
286 unsigned expectedBBitSize = uArchInstruction->getPackedFormatBitSizeB();
288 if (aPackedBitWidth % expectedABitSize != 0)
289 return rewriter.notifyMatchFailure(
291 "A operand packed bit width must be a multiple of uArch packed "
292 "format requirement");
293 if (bPackedBitWidth % expectedBBitSize != 0)
294 return rewriter.notifyMatchFailure(
296 "B operand packed bit width must be a multiple of uArch packed "
297 "format requirement");
301 auto newOp = xegpu::DpasOp::create(
302 rewriter, op->getLoc(), laneResultTyOrFailure.value(),
304 laneATypeOrFailure.value()),
306 laneBTypeOrFailure.value()),
308 laneResultTyOrFailure.value()),
312 rewriter.replaceOp(op, castValueTo(rewriter, newOp.getResult(),
313 expectedLaneResultTyOrFailure.value()));
326 ConversionPatternRewriter &rewriter)
const override {
333 return rewriter.notifyMatchFailure(
334 op,
"operation result is not a vector type");
336 xegpu::DistributeLayoutAttr layout =
338 if (!layout || !layout.isForSubgroup())
339 return rewriter.notifyMatchFailure(
340 op,
"operation result does not have subgroup distribute layout");
342 auto laneShapeOrFailure =
345 if (failed(laneShapeOrFailure))
346 return rewriter.notifyMatchFailure(
347 op,
"unable to compute lane vector type from the layout");
349 VectorType newResultType = laneShapeOrFailure.value();
355 if (!isa<xegpu::DistributeLayoutAttr>(attr.getValue()))
359 Operation *newOp = rewriter.create(state);
361 rewriter.replaceOp(op, newOp->
getResult(0));
373struct SgToLaneArithConstant :
public OpConversionPattern<arith::ConstantOp> {
374 using OpConversionPattern<arith::ConstantOp>::OpConversionPattern;
377 matchAndRewrite(arith::ConstantOp op, OpAdaptor adaptor,
378 ConversionPatternRewriter &rewriter)
const override {
379 auto resultType = dyn_cast<VectorType>(op.getType());
384 auto denseAttr = dyn_cast<DenseElementsAttr>(op.getValue());
386 return rewriter.notifyMatchFailure(
387 op,
"only dense vector constants are supported");
389 xegpu::DistributeLayoutAttr layout =
391 if (!layout || !layout.isForSubgroup())
392 return rewriter.notifyMatchFailure(
393 op,
"operation result does not have subgroup distribute layout");
395 auto laneShapeOrFailure =
398 if (failed(laneShapeOrFailure))
399 return rewriter.notifyMatchFailure(
400 op,
"unable to compute lane vector type from the layout");
402 VectorType newResultType = laneShapeOrFailure.value();
407 if (denseAttr.isSplat()) {
408 auto scalarValue = denseAttr.getSplatValue<
Attribute>();
411 arith::ConstantOp::create(rewriter, loc, newResultType, newDenseAttr);
412 rewriter.replaceOp(op, newOp.getResult());
419 arith::ConstantOp::create(rewriter, loc, resultType, denseAttr);
421 Value laneId = gpu::LaneIdOp::create(rewriter, loc, rewriter.getIndexType(),
422 mlir::IntegerAttr());
423 auto maybeCoordsVec = layout.computeDistributedCoords(
424 rewriter, loc, laneId, resultType.getShape());
425 if (failed(maybeCoordsVec))
426 return rewriter.notifyMatchFailure(
427 op,
"failed to compute distributed coordinates from layout");
432 int64_t rank = newResultType.getRank();
438 for (
int64_t d = 0; d < rank; d++)
439 blockGridShape[d] = distShape[d] / laneData[d];
442 auto blockType = VectorType::get(laneData, newResultType.getElementType());
447 rewriter, loc, newResultType, rewriter.getZeroAttr(newResultType));
449 for (
auto [blockIdx, blockStart] : llvm::enumerate(coordsVec)) {
457 for (
int64_t d = 0; d < rank; d++)
459 rewriter, loc, blockStart[d],
461 blockElems.push_back(vector::ExtractOp::create(
462 rewriter, loc, fullConst.getResult(), pos));
469 vector::FromElementsOp::create(rewriter, loc, blockType, blockElems);
473 for (
int64_t d = 0; d < rank; d++)
474 offsets[d] = blockGridPos[d] * laneData[d];
475 result = vector::InsertStridedSliceOp::create(rewriter, loc, block,
476 result, offsets, strides);
479 rewriter.replaceOp(op,
result);
485struct SgToLanePrefetchNd :
public OpConversionPattern<xegpu::PrefetchNdOp> {
486 using OpConversionPattern<xegpu::PrefetchNdOp>::OpConversionPattern;
489 matchAndRewrite(xegpu::PrefetchNdOp op, OpAdaptor adaptor,
490 ConversionPatternRewriter &rewriter)
const override {
491 xegpu::DistributeLayoutAttr layout = op.getAnchorLayout();
496 xegpu::PrefetchNdOp::create(rewriter, op.getLoc(), adaptor.getTensorDesc(),
497 op.getMixedOffsets(), op.getL1HintAttr(),
498 op.getL2HintAttr(), op.getL3HintAttr(),
500 rewriter.eraseOp(op);
530struct SgToLaneLoadGather :
public OpConversionPattern<xegpu::LoadGatherOp> {
531 using OpConversionPattern<xegpu::LoadGatherOp>::OpConversionPattern;
534 matchAndRewrite(xegpu::LoadGatherOp op, OpAdaptor adaptor,
535 ConversionPatternRewriter &rewriter)
const override {
536 xegpu::DistributeLayoutAttr layout = op.getAnchorLayout();
540 VectorType origResultTy = op.getValueType();
545 int effectiveVecRank = 1;
548 shape.take_front(origResultTy.getRank() - effectiveVecRank),
549 [](
int64_t d) { return d != 1; }))
550 return rewriter.notifyMatchFailure(
551 op,
"Only unit dimensions allowed for the leading "
552 "dimensions of the load vector!");
554 auto distResultTyOrFailure =
556 if (failed(distResultTyOrFailure))
557 return rewriter.notifyMatchFailure(
558 op,
"unable to compute expected lane vector type from lane layout");
560 VectorType distResultTy = distResultTyOrFailure.value();
561 VectorType distResultTy1D = VectorType::get({distResultTy.getNumElements()},
562 distResultTy.getElementType());
565 Value distOffsets = adaptor.getOffsets();
566 auto distOffsetsTy = cast<VectorType>(distOffsets.
getType());
567 VectorType offsetsTy1D = VectorType::get({distOffsetsTy.getNumElements()},
568 distOffsetsTy.getElementType());
569 distOffsets = castValueTo(
572 Value distMask = adaptor.getMask();
573 auto distMaskTy = cast<VectorType>(distMask.
getType());
574 VectorType maskTy1D = VectorType::get({distMaskTy.getNumElements()},
575 distMaskTy.getElementType());
579 Value distSource = adaptor.getSource();
580 auto newOp = xegpu::LoadGatherOp::create(
581 rewriter, op.getLoc(), distResultTy1D, distSource, distOffsets,
582 distMask, op.getL1HintAttr(), op.getL2HintAttr(), op.getL3HintAttr(),
586 if (distResultTy1D != distResultTy)
589 rewriter.replaceOp(op,
result);
598struct SgToLaneVectorReduction
599 :
public OpConversionPattern<vector::ReductionOp> {
600 using OpConversionPattern<vector::ReductionOp>::OpConversionPattern;
603 matchAndRewrite(vector::ReductionOp op, OpAdaptor adaptor,
604 ConversionPatternRewriter &rewriter)
const override {
608 if (!layout || !layout.isForSubgroup())
611 VectorType srcVecType = op.getSourceVectorType();
613 if (srcVecType.getRank() != 1)
614 return rewriter.notifyMatchFailure(
615 op,
"Only rank 1 reductions can be distributed.");
617 if (layout.getRank() != srcVecType.getRank())
618 return rewriter.notifyMatchFailure(
619 op,
"Layout rank does not match vector rank.");
622 int64_t sgSize = layout.getEffectiveLaneLayoutAsInt()[0];
626 return rewriter.notifyMatchFailure(
627 op,
"xegpu::ReductionOp require target attribute attached to "
628 "determine subgroup size");
631 if (sgSize != uArch->getSubgroupSize() ||
632 srcVecType.getShape()[0] % sgSize != 0)
633 return rewriter.notifyMatchFailure(op,
634 "Invalid layout or reduction vector "
635 "dimension must match subgroup size.");
637 if (!op.getType().isIntOrFloat())
638 return rewriter.notifyMatchFailure(
639 op,
"Reduction distribution currently only supports floats and "
643 Value laneValVec = adaptor.getVector();
647 op.getLoc(), rewriter, laneValVec, op.getKind(), sgSize);
650 if (adaptor.getAcc())
652 rewriter, op.getLoc(), op.getKind(), fullReduce, adaptor.getAcc());
654 rewriter.replaceOp(op, fullReduce);
663struct SgToLaneMultiDimReduction
664 :
public OpConversionPattern<vector::MultiDimReductionOp> {
665 using OpConversionPattern<vector::MultiDimReductionOp>::OpConversionPattern;
668 matchAndRewrite(vector::MultiDimReductionOp op, OpAdaptor adaptor,
669 ConversionPatternRewriter &rewriter)
const override {
672 assert(reductionDims.size() == 1 &&
673 "Expecting single reduction dimension for subgroup multi "
676 VectorType sourceType = op.getSourceVectorType();
677 int64_t rank = sourceType.getRank();
680 if (llvm::any_of(
shape.take_front(rank - 2),
681 [](
int64_t d) { return d != 1; }))
682 return rewriter.notifyMatchFailure(
683 op,
"only unit leading dimensions are supported for "
684 "multi_reduction with rank > 2");
688 if (op.getType().isIntOrFloat()) {
689 auto reductionDim = reductionDims[0];
690 VectorType origSourceType = op.getSourceVectorType();
691 int64_t reductionDimSize = origSourceType.getShape()[reductionDim];
695 op.getKind(), reductionDimSize);
697 if (adaptor.getAcc())
699 result, adaptor.getAcc());
700 }
else if (isReductionLaneLocal(op)) {
704 auto reductionDim = reductionDims[0];
708 reductionDim, op.getLoc(), rewriter);
710 auto reductionDim = reductionDims[0];
711 VectorType sourceType = op.getSourceVectorType();
712 int64_t reductionDimSize = sourceType.getShape()[reductionDim];
716 reductionDim, reductionDimSize, op.getLoc(), rewriter);
718 rewriter.replaceOp(op,
result);
727 ConversionPatternRewriter &rewriter,
Location loc,
730 Value laneId = gpu::LaneIdOp::create(rewriter, loc, rewriter.getIndexType(),
731 mlir::IntegerAttr());
733 layout.computeDistributedCoords(rewriter, loc, laneId, payloadShape);
734 if (failed(maybeCoords))
736 assert(maybeCoords.value().size() == 1 &&
737 "Expected one set of distributed offsets");
741 return llvm::map_to_vector(ofrVec, llvm::CastTo<Value>);
745struct SgToLaneLoadMatrix :
public OpConversionPattern<xegpu::LoadMatrixOp> {
746 using OpConversionPattern<xegpu::LoadMatrixOp>::OpConversionPattern;
749 matchAndRewrite(xegpu::LoadMatrixOp op, OpAdaptor adaptor,
750 ConversionPatternRewriter &rewriter)
const override {
751 auto layout = op.getLayoutAttr();
756 VectorType sgPayloadTy = dyn_cast<VectorType>(op.getResult().getType());
758 return rewriter.notifyMatchFailure(
759 op,
"the matrix op payload must be a vector type");
761 auto loc = op.getLoc();
762 auto offsets = op.getMixedOffsets();
764 return rewriter.notifyMatchFailure(op,
"the load op must have offsets");
766 FailureOr<VectorType> distPayloadTyOrFailure =
767 getDistVecTypeBasedOnLaneLayout(layout, sgPayloadTy);
768 if (failed(distPayloadTyOrFailure))
769 return rewriter.notifyMatchFailure(
770 op,
"Failed to distribute matrix op payload based on layout.");
776 if (!op.getSubgroupBlockIoAttr()) {
777 newCoords = computeDistributedCoordsForMatrixOp(
778 rewriter, loc, layout, sgPayloadTy.getShape(), offsetsAsValues);
779 if (newCoords.empty())
780 return rewriter.notifyMatchFailure(
781 op,
"Failed to compute distributed coordinates.");
785 ShapedType::kDynamic);
787 rewriter.getDenseI64ArrayAttr(newConstOffsets);
789 auto newOp = xegpu::LoadMatrixOp::create(
790 rewriter, loc, *distPayloadTyOrFailure, adaptor.getMemDesc(),
791 ValueRange(newCoords), newConstOffsetsAttr, op.getSubgroupBlockIoAttr(),
792 xegpu::DistributeLayoutAttr{});
793 rewriter.replaceOp(op, newOp.getResult());
799struct SgToLaneVectorTranspose
800 :
public OpConversionPattern<vector::TransposeOp> {
801 using OpConversionPattern<vector::TransposeOp>::OpConversionPattern;
804 matchAndRewrite(vector::TransposeOp op, OpAdaptor adaptor,
805 ConversionPatternRewriter &rewriter)
const override {
806 xegpu::DistributeLayoutAttr sourceLayout =
808 xegpu::DistributeLayoutAttr resultLayout =
810 if (!sourceLayout || !resultLayout)
811 return rewriter.notifyMatchFailure(
812 op,
"the source or result vector of the transpose op lacks layout "
816 if (!resultLayout.isTransposeOf(sourceLayout, perm,
818 return rewriter.notifyMatchFailure(
819 op,
"the source or result vector layouts must be transposes of "
821 FailureOr<VectorType> distributedResultTypeOrFailure =
822 getDistVecTypeBasedOnLaneLayout(resultLayout, op.getResultVectorType());
823 if (failed(distributedResultTypeOrFailure))
824 return rewriter.notifyMatchFailure(
825 op,
"Failed to distribute the result vector type in "
826 "vector::Transpose op");
827 auto newOp = vector::TransposeOp::create(rewriter, op.getLoc(),
828 adaptor.getVector(), perm);
829 rewriter.replaceOp(op, castValueTo(rewriter, newOp.getResult(),
830 distributedResultTypeOrFailure.value()));
837struct SgToLaneVectorBitcast :
public OpConversionPattern<vector::BitCastOp> {
838 using OpConversionPattern<vector::BitCastOp>::OpConversionPattern;
841 matchAndRewrite(vector::BitCastOp op, OpAdaptor adaptor,
842 ConversionPatternRewriter &rewriter)
const override {
843 xegpu::DistributeLayoutAttr resultLayout =
846 return rewriter.notifyMatchFailure(
847 op,
"result vector of the bitcast op lacks layout attribute");
848 FailureOr<VectorType> distributedResultTypeOrFailure =
849 getDistVecTypeBasedOnLaneLayout(resultLayout, op.getResultVectorType());
850 if (failed(distributedResultTypeOrFailure))
851 return rewriter.notifyMatchFailure(
852 op,
"Failed to distribute the result vector type in "
853 "vector::BitCast op");
854 auto newOp = vector::BitCastOp::create(
855 rewriter, op.getLoc(), distributedResultTypeOrFailure.value(),
856 adaptor.getSource());
857 rewriter.replaceOp(op, newOp.getResult());
900template <
typename OpType,
901 typename = std::enable_if_t<llvm::is_one_of<
902 OpType, vector::CreateMaskOp, vector::ConstantMaskOp>::value>>
903struct SgToLaneCreateMask :
public OpConversionPattern<OpType> {
904 using OpConversionPattern<OpType>::OpConversionPattern;
907 matchAndRewrite(OpType op,
typename OpType::Adaptor adaptor,
908 ConversionPatternRewriter &rewriter)
const override {
909 xegpu::DistributeLayoutAttr layout =
911 if (!layout || !layout.isForSubgroup())
912 return rewriter.notifyMatchFailure(
913 op,
"operation result does not have subgroup distribute layout");
915 VectorType origType = op.getType();
916 FailureOr<VectorType> distTypeOrFailure =
917 getDistVecTypeBasedOnLaneLayout(layout, origType);
918 if (failed(distTypeOrFailure))
919 return rewriter.notifyMatchFailure(
920 op,
"unable to compute lane vector type from the layout");
922 VectorType distType = distTypeOrFailure.value();
927 if constexpr (std::is_same_v<OpType, vector::CreateMaskOp>) {
928 origBounds.append(op.getOperands().begin(), op.getOperands().end());
930 auto dimSizes = op.getMaskDimSizesAttr().asArrayRef();
931 for (
auto dimSize : dimSizes)
932 origBounds.push_back(
939 Value laneId = gpu::LaneIdOp::create(rewriter, loc, rewriter.getIndexType(),
940 mlir::IntegerAttr());
941 auto maybeCoordsVec =
942 layout.computeDistributedCoords(rewriter, loc, laneId, origShape);
943 if (failed(maybeCoordsVec))
944 return rewriter.notifyMatchFailure(
945 op,
"failed to compute distributed coordinates from layout");
950 int64_t rank = distType.getRank();
951 int64_t numElements = distType.getNumElements();
953 if (
static_cast<int64_t>(laneData.size()) != rank ||
955 return rewriter.notifyMatchFailure(
956 op,
"lane_data does not tile the distributed vector");
959 assert(
static_cast<int64_t>(laneDataCoords.size()) *
962 "number of coordinate sets must match number of lane_data blocks");
966 layout.computeStaticDistributedCoords(0, origShape);
967 if (staticLaneDataOffsetsOrig.empty() ||
968 staticLaneDataOffsetsOrig.size() != laneDataCoords.size())
969 return rewriter.notifyMatchFailure(
970 op,
"static and dynamic coordinates disagree on the distribution "
985 staticLaneDataOffsetsOrig[linearLaneDataIdx++];
991 for (
int64_t d = 0; d < rank; d++)
992 elementOffsetDist[d] =
993 laneDataOffsetDist[d] + elemOffsetInLaneData[d];
994 int64_t elemLinearizedIdxDist =
995 linearize(elementOffsetDist, distStrides);
996 for (
int64_t d = 0; d < rank; d++)
997 staticElemOffset[d][elemLinearizedIdxDist] =
998 staticLaneDataOffsetOrig[d] + elemOffsetInLaneData[d];
1007 auto flatIndexType = VectorType::get(numElements, rewriter.getIndexType());
1009 for (
int64_t d = 0; d < rank; d++) {
1011 if (constBound && *constBound >= origShape[d])
1014 Value materializedStaticOffset = arith::ConstantOp::create(
1015 rewriter, loc, rewriter.getIndexVectorAttr(staticElemOffset[d]));
1018 Value validDimExtent = arith::SubIOp::create(rewriter, loc, origBounds[d],
1019 laneDataCoords[0][d]);
1021 Value validDimExtentPerElement = vector::BroadcastOp::create(
1022 rewriter, loc, flatIndexType, validDimExtent);
1023 Value elementMaskInDim = arith::CmpIOp::create(
1024 rewriter, loc, arith::CmpIPredicate::slt, materializedStaticOffset,
1025 validDimExtentPerElement);
1028 inBounds = inBounds ? arith::AndIOp::create(rewriter, loc, inBounds,
1036 rewriter.replaceOp(op, arith::ConstantOp::create(
1037 rewriter, loc, distType,
1042 rewriter.createOrFold<vector::ShapeCastOp>(loc, distType, inBounds);
1043 rewriter.replaceOp(op, resMask);
1049struct SgToLaneStoreMatrix :
public OpConversionPattern<xegpu::StoreMatrixOp> {
1050 using OpConversionPattern<xegpu::StoreMatrixOp>::OpConversionPattern;
1053 matchAndRewrite(xegpu::StoreMatrixOp op, OpAdaptor adaptor,
1054 ConversionPatternRewriter &rewriter)
const override {
1055 auto layout = op.getLayoutAttr();
1060 VectorType sgPayloadTy = dyn_cast<VectorType>(op.getData().getType());
1062 return rewriter.notifyMatchFailure(
1063 op,
"the matrix op payload must be a vector type");
1065 auto loc = op.getLoc();
1066 auto offsets = op.getMixedOffsets();
1067 if (offsets.empty())
1068 return rewriter.notifyMatchFailure(op,
"the store op must have offsets");
1070 FailureOr<VectorType> distPayloadTyOrFailure =
1071 getDistVecTypeBasedOnLaneLayout(layout, sgPayloadTy);
1072 if (failed(distPayloadTyOrFailure))
1073 return rewriter.notifyMatchFailure(
1074 op,
"Failed to distribute matrix op payload based on layout.");
1080 if (!op.getSubgroupBlockIoAttr()) {
1081 newCoords = computeDistributedCoordsForMatrixOp(
1082 rewriter, loc, layout, sgPayloadTy.getShape(), offsetsAsValues);
1083 if (newCoords.empty())
1084 return rewriter.notifyMatchFailure(
1085 op,
"Failed to compute distributed coordinates.");
1089 ShapedType::kDynamic);
1091 rewriter.getDenseI64ArrayAttr(newConstOffsets);
1093 xegpu::StoreMatrixOp::create(
1096 distPayloadTyOrFailure.value()),
1097 adaptor.getMemDesc(),
ValueRange(newCoords), newConstOffsetsAttr,
1098 op.getSubgroupBlockIoAttr(), xegpu::DistributeLayoutAttr{});
1099 rewriter.eraseOp(op);
1130struct SgToLaneStoreScatter
1131 :
public OpConversionPattern<xegpu::StoreScatterOp> {
1132 using OpConversionPattern<xegpu::StoreScatterOp>::OpConversionPattern;
1135 matchAndRewrite(xegpu::StoreScatterOp op, OpAdaptor adaptor,
1136 ConversionPatternRewriter &rewriter)
const override {
1137 xegpu::DistributeLayoutAttr layout = op.getAnchorLayout();
1141 VectorType origValueTy = op.getValueType();
1146 int effectiveVecRank = 1;
1148 if (llvm::any_of(
shape.take_front(origValueTy.getRank() - effectiveVecRank),
1149 [](
int64_t d) { return d != 1; }))
1150 return rewriter.notifyMatchFailure(
1151 op,
"Only unit dimensions allowed for the leading "
1152 "dimensions of the store vector!");
1154 auto distValueTyOrFailure =
1156 if (failed(distValueTyOrFailure))
1157 return rewriter.notifyMatchFailure(
1158 op,
"unable to compute expected lane vector type from lane layout");
1160 VectorType distValueTy = distValueTyOrFailure.value();
1161 VectorType distValueTy1D = VectorType::get({distValueTy.getNumElements()},
1162 distValueTy.getElementType());
1164 Value distValue = adaptor.getValue();
1165 if (distValue.
getType() != distValueTy1D)
1170 Value distOffsets = adaptor.getOffsets();
1171 auto distOffsetsTy = cast<VectorType>(distOffsets.
getType());
1172 VectorType offsetsTy1D = VectorType::get({distOffsetsTy.getNumElements()},
1173 distOffsetsTy.getElementType());
1174 distOffsets = castValueTo(
1177 Value distMask = adaptor.getMask();
1178 auto distMaskTy = cast<VectorType>(distMask.
getType());
1179 VectorType maskTy1D = VectorType::get({distMaskTy.getNumElements()},
1180 distMaskTy.getElementType());
1184 Value distDest = adaptor.getDest();
1185 xegpu::StoreScatterOp::create(rewriter, op.getLoc(), distValue, distDest,
1186 distOffsets, distMask, op.getL1HintAttr(),
1187 op.getL2HintAttr(), op.getL3HintAttr(),
1190 rewriter.eraseOp(op);
1199struct SgToLaneVectorStep :
public OpConversionPattern<vector::StepOp> {
1200 using OpConversionPattern<vector::StepOp>::OpConversionPattern;
1203 matchAndRewrite(vector::StepOp op, OpAdaptor adaptor,
1204 ConversionPatternRewriter &rewriter)
const override {
1205 xegpu::DistributeLayoutAttr resultLayout =
1207 if (!resultLayout || !resultLayout.isForSubgroup())
1208 return rewriter.notifyMatchFailure(
1209 op,
"the result vector of the step op lacks subgroup layout");
1211 auto loc = op.getLoc();
1212 auto stepResultVecTy = op.getResult().getType();
1213 auto laneShapeOrFailure =
1215 if (failed(laneShapeOrFailure))
1216 return rewriter.notifyMatchFailure(
1217 op,
"unable to compute lane vector type from the layout");
1218 VectorType newVecTy = laneShapeOrFailure.value();
1220 Value laneId = gpu::LaneIdOp::create(rewriter, loc, rewriter.getIndexType(),
1221 mlir::IntegerAttr());
1222 auto laneDataBlockCoords = resultLayout.computeDistributedCoords(
1223 rewriter, loc, laneId, stepResultVecTy.getShape());
1224 if (failed(laneDataBlockCoords))
1225 return rewriter.notifyMatchFailure(
1226 op,
"failed to compute lane data block coordinates");
1228 auto laneDataBlockCoordsVec = laneDataBlockCoords.value();
1229 auto laneDataBlockLength = resultLayout.getEffectiveLaneDataAsInt()[0];
1230 assert(
static_cast<int64_t>(laneDataBlockCoordsVec.size()) ==
1231 newVecTy.getNumElements() / laneDataBlockLength);
1240 for (
auto &laneDataBlockCoords : laneDataBlockCoordsVec) {
1241 auto laneDataBlockStartCoord = laneDataBlockCoords[0];
1242 stepVals.push_back(laneDataBlockStartCoord);
1243 for (
int i = 1; i < laneDataBlockLength; ++i) {
1245 stepVals.push_back(arith::AddIOp::create(
1246 rewriter, loc, laneDataBlockStartCoord, offset));
1249 assert(
static_cast<int64_t>(stepVals.size()) == newVecTy.getNumElements() &&
1250 "Expecting the number of step values to match the number of "
1251 "elements in the vector");
1253 vector::FromElementsOp::create(rewriter, loc, newVecTy, stepVals);
1254 rewriter.replaceOp(op, stepOpVal);
1261struct SgToLaneVectorExtract :
public OpConversionPattern<vector::ExtractOp> {
1262 using OpConversionPattern<vector::ExtractOp>::OpConversionPattern;
1265 matchAndRewrite(vector::ExtractOp op, OpAdaptor adaptor,
1266 ConversionPatternRewriter &rewriter)
const override {
1268 auto resultType = dyn_cast<VectorType>(op.getType());
1270 return rewriter.notifyMatchFailure(op,
"scalar extract not supported");
1272 xegpu::DistributeLayoutAttr layout =
1274 if (!layout || !layout.isForSubgroup())
1279 auto laneLayout = layout.getEffectiveLaneLayoutAsInt();
1281 [](
int64_t v) {
return v != 1; }))
1282 return rewriter.notifyMatchFailure(
1283 op,
"only innermost dimension distribution is supported for "
1286 auto newOp = vector::ExtractOp::create(
1287 rewriter, op.getLoc(), adaptor.getSource(), op.getMixedPosition());
1288 rewriter.replaceOp(op, newOp.getResult());
1294struct SgToLaneVectorShapeCast
1295 :
public OpConversionPattern<vector::ShapeCastOp> {
1296 using OpConversionPattern<vector::ShapeCastOp>::OpConversionPattern;
1299 matchAndRewrite(vector::ShapeCastOp op, OpAdaptor adaptor,
1300 ConversionPatternRewriter &rewriter)
const override {
1301 xegpu::DistributeLayoutAttr resultLayout =
1303 if (!resultLayout || !resultLayout.isForSubgroup())
1304 return rewriter.notifyMatchFailure(
1305 op,
"the result vector of the shape_cast op lacks subgroup layout");
1308 resultLayout, op.getResultVectorType());
1309 if (failed(resultDistTypeOrFailure))
1310 return rewriter.notifyMatchFailure(
1311 op,
"failed to get distributed vector type for result");
1313 Value source = adaptor.getSource();
1314 auto newShapeCast = vector::ShapeCastOp::create(
1315 rewriter, op.getLoc(), resultDistTypeOrFailure.value(), source);
1316 rewriter.replaceOp(op, newShapeCast);
1324struct SgToLaneVectorExtractStridedSlice
1325 :
public OpConversionPattern<vector::ExtractStridedSliceOp> {
1326 using OpConversionPattern<vector::ExtractStridedSliceOp>::OpConversionPattern;
1329 matchAndRewrite(vector::ExtractStridedSliceOp op, OpAdaptor adaptor,
1330 ConversionPatternRewriter &rewriter)
const override {
1331 xegpu::DistributeLayoutAttr resultLayout =
1333 if (!resultLayout || !resultLayout.isForSubgroup())
1336 VectorType resultType = op.getType();
1337 auto distResultTyOrFailure =
1339 if (failed(distResultTyOrFailure))
1340 return rewriter.notifyMatchFailure(
1341 op,
"unable to compute distributed vector type from lane layout");
1342 VectorType distResultTy = *distResultTyOrFailure;
1345 getDistributedDims(resultType, distResultTy);
1348 int64_t sourceRank = op.getSourceVectorType().getRank();
1350 llvm::map_to_vector(op.getSizes(), [](
Attribute attr) { return attr; });
1352 op.getOffsets(), [](
Attribute attr) { return attr; });
1354 op.getStrides(), [](
Attribute attr) { return attr; });
1355 for (
int64_t i = op.getSizes().size(); i < sourceRank; ++i) {
1356 updatedSizes.push_back(
1357 rewriter.getI64IntegerAttr(op.getSourceVectorType().getDimSize(i)));
1358 updatedOffsets.push_back(rewriter.getI64IntegerAttr(0));
1359 updatedStrides.push_back(rewriter.getI64IntegerAttr(1));
1364 if (!distributedDims.empty()) {
1366 if (!sourceLayout || sourceLayout.getEffectiveLaneLayoutAsInt().empty())
1367 return rewriter.notifyMatchFailure(
1368 op,
"source of extract_strided_slice lacks distribution layout");
1370 sourceLayout.getEffectiveLaneLayoutAsInt();
1373 for (
int64_t distDim : distributedDims) {
1374 int64_t lanes = laneLayout[distDim];
1375 if (lanes == 0 || sourceShape[distDim] % lanes != 0)
1376 return rewriter.notifyMatchFailure(
1377 op,
"source size along a distributed dim is not a multiple of "
1380 cast<IntegerAttr>(updatedOffsets[distDim]).getInt();
1381 if (distrDimOffset % (lanes * laneData[distDim]) != 0)
1382 return rewriter.notifyMatchFailure(
1383 op,
"offset along a distributed dim is not a multiple of its "
1385 updatedSizes[distDim] =
1386 rewriter.getI64IntegerAttr(distResultTy.getDimSize(distDim));
1387 updatedOffsets[distDim] =
1388 rewriter.getI64IntegerAttr(distrDimOffset / lanes);
1392 auto newOp = vector::ExtractStridedSliceOp::create(
1393 rewriter, op.getLoc(), distResultTy, adaptor.getSource(),
1394 ArrayAttr::get(rewriter.getContext(), updatedOffsets),
1395 ArrayAttr::get(rewriter.getContext(), updatedSizes),
1396 ArrayAttr::get(rewriter.getContext(), updatedStrides));
1397 rewriter.replaceOp(op, newOp.getResult());
1459struct SgToLaneBroadcast :
public OpConversionPattern<vector::BroadcastOp> {
1460 using OpConversionPattern<vector::BroadcastOp>::OpConversionPattern;
1463 matchAndRewrite(vector::BroadcastOp op, OpAdaptor adaptor,
1464 ConversionPatternRewriter &rewriter)
const override {
1465 xegpu::DistributeLayoutAttr resultLayout =
1467 if (!resultLayout || !resultLayout.isForSubgroup())
1468 return rewriter.notifyMatchFailure(
1469 op,
"result does not have subgroup distribute layout");
1471 VectorType destType = op.getResultVectorType();
1472 VectorType sourceType = dyn_cast<VectorType>(op.getSourceType());
1474 xegpu::DistributeLayoutAttr sourceLayout =
1478 int64_t rankDiff = destType.getRank() - sourceType.getRank();
1481 if (!sourceLayout || !sourceLayout.isSliceOf(resultLayout))
1483 "broadcast source layout must be a slice of result layout");
1484 }
else if (rankDiff == 0) {
1486 auto broadcastUnitDimsSet = op.computeBroadcastedUnitDims();
1488 broadcastUnitDimsSet.end());
1489 assert(sourceLayout.isEqualTo(
1490 sourceLayout.setUnitDimData(broadcastUnitDims)) &&
1491 "The sg_data for unit dimensions should be set as 1");
1492 sourceLayout = sourceLayout.setUnitDimLayout(broadcastUnitDims);
1497 return rewriter.notifyMatchFailure(
1498 op,
"broadcast from scalar must not have a layout attribute");
1503 if (failed(destDistType))
1504 return rewriter.notifyMatchFailure(
1505 op,
"failed to distribute the result vector type");
1507 Value source = adaptor.getSource();
1509 if (source.
getType() == destDistType.value()) {
1510 rewriter.replaceOp(op, source);
1514 auto newOp = vector::BroadcastOp::create(rewriter, op.getLoc(),
1515 destDistType.value(), source);
1516 rewriter.replaceOp(op, newOp);
1524struct SgToLaneVectorInsertStridedSlice
1525 :
public OpConversionPattern<vector::InsertStridedSliceOp> {
1526 using OpConversionPattern<vector::InsertStridedSliceOp>::OpConversionPattern;
1529 matchAndRewrite(vector::InsertStridedSliceOp op, OpAdaptor adaptor,
1530 ConversionPatternRewriter &rewriter)
const override {
1531 xegpu::DistributeLayoutAttr resultLayout =
1533 if (!resultLayout || !resultLayout.isForSubgroup())
1536 VectorType destType = op.getDestVectorType();
1537 auto distDestTyOrFailure =
1539 if (failed(distDestTyOrFailure))
1540 return rewriter.notifyMatchFailure(
1541 op,
"unable to compute distributed vector type from lane layout");
1542 VectorType distDestTy = *distDestTyOrFailure;
1545 getDistributedDims(destType, distDestTy);
1548 op.getOffsets(), [](
Attribute attr) { return attr; });
1550 if (!destDistributedDims.empty()) {
1551 if (destDistributedDims.size() != 1)
1552 return rewriter.notifyMatchFailure(
1553 op,
"only single dimension distribution is supported");
1554 int64_t destDistDim = destDistributedDims[0];
1556 VectorType srcType = op.getSourceVectorType();
1559 destDistDim - (destType.getRank() - srcType.getRank());
1560 if (sourceDistDim < 0)
1561 return rewriter.notifyMatchFailure(
1562 op,
"distributed dimension must be in the last k dims of dest");
1566 if (!destLayout || !sourceLayout ||
1567 destLayout.getEffectiveLaneLayoutAsInt().empty() ||
1568 sourceLayout.getEffectiveLaneLayoutAsInt().empty())
1569 return rewriter.notifyMatchFailure(
1570 op,
"source or dest of insert_strided_slice lacks distribution "
1573 auto destLaneData = destLayout.getEffectiveLaneDataAsInt();
1574 auto sourceLaneData = sourceLayout.getEffectiveLaneDataAsInt();
1577 if ((destDistDim <
static_cast<int64_t>(destLaneData.size()) &&
1578 destLaneData[destDistDim] != 1) ||
1579 (sourceDistDim <
static_cast<int64_t>(sourceLaneData.size()) &&
1580 sourceLaneData[sourceDistDim] != 1))
1581 return rewriter.notifyMatchFailure(
1582 op,
"expecting unit lane data along the distributed dimension");
1586 auto destLaneLayout = destLayout.getEffectiveLaneLayoutAsInt();
1588 destDistDim < static_cast<int64_t>(destLaneLayout.size())
1589 ? destLaneLayout[destDistDim]
1592 int64_t srcDistrDimSize = srcType.getDimSize(sourceDistDim);
1593 if (srcDistrDimSize % numLanesAlongDim != 0)
1594 return rewriter.notifyMatchFailure(
1595 op,
"source distributed dim size is not a multiple of "
1596 "the number of lanes along that dimension");
1599 cast<IntegerAttr>(op.getOffsets()[destDistDim]).getInt();
1600 if (destDistrDimOffset % numLanesAlongDim != 0)
1601 return rewriter.notifyMatchFailure(
1602 op,
"offset along distributed dim is not a multiple of "
1603 "the number of lanes along that dimension");
1605 updatedOffsets[destDistDim] =
1606 rewriter.getI64IntegerAttr(destDistrDimOffset / numLanesAlongDim);
1609 auto newOp = vector::InsertStridedSliceOp::create(
1610 rewriter, op.getLoc(), distDestTy, adaptor.getValueToStore(),
1612 ArrayAttr::get(rewriter.getContext(), updatedOffsets), op.getStrides());
1613 rewriter.replaceOp(op, newOp.getResult());
1620struct SgToLaneVectorInsert :
public OpConversionPattern<vector::InsertOp> {
1621 using OpConversionPattern<vector::InsertOp>::OpConversionPattern;
1624 matchAndRewrite(vector::InsertOp op, OpAdaptor adaptor,
1625 ConversionPatternRewriter &rewriter)
const override {
1627 auto valueType = dyn_cast<VectorType>(op.getValueToStoreType());
1629 return rewriter.notifyMatchFailure(op,
"scalar insert not supported");
1631 xegpu::DistributeLayoutAttr layout =
1633 if (!layout || !layout.isForSubgroup())
1638 auto laneLayout = layout.getEffectiveLaneLayoutAsInt();
1640 [](
int64_t v) {
return v != 1; }))
1641 return rewriter.notifyMatchFailure(
1642 op,
"only innermost dimension distribution is supported for "
1645 auto newOp = vector::InsertOp::create(
1646 rewriter, op.getLoc(), adaptor.getValueToStore(), adaptor.getDest(),
1647 op.getMixedPosition());
1648 rewriter.replaceOp(op, newOp.getResult());
1666static FailureOr<Value>
1667shuffleDataAsLaneLayoutChange(ConversionPatternRewriter &rewriter,
Location loc,
1670 VectorType srcTy = dyn_cast<VectorType>(src.
getType());
1671 if (!srcTy || srcTy.getRank() != 2)
1674 if (targetLaneNum <= 0 || currentLaneNum != targetLaneNum * 2)
1678 srcTy.getNumElements() * srcTy.getElementTypeBitWidth();
1679 if (vectorBitWidth % 32 != 0)
1691 Type shuffleElemTy = rewriter.getI32Type();
1692 int64_t numShuffles = vectorBitWidth / 32;
1693 VectorType shuffleBundleTy = VectorType::get({numShuffles}, shuffleElemTy);
1695 Value temp = arith::ConstantOp::create(
1698 IntegerAttr::get(shuffleElemTy, 0)));
1699 VectorType flatSrcTy =
1700 VectorType::get({srcTy.getNumElements()}, srcTy.getElementType());
1701 Value flatSrc = vector::ShapeCastOp::create(rewriter, loc, flatSrcTy, src);
1702 Value shuffleBundle =
1703 vector::BitCastOp::create(rewriter, loc, shuffleBundleTy, flatSrc);
1704 for (
int64_t i = 0; i < numShuffles; i++) {
1706 vector::ExtractOp::create(rewriter, loc, shuffleBundle, i);
1707 shuffleElem = gpu::ShuffleOp::create(rewriter, loc, shuffleElem, 0,
1708 targetLaneNum, gpu::ShuffleMode::UP)
1710 temp = vector::InsertOp::create(rewriter, loc, shuffleElem, temp, i);
1712 temp = vector::BitCastOp::create(rewriter, loc, flatSrcTy, temp);
1713 temp = vector::ShapeCastOp::create(rewriter, loc, srcTy, temp);
1718 Value res = vector::ShuffleOp::create(rewriter, loc, src, temp,
indices);
1729static FailureOr<Value> repackLaneData(ConversionPatternRewriter &rewriter,
1733 auto srcTy = dyn_cast<VectorType>(src.
getType());
1736 int64_t rank = srcTy.getRank();
1737 Type elemTy = srcTy.getElementType();
1738 int64_t k = srcTy.getShape()[repackDim];
1740 bool roundRobinToContig = inputData == 1 && targetData == k;
1741 bool contigToRoundRobin = inputData == k && targetData == 1;
1742 if (!roundRobinToContig && !contigToRoundRobin)
1747 xegpu::LaneShuffleMode mode = roundRobinToContig
1748 ? xegpu::LaneShuffleMode::Pack
1749 : xegpu::LaneShuffleMode::Unpack;
1750 VectorType runTy = VectorType::get({k}, elemTy);
1754 if (srcTy.getNumElements() == k) {
1757 xegpu::LaneShuffleOp::create(rewriter, loc, runTy, src, mode));
1758 Value flat = vector::ShapeCastOp::create(rewriter, loc, runTy, src);
1760 xegpu::LaneShuffleOp::create(rewriter, loc, runTy, flat, mode);
1761 return Value(vector::ShapeCastOp::create(rewriter, loc, srcTy, shuffled));
1766 if (repackDim == rank - 1) {
1770 Value result = arith::ConstantOp::create(rewriter, loc, srcTy,
1771 rewriter.getZeroAttr(srcTy));
1772 for (
int64_t i = 0; i < numRuns; ++i) {
1774 Value run = vector::ExtractOp::create(rewriter, loc, src, pos);
1776 xegpu::LaneShuffleOp::create(rewriter, loc, runTy, run, mode);
1777 result = vector::InsertOp::create(rewriter, loc, shuffled,
result, pos);
1787 for (
int64_t d = 0; d < rank; ++d)
1788 if (d != repackDim) {
1789 keptShape.push_back(srcTy.getShape()[d]);
1790 keptDims.push_back(d);
1795 sliceSizes[repackDim] = k;
1797 VectorType sliceTy = VectorType::get(sliceSizes, elemTy);
1798 Value result = arith::ConstantOp::create(rewriter, loc, srcTy,
1799 rewriter.getZeroAttr(srcTy));
1800 for (
int64_t i = 0; i < numRuns; ++i) {
1803 for (
auto [dim, coord] : llvm::zip_equal(keptDims, keptPos))
1804 offsets[dim] = coord;
1805 Value slice = vector::ExtractStridedSliceOp::create(
1806 rewriter, loc, src, offsets, sliceSizes, sliceStrides);
1807 Value run = vector::ShapeCastOp::create(rewriter, loc, runTy, slice);
1809 xegpu::LaneShuffleOp::create(rewriter, loc, runTy, run, mode);
1810 Value repackedSlice =
1811 vector::ShapeCastOp::create(rewriter, loc, sliceTy, repacked);
1812 result = vector::InsertStridedSliceOp::create(
1813 rewriter, loc, repackedSlice,
result, offsets, sliceStrides);
1819struct SgToLaneConvertLayout
1820 :
public OpConversionPattern<xegpu::ConvertLayoutOp> {
1821 using OpConversionPattern<xegpu::ConvertLayoutOp>::OpConversionPattern;
1824 matchAndRewrite(xegpu::ConvertLayoutOp op, OpAdaptor adaptor,
1825 ConversionPatternRewriter &rewriter)
const override {
1826 auto inputLayout = op.getEffectiveInputLayout();
1827 auto targetLayout = op.getTargetLayoutAttr();
1828 Type valType = op.getResult().getType();
1831 rewriter.replaceOp(op, op.getSource());
1835 auto resShape = cast<VectorType>(valType).getShape();
1840 if (inputLayout.isCompatibleWith(targetLayout, resShapeVec,
1842 rewriter.replaceOp(op, adaptor.getSource());
1853 if (inputLayout.getEffectiveOrderAsInt() ==
1854 targetLayout.getEffectiveOrderAsInt() &&
1855 inputLayout.getRank() == 2 && targetLayout.getRank() == 2) {
1856 auto laneLayout = inputLayout.getEffectiveLaneLayoutAsInt();
1857 auto targetLaneLayout = targetLayout.getEffectiveLaneLayoutAsInt();
1858 auto laneData = inputLayout.getEffectiveLaneDataAsInt();
1859 auto targetLaneData = targetLayout.getEffectiveLaneDataAsInt();
1860 if (laneLayout.size() == 2 && targetLaneLayout.size() == 2 &&
1861 laneData == targetLaneData && laneLayout[1] == 1 &&
1862 targetLaneLayout[1] == 1 && laneLayout[0] > 1 &&
1863 laneLayout[0] != targetLaneLayout[0]) {
1864 FailureOr<Value> res = shuffleDataAsLaneLayoutChange(
1865 rewriter, op.getLoc(), adaptor.getSource(), laneLayout[0],
1866 targetLaneLayout[0]);
1867 if (succeeded(res)) {
1868 rewriter.replaceOp(op, *res);
1880 if (inputLayout.getEffectiveOrderAsInt() ==
1881 targetLayout.getEffectiveOrderAsInt() &&
1882 inputLayout.getEffectiveLaneLayoutAsInt() ==
1883 targetLayout.getEffectiveLaneLayoutAsInt()) {
1884 auto laneLayout = inputLayout.getEffectiveLaneLayoutAsInt();
1885 auto laneData = inputLayout.getEffectiveLaneDataAsInt();
1886 auto targetLaneData = targetLayout.getEffectiveLaneDataAsInt();
1889 int64_t rank = laneData.size();
1891 bool multipleChanged =
false;
1892 for (
int64_t d = 0; d < rank; ++d)
1893 if (laneData[d] != targetLaneData[d]) {
1894 if (repackDim != -1)
1895 multipleChanged =
true;
1901 int64_t otherDim = repackDim == rank - 1 ? rank - 2 : rank - 1;
1902 bool laneLayoutOk = repackDim != -1 && laneLayout[repackDim] != 1 &&
1903 (rank < 2 || laneLayout[otherDim] == 1);
1907 if (repackDim != -1 && repackDim >= rank - 2 && !multipleChanged &&
1909 FailureOr<Value> res = repackLaneData(
1910 rewriter, op.getLoc(), adaptor.getSource(), repackDim,
1911 laneData[repackDim], targetLaneData[repackDim]);
1912 if (succeeded(res)) {
1913 rewriter.replaceOp(op, *res);
1919 return rewriter.notifyMatchFailure(
1920 op,
"lowering incompatible convert_layout not yet supported");
1928struct LaneFragmentPiece {
1933static LaneFragmentPiece getLaneFragmentPiece(
int64_t unit,
1939 for (
int64_t d = 0; d < rank; ++d)
1940 units[d] =
shape[d] / std::min(
shape[d], laneLayout[d] * laneData[d]);
1942 LaneFragmentPiece piece;
1943 for (
int64_t d = 0; d < rank; ++d) {
1944 piece.offsets.push_back(unitCoords[d] * laneData[d]);
1945 piece.sizes.push_back(laneData[d]);
1964struct SgToLaneConvertLayoutViaSLM
1965 :
public OpConversionPattern<xegpu::ConvertLayoutOp> {
1966 using OpConversionPattern<xegpu::ConvertLayoutOp>::OpConversionPattern;
1969 matchAndRewrite(xegpu::ConvertLayoutOp op, OpAdaptor adaptor,
1970 ConversionPatternRewriter &rewriter)
const override {
1971 xegpu::DistributeLayoutAttr inputLayout = op.getEffectiveInputLayout();
1972 xegpu::DistributeLayoutAttr targetLayout = op.getTargetLayoutAttr();
1973 if (!inputLayout || !targetLayout || !inputLayout.isForSubgroup() ||
1974 !targetLayout.isForSubgroup())
1975 return rewriter.notifyMatchFailure(op,
"both layouts must be lane level");
1977 auto valueTy = dyn_cast<VectorType>(op.getResult().getType());
1979 return rewriter.notifyMatchFailure(op,
"value type must be a vector");
1980 Type elemTy = valueTy.getElementType();
1984 return rewriter.notifyMatchFailure(
1985 op,
"element type must be a whole number of bytes");
1988 int64_t rank = sgShape.size();
1989 if (rank != inputLayout.getRank() || rank != targetLayout.getRank())
1990 return rewriter.notifyMatchFailure(
1991 op,
"both layouts must have the rank of the value");
1993 FailureOr<VectorType> distInputTy =
1995 FailureOr<VectorType> distTargetTy =
1997 if (failed(distInputTy) || failed(distTargetTy))
1998 return rewriter.notifyMatchFailure(
1999 op,
"value type must be distributable by both layouts");
2004 return rewriter.notifyMatchFailure(
2005 op,
"target attribute is required to determine the subgroup size");
2006 FailureOr<int64_t> numSubgroups =
2008 if (failed(numSubgroups))
2009 return rewriter.notifyMatchFailure(
2010 op,
"the scratch holds one tile per subgroup, so @known_block_size "
2011 "must be attached to the kernel, with power-of-two dimensions "
2012 "covering at least one subgroup");
2018 slmShape[0] *= *numSubgroups;
2021 auto slmTy = MemRefType::get({slmBytes}, rewriter.getI8Type(), {}, 3);
2022 Value slm = memref::AllocaOp::create(rewriter, loc, slmTy);
2023 Value memDesc = xegpu::CreateMemDescOp::create(
2025 xegpu::MemDescType::get(rewriter.getContext(), slmShape, elemTy,
2030 Value sgId = gpu::SubgroupIdOp::create(rewriter, loc,
2031 rewriter.getIndexType(),
nullptr);
2033 base[0] = arith::MulIOp::create(
2034 rewriter, loc, sgId,
2036 for (
int64_t d = 1; d < rank; ++d)
2039 Value laneId = gpu::LaneIdOp::create(rewriter, loc, rewriter.getIndexType(),
2040 mlir::IntegerAttr());
2041 auto dynamicOffsets = rewriter.getDenseI64ArrayAttr(
2047 for (
auto [coord,
b] : llvm::zip_equal(coords, base))
2048 offsets.push_back(arith::AddIOp::create(rewriter, loc, coord,
b));
2054 inputLayout.computeDistributedCoords(rewriter, loc, laneId, sgShape);
2055 if (failed(storeCoords))
2056 return rewriter.notifyMatchFailure(
2057 op,
"failed to compute the input_layout coordinates");
2059 inputLayout.getEffectiveLaneLayoutAsInt();
2061 inputLayout.getEffectiveLaneDataAsInt();
2065 for (
auto [unit, coords] : llvm::enumerate(*storeCoords)) {
2066 LaneFragmentPiece piece =
2067 getLaneFragmentPiece(unit, sgShape, inputLaneLayout, inputLaneData);
2068 Value data = fragment;
2070 data = vector::ExtractStridedSliceOp::create(
2071 rewriter, loc, fragment, piece.offsets, piece.sizes,
2073 xegpu::StoreMatrixOp::create(rewriter, loc,
TypeRange{}, data, memDesc,
2074 ValueRange(withBase(coords)), dynamicOffsets,
2075 nullptr, xegpu::DistributeLayoutAttr{});
2081 xegpu::FenceOp::create(rewriter, loc, xegpu::MemorySpace::SLM,
2082 xegpu::FenceScope::Workgroup);
2086 targetLayout.computeDistributedCoords(rewriter, loc, laneId, sgShape);
2087 if (failed(loadCoords))
2088 return rewriter.notifyMatchFailure(
2089 op,
"failed to compute the target_layout coordinates");
2091 targetLayout.getEffectiveLaneLayoutAsInt();
2093 targetLayout.getEffectiveLaneDataAsInt();
2095 if (loadCoords->size() > 1)
2096 result = arith::ConstantOp::create(rewriter, loc, *distTargetTy,
2097 rewriter.getZeroAttr(*distTargetTy));
2098 for (
auto [unit, coords] : llvm::enumerate(*loadCoords)) {
2099 LaneFragmentPiece piece =
2100 getLaneFragmentPiece(unit, sgShape, targetLaneLayout, targetLaneData);
2101 auto pieceTy = VectorType::get(piece.sizes, elemTy);
2102 Value loaded = xegpu::LoadMatrixOp::create(
2103 rewriter, loc, pieceTy, memDesc,
ValueRange(withBase(coords)),
2104 dynamicOffsets,
nullptr, xegpu::DistributeLayoutAttr{});
2109 result = vector::InsertStridedSliceOp::create(
2110 rewriter, loc, loaded,
result, piece.offsets,
2114 rewriter.replaceOp(op, castValueTo(rewriter,
2123static bool hasDefaultOrderAndUnitLaneData(xegpu::DistributeLayoutAttr layout) {
2124 if (layout.getRank() != 2)
2132struct ElementLaneRedistribution {
2134 VectorType valueType;
2136 VectorType distributedInput;
2137 VectorType distributedTarget;
2151static FailureOr<ElementLaneRedistribution>
2152matchElementLaneRedistribution(xegpu::ConvertLayoutOp op,
2153 ConversionPatternRewriter &rewriter) {
2154 xegpu::DistributeLayoutAttr inputLayout = op.getEffectiveInputLayout();
2155 xegpu::DistributeLayoutAttr targetLayout = op.getTargetLayoutAttr();
2156 if (!hasDefaultOrderAndUnitLaneData(inputLayout) ||
2157 !hasDefaultOrderAndUnitLaneData(targetLayout))
2158 return rewriter.notifyMatchFailure(
2159 op,
"both layouts must be rank 2 with effective lane_data [1, 1] and "
2160 "effective order [1, 0]");
2162 auto valueType = dyn_cast<VectorType>(op.getResult().getType());
2164 return rewriter.notifyMatchFailure(op,
"value type must be a vector");
2166 FailureOr<VectorType> distributedInput =
2168 FailureOr<VectorType> distributedTarget =
2170 if (failed(distributedInput) || failed(distributedTarget))
2171 return rewriter.notifyMatchFailure(
2172 op,
"value type must be distributable by both layouts");
2177 return rewriter.notifyMatchFailure(
2178 op,
"target attribute is required to determine the subgroup size");
2180 ElementLaneRedistribution redistribution;
2181 redistribution.valueType = valueType;
2182 redistribution.distributedInput = *distributedInput;
2183 redistribution.distributedTarget = *distributedTarget;
2184 redistribution.inputLaneLayout = inputLayout.getEffectiveLaneLayoutAsInt();
2185 redistribution.targetLaneLayout = targetLayout.getEffectiveLaneLayoutAsInt();
2186 redistribution.subgroupSize = uArch->getSubgroupSize();
2187 return redistribution;
2202static FailureOr<int64_t> getDistributedDimLaneStride(xegpu::SliceAttr slice) {
2203 xegpu::SliceAttr flattened = slice.flatten();
2204 auto parent = dyn_cast<xegpu::LayoutAttr>(flattened.getParent());
2210 if (parentLaneLayout.size() != parentOrder.size())
2214 std::optional<int64_t> distributedDim;
2215 for (
int64_t dim = 0, rank =
static_cast<int64_t>(parentLaneLayout.size());
2216 dim < rank; ++dim) {
2217 if (llvm::is_contained(slicedDims, dim) || parentLaneLayout[dim] == 1)
2221 distributedDim = dim;
2223 if (!distributedDim)
2227 for (
int64_t dim : parentOrder) {
2228 if (dim == *distributedDim)
2230 stride *= parentLaneLayout[dim];
2246struct SgToLaneConvertLayoutBroadcastExtract
2247 :
public OpConversionPattern<xegpu::ConvertLayoutOp> {
2248 using OpConversionPattern<xegpu::ConvertLayoutOp>::OpConversionPattern;
2251 matchAndRewrite(xegpu::ConvertLayoutOp op, OpAdaptor adaptor,
2252 ConversionPatternRewriter &rewriter)
const override {
2253 xegpu::DistributeLayoutAttr inputLayout = op.getEffectiveInputLayout();
2254 xegpu::DistributeLayoutAttr targetLayout = op.getTargetLayoutAttr();
2256 if (!isa<xegpu::SliceAttr>(inputLayout))
2257 return rewriter.notifyMatchFailure(op,
2258 "input_layout must be #xegpu.slice");
2259 if (!isa<xegpu::LayoutAttr>(targetLayout))
2260 return rewriter.notifyMatchFailure(op,
2261 "target_layout must be #xegpu.layout");
2263 FailureOr<ElementLaneRedistribution> redistribution =
2264 matchElementLaneRedistribution(op, rewriter);
2265 if (failed(redistribution))
2269 return rewriter.notifyMatchFailure(
2270 op,
"input_layout effective lane_layout must be [1, 1]");
2272 if (targetLaneLayout[0] <= 1 || targetLaneLayout[1] != 1)
2273 return rewriter.notifyMatchFailure(
2274 op,
"target_layout effective lane_layout must be [n, 1] with n > 1");
2275 if (redistribution->distributedInput != redistribution->valueType)
2276 return rewriter.notifyMatchFailure(
2277 op,
"distributed input_layout type must equal the value type");
2279 return rewriter.notifyMatchFailure(
2280 op,
"distributed target_layout type must be vector<1x1>");
2281 if (targetLaneLayout[0] > redistribution->subgroupSize)
2282 return rewriter.notifyMatchFailure(
2283 op,
"target_layout effective lane_layout[0] must not exceed the "
2286 VectorType valueType = redistribution->valueType;
2290 redistribution->distributedInput);
2291 auto flatType = VectorType::get({valueType.getNumElements()},
2292 valueType.getElementType());
2293 Value flat = vector::ShapeCastOp::create(rewriter, loc, flatType, src);
2294 Value laneId = gpu::LaneIdOp::create(rewriter, loc, rewriter.getIndexType(),
2295 mlir::IntegerAttr());
2298 Value row = arith::RemUIOp::create(rewriter, loc, laneId, rowCount);
2299 Value element = vector::ExtractOp::create(rewriter, loc, flat,
2301 rewriter.replaceOpWithNewOp<vector::FromElementsOp>(
2302 op, redistribution->distributedTarget, element);
2348struct SgToLaneConvertLayoutPartialBroadcastExtractShuffle
2349 :
public OpConversionPattern<xegpu::ConvertLayoutOp> {
2350 using OpConversionPattern<xegpu::ConvertLayoutOp>::OpConversionPattern;
2353 matchAndRewrite(xegpu::ConvertLayoutOp op, OpAdaptor adaptor,
2354 ConversionPatternRewriter &rewriter)
const override {
2355 xegpu::DistributeLayoutAttr inputLayout = op.getEffectiveInputLayout();
2356 xegpu::DistributeLayoutAttr targetLayout = op.getTargetLayoutAttr();
2358 auto inputSlice = dyn_cast<xegpu::SliceAttr>(inputLayout);
2360 return rewriter.notifyMatchFailure(op,
2361 "input_layout must be #xegpu.slice");
2362 if (!isa<xegpu::LayoutAttr>(targetLayout))
2363 return rewriter.notifyMatchFailure(op,
2364 "target_layout must be #xegpu.layout");
2366 FailureOr<ElementLaneRedistribution> redistribution =
2367 matchElementLaneRedistribution(op, rewriter);
2368 if (failed(redistribution))
2371 VectorType valueType = redistribution->valueType;
2372 Type elementType = valueType.getElementType();
2374 !llvm::is_contained({8u, 16u, 32u, 64u},
2376 return rewriter.notifyMatchFailure(
2377 op,
"element type must be an int or float of bit width 8, 16, 32 or "
2378 "64 to be carried by gpu.shuffle");
2381 return rewriter.notifyMatchFailure(
2382 op,
"input_layout effective lane_layout must be [1, 2]");
2384 if (targetLaneLayout[0] <= 1 || targetLaneLayout[1] != 1)
2385 return rewriter.notifyMatchFailure(
2386 op,
"target_layout effective lane_layout must be [n, 1] with n > 1");
2387 if (redistribution->distributedInput.getShape() !=
2389 return rewriter.notifyMatchFailure(
2390 op,
"distributed input_layout type must be vector<shape[0]x1>");
2392 return rewriter.notifyMatchFailure(
2393 op,
"distributed target_layout type must be vector<1x2>");
2395 FailureOr<int64_t> stride = getDistributedDimLaneStride(inputSlice);
2397 return rewriter.notifyMatchFailure(
2398 op,
"input_layout parent must have exactly one non-sliced dimension "
2399 "with lane_layout extent greater than one");
2402 if (*stride * redistribution->inputLaneLayout[1] !=
2403 redistribution->subgroupSize)
2404 return rewriter.notifyMatchFailure(
2405 op,
"input_layout distributed dimension lane stride times its extent "
2406 "must equal the subgroup size");
2410 if (*stride % targetLaneLayout[0] != 0)
2411 return rewriter.notifyMatchFailure(
2412 op,
"target_layout effective lane_layout[0] must divide the "
2413 "input_layout distributed dimension lane stride");
2418 redistribution->distributedInput);
2420 VectorType::get({redistribution->distributedInput.getNumElements()},
2421 valueType.getElementType());
2422 Value flat = vector::ShapeCastOp::create(rewriter, loc, flatType, src);
2423 Value laneId = gpu::LaneIdOp::create(rewriter, loc, rewriter.getIndexType(),
2424 mlir::IntegerAttr());
2427 Value row = arith::RemUIOp::create(rewriter, loc, laneId, rowCount);
2428 Value own = vector::ExtractOp::create(rewriter, loc, flat,
2431 Type i32Type = rewriter.getI32Type();
2433 redistribution->subgroupSize);
2434 Value rowI32 = arith::IndexCastOp::create(rewriter, loc, i32Type, row);
2437 Value otherOwner = arith::AddIOp::create(rewriter, loc, rowI32, strideI32);
2438 Value column0 = gpu::ShuffleOp::create(rewriter, loc, own, rowI32, width,
2439 gpu::ShuffleMode::IDX)
2440 .getShuffleResult();
2441 Value column1 = gpu::ShuffleOp::create(rewriter, loc, own, otherOwner,
2442 width, gpu::ShuffleMode::IDX)
2443 .getShuffleResult();
2444 rewriter.replaceOpWithNewOp<vector::FromElementsOp>(
2445 op, redistribution->distributedTarget,
ValueRange{column0, column1});
2460struct SgToLaneConvertLayoutDeinterleaveSelect
2461 :
public OpConversionPattern<xegpu::ConvertLayoutOp> {
2462 using OpConversionPattern<xegpu::ConvertLayoutOp>::OpConversionPattern;
2465 matchAndRewrite(xegpu::ConvertLayoutOp op, OpAdaptor adaptor,
2466 ConversionPatternRewriter &rewriter)
const override {
2467 xegpu::DistributeLayoutAttr inputLayout = op.getEffectiveInputLayout();
2468 xegpu::DistributeLayoutAttr targetLayout = op.getTargetLayoutAttr();
2470 if (!isa<xegpu::SliceAttr>(inputLayout))
2471 return rewriter.notifyMatchFailure(op,
2472 "input_layout must be #xegpu.slice");
2473 auto targetSlice = dyn_cast<xegpu::SliceAttr>(targetLayout);
2475 return rewriter.notifyMatchFailure(op,
2476 "target_layout must be #xegpu.slice");
2478 FailureOr<ElementLaneRedistribution> redistribution =
2479 matchElementLaneRedistribution(op, rewriter);
2480 if (failed(redistribution))
2483 VectorType valueType = redistribution->valueType;
2485 return rewriter.notifyMatchFailure(
2486 op,
"input_layout effective lane_layout must be [1, 1]");
2489 return rewriter.notifyMatchFailure(
2490 op,
"target_layout effective lane_layout must be [1, 2]");
2491 if (redistribution->distributedInput != valueType)
2492 return rewriter.notifyMatchFailure(
2493 op,
"distributed input_layout type must equal the value type");
2494 if (redistribution->distributedTarget.getShape() !=
2496 return rewriter.notifyMatchFailure(
2497 op,
"distributed target_layout type must be vector<shape[0]x1>");
2499 FailureOr<int64_t> stride = getDistributedDimLaneStride(targetSlice);
2501 return rewriter.notifyMatchFailure(
2502 op,
"target_layout parent must have exactly one non-sliced dimension "
2503 "with lane_layout extent greater than one");
2506 if (*stride * targetLaneLayout[1] != redistribution->subgroupSize)
2507 return rewriter.notifyMatchFailure(
2508 op,
"target_layout distributed dimension lane stride times its "
2509 "extent must equal the subgroup size");
2514 redistribution->distributedInput);
2515 auto flatType = VectorType::get({valueType.getNumElements()},
2516 valueType.getElementType());
2517 Value flat = vector::ShapeCastOp::create(rewriter, loc, flatType, src);
2518 auto deinterleaved = vector::DeinterleaveOp::create(rewriter, loc, flat);
2519 Value laneId = gpu::LaneIdOp::create(rewriter, loc, rewriter.getIndexType(),
2520 mlir::IntegerAttr());
2523 Value group = arith::DivUIOp::create(rewriter, loc, laneId, strideVal);
2524 Value isFirstGroup = arith::CmpIOp::create(
2525 rewriter, loc, arith::CmpIPredicate::eq, group, zero);
2526 Value selected = arith::SelectOp::create(rewriter, loc, isFirstGroup,
2527 deinterleaved.getRes1(),
2528 deinterleaved.getRes2());
2529 rewriter.replaceOpWithNewOp<vector::ShapeCastOp>(
2530 op, redistribution->distributedTarget, selected);
2536struct SgToLaneVectorInterleave
2537 :
public OpConversionPattern<vector::InterleaveOp> {
2538 using OpConversionPattern<vector::InterleaveOp>::OpConversionPattern;
2541 matchAndRewrite(vector::InterleaveOp op, OpAdaptor adaptor,
2542 ConversionPatternRewriter &rewriter)
const override {
2544 auto newOp = vector::InterleaveOp::create(
2545 rewriter, op.getLoc(), adaptor.getLhs(), adaptor.getRhs());
2546 rewriter.replaceOp(op, newOp.getResult());
2552struct SgToLaneVectorDeinterleave
2553 :
public OpConversionPattern<vector::DeinterleaveOp> {
2554 using OpConversionPattern<vector::DeinterleaveOp>::OpConversionPattern;
2557 matchAndRewrite(vector::DeinterleaveOp op, OpAdaptor adaptor,
2558 ConversionPatternRewriter &rewriter)
const override {
2560 auto newOp = vector::DeinterleaveOp::create(rewriter, op.getLoc(),
2561 adaptor.getSource());
2562 rewriter.replaceOp(op, newOp.getResults());
2567struct SgToLaneDpasMx :
public OpConversionPattern<xegpu::DpasMxOp> {
2568 using OpConversionPattern<xegpu::DpasMxOp>::OpConversionPattern;
2571 matchAndRewrite(xegpu::DpasMxOp op, OpAdaptor adaptor,
2572 ConversionPatternRewriter &rewriter)
const override {
2577 if (!uArch->isSupportedInstruction(
2579 return rewriter.notifyMatchFailure(
2580 op,
"target uArch does not support scaled subgroup mma");
2582 auto layoutA = cast<xegpu::LayoutAttr>(op.getLayoutAAttr());
2583 auto layoutB = cast<xegpu::LayoutAttr>(op.getLayoutBAttr());
2584 auto layoutCd = cast<xegpu::LayoutAttr>(op.getLayoutCdAttr());
2585 if (!layoutA || !layoutB || !layoutCd)
2586 return rewriter.notifyMatchFailure(
2587 op,
"missing required layout attributes for DpasMxOp distribution");
2590 auto expected1DTypeResult =
2592 auto expected1DTypeA =
2594 auto expected1DTypeB =
2597 VectorType expected1DTypeScaleA, expected1DTypeScaleB;
2598 if (op.getScaleA()) {
2599 auto layoutScaleA = cast<xegpu::LayoutAttr>(op.getLayoutAScaleAttr());
2601 cast<VectorType>(op.getScaleA().getType()), layoutScaleA);
2602 if (failed(expected1DTypeScaleAOrFailure))
2603 return rewriter.notifyMatchFailure(
2604 op,
"failed to calculate expected 1D vector type for scale A");
2605 expected1DTypeScaleA = expected1DTypeScaleAOrFailure.value();
2607 if (op.getScaleB()) {
2608 auto layoutScaleB = cast<xegpu::LayoutAttr>(op.getLayoutBScaleAttr());
2610 cast<VectorType>(op.getScaleB().getType()), layoutScaleB);
2611 if (failed(expected1DTypeScaleBOrFailure))
2612 return rewriter.notifyMatchFailure(
2613 op,
"failed to calculate expected 1D vector type for scale B");
2614 expected1DTypeScaleB = expected1DTypeScaleBOrFailure.value();
2617 auto expectedNDTypeResult =
2619 if (failed(expected1DTypeResult) || failed(expected1DTypeA) ||
2620 failed(expected1DTypeB))
2621 return rewriter.notifyMatchFailure(
2623 "failed to calculate supported workitem 1D vector types for DpasOp "
2625 if (failed(expectedNDTypeResult))
2626 return rewriter.notifyMatchFailure(
2627 op,
"unable to compute expected workitem vector type for DpasOp from "
2631 const auto *uArchInstruction = dyn_cast<
2634 assert(uArchInstruction);
2635 auto wiAType = expected1DTypeA.value();
2636 auto wiBType = expected1DTypeB.value();
2638 unsigned aPackedBitWidth =
2639 wiAType.getElementTypeBitWidth() * wiAType.getNumElements();
2640 unsigned bPackedBitWidth =
2641 wiBType.getElementTypeBitWidth() * wiBType.getNumElements();
2642 if (aPackedBitWidth % uArchInstruction->getPackedFormatBitSizeA())
2643 return rewriter.notifyMatchFailure(
2644 op,
"A operand packed bit width must be a multiple of uArch packed "
2645 "format requirement");
2646 if (bPackedBitWidth % uArchInstruction->getPackedFormatBitSizeB())
2647 return rewriter.notifyMatchFailure(
2648 op,
"B operand packed bit width must be a multiple of uArch packed "
2649 "format requirement");
2651 auto newOp = xegpu::DpasMxOp::create(
2652 rewriter, op->getLoc(), expected1DTypeResult.value(),
2654 expected1DTypeA.value()),
2656 expected1DTypeB.value()),
2658 ? castValueTo(rewriter,
2660 expected1DTypeResult.value())
2664 ? castValueTo(rewriter,
2666 expected1DTypeScaleA)
2669 ? castValueTo(rewriter,
2671 expected1DTypeScaleB)
2677 rewriter.replaceOp(op, castValueTo(rewriter, newOp.getResult(),
2678 expectedNDTypeResult.value()));
2683struct XeGPUSgToLaneDistributePass
2684 :
public xegpu::impl::XeGPUSgToLaneDistributeBase<
2685 XeGPUSgToLaneDistributePass> {
2686 void runOnOperation()
override;
2691void XeGPUSgToLaneDistributePass::runOnOperation() {
2694 Operation *root = getOperation();
2696 signalPassFailure();
2701 llvm::SmallSetVector<UnrealizedConversionCastOp, 8> existingCasts;
2703 [&](UnrealizedConversionCastOp castOp) { existingCasts.insert(castOp); });
2709 TypeConverter typeConverter;
2713 auto materializeCast = [](OpBuilder &builder, Type type,
ValueRange inputs,
2714 Location loc) -> Value {
2715 return UnrealizedConversionCastOp::create(builder, loc, type, inputs)
2718 typeConverter.addSourceMaterialization(materializeCast);
2719 typeConverter.addTargetMaterialization(materializeCast);
2724 typeConverter, patterns,
target, root);
2725 target.addLegalOp<UnrealizedConversionCastOp>();
2726 if (
failed(applyPartialConversion(root,
target, std::move(patterns))))
2727 return signalPassFailure();
2738 typeConverter.addConversion([](
Type type) ->
Type {
return type; });
2740 typeConverter.addConversion([](TensorDescType type) ->
Type {
2741 if (type.getLayoutAttr()) {
2742 return type.dropLayouts();
2750 auto getSubShapeAndCount = [](VectorType vecTy,
2751 xegpu::DistributeLayoutAttr layout)
2754 if (failed(distTyOrFailure))
2761 std::move(loopArgTypes));
2769 target.addDynamicallyLegalOp<xegpu::CreateNdDescOp>(
2770 [&](xegpu::CreateNdDescOp op) {
return !op.getType().getLayoutAttr(); });
2772 target.addDynamicallyLegalDialect<xegpu::XeGPUDialect>([](
Operation *op) {
2773 if (isa<xegpu::ConvertLayoutOp>(op))
2775 auto anchorOp = dyn_cast<AnchorLayoutInterface>(op);
2778 return !anchorOp.getAnchorLayout();
2781 target.addDynamicallyLegalOp<arith::ConstantOp>(
2782 [=](arith::ConstantOp op) ->
bool {
2784 if (!isa<VectorType>(op.getResult().getType()))
2790 target.addDynamicallyLegalDialect<math::MathDialect, arith::ArithDialect>(
2791 [=](
Operation *op) -> std::optional<bool> {
2796 if (op->getNumResults() != 1)
2799 VectorType resultType =
2800 dyn_cast<VectorType>(op->getResult(0).getType());
2805 for (
Value operand : op->getOperands()) {
2806 VectorType operandType = dyn_cast<VectorType>(operand.getType());
2807 if (!operandType || operandType.getShape() != resultType.getShape()) {
2815 target.addDynamicallyLegalOp<vector::ReductionOp>(
2816 [=](vector::ReductionOp op) ->
bool {
2821 target.addDynamicallyLegalOp<vector::MultiDimReductionOp>(
2822 [=](vector::MultiDimReductionOp op) ->
bool {
2823 return !isValidSubgroupMultiReductionOp(op);
2825 target.addDynamicallyLegalOp<vector::CreateMaskOp, vector::ConstantMaskOp,
2826 vector::TransposeOp, vector::BitCastOp,
2827 vector::ShapeCastOp, vector::StepOp,
2828 vector::BroadcastOp>([=](
Operation *op) ->
bool {
2831 target.addDynamicallyLegalOp<vector::ExtractOp>(
2832 [=](vector::ExtractOp op) ->
bool {
2833 if (!isa<VectorType>(op.getType()))
2837 target.addDynamicallyLegalOp<vector::InsertOp>(
2838 [=](vector::InsertOp op) ->
bool {
2841 target.addDynamicallyLegalOp<vector::ExtractStridedSliceOp>(
2842 [=](vector::ExtractStridedSliceOp op) ->
bool {
2845 target.addDynamicallyLegalOp<vector::InsertStridedSliceOp>(
2846 [=](vector::InsertStridedSliceOp op) ->
bool {
2849 target.addDynamicallyLegalOp<vector::InterleaveOp, vector::DeinterleaveOp>(
2853 target.markUnknownOpDynamicallyLegal([](
Operation *op) {
return true; });
2855 .
add<SgToLaneCreateNdDesc, SgToLaneLoadNd, SgToLaneStoreNd, SgToLaneDpas,
2856 SgToLaneElementWise, SgToLaneArithConstant, SgToLanePrefetchNd,
2857 SgToLaneLoadGather, SgToLaneStoreScatter, SgToLaneVectorReduction,
2858 SgToLaneMultiDimReduction, SgToLaneVectorExtract,
2859 SgToLaneVectorInsert, SgToLaneVectorExtractStridedSlice,
2860 SgToLaneVectorInsertStridedSlice, SgToLaneLoadMatrix,
2861 SgToLaneStoreMatrix, SgToLaneConvertLayout, SgToLaneVectorTranspose,
2862 SgToLaneVectorBitcast, SgToLaneVectorStep, SgToLaneVectorShapeCast,
2863 SgToLaneBroadcast, SgToLaneCreateMask<vector::CreateMaskOp>,
2864 SgToLaneCreateMask<vector::ConstantMaskOp>,
2865 SgToLaneVectorDeinterleave, SgToLaneVectorInterleave, SgToLaneDpasMx,
2866 SgToLaneConvertLayoutBroadcastExtract,
2867 SgToLaneConvertLayoutPartialBroadcastExtractShuffle,
2868 SgToLaneConvertLayoutDeinterleaveSelect>(typeConverter,
2872 patterns.
add<SgToLaneConvertLayoutViaSLM>(typeConverter,
Attributes are known-constant values of operations.
static DenseElementsAttr get(ShapedType type, ArrayRef< Attribute > values)
Constructs a dense elements attribute from an array of element values.
This class defines the main interface for locations in MLIR and acts as a non-nullable wrapper around...
MLIRContext is the top-level object for a collection of MLIR operations.
Operation is the basic unit of execution within MLIR.
OpResult getResult(unsigned idx)
Get the 'idx'th result of this operation.
Location getLoc()
The source location the operation was defined or derived from.
Attribute getPropertiesAsAttribute()
Return the properties converted to an attribute.
OperationName getName()
The name of an operation is the key identifier for it.
DictionaryAttr getDiscardableAttrDictionary()
Return all of the discardable attributes on this operation as a DictionaryAttr.
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),...
unsigned getNumResults()
Return the number of results held by this operation.
This class represents the benefit of a pattern match in a unitless scheme that ranges from 0 (very li...
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.
A range-style iterator that allows for iterating over the offsets of all potential tiles of size tile...
This class provides an abstraction over the various different ranges of value types.
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
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 provides an abstraction over the different types of ranges over Values.
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.
static ConstantIndexOp create(OpBuilder &builder, Location location, int64_t value)
static ConstantIntOp create(OpBuilder &builder, Location location, int64_t value, unsigned width)
static DenseArrayAttrImpl get(MLIRContext *context, ArrayRef< int64_t > content)
bool hasElementwiseMappableTraits(Operation *op)
Together, Elementwise, Scalarizable, Vectorizable, and Tensorizable provide an easy way for scalar op...
void populateSCFStructuralTypeConversionsAndLegality(const TypeConverter &typeConverter, RewritePatternSet &patterns, ConversionTarget &target, PatternBenefit benefit=1)
Populates patterns for SCF structural type conversions and sets up the provided ConversionTarget with...
Value makeArithReduction(OpBuilder &b, Location loc, CombiningKind kind, Value v1, Value acc, arith::FastMathFlagsAttr fastmath=nullptr, Value mask=nullptr)
Returns the result value of reducing two scalar/vector values with the corresponding arith operation.
SmallVector< Value > getAsValues(OpBuilder &builder, Location loc, ArrayRef< OpFoldResult > foldResults)
Convert foldResults into Values.
@ SubgroupMatrixMultiplyAcc
@ SubgroupScaledMatrixMultiplyAcc
const uArch * getUArch(llvm::StringRef archName)
void populateXeGPUSgToLaneDistributeTypeConversionAndLegality(TypeConverter &typeConverter, RewritePatternSet &patterns, ConversionTarget &target, Operation *topLevelOp)
Defines type conversions and legality for XeGPU subgroup to lane distribution and appends the require...
bool requirePacked(const DistributeLayoutAttr layout)
Helper function to check if the layout is packed.
void removeTemporaryLayoutAttrs(Operation *op)
Removes the temporary layout attributes for each OpOperand and OpResult of the given operation.
Value subgroupReduction(Location loc, OpBuilder &builder, Value input, vector::CombiningKind kind, uint32_t size)
Given an input value representing per-lane data, this function returns the result after performing a ...
bool recoverTemporaryLayouts(Operation *rootOp)
Attach layout attributes to all vector-type operands of operations within the given operation's neste...
FailureOr< int64_t > getNumSubgroupsFromBlockSize(Operation *op, int64_t subgroupSize)
Returns the number of subgroups the kernel enclosing op runs, derived from the known_block_size of it...
FailureOr< VectorType > getDistVecTypeBasedOnLaneLayout(DistributeLayoutAttr layout, VectorType originalType)
Helper function to get distributed vector type for a source vector type according to the lane_layout.
Value lowerToVectorReductions(TypedValue< VectorType > src, TypedValue< VectorType > acc, vector::CombiningKind kind, int64_t reductionDim, Location loc, PatternRewriter &rewriter)
Given a src and an acc argumments from a vector::MultiDimReductionOp, lower to a set of vector::Reduc...
bool requireTranspose(const DistributeLayoutAttr layout, const uArch::uArch *uArch)
Helper function to check if the layout requires a transpose effect.
DistributeLayoutAttr getDistributeLayoutAttr(const Value value)
Retrieves the DistributeLayoutAttr associated with a given Value, or nullptr if none is found.
DenseMap< Value, SmallVector< Type > > precomputeLoopBlockArgTypes(Operation *topLevelOp, SubShapeAndCountFn getSubShapeAndCount)
Pre-computes distributed VectorType mappings for every value carried through an SCF loop under topLev...
std::optional< std::string > getChipStr(Operation *op)
Retrieves the chip string from the XeVM target attribute of the parent GPU module operation.
void addVectorTypeConversion(TypeConverter &converter, SubShapeAndCountFn getSubShapeAndCount, DenseMap< Value, SmallVector< Type > > loopArgTypes)
Adds a context-aware VectorType conversion to converter (1:1 shape-changing or 1:N,...
DistributeLayoutAttr getTemporaryLayout(const T &operandOrResult)
get and set distribute layout attribute for non-anchor operations (and offsets/masks of load/store op...
void populateXeGPUSgToLaneDistributeTypeConversions(TypeConverter &typeConverter, Operation *topLevelOp)
Define only the type conversions needed for XeGPU subgroup to lane distribution.
Value lowerCrossLaneReductionToShuffles(TypedValue< VectorType > src, TypedValue< VectorType > acc, vector::CombiningKind kind, int64_t reductionDim, int64_t reductionSize, Location loc, PatternRewriter &rewriter)
Lowers cross-lane reductions to shuffle operations on a 2D vector.
void cleanupUnrealizedConversionCasts(Operation *root, const llvm::SmallSetVector< UnrealizedConversionCastOp, 8 > &existingCasts)
Cleans up UnrealizedConversionCastOps inserted during SCF structural type conversion and/or XeGPU unr...
SmallVector< OpFoldResult > addWithRightAligned(OpBuilder &builder, Location loc, ArrayRef< OpFoldResult > lhs, ArrayRef< OpFoldResult > rhs)
Generates element-wise addition ops of two arrays with automatic alignment.
FailureOr< VectorType > getDistributedVectorType(xegpu::TensorDescType tdescTy)
If tensor descriptor has a layout attribute it is used in SIMT mode.
Include the generated interface declarations.
detail::DenseArrayAttrImpl< int64_t > DenseI64ArrayAttr
std::optional< int64_t > getConstantIntValue(OpFoldResult ofr)
If ofr is a constant integer or an IntegerAttr, return the integer.
SmallVector< int64_t > computeStrides(ArrayRef< int64_t > sizes)
SmallVector< int64_t > delinearize(int64_t linearIndex, ArrayRef< int64_t > strides)
Given the strides together with a linear index in the dimension space, return the vector-space offset...
int64_t computeProduct(ArrayRef< int64_t > basis)
Self-explicit.
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.
OpFoldResult getAsOpFoldResult(Value val)
Given a value, try to extract a constant Attribute.
int64_t linearize(ArrayRef< int64_t > offsets, ArrayRef< int64_t > basis)
Return the linearized index of 'offsets' w.r.t.
std::optional< SmallVector< int64_t > > computeShapeRatio(ArrayRef< int64_t > shape, ArrayRef< int64_t > subShape)
Return the multi-dimensional integral ratio of subShape to the trailing dimensions of shape.
This represents an operation in an abstracted form, suitable for use with the builder APIs.
void addOperands(ValueRange newOperands)
void addAttribute(StringRef name, Attribute attr)
Add an attribute with the specified name.
void addTypes(ArrayRef< Type > newTypes)
Attribute propertiesAttr
This Attribute is used to opaquely construct the properties of the operation.