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;
832 rewriter.getI64IntegerAttr(op.getChunkSize().value_or(1));
833 VectorType newTy = VectorType::get(sgShape, resultType.getElementType());
834 for (
auto [offsets, mask] :
835 llvm::zip(adaptor.getOffsets(), adaptor.getMask())) {
836 auto newLayout = layout.dropSgLayoutAndData();
837 auto newLoadOp = xegpu::LoadGatherOp::create(
838 rewriter, loc, newTy, op.getSource(), offsets, mask, chunkSizeAttr,
839 op.getL1HintAttr(), op.getL2HintAttr(), op.getL3HintAttr(), newLayout,
841 newLoadOps.push_back(newLoadOp);
843 rewriter.replaceOpWithMultiple(op, {newLoadOps});
850struct WgToSgStoreScatterOp
851 :
public OpConversionPattern<xegpu::StoreScatterOp> {
852 using OpConversionPattern<xegpu::StoreScatterOp>::OpConversionPattern;
854 matchAndRewrite(xegpu::StoreScatterOp op, OneToNOpAdaptor adaptor,
855 ConversionPatternRewriter &rewriter)
const override {
857 Location loc = op.getLoc();
858 VectorType valueType = dyn_cast<VectorType>(op.getValue().getType());
862 xegpu::DistributeLayoutAttr layout = op.getLayoutAttr();
864 if (!layout || !layout.isForWorkgroup())
868 auto offsetsVecType =
869 dyn_cast<VectorType>(adaptor.getOffsets().front().getType());
871 dyn_cast<VectorType>(adaptor.getMask().front().getType());
872 if (!offsetsVecType || !maskVecType ||
873 offsetsVecType.getShape() != maskVecType.getShape()) {
874 return rewriter.notifyMatchFailure(op,
875 "offsets have not been distributed");
878 auto chunkSizeOpt = op.getChunkSize();
879 int64_t chunkSize = chunkSizeOpt ?
static_cast<int64_t
>(*chunkSizeOpt) : 1;
880 auto chunkSizeAttr = rewriter.getI64IntegerAttr(chunkSize);
881 for (
auto [val, offs, mask] : llvm::zip(
882 adaptor.getValue(), adaptor.getOffsets(), adaptor.getMask())) {
883 xegpu::StoreScatterOp::create(rewriter, loc, val, op.getDest(), offs,
884 mask, chunkSizeAttr, op.getL1HintAttr(),
885 op.getL2HintAttr(), op.getL3HintAttr(),
886 layout.dropSgLayoutAndData(),
889 rewriter.eraseOp(op);
894struct WgToSgLoadMatrixOp :
public OpConversionPattern<xegpu::LoadMatrixOp> {
895 using OpConversionPattern<xegpu::LoadMatrixOp>::OpConversionPattern;
897 matchAndRewrite(xegpu::LoadMatrixOp op, OneToNOpAdaptor adaptor,
898 ConversionPatternRewriter &rewriter)
const override {
900 SmallVector<SmallVector<OpFoldResult>> offsetsList;
901 if (
failed(genOffsetsList(rewriter, op, offsetsList)))
904 ArrayRef<int64_t> wgShape = op.getDataShape();
905 VectorType valueTy = llvm::dyn_cast<VectorType>(op.getRes().getType());
906 assert(valueTy &&
"the value type must be vector type!");
907 Type elemTy = valueTy.getElementType();
909 xegpu::DistributeLayoutAttr layout = op.getLayoutAttr();
910 SmallVector<int64_t> sgShape = getSgShapeAndCount(wgShape, layout).first;
911 VectorType newResTy = VectorType::get(sgShape, elemTy);
912 SmallVector<Value> newOps;
913 for (
auto offsets : offsetsList) {
914 auto newOp = xegpu::LoadMatrixOp::create(rewriter, op.getLoc(), newResTy,
915 op.getMemDesc(), offsets,
916 layout.dropSgLayoutAndData());
917 newOps.push_back(newOp);
919 rewriter.replaceOpWithMultiple(op, {newOps});
925struct WgToSgStoreMatrixOp :
public OpConversionPattern<xegpu::StoreMatrixOp> {
926 using OpConversionPattern<xegpu::StoreMatrixOp>::OpConversionPattern;
928 matchAndRewrite(xegpu::StoreMatrixOp op, OneToNOpAdaptor adaptor,
929 ConversionPatternRewriter &rewriter)
const override {
931 SmallVector<SmallVector<OpFoldResult>> offsetsList;
932 if (
failed(genOffsetsList(rewriter, op, offsetsList)))
935 xegpu::DistributeLayoutAttr layout = op.getLayoutAttr();
936 for (
auto [v, offsets] : llvm::zip(adaptor.getData(), offsetsList))
937 xegpu::StoreMatrixOp::create(rewriter, op.getLoc(), v, op.getMemDesc(),
938 offsets, layout.dropSgLayoutAndData());
939 rewriter.eraseOp(op);
945struct WgToSgVectorStepOp :
public OpConversionPattern<vector::StepOp> {
946 using OpConversionPattern<vector::StepOp>::OpConversionPattern;
948 matchAndRewrite(vector::StepOp op, OneToNOpAdaptor adaptor,
949 ConversionPatternRewriter &rewriter)
const override {
950 xegpu::DistributeLayoutAttr layout =
952 if (!layout || !layout.isForWorkgroup())
955 Location loc = op.getLoc();
956 VectorType type = op.getResult().getType();
957 auto wgShape = type.getShape();
958 std::optional<SmallVector<int64_t>> sgShape =
959 getSgShapeAndCount(wgShape, layout).first;
964 gpu::SubgroupIdOp::create(rewriter, loc,
nullptr);
966 layout.computeDistributedCoords(rewriter, loc, sgId, wgShape);
970 VectorType newTy = type.cloneWith(*sgShape, type.getElementType());
971 auto steps = vector::StepOp::create(rewriter, loc, newTy);
972 SmallVector<Value> newOps;
973 for (
auto offsets : *sgOffsets) {
976 vector::BroadcastOp::create(rewriter, loc, newTy, offsets[0]);
978 arith::AddIOp::create(rewriter, loc, steps, bcastOffset);
979 newOps.push_back(finalSteps);
982 rewriter.replaceOpWithMultiple(op, {newOps});
988struct WgToSgVectorShapeCastOp
989 :
public OpConversionPattern<vector::ShapeCastOp> {
990 using OpConversionPattern<vector::ShapeCastOp>::OpConversionPattern;
993 matchAndRewrite(vector::ShapeCastOp op, OneToNOpAdaptor adaptor,
994 ConversionPatternRewriter &rewriter)
const override {
996 VectorType resultType = dyn_cast<VectorType>(op.getResult().getType());
1000 ArrayRef<int64_t> wgShape = resultType.getShape();
1001 xegpu::DistributeLayoutAttr layout =
1003 if (!layout || !layout.isForWorkgroup())
1008 auto srcType = dyn_cast<VectorType>(op.getSource().getType());
1012 ArrayRef<int64_t> srcShape = srcType.getShape();
1014 xegpu::DistributeLayoutAttr layoutToDistribute = layout;
1015 SmallVector<int64_t> expandedUnitDims;
1017 xegpu::DistributeLayoutAttr sourceLayout =
1020 if (!sourceLayout.isSliceOf(layout))
1021 return rewriter.notifyMatchFailure(
1022 op,
"The ShapeCast op only expands dimensions, the input layout "
1023 "must be a slice of the result layout.");
1025 assert(layoutToDistribute.isEqualTo(
1026 layoutToDistribute.setUnitDimData(expandedUnitDims)) &&
1027 "The sg_data for unit dimensions should be set as 1");
1030 SmallVector<int64_t> sgShape =
1031 getSgShapeAndCount(wgShape, layoutToDistribute).first;
1032 VectorType newResultType =
1033 VectorType::get(sgShape, resultType.getElementType());
1035 SmallVector<Value> newShapeCastOps;
1036 for (
auto src : adaptor.getSource()) {
1037 auto newShapeCast = vector::ShapeCastOp::create(rewriter, op.getLoc(),
1038 newResultType, src);
1039 newShapeCastOps.push_back(newShapeCast.getResult());
1042 rewriter.replaceOpWithMultiple(op, {newShapeCastOps});
1079struct WgToSgMultiDimReductionOp
1080 :
public OpConversionPattern<vector::MultiDimReductionOp> {
1081 using OpConversionPattern<vector::MultiDimReductionOp>::OpConversionPattern;
1084 matchAndRewrite(vector::MultiDimReductionOp op, OneToNOpAdaptor adaptor,
1085 ConversionPatternRewriter &rewriter)
const override {
1086 Location loc = op.getLoc();
1088 VectorType srcType = op.getSourceVectorType();
1089 Type resultTy = op.getResult().getType();
1090 VectorType dstVecType = dyn_cast<VectorType>(resultTy);
1091 bool isScalarResult = !dstVecType;
1093 auto originalSrcShape = srcType.getShape();
1094 Type elemTy = srcType.getElementType();
1096 xegpu::DistributeLayoutAttr layout =
1098 if (!layout || !layout.isForWorkgroup())
1101 auto reductionDims = llvm::to_vector(op.getReductionDims());
1104 SmallVector<int64_t> sgLayout;
1105 SmallVector<int64_t> sgData;
1106 xegpu::DistributeLayoutAttr parentLayout;
1107 if (
auto sliceAttr = dyn_cast<xegpu::SliceAttr>(layout)) {
1108 parentLayout = sliceAttr.getParent();
1109 sgLayout = parentLayout.getEffectiveSgLayoutAsInt();
1110 sgData = parentLayout.getEffectiveSgDataAsInt();
1112 return rewriter.notifyMatchFailure(
1113 op,
"Reduction should have SliceAttr layout");
1116 SmallVector<Value> localReductions;
1117 auto sgSrcs = adaptor.getSource();
1118 auto sgSrcType = dyn_cast<VectorType>(sgSrcs.front().getType());
1119 SmallVector<int64_t> sgSrcShape(sgSrcType.getShape().begin(),
1120 sgSrcType.getShape().end());
1127 auto originalDstShape = dstVecType.getShape();
1128 SmallVector<int64_t> sgDstShape =
1129 getSgShapeAndCount(originalDstShape, layout).first;
1130 sgDstType = VectorType::get(sgDstShape, elemTy);
1135 for (
auto sgSrc : sgSrcs) {
1138 rewriter, loc, sgDstType, op.getKind());
1140 auto localReduce = vector::MultiDimReductionOp::create(
1141 rewriter, loc, sgDstType, op.getKind(), sgSrc, neutralLocalAcc,
1143 localReductions.push_back(localReduce.getResult());
1147 SmallVector<int64_t> crossSgReductionDims;
1148 for (int64_t reductionDim : reductionDims) {
1149 bool needsCrossSubgroupReduction =
1150 (sgLayout[reductionDim] > 1) &&
1151 (sgData[reductionDim] < originalSrcShape[reductionDim]);
1153 if (needsCrossSubgroupReduction) {
1154 crossSgReductionDims.push_back(reductionDim);
1159 if (crossSgReductionDims.empty()) {
1160 SmallVector<Value> results;
1161 for (
auto localResult : localReductions) {
1163 rewriter, loc, op.getKind(), localResult, adaptor.getAcc()[0]);
1164 results.push_back(finalResult);
1166 rewriter.replaceOpWithMultiple(op, {results});
1171 auto slmStoreDataShape = sgSrcShape;
1172 for (int64_t dim : reductionDims)
1173 slmStoreDataShape[dim] = 1;
1174 VectorType slmStoreDataType = VectorType::get(slmStoreDataShape, elemTy);
1175 SmallVector<Value> slmStoreData;
1176 for (
auto localResult : localReductions) {
1177 if (isScalarResult) {
1179 slmStoreData.push_back(vector::BroadcastOp::create(
1180 rewriter, loc, slmStoreDataType, localResult));
1182 slmStoreData.push_back(vector::ShapeCastOp::create(
1183 rewriter, loc, slmStoreDataType, localResult));
1187 SmallVector<int64_t> slmShape(originalSrcShape.begin(),
1188 originalSrcShape.end());
1189 SmallVector<int> slmSgData(sgData.begin(), sgData.end());
1190 SmallVector<int> slmSgLayout(sgLayout.begin(), sgLayout.end());
1191 for (
int dim : reductionDims) {
1192 slmShape[dim] = sgLayout[dim];
1195 xegpu::LayoutAttr slmStoreLayout =
1196 xegpu::LayoutAttr::get(rewriter.getContext(), slmSgLayout, slmSgData);
1200 auto bytesPerElement = bitWidth / 8;
1202 auto slmTy = MemRefType::get({slmSize}, rewriter.getI8Type(), {}, 3);
1203 auto slm = memref::AllocaOp::create(rewriter, loc, slmTy);
1205 auto memDescType = xegpu::MemDescType::get(rewriter.getContext(), slmShape,
1208 xegpu::CreateMemDescOp::create(rewriter, loc, memDescType, slm);
1211 auto sgId = gpu::SubgroupIdOp::create(rewriter, loc,
1212 rewriter.getIndexType(),
nullptr);
1214 auto slmStoreCoords =
1215 slmStoreLayout.computeDistributedCoords(rewriter, loc, sgId, slmShape);
1216 if (
failed(slmStoreCoords))
1218 for (
auto [data, coord] : llvm::zip(slmStoreData, *slmStoreCoords)) {
1219 SmallVector<OpFoldResult> coordOfr(coord.begin(), coord.end());
1220 xegpu::StoreMatrixOp::create(rewriter, loc, data, memDesc.getResult(),
1225 gpu::BarrierOp::create(rewriter, loc);
1228 SmallVector<int64_t> slmLoadDataShape(sgSrcShape.begin(), sgSrcShape.end());
1229 for (int64_t dim : reductionDims) {
1230 slmLoadDataShape[dim] = slmShape[dim];
1231 slmSgData[dim] = slmShape[dim];
1233 xegpu::LayoutAttr slmLoadLayout =
1234 xegpu::LayoutAttr::get(rewriter.getContext(), slmSgLayout, slmSgData);
1235 auto slmLoadCoords =
1236 slmLoadLayout.computeDistributedCoords(rewriter, loc, sgId, slmShape);
1237 if (
failed(slmLoadCoords))
1240 VectorType slmLoadType = VectorType::get(slmLoadDataShape, elemTy);
1241 SmallVector<Value> slmLoadData;
1242 for (
auto coord : *slmLoadCoords) {
1243 SmallVector<OpFoldResult> coordOfr(coord.begin(), coord.end());
1244 slmLoadData.push_back(xegpu::LoadMatrixOp::create(
1245 rewriter, loc, slmLoadType, memDesc.getResult(), coordOfr,
1252 rewriter, loc, sgDstType, op.getKind());
1254 SmallVector<Value> finalResults;
1255 for (
size_t i = 0; i < slmLoadData.size(); ++i) {
1256 auto loaded = slmLoadData[i];
1257 auto finalReduce = vector::MultiDimReductionOp::create(
1258 rewriter, loc, sgDstType, op.getKind(), loaded, neutralFinalAcc,
1261 rewriter, loc, op.getKind(), finalReduce.getResult(),
1262 adaptor.getAcc()[i]));
1264 rewriter.replaceOpWithMultiple(op, {finalResults});
1270struct WgToSgVectorTransposeOp
1271 :
public OpConversionPattern<vector::TransposeOp> {
1272 using OpConversionPattern<vector::TransposeOp>::OpConversionPattern;
1275 matchAndRewrite(vector::TransposeOp op, OneToNOpAdaptor adaptor,
1276 ConversionPatternRewriter &rewriter)
const override {
1277 VectorType resultType = op.getResultVectorType();
1279 ArrayRef<int64_t> wgShape = resultType.getShape();
1280 xegpu::DistributeLayoutAttr layout =
1282 if (!layout || !layout.isForWorkgroup())
1284 xegpu::DistributeLayoutAttr sourceLayout =
1286 if (!sourceLayout || !sourceLayout.isForWorkgroup())
1289 SmallVector<int64_t> sourceSgLayout =
1290 sourceLayout.getEffectiveSgLayoutAsInt();
1291 SmallVector<int64_t> resultSgLayout = layout.getEffectiveSgLayoutAsInt();
1293 ArrayRef<int64_t> permutation = op.getPermutation();
1294 size_t permutationSize = permutation.size();
1295 if (sourceSgLayout.size() != permutationSize ||
1296 resultSgLayout.size() != permutationSize) {
1297 return rewriter.notifyMatchFailure(
1298 op,
"Layouts and permutation must have the same rank");
1303 if (!layout.isTransposeOf(sourceLayout, permutation,
1304 xegpu::LayoutKind::Subgroup))
1305 return rewriter.notifyMatchFailure(
1306 op,
"Result layout is not a valid transpose of source layout "
1307 "according to permutation");
1309 SmallVector<int64_t> sgShape = getSgShapeAndCount(wgShape, layout).first;
1310 VectorType newResultType =
1311 VectorType::get(sgShape, resultType.getElementType());
1313 SmallVector<Value> newTransposeOps;
1314 for (
auto src : adaptor.getVector()) {
1315 auto newTranspose = vector::TransposeOp::create(
1316 rewriter, op.getLoc(), newResultType, src, permutation);
1317 newTransposeOps.push_back(newTranspose.getResult());
1319 rewriter.replaceOpWithMultiple(op, {newTransposeOps});
1325template <
typename MaskOpType>
1326struct WgToSgVectorMaskOp :
public OpConversionPattern<MaskOpType> {
1327 using OpConversionPattern<MaskOpType>::OpConversionPattern;
1329 LogicalResult matchAndRewrite(
1331 typename OpConversionPattern<MaskOpType>::OneToNOpAdaptor adaptor,
1332 ConversionPatternRewriter &rewriter)
const override {
1333 xegpu::DistributeLayoutAttr layout =
1335 if (!layout || !layout.isForWorkgroup())
1338 Location loc = op.getLoc();
1339 VectorType type = op.getResult().getType();
1340 auto wgShape = type.getShape();
1342 SmallVector<Value> wgMaskDimSizes;
1343 if constexpr (std::is_same_v<MaskOpType, vector::ConstantMaskOp>) {
1344 for (int64_t maskSize : op.getMaskDimSizes()) {
1345 wgMaskDimSizes.push_back(
1348 }
else if constexpr (std::is_same_v<MaskOpType, vector::CreateMaskOp>) {
1349 wgMaskDimSizes = llvm::to_vector(op.getOperands());
1353 gpu::SubgroupIdOp::create(rewriter, loc,
nullptr);
1355 layout.computeDistributedCoords(rewriter, loc, sgId, wgShape);
1359 SmallVector<int64_t> sgShape = getSgShapeAndCount(wgShape, layout).first;
1360 VectorType resultType = VectorType::get(sgShape, type.getElementType());
1364 SmallVector<Value> newCreateMaskOps;
1365 for (
auto offsetSet : *sgOffsets) {
1366 SmallVector<Value> maskOperands;
1368 for (
auto [i, wgMaskDimSize] : llvm::enumerate(wgMaskDimSizes)) {
1371 Value offset = offsetSet[i];
1372 Value adjustedMaskSize =
1373 arith::SubIOp::create(rewriter, loc, wgMaskDimSize, offset);
1376 arith::MaxSIOp::create(rewriter, loc, adjustedMaskSize, zero);
1378 arith::MinSIOp::create(rewriter, loc, nonNegative, dimSizeVal);
1379 maskOperands.push_back(sgMaskSize);
1382 auto newCreateMaskOp =
1383 vector::CreateMaskOp::create(rewriter, loc, resultType, maskOperands);
1384 newCreateMaskOps.push_back(newCreateMaskOp.getResult());
1387 rewriter.replaceOpWithMultiple(op, {newCreateMaskOps});
1392using WgToSgVectorConstantMaskOp = WgToSgVectorMaskOp<vector::ConstantMaskOp>;
1393using WgToSgVectorCreateMaskOp = WgToSgVectorMaskOp<vector::CreateMaskOp>;
1396struct WgToSgVectorBitCastOp :
public OpConversionPattern<vector::BitCastOp> {
1397 using OpConversionPattern<vector::BitCastOp>::OpConversionPattern;
1400 matchAndRewrite(vector::BitCastOp op, OneToNOpAdaptor adaptor,
1401 ConversionPatternRewriter &rewriter)
const override {
1402 VectorType resultType = op.getResultVectorType();
1404 ArrayRef<int64_t> wgShape = resultType.getShape();
1405 xegpu::DistributeLayoutAttr layout =
1407 if (!layout || !layout.isForWorkgroup())
1410 SmallVector<int64_t> sgShape = getSgShapeAndCount(wgShape, layout).first;
1411 VectorType newResultType =
1412 VectorType::get(sgShape, resultType.getElementType());
1414 SmallVector<Value> newBitCastOps;
1415 for (
auto src : adaptor.getSource()) {
1417 vector::BitCastOp::create(rewriter, op.getLoc(), newResultType, src);
1418 newBitCastOps.push_back(newBitCast.getResult());
1421 rewriter.replaceOpWithMultiple(op, {newBitCastOps});
1427struct WgToSgVectorInterleaveOp
1428 :
public OpConversionPattern<vector::InterleaveOp> {
1429 using OpConversionPattern<vector::InterleaveOp>::OpConversionPattern;
1432 matchAndRewrite(vector::InterleaveOp op, OneToNOpAdaptor adaptor,
1433 ConversionPatternRewriter &rewriter)
const override {
1434 VectorType resultType = op.getResultVectorType();
1436 ArrayRef<int64_t> wgShape = resultType.getShape();
1437 xegpu::DistributeLayoutAttr layout =
1439 if (!layout || !layout.isForWorkgroup())
1442 SmallVector<int64_t> sgShape = getSgShapeAndCount(wgShape, layout).first;
1443 VectorType newResultType =
1444 VectorType::get(sgShape, resultType.getElementType());
1446 SmallVector<Value> newInterleaveOps;
1449 for (
auto [
lhs,
rhs] : llvm::zip(adaptor.getLhs(), adaptor.getRhs())) {
1450 auto newInterleave = vector::InterleaveOp::create(
1451 rewriter, op.getLoc(), newResultType,
lhs,
rhs);
1452 newInterleaveOps.push_back(newInterleave.getResult());
1455 rewriter.replaceOpWithMultiple(op, {newInterleaveOps});
1461struct WgToSgVectorDeinterleaveOp
1462 :
public OpConversionPattern<vector::DeinterleaveOp> {
1463 using OpConversionPattern<vector::DeinterleaveOp>::OpConversionPattern;
1466 matchAndRewrite(vector::DeinterleaveOp op, OneToNOpAdaptor adaptor,
1467 ConversionPatternRewriter &rewriter)
const override {
1468 SmallVector<Value> newRes1Ops;
1469 SmallVector<Value> newRes2Ops;
1471 for (
auto src : adaptor.getSource()) {
1472 auto newDeinterleave =
1473 vector::DeinterleaveOp::create(rewriter, op.getLoc(), src);
1474 newRes1Ops.push_back(newDeinterleave.getRes1());
1475 newRes2Ops.push_back(newDeinterleave.getRes2());
1478 SmallVector<SmallVector<Value>> results = {newRes1Ops, newRes2Ops};
1479 rewriter.replaceOpWithMultiple(op, results);
1491 converter.addConversion([](
Type type) ->
Type {
return type; });
1494 converter.addConversion(
1495 [](xegpu::TensorDescType type,
1497 xegpu::DistributeLayoutAttr layout = type.getLayoutAttr();
1498 if (!layout || !layout.isForWorkgroup())
1499 return std::nullopt;
1501 Type elemTy = type.getElementType();
1506 std::tie(subShape, count) = getSgShapeAndCount(
shape, layout);
1508 layout = layout.dropSgLayoutAndData();
1510 auto newTy = xegpu::TensorDescType::get(
1511 type.
getContext(), subShape, elemTy, type.getEncoding(), layout);
1512 result.append(count, newTy);
1518 auto getSubShapeAndCount = [](VectorType vecTy,
1519 xegpu::DistributeLayoutAttr layout)
1521 if (!layout.isForWorkgroup())
1523 return getSgShapeAndCount(vecTy.getShape(), layout);
1528 std::move(loopArgTypes));
1532 patterns.
add<WgToSgCreateNdOp, WgToSgLoadNdOp, WgToSgStoreNdOp, WgToSgDpasOp,
1533 WgToSgDpasMxOp, WgToSgPrefetchNdOp, WgToSgElementwiseOp,
1534 WgToSgVectorBroadcastOp, WgToSgConvertLayoutOp,
1535 WgToSgArithConstantOp, WgToSgLoadGatherOp, WgToSgStoreScatterOp,
1536 WgToSgLoadMatrixOp, WgToSgStoreMatrixOp, WgToSgVectorStepOp,
1537 WgToSgVectorShapeCastOp, WgToSgMultiDimReductionOp,
1538 WgToSgVectorTransposeOp, WgToSgVectorConstantMaskOp,
1539 WgToSgVectorCreateMaskOp, WgToSgVectorBitCastOp,
1540 WgToSgVectorInterleaveOp, WgToSgVectorDeinterleaveOp>(
1547struct XeGPUWgToSgDistributePass
1548 :
public xegpu::impl::XeGPUWgToSgDistributeBase<XeGPUWgToSgDistributePass> {
1549 void runOnOperation()
override;
1553void XeGPUWgToSgDistributePass::runOnOperation() {
1555 Operation *op = getOperation();
1557 signalPassFailure();
1562 llvm::SmallSetVector<UnrealizedConversionCastOp, 8> existingCasts;
1563 getOperation()->walk(
1564 [&](UnrealizedConversionCastOp castOp) { existingCasts.insert(castOp); });
1571 RewritePatternSet patterns(ctx);
1572 ConversionTarget
target(*ctx);
1573 TypeConverter converter;
1576 auto materializeCast = [](OpBuilder &builder, Type type,
ValueRange inputs,
1577 Location loc) -> Value {
1578 return UnrealizedConversionCastOp::create(builder, loc, type, inputs)
1581 converter.addSourceMaterialization(materializeCast);
1582 converter.addTargetMaterialization(materializeCast);
1586 auto getTensorDescType = [](Operation *op) -> xegpu::TensorDescType {
1587 if (
auto createOp = dyn_cast<xegpu::CreateNdDescOp>(op))
1588 return createOp.getType();
1589 if (
auto loadOp = dyn_cast<xegpu::LoadNdOp>(op))
1590 return loadOp.getTensorDescType();
1591 if (
auto storeOp = dyn_cast<xegpu::StoreNdOp>(op))
1592 return storeOp.getTensorDescType();
1593 if (
auto prefetchOp = dyn_cast<xegpu::PrefetchNdOp>(op))
1594 return prefetchOp.getTensorDescType();
1595 return xegpu::TensorDescType();
1598 auto isLegal = [&](xegpu::DistributeLayoutAttr layout) ->
bool {
1599 return !layout || !layout.isForWorkgroup();
1602 target.addDynamicallyLegalOp<xegpu::CreateNdDescOp, xegpu::LoadNdOp,
1603 xegpu::StoreNdOp, xegpu::PrefetchNdOp>(
1604 [=](Operation *op) ->
bool {
1605 auto tdescTy = getTensorDescType(op);
1606 auto layout = dyn_cast_if_present<xegpu::DistributeLayoutAttr>(
1607 tdescTy.getLayout());
1608 return isLegal(layout);
1611 target.addDynamicallyLegalOp<xegpu::DpasOp>([=](xegpu::DpasOp op) ->
bool {
1612 auto layout = op.getLayoutCdAttr();
1613 return isLegal(layout);
1616 target.addDynamicallyLegalOp<xegpu::DpasMxOp>(
1617 [=](xegpu::DpasMxOp op) ->
bool {
1618 auto layout = op.getLayoutCdAttr();
1619 return isLegal(layout);
1622 target.addDynamicallyLegalOp<xegpu::LoadMatrixOp>(
1623 [=](xegpu::LoadMatrixOp op) ->
bool {
1624 return isLegal(op.getLayoutAttr());
1627 target.addDynamicallyLegalOp<xegpu::StoreMatrixOp>(
1628 [=](xegpu::StoreMatrixOp op) ->
bool {
1629 return isLegal(op.getLayoutAttr());
1632 target.addDynamicallyLegalOp<arith::ConstantOp>(
1633 [=](arith::ConstantOp op) ->
bool {
1634 auto vecType = dyn_cast<VectorType>(op.getType());
1640 return isLegal(layout);
1643 target.addDynamicallyLegalOp<
1644 vector::ShapeCastOp, vector::StepOp, vector::TransposeOp,
1645 vector::BroadcastOp, vector::MultiDimReductionOp, vector::ConstantMaskOp,
1646 vector::CreateMaskOp, vector::BitCastOp, vector::InterleaveOp,
1647 vector::DeinterleaveOp>([=](Operation *op) ->
bool {
1651 return isLegal(layout);
1654 target.addDynamicallyLegalOp<xegpu::LoadGatherOp>(
1655 [=](xegpu::LoadGatherOp op) ->
bool {
1656 auto layout = op.getLayoutAttr();
1657 return isLegal(layout);
1660 target.addDynamicallyLegalOp<xegpu::StoreScatterOp>(
1661 [=](xegpu::StoreScatterOp op) ->
bool {
1662 auto layout = op.getLayoutAttr();
1663 return isLegal(layout);
1666 target.addDynamicallyLegalOp<xegpu::ConvertLayoutOp>(
1667 [=](xegpu::ConvertLayoutOp op) ->
bool {
1668 return isLegal(op.getEffectiveInputLayout()) &&
1669 isLegal(op.getTargetLayout());
1672 target.addDynamicallyLegalDialect<math::MathDialect, arith::ArithDialect>(
1673 [=](Operation *op) -> std::optional<bool> {
1678 VectorType resultType =
1686 VectorType operandType = dyn_cast<VectorType>(operand.getType());
1687 if (!operandType || operandType.getShape() != resultType.getShape()) {
1692 xegpu::DistributeLayoutAttr layout =
1694 return isLegal(layout);
1697 target.addLegalOp<UnrealizedConversionCastOp>();
1699 target.markUnknownOpDynamicallyLegal([](Operation *) {
return true; });
1705 applyPartialConversion(getOperation(),
target, std::move(patterns))))
1706 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.