[Mlir-commits] [mlir] c05e809 - [MLIR][XeGPU] Support Layout propagation for interleave and deintereleave op (#194966)
llvmlistbot at llvm.org
llvmlistbot at llvm.org
Tue May 5 08:02:09 PDT 2026
Author: Jianhui Li
Date: 2026-05-05T08:02:04-07:00
New Revision: c05e809430c61a22a48089205e15f93949ed7071
URL: https://github.com/llvm/llvm-project/commit/c05e809430c61a22a48089205e15f93949ed7071
DIFF: https://github.com/llvm/llvm-project/commit/c05e809430c61a22a48089205e15f93949ed7071.diff
LOG: [MLIR][XeGPU] Support Layout propagation for interleave and deintereleave op (#194966)
Enable propagation of interleave and deinterleave with their own
propagation rules.
---------
Co-authored-by: Claude Sonnet 4.5 <noreply at anthropic.com>
Added:
Modified:
mlir/include/mlir/Dialect/XeGPU/Transforms/XeGPULayoutImpl.h
mlir/lib/Dialect/XeGPU/Transforms/XeGPULayoutImpl.cpp
mlir/lib/Dialect/XeGPU/Transforms/XeGPUPropagateLayout.cpp
mlir/test/Dialect/XeGPU/propagate-layout.mlir
Removed:
################################################################################
diff --git a/mlir/include/mlir/Dialect/XeGPU/Transforms/XeGPULayoutImpl.h b/mlir/include/mlir/Dialect/XeGPU/Transforms/XeGPULayoutImpl.h
index 2dd8d9f610faf..694e79515ace4 100644
--- a/mlir/include/mlir/Dialect/XeGPU/Transforms/XeGPULayoutImpl.h
+++ b/mlir/include/mlir/Dialect/XeGPU/Transforms/XeGPULayoutImpl.h
@@ -103,6 +103,16 @@ DistributeLayoutAttr inferBitCastSourceLayout(DistributeLayoutAttr resLayout,
int resElemTyBitWidth,
int srcElemTyBitWidth);
+/// Infers the source layout attribute for an interleave operation given the
+/// result layout attribute. Interleave doubles the innermost dimension size.
+DistributeLayoutAttr
+inferInterleaveSourceLayout(DistributeLayoutAttr resLayout);
+
+/// Infers the source layout attribute for a deinterleave operation given the
+/// result layout attribute. Deinterleave halves the innermost dimension size.
+DistributeLayoutAttr
+inferDeinterleaveSourceLayout(DistributeLayoutAttr resLayout);
+
/// Infers the source layout attribute for a shape cast operation given the
/// result layout attribute, result shape, and source shape.
DistributeLayoutAttr inferShapeCastSourceLayout(DistributeLayoutAttr resLayout,
@@ -161,6 +171,14 @@ DistributeLayoutAttr setupBitCastResultLayout(
LayoutKind layoutKind, VectorType srcVectorTy, VectorType resVectorTy,
DistributeLayoutAttr consumerLayout, const uArch::uArch *uArch);
+/// Sets up the result layout for an interleave operation to ensure the source
+/// layout can be safely derived. Interleave doubles the innermost dimension,
+/// so the result layout must ensure that laneData is at least 2 (or a multiple
+/// of 2), and instData must be divisible by innermostDimLaneLayout * 2.
+DistributeLayoutAttr setupInterleaveResultLayout(
+ LayoutKind layoutKind, VectorType srcVectorTy, VectorType resVectorTy,
+ DistributeLayoutAttr consumerLayout, const uArch::uArch *uArch);
+
/// Sets up the result layout for an insert strided slice operation.
/// Creates a result layout based on the specified layout kind (InstData or
/// Lane).
diff --git a/mlir/lib/Dialect/XeGPU/Transforms/XeGPULayoutImpl.cpp b/mlir/lib/Dialect/XeGPU/Transforms/XeGPULayoutImpl.cpp
index f91e80823c2e9..9d187ac143ac9 100644
--- a/mlir/lib/Dialect/XeGPU/Transforms/XeGPULayoutImpl.cpp
+++ b/mlir/lib/Dialect/XeGPU/Transforms/XeGPULayoutImpl.cpp
@@ -458,6 +458,77 @@ xegpu::inferBitCastSourceLayout(xegpu::DistributeLayoutAttr resLayout,
return finalSrcLayout;
}
+/// Infers the source layout attribute for an interleave operation given the
+/// result layout attribute. Interleave doubles the size of the innermost
+/// dimension, so the layout inference is similar to bitcast where the source
+/// element type is larger than the result element type (ratio = 2).
+xegpu::DistributeLayoutAttr
+xegpu::inferInterleaveSourceLayout(xegpu::DistributeLayoutAttr resLayout) {
+
+ SmallVector<int64_t> sgData = resLayout.getEffectiveSgDataAsInt();
+ SmallVector<int64_t> instData = resLayout.getEffectiveInstDataAsInt();
+ SmallVector<int64_t> laneData = resLayout.getEffectiveLaneDataAsInt();
+ size_t sgDataSize = sgData.size();
+ size_t instDataSize = instData.size();
+ size_t laneDataSize = laneData.size();
+ int64_t sgDataValue = -1;
+ int64_t instDataValue = -1;
+ int64_t laneDataValue = -1;
+ int64_t dim = resLayout.getRank() - 1;
+
+ // Interleave doubles the innermost dimension, so we need to halve the
+ // layout values (similar to bitcast with ratio = 2)
+ constexpr int ratio = 2;
+ if (sgDataSize) {
+ assert((sgData.back() % ratio) == 0 &&
+ "sgData not divisible by interleave ratio");
+ sgDataValue = sgData.back() / ratio;
+ }
+ if (instDataSize) {
+ assert((instData.back() % ratio) == 0 &&
+ "instData not divisible by interleave ratio");
+ instDataValue = instData.back() / ratio;
+ }
+ if (laneDataSize) {
+ assert((laneData.back() % ratio) == 0 &&
+ "laneData not divisible by interleave ratio");
+ laneDataValue = laneData.back() / ratio;
+ }
+
+ return resLayout.setDimData(dim, sgDataValue, instDataValue, laneDataValue);
+}
+
+/// Infers the source layout attribute for a deinterleave operation given the
+/// result layout attribute. Deinterleave halves the size of the innermost
+/// dimension, so the layout inference is similar to bitcast where the source
+/// element type is smaller than the result element type (ratio = 2).
+xegpu::DistributeLayoutAttr
+xegpu::inferDeinterleaveSourceLayout(xegpu::DistributeLayoutAttr resLayout) {
+
+ SmallVector<int64_t> sgData = resLayout.getEffectiveSgDataAsInt();
+ SmallVector<int64_t> instData = resLayout.getEffectiveInstDataAsInt();
+ SmallVector<int64_t> laneData = resLayout.getEffectiveLaneDataAsInt();
+ size_t sgDataSize = sgData.size();
+ size_t instDataSize = instData.size();
+ size_t laneDataSize = laneData.size();
+ int64_t sgDataValue = -1;
+ int64_t instDataValue = -1;
+ int64_t laneDataValue = -1;
+ int64_t dim = resLayout.getRank() - 1;
+
+ // Deinterleave halves the innermost dimension, so we need to double the
+ // layout values (similar to bitcast with ratio = 2)
+ constexpr int ratio = 2;
+ if (sgDataSize)
+ sgDataValue = sgData.back() * ratio;
+ if (instDataSize)
+ instDataValue = instData.back() * ratio;
+ if (laneDataSize)
+ laneDataValue = laneData.back() * ratio;
+
+ return resLayout.setDimData(dim, sgDataValue, instDataValue, laneDataValue);
+}
+
/// Infers the source layout attribute for an insert strided slice operation
/// given the result layout attribute, result shape, and source shape. Removes
/// leading dimensions from the result layout to match the source shape size.
@@ -877,6 +948,69 @@ xegpu::DistributeLayoutAttr xegpu::setupBitCastResultLayout(
return consumerLayout;
}
+/// Sets up the result layout for an interleave operation to ensure the source
+/// layout can be safely derived. Interleave doubles the innermost dimension,
+/// so the result layout must ensure that laneData is a multiple
+/// of 2, and instData must be divisible by innermostDimLaneLayout * 2.
+///
+/// Example:
+/// Interleave: vector<128x256xf4> -> vector<128x512xf4>
+/// Consumer layout: laneLayout=[1, 16], laneData=[1, 4], instData=[1, 64]
+/// Result layout adjustment to ensure source can be safely inferred:
+/// - laneData must be >= 2 and multiple of 2 (so source = laneData/2 is
+/// valid)
+/// - instData must be divisible by (16 * 2 = 32) (so source = instData/2 is
+/// valid)
+/// - Adjusted instData: ensure (instData % 32 == 0)
+///
+xegpu::DistributeLayoutAttr xegpu::setupInterleaveResultLayout(
+ xegpu::LayoutKind layoutKind, VectorType srcVecTy, VectorType resVecTy,
+ DistributeLayoutAttr consumerLayout, const xegpu::uArch::uArch *uArch) {
+
+ ArrayRef<int64_t> srcShape = srcVecTy.getShape();
+ SmallVector<int64_t> sgData = consumerLayout.getEffectiveSgDataAsInt();
+ SmallVector<int64_t> instData = consumerLayout.getEffectiveInstDataAsInt();
+ SmallVector<int64_t> laneData = consumerLayout.getEffectiveLaneDataAsInt();
+
+ assert(consumerLayout.getRank() == static_cast<int64_t>(srcShape.size()) &&
+ "consumer layout rank must match source shape rank");
+ const size_t innerMostDim = srcShape.size() - 1;
+ int64_t sgDataValue = -1;
+ int64_t instDataValue = -1;
+ int64_t laneDataValue = -1;
+
+ // Interleave doubles the innermost dimension (ratio = 2)
+ constexpr int ratio = 2;
+ int innermostDimLaneLayout = uArch->getSubgroupSize();
+
+ if (layoutKind == xegpu::LayoutKind::Subgroup) {
+ sgDataValue = sgData[innerMostDim];
+ // Ensure sgDataValue is divisible by ratio so source sgData can be inferred
+ while ((sgDataValue <= srcShape[innerMostDim]) &&
+ (sgDataValue % ratio != 0))
+ sgDataValue *= ratio;
+ } else if (layoutKind == xegpu::LayoutKind::InstData) {
+ instDataValue = instData[innerMostDim];
+ // Adjust instDataValue so it can be divided by (innermostDimLaneLayout *
+ // ratio) when inferring the source layout
+ while ((instDataValue <= srcShape[innerMostDim]) &&
+ (instDataValue % (innermostDimLaneLayout * ratio) != 0))
+ instDataValue *= ratio;
+ assert((srcShape[innerMostDim] % instDataValue) == 0 &&
+ "srcShape, instData, and laneLayout for innermost must be 2^n!");
+ } else if (layoutKind == xegpu::LayoutKind::Lane) {
+ laneDataValue = laneData[innerMostDim];
+ // Ensure laneDataValue is at least 2 and divisible by ratio
+ // so that source laneData = laneDataValue/2 is valid
+ while ((laneDataValue <= srcShape[innerMostDim]) &&
+ (laneDataValue % ratio != 0))
+ laneDataValue *= ratio;
+ }
+
+ return consumerLayout.setDimData(innerMostDim, sgDataValue, instDataValue,
+ laneDataValue);
+}
+
/// Sets up the result layout for an insert strided slice operation.
/// Creates a result layout based on the specified layout kind (InstData or
/// Lane).
@@ -1605,6 +1739,27 @@ xegpu::inferSourceLayoutFromResult(OpOperand &operand,
transpose.getPermutation());
}
+ // For vector::BitCastOp, infer source layout from result layout using
+ // element type bitwidths.
+ if (auto bitcast = dyn_cast<vector::BitCastOp>(op)) {
+ int resElemBitWidth =
+ bitcast.getResultVectorType().getElementType().getIntOrFloatBitWidth();
+ int srcElemBitWidth =
+ bitcast.getSourceVectorType().getElementType().getIntOrFloatBitWidth();
+ return xegpu::inferBitCastSourceLayout(resLayout, resElemBitWidth,
+ srcElemBitWidth);
+ }
+
+ // for vector::interleave
+ if (auto interleave = dyn_cast<vector::InterleaveOp>(op)) {
+ return xegpu::inferInterleaveSourceLayout(resLayout);
+ }
+
+ // for vector::deinterleave
+ if (auto deinterleave = dyn_cast<vector::DeinterleaveOp>(op)) {
+ return xegpu::inferDeinterleaveSourceLayout(resLayout);
+ }
+
// For vector::ExtractStridedSliceOp, simply return result layout
if (dyn_cast<vector::ExtractStridedSliceOp>(op))
return resLayout;
diff --git a/mlir/lib/Dialect/XeGPU/Transforms/XeGPUPropagateLayout.cpp b/mlir/lib/Dialect/XeGPU/Transforms/XeGPUPropagateLayout.cpp
index a5776ebce2e95..a9236187eb77b 100644
--- a/mlir/lib/Dialect/XeGPU/Transforms/XeGPUPropagateLayout.cpp
+++ b/mlir/lib/Dialect/XeGPU/Transforms/XeGPUPropagateLayout.cpp
@@ -347,6 +347,14 @@ class LayoutInfoPropagation
ArrayRef<LayoutInfoLattice *> operands,
ArrayRef<const LayoutInfoLattice *> results);
+ void visitVectorInterleaveOp(vector::InterleaveOp interleave,
+ ArrayRef<LayoutInfoLattice *> operands,
+ ArrayRef<const LayoutInfoLattice *> results);
+
+ void visitVectorDeinterleaveOp(vector::DeinterleaveOp deinterleave,
+ ArrayRef<LayoutInfoLattice *> operands,
+ ArrayRef<const LayoutInfoLattice *> results);
+
void visitPrefetchNdOp(xegpu::PrefetchNdOp prefetch,
ArrayRef<LayoutInfoLattice *> operands,
ArrayRef<const LayoutInfoLattice *> results);
@@ -453,6 +461,12 @@ LogicalResult LayoutInfoPropagation::visitOperation(
.Case([&](vector::BitCastOp bitcastOp) {
visitVectorBitcastOp(bitcastOp, operands, results);
})
+ .Case([&](vector::InterleaveOp interleaveOp) {
+ visitVectorInterleaveOp(interleaveOp, operands, results);
+ })
+ .Case([&](vector::DeinterleaveOp deinterleaveOp) {
+ visitVectorDeinterleaveOp(deinterleaveOp, operands, results);
+ })
.Case([&](vector::MultiDimReductionOp reductionOp) {
visitVectorMultiReductionOp(reductionOp, operands, results);
})
@@ -1064,10 +1078,12 @@ void LayoutInfoPropagation::visitTransposeOp(
LayoutInfo resultLayout = results[0]->getValue();
if (!resultLayout.isAssigned())
return;
+
auto consumerLayoutAttr =
dyn_cast<xegpu::DistributeLayoutAttr>(resultLayout.get());
auto srcLayoutAttr = xegpu::inferTransposeSourceLayout(
consumerLayoutAttr, transpose.getPermutation());
+
// Propagate the new layout to the vector operand.
propagateIfChanged(operands[0], operands[0]->meet(LayoutInfo(srcLayoutAttr)));
}
@@ -1105,6 +1121,63 @@ void LayoutInfoPropagation::visitVectorBitcastOp(
propagateIfChanged(operands[0], operands[0]->meet(LayoutInfo(srcLayoutAttr)));
}
+/// For vector::InterleaveOp, the result has double the innermost dimension size
+/// compared to each source operand. The layout is propagated from result to
+/// sources, adjusting for the 2x size increase.
+void LayoutInfoPropagation::visitVectorInterleaveOp(
+ vector::InterleaveOp interleave, ArrayRef<LayoutInfoLattice *> operands,
+ ArrayRef<const LayoutInfoLattice *> results) {
+ // Need the layout of interleave result to propagate to the operands.
+ LayoutInfo resLayoutInfo = results[0]->getValue();
+ if (!resLayoutInfo.isAssigned())
+ return;
+
+ auto srcVecType = interleave.getSourceVectorType();
+ auto resVecType = interleave.getResultVectorType();
+
+ auto consumerLayoutAttr =
+ dyn_cast<xegpu::DistributeLayoutAttr>(resLayoutInfo.get());
+ const uArch *uArch = getUArch(xegpu::getChipStr(interleave).value_or(""));
+ if (!uArch)
+ return;
+
+ // Setup the result layout to ensure the source layout can be safely derived
+ auto requiredResLayoutAttr = setupInterleaveResultLayout(
+ layoutKind, srcVecType, resVecType, consumerLayoutAttr, uArch);
+
+ xegpu::setTemporaryLayout(interleave->getResult(0), requiredResLayoutAttr);
+
+ // Derive the source layout from the result layout (halve the innermost dim)
+ auto srcLayoutAttr =
+ xegpu::inferInterleaveSourceLayout(requiredResLayoutAttr);
+
+ // Both operands (lhs and rhs) get the same source layout
+ propagateIfChanged(operands[0], operands[0]->meet(LayoutInfo(srcLayoutAttr)));
+ propagateIfChanged(operands[1], operands[1]->meet(LayoutInfo(srcLayoutAttr)));
+}
+
+/// For vector::DeinterleaveOp, the source has double the innermost dimension
+/// size compared to each result. The layout is propagated from results to
+/// source, adjusting for the 2x size decrease in results.
+void LayoutInfoPropagation::visitVectorDeinterleaveOp(
+ vector::DeinterleaveOp deinterleave, ArrayRef<LayoutInfoLattice *> operands,
+ ArrayRef<const LayoutInfoLattice *> results) {
+ // Need the layout of deinterleave results to propagate to the operand.
+ // Use the first result's layout (both results should have the same layout)
+ LayoutInfo resLayoutInfo = results[0]->getValue();
+ if (!resLayoutInfo.isAssigned())
+ return;
+
+ auto consumerLayoutAttr =
+ dyn_cast<xegpu::DistributeLayoutAttr>(resLayoutInfo.get());
+
+ // Derive the source layout from the result layout (double the innermost dim)
+ // No setup function needed - just infer directly
+ auto srcLayoutAttr = xegpu::inferDeinterleaveSourceLayout(consumerLayoutAttr);
+
+ propagateIfChanged(operands[0], operands[0]->meet(LayoutInfo(srcLayoutAttr)));
+}
+
void LayoutInfoPropagation::visitInsertStridedSliceOp(
vector::InsertStridedSliceOp insertStridedSlice,
ArrayRef<LayoutInfoLattice *> operands,
diff --git a/mlir/test/Dialect/XeGPU/propagate-layout.mlir b/mlir/test/Dialect/XeGPU/propagate-layout.mlir
index 72d066d516540..fe4637170bd15 100644
--- a/mlir/test/Dialect/XeGPU/propagate-layout.mlir
+++ b/mlir/test/Dialect/XeGPU/propagate-layout.mlir
@@ -1064,3 +1064,45 @@ func.func @dpas_mx_fp4(%arg0: memref<8x64xf4E2M1FN>, %arg1: memref<64x16xf4E2M1F
return
}
}
+
+// -----
+gpu.module @test {
+// CHECK-LABEL: func.func @vector_interleave_f16(
+// CHECK: %[[LOAD1:.*]] = xegpu.load_nd %{{.*}} <{layout = #xegpu.layout<lane_layout = [1, 16], lane_data = [1, 1]>}>
+// CHECK-SAME: !xegpu.tensor_desc<8x16xf16, #xegpu.layout<lane_layout = [1, 16], lane_data = [1, 1]>> -> vector<8x16xf16>
+// CHECK: %[[LOAD2:.*]] = xegpu.load_nd %{{.*}} <{layout = #xegpu.layout<lane_layout = [1, 16], lane_data = [1, 1]>}>
+// CHECK-SAME: !xegpu.tensor_desc<8x16xf16, #xegpu.layout<lane_layout = [1, 16], lane_data = [1, 1]>> -> vector<8x16xf16>
+// CHECK-NEXT: %{{.*}} = vector.interleave %[[LOAD1]], %[[LOAD2]] {layout_result_0 = #xegpu.layout<lane_layout = [1, 16], lane_data = [1, 2]>}
+// CHECK-SAME: vector<8x16xf16> -> vector<8x32xf16>
+func.func @vector_interleave_f16(%arg0: memref<8x16xf16>, %arg1: memref<8x16xf16>, %arg2: memref<8x32xf16>) {
+ %c0 = arith.constant 0 : index
+ %0 = xegpu.create_nd_tdesc %arg0 : memref<8x16xf16> -> !xegpu.tensor_desc<8x16xf16>
+ %1 = xegpu.create_nd_tdesc %arg1 : memref<8x16xf16> -> !xegpu.tensor_desc<8x16xf16>
+ %2 = xegpu.load_nd %0[0, 0] : !xegpu.tensor_desc<8x16xf16> -> vector<8x16xf16>
+ %3 = xegpu.load_nd %1[0, 0] : !xegpu.tensor_desc<8x16xf16> -> vector<8x16xf16>
+ %4 = vector.interleave %2, %3 : vector<8x16xf16> -> vector<8x32xf16>
+ %5 = xegpu.create_nd_tdesc %arg2 : memref<8x32xf16> -> !xegpu.tensor_desc<8x32xf16>
+ xegpu.store_nd %4, %5[0, 0] : vector<8x32xf16>, !xegpu.tensor_desc<8x32xf16>
+ return
+}
+}
+
+// -----
+gpu.module @test {
+// CHECK-LABEL: func.func @vector_deinterleave_f16(
+// CHECK: %[[LOAD:.*]] = xegpu.load_nd %{{.*}} <{layout = #xegpu.layout<lane_layout = [1, 16], lane_data = [1, 2]>}>
+// CHECK-SAME: !xegpu.tensor_desc<8x32xf16, #xegpu.layout<lane_layout = [1, 16], lane_data = [1, 2]>> -> vector<8x32xf16>
+// CHECK-NEXT: %{{.*}}, %{{.*}} = vector.deinterleave %[[LOAD]] {layout_result_0 = #xegpu.layout<lane_layout = [1, 16], lane_data = [1, 1]>, layout_result_1 = #xegpu.layout<lane_layout = [1, 16], lane_data = [1, 1]>}
+// CHECK-SAME: vector<8x32xf16> -> vector<8x16xf16>
+func.func @vector_deinterleave_f16(%arg0: memref<8x32xf16>, %arg1: memref<8x16xf16>, %arg2: memref<8x16xf16>) {
+ %c0 = arith.constant 0 : index
+ %0 = xegpu.create_nd_tdesc %arg0 : memref<8x32xf16> -> !xegpu.tensor_desc<8x32xf16>
+ %1 = xegpu.load_nd %0[0, 0] : !xegpu.tensor_desc<8x32xf16> -> vector<8x32xf16>
+ %2:2 = vector.deinterleave %1 : vector<8x32xf16> -> vector<8x16xf16>
+ %3 = xegpu.create_nd_tdesc %arg1 : memref<8x16xf16> -> !xegpu.tensor_desc<8x16xf16>
+ %4 = xegpu.create_nd_tdesc %arg2 : memref<8x16xf16> -> !xegpu.tensor_desc<8x16xf16>
+ xegpu.store_nd %2#0, %3[0, 0] : vector<8x16xf16>, !xegpu.tensor_desc<8x16xf16>
+ xegpu.store_nd %2#1, %4[0, 0] : vector<8x16xf16>, !xegpu.tensor_desc<8x16xf16>
+ return
+}
+}
More information about the Mlir-commits
mailing list