[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