[Mlir-commits] [mlir] [MLIR][XeGPU] Relax LoadStoreMatrixToXeVMPattern to accept non-distributed load_matrix/store_matrix (PR #210595)

llvmlistbot at llvm.org llvmlistbot at llvm.org
Sun Jul 19 09:26:13 PDT 2026


https://github.com/Naaz30 updated https://github.com/llvm/llvm-project/pull/210595

>From 54feb26eb7c8a206e3440d3eddbfddcb612ff92f Mon Sep 17 00:00:00 2001
From: Naazni Yahya <naazniyahya at Naaznis-MacBook-Air.local>
Date: Sat, 18 Jul 2026 02:05:32 +0530
Subject: [PATCH 1/5] [MLIR][Conversion][XeGPU][XeVM] Added xegpu load/store
 flattening(#208902)

---
 .../lib/Conversion/XeGPUToXeVM/XeGPUToXeVM.cpp | 18 +++++++++++-------
 1 file changed, 11 insertions(+), 7 deletions(-)

diff --git a/mlir/lib/Conversion/XeGPUToXeVM/XeGPUToXeVM.cpp b/mlir/lib/Conversion/XeGPUToXeVM/XeGPUToXeVM.cpp
index 6144d7c0c1a15..645fb6494e19c 100644
--- a/mlir/lib/Conversion/XeGPUToXeVM/XeGPUToXeVM.cpp
+++ b/mlir/lib/Conversion/XeGPUToXeVM/XeGPUToXeVM.cpp
@@ -719,9 +719,11 @@ class LoadStoreMatrixToXeVMPattern : public OpConversionPattern<OpType> {
       // Some transforms may leave unit dimension in the 2D vector, adaptors do
       // not catch it for results.
       if (auto vecType = dyn_cast<VectorType>(resType)) {
-        assert(llvm::count_if(vecType.getShape(),
-                              [](int64_t d) { return d != 1; }) <= 1 &&
-               "Expected either 1D vector or nD with unit dimensions");
+        // Flatten to 1D
+        // Accepts genuine multi-dim tiles 
+        // (e.g. a non-distributed `vector<4x8xf32>` result) by reinterpreting the whole
+        // tile as one flat, contiguous run. This is only valid when the
+        // underlying `mem_desc` region is actually contiguous in memory
         resType = VectorType::get({vecType.getNumElements()},
                                   vecType.getElementType());
       }
@@ -787,9 +789,11 @@ class LoadStoreMatrixToXeVMPattern : public OpConversionPattern<OpType> {
 
     if (valOrResVecTy.getNumElements() >= 1) {
       auto chipOpt = xegpu::getChipStr(op);
-      if (!chipOpt ||
-          (*chipOpt != "pvc" && *chipOpt != "bmg" && *chipOpt != "cri")) {
-        // the lowering for chunk load only works for pvc, bmg or cri
+      // Only reject an explicitly unsupported chip; a missing gpu.module/target
+      // (chipOpt == nullopt) is no longer treated as a hard failure, so this
+      // path also works for standalone/non-GPU-kernel-context IR.
+      if (chipOpt && *chipOpt != "pvc" && *chipOpt != "bmg" &&
+          *chipOpt != "cri") {
         return rewriter.notifyMatchFailure(
             op, "The lowering is specific to pvc, bmg or cri.");
       }
@@ -1626,4 +1630,4 @@ void mlir::populateXeGPUToXeVMConversionPatterns(
   patterns.add<DpasMxToXeVMPattern>(typeConverter, patterns.getContext());
   patterns.add<ExtfToXeVMPattern, TruncfToXeVMPattern>(typeConverter,
                                                        patterns.getContext());
-}
+}
\ No newline at end of file

>From b8c1e1efbc222e0bb37782daa19ddb234e63eb01 Mon Sep 17 00:00:00 2001
From: Naazni Yahya <naazniyahya at Naaznis-MacBook-Air.local>
Date: Sun, 19 Jul 2026 19:00:05 +0530
Subject: [PATCH 2/5] MLIR][Conversion][XeGPU][XeVM] Added xegpu load/store
 flattening(#208902)

---
 .../Conversion/XeGPUToXeVM/XeGPUToXeVM.cpp    | 10 ++---
 .../XeGPUToXeVM/loadstore_matrix.mlir         | 41 +++++++++++++++----
 2 files changed, 35 insertions(+), 16 deletions(-)

diff --git a/mlir/lib/Conversion/XeGPUToXeVM/XeGPUToXeVM.cpp b/mlir/lib/Conversion/XeGPUToXeVM/XeGPUToXeVM.cpp
index 645fb6494e19c..03808d2d0f8ff 100644
--- a/mlir/lib/Conversion/XeGPUToXeVM/XeGPUToXeVM.cpp
+++ b/mlir/lib/Conversion/XeGPUToXeVM/XeGPUToXeVM.cpp
@@ -720,10 +720,8 @@ class LoadStoreMatrixToXeVMPattern : public OpConversionPattern<OpType> {
       // not catch it for results.
       if (auto vecType = dyn_cast<VectorType>(resType)) {
         // Flatten to 1D
-        // Accepts genuine multi-dim tiles 
-        // (e.g. a non-distributed `vector<4x8xf32>` result) by reinterpreting the whole
-        // tile as one flat, contiguous run. This is only valid when the
-        // underlying `mem_desc` region is actually contiguous in memory
+        // Accepts multi-dim tile as one flat, contiguous run. 
+        // This is only valid when the underlying mem_desc region is contiguous in memory
         resType = VectorType::get({vecType.getNumElements()},
                                   vecType.getElementType());
       }
@@ -789,9 +787,7 @@ class LoadStoreMatrixToXeVMPattern : public OpConversionPattern<OpType> {
 
     if (valOrResVecTy.getNumElements() >= 1) {
       auto chipOpt = xegpu::getChipStr(op);
-      // Only reject an explicitly unsupported chip; a missing gpu.module/target
-      // (chipOpt == nullopt) is no longer treated as a hard failure, so this
-      // path also works for standalone/non-GPU-kernel-context IR.
+      // reject an explicitly unsupported chip
       if (chipOpt && *chipOpt != "pvc" && *chipOpt != "bmg" &&
           *chipOpt != "cri") {
         return rewriter.notifyMatchFailure(
diff --git a/mlir/test/Conversion/XeGPUToXeVM/loadstore_matrix.mlir b/mlir/test/Conversion/XeGPUToXeVM/loadstore_matrix.mlir
index 07fb09fa2c24b..50f7263b65e05 100644
--- a/mlir/test/Conversion/XeGPUToXeVM/loadstore_matrix.mlir
+++ b/mlir/test/Conversion/XeGPUToXeVM/loadstore_matrix.mlir
@@ -1,10 +1,11 @@
 // RUN: mlir-opt  -split-input-file -convert-xegpu-to-xevm %s | FileCheck %s
 
-gpu.module @test_kernel [#xevm.target<chip = "pvc">] {
+gpu.module @test_kernel[#xevm.target<chip = "pvc">] {
 
   // e.g. for mem_desc<32x32xf16, @strides=[1, 16]>
-  // its memory layout tuple is (blocked shape = [1,1,32,32],strides=[1024,1024,32,1])
-  //CHECK-LABEL: load_store_matrix_plain
+  // its memory layout tuple is (blocked shape =
+  // [1,1,32,32],strides=[1024,1024,32,1])
+  // CHECK-LABEL: load_store_matrix_plain
   gpu.func @load_store_matrix_plain(%arg0: memref<4096xi8, 3>) -> f32 {
 
     //CHECK: %[[INTPTR:.*]] = memref.extract_aligned_pointer_as_index %arg0 : memref<4096xi8, 3> -> index
@@ -324,12 +325,34 @@ gpu.module @test_kernel [#xevm.target<chip = "pvc">] {
 
   //CHECK-LABEL: load_matrix_f8e8m0_vector
   gpu.func @load_matrix_f8e8m0_vector(%arg0: memref<1024xi8, 3>) -> vector<8xf8E8M0FNU> {
-    %c0 = arith.constant 0 : index
-    %0 = xegpu.create_mem_desc %arg0 : memref<1024xi8, 3> -> !xegpu.mem_desc<32x32xf8E8M0FNU>
-    //CHECK: %[[LOADEDV:.*]] = llvm.load %{{.*}} : !llvm.ptr<3> -> vector<8xi8>
-    //CHECK: vector.bitcast %[[LOADEDV]] : vector<8xi8> to vector<8xf8E8M0FNU>
-    %1 = xegpu.load_matrix %0[%c0, %c0]: !xegpu.mem_desc<32x32xf8E8M0FNU>, index, index -> vector<8xf8E8M0FNU>
-    gpu.return %1 : vector<8xf8E8M0FNU>
+    % c0 = arith.constant 0 : index % 0 =
+               xegpu.create_mem_desc %
+               arg0 : memref<1024xi8, 3>->!xegpu.mem_desc<32x32xf8E8M0FNU>
+               // CHECK: %[[LOADEDV:.*]] = llvm.load %{{.*}} : !llvm.ptr<3> ->
+               // vector<8xi8> CHECK: vector.bitcast %[[LOADEDV]] : vector<8xi8>
+               // to vector<8xf8E8M0FNU>
+               % 1 = xegpu.load_matrix % 0 [% c0, % c0]
+        : !xegpu.mem_desc<32x32xf8E8M0FNU>,
+      index, index->vector<8xf8E8M0FNU> gpu.return % 1 : vector<8xf8E8M0FNU>
   }
+}
 
+// -----
+
+// CHECK-LABEL: func.func @m(
+// CHECK: memref.extract_aligned_pointer_as_index
+// CHECK: llvm.inttoptr
+// CHECK: llvm.load
+// CHECK-SAME: vector<32xf32>
+
+module {
+  func.func @m(% arg0 : index) {
+    % alloca =
+        memref.alloca()
+        : memref<4x8xf32, 3> %
+        0 = xegpu.create_mem_desc %
+            alloca : memref<4x8xf32, 3>->!xegpu.mem_desc<4x8xf32> % 1 =
+                xegpu.load_matrix %
+                0 [0, 0] : !xegpu.mem_desc<4x8xf32>->vector<4x8xf32> return
+  }
 }

>From 232df7d956d8cc04c4e8e0a28e4f4ae6a2feb3d6 Mon Sep 17 00:00:00 2001
From: Naazni Yahya <naazniyahya at Naaznis-MacBook-Air.local>
Date: Sun, 19 Jul 2026 19:31:36 +0530
Subject: [PATCH 3/5] [MLIR][Conversion][XeGPU][XeVM] Added xegpu load/store
 flattening(llvm#208902)

---
 mlir/lib/Conversion/XeGPUToXeVM/XeGPUToXeVM.cpp        | 5 +++--
 mlir/test/Conversion/XeGPUToXeVM/loadstore_matrix.mlir | 8 ++------
 2 files changed, 5 insertions(+), 8 deletions(-)

diff --git a/mlir/lib/Conversion/XeGPUToXeVM/XeGPUToXeVM.cpp b/mlir/lib/Conversion/XeGPUToXeVM/XeGPUToXeVM.cpp
index 03808d2d0f8ff..fae578eefd107 100644
--- a/mlir/lib/Conversion/XeGPUToXeVM/XeGPUToXeVM.cpp
+++ b/mlir/lib/Conversion/XeGPUToXeVM/XeGPUToXeVM.cpp
@@ -720,8 +720,9 @@ class LoadStoreMatrixToXeVMPattern : public OpConversionPattern<OpType> {
       // not catch it for results.
       if (auto vecType = dyn_cast<VectorType>(resType)) {
         // Flatten to 1D
-        // Accepts multi-dim tile as one flat, contiguous run. 
-        // This is only valid when the underlying mem_desc region is contiguous in memory
+        // Accepts multi-dim tile as one flat, contiguous run.
+        // This is only valid when the underlying mem_desc region is contiguous
+        // in memory
         resType = VectorType::get({vecType.getNumElements()},
                                   vecType.getElementType());
       }
diff --git a/mlir/test/Conversion/XeGPUToXeVM/loadstore_matrix.mlir b/mlir/test/Conversion/XeGPUToXeVM/loadstore_matrix.mlir
index 50f7263b65e05..040900d5f40f3 100644
--- a/mlir/test/Conversion/XeGPUToXeVM/loadstore_matrix.mlir
+++ b/mlir/test/Conversion/XeGPUToXeVM/loadstore_matrix.mlir
@@ -335,9 +335,6 @@ gpu.module @test_kernel[#xevm.target<chip = "pvc">] {
         : !xegpu.mem_desc<32x32xf8E8M0FNU>,
       index, index->vector<8xf8E8M0FNU> gpu.return % 1 : vector<8xf8E8M0FNU>
   }
-}
-
-// -----
 
 // CHECK-LABEL: func.func @m(
 // CHECK: memref.extract_aligned_pointer_as_index
@@ -345,8 +342,7 @@ gpu.module @test_kernel[#xevm.target<chip = "pvc">] {
 // CHECK: llvm.load
 // CHECK-SAME: vector<32xf32>
 
-module {
-  func.func @m(% arg0 : index) {
+func.func @m(%arg0: index) {
     % alloca =
         memref.alloca()
         : memref<4x8xf32, 3> %
@@ -355,4 +351,4 @@ module {
                 xegpu.load_matrix %
                 0 [0, 0] : !xegpu.mem_desc<4x8xf32>->vector<4x8xf32> return
   }
-}
+}
\ No newline at end of file

>From 53f03a3dcc1adf30f3e4bf87324afd04759ccd78 Mon Sep 17 00:00:00 2001
From: Naazni Yahya <naazniyahya at Naaznis-MacBook-Air.local>
Date: Sun, 19 Jul 2026 21:04:52 +0530
Subject: [PATCH 4/5] [MLIR][Conversion][XeGPU][XeVM] Added xegpu load/store
 flattening(llvm#208902)

---
 .../Conversion/XeGPUToXeVM/loadstore_matrix.mlir | 16 ----------------
 1 file changed, 16 deletions(-)

diff --git a/mlir/test/Conversion/XeGPUToXeVM/loadstore_matrix.mlir b/mlir/test/Conversion/XeGPUToXeVM/loadstore_matrix.mlir
index 040900d5f40f3..1ea11f0628754 100644
--- a/mlir/test/Conversion/XeGPUToXeVM/loadstore_matrix.mlir
+++ b/mlir/test/Conversion/XeGPUToXeVM/loadstore_matrix.mlir
@@ -335,20 +335,4 @@ gpu.module @test_kernel[#xevm.target<chip = "pvc">] {
         : !xegpu.mem_desc<32x32xf8E8M0FNU>,
       index, index->vector<8xf8E8M0FNU> gpu.return % 1 : vector<8xf8E8M0FNU>
   }
-
-// CHECK-LABEL: func.func @m(
-// CHECK: memref.extract_aligned_pointer_as_index
-// CHECK: llvm.inttoptr
-// CHECK: llvm.load
-// CHECK-SAME: vector<32xf32>
-
-func.func @m(%arg0: index) {
-    % alloca =
-        memref.alloca()
-        : memref<4x8xf32, 3> %
-        0 = xegpu.create_mem_desc %
-            alloca : memref<4x8xf32, 3>->!xegpu.mem_desc<4x8xf32> % 1 =
-                xegpu.load_matrix %
-                0 [0, 0] : !xegpu.mem_desc<4x8xf32>->vector<4x8xf32> return
-  }
 }
\ No newline at end of file

>From 94617e6e8d43b27fb89154c70b58fb2fb03f72bf Mon Sep 17 00:00:00 2001
From: Naazni Yahya <naazniyahya at Naaznis-MacBook-Air.local>
Date: Sun, 19 Jul 2026 21:55:58 +0530
Subject: [PATCH 5/5] [MLIR][Conversion][XeGPU][XeVM] Added xegpu load/store
 flattening(llvm#208902)

---
 .../XeGPUToXeVM/loadstore_matrix.mlir         | 23 ++++++++-----------
 1 file changed, 10 insertions(+), 13 deletions(-)

diff --git a/mlir/test/Conversion/XeGPUToXeVM/loadstore_matrix.mlir b/mlir/test/Conversion/XeGPUToXeVM/loadstore_matrix.mlir
index 1ea11f0628754..3307e10451008 100644
--- a/mlir/test/Conversion/XeGPUToXeVM/loadstore_matrix.mlir
+++ b/mlir/test/Conversion/XeGPUToXeVM/loadstore_matrix.mlir
@@ -1,11 +1,10 @@
 // RUN: mlir-opt  -split-input-file -convert-xegpu-to-xevm %s | FileCheck %s
 
-gpu.module @test_kernel[#xevm.target<chip = "pvc">] {
+gpu.module @test_kernel [#xevm.target<chip = "pvc">] {
 
   // e.g. for mem_desc<32x32xf16, @strides=[1, 16]>
-  // its memory layout tuple is (blocked shape =
-  // [1,1,32,32],strides=[1024,1024,32,1])
-  // CHECK-LABEL: load_store_matrix_plain
+  // its memory layout tuple is (blocked shape = [1,1,32,32],strides=[1024,1024,32,1])
+  //CHECK-LABEL: load_store_matrix_plain
   gpu.func @load_store_matrix_plain(%arg0: memref<4096xi8, 3>) -> f32 {
 
     //CHECK: %[[INTPTR:.*]] = memref.extract_aligned_pointer_as_index %arg0 : memref<4096xi8, 3> -> index
@@ -325,14 +324,12 @@ gpu.module @test_kernel[#xevm.target<chip = "pvc">] {
 
   //CHECK-LABEL: load_matrix_f8e8m0_vector
   gpu.func @load_matrix_f8e8m0_vector(%arg0: memref<1024xi8, 3>) -> vector<8xf8E8M0FNU> {
-    % c0 = arith.constant 0 : index % 0 =
-               xegpu.create_mem_desc %
-               arg0 : memref<1024xi8, 3>->!xegpu.mem_desc<32x32xf8E8M0FNU>
-               // CHECK: %[[LOADEDV:.*]] = llvm.load %{{.*}} : !llvm.ptr<3> ->
-               // vector<8xi8> CHECK: vector.bitcast %[[LOADEDV]] : vector<8xi8>
-               // to vector<8xf8E8M0FNU>
-               % 1 = xegpu.load_matrix % 0 [% c0, % c0]
-        : !xegpu.mem_desc<32x32xf8E8M0FNU>,
-      index, index->vector<8xf8E8M0FNU> gpu.return % 1 : vector<8xf8E8M0FNU>
+    %c0 = arith.constant 0 : index
+    %0 = xegpu.create_mem_desc %arg0 : memref<1024xi8, 3> -> !xegpu.mem_desc<32x32xf8E8M0FNU>
+    //CHECK: %[[LOADEDV:.*]] = llvm.load %{{.*}} : !llvm.ptr<3> -> vector<8xi8>
+    //CHECK: vector.bitcast %[[LOADEDV]] : vector<8xi8> to vector<8xf8E8M0FNU>
+    %1 = xegpu.load_matrix %0[%c0, %c0]: !xegpu.mem_desc<32x32xf8E8M0FNU>, index, index -> vector<8xf8E8M0FNU>
+    gpu.return %1 : vector<8xf8E8M0FNU>
   }
+
 }
\ No newline at end of file



More information about the Mlir-commits mailing list