[Mlir-commits] [mlir] [mlir][xegpu] Allow create_mem_desc from ND memref (PR #211836)
llvmlistbot at llvm.org
llvmlistbot at llvm.org
Fri Jul 24 08:45:07 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>
Relax the create_mem_desc source operand constraint to accept a statically shaped shared-memory memref of any rank, replacing the 1D/2D-only StaticShared{1,2}DMemRefOf classes with a rank-agnostic StaticSharedMemRefOf.
Add a verifier requiring the source memref to be contiguous row-major, update the op documentation, and add valid/invalid lit tests.
assisted-by-claude
---
Full diff: https://github.com/llvm/llvm-project/pull/211836.diff
6 Files Affected:
- (modified) mlir/include/mlir/Dialect/XeGPU/IR/XeGPUOps.td (+3-2)
- (modified) mlir/include/mlir/Dialect/XeGPU/IR/XeGPUTypes.td (+3-8)
- (modified) mlir/lib/Dialect/XeGPU/IR/CMakeLists.txt (+1)
- (modified) mlir/lib/Dialect/XeGPU/IR/XeGPUOps.cpp (+11)
- (modified) mlir/test/Dialect/XeGPU/invalid.mlir (+10-1)
- (modified) mlir/test/Dialect/XeGPU/ops.mlir (+9)
``````````diff
diff --git a/mlir/include/mlir/Dialect/XeGPU/IR/XeGPUOps.td b/mlir/include/mlir/Dialect/XeGPU/IR/XeGPUOps.td
index 7f8389a6acc47..52785fa310aa4 100644
--- a/mlir/include/mlir/Dialect/XeGPU/IR/XeGPUOps.td
+++ b/mlir/include/mlir/Dialect/XeGPU/IR/XeGPUOps.td
@@ -1277,7 +1277,7 @@ def XeGPU_CreateMemDescOp: XeGPU_Op<"create_mem_desc", [Pure,
as the underlying shared local memory.
Arguments:
- - `source` : 1D or 2D statically shaped memref, representing the raw SLM buffer. The provided memref must be contiguous.
+ - `source` : a statically shaped memref of any rank, representing the raw SLM buffer. The provided memref must be contiguous.
Results:
- `mem_desc` : the memory descriptor (1D or higher).
@@ -1296,9 +1296,10 @@ def XeGPU_CreateMemDescOp: XeGPU_Op<"create_mem_desc", [Pure,
```
}];
- let arguments = (ins AnyTypeOf<[StaticShared1DMemRefOf<[XeGPU_ScalarType]>, StaticShared2DMemRefOf<[XeGPU_ScalarType]>]>:$source);
+ let arguments = (ins StaticSharedMemRefOf<[XeGPU_ScalarType]>:$source);
let results = (outs XeGPU_MemDesc:$mem_desc);
let assemblyFormat = "$source prop-dict attr-dict `` `:` type($source) `->` qualified(type($mem_desc))";
+ let hasVerifier = 1;
}
def XeGPU_LoadMatrixOp: XeGPU_Op<"load_matrix", [MemoryEffects<[MemRead]>,
diff --git a/mlir/include/mlir/Dialect/XeGPU/IR/XeGPUTypes.td b/mlir/include/mlir/Dialect/XeGPU/IR/XeGPUTypes.td
index 0423303c23493..78c29ae9e7f43 100644
--- a/mlir/include/mlir/Dialect/XeGPU/IR/XeGPUTypes.td
+++ b/mlir/include/mlir/Dialect/XeGPU/IR/XeGPUTypes.td
@@ -40,14 +40,9 @@ class XeGPUTypeDef<string name, string typeMnemonic, list<Trait> traits = [],
}
def isSharedPred : CPred<"XeGPUDialect::isSharedMemory(llvm::cast<mlir::MemRefType>($_self))">;
-class StaticShared1DMemRefOf<list<Type> allowedTypes> :
- ConfinedType<MemRefRankOf<allowedTypes, [1]>, [HasStaticShapePred, isSharedPred],
- "reside in share memory and statically 1d shaped " # MemRefOf<allowedTypes>.summary # " ",
- "mlir::MemRefType">;
-
-class StaticShared2DMemRefOf<list<Type> allowedTypes>:
- ConfinedType<MemRefRankOf<allowedTypes, [2]>, [HasStaticShapePred, isSharedPred],
- "reside in share memory and statically 2d shaped " # MemRefOf<allowedTypes>.summary # " ",
+class StaticSharedMemRefOf<list<Type> allowedTypes>:
+ ConfinedType<MemRefOf<allowedTypes>, [HasStaticShapePred, isSharedPred],
+ "reside in share memory and statically shaped " # MemRefOf<allowedTypes>.summary # " ",
"mlir::MemRefType">;
def XeGPU_TensorDesc: XeGPUTypeDef<"TensorDesc", "tensor_desc",
diff --git a/mlir/lib/Dialect/XeGPU/IR/CMakeLists.txt b/mlir/lib/Dialect/XeGPU/IR/CMakeLists.txt
index 7869a28dfed57..4ea73cda6de95 100644
--- a/mlir/lib/Dialect/XeGPU/IR/CMakeLists.txt
+++ b/mlir/lib/Dialect/XeGPU/IR/CMakeLists.txt
@@ -18,6 +18,7 @@ add_mlir_dialect_library(MLIRXeGPUDialect
MLIRArithUtils
MLIRDialectUtils
MLIRGPUDialect
+ MLIRMemRefUtils
MLIRXeVMDialect
MLIRIR
MLIRViewLikeInterface
diff --git a/mlir/lib/Dialect/XeGPU/IR/XeGPUOps.cpp b/mlir/lib/Dialect/XeGPU/IR/XeGPUOps.cpp
index 2ffe883eb0d9a..bfbb8b08bf173 100644
--- a/mlir/lib/Dialect/XeGPU/IR/XeGPUOps.cpp
+++ b/mlir/lib/Dialect/XeGPU/IR/XeGPUOps.cpp
@@ -7,6 +7,7 @@
//===----------------------------------------------------------------------===//
#include "mlir/Dialect/Arith/Utils/Utils.h"
+#include "mlir/Dialect/MemRef/Utils/MemRefUtils.h"
#include "mlir/Dialect/Utils/IndexingUtils.h"
#include "mlir/Dialect/Utils/StaticValueUtils.h"
#include "mlir/Dialect/XeGPU/IR/XeGPU.h"
@@ -195,6 +196,16 @@ IsValidMatrixOpParams(VectorType dataTy, MemDescType mdescTy,
return success();
}
+//===----------------------------------------------------------------------===//
+// XeGPU_CreateMemDescOp
+//===----------------------------------------------------------------------===//
+LogicalResult CreateMemDescOp::verify() {
+ auto srcTy = getSource().getType();
+ if (!memref::isStaticShapeAndContiguousRowMajor(srcTy))
+ return emitOpError("source memref must be contiguous.");
+ return success();
+}
+
//===----------------------------------------------------------------------===//
// XeGPU_CreateNdDescOp
//===----------------------------------------------------------------------===//
diff --git a/mlir/test/Dialect/XeGPU/invalid.mlir b/mlir/test/Dialect/XeGPU/invalid.mlir
index d5d4950fe7d7e..ac87045de7012 100644
--- a/mlir/test/Dialect/XeGPU/invalid.mlir
+++ b/mlir/test/Dialect/XeGPU/invalid.mlir
@@ -607,7 +607,7 @@ func.func @slice_attr_repeat_dim() {
// -----
func.func @create_mem_desc_non_slm() {
%m = memref.alloca() {alignment = 1024} : memref<2048xi8, 1>
- // expected-error at +1 {{operand #0 must be reside in share memory and statically 1d shaped memref }}
+ // expected-error at +1 {{operand #0 must be reside in share memory and statically shaped memref }}
%mem_desc = xegpu.create_mem_desc %m : memref<2048xi8, 1> -> !xegpu.mem_desc<16x64xf16>
return
}
@@ -805,3 +805,12 @@ func.func @contiguity_does_not_divide(%src: i64, %offset: vector<6xindex>, %mask
: i64, vector<6xindex>, vector<6xi1> -> vector<6xf32>
return
}
+
+// -----
+func.func @create_mem_desc_non_contiguous() {
+ %m = memref.alloca() {alignment = 1024} : memref<32x64xf16, 3>
+ %m_sub = memref.subview %m[0, 0][16, 32][1, 1] : memref<32x64xf16, 3> to memref<16x32xf16, strided<[64, 1]>, 3>
+ // expected-error at +1 {{source memref must be contiguous.}}
+ %mem_desc = xegpu.create_mem_desc %m_sub : memref<16x32xf16, strided<[64, 1]>, 3> -> !xegpu.mem_desc<16x32xf16>
+ return
+}
diff --git a/mlir/test/Dialect/XeGPU/ops.mlir b/mlir/test/Dialect/XeGPU/ops.mlir
index 6cffa3eec369b..5591eea00a2ea 100644
--- a/mlir/test/Dialect/XeGPU/ops.mlir
+++ b/mlir/test/Dialect/XeGPU/ops.mlir
@@ -576,6 +576,15 @@ gpu.func @create_mem_desc_with_stride_from_2d_memref() {
gpu.return
}
+// CHECK-LABEL: gpu.func @create_mem_desc_from_3d_memref({{.*}}) {
+gpu.func @create_mem_desc_from_3d_memref() {
+ //CHECK: [[alloc:%.+]] = memref.alloca() {alignment = 1024 : i64} : memref<1x16x64xf16, 3>
+ //CHECK: [[mdesc:%.+]] = xegpu.create_mem_desc [[alloc]] : memref<1x16x64xf16, 3> -> !xegpu.mem_desc<1x16x64xf16>
+ %m = memref.alloca() {alignment = 1024} : memref<1x16x64xf16, 3>
+ %mem_desc = xegpu.create_mem_desc %m : memref<1x16x64xf16, 3> -> !xegpu.mem_desc<1x16x64xf16>
+ gpu.return
+}
+
// CHECK: gpu.func @load_matrix([[ARG0:%.+]]: !xegpu.mem_desc<16x64xf16>)
gpu.func @load_matrix(%arg0: !xegpu.mem_desc<16x64xf16>) {
// CHECK: xegpu.load_matrix [[ARG0]][8, 8] : !xegpu.mem_desc<16x64xf16> -> vector<8x16xf16>
``````````
</details>
https://github.com/llvm/llvm-project/pull/211836
More information about the Mlir-commits
mailing list