[Mlir-commits] [mlir] [MLIR][XeGPU] Adding Layout Utility inferMaskOffsetLayoutForScatterIO (PR #191573)
Jianhui Li
llvmlistbot at llvm.org
Fri Apr 10 16:48:48 PDT 2026
https://github.com/Jianhui-Li updated https://github.com/llvm/llvm-project/pull/191573
>From 288e1188280ca26acaff1d9cb6a20a6248de9a4e Mon Sep 17 00:00:00 2001
From: Jianhui Li <jian.hui.li at intel.com>
Date: Fri, 10 Apr 2026 23:38:53 +0000
Subject: [PATCH 1/2] refactor with inferMaskOffsetLayoutForScatterIO layout
utility
---
.../XeGPU/Transforms/XeGPULayoutImpl.h | 6 ++
mlir/lib/Dialect/XeGPU/IR/XeGPUDialect.cpp | 2 +-
.../XeGPU/Transforms/XeGPULayoutImpl.cpp | 11 +++
.../XeGPU/Transforms/XeGPUPropagateLayout.cpp | 36 ++--------
.../Transforms/XeGPUSubgroupDistribute.cpp | 13 ++--
.../XeGPU/subgroup-distribute-unit.mlir | 69 +++++--------------
6 files changed, 47 insertions(+), 90 deletions(-)
diff --git a/mlir/include/mlir/Dialect/XeGPU/Transforms/XeGPULayoutImpl.h b/mlir/include/mlir/Dialect/XeGPU/Transforms/XeGPULayoutImpl.h
index 9cf9a8705209b..2172a24bb7a59 100644
--- a/mlir/include/mlir/Dialect/XeGPU/Transforms/XeGPULayoutImpl.h
+++ b/mlir/include/mlir/Dialect/XeGPU/Transforms/XeGPULayoutImpl.h
@@ -112,6 +112,12 @@ inferInsertStridedSliceSourceLayout(DistributeLayoutAttr resLayout,
ArrayRef<int64_t> resShape,
ArrayRef<int64_t> srcShape);
+/// Infers the layout attribute for mask and offset operand for Chunked load
+/// and store, given the anchor layout attribute for the value being load/store.
+DistributeLayoutAttr
+inferMaskOffsetLayoutForScatterIO(DistributeLayoutAttr payloadLayout,
+ int chunkSize);
+
/// Sets up layout for Multi-Reduction operations by creating a SliceAttr for
/// the result.
///
diff --git a/mlir/lib/Dialect/XeGPU/IR/XeGPUDialect.cpp b/mlir/lib/Dialect/XeGPU/IR/XeGPUDialect.cpp
index 64c56b5adf5d7..eaa43c02946d8 100644
--- a/mlir/lib/Dialect/XeGPU/IR/XeGPUDialect.cpp
+++ b/mlir/lib/Dialect/XeGPU/IR/XeGPUDialect.cpp
@@ -585,7 +585,7 @@ DistributeLayoutAttr LayoutAttr::dropDims(SmallVector<int64_t> dimGroup) {
int64_t offset = llvm::count_if(dimGroup, [&](int64_t s) { return s < d; });
newOrder.push_back(d - offset);
}
- if (sgLayout.empty() && laneLayout.empty())
+ if ((sgLayout.empty() && laneLayout.empty()) || newOrder.size() == 1)
newOrder.clear();
auto toAttr = [&](ArrayRef<int64_t> v) -> DenseI32ArrayAttr {
diff --git a/mlir/lib/Dialect/XeGPU/Transforms/XeGPULayoutImpl.cpp b/mlir/lib/Dialect/XeGPU/Transforms/XeGPULayoutImpl.cpp
index c3c40ceb4c6ae..ffbd3b497aae8 100644
--- a/mlir/lib/Dialect/XeGPU/Transforms/XeGPULayoutImpl.cpp
+++ b/mlir/lib/Dialect/XeGPU/Transforms/XeGPULayoutImpl.cpp
@@ -381,6 +381,17 @@ xegpu::inferShapeCastSourceLayout(xegpu::DistributeLayoutAttr resLayout,
return nullptr;
}
+/// Infers the layout attribute for mask and offset operand for Chunked load
+/// and store, given the anchor layout attribute for the value being load/store.
+xegpu::DistributeLayoutAttr xegpu::inferMaskOffsetLayoutForScatterIO(
+ xegpu::DistributeLayoutAttr payloadLayout, int chunkSize) {
+ auto rank = payloadLayout.getRank();
+ if (chunkSize > 1)
+ return payloadLayout.dropDims(
+ llvm::to_vector(llvm::seq<int64_t>(rank - 1, rank)));
+ return payloadLayout;
+}
+
/// Sets up layout for reduction operations by creating a SliceAttr for the
/// result.
///
diff --git a/mlir/lib/Dialect/XeGPU/Transforms/XeGPUPropagateLayout.cpp b/mlir/lib/Dialect/XeGPU/Transforms/XeGPUPropagateLayout.cpp
index 4c30dacae8850..ca09fc7251814 100644
--- a/mlir/lib/Dialect/XeGPU/Transforms/XeGPUPropagateLayout.cpp
+++ b/mlir/lib/Dialect/XeGPU/Transforms/XeGPUPropagateLayout.cpp
@@ -1049,21 +1049,9 @@ void LayoutInfoPropagation::visitLoadGatherOp(
load.setLayoutAttr(requiredAnchorLayoutAttr);
}
- auto maskLayoutAttr = requiredAnchorLayoutAttr;
- // Special handling mask layout for chunked ops: Enforce the default xegpu 1D
- // layout for mask.
- if (chunkSize > 1) {
- if (layoutKind == xegpu::LayoutKind::InstData)
- maskLayoutAttr =
- xegpu::LayoutAttr::get(load->getContext(), {subgroupSize});
- else if (layoutKind == xegpu::LayoutKind::Lane)
- maskLayoutAttr =
- xegpu::LayoutAttr::get(load->getContext(), {subgroupSize}, {1});
- else
- assert(false &&
- "chunked StoreScatterOp should not be used at workgroup level");
- }
-
+ assert((chunkSize <= 1) || (layoutKind != xegpu::LayoutKind::Subgroup));
+ auto maskLayoutAttr = xegpu::inferMaskOffsetLayoutForScatterIO(
+ requiredAnchorLayoutAttr, chunkSize);
LayoutInfo maskLayoutInfo = LayoutInfo(maskLayoutAttr);
auto loadLayoutInfo = LayoutInfo(requiredAnchorLayoutAttr);
@@ -1122,21 +1110,9 @@ void LayoutInfoPropagation::visitStoreScatterOp(
}
LayoutInfo srcLayoutInfo = LayoutInfo(requiredAnchorLayoutAttr);
- auto maskLayoutAttr = requiredAnchorLayoutAttr;
- // Special handling mask layout for chunked ops: Enforce the default xegpu 1D
- // layout for mask.
- if (chunkSize > 1) {
- if (layoutKind == xegpu::LayoutKind::InstData)
- maskLayoutAttr =
- xegpu::LayoutAttr::get(storeScatter->getContext(), {subgroupSize});
- else if (layoutKind == xegpu::LayoutKind::Lane)
- maskLayoutAttr = xegpu::LayoutAttr::get(storeScatter->getContext(),
- {subgroupSize}, {1});
- else
- assert(false &&
- "chunked StoreScatterOp should not be used at workgroup level");
- }
-
+ assert((chunkSize <= 1) || (layoutKind != xegpu::LayoutKind::Subgroup));
+ auto maskLayoutAttr = xegpu::inferMaskOffsetLayoutForScatterIO(
+ requiredAnchorLayoutAttr, chunkSize);
LayoutInfo maskLayoutInfo = LayoutInfo(maskLayoutAttr);
// Propagate the payload operand layout
diff --git a/mlir/lib/Dialect/XeGPU/Transforms/XeGPUSubgroupDistribute.cpp b/mlir/lib/Dialect/XeGPU/Transforms/XeGPUSubgroupDistribute.cpp
index d8ce24ddd5cb0..ca454e632a3ea 100644
--- a/mlir/lib/Dialect/XeGPU/Transforms/XeGPUSubgroupDistribute.cpp
+++ b/mlir/lib/Dialect/XeGPU/Transforms/XeGPUSubgroupDistribute.cpp
@@ -825,12 +825,10 @@ struct StoreDistribution final : public gpu::WarpDistributionPattern {
}
}
- auto layoutPayload =
- xegpu::getTemporaryLayout(storeScatterOp->getOpOperand(0));
+ auto layoutPayload = storeScatterOp.getLayoutAttr();
auto layoutOffsets =
- xegpu::getTemporaryLayout(storeScatterOp->getOpOperand(2));
- auto layoutMask =
- xegpu::getTemporaryLayout(storeScatterOp->getOpOperand(3));
+ xegpu::inferMaskOffsetLayoutForScatterIO(layoutPayload, chunkSize);
+ auto layoutMask = layoutOffsets;
FailureOr<VectorType> distStoreVecByWarpOpOrFailure =
getDistVecTypeBasedOnLaneLayout(layoutPayload, storeVecTy);
@@ -1132,9 +1130,10 @@ struct LoadDistribution final : public gpu::WarpDistributionPattern {
}
}
+ auto layoutPayload = loadGatherOp.getLayoutAttr();
auto layoutOffsets =
- xegpu::getTemporaryLayout(loadGatherOp->getOpOperand(1));
- auto layoutMask = xegpu::getTemporaryLayout(loadGatherOp->getOpOperand(2));
+ xegpu::inferMaskOffsetLayoutForScatterIO(layoutPayload, chunkSize);
+ auto layoutMask = layoutOffsets;
FailureOr<VectorType> distOffsetsByWarpOpOrFailure =
getDistVecTypeBasedOnLaneLayout(layoutOffsets, offsetsTy);
diff --git a/mlir/test/Dialect/XeGPU/subgroup-distribute-unit.mlir b/mlir/test/Dialect/XeGPU/subgroup-distribute-unit.mlir
index 18ebc09caa5aa..27c5bd497b948 100644
--- a/mlir/test/Dialect/XeGPU/subgroup-distribute-unit.mlir
+++ b/mlir/test/Dialect/XeGPU/subgroup-distribute-unit.mlir
@@ -469,10 +469,9 @@ gpu.func @vector_multi_reduction_3d_trivial_reduction(%laneid: index) {
gpu.return
}
-
// CHECK-LABEL: gpu.func @scatter_ops_chunksize({{.*}}) {
-// CHECK: %[[OFFSETS:.*]] = arith.constant {{.*}} dense<12> : vector<16xindex>
-// CHECK: %[[MASKS:.*]] = arith.constant {{.*}} dense<true> : vector<16xi1>
+// CHECK: %[[OFFSETS:.*]] = arith.constant dense<12> : vector<16xindex>
+// CHECK: %[[MASKS:.*]] = arith.constant dense<true> : vector<16xi1>
// CHECK: %[[W:.*]]:4 = gpu.warp_execute_on_lane_0(%{{.*}})[16]
// CHECK-SAME: -> (vector<1x8xf16>, memref<256xf16>, vector<1xindex>, vector<1xi1>) {
// CHECK: gpu.yield %{{.*}}, %{{.*}}, %[[OFFSETS]], %[[MASKS]] :
@@ -484,34 +483,19 @@ gpu.func @vector_multi_reduction_3d_trivial_reduction(%laneid: index) {
// CHECK-SAME: : vector<8xf16>, memref<256xf16>, vector<1xindex>, vector<1xi1>
gpu.func @scatter_ops_chunksize(%laneid: index, %src: memref<256xf16>) {
gpu.warp_execute_on_lane_0(%laneid)[16] {
- %1 = arith.constant
- {layout_result_0 = #xegpu.layout<lane_layout = [16], lane_data = [1]>}
- dense<1>: vector<16xi1>
- %offset = arith.constant
- {layout_result_0 = #xegpu.layout<lane_layout = [16], lane_data = [1]>}
- dense<12> : vector<16xindex>
- %3 = xegpu.load %src[%offset], %1 <{chunk_size=8}>
- {
- layout_operand_1 = #xegpu.layout<lane_layout = [16], lane_data = [1]>,
- layout_operand_2 = #xegpu.layout<lane_layout = [16], lane_data = [1]>,
- layout_result_0 = #xegpu.layout<lane_layout = [16, 1], lane_data = [1, 2]>
- }
+ %1 = arith.constant dense<1>: vector<16xi1>
+ %offset = arith.constant dense<12> : vector<16xindex>
+ %3 = xegpu.load %src[%offset], %1 <{chunk_size=8, layout = #xegpu.layout<lane_layout = [16, 1], lane_data = [1, 2]>}>
: memref<256xf16>, vector<16xindex>, vector<16xi1> -> vector<16x8xf16>
- xegpu.store %3, %src[%offset], %1 <{chunk_size=8}>
- {
- layout_operand_0 = #xegpu.layout<lane_layout = [16, 1], lane_data = [1, 2]>,
- layout_operand_2 = #xegpu.layout<lane_layout = [16], lane_data = [1]>,
- layout_operand_3 = #xegpu.layout<lane_layout = [16], lane_data = [1]>
- }
+ xegpu.store %3, %src[%offset], %1 <{chunk_size=8, layout = #xegpu.layout<lane_layout = [16, 1], lane_data = [1, 2]>}>
: vector<16x8xf16>, memref<256xf16>, vector<16xindex>, vector<16xi1>
}
gpu.return
}
-
// CHECK-LABEL: gpu.func @scatter_ops({{.*}}) {
-// CHECK: %[[OFFSETS:.*]] = arith.constant {{.*}} dense<12> : vector<16xindex>
-// CHECK: %[[MASKS:.*]] = arith.constant {{.*}} dense<true> : vector<16xi1>
+// CHECK: %[[OFFSETS:.*]] = arith.constant dense<12> : vector<16xindex>
+// CHECK: %[[MASKS:.*]] = arith.constant dense<true> : vector<16xi1>
// CHECK: %[[W:.*]]:4 = gpu.warp_execute_on_lane_0(%{{.*}})[16]
// CHECK-SAME: -> (vector<1xf16>, memref<256xf16>, vector<1xindex>, vector<1xi1>) {
// CHECK: gpu.yield %{{.*}}, %{{.*}}, %[[OFFSETS]], %[[MASKS]]
@@ -523,23 +507,15 @@ gpu.func @scatter_ops_chunksize(%laneid: index, %src: memref<256xf16>) {
// CHECK-SAME: : vector<1xf16>, memref<256xf16>, vector<1xindex>, vector<1xi1>
gpu.func @scatter_ops(%src: memref<256xf16>, %laneid: index) {
gpu.warp_execute_on_lane_0(%laneid)[16] {
- %1 = arith.constant
- {layout_result_0 = #xegpu.layout<lane_layout = [16], lane_data = [1]>}
- dense<1> : vector<16xi1>
- %offset = arith.constant
- {layout_result_0 = #xegpu.layout<lane_layout = [16], lane_data = [1]>}
- dense<12> : vector<16xindex>
+ %1 = arith.constant dense<1> : vector<16xi1>
+ %offset = arith.constant dense<12> : vector<16xindex>
%3 = xegpu.load %src[%offset], %1
{
- layout_operand_1 = #xegpu.layout<lane_layout = [16], lane_data = [1]>,
- layout_operand_2 = #xegpu.layout<lane_layout = [16], lane_data = [1]>,
- layout_result_0 = #xegpu.layout<lane_layout = [16], lane_data = [1]>
+ layout = #xegpu.layout<lane_layout = [16], lane_data = [1]>
} : memref<256xf16>, vector<16xindex>, vector<16xi1> -> vector<16xf16>
xegpu.store %3, %src[%offset], %1
{
- layout_operand_0 = #xegpu.layout<lane_layout = [16], lane_data = [1]>,
- layout_operand_2 = #xegpu.layout<lane_layout = [16], lane_data = [1]>,
- layout_operand_3 = #xegpu.layout<lane_layout = [16], lane_data = [1]>
+ layout = #xegpu.layout<lane_layout = [16], lane_data = [1]>
}
: vector<16xf16>, memref<256xf16>, vector<16xindex>, vector<16xi1>
}
@@ -547,8 +523,8 @@ gpu.func @scatter_ops(%src: memref<256xf16>, %laneid: index) {
}
// CHECK-LABEL: gpu.func @scatter_ops_with_leading_dims({{.*}}) {
-// CHECK: %[[OFFSETS:.*]] = arith.constant {{.*}} dense<12> : vector<1x1x16xindex>
-// CHECK: %[[MASKS:.*]] = arith.constant {{.*}} dense<true> : vector<1x1x16xi1>
+// CHECK: %[[OFFSETS:.*]] = arith.constant dense<12> : vector<1x1x16xindex>
+// CHECK: %[[MASKS:.*]] = arith.constant dense<true> : vector<1x1x16xi1>
// CHECK: %[[W:.*]]:4 = gpu.warp_execute_on_lane_0(%{{.*}})[16]
// CHECK-SAME: -> (vector<1x1x1xf16>, memref<256xf16>, vector<1x1x1xindex>, vector<1x1x1xi1>) {
// CHECK: gpu.yield %{{.*}}, %{{.*}}, %[[OFFSETS]], %[[MASKS]]
@@ -563,23 +539,12 @@ gpu.func @scatter_ops(%src: memref<256xf16>, %laneid: index) {
gpu.func @scatter_ops_with_leading_dims(%src: memref<256xf16>, %laneid: index) {
gpu.warp_execute_on_lane_0(%laneid)[16] {
%1 = arith.constant
- {layout_result_0 = #xegpu.layout<lane_layout = [1, 1, 16], lane_data = [1, 1, 1]>}
dense<1> : vector<1x1x16xi1>
%offset = arith.constant
- {layout_result_0 = #xegpu.layout<lane_layout = [1, 1, 16], lane_data = [1, 1, 1]>}
dense<12> : vector<1x1x16xindex>
- %3 = xegpu.load %src[%offset], %1
- {
- layout_operand_1 = #xegpu.layout<lane_layout = [1, 1, 16], lane_data = [1, 1, 1]>,
- layout_operand_2 = #xegpu.layout<lane_layout = [1, 1, 16], lane_data = [1, 1, 1]>,
- layout_result_0 = #xegpu.layout<lane_layout = [1, 1, 16], lane_data = [1, 1, 1]>
- } : memref<256xf16>, vector<1x1x16xindex>, vector<1x1x16xi1> -> vector<1x1x16xf16>
- xegpu.store %3, %src[%offset], %1
- {
- layout_operand_0 = #xegpu.layout<lane_layout = [1, 1, 16], lane_data = [1, 1, 1]>,
- layout_operand_2 = #xegpu.layout<lane_layout = [1, 1, 16], lane_data = [1, 1, 1]>,
- layout_operand_3 = #xegpu.layout<lane_layout = [1, 1, 16], lane_data = [1, 1, 1]>
- }
+ %3 = xegpu.load %src[%offset], %1 {layout = #xegpu.layout<lane_layout = [1, 1, 16], lane_data = [1, 1, 1]>}
+ : memref<256xf16>, vector<1x1x16xindex>, vector<1x1x16xi1> -> vector<1x1x16xf16>
+ xegpu.store %3, %src[%offset], %1 { layout = #xegpu.layout<lane_layout = [1, 1, 16], lane_data = [1, 1, 1]>}
: vector<1x1x16xf16>, memref<256xf16>, vector<1x1x16xindex>, vector<1x1x16xi1>
}
gpu.return
>From 77465b9b4984ebbd0dce5c6555078c4375c62d9d Mon Sep 17 00:00:00 2001
From: Jianhui Li <jian.hui.li at intel.com>
Date: Fri, 10 Apr 2026 23:48:35 +0000
Subject: [PATCH 2/2] remove unused variables
---
mlir/lib/Dialect/XeGPU/Transforms/XeGPUPropagateLayout.cpp | 2 --
1 file changed, 2 deletions(-)
diff --git a/mlir/lib/Dialect/XeGPU/Transforms/XeGPUPropagateLayout.cpp b/mlir/lib/Dialect/XeGPU/Transforms/XeGPUPropagateLayout.cpp
index ca09fc7251814..1748bc27b87c8 100644
--- a/mlir/lib/Dialect/XeGPU/Transforms/XeGPUPropagateLayout.cpp
+++ b/mlir/lib/Dialect/XeGPU/Transforms/XeGPUPropagateLayout.cpp
@@ -1027,7 +1027,6 @@ void LayoutInfoPropagation::visitLoadGatherOp(
const uArch *uArch = getUArch(getChipStr(load).value_or(""));
if (!uArch)
return;
- auto subgroupSize = uArch->getSubgroupSize();
VectorType resVecTy = load.getValueType();
int chunkSize = load.getChunkSize().value_or(1);
@@ -1093,7 +1092,6 @@ void LayoutInfoPropagation::visitStoreScatterOp(
const uArch *uArch = getUArch(getChipStr(storeScatter).value_or(""));
if (!uArch)
return;
- auto subgroupSize = uArch->getSubgroupSize();
VectorType srcVecTy = storeScatter.getValueType();
int chunkSize = storeScatter.getChunkSize().value_or(1);
More information about the Mlir-commits
mailing list