[Mlir-commits] [mlir] [MLIR][XeGPU] Support Layout propagation for interleave and deintereleave op (PR #194966)
Artem Kroviakov
llvmlistbot at llvm.org
Thu Apr 30 03:04:12 PDT 2026
================
@@ -877,6 +956,71 @@ 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 at least 2 (or 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");
+ size_t dim = srcShape.size() - 1;
+ int64_t sgDataValue = -1;
+ int64_t instDataValue = -1;
+ int64_t laneDataValue = -1;
+ const int subgroupSize = uArch->getSubgroupSize();
+
+ // Interleave doubles the innermost dimension (ratio = 2)
+ int ratio = 2;
+ int innermostDimLaneLayout = subgroupSize;
+
+ if (layoutKind == xegpu::LayoutKind::Subgroup) {
+ sgDataValue = sgData[dim];
+ // Ensure sgDataValue is divisible by ratio so source sgData can be inferred
+ while ((sgDataValue <= srcShape[dim]) && (sgDataValue % ratio != 0))
+ sgDataValue *= 2;
----------------
akroviakov wrote:
```suggestion
sgDataValue *= ratio;
```
Does 2 refer to the ratio?
https://github.com/llvm/llvm-project/pull/194966
More information about the Mlir-commits
mailing list