[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