[Mlir-commits] [mlir] [MLIR][XeGPU] Add wg-to-sg distirbution for dpasmx, bitcast, interleave, and deinterleave (PR #194985)
llvmlistbot at llvm.org
llvmlistbot at llvm.org
Wed Apr 29 17:13:25 PDT 2026
llvmorg-github-actions[bot] wrote:
<!--LLVM PR SUMMARY COMMENT-->
@llvm/pr-subscribers-mlir-gpu
Author: Jianhui Li (Jianhui-Li)
<details>
<summary>Changes</summary>
As title.
Assisted by Claude
---
Patch is 21.02 KiB, truncated to 20.00 KiB below, full version: https://github.com/llvm/llvm-project/pull/194985.diff
5 Files Affected:
- (modified) mlir/include/mlir/Dialect/XeGPU/Transforms/XeGPULayoutImpl.h (-6)
- (modified) mlir/lib/Dialect/XeGPU/Transforms/XeGPULayoutImpl.cpp (-14)
- (modified) mlir/lib/Dialect/XeGPU/Transforms/XeGPUWgToSgDistribute.cpp (+181-11)
- (modified) mlir/lib/Dialect/XeGPU/Utils/XeGPUUtils.cpp (+26)
- (modified) mlir/test/Dialect/XeGPU/xegpu-wg-to-sg.mlir (+76)
``````````diff
diff --git a/mlir/include/mlir/Dialect/XeGPU/Transforms/XeGPULayoutImpl.h b/mlir/include/mlir/Dialect/XeGPU/Transforms/XeGPULayoutImpl.h
index 2dd8d9f610faf..5699744f8218e 100644
--- a/mlir/include/mlir/Dialect/XeGPU/Transforms/XeGPULayoutImpl.h
+++ b/mlir/include/mlir/Dialect/XeGPU/Transforms/XeGPULayoutImpl.h
@@ -39,12 +39,6 @@ LogicalResult propagateLayouts(OpBuilder &builder, Operation *target,
LogicalResult resolveLayoutConflicts(Operation *target);
-/// [to-be-deprecated] Set the DistributeLayoutAttr for each OpOperand and
-/// OpResult of of the given operation. If the operation contains regions, it is
-/// also applied recursively to the contained operations operation.
-/// TODO: To be replaced by recoverTemporaryLayouts()
-void recoverTemporaryLayoutsDeprecated(Operation *op);
-
/// Attach layout attributes to all vector-type operands of operations within
/// the given operation's nested region. Reports an error if any vector operand
/// lacks a layout attribute.
diff --git a/mlir/lib/Dialect/XeGPU/Transforms/XeGPULayoutImpl.cpp b/mlir/lib/Dialect/XeGPU/Transforms/XeGPULayoutImpl.cpp
index d3925c40f9123..c7ae5198933f1 100644
--- a/mlir/lib/Dialect/XeGPU/Transforms/XeGPULayoutImpl.cpp
+++ b/mlir/lib/Dialect/XeGPU/Transforms/XeGPULayoutImpl.cpp
@@ -33,20 +33,6 @@
using namespace mlir;
-void xegpu::recoverTemporaryLayoutsDeprecated(Operation *op) {
- op->walk([&](Operation *nestOp) {
- for (OpOperand &opr : nestOp->getOpOperands()) {
- auto layout = getDistributeLayoutAttr(opr.get());
- setDistributeLayoutAttr(opr, layout);
- }
-
- for (OpResult result : nestOp->getOpResults()) {
- auto layout = getDistributeLayoutAttr(result);
- setDistributeLayoutAttr(result, layout);
- }
- });
-}
-
SmallVector<NamedAttribute>
xegpu::dropSgLayoutAndDataOnAttrs(ArrayRef<NamedAttribute> attrs) {
SmallVector<NamedAttribute> out;
diff --git a/mlir/lib/Dialect/XeGPU/Transforms/XeGPUWgToSgDistribute.cpp b/mlir/lib/Dialect/XeGPU/Transforms/XeGPUWgToSgDistribute.cpp
index 8aa0758943cd1..1706bab27fe29 100644
--- a/mlir/lib/Dialect/XeGPU/Transforms/XeGPUWgToSgDistribute.cpp
+++ b/mlir/lib/Dialect/XeGPU/Transforms/XeGPUWgToSgDistribute.cpp
@@ -343,6 +343,65 @@ struct WgToSgDpasOp : public OpConversionPattern<xegpu::DpasOp> {
}
};
+/// This pattern transforms the DpasMxOp to work at subgroup level.
+struct WgToSgDpasMxOp : public OpConversionPattern<xegpu::DpasMxOp> {
+ using OpConversionPattern<xegpu::DpasMxOp>::OpConversionPattern;
+ LogicalResult
+ matchAndRewrite(xegpu::DpasMxOp op, OneToNOpAdaptor adaptor,
+ ConversionPatternRewriter &rewriter) const override {
+
+ Location loc = op.getLoc();
+ VectorType resultTy = op.getResult().getType();
+
+ if (resultTy.getRank() != 2)
+ return failure();
+
+ auto layoutCd = op.getLayoutCdAttr();
+ auto layoutA = op.getLayoutAAttr();
+ auto layoutB = op.getLayoutBAttr();
+ auto layoutAScale = op.getLayoutAScaleAttr();
+ auto layoutBScale = op.getLayoutBScaleAttr();
+
+ if (!layoutCd || !layoutA || !layoutB || !layoutAScale || !layoutBScale)
+ return failure();
+
+ size_t index_c = 0;
+ SmallVector<Value> newDpasMxOps;
+ for (auto [index_a, aVec] : llvm::enumerate(adaptor.getA())) {
+ for (auto [index_b, bVec] : llvm::enumerate(adaptor.getB())) {
+ Value accVal;
+ if (op.getAcc()) {
+ accVal = adaptor.getAcc()[index_c++];
+ }
+ Value scaleAVal;
+ if (op.getScaleA()) {
+ scaleAVal = adaptor.getScaleA()[index_a];
+ }
+ Value scaleBVal;
+ if (op.getScaleB()) {
+ scaleBVal = adaptor.getScaleB()[index_b];
+ }
+
+ ArrayRef<int64_t> aVecShape =
+ llvm::cast<VectorType>(aVec.getType()).getShape();
+ ArrayRef<int64_t> bVecShape =
+ llvm::cast<VectorType>(bVec.getType()).getShape();
+ VectorType resTy = VectorType::get({aVecShape[0], bVecShape[1]},
+ resultTy.getElementType());
+ auto newDpasMxOp = xegpu::DpasMxOp::create(
+ rewriter, loc, resTy, aVec, bVec, accVal, scaleAVal, scaleBVal,
+ layoutA.dropSgLayoutAndData(), layoutB.dropSgLayoutAndData(),
+ layoutCd.dropSgLayoutAndData(), layoutAScale.dropSgLayoutAndData(),
+ layoutBScale.dropSgLayoutAndData());
+
+ newDpasMxOps.push_back(newDpasMxOp);
+ }
+ }
+ rewriter.replaceOpWithMultiple(op, {newDpasMxOps});
+ return success();
+ }
+};
+
/// This pattern transforms vector.broadcast ops to work at subgroup level.
struct WgToSgVectorBroadcastOp
: public OpConversionPattern<vector::BroadcastOp> {
@@ -1403,19 +1462,116 @@ struct WgToSgVectorMaskOp : public OpConversionPattern<MaskOpType> {
using WgToSgVectorConstantMaskOp = WgToSgVectorMaskOp<vector::ConstantMaskOp>;
using WgToSgVectorCreateMaskOp = WgToSgVectorMaskOp<vector::CreateMaskOp>;
+
+// This pattern transforms vector.bitcast ops to work at subgroup level.
+struct WgToSgVectorBitCastOp : public OpConversionPattern<vector::BitCastOp> {
+ using OpConversionPattern<vector::BitCastOp>::OpConversionPattern;
+
+ LogicalResult
+ matchAndRewrite(vector::BitCastOp op, OneToNOpAdaptor adaptor,
+ ConversionPatternRewriter &rewriter) const override {
+ VectorType resultType = op.getResultVectorType();
+
+ ArrayRef<int64_t> wgShape = resultType.getShape();
+ xegpu::DistributeLayoutAttr layout =
+ xegpu::getTemporaryLayout(dyn_cast<OpResult>(op.getResult()));
+ if (!layout || !layout.isForWorkgroup())
+ return failure();
+
+ SmallVector<int64_t> sgShape = getSgShapeAndCount(wgShape, layout).first;
+ VectorType newResultType =
+ VectorType::get(sgShape, resultType.getElementType());
+
+ SmallVector<Value> newBitCastOps;
+ for (auto src : adaptor.getSource()) {
+ auto newBitCast =
+ vector::BitCastOp::create(rewriter, op.getLoc(), newResultType, src);
+ newBitCastOps.push_back(newBitCast.getResult());
+ }
+
+ rewriter.replaceOpWithMultiple(op, {newBitCastOps});
+ return success();
+ }
+};
+
+// This pattern transforms vector.interleave ops to work at subgroup level.
+struct WgToSgVectorInterleaveOp
+ : public OpConversionPattern<vector::InterleaveOp> {
+ using OpConversionPattern<vector::InterleaveOp>::OpConversionPattern;
+
+ LogicalResult
+ matchAndRewrite(vector::InterleaveOp op, OneToNOpAdaptor adaptor,
+ ConversionPatternRewriter &rewriter) const override {
+ VectorType resultType = op.getResultVectorType();
+
+ ArrayRef<int64_t> wgShape = resultType.getShape();
+ xegpu::DistributeLayoutAttr layout =
+ xegpu::getTemporaryLayout(dyn_cast<OpResult>(op.getResult()));
+ if (!layout || !layout.isForWorkgroup())
+ return failure();
+
+ SmallVector<int64_t> sgShape = getSgShapeAndCount(wgShape, layout).first;
+ VectorType newResultType =
+ VectorType::get(sgShape, resultType.getElementType());
+
+ SmallVector<Value> newInterleaveOps;
+ // Interleave operates pairwise: each lhs value is interleaved with
+ // corresponding rhs value
+ for (auto [lhs, rhs] : llvm::zip(adaptor.getLhs(), adaptor.getRhs())) {
+ auto newInterleave = vector::InterleaveOp::create(
+ rewriter, op.getLoc(), newResultType, lhs, rhs);
+ newInterleaveOps.push_back(newInterleave.getResult());
+ }
+
+ rewriter.replaceOpWithMultiple(op, {newInterleaveOps});
+ return success();
+ }
+};
+
+// This pattern transforms vector.deinterleave ops to work at subgroup level.
+struct WgToSgVectorDeinterleaveOp
+ : public OpConversionPattern<vector::DeinterleaveOp> {
+ using OpConversionPattern<vector::DeinterleaveOp>::OpConversionPattern;
+
+ LogicalResult
+ matchAndRewrite(vector::DeinterleaveOp op, OneToNOpAdaptor adaptor,
+ ConversionPatternRewriter &rewriter) const override {
+ xegpu::DistributeLayoutAttr layout =
+ xegpu::getTemporaryLayout(dyn_cast<OpResult>(op.getRes1()));
+ if (!layout || !layout.isForWorkgroup())
+ return failure();
+
+ SmallVector<Value> newRes1Ops;
+ SmallVector<Value> newRes2Ops;
+
+ for (auto src : adaptor.getSource()) {
+ auto newDeinterleave =
+ vector::DeinterleaveOp::create(rewriter, op.getLoc(), src);
+ newRes1Ops.push_back(newDeinterleave.getRes1());
+ newRes2Ops.push_back(newDeinterleave.getRes2());
+ }
+
+ SmallVector<SmallVector<Value>> results = {newRes1Ops, newRes2Ops};
+ rewriter.replaceOpWithMultiple(op, results);
+ return success();
+ }
+};
+
} // namespace
namespace mlir {
namespace xegpu {
void populateXeGPUWgToSgDistributePatterns(RewritePatternSet &patterns) {
patterns.add<WgToSgCreateNdOp, WgToSgLoadNdOp, WgToSgStoreNdOp, WgToSgDpasOp,
- WgToSgPrefetchNdOp, UnrealizedConversionCastOpPattern,
- WgToSgElementwiseOp, WgToSgVectorBroadcastOp,
- WgToSgConvertLayoutOp, WgToSgArithConstantOp, WgToSgLoadGatherOp,
- WgToSgStoreScatterOp, WgToSgLoadMatrixOp, WgToSgStoreMatrixOp,
- WgToSgVectorStepOp, WgToSgVectorShapeCastOp,
- WgToSgMultiDimReductionOp, WgToSgVectorTransposeOp,
- WgToSgVectorConstantMaskOp, WgToSgVectorCreateMaskOp>(
+ WgToSgDpasMxOp, WgToSgPrefetchNdOp,
+ UnrealizedConversionCastOpPattern, WgToSgElementwiseOp,
+ WgToSgVectorBroadcastOp, WgToSgConvertLayoutOp,
+ WgToSgArithConstantOp, WgToSgLoadGatherOp, WgToSgStoreScatterOp,
+ WgToSgLoadMatrixOp, WgToSgStoreMatrixOp, WgToSgVectorStepOp,
+ WgToSgVectorShapeCastOp, WgToSgMultiDimReductionOp,
+ WgToSgVectorTransposeOp, WgToSgVectorConstantMaskOp,
+ WgToSgVectorCreateMaskOp, WgToSgVectorBitCastOp,
+ WgToSgVectorInterleaveOp, WgToSgVectorDeinterleaveOp>(
patterns.getContext());
}
} // namespace xegpu
@@ -1539,6 +1695,12 @@ void XeGPUWgToSgDistributePass::runOnOperation() {
return isLegal(layout);
});
+ target.addDynamicallyLegalOp<xegpu::DpasMxOp>(
+ [=](xegpu::DpasMxOp op) -> bool {
+ auto layout = op.getLayoutCdAttr();
+ return isLegal(layout);
+ });
+
target.addDynamicallyLegalOp<xegpu::LoadMatrixOp>(
[=](xegpu::LoadMatrixOp op) -> bool {
return isLegal(op.getLayoutAttr());
@@ -1560,10 +1722,10 @@ void XeGPUWgToSgDistributePass::runOnOperation() {
return isLegal(layout);
});
- target.addDynamicallyLegalOp<vector::ShapeCastOp, vector::StepOp,
- vector::TransposeOp, vector::BroadcastOp,
- vector::MultiDimReductionOp,
- vector::ConstantMaskOp, vector::CreateMaskOp>(
+ target.addDynamicallyLegalOp<
+ vector::ShapeCastOp, vector::StepOp, vector::TransposeOp,
+ vector::BroadcastOp, vector::MultiDimReductionOp, vector::ConstantMaskOp,
+ vector::CreateMaskOp, vector::BitCastOp, vector::InterleaveOp>(
[=](Operation *op) -> bool {
// Check for either a SliceAttr or LayoutAttr on the result.
auto layout =
@@ -1571,6 +1733,14 @@ void XeGPUWgToSgDistributePass::runOnOperation() {
return isLegal(layout);
});
+ target.addDynamicallyLegalOp<vector::DeinterleaveOp>(
+ [=](vector::DeinterleaveOp op) -> bool {
+ // DeinterleaveOp has two results, check the first one
+ auto layout =
+ xegpu::getTemporaryLayout(dyn_cast<OpResult>(op.getRes1()));
+ return isLegal(layout);
+ });
+
target.addDynamicallyLegalOp<xegpu::LoadGatherOp>(
[=](xegpu::LoadGatherOp op) -> bool {
auto layout = op.getLayoutAttr();
diff --git a/mlir/lib/Dialect/XeGPU/Utils/XeGPUUtils.cpp b/mlir/lib/Dialect/XeGPU/Utils/XeGPUUtils.cpp
index 2d1ce6eea17aa..fb0b5849fa46e 100644
--- a/mlir/lib/Dialect/XeGPU/Utils/XeGPUUtils.cpp
+++ b/mlir/lib/Dialect/XeGPU/Utils/XeGPUUtils.cpp
@@ -183,6 +183,32 @@ xegpu::getDistributeLayoutAttr(const OpOperand &opr) {
return dpasOp.getLayoutCdAttr();
}
}
+ if (auto dpasMxOp = dyn_cast<xegpu::DpasMxOp>(op)) {
+ // DpasMxOp has operands: a, b, optional acc, optional scale_a, optional scale_b
+ // Use AttrSizedOperandSegments to determine which operand this is
+ auto segmentSizesAttr = dpasMxOp->getAttrOfType<DenseI32ArrayAttr>(
+ dpasMxOp.getOperandSegmentSizesAttrName());
+ if (!segmentSizesAttr)
+ return nullptr;
+
+ auto segmentSizes = segmentSizesAttr.asArrayRef();
+ unsigned aSize = segmentSizes[0];
+ unsigned bSize = segmentSizes[1];
+ unsigned accSize = segmentSizes[2];
+ unsigned scaleASize = segmentSizes[3];
+
+ if (idx < aSize) {
+ return dpasMxOp.getLayoutAAttr();
+ } else if (idx < aSize + bSize) {
+ return dpasMxOp.getLayoutBAttr();
+ } else if (idx < aSize + bSize + accSize) {
+ return dpasMxOp.getLayoutCdAttr();
+ } else if (idx < aSize + bSize + accSize + scaleASize) {
+ return dpasMxOp.getLayoutAScaleAttr();
+ } else {
+ return dpasMxOp.getLayoutBScaleAttr();
+ }
+ }
if (auto convertOp = dyn_cast<xegpu::ConvertLayoutOp>(op)) {
return convertOp.getInputLayoutAttr();
}
diff --git a/mlir/test/Dialect/XeGPU/xegpu-wg-to-sg.mlir b/mlir/test/Dialect/XeGPU/xegpu-wg-to-sg.mlir
index f2cc05808ed12..fe5e9d014fb87 100644
--- a/mlir/test/Dialect/XeGPU/xegpu-wg-to-sg.mlir
+++ b/mlir/test/Dialect/XeGPU/xegpu-wg-to-sg.mlir
@@ -93,6 +93,40 @@ gpu.module @test_distribution {
gpu.return
}
+ // CHECK-LABEL: dpas_mx
+ gpu.func @dpas_mx(%a: memref<128x128xf8E5M2>, %b: memref<128x128xf8E5M2>, %a_scale: memref<128x4xf8E8M0FNU>, %b_scale: memref<4x128xf8E8M0FNU>) {
+ // CHECK: %[[DPAS_MX:.*]] = xegpu.dpas_mx %{{.*}}, %{{.*}}, %{{.*}} scale_a = %{{.*}} scale_b = %{{.*}} {layout_a = #xegpu.layout<lane_layout = [1, 16], lane_data = [1, 2]>, layout_a_scale = #xegpu.layout<lane_layout = [16, 1], lane_data = [1, 1]>, layout_b = #xegpu.layout<lane_layout = [1, 16], lane_data = [2, 1]>, layout_b_scale = #xegpu.layout<lane_layout = [1, 16], lane_data = [1, 1]>, layout_cd = #xegpu.layout<lane_layout = [1, 16], lane_data = [1, 1]>} : vector<16x128xf8E5M2>, vector<128x16xf8E5M2>, vector<16x16xbf16>, vector<16x4xf8E8M0FNU>, vector<4x16xf8E8M0FNU> -> vector<16x16xbf16>
+ %tdesc_a = xegpu.create_nd_tdesc %a : memref<128x128xf8E5M2>
+ -> !xegpu.tensor_desc<128x128xf8E5M2, #xegpu.layout<sg_layout = [8, 8], sg_data = [16, 128], lane_layout = [1, 16], lane_data = [1, 2]>>
+ %load_a = xegpu.load_nd %tdesc_a[0, 0] {layout = #xegpu.layout<sg_layout = [8, 8], sg_data = [16, 128], lane_layout = [1, 16], lane_data = [1, 2]>}
+ : !xegpu.tensor_desc<128x128xf8E5M2, #xegpu.layout<sg_layout = [8, 8], sg_data = [16, 128], lane_layout = [1, 16], lane_data = [1, 2]>>
+ -> vector<128x128xf8E5M2>
+ %tdesc_b = xegpu.create_nd_tdesc %b : memref<128x128xf8E5M2>
+ -> !xegpu.tensor_desc<128x128xf8E5M2, #xegpu.layout<sg_layout = [8, 8], sg_data = [128, 16], lane_layout = [1, 16], lane_data = [2, 1]>>
+ %load_b = xegpu.load_nd %tdesc_b[0, 0] {layout = #xegpu.layout<sg_layout = [8, 8], sg_data = [128, 16], lane_layout = [1, 16], lane_data = [2, 1]>}
+ : !xegpu.tensor_desc<128x128xf8E5M2, #xegpu.layout<sg_layout = [8, 8], sg_data = [128, 16], lane_layout = [1, 16], lane_data = [2, 1]>>
+ -> vector<128x128xf8E5M2>
+ %tdesc_a_scale = xegpu.create_nd_tdesc %a_scale : memref<128x4xf8E8M0FNU>
+ -> !xegpu.tensor_desc<128x4xf8E8M0FNU, #xegpu.layout<sg_layout = [8, 8], sg_data = [16, 4], lane_layout = [16, 1], lane_data = [1, 1]>>
+ %load_a_scale = xegpu.load_nd %tdesc_a_scale[0, 0] {layout = #xegpu.layout<sg_layout = [8, 8], sg_data = [16, 4], lane_layout = [16, 1], lane_data = [1, 1]>}
+ : !xegpu.tensor_desc<128x4xf8E8M0FNU, #xegpu.layout<sg_layout = [8, 8], sg_data = [16, 4], lane_layout = [16, 1], lane_data = [1, 1]>>
+ -> vector<128x4xf8E8M0FNU>
+ %tdesc_b_scale = xegpu.create_nd_tdesc %b_scale : memref<4x128xf8E8M0FNU>
+ -> !xegpu.tensor_desc<4x128xf8E8M0FNU, #xegpu.layout<sg_layout = [8, 8], sg_data = [4, 16], lane_layout = [1, 16], lane_data = [1, 1]>>
+ %load_b_scale = xegpu.load_nd %tdesc_b_scale[0, 0] {layout = #xegpu.layout<sg_layout = [8, 8], sg_data = [4, 16], lane_layout = [1, 16], lane_data = [1, 1]>}
+ : !xegpu.tensor_desc<4x128xf8E8M0FNU, #xegpu.layout<sg_layout = [8, 8], sg_data = [4, 16], lane_layout = [1, 16], lane_data = [1, 1]>>
+ -> vector<4x128xf8E8M0FNU>
+ %cst = arith.constant dense<0.0> : vector<128x128xbf16>
+ %dpas_mx = xegpu.dpas_mx %load_a, %load_b, %cst scale_a = %load_a_scale scale_b = %load_b_scale
+ {layout_a = #xegpu.layout<sg_layout = [8, 8], sg_data = [16, 128], lane_layout = [1, 16], lane_data = [1, 2]>,
+ layout_b = #xegpu.layout<sg_layout = [8, 8], sg_data = [128, 16], lane_layout = [1, 16], lane_data = [2, 1]>,
+ layout_cd = #xegpu.layout<sg_layout = [8, 8], sg_data = [16, 16], lane_layout = [1, 16], lane_data = [1, 1]>,
+ layout_a_scale = #xegpu.layout<sg_layout = [8, 8], sg_data = [16, 4], lane_layout = [16, 1], lane_data = [1, 1]>,
+ layout_b_scale = #xegpu.layout<sg_layout = [8, 8], sg_data = [4, 16], lane_layout = [1, 16], lane_data = [1, 1]>}
+ : vector<128x128xf8E5M2>, vector<128x128xf8E5M2>, vector<128x128xbf16>, vector<128x4xf8E8M0FNU>, vector<4x128xf8E8M0FNU> -> vector<128x128xbf16>
+ gpu.return
+ }
+
// CHECK-LABEL: dpas_no_sg_data
gpu.func @dpas_no_sg_data(%a: memref<128x128xf16>, %b: memref<128x128xf16>) {
// CHECK: %[[DPAS:.*]] = xegpu.dpas %{{.*}}, %{{.*}} {layout_a = #xegpu.layout<lane_layout = [1, 16], lane_data = [1, 1], order = [1, 0]>, layout_b = #xegpu.layout<lane_layout = [1, 16], lane_data = [2, 1], order = [1, 0]>, layout_cd = #xegpu.layout<lane_layout = [1, 16], lane_data = [1, 1], order = [1, 0]>} : vector<16x128xf16>, vector<128x16xf16> -> vector<16x16xf32>
@@ -1123,6 +1157,48 @@ gpu.module @test_distribution {
gpu.return
}
+ // CHECK-LABEL: @bitcast_distribution
+ gpu.func @bitcast_distribution(%src: memref<256x128xf32>) {
+ %tdesc = xegpu.create_nd_tdesc %src : memref<256x128xf32>
+ -> !xegpu.tensor_desc<256x128xf32, #xegpu.layout<sg_layout = [8, 4], sg_data = [32, 32], lane_layout = [1, 16], lane_data = [1, 1]>>
+ %load = xegpu.load_nd %tdesc[0, 0] {layout = #xegpu.layout<sg_layout = [8, 4], sg_data = [32, 32], lane_layout = [1, 16], lane_data = [1, 1]>}
+ : !xegpu.tensor_desc<256x128xf32, #xegpu.layout<sg_layout = [8, 4], sg_data = [32, 32], lane_layout = [1, 16], lane_data = [1, 1]>>
+ -> vector<256x128xf32>
+ // CHECK: vector.bitcast {{.*}} : vector<32x32xf32> to vector<32x64xi16>
+ %bitcast = vector.bitcast %load {layout_result_0 = #xegpu.layout<sg_layout = [8, 4], sg_data = [32, 64], lane_layout = [1, 16], lane_data = [1, 1]>}
+ : vector<256x128xf32> to vector<256x256xi16>
+ gpu.return
+ }
+
+ // CHECK-LABEL: @interleave_distribution
+ gpu.func @interleave_distribution(%src: memref<256x128xf32>) {
+ %tdesc = xegpu.create_nd_tdesc %src : memref<256x128xf32>
+ -> !xegpu.tensor_desc<256x128xf32, #xegpu.layout<sg_layout = [8, 4], sg_data = [32, 32], lane_layout = [1, 16], lane_data = [1, 1]>>
+ %load1 = xegpu.load_nd %tdesc[0, 0] {layout = #xegpu.layout<sg_layout = [8, 4], sg_data = [32, 32], lane_layout = [1, 16], lane_data = [1, 1]>}
+ : !xegpu.tensor_desc<256x128xf32, #xegpu.layout<sg_layout = [8, 4], sg_data = [32, 32], lane_layout = [1, 16], lane_data = [1, 1]>>
+ -> vector<256x128xf32>
+ %load2 = xegpu.load_nd %tdesc[0, 0] {layout = #xegpu.layout<sg_layout = [8, 4], sg_data = [32, 32], lane_layout = [1, 16], lane_data = [1, 1]>}
+ : !xegpu.tensor_desc<256x128xf32, #xegpu.layout<sg_layout = [8, 4], sg_data = [32, 32], lane_layout = [1, 16], lane_data = [1, 1]>>
+ -> vector<256x128xf32>
+ // CHECK: vector.interleave {{.*}}, {{.*}} : vector<32x32xf32>
+ %interleave = vector.interleave %load1, %load2 {layout_result_0 = #xegpu.layout<sg_layout = [8, 4], sg_data = [32, 64], lane_layout = [1, 16], lane_data = [1, 1]>}
+ : vector<256x128xf32> -> vector<256x256xf32>
+ gpu.return
+ }
+
+ // CHECK-LABEL: @deinterleave_distribution
+ gpu.func...
[truncated]
``````````
</details>
https://github.com/llvm/llvm-project/pull/194985
More information about the Mlir-commits
mailing list