25#include "llvm/ADT/SetVector.h"
30#define GEN_PASS_DEF_XEGPUWGTOSGDISTRIBUTE
31#include "mlir/Dialect/XeGPU/Transforms/Passes.h.inc"
40static xegpu::RangeAttr getRangeSpecAttr(
Operation *op) {
43 if (
auto attr = llvm::dyn_cast_if_present<xegpu::RangeAttr>(
51static std::pair<SmallVector<int64_t>,
int>
53 xegpu::DistributeLayoutAttr layout) {
56 auto distributedShape = layout.computeDistributedShape(
58 if (
failed(distributedShape))
59 return std::make_pair(sgShape, count);
60 auto sgData = layout.getEffectiveSgDataAsInt();
62 return std::make_pair(sgData, count);
69template <
typename OpType,
70 typename = std::enable_if_t<llvm::is_one_of<
71 OpType, xegpu::LoadNdOp, xegpu::StoreNdOp, xegpu::PrefetchNdOp,
72 xegpu::LoadMatrixOp, xegpu::StoreMatrixOp>::value>>
74genOffsetsList(ConversionPatternRewriter &rewriter, OpType op,
79 if (origOffsets.empty())
83 xegpu::DistributeLayoutAttr layout;
84 if constexpr (std::is_same_v<OpType, xegpu::LoadMatrixOp> ||
85 std::is_same_v<OpType, xegpu::StoreMatrixOp>) {
86 layout = op.getLayoutAttr();
88 layout = op.getDescLayoutAttr();
92 if (!layout || !layout.isForWorkgroup())
96 gpu::SubgroupIdOp::create(rewriter, loc,
nullptr);
99 xegpu::RangeAttr sgIdRange = getRangeSpecAttr(op);
101 int64_t startOfRange = sgIdRange.getStart().getInt();
102 int64_t endOfRange = sgIdRange.getEnd().getInt();
104 if (layout.getNumSubgroups() != endOfRange - startOfRange)
105 return rewriter.notifyMatchFailure(
106 op,
"sg_layout size must match the sg_id_range");
108 if (startOfRange > 0) {
109 Value startOfRangeVal =
111 sgId = index::SubOp::create(rewriter, loc, sgId, startOfRangeVal);
118 auto maybeDescOffsets =
119 layout.computeDistributedCoords(rewriter, loc, sgId, wgShape);
120 if (
failed(maybeDescOffsets))
125 for (
const auto &sgOffsets : *maybeDescOffsets) {
128 offsetsList.push_back(std::move(newOffsets));
182struct WgToSgCreateNdOp :
public OpConversionPattern<xegpu::CreateNdDescOp> {
183 using OpConversionPattern<xegpu::CreateNdDescOp>::OpConversionPattern;
186 matchAndRewrite(xegpu::CreateNdDescOp op, OneToNOpAdaptor adaptor,
187 ConversionPatternRewriter &rewriter)
const override {
189 Location loc = op.getLoc();
191 xegpu::TensorDescType tdescTy = op.getType();
192 auto layout = dyn_cast<xegpu::DistributeLayoutAttr>(tdescTy.getLayout());
193 if (!layout || !layout.isForWorkgroup())
196 Type elemTy = tdescTy.getElementType();
197 ArrayRef<int64_t> wgShape = tdescTy.getShape();
199 SmallVector<int64_t> sgShape;
201 std::tie(sgShape, count) = getSgShapeAndCount(wgShape, layout);
202 xegpu::TensorDescType newTdescTy =
203 xegpu::TensorDescType::get(ctx, sgShape, elemTy, tdescTy.getEncoding(),
204 layout.dropSgLayoutAndData());
206 Value src = op.getSource();
207 SmallVector<Value> newCreateNdOps(count);
208 std::generate(newCreateNdOps.begin(), newCreateNdOps.end(), [&]() -> Value {
209 if (isa<MemRefType>(src.getType()))
210 return xegpu::CreateNdDescOp::create(rewriter, loc, newTdescTy,
211 cast<TypedValue<MemRefType>>(src));
212 return xegpu::CreateNdDescOp::create(rewriter, loc, newTdescTy, src,
214 op.getMixedStrides());
217 rewriter.replaceOpWithMultiple(op, {newCreateNdOps});
223struct WgToSgLoadNdOp :
public OpConversionPattern<xegpu::LoadNdOp> {
224 using OpConversionPattern<xegpu::LoadNdOp>::OpConversionPattern;
226 matchAndRewrite(xegpu::LoadNdOp op, OneToNOpAdaptor adaptor,
227 ConversionPatternRewriter &rewriter)
const override {
229 SmallVector<SmallVector<OpFoldResult>> offsetsList;
230 if (
failed(genOffsetsList(rewriter, op, offsetsList)))
233 xegpu::DistributeLayoutAttr layout = op.getLayoutAttr();
235 layout = layout.dropSgLayoutAndData();
236 SmallVector<Value> newOps;
237 for (
auto [tdesc, offsets] :
238 llvm::zip(adaptor.getTensorDesc(), offsetsList)) {
239 auto tdescTy = dyn_cast<xegpu::TensorDescType>(tdesc.getType());
240 VectorType newResTy =
241 VectorType::get(tdescTy.getShape(), tdescTy.getElementType());
242 auto newOp = xegpu::LoadNdOp::create(
243 rewriter, op.getLoc(), newResTy, tdesc, offsets,
244 nullptr,
nullptr, op.getL1HintAttr(),
245 op.getL2HintAttr(), op.getL3HintAttr(), layout);
246 newOps.push_back(newOp);
248 rewriter.replaceOpWithMultiple(op, {newOps});
255struct WgToSgStoreNdOp :
public OpConversionPattern<xegpu::StoreNdOp> {
256 using OpConversionPattern<xegpu::StoreNdOp>::OpConversionPattern;
258 matchAndRewrite(xegpu::StoreNdOp op, OneToNOpAdaptor adaptor,
259 ConversionPatternRewriter &rewriter)
const override {
260 SmallVector<SmallVector<OpFoldResult>> offsetsList;
261 if (
failed(genOffsetsList(rewriter, op, offsetsList)))
264 xegpu::DistributeLayoutAttr layout = op.getLayoutAttr();
266 layout = layout.dropSgLayoutAndData();
267 for (
auto [v, tdesc, offsets] :
268 llvm::zip(adaptor.getValue(), adaptor.getTensorDesc(), offsetsList)) {
269 xegpu::StoreNdOp::create(rewriter, op.getLoc(), v, tdesc, offsets,
270 op.getL1HintAttr(), op.getL2HintAttr(),
271 op.getL3HintAttr(), layout);
273 rewriter.eraseOp(op);
280struct WgToSgPrefetchNdOp :
public OpConversionPattern<xegpu::PrefetchNdOp> {
281 using OpConversionPattern<xegpu::PrefetchNdOp>::OpConversionPattern;
283 matchAndRewrite(xegpu::PrefetchNdOp op, OneToNOpAdaptor adaptor,
284 ConversionPatternRewriter &rewriter)
const override {
285 SmallVector<SmallVector<OpFoldResult>> offsetsList;
286 if (
failed(genOffsetsList(rewriter, op, offsetsList)))
289 xegpu::DistributeLayoutAttr layout = op.getLayoutAttr();
291 layout = layout.dropSgLayoutAndData();
292 for (
auto [tdesc, offsets] :
293 llvm::zip(adaptor.getTensorDesc(), offsetsList)) {
294 xegpu::PrefetchNdOp::create(rewriter, op.getLoc(), tdesc, offsets,
295 op.getL1HintAttr(), op.getL2HintAttr(),
296 op.getL3HintAttr(), layout);
298 rewriter.eraseOp(op);
305struct WgToSgDpasOp :
public OpConversionPattern<xegpu::DpasOp> {
306 using OpConversionPattern<xegpu::DpasOp>::OpConversionPattern;
308 matchAndRewrite(xegpu::DpasOp op, OneToNOpAdaptor adaptor,
309 ConversionPatternRewriter &rewriter)
const override {
310 Location loc = op.getLoc();
311 VectorType resultTy = op.getResult().getType();
312 if (resultTy.getRank() < 2)
315 auto layoutCd = op.getLayoutCdAttr();
316 auto layoutA = op.getLayoutAAttr();
317 auto layoutB = op.getLayoutBAttr();
318 if (!layoutCd || !layoutA || !layoutB)
321 SmallVector<Value> newDpasOps;
322 for (
auto aVec : adaptor.getLhs()) {
323 for (
auto bVec : adaptor.getRhs()) {
327 tmpC = adaptor.getAcc()[i++];
329 ArrayRef<int64_t> aVecShape =
330 cast<VectorType>(aVec.getType()).getShape();
331 ArrayRef<int64_t> bVecShape =
332 cast<VectorType>(bVec.getType()).getShape();
335 SmallVector<int64_t> resShape(aVecShape.drop_back(2));
336 resShape.push_back(aVecShape[aVecShape.size() - 2]);
337 resShape.push_back(bVecShape[bVecShape.size() - 1]);
338 VectorType resTy = VectorType::get(resShape, resultTy.getElementType());
339 auto newDpasOp = xegpu::DpasOp::create(
340 rewriter, loc, resTy, aVec, bVec, tmpC,
341 nullptr,
nullptr,
nullptr);
342 newDpasOp.setLayoutCdAttr(layoutCd.dropSgLayoutAndData());
343 newDpasOp.setLayoutAAttr(layoutA.dropSgLayoutAndData());
344 newDpasOp.setLayoutBAttr(layoutB.dropSgLayoutAndData());
346 newDpasOps.push_back(newDpasOp);
349 rewriter.replaceOpWithMultiple(op, {newDpasOps});
355struct WgToSgDpasMxOp :
public OpConversionPattern<xegpu::DpasMxOp> {
356 using OpConversionPattern<xegpu::DpasMxOp>::OpConversionPattern;
358 matchAndRewrite(xegpu::DpasMxOp op, OneToNOpAdaptor adaptor,
359 ConversionPatternRewriter &rewriter)
const override {
361 Location loc = op.getLoc();
362 VectorType resultTy = op.getResult().getType();
364 if (resultTy.getRank() < 2)
367 auto layoutCd = op.getLayoutCdAttr();
368 auto layoutA = op.getLayoutAAttr();
369 auto layoutB = op.getLayoutBAttr();
370 auto layoutAScale = op.getLayoutAScaleAttr();
371 auto layoutBScale = op.getLayoutBScaleAttr();
373 if (!layoutCd || !layoutA || !layoutB || !layoutAScale || !layoutBScale)
377 SmallVector<Value> newDpasMxOps;
378 for (
auto [index_a, aVec] : llvm::enumerate(adaptor.getA())) {
379 for (
auto [index_b, bVec] : llvm::enumerate(adaptor.getB())) {
380 Value accVal = (op.getAcc()) ? adaptor.getAcc()[index_c++] : Value();
382 (op.getScaleA()) ? adaptor.getScaleA()[index_a] : Value();
384 (op.getScaleB()) ? adaptor.getScaleB()[index_b] : Value();
386 ArrayRef<int64_t> aVecShape =
387 cast<VectorType>(aVec.getType()).getShape();
388 ArrayRef<int64_t> bVecShape =
389 cast<VectorType>(bVec.getType()).getShape();
391 SmallVector<int64_t> resShape(aVecShape.drop_back(2));
392 resShape.push_back(aVecShape[aVecShape.size() - 2]);
393 resShape.push_back(bVecShape[bVecShape.size() - 1]);
394 VectorType resTy = VectorType::get(resShape, resultTy.getElementType());
395 auto newDpasMxOp = xegpu::DpasMxOp::create(
396 rewriter, loc, resTy, aVec, bVec, accVal, scaleAVal, scaleBVal,
397 layoutA.dropSgLayoutAndData(), layoutB.dropSgLayoutAndData(),
398 layoutCd.dropSgLayoutAndData(), layoutAScale.dropSgLayoutAndData(),
399 layoutBScale.dropSgLayoutAndData());
401 newDpasMxOps.push_back(newDpasMxOp);
404 rewriter.replaceOpWithMultiple(op, {newDpasMxOps});
410struct WgToSgVectorBroadcastOp
411 :
public OpConversionPattern<vector::BroadcastOp> {
412 using OpConversionPattern<vector::BroadcastOp>::OpConversionPattern;
415 matchAndRewrite(vector::BroadcastOp op, OneToNOpAdaptor adaptor,
416 ConversionPatternRewriter &rewriter)
const override {
418 VectorType resultType = op.getResult().getType();
419 ArrayRef<int64_t> wgShape = resultType.getShape();
421 xegpu::DistributeLayoutAttr layout =
423 if (!layout || !layout.isForWorkgroup())
426 SmallVector<int64_t> sgShape;
428 std::tie(sgShape, count) = getSgShapeAndCount(wgShape, layout);
429 VectorType newResultType =
430 VectorType::get(sgShape, resultType.getElementType());
432 SmallVector<Value> newBroadcastOps;
433 auto distSource = adaptor.getOperands().front();
434 int numDistributions = count / distSource.size();
435 for (
int i = 0; i < numDistributions; ++i) {
436 for (
auto operand : distSource) {
437 auto newBroadcast = vector::BroadcastOp::create(rewriter, op.getLoc(),
438 newResultType, operand);
440 newBroadcastOps.push_back(newBroadcast.getResult());
443 rewriter.replaceOpWithMultiple(op, {newBroadcastOps});
450 WgToSgElementwiseOp(MLIRContext *ctx)
451 : ConversionPattern(MatchAnyOpTypeTag(), 1, ctx) {}
454 matchAndRewrite(Operation *op, ArrayRef<ValueRange> operands,
455 ConversionPatternRewriter &rewriter)
const override {
461 assert(resultType &&
"Expected result to be a VectorType");
463 ArrayRef<int64_t> wgShape = resultType.getShape();
465 xegpu::DistributeLayoutAttr layout =
467 if (!layout || !layout.isForWorkgroup())
470 SmallVector<int64_t> sgShape = getSgShapeAndCount(wgShape, layout).first;
472 size_t numVariants = operands.empty() ? 0 : operands.front().size();
474 if (llvm::any_of(operands, [&](
const ValueRange &operandVec) {
475 return operandVec.size() != numVariants;
479 SmallVector<Value> newResults;
480 VectorType newResultType =
481 VectorType::get(sgShape, resultType.getElementType());
483 for (
size_t i = 0; i < numVariants; ++i) {
484 SmallVector<Value> opOperands;
485 for (
auto &operandVec : operands)
486 opOperands.push_back(operandVec[i]);
489 state.addOperands(opOperands);
490 state.addTypes(newResultType);
493 Operation *newOp = rewriter.create(state);
495 newResults.push_back(newOp->
getResult(0));
498 rewriter.replaceOpWithMultiple(op, {newResults});
529struct WgToSgConvertLayoutOp
530 :
public OpConversionPattern<xegpu::ConvertLayoutOp> {
531 using OpConversionPattern<xegpu::ConvertLayoutOp>::OpConversionPattern;
534 matchAndRewrite(xegpu::ConvertLayoutOp op, OneToNOpAdaptor adaptor,
535 ConversionPatternRewriter &rewriter)
const override {
536 Location loc = op.getLoc();
537 auto inputLayout = op.getEffectiveInputLayout();
538 auto targetLayout = op.getTargetLayout();
540 if (!inputLayout || !targetLayout || !inputLayout.isForWorkgroup() ||
541 !targetLayout.isForWorkgroup())
542 return rewriter.notifyMatchFailure(
543 op,
"Input and target layouts must have subgroup layout");
545 Type resultType = op.getResult().getType();
547 rewriter.replaceOp(op, op.getSource());
548 assert(!inputLayout.dropSgLayoutAndData() &&
549 !targetLayout.dropSgLayoutAndData() &&
550 "unexpected layout attributes for scalar type");
554 ArrayRef<int64_t> wgShape = cast<VectorType>(resultType).getShape();
555 SmallVector<int64_t> inputSgLayout =
556 inputLayout.getEffectiveSgLayoutAsInt();
557 SmallVector<int64_t> inputSgData = inputLayout.getEffectiveSgDataAsInt();
558 SmallVector<int64_t> targetSgLayout =
559 targetLayout.getEffectiveSgLayoutAsInt();
560 SmallVector<int64_t> targetSgData = targetLayout.getEffectiveSgDataAsInt();
563 SmallVector<int64_t> wgShapeVec(wgShape.begin(), wgShape.end());
564 if (inputLayout.isCompatibleWith(targetLayout, wgShapeVec,
565 xegpu::LayoutKind::Subgroup)) {
566 inputLayout = inputLayout.dropSgLayoutAndData();
567 targetLayout = targetLayout.dropSgLayoutAndData();
569 SmallVector<Value> newOps(adaptor.getSource());
570 if (inputLayout && targetLayout) {
571 for (
auto [i, src] : llvm::enumerate(adaptor.getSource())) {
572 auto newOp = xegpu::ConvertLayoutOp::create(
573 rewriter, loc, src.
getType(), src, inputLayout, targetLayout);
577 rewriter.replaceOpWithMultiple(op, {newOps});
582 Type elemTy = cast<VectorType>(op.getSource().getType()).getElementType();
584 SmallVector<int64_t> slmShape = llvm::to_vector(wgShape);
588 auto bytesPerElement = bitWidth / 8;
592 auto slmTy = MemRefType::get({slmSize}, rewriter.getI8Type(), {}, 3);
593 auto slm = memref::AllocaOp::create(rewriter, loc, slmTy);
595 auto memDescType = xegpu::MemDescType::get(rewriter.getContext(), slmShape,
598 xegpu::CreateMemDescOp::create(rewriter, loc, memDescType, slm);
600 auto sgId = gpu::SubgroupIdOp::create(rewriter, loc,
601 rewriter.getIndexType(),
nullptr);
604 auto storeCoords = inputLayout.computeDistributedCoords(
605 rewriter, loc, sgId.getResult(), wgShape);
610 for (
auto [src, coords] : llvm::zip(adaptor.getSource(), *storeCoords)) {
611 SmallVector<OpFoldResult> storeMatrixOffsets;
612 for (Value coord : coords) {
613 storeMatrixOffsets.push_back(coord);
615 xegpu::StoreMatrixOp::create(rewriter, loc, src, memDesc.getResult(),
616 storeMatrixOffsets,
nullptr );
619 gpu::BarrierOp::create(rewriter, loc);
622 auto loadCoords = targetLayout.computeDistributedCoords(
623 rewriter, loc, sgId.getResult(), wgShape);
627 VectorType loadType = VectorType::get(targetSgData, elemTy);
630 SmallVector<Value> finalResults;
631 for (
auto coords : *loadCoords) {
632 SmallVector<OpFoldResult> loadMatrixOffsets;
633 for (Value coord : coords) {
634 loadMatrixOffsets.push_back(coord);
636 auto loadOp = xegpu::LoadMatrixOp::create(
637 rewriter, loc, loadType, memDesc.getResult(), loadMatrixOffsets,
638 targetLayout.dropSgLayoutAndData());
640 finalResults.push_back(loadOp.getResult());
643 rewriter.replaceOpWithMultiple(op, {finalResults});
649struct WgToSgArithConstantOp :
public OpConversionPattern<arith::ConstantOp> {
650 using OpConversionPattern<arith::ConstantOp>::OpConversionPattern;
653 matchAndRewrite(arith::ConstantOp op, OneToNOpAdaptor adaptor,
654 ConversionPatternRewriter &rewriter)
const override {
655 auto vecAttr = dyn_cast<DenseElementsAttr>(op.getValue());
656 auto vecType = dyn_cast<VectorType>(op.getType());
657 if (!vecAttr || !vecType)
660 xegpu::DistributeLayoutAttr layout =
662 if (!layout || !layout.isForWorkgroup())
665 ArrayRef<int64_t> wgShape = vecType.getShape();
666 SmallVector<int64_t> sgShape;
668 std::tie(sgShape, count) = getSgShapeAndCount(wgShape, layout);
670 auto newType = VectorType::get(sgShape, vecType.getElementType());
671 Location loc = op.getLoc();
672 auto eltType = vecType.getElementType();
674 if (vecAttr.isSplat()) {
676 Attribute singleVal = vecAttr.getSplatValue<Attribute>();
678 SmallVector<Value> newConstOps;
679 for (
int i = 0; i < count; ++i) {
680 auto cstOp = arith::ConstantOp::create(rewriter, loc, newType, sgAttr);
681 newConstOps.push_back(cstOp);
683 rewriter.replaceOpWithMultiple(op, {newConstOps});
685 }
else if (sgShape == wgShape) {
688 arith::ConstantOp::create(rewriter, op.getLoc(), vecType, vecAttr);
689 rewriter.replaceOp(op, newConstOp);
695 if (!eltType.isIndex())
696 return rewriter.notifyMatchFailure(
697 op,
"Unsupported element type for non-splat constant op.");
699 if (wgShape.size() > 2)
700 return rewriter.notifyMatchFailure(
701 op,
"Only 1D & 2D vector constant supported");
703 SmallVector<Attribute> values(vecAttr.getValues<Attribute>());
704 int64_t rowStride = 0, colStride = 0;
705 int64_t rows = wgShape.size() == 1 ? 1 : wgShape[0];
706 int64_t cols = wgShape.size() == 1 ? wgShape[0] : wgShape[1];
710 colStride = cast<IntegerAttr>(values[1]).getInt() -
711 cast<IntegerAttr>(values[0]).getInt();
714 rowStride = cast<IntegerAttr>(values[cols]).getInt() -
715 cast<IntegerAttr>(values[0]).getInt();
718 for (int64_t r = 0; r < rows; ++r) {
719 for (int64_t c = 0; c < cols; ++c) {
720 int64_t idx = r * cols + c;
722 if (c > 0 && cols > 1) {
723 int64_t prevIdx = r * cols + (c - 1);
724 int64_t diff = cast<IntegerAttr>(values[idx]).getInt() -
725 cast<IntegerAttr>(values[prevIdx]).getInt();
726 if (diff != colStride)
727 return rewriter.notifyMatchFailure(
728 op,
"Non-constant column stride in constant op.");
731 if (r > 0 && rows > 1) {
732 int64_t prevIdx = (r - 1) * cols + c;
733 int64_t diff = cast<IntegerAttr>(values[idx]).getInt() -
734 cast<IntegerAttr>(values[prevIdx]).getInt();
735 if (diff != rowStride)
736 return rewriter.notifyMatchFailure(
737 op,
"Non-constant row stride in constant op.");
745 SmallVector<Attribute> baseTileValues;
746 int baseTileCols = sgShape[sgShape.size() - 1];
747 int64_t baseTileRows = sgShape.size() == 1 ? 1 : sgShape[0];
748 for (int64_t r = 0; r < baseTileRows; ++r) {
749 for (int64_t c = 0; c < baseTileCols; ++c) {
750 baseTileValues.push_back(values[r * cols + c]);
756 auto baseConstVec = arith::ConstantOp::create(rewriter, loc, tileAttr);
760 gpu::SubgroupIdOp::create(rewriter, loc,
nullptr);
762 layout.computeDistributedCoords(rewriter, loc, sgId, wgShape);
766 SmallVector<Value, 2> strideConsts;
767 strideConsts.push_back(
771 strideConsts.begin(),
774 SmallVector<Value> newConstOps;
775 for (
auto offsets : *sgOffsets) {
778 for (
size_t i = 0; i < strideConsts.size(); ++i) {
780 arith::MulIOp::create(rewriter, loc, rewriter.getIndexType(),
781 offsets[i], strideConsts[i]);
782 mulOffset = arith::AddIOp::create(
783 rewriter, loc, rewriter.getIndexType(), mulOffset,
mul);
786 auto bcastOffset = vector::BroadcastOp::create(
787 rewriter, loc, baseConstVec.getType(), mulOffset);
789 arith::AddIOp::create(rewriter, loc, baseConstVec, bcastOffset);
790 newConstOps.push_back(finalConst);
792 rewriter.replaceOpWithMultiple(op, {newConstOps});
800struct WgToSgLoadGatherOp :
public OpConversionPattern<xegpu::LoadGatherOp> {
801 using OpConversionPattern<xegpu::LoadGatherOp>::OpConversionPattern;
803 matchAndRewrite(xegpu::LoadGatherOp op, OneToNOpAdaptor adaptor,
804 ConversionPatternRewriter &rewriter)
const override {
806 Location loc = op.getLoc();
807 VectorType resultType = dyn_cast<VectorType>(op.getResult().getType());
810 ArrayRef<int64_t> wgShape = resultType.getShape();
812 xegpu::DistributeLayoutAttr layout = op.getLayoutAttr();
814 if (!layout || !layout.isForWorkgroup())
817 SmallVector<int64_t> sgShape = getSgShapeAndCount(wgShape, layout).first;
820 auto offsetsVecType =
821 dyn_cast<VectorType>(adaptor.getOffsets().front().getType());
823 dyn_cast<VectorType>(adaptor.getMask().front().getType());
824 if (!offsetsVecType || !maskVecType ||
825 offsetsVecType.getShape() != maskVecType.getShape()) {
826 return rewriter.notifyMatchFailure(op,
827 "offsets have not been distributed");
830 SmallVector<Value> newLoadOps;
831 VectorType newTy = VectorType::get(sgShape, resultType.getElementType());
832 for (
auto [offsets, mask] :
833 llvm::zip(adaptor.getOffsets(), adaptor.getMask())) {
834 auto newLayout = layout.dropSgLayoutAndData();
835 auto newLoadOp = xegpu::LoadGatherOp::create(
836 rewriter, loc, newTy, op.getSource(), offsets, mask,
837 op.getL1HintAttr(), op.getL2HintAttr(), op.getL3HintAttr(), newLayout,
839 newLoadOps.push_back(newLoadOp);
841 rewriter.replaceOpWithMultiple(op, {newLoadOps});
848struct WgToSgStoreScatterOp
849 :
public OpConversionPattern<xegpu::StoreScatterOp> {
850 using OpConversionPattern<xegpu::StoreScatterOp>::OpConversionPattern;
852 matchAndRewrite(xegpu::StoreScatterOp op, OneToNOpAdaptor adaptor,
853 ConversionPatternRewriter &rewriter)
const override {
855 Location loc = op.getLoc();
856 VectorType valueType = dyn_cast<VectorType>(op.getValue().getType());
860 xegpu::DistributeLayoutAttr layout = op.getLayoutAttr();
862 if (!layout || !layout.isForWorkgroup())
866 auto offsetsVecType =
867 dyn_cast<VectorType>(adaptor.getOffsets().front().getType());
869 dyn_cast<VectorType>(adaptor.getMask().front().getType());
870 if (!offsetsVecType || !maskVecType ||
871 offsetsVecType.getShape() != maskVecType.getShape()) {
872 return rewriter.notifyMatchFailure(op,
873 "offsets have not been distributed");
876 for (
auto [val, offs, mask] : llvm::zip(
877 adaptor.getValue(), adaptor.getOffsets(), adaptor.getMask())) {
878 xegpu::StoreScatterOp::create(
879 rewriter, loc, val, op.getDest(), offs, mask, op.getL1HintAttr(),
880 op.getL2HintAttr(), op.getL3HintAttr(), layout.dropSgLayoutAndData(),
883 rewriter.eraseOp(op);
888struct WgToSgLoadMatrixOp :
public OpConversionPattern<xegpu::LoadMatrixOp> {
889 using OpConversionPattern<xegpu::LoadMatrixOp>::OpConversionPattern;
891 matchAndRewrite(xegpu::LoadMatrixOp op, OneToNOpAdaptor adaptor,
892 ConversionPatternRewriter &rewriter)
const override {
894 SmallVector<SmallVector<OpFoldResult>> offsetsList;
895 if (
failed(genOffsetsList(rewriter, op, offsetsList)))
898 ArrayRef<int64_t> wgShape = op.getDataShape();
899 VectorType valueTy = llvm::dyn_cast<VectorType>(op.getRes().getType());
900 assert(valueTy &&
"the value type must be vector type!");
901 Type elemTy = valueTy.getElementType();
903 xegpu::DistributeLayoutAttr layout = op.getLayoutAttr();
904 SmallVector<int64_t> sgShape = getSgShapeAndCount(wgShape, layout).first;
905 VectorType newResTy = VectorType::get(sgShape, elemTy);
906 SmallVector<Value> newOps;
907 for (
auto offsets : offsetsList) {
908 auto newOp = xegpu::LoadMatrixOp::create(rewriter, op.getLoc(), newResTy,
909 op.getMemDesc(), offsets,
910 layout.dropSgLayoutAndData());
911 newOps.push_back(newOp);
913 rewriter.replaceOpWithMultiple(op, {newOps});
919struct WgToSgStoreMatrixOp :
public OpConversionPattern<xegpu::StoreMatrixOp> {
920 using OpConversionPattern<xegpu::StoreMatrixOp>::OpConversionPattern;
922 matchAndRewrite(xegpu::StoreMatrixOp op, OneToNOpAdaptor adaptor,
923 ConversionPatternRewriter &rewriter)
const override {
925 SmallVector<SmallVector<OpFoldResult>> offsetsList;
926 if (
failed(genOffsetsList(rewriter, op, offsetsList)))
929 xegpu::DistributeLayoutAttr layout = op.getLayoutAttr();
930 for (
auto [v, offsets] : llvm::zip(adaptor.getData(), offsetsList))
931 xegpu::StoreMatrixOp::create(rewriter, op.getLoc(), v, op.getMemDesc(),
932 offsets, layout.dropSgLayoutAndData());
933 rewriter.eraseOp(op);
939struct WgToSgVectorStepOp :
public OpConversionPattern<vector::StepOp> {
940 using OpConversionPattern<vector::StepOp>::OpConversionPattern;
942 matchAndRewrite(vector::StepOp op, OneToNOpAdaptor adaptor,
943 ConversionPatternRewriter &rewriter)
const override {
944 xegpu::DistributeLayoutAttr layout =
946 if (!layout || !layout.isForWorkgroup())
949 Location loc = op.getLoc();
950 VectorType type = op.getResult().getType();
951 auto wgShape = type.getShape();
952 std::optional<SmallVector<int64_t>> sgShape =
953 getSgShapeAndCount(wgShape, layout).first;
958 gpu::SubgroupIdOp::create(rewriter, loc,
nullptr);
960 layout.computeDistributedCoords(rewriter, loc, sgId, wgShape);
964 VectorType newTy = type.cloneWith(*sgShape, type.getElementType());
965 auto steps = vector::StepOp::create(rewriter, loc, newTy);
966 SmallVector<Value> newOps;
967 for (
auto offsets : *sgOffsets) {
970 vector::BroadcastOp::create(rewriter, loc, newTy, offsets[0]);
972 arith::AddIOp::create(rewriter, loc, steps, bcastOffset);
973 newOps.push_back(finalSteps);
976 rewriter.replaceOpWithMultiple(op, {newOps});
982struct WgToSgVectorShapeCastOp
983 :
public OpConversionPattern<vector::ShapeCastOp> {
984 using OpConversionPattern<vector::ShapeCastOp>::OpConversionPattern;
987 matchAndRewrite(vector::ShapeCastOp op, OneToNOpAdaptor adaptor,
988 ConversionPatternRewriter &rewriter)
const override {
990 VectorType resultType = dyn_cast<VectorType>(op.getResult().getType());
994 ArrayRef<int64_t> wgShape = resultType.getShape();
995 xegpu::DistributeLayoutAttr layout =
997 if (!layout || !layout.isForWorkgroup())
1002 auto srcType = dyn_cast<VectorType>(op.getSource().getType());
1006 ArrayRef<int64_t> srcShape = srcType.getShape();
1008 xegpu::DistributeLayoutAttr layoutToDistribute = layout;
1009 SmallVector<int64_t> expandedUnitDims;
1011 xegpu::DistributeLayoutAttr sourceLayout =
1014 if (!sourceLayout.isSliceOf(layout))
1015 return rewriter.notifyMatchFailure(
1016 op,
"The ShapeCast op only expands dimensions, the input layout "
1017 "must be a slice of the result layout.");
1019 assert(layoutToDistribute.isEqualTo(
1020 layoutToDistribute.setUnitDimData(expandedUnitDims)) &&
1021 "The sg_data for unit dimensions should be set as 1");
1024 SmallVector<int64_t> sgShape =
1025 getSgShapeAndCount(wgShape, layoutToDistribute).first;
1026 VectorType newResultType =
1027 VectorType::get(sgShape, resultType.getElementType());
1029 SmallVector<Value> newShapeCastOps;
1030 for (
auto src : adaptor.getSource()) {
1031 auto newShapeCast = vector::ShapeCastOp::create(rewriter, op.getLoc(),
1032 newResultType, src);
1033 newShapeCastOps.push_back(newShapeCast.getResult());
1036 rewriter.replaceOpWithMultiple(op, {newShapeCastOps});
1073struct WgToSgMultiDimReductionOp
1074 :
public OpConversionPattern<vector::MultiDimReductionOp> {
1075 using OpConversionPattern<vector::MultiDimReductionOp>::OpConversionPattern;
1078 matchAndRewrite(vector::MultiDimReductionOp op, OneToNOpAdaptor adaptor,
1079 ConversionPatternRewriter &rewriter)
const override {
1080 Location loc = op.getLoc();
1082 VectorType srcType = op.getSourceVectorType();
1083 Type resultTy = op.getResult().getType();
1084 VectorType dstVecType = dyn_cast<VectorType>(resultTy);
1085 bool isScalarResult = !dstVecType;
1087 auto originalSrcShape = srcType.getShape();
1088 Type elemTy = srcType.getElementType();
1090 xegpu::DistributeLayoutAttr layout =
1092 if (!layout || !layout.isForWorkgroup())
1095 auto reductionDims = llvm::to_vector(op.getReductionDims());
1098 SmallVector<int64_t> sgLayout;
1099 SmallVector<int64_t> sgData;
1100 xegpu::DistributeLayoutAttr parentLayout;
1101 if (
auto sliceAttr = dyn_cast<xegpu::SliceAttr>(layout)) {
1102 parentLayout = sliceAttr.getParent();
1103 sgLayout = parentLayout.getEffectiveSgLayoutAsInt();
1104 sgData = parentLayout.getEffectiveSgDataAsInt();
1106 return rewriter.notifyMatchFailure(
1107 op,
"Reduction should have SliceAttr layout");
1110 SmallVector<Value> localReductions;
1111 auto sgSrcs = adaptor.getSource();
1112 auto sgSrcType = dyn_cast<VectorType>(sgSrcs.front().getType());
1113 SmallVector<int64_t> sgSrcShape(sgSrcType.getShape().begin(),
1114 sgSrcType.getShape().end());
1121 auto originalDstShape = dstVecType.getShape();
1122 SmallVector<int64_t> sgDstShape =
1123 getSgShapeAndCount(originalDstShape, layout).first;
1124 sgDstType = VectorType::get(sgDstShape, elemTy);
1129 for (
auto sgSrc : sgSrcs) {
1132 rewriter, loc, sgDstType, op.getKind());
1134 auto localReduce = vector::MultiDimReductionOp::create(
1135 rewriter, loc, sgDstType, op.getKind(), sgSrc, neutralLocalAcc,
1137 localReductions.push_back(localReduce.getResult());
1141 SmallVector<int64_t> crossSgReductionDims;
1142 for (int64_t reductionDim : reductionDims) {
1143 bool needsCrossSubgroupReduction =
1144 (sgLayout[reductionDim] > 1) &&
1145 (sgData[reductionDim] < originalSrcShape[reductionDim]);
1147 if (needsCrossSubgroupReduction) {
1148 crossSgReductionDims.push_back(reductionDim);
1153 if (crossSgReductionDims.empty()) {
1154 SmallVector<Value> results;
1155 for (
auto localResult : localReductions) {
1157 rewriter, loc, op.getKind(), localResult, adaptor.getAcc()[0]);
1158 results.push_back(finalResult);
1160 rewriter.replaceOpWithMultiple(op, {results});
1165 auto slmStoreDataShape = sgSrcShape;
1166 for (int64_t dim : reductionDims)
1167 slmStoreDataShape[dim] = 1;
1168 VectorType slmStoreDataType = VectorType::get(slmStoreDataShape, elemTy);
1169 SmallVector<Value> slmStoreData;
1170 for (
auto localResult : localReductions) {
1171 if (isScalarResult) {
1173 slmStoreData.push_back(vector::BroadcastOp::create(
1174 rewriter, loc, slmStoreDataType, localResult));
1176 slmStoreData.push_back(vector::ShapeCastOp::create(
1177 rewriter, loc, slmStoreDataType, localResult));
1181 SmallVector<int64_t> slmShape(originalSrcShape.begin(),
1182 originalSrcShape.end());
1183 SmallVector<int> slmSgData(sgData.begin(), sgData.end());
1184 SmallVector<int> slmSgLayout(sgLayout.begin(), sgLayout.end());
1185 for (
int dim : reductionDims) {
1186 slmShape[dim] = sgLayout[dim];
1189 xegpu::LayoutAttr slmStoreLayout =
1190 xegpu::LayoutAttr::get(rewriter.getContext(), slmSgLayout, slmSgData);
1194 auto bytesPerElement = bitWidth / 8;
1196 auto slmTy = MemRefType::get({slmSize}, rewriter.getI8Type(), {}, 3);
1197 auto slm = memref::AllocaOp::create(rewriter, loc, slmTy);
1199 auto memDescType = xegpu::MemDescType::get(rewriter.getContext(), slmShape,
1202 xegpu::CreateMemDescOp::create(rewriter, loc, memDescType, slm);
1205 auto sgId = gpu::SubgroupIdOp::create(rewriter, loc,
1206 rewriter.getIndexType(),
nullptr);
1208 auto slmStoreCoords =
1209 slmStoreLayout.computeDistributedCoords(rewriter, loc, sgId, slmShape);
1210 if (
failed(slmStoreCoords))
1212 for (
auto [data, coord] : llvm::zip(slmStoreData, *slmStoreCoords)) {
1213 SmallVector<OpFoldResult> coordOfr(coord.begin(), coord.end());
1214 xegpu::StoreMatrixOp::create(rewriter, loc, data, memDesc.getResult(),
1219 gpu::BarrierOp::create(rewriter, loc);
1222 SmallVector<int64_t> slmLoadDataShape(sgSrcShape.begin(), sgSrcShape.end());
1223 for (int64_t dim : reductionDims) {
1224 slmLoadDataShape[dim] = slmShape[dim];
1225 slmSgData[dim] = slmShape[dim];
1227 xegpu::LayoutAttr slmLoadLayout =
1228 xegpu::LayoutAttr::get(rewriter.getContext(), slmSgLayout, slmSgData);
1229 auto slmLoadCoords =
1230 slmLoadLayout.computeDistributedCoords(rewriter, loc, sgId, slmShape);
1231 if (
failed(slmLoadCoords))
1234 VectorType slmLoadType = VectorType::get(slmLoadDataShape, elemTy);
1235 SmallVector<Value> slmLoadData;
1236 for (
auto coord : *slmLoadCoords) {
1237 SmallVector<OpFoldResult> coordOfr(coord.begin(), coord.end());
1238 slmLoadData.push_back(xegpu::LoadMatrixOp::create(
1239 rewriter, loc, slmLoadType, memDesc.getResult(), coordOfr,
1246 rewriter, loc, sgDstType, op.getKind());
1248 SmallVector<Value> finalResults;
1249 for (
size_t i = 0; i < slmLoadData.size(); ++i) {
1250 auto loaded = slmLoadData[i];
1251 auto finalReduce = vector::MultiDimReductionOp::create(
1252 rewriter, loc, sgDstType, op.getKind(), loaded, neutralFinalAcc,
1255 rewriter, loc, op.getKind(), finalReduce.getResult(),
1256 adaptor.getAcc()[i]));
1258 rewriter.replaceOpWithMultiple(op, {finalResults});
1264struct WgToSgVectorTransposeOp
1265 :
public OpConversionPattern<vector::TransposeOp> {
1266 using OpConversionPattern<vector::TransposeOp>::OpConversionPattern;
1269 matchAndRewrite(vector::TransposeOp op, OneToNOpAdaptor adaptor,
1270 ConversionPatternRewriter &rewriter)
const override {
1271 VectorType resultType = op.getResultVectorType();
1273 ArrayRef<int64_t> wgShape = resultType.getShape();
1274 xegpu::DistributeLayoutAttr layout =
1276 if (!layout || !layout.isForWorkgroup())
1278 xegpu::DistributeLayoutAttr sourceLayout =
1280 if (!sourceLayout || !sourceLayout.isForWorkgroup())
1283 SmallVector<int64_t> sourceSgLayout =
1284 sourceLayout.getEffectiveSgLayoutAsInt();
1285 SmallVector<int64_t> resultSgLayout = layout.getEffectiveSgLayoutAsInt();
1287 ArrayRef<int64_t> permutation = op.getPermutation();
1288 size_t permutationSize = permutation.size();
1289 if (sourceSgLayout.size() != permutationSize ||
1290 resultSgLayout.size() != permutationSize) {
1291 return rewriter.notifyMatchFailure(
1292 op,
"Layouts and permutation must have the same rank");
1297 if (!layout.isTransposeOf(sourceLayout, permutation,
1298 xegpu::LayoutKind::Subgroup))
1299 return rewriter.notifyMatchFailure(
1300 op,
"Result layout is not a valid transpose of source layout "
1301 "according to permutation");
1303 SmallVector<int64_t> sgShape = getSgShapeAndCount(wgShape, layout).first;
1304 VectorType newResultType =
1305 VectorType::get(sgShape, resultType.getElementType());
1307 SmallVector<Value> newTransposeOps;
1308 for (
auto src : adaptor.getVector()) {
1309 auto newTranspose = vector::TransposeOp::create(
1310 rewriter, op.getLoc(), newResultType, src, permutation);
1311 newTransposeOps.push_back(newTranspose.getResult());
1313 rewriter.replaceOpWithMultiple(op, {newTransposeOps});
1319template <
typename MaskOpType>
1320struct WgToSgVectorMaskOp :
public OpConversionPattern<MaskOpType> {
1321 using OpConversionPattern<MaskOpType>::OpConversionPattern;
1323 LogicalResult matchAndRewrite(
1325 typename OpConversionPattern<MaskOpType>::OneToNOpAdaptor adaptor,
1326 ConversionPatternRewriter &rewriter)
const override {
1327 xegpu::DistributeLayoutAttr layout =
1329 if (!layout || !layout.isForWorkgroup())
1332 Location loc = op.getLoc();
1333 VectorType type = op.getResult().getType();
1334 auto wgShape = type.getShape();
1336 SmallVector<Value> wgMaskDimSizes;
1337 if constexpr (std::is_same_v<MaskOpType, vector::ConstantMaskOp>) {
1338 for (int64_t maskSize : op.getMaskDimSizes()) {
1339 wgMaskDimSizes.push_back(
1342 }
else if constexpr (std::is_same_v<MaskOpType, vector::CreateMaskOp>) {
1343 wgMaskDimSizes = llvm::to_vector(op.getOperands());
1347 gpu::SubgroupIdOp::create(rewriter, loc,
nullptr);
1349 layout.computeDistributedCoords(rewriter, loc, sgId, wgShape);
1353 SmallVector<int64_t> sgShape = getSgShapeAndCount(wgShape, layout).first;
1354 VectorType resultType = VectorType::get(sgShape, type.getElementType());
1358 SmallVector<Value> newCreateMaskOps;
1359 for (
auto offsetSet : *sgOffsets) {
1360 SmallVector<Value> maskOperands;
1362 for (
auto [i, wgMaskDimSize] : llvm::enumerate(wgMaskDimSizes)) {
1365 Value offset = offsetSet[i];
1366 Value adjustedMaskSize =
1367 arith::SubIOp::create(rewriter, loc, wgMaskDimSize, offset);
1370 arith::MaxSIOp::create(rewriter, loc, adjustedMaskSize, zero);
1372 arith::MinSIOp::create(rewriter, loc, nonNegative, dimSizeVal);
1373 maskOperands.push_back(sgMaskSize);
1376 auto newCreateMaskOp =
1377 vector::CreateMaskOp::create(rewriter, loc, resultType, maskOperands);
1378 newCreateMaskOps.push_back(newCreateMaskOp.getResult());
1381 rewriter.replaceOpWithMultiple(op, {newCreateMaskOps});
1386using WgToSgVectorConstantMaskOp = WgToSgVectorMaskOp<vector::ConstantMaskOp>;
1387using WgToSgVectorCreateMaskOp = WgToSgVectorMaskOp<vector::CreateMaskOp>;
1390struct WgToSgVectorBitCastOp :
public OpConversionPattern<vector::BitCastOp> {
1391 using OpConversionPattern<vector::BitCastOp>::OpConversionPattern;
1394 matchAndRewrite(vector::BitCastOp op, OneToNOpAdaptor adaptor,
1395 ConversionPatternRewriter &rewriter)
const override {
1396 VectorType resultType = op.getResultVectorType();
1398 ArrayRef<int64_t> wgShape = resultType.getShape();
1399 xegpu::DistributeLayoutAttr layout =
1401 if (!layout || !layout.isForWorkgroup())
1404 SmallVector<int64_t> sgShape = getSgShapeAndCount(wgShape, layout).first;
1405 VectorType newResultType =
1406 VectorType::get(sgShape, resultType.getElementType());
1408 SmallVector<Value> newBitCastOps;
1409 for (
auto src : adaptor.getSource()) {
1411 vector::BitCastOp::create(rewriter, op.getLoc(), newResultType, src);
1412 newBitCastOps.push_back(newBitCast.getResult());
1415 rewriter.replaceOpWithMultiple(op, {newBitCastOps});
1421struct WgToSgVectorInterleaveOp
1422 :
public OpConversionPattern<vector::InterleaveOp> {
1423 using OpConversionPattern<vector::InterleaveOp>::OpConversionPattern;
1426 matchAndRewrite(vector::InterleaveOp op, OneToNOpAdaptor adaptor,
1427 ConversionPatternRewriter &rewriter)
const override {
1428 VectorType resultType = op.getResultVectorType();
1430 ArrayRef<int64_t> wgShape = resultType.getShape();
1431 xegpu::DistributeLayoutAttr layout =
1433 if (!layout || !layout.isForWorkgroup())
1436 SmallVector<int64_t> sgShape = getSgShapeAndCount(wgShape, layout).first;
1437 VectorType newResultType =
1438 VectorType::get(sgShape, resultType.getElementType());
1440 SmallVector<Value> newInterleaveOps;
1443 for (
auto [
lhs,
rhs] : llvm::zip(adaptor.getLhs(), adaptor.getRhs())) {
1444 auto newInterleave = vector::InterleaveOp::create(
1445 rewriter, op.getLoc(), newResultType,
lhs,
rhs);
1446 newInterleaveOps.push_back(newInterleave.getResult());
1449 rewriter.replaceOpWithMultiple(op, {newInterleaveOps});
1455struct WgToSgVectorDeinterleaveOp
1456 :
public OpConversionPattern<vector::DeinterleaveOp> {
1457 using OpConversionPattern<vector::DeinterleaveOp>::OpConversionPattern;
1460 matchAndRewrite(vector::DeinterleaveOp op, OneToNOpAdaptor adaptor,
1461 ConversionPatternRewriter &rewriter)
const override {
1462 SmallVector<Value> newRes1Ops;
1463 SmallVector<Value> newRes2Ops;
1465 for (
auto src : adaptor.getSource()) {
1466 auto newDeinterleave =
1467 vector::DeinterleaveOp::create(rewriter, op.getLoc(), src);
1468 newRes1Ops.push_back(newDeinterleave.getRes1());
1469 newRes2Ops.push_back(newDeinterleave.getRes2());
1472 SmallVector<SmallVector<Value>> results = {newRes1Ops, newRes2Ops};
1473 rewriter.replaceOpWithMultiple(op, results);
1485 converter.addConversion([](
Type type) ->
Type {
return type; });
1488 converter.addConversion(
1489 [](xegpu::TensorDescType type,
1491 xegpu::DistributeLayoutAttr layout = type.getLayoutAttr();
1492 if (!layout || !layout.isForWorkgroup())
1493 return std::nullopt;
1495 Type elemTy = type.getElementType();
1500 std::tie(subShape, count) = getSgShapeAndCount(
shape, layout);
1502 layout = layout.dropSgLayoutAndData();
1504 auto newTy = xegpu::TensorDescType::get(
1505 type.
getContext(), subShape, elemTy, type.getEncoding(), layout);
1506 result.append(count, newTy);
1512 auto getSubShapeAndCount = [](VectorType vecTy,
1513 xegpu::DistributeLayoutAttr layout)
1515 if (!layout.isForWorkgroup())
1517 return getSgShapeAndCount(vecTy.getShape(), layout);
1522 std::move(loopArgTypes));
1526 patterns.
add<WgToSgCreateNdOp, WgToSgLoadNdOp, WgToSgStoreNdOp, WgToSgDpasOp,
1527 WgToSgDpasMxOp, WgToSgPrefetchNdOp, WgToSgElementwiseOp,
1528 WgToSgVectorBroadcastOp, WgToSgConvertLayoutOp,
1529 WgToSgArithConstantOp, WgToSgLoadGatherOp, WgToSgStoreScatterOp,
1530 WgToSgLoadMatrixOp, WgToSgStoreMatrixOp, WgToSgVectorStepOp,
1531 WgToSgVectorShapeCastOp, WgToSgMultiDimReductionOp,
1532 WgToSgVectorTransposeOp, WgToSgVectorConstantMaskOp,
1533 WgToSgVectorCreateMaskOp, WgToSgVectorBitCastOp,
1534 WgToSgVectorInterleaveOp, WgToSgVectorDeinterleaveOp>(
1541struct XeGPUWgToSgDistributePass
1542 :
public xegpu::impl::XeGPUWgToSgDistributeBase<XeGPUWgToSgDistributePass> {
1543 void runOnOperation()
override;
1547void XeGPUWgToSgDistributePass::runOnOperation() {
1549 Operation *op = getOperation();
1551 signalPassFailure();
1556 llvm::SmallSetVector<UnrealizedConversionCastOp, 8> existingCasts;
1557 getOperation()->walk(
1558 [&](UnrealizedConversionCastOp castOp) { existingCasts.insert(castOp); });
1565 RewritePatternSet patterns(ctx);
1566 ConversionTarget
target(*ctx);
1567 TypeConverter converter;
1570 auto materializeCast = [](OpBuilder &builder, Type type,
ValueRange inputs,
1571 Location loc) -> Value {
1572 return UnrealizedConversionCastOp::create(builder, loc, type, inputs)
1575 converter.addSourceMaterialization(materializeCast);
1576 converter.addTargetMaterialization(materializeCast);
1580 auto getTensorDescType = [](Operation *op) -> xegpu::TensorDescType {
1581 if (
auto createOp = dyn_cast<xegpu::CreateNdDescOp>(op))
1582 return createOp.getType();
1583 if (
auto loadOp = dyn_cast<xegpu::LoadNdOp>(op))
1584 return loadOp.getTensorDescType();
1585 if (
auto storeOp = dyn_cast<xegpu::StoreNdOp>(op))
1586 return storeOp.getTensorDescType();
1587 if (
auto prefetchOp = dyn_cast<xegpu::PrefetchNdOp>(op))
1588 return prefetchOp.getTensorDescType();
1589 return xegpu::TensorDescType();
1592 auto isLegal = [&](xegpu::DistributeLayoutAttr layout) ->
bool {
1593 return !layout || !layout.isForWorkgroup();
1596 target.addDynamicallyLegalOp<xegpu::CreateNdDescOp, xegpu::LoadNdOp,
1597 xegpu::StoreNdOp, xegpu::PrefetchNdOp>(
1598 [=](Operation *op) ->
bool {
1599 auto tdescTy = getTensorDescType(op);
1600 auto layout = dyn_cast_if_present<xegpu::DistributeLayoutAttr>(
1601 tdescTy.getLayout());
1602 return isLegal(layout);
1605 target.addDynamicallyLegalOp<xegpu::DpasOp>([=](xegpu::DpasOp op) ->
bool {
1606 auto layout = op.getLayoutCdAttr();
1607 return isLegal(layout);
1610 target.addDynamicallyLegalOp<xegpu::DpasMxOp>(
1611 [=](xegpu::DpasMxOp op) ->
bool {
1612 auto layout = op.getLayoutCdAttr();
1613 return isLegal(layout);
1616 target.addDynamicallyLegalOp<xegpu::LoadMatrixOp>(
1617 [=](xegpu::LoadMatrixOp op) ->
bool {
1618 return isLegal(op.getLayoutAttr());
1621 target.addDynamicallyLegalOp<xegpu::StoreMatrixOp>(
1622 [=](xegpu::StoreMatrixOp op) ->
bool {
1623 return isLegal(op.getLayoutAttr());
1626 target.addDynamicallyLegalOp<arith::ConstantOp>(
1627 [=](arith::ConstantOp op) ->
bool {
1628 auto vecType = dyn_cast<VectorType>(op.getType());
1634 return isLegal(layout);
1637 target.addDynamicallyLegalOp<
1638 vector::ShapeCastOp, vector::StepOp, vector::TransposeOp,
1639 vector::BroadcastOp, vector::MultiDimReductionOp, vector::ConstantMaskOp,
1640 vector::CreateMaskOp, vector::BitCastOp, vector::InterleaveOp,
1641 vector::DeinterleaveOp>([=](Operation *op) ->
bool {
1645 return isLegal(layout);
1648 target.addDynamicallyLegalOp<xegpu::LoadGatherOp>(
1649 [=](xegpu::LoadGatherOp op) ->
bool {
1650 auto layout = op.getLayoutAttr();
1651 return isLegal(layout);
1654 target.addDynamicallyLegalOp<xegpu::StoreScatterOp>(
1655 [=](xegpu::StoreScatterOp op) ->
bool {
1656 auto layout = op.getLayoutAttr();
1657 return isLegal(layout);
1660 target.addDynamicallyLegalOp<xegpu::ConvertLayoutOp>(
1661 [=](xegpu::ConvertLayoutOp op) ->
bool {
1662 return isLegal(op.getEffectiveInputLayout()) &&
1663 isLegal(op.getTargetLayout());
1666 target.addDynamicallyLegalDialect<math::MathDialect, arith::ArithDialect>(
1667 [=](Operation *op) -> std::optional<bool> {
1672 VectorType resultType =
1680 VectorType operandType = dyn_cast<VectorType>(operand.getType());
1681 if (!operandType || operandType.getShape() != resultType.getShape()) {
1686 xegpu::DistributeLayoutAttr layout =
1688 return isLegal(layout);
1691 target.addLegalOp<UnrealizedConversionCastOp>();
1693 target.markUnknownOpDynamicallyLegal([](Operation *) {
return true; });
1699 applyPartialConversion(getOperation(),
target, std::move(patterns))))
1700 return signalPassFailure();
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 * getContext() const
Return the context this location is uniqued in.
Operation is the basic unit of execution within MLIR.
Attribute getDiscardableAttr(StringRef name)
Access a discardable attribute by name, returns a null Attribute if the discardable attribute does no...
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.
OpTy getParentOfType()
Return the closest surrounding parent operation that is of type 'OpTy'.
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.
operand_range getOperands()
Returns an iterator on the underlying Value's.
unsigned getNumResults()
Return the number of results held by this operation.
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.
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
MLIRContext * getContext() const
Return the MLIRContext in which this type was uniqued.
bool isIntOrFloat() const
Return true if this is an integer (of any signedness) or a float type.
unsigned getIntOrFloatBitWidth() const
Return the bit width of an integer or a float type, assert failure on other types.
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
Type getType() const
Return the type of this value.
static ConstantIndexOp create(OpBuilder &builder, Location location, int64_t value)
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.
void removeTemporaryLayoutAttrs(Operation *op)
Removes the temporary layout attributes for each OpOperand and OpResult of the given operation.
void populateXeGPUWgToSgDistributeTypeConversions(TypeConverter &converter, Operation *topLevelOp)
Define the type conversions needed for XeGPU workgroup to subgroup distribution.
Value createReductionNeutralValue(OpBuilder &builder, Location loc, Type type, vector::CombiningKind kind)
Creates a constant filled with the neutral (identity) value for the given reduction kind.
bool matchUnitDimExpansion(ArrayRef< int64_t > src, ArrayRef< int64_t > dst, SmallVector< int64_t > &expandedUnitDims)
bool recoverTemporaryLayouts(Operation *rootOp)
Attach layout attributes to all vector-type operands of operations within the given operation's neste...
DenseMap< Value, SmallVector< Type > > precomputeLoopBlockArgTypes(Operation *topLevelOp, SubShapeAndCountFn getSubShapeAndCount)
Pre-computes distributed VectorType mappings for every value carried through an SCF loop under topLev...
void populateXeGPUWgToSgDistributePatterns(RewritePatternSet &patterns)
Appends patterns for XeGPU workgroup to subgroup distribution into patterns.
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 removeLayoutAttrs(Operation *op)
Removes the DistributeLayoutAttr for each OpOperand and OpResult of the given operation if they exist...
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.
Include the generated interface declarations.
int64_t computeProduct(ArrayRef< int64_t > basis)
Self-explicit.
OpFoldResult getAsOpFoldResult(Value val)
Given a value, try to extract a constant Attribute.