[Mlir-commits] [mlir] [MLIR][XeGPU] Add unrolling/blocking support for 3D+ batched operations (PR #201725)

Jianhui Li llvmlistbot at llvm.org
Fri Jun 5 10:59:21 PDT 2026


https://github.com/Jianhui-Li updated https://github.com/llvm/llvm-project/pull/201725

>From 86a3fd755beb2e1a4d6938a356453d464ad47375 Mon Sep 17 00:00:00 2001
From: Jianhui Li <jian.hui.li at intel.com>
Date: Tue, 26 May 2026 21:45:15 +0000
Subject: [PATCH 1/4] [MLIR][XeGPU] Add unrolling/blocking support for 3D+
 operations

Add complete transform pass and lowering support for 3D+ batched operations,
building on the operation definition extensions for load_nd/store_nd/
prefetch_nd/dpas/dpas_mx. This enables end-to-end compilation of batched
GEMM workloads (e.g., [4, 64, 32] x [4, 32, 64] -> [4, 64, 64]).

Key changes:

Transform passes (XeGPUUnroll.cpp):
- Implement 3D batch unrolling using memref.subview to handle batch offsets
- UnrollCreateNdOp: For rank > 2 with memref source, create per-batch
  memref.subview slices and corresponding create_nd_tdesc ops
- UnrollLoadNdOp/StoreNdOp/PrefetchNdOp: Iterate over batch dimension then
  inner 2D tile offsets, reusing batch tdescs across inner tiles
- UnrollDpasOp/UnrollDpasMxOp: Add outer batch loop before M/K/N unrolling
- Remove returnSingleType parameter, compute correct batch counts directly
- Use targetShape (inst_data) for tdesc window size

Blocking pass (XeGPUBlocking.cpp):
- Update getUnrolledTypes() to handle batch dimensions in TensorDescType
- Extend inst_data extraction for >2D vectors with leading batch dims
- Fix getDpasInstDataVectors() to use correct M dimension for >2D types

Layout propagation (XeGPULayoutImpl.cpp, XeGPUPropagateLayout.cpp):
- Extend createScaleLayout() to handle rank >= 2
- Fix store_nd inst_data propagation for rank > 2
- Update lane layout logic to handle >2D vectors with leading unit dims

Distribution (XeGPUWgToSgDistribute.cpp):
- Extend WgToSgDpasOp and WgToSgDpasMxOp to handle >2D result shapes

Lowering (XeGPUToXeVM.cpp):
- Keep batch offset computation in unroll pass via memref.subview
- Maintain 2D lowering logic (last 2 dims for H/W offsets)

Utilities (XeGPUUtils.cpp):
- Update vector/tdesc creation for batch-aware shapes

Tests:
- Add xegpu-blocking.mlir tests for 3D batch dpas and dpas_mx unrolling
- Add simple_3d_gemm.mlir integration test (4-batch GEMM)
- Add simple_3d_mxfp_gemm.mlir integration test (XFAIL due to pre-existing
  dpas_mx lane propagation issue on cri chip)
- Update invalid.mlir test expectations

This PR depends on:
- [MLIR][XeGPU] Extend op definitions to support 3D+: load_nd, store_nd, prefetch_nd
- [MLIR][XeGPU] Extend op definitions to support 3D+: dpas, dpas_mx

Co-Authored-By: Claude Sonnet 4.5 <noreply at anthropic.com>
---
 .../include/mlir/Dialect/XeGPU/IR/XeGPUOps.td |   4 +-
 .../mlir/Dialect/XeGPU/IR/XeGPUTypes.td       |   4 +-
 .../Dialect/XeGPU/Transforms/Transforms.h     |   6 +-
 .../Conversion/XeGPUToXeVM/XeGPUToXeVM.cpp    |  33 +-
 mlir/lib/Dialect/XeGPU/IR/XeGPUOps.cpp        | 123 +++--
 .../XeGPU/Transforms/XeGPUBlocking.cpp        | 148 ++++--
 .../XeGPU/Transforms/XeGPULayoutImpl.cpp      |  12 +-
 .../XeGPU/Transforms/XeGPUPropagateLayout.cpp |  33 +-
 .../Dialect/XeGPU/Transforms/XeGPUUnroll.cpp  | 430 +++++++++++++-----
 .../Transforms/XeGPUWgToSgDistribute.cpp      |  21 +-
 mlir/lib/Dialect/XeGPU/Utils/XeGPUUtils.cpp   |  24 +-
 mlir/test/Dialect/XeGPU/invalid.mlir          |  18 +-
 mlir/test/Dialect/XeGPU/xegpu-blocking.mlir   |  69 +++
 .../Dialect/XeGPU/WG/simple_3d_gemm.mlir      |  55 +++
 .../Dialect/XeGPU/WG/simple_3d_mxfp_gemm.mlir |  55 +++
 .../lib/Dialect/XeGPU/TestXeGPUTransforms.cpp |  14 +-
 16 files changed, 787 insertions(+), 262 deletions(-)
 create mode 100644 mlir/test/Integration/Dialect/XeGPU/WG/simple_3d_gemm.mlir
 create mode 100644 mlir/test/Integration/Dialect/XeGPU/WG/simple_3d_mxfp_gemm.mlir

diff --git a/mlir/include/mlir/Dialect/XeGPU/IR/XeGPUOps.td b/mlir/include/mlir/Dialect/XeGPU/IR/XeGPUOps.td
index fce14999b4011..dcb51352edf13 100644
--- a/mlir/include/mlir/Dialect/XeGPU/IR/XeGPUOps.td
+++ b/mlir/include/mlir/Dialect/XeGPU/IR/XeGPUOps.td
@@ -1460,9 +1460,9 @@ def XeGPU_DpasMxOp : XeGPU_Op<"dpas_mx", [Pure, AttrSizedOperandSegments,
   let arguments = (ins XeGPU_DpasOprType:$a, XeGPU_DpasOprType:$b,
       Optional<XeGPU_DpasResType>:$acc,
       Optional<AnyTypeOf<[F8E8M0FNU,
-                          VectorOfRankAndType<[1, 2], [F8E8M0FNU]>]>>:$scale_a,
+                          VectorOfRankAndType<[1, 2, 3, 4], [F8E8M0FNU]>]>>:$scale_a,
       Optional<AnyTypeOf<[F8E8M0FNU,
-                          VectorOfRankAndType<[1, 2], [F8E8M0FNU]>]>>:$scale_b,
+                          VectorOfRankAndType<[1, 2, 3, 4], [F8E8M0FNU]>]>>:$scale_b,
       OptionalAttr<DistributeLayoutAttr>:$layout_a,
       OptionalAttr<DistributeLayoutAttr>:$layout_b,
       OptionalAttr<DistributeLayoutAttr>:$layout_cd,
diff --git a/mlir/include/mlir/Dialect/XeGPU/IR/XeGPUTypes.td b/mlir/include/mlir/Dialect/XeGPU/IR/XeGPUTypes.td
index 95a3e22cf803b..0423303c23493 100644
--- a/mlir/include/mlir/Dialect/XeGPU/IR/XeGPUTypes.td
+++ b/mlir/include/mlir/Dialect/XeGPU/IR/XeGPUTypes.td
@@ -21,8 +21,8 @@ def XeGPU_ScalarType: AnyTypeOf<[XeGPU_IntType, XeGPU_FloatType]>;
 def XeGPU_PointerType : AnyTypeOf<[UI64, UI32, I64, I32]>;
 def XeGPU_BaseAddrType
     : AnyTypeOf<[Non0RankedMemRefOf<[XeGPU_ScalarType]>, XeGPU_PointerType]>;
-def XeGPU_DpasOprType: FixedVectorOfRankAndType<[1, 2, 3], [XeGPU_ScalarType]>;
-def XeGPU_DpasResType: FixedVectorOfRankAndType<[1, 2], [XeGPU_ScalarType]>;
+def XeGPU_DpasOprType: FixedVectorOfRankAndType<[1, 2, 3, 4], [XeGPU_ScalarType]>;
+def XeGPU_DpasResType: FixedVectorOfRankAndType<[1, 2, 3, 4], [XeGPU_ScalarType]>;
 def XeGPU_OffsetType: FixedVectorOfNonZeroRankOf<[Index]>;
 def XeGPU_MaskType: FixedVectorOfNonZeroRankOf<[I1]>;
 def XeGPU_ValueType: VectorOfRankAndType<[1,2,3,4,5,6,7,8], [XeGPU_ScalarType]>;
diff --git a/mlir/include/mlir/Dialect/XeGPU/Transforms/Transforms.h b/mlir/include/mlir/Dialect/XeGPU/Transforms/Transforms.h
index 919a69908bdce..39adb0f0158c4 100644
--- a/mlir/include/mlir/Dialect/XeGPU/Transforms/Transforms.h
+++ b/mlir/include/mlir/Dialect/XeGPU/Transforms/Transforms.h
@@ -49,11 +49,9 @@ struct UnrollOptions {
 
   /// Function that converts a ShapedType (TensorDescType or VectorType)
   /// into the unrolled type based on the tileShape. It returns a vector of
-  /// types representing the unrolled types for simplicity. When
-  /// `returnSingleType` is true, it returns a vector containing only one single
-  /// unrolled type.
+  /// types representing the unrolled types for simplicity.
   using UnrolledTypeFnType = std::function<SmallVector<Type>(
-      ShapedType type, ArrayRef<int64_t> tileShape, bool returnSingleType)>;
+      ShapedType type, ArrayRef<int64_t> tileShape)>;
   UnrolledTypeFnType getUnrolledTypes = nullptr;
   UnrollOptions &setUnrolledTypesFn(UnrolledTypeFnType fn) {
     getUnrolledTypes = std::move(fn);
diff --git a/mlir/lib/Conversion/XeGPUToXeVM/XeGPUToXeVM.cpp b/mlir/lib/Conversion/XeGPUToXeVM/XeGPUToXeVM.cpp
index cf19d1f014c7b..08f00ab325070 100644
--- a/mlir/lib/Conversion/XeGPUToXeVM/XeGPUToXeVM.cpp
+++ b/mlir/lib/Conversion/XeGPUToXeVM/XeGPUToXeVM.cpp
@@ -54,7 +54,7 @@ enum class NdTdescOffset : uint32_t {
   BasePtr = 0,    // Base pointer (i64)
   BaseShapeW = 2, // Base shape width (i32)
   BaseShapeH = 3, // Base shape height (i32)
-  BasePitch = 4,  // Base pitch (i32)
+  BasePitch = 4,  // Base pitch/stride of dim rank-2 (i32)
 };
 
 static int32_t getNumericXeVMAddrSpace(xegpu::MemorySpace xeGpuMemspace) {
@@ -240,11 +240,12 @@ class CreateNdDescToXeVMPattern
       val = getValueOrCreateCastToIndexLike(rewriter, loc, payloadElemTy, val);
       return val;
     };
-    // Get shape values from op fold results.
-    baseShapeW = createOffset(mixedSizes, 1);
-    baseShapeH = createOffset(mixedSizes, 0);
-    // Get pitch value from op fold results.
-    Value basePitch = createOffset(mixedStrides, 0);
+    // For ND descriptors, the last 2 dimensions are the 2D tile (H, W).
+    // Any leading dimensions are batch dims with associated strides.
+    baseShapeW = createOffset(mixedSizes, rank - 1);
+    baseShapeH = createOffset(mixedSizes, rank - 2);
+    // Pitch is the stride of dim rank-2 (the row stride of the 2D tile).
+    Value basePitch = createOffset(mixedStrides, rank - 2);
     // Populate payload.
     Value payLoadAsI64 =
         vector::BitCastOp::create(rewriter, loc, payloadI64Ty, payload);
@@ -343,7 +344,7 @@ class LoadStorePrefetchNdToXeVMPattern : public OpConversionPattern<OpType> {
     // Get address space from tensor descriptor memory space.
     auto ptrTypeLLVM = LLVM::LLVMPointerType::get(
         ctxt, getNumericXeVMAddrSpace(tdescTy.getMemorySpace()));
-    if (tileRank == 2) {
+    if (tileRank >= 2) {
       // Compute element byte size.
       Value elemByteSize = arith::ConstantIntOp::create(
           rewriter, loc, rewriter.getI32Type(), elemBitSize / 8);
@@ -359,14 +360,16 @@ class LoadStorePrefetchNdToXeVMPattern : public OpConversionPattern<OpType> {
           rewriter, loc, tdesc, static_cast<int>(NdTdescOffset::BaseShapeH));
       Value basePitch = vector::ExtractOp::create(
           rewriter, loc, tdesc, static_cast<int>(NdTdescOffset::BasePitch));
-      // Offsets are provided by the op.
-      // convert them to i32.
-      Value offsetW =
-          getValueOrCreateConstantIntOp(rewriter, loc, mixedOffsets[1]);
+
+      // For rank > 2, leading (batch) dim offsets should be 0 after unrolling
+      // (batch is baked into the base pointer via memref.subview during
+      // blocking). Use only the last 2 offsets for the 2D block operation.
+      Value offsetW = getValueOrCreateConstantIntOp(rewriter, loc,
+                                                    mixedOffsets[tileRank - 1]);
       offsetW = getValueOrCreateCastToIndexLike(rewriter, loc,
                                                 rewriter.getI32Type(), offsetW);
-      Value offsetH =
-          getValueOrCreateConstantIntOp(rewriter, loc, mixedOffsets[0]);
+      Value offsetH = getValueOrCreateConstantIntOp(rewriter, loc,
+                                                    mixedOffsets[tileRank - 2]);
       offsetH = getValueOrCreateCastToIndexLike(rewriter, loc,
                                                 rewriter.getI32Type(), offsetH);
       // Convert base pointer (i64) to LLVM pointer type.
@@ -393,8 +396,8 @@ class LoadStorePrefetchNdToXeVMPattern : public OpConversionPattern<OpType> {
         offsetW =
             arith::ShRSIOp::create(rewriter, loc, offsetW, wScaleFactorValLog2);
       }
-      // Get tile height from the tensor descriptor type.
-      auto tileH = tdescTy.getDimSize(0);
+      // Get tile height from the tensor descriptor type (second-to-last dim).
+      auto tileH = tdescTy.getDimSize(tileRank - 2);
       // Get vblocks from the tensor descriptor type.
       int32_t vblocks = tdescTy.getArrayLength();
       if constexpr (std::is_same_v<OpType, xegpu::StoreNdOp>) {
diff --git a/mlir/lib/Dialect/XeGPU/IR/XeGPUOps.cpp b/mlir/lib/Dialect/XeGPU/IR/XeGPUOps.cpp
index 56d340141ee8f..aecdc39d8351e 100644
--- a/mlir/lib/Dialect/XeGPU/IR/XeGPUOps.cpp
+++ b/mlir/lib/Dialect/XeGPU/IR/XeGPUOps.cpp
@@ -344,9 +344,6 @@ LogicalResult LoadNdOp::verify() {
   auto tdescTy = getTensorDescType();
   auto valueTy = getType();
 
-  if (tdescTy.getRank() > 2)
-    return emitOpError("Expects a 1D or 2D TensorDesc.\n");
-
   if (!valueTy)
     return emitOpError("Invalid result, it should be a VectorType.\n");
 
@@ -466,9 +463,6 @@ LogicalResult StoreNdOp::verify() {
   auto dstTy = getTensorDescType(); // Tile
   auto valTy = getValueType();      // Vector
 
-  if (dstTy.getRank() > 2)
-    return emitOpError("Expects a 1D or 2D TensorDesc.\n");
-
   if (!valTy)
     return emitOpError("Expecting a VectorType result.\n");
 
@@ -720,31 +714,78 @@ static LogicalResult verifyDpasDimensions(Operation *op,
   if (aRank == 1 && bRank == 1 && resRank == 1)
     return success();
 
-  // Validate A and B are 2D
-  if (aRank != 2)
-    return op->emitOpError("A operand must be a 2D vector.");
-  if (bRank < 2 || bRank > 3)
-    return op->emitOpError("B operand must be a 2D or 3D vector.");
-  if (resRank != 2)
-    return op->emitOpError("Result must be a 2D vector.");
+  // A must be at least 2D, B must be 2D or 3D (innermost dims), result at
+  // least 2D.
+  if (aRank < 2)
+    return op->emitOpError("A operand must be at least a 2D vector.");
+  if (bRank < 2)
+    return op->emitOpError("B operand must be at least a 2D vector.");
+  if (resRank < 2)
+    return op->emitOpError("Result must be at least a 2D vector.");
+
+  // Determine batch dimensions. For A[batch..., M, K], B[batch..., K, N] (or
+  // B[batch..., K/vnni, N, vnni]), result[batch..., M, N].
+  // B may have one extra trailing dim for VNNI packing (3D innermost).
+  // Determine how many trailing dims are the "core" matmul dims.
+  // A core dims: last 2 (M, K)
+  // B core dims: last 2 (K, N) or last 3 (K/vnni, N, vnni) for packed
+  // Result core dims: last 2 (M, N)
+  int64_t aBatchRank = aRank - 2;
+  int64_t resBatchRank = resRank - 2;
+
+  // B can have an extra trailing dim for VNNI packing. Determine B's batch
+  // rank: if bRank > aRank, the extra dim is the VNNI packing dim.
+  bool bPacked = (bRank == aRank + 1);
+  int64_t bBatchRank = bPacked ? bRank - 3 : bRank - 2;
+
+  // Batch ranks must match across A, B, and result.
+  if (aBatchRank != bBatchRank || aBatchRank != resBatchRank)
+    return op->emitOpError("Batch dimension rank mismatch among A, B, and "
+                           "result.");
+
+  // Verify batch dimensions match.
+  for (int64_t i = 0; i < aBatchRank; ++i) {
+    if (aShape[i] != resShape[i])
+      return op->emitOpError("Batch dimension mismatch at dim ")
+             << i << ": A has " << aShape[i] << " but result has "
+             << resShape[i] << ".";
+    if (aShape[i] != bShape[i])
+      return op->emitOpError("Batch dimension mismatch at dim ")
+             << i << ": A has " << aShape[i] << " but B has " << bShape[i]
+             << ".";
+  }
 
-  // Calculate effective K dimension for B (handle 3D packed case)
-  int64_t bK = bRank == 3 ? bShape[0] * bShape[2] : bShape[0];
+  // Now verify the core matmul dimensions (last 2 of A and result, last 2 or
+  // 3 of B).
+  int64_t aM = aShape[aBatchRank];
+  int64_t aK = aShape[aBatchRank + 1];
+  int64_t resM = resShape[resBatchRank];
+  int64_t resN = resShape[resBatchRank + 1];
+
+  // Calculate effective K dimension for B (handle packed case)
+  int64_t bK, bN;
+  if (bPacked) {
+    bK = bShape[bBatchRank] * bShape[bBatchRank + 2];
+    bN = bShape[bBatchRank + 1];
+  } else {
+    bK = bShape[bBatchRank];
+    bN = bShape[bBatchRank + 1];
+  }
 
   // Verify K dimension match between A and B
-  if (bK != aShape[1])
+  if (bK != aK)
     return op->emitOpError("K-dimension mismatch: A has K=")
-           << aShape[1] << " but B has K=" << bK << ".";
+           << aK << " but B has K=" << bK << ".";
 
   // Verify M dimension match between A and result
-  if (aShape[0] != resShape[0])
+  if (aM != resM)
     return op->emitOpError("M-dimension mismatch: A has M=")
-           << aShape[0] << " but result has M=" << resShape[0] << ".";
+           << aM << " but result has M=" << resM << ".";
 
   // Verify N dimension match between B and result
-  if (bShape[1] != resShape[1])
+  if (bN != resN)
     return op->emitOpError("N-dimension mismatch: B has N=")
-           << bShape[1] << " but result has N=" << resShape[1] << ".";
+           << bN << " but result has N=" << resN << ".";
 
   return success();
 }
@@ -907,6 +948,9 @@ LogicalResult DpasMxOp::verify() {
   if (failed(verifyDpasDimensions(*this, aShape, bShape, resShape)))
     return failure();
 
+  // Determine batch rank from A operand.
+  int64_t aBatchRank = aShape.size() - 2;
+
   // Validate scale_a if present
   if (getScaleA()) {
     auto scaleAVecType = dyn_cast<VectorType>(getScaleAType());
@@ -914,19 +958,20 @@ LogicalResult DpasMxOp::verify() {
     if (scaleAVecType && scaleAVecType.getRank() > 1) {
       auto scaleAShape = scaleAVecType.getShape();
 
-      if (scaleAVecType.getRank() != 2)
-        return emitOpError("Scale A must be a 2D vector when not a scalar.");
+      if (scaleAVecType.getRank() < 2)
+        return emitOpError("Scale A must be at least a 2D vector when not a "
+                           "scalar.");
 
       // Verify layout distributability for scale_a
       if (failed(verifyLayoutDistributable(*this, getLayoutAScale(),
                                            scaleAShape, "ScaleA")))
         return failure();
 
-      // Validate M dimension: scale_a[0] must match a[0]
-      if (scaleAShape[0] != aShape[0])
+      // Validate M dimension: scale_a's M must match A's M (last-1 dim)
+      if (scaleAShape[scaleAShape.size() - 2] != aShape[aBatchRank])
         return emitOpError("Scale A M dimension [")
-               << scaleAShape[0] << "] must match A M dimension [" << aShape[0]
-               << "].";
+               << scaleAShape[scaleAShape.size() - 2]
+               << "] must match A M dimension [" << aShape[aBatchRank] << "].";
     }
   }
 
@@ -937,19 +982,21 @@ LogicalResult DpasMxOp::verify() {
     if (scaleBVecType && scaleBVecType.getRank() > 1) {
       auto scaleBShape = scaleBVecType.getShape();
 
-      if (scaleBVecType.getRank() != 2)
-        return emitOpError("Scale B must be a 2D vector when not a scalar.");
+      if (scaleBVecType.getRank() < 2)
+        return emitOpError("Scale B must be at least a 2D vector when not a "
+                           "scalar.");
 
       // Verify layout distributability for scale_b
       if (failed(verifyLayoutDistributable(*this, getLayoutBScale(),
                                            scaleBShape, "ScaleB")))
         return failure();
 
-      // Validate N dimension: scale_b[1] must match b[1]
-      if (scaleBShape[1] != bShape[1])
+      // Validate N dimension: scale_b's N (last dim) must match B's N (last
+      // dim)
+      if (scaleBShape.back() != bShape.back())
         return emitOpError("Scale B N dimension [")
-               << scaleBShape[1] << "] must match B N dimension [" << bShape[1]
-               << "].";
+               << scaleBShape.back() << "] must match B N dimension ["
+               << bShape.back() << "].";
     }
   }
 
@@ -964,12 +1011,12 @@ LogicalResult DpasMxOp::verify() {
       auto scaleAShape = scaleAVecType.getShape();
       auto scaleBShape = scaleBVecType.getShape();
 
-      // Validate scale K dimension compatibility: scale_a[1] must match
-      // scale_b[0]
-      if (scaleAShape[1] != scaleBShape[0])
+      // Validate scale K dimension compatibility: scale_a's last dim must
+      // match scale_b's second-to-last dim
+      if (scaleAShape.back() != scaleBShape[scaleBShape.size() - 2])
         return emitOpError("Scale K dimension mismatch: scale_a has K=")
-               << scaleAShape[1] << " but scale_b has K=" << scaleBShape[0]
-               << ".";
+               << scaleAShape.back() << " but scale_b has K="
+               << scaleBShape[scaleBShape.size() - 2] << ".";
     }
   }
 
diff --git a/mlir/lib/Dialect/XeGPU/Transforms/XeGPUBlocking.cpp b/mlir/lib/Dialect/XeGPU/Transforms/XeGPUBlocking.cpp
index 4e3695adc4897..e97579885bdd7 100644
--- a/mlir/lib/Dialect/XeGPU/Transforms/XeGPUBlocking.cpp
+++ b/mlir/lib/Dialect/XeGPU/Transforms/XeGPUBlocking.cpp
@@ -44,36 +44,57 @@ resolveUnrealizedConversionCastOp(UnrealizedConversionCastOp castOp) {
   ValueRange inputs = castOp.getInputs();
   ValueRange outputs = castOp.getOutputs();
 
-  auto hasIdenticalVectorTypes = [](ValueRange values) {
+  auto hasIdenticalVectorOrTdescTypes = [](ValueRange values) {
     auto types = values.getTypes();
     return llvm::all_of(types, [&](Type type) {
-      return isa<VectorType>(type) && type == types.front();
+      return (isa<VectorType>(type) || isa<xegpu::TensorDescType>(type)) &&
+             type == types.front();
     });
   };
 
   // We only interest in the case where all inputs and outputs have the
   // identical VectorTypes
-  if (!hasIdenticalVectorTypes(inputs) || !hasIdenticalVectorTypes(outputs)) {
+  if (!hasIdenticalVectorOrTdescTypes(inputs) ||
+      !hasIdenticalVectorOrTdescTypes(outputs)) {
     LDBG() << "skip unrealized conversion cast op not emulating pack/unpack.";
     return;
   }
 
   VectorType outputTy = dyn_cast<VectorType>(outputs[0].getType());
-  OpBuilder builder(castOp);
-  if (inputs.size() > 1 && outputs.size() == 1) {
-    // the castOp is emulating an unpack op
-    ArrayRef<int64_t> shape = outputTy.getShape();
-    Value result = xegpu::createVectorWithShapeFromValues(
-        builder, castOp.getLoc(), inputs, shape);
-    castOp->replaceAllUsesWith(ValueRange(result));
-    castOp->erase();
-  } else if (castOp.getNumResults() > 1 && castOp.getNumOperands() == 1) {
-    // the castOp is emulating a pack op
-    ArrayRef<int64_t> tileShape = outputTy.getShape();
-    SmallVector<Value> results = xegpu::extractVectorsWithShapeFromValue(
-        builder, castOp.getLoc(), inputs[0], tileShape);
-    castOp->replaceAllUsesWith(results);
-    castOp->erase();
+  if (outputTy) {
+    OpBuilder builder(castOp);
+    if (inputs.size() > 1 && outputs.size() == 1) {
+      // the castOp is emulating an unpack op
+      ArrayRef<int64_t> shape = outputTy.getShape();
+      Value result = xegpu::createVectorWithShapeFromValues(
+          builder, castOp.getLoc(), inputs, shape);
+      castOp->replaceAllUsesWith(ValueRange(result));
+      castOp->erase();
+    } else if (castOp.getNumResults() > 1 && castOp.getNumOperands() == 1) {
+      // the castOp is emulating a pack op
+      ArrayRef<int64_t> tileShape = outputTy.getShape();
+      SmallVector<Value> results = xegpu::extractVectorsWithShapeFromValue(
+          builder, castOp.getLoc(), inputs[0], tileShape);
+      castOp->replaceAllUsesWith(results);
+      castOp->erase();
+    }
+  } else {
+    if (castOp.getNumResults() > 1 && castOp.getNumOperands() == 1) {
+      if (auto prevCastOp =
+              inputs[0].getDefiningOp<UnrealizedConversionCastOp>()) {
+        if (prevCastOp.getNumResults() == 1 &&
+            prevCastOp.getNumOperands() > 1 &&
+            prevCastOp.getOutputs()[0].getType() ==
+                castOp.getInputs()[0].getType()) {
+          castOp->replaceAllUsesWith(prevCastOp.getInputs());
+          castOp->erase();
+          prevCastOp->erase();
+          return;
+        }
+      }
+    }
+
+    LDBG() << "skip unrealized conversion cast op not emulating pack/unpack.";
   }
 }
 
@@ -170,11 +191,24 @@ XeGPUBlockingPass::getTileShape(Operation *op) const {
     std::optional<SmallVector<int64_t>> bTile =
         getTileShape(op->getOpOperand(1));
 
-    if (!aTile || aTile->size() != 2 || !bTile || bTile->size() != 2)
+    if (!aTile || aTile->size() < 2 || !bTile || bTile->size() < 2)
+      return std::nullopt;
+
+    // Both must have the same number of batch dimensions.
+    int64_t aBatchRank = aTile->size() - 2;
+    int64_t bBatchRank = bTile->size() - 2;
+    if (aBatchRank != bBatchRank)
       return std::nullopt;
 
-    // semantic check for A and B
-    if ((*aTile)[1] != (*bTile)[0])
+    // Batch dimensions must match.
+    for (int64_t i = 0; i < aBatchRank; ++i) {
+      if ((*aTile)[i] != (*bTile)[i])
+        return std::nullopt;
+    }
+
+    // semantic check for A and B: K dimension must match
+    // A[..., M, K] x B[..., K, N]
+    if ((*aTile).back() != (*bTile)[bBatchRank])
       return std::nullopt;
 
     return std::make_pair(*aTile, *bTile);
@@ -189,8 +223,15 @@ XeGPUBlockingPass::getTileShape(Operation *op) const {
 
     std::optional<SmallVector<int64_t>> cTile =
         getTileShape(op->getOpOperand(cOperandIdx));
-    int64_t expectedCTile[2] = {aTile[0], bTile[1]};
-    if (!cTile || !llvm::equal(*cTile, expectedCTile))
+    if (!cTile)
+      return false;
+    // Expected C tile: batch dims from A + [M, N]
+    int64_t aBatchRank = aTile.size() - 2;
+    SmallVector<int64_t> expectedCTile(aTile.begin(),
+                                       aTile.begin() + aBatchRank);
+    expectedCTile.push_back(aTile[aBatchRank]); // M from A
+    expectedCTile.push_back(bTile.back());      // N from B
+    if (!llvm::equal(*cTile, expectedCTile))
       return false;
     return true;
   };
@@ -202,16 +243,18 @@ XeGPUBlockingPass::getTileShape(Operation *op) const {
     std::optional<SmallVector<int64_t>> aScaleTile =
         getTileShape(op->getOpOperand(scaleAOperandIdx));
 
-    if (!aScaleTile || aScaleTile->size() != 2)
+    if (!aScaleTile || aScaleTile->size() < 2)
       return std::nullopt;
 
-    // Validate scale_a tile: [M_tile, K_scale]
-    // M dimension must match A's M dimension
-    if ((*aScaleTile)[0] != aTile[0])
+    // Validate scale_a tile: [batch..., M_tile, K_scale]
+    // M dimension (second-to-last) must match A's M dimension
+    int64_t scaleRank = aScaleTile->size();
+    int64_t aBatchRank = aTile.size() - 2;
+    if ((*aScaleTile)[scaleRank - 2] != aTile[aBatchRank])
       return std::nullopt;
 
-    // Return the K scale factor
-    return (*aScaleTile)[1];
+    // Return the K scale factor (last dim)
+    return (*aScaleTile)[scaleRank - 1];
   };
 
   // Helper lambda to validate scale B tile for DpasMxOp
@@ -221,16 +264,17 @@ XeGPUBlockingPass::getTileShape(Operation *op) const {
     std::optional<SmallVector<int64_t>> bScaleTile =
         getTileShape(op->getOpOperand(scaleBOperandIdx));
 
-    if (!bScaleTile || bScaleTile->size() != 2)
+    if (!bScaleTile || bScaleTile->size() < 2)
       return std::nullopt;
 
-    // Validate scale_b tile: [K_scale, N_tile]
-    // N dimension must match B's N dimension
-    if ((*bScaleTile)[1] != bTile[1])
+    // Validate scale_b tile: [batch..., K_scale, N_tile]
+    // N dimension (last) must match B's N dimension (last)
+    if (bScaleTile->back() != bTile.back())
       return std::nullopt;
 
-    // Return the K scale factor
-    return (*bScaleTile)[0];
+    // Return the K scale factor (second-to-last dim)
+    int64_t scaleRank = bScaleTile->size();
+    return (*bScaleTile)[scaleRank - 2];
   };
 
   if (isa<xegpu::DpasOp>(op)) {
@@ -244,7 +288,13 @@ XeGPUBlockingPass::getTileShape(Operation *op) const {
     if (!validateCTile(op, 2, aTile, bTile))
       return std::nullopt;
 
-    return SmallVector<int64_t>({aTile[0], aTile[1], bTile[1]});
+    // Return [batch..., M, K, N] as the target shape for unrolling.
+    int64_t aBatchRank = aTile.size() - 2;
+    SmallVector<int64_t> tileShape(aTile.begin(), aTile.begin() + aBatchRank);
+    tileShape.push_back(aTile[aBatchRank]);     // M
+    tileShape.push_back(aTile[aBatchRank + 1]); // K
+    tileShape.push_back(bTile.back());          // N
+    return tileShape;
   }
 
   if (auto dpasMxOp = dyn_cast<xegpu::DpasMxOp>(op)) {
@@ -292,7 +342,14 @@ XeGPUBlockingPass::getTileShape(Operation *op) const {
       kScaleFactor = *scaleBFactor;
     }
 
-    return SmallVector<int64_t>({aTile[0], aTile[1], bTile[1], kScaleFactor});
+    // Return [batch..., M, K, N, S] as the target shape for unrolling.
+    int64_t aBatchRank = aTile.size() - 2;
+    SmallVector<int64_t> tileShape(aTile.begin(), aTile.begin() + aBatchRank);
+    tileShape.push_back(aTile[aBatchRank]);     // M
+    tileShape.push_back(aTile[aBatchRank + 1]); // K
+    tileShape.push_back(bTile.back());          // N
+    tileShape.push_back(kScaleFactor);          // S
+    return tileShape;
   }
 
   if (OpTrait::hasElementwiseMappableTraits(op) && op->getNumResults() == 1)
@@ -432,24 +489,25 @@ void XeGPUBlockingPass::runOnOperation() {
 
   options.setNativeShapeFn([&](Operation *op) { return getTileShape(op); });
 
-  options.setUnrolledTypesFn([&](ShapedType type, ArrayRef<int64_t> tileShape,
-                                 bool returnSingleType = false) {
+  options.setUnrolledTypesFn([&](ShapedType type, ArrayRef<int64_t> tileShape) {
     Type elemTy = type.getElementType();
-    Type newTy;
 
     if (auto tdescTy = dyn_cast<xegpu::TensorDescType>(type)) {
 
       Attribute encoding = tdescTy.getEncoding();
 
-      newTy =
+      xegpu::TensorDescType newTy =
           xegpu::TensorDescType::get(ctx, tileShape, elemTy, encoding,
                                      tdescTy.getLayoutAttr().dropInstData());
-    } else {
-      newTy = VectorType::get(tileShape, elemTy);
+      // compute the product of batch (higher) dimensions
+      ArrayRef<int64_t> shape = type.getShape();
+      int64_t batchCount =
+          shape.size() > 2 ? computeProduct(shape.drop_back(2)) : 1;
+      return SmallVector<Type>(batchCount, newTy);
     }
+    Type newTy;
+    newTy = VectorType::get(tileShape, elemTy);
 
-    if (returnSingleType)
-      return SmallVector<Type>{newTy};
     std::optional<SmallVector<int64_t>> ratio =
         computeShapeRatio(type.getShape(), tileShape);
     assert(ratio && "The shape of the type must be a multiple of tileShape.");
diff --git a/mlir/lib/Dialect/XeGPU/Transforms/XeGPULayoutImpl.cpp b/mlir/lib/Dialect/XeGPU/Transforms/XeGPULayoutImpl.cpp
index 2e63439321988..501c65456cc5e 100644
--- a/mlir/lib/Dialect/XeGPU/Transforms/XeGPULayoutImpl.cpp
+++ b/mlir/lib/Dialect/XeGPU/Transforms/XeGPULayoutImpl.cpp
@@ -1387,9 +1387,9 @@ template <typename RankedTy>
 static xegpu::LayoutAttr getDefaultLaneLayout2DBlockIo(
     RankedTy ty, const xegpu::uArch::uArch *uArch,
     std::optional<unsigned> packingSize = std::nullopt, bool vnni = false) {
-  // Expecting a 1D or 2D vector.
-  assert(((ty.getRank() == 1 && !vnni) || ty.getRank() == 2) &&
-         "Expected 1D non-vnni or 2D vector.");
+  // Expecting at least 1D vector. For rank > 2, leading dims are batch dims.
+  assert(((ty.getRank() >= 1 && !vnni) || ty.getRank() >= 2) &&
+         "Expected at least 1D non-vnni or 2D vector.");
   // Expecting int or float element type.
   assert(ty.getElementType().isIntOrFloat() &&
          "Expected int or float element type.");
@@ -1463,11 +1463,13 @@ getDpasInstDataVectors(VectorType aTy, VectorType bTy, VectorType cdTy,
         dyn_cast<xegpu::uArch::SubgroupMatrixMultiplyAcc>(uArch->getInstruction(
             xegpu::uArch::InstructionKind::SubgroupMatrixMultiplyAcc));
 
-  const unsigned dataALen = aTy.getShape().front();
+  // M dimension is the second-to-last dim of A (handles batch dims).
+  const unsigned dataALen = aTy.getShape()[aTy.getRank() - 2];
   auto supportedALen = uArchInstruction->getSupportedM(aTy.getElementType());
   const int maxALen =
       xegpu::getLargestDivisor(dataALen, ArrayRef<unsigned>(supportedALen));
 
+  // N dimension is the last dim of B.
   const unsigned dataBLen = bTy.getShape().back();
   auto supportedBLen = uArchInstruction->getSupportedN(bTy.getElementType());
   const int maxBLen =
@@ -1656,7 +1658,7 @@ createScaleLayout(mlir::MLIRContext *context, VectorType matrixTy,
               xegpu::uArch::InstructionKind::SubgroupScaledMatrixMultiplyAcc));
 
   int64_t rank = matrixLayout.getRank();
-  assert(rank == 2 && "dpas layouts must be two dimensions");
+  assert(rank >= 2 && "dpas layouts must be at least two dimensions");
 
   SmallVector<int64_t> sgLayout = matrixLayout.getEffectiveSgLayoutAsInt();
   SmallVector<int64_t> sgData = matrixLayout.getEffectiveSgDataAsInt();
diff --git a/mlir/lib/Dialect/XeGPU/Transforms/XeGPUPropagateLayout.cpp b/mlir/lib/Dialect/XeGPU/Transforms/XeGPUPropagateLayout.cpp
index c504ad3342d01..b19c47dd14d28 100644
--- a/mlir/lib/Dialect/XeGPU/Transforms/XeGPUPropagateLayout.cpp
+++ b/mlir/lib/Dialect/XeGPU/Transforms/XeGPUPropagateLayout.cpp
@@ -266,26 +266,31 @@ struct LayoutInfoLattice : public Lattice<LayoutInfo> {
 /// Helper Function to get the default layout for uniform values like constants.
 /// For 1D vector, lane_layout is [subgroupSize] and lane_data is [1].
 /// For 2D vector, lane_layout is [1, subgroupSize] and lane_data is [1, 1].
+/// For ND vector (N>2), leading dims get unit lane_layout and lane_data.
 static LayoutInfo getDefaultSIMTLayoutInfo(mlir::MLIRContext *ctx,
                                            unsigned rank,
                                            const xegpu::uArch::uArch *uArch) {
-  assert((rank == 1 || rank == 2) && "Expected 1D or 2D vector.");
+  assert(rank >= 1 && "Expected at least 1D vector.");
   if (rank == 1) {
     return LayoutInfo(
         xegpu::LayoutAttr::get(ctx, {uArch->getSubgroupSize()}, {1}));
   }
-  return LayoutInfo(
-      xegpu::LayoutAttr::get(ctx, {1, uArch->getSubgroupSize()}, {1, 1}));
+  // For rank >= 2, lane_layout is [1, ..., 1, subgroupSize] and
+  // lane_data is [1, ..., 1, 1].
+  SmallVector<int32_t> laneLayout(rank, 1);
+  SmallVector<int32_t> laneData(rank, 1);
+  laneLayout[rank - 1] = uArch->getSubgroupSize();
+  return LayoutInfo(xegpu::LayoutAttr::get(ctx, laneLayout, laneData));
 }
 
 /// Helper to get the default layout for 2D block operations.
+/// For ND (N>2) types, leading dimensions get unit layout/data values.
 template <typename Ty>
 static LayoutInfo getSIMTLayoutInfoBlockIO(Ty ty,
                                            const xegpu::uArch::uArch *uArch,
                                            unsigned packingSize) {
-  // Expecting a 1D or 2D vector.
-  assert((ty.getRank() == 1 || ty.getRank() == 2) &&
-         "Expected 1D or 2D vector.");
+  // Expecting at least 1D.
+  assert(ty.getRank() >= 1 && "Expected at least 1D vector.");
   // Expecting int or float element type.
   assert(ty.getElementType().isIntOrFloat() &&
          "Expected int or float element type.");
@@ -295,8 +300,14 @@ static LayoutInfo getSIMTLayoutInfoBlockIO(Ty ty,
   // Packing factor is determined by the element type bitwidth.
   unsigned bitwidth = ty.getElementType().getIntOrFloatBitWidth();
   int packingFactor = bitwidth < packingSize ? packingSize / bitwidth : 1;
-  return LayoutInfo(xegpu::LayoutAttr::get(
-      ty.getContext(), {1, uArch->getSubgroupSize()}, {1, packingFactor}));
+  // For rank >= 2, distribute along the last dimension with leading units.
+  unsigned rank = ty.getRank();
+  SmallVector<int32_t> laneLayout(rank, 1);
+  SmallVector<int32_t> laneData(rank, 1);
+  laneLayout[rank - 1] = uArch->getSubgroupSize();
+  laneData[rank - 1] = packingFactor;
+  return LayoutInfo(xegpu::LayoutAttr::get(ty.getContext(), laneLayout,
+                                           laneData));
 }
 
 //===----------------------------------------------------------------------===//
@@ -983,7 +994,11 @@ void LayoutInfoPropagation::visitStoreNdOp(
       if (instHeight == -1)
         store.emitWarning(
             "No suitable instruction multiple found for the given shape.");
-      instData = {instHeight, instWidth};
+      // For rank > 2, prepend unit dims for batch dimensions.
+      for (int64_t d = 0; d < dataTy.getRank() - 2; ++d)
+        instData.push_back(1);
+      instData.push_back(instHeight);
+      instData.push_back(instWidth);
     }
 
     if (layoutKind == xegpu::LayoutKind::InstData)
diff --git a/mlir/lib/Dialect/XeGPU/Transforms/XeGPUUnroll.cpp b/mlir/lib/Dialect/XeGPU/Transforms/XeGPUUnroll.cpp
index aab36b79845e4..95ad1808dac07 100644
--- a/mlir/lib/Dialect/XeGPU/Transforms/XeGPUUnroll.cpp
+++ b/mlir/lib/Dialect/XeGPU/Transforms/XeGPUUnroll.cpp
@@ -13,6 +13,7 @@
 //===----------------------------------------------------------------------===//
 
 #include "mlir/Dialect/Arith/IR/Arith.h"
+#include "mlir/Dialect/MemRef/IR/MemRef.h"
 #include "mlir/Dialect/Utils/IndexingUtils.h"
 #include "mlir/Dialect/Vector/IR/VectorOps.h"
 #include "mlir/Dialect/XeGPU/IR/XeGPU.h"
@@ -59,9 +60,8 @@ struct UnrollPattern : public OpRewritePattern<SourceOp> {
   }
 
   SmallVector<Type> getUnrolledTypes(ShapedType type,
-                                     ArrayRef<int64_t> tileShape,
-                                     bool returnSingleType = false) const {
-    return options.getUnrolledTypes(type, tileShape, returnSingleType);
+                                     ArrayRef<int64_t> tileShape) const {
+    return options.getUnrolledTypes(type, tileShape);
   }
 
   /// Emulate the the unpack behavior using insert_strided_slice for VectorType
@@ -195,16 +195,59 @@ struct UnrollCreateNdOp : public UnrollPattern<xegpu::CreateNdDescOp> {
     if (!targetShape)
       return failure();
 
+    int64_t rank = tdescTy.getRank();
+    int64_t batchRank = rank - 2;
+
+    // For rank <= 2 or non-memref source: existing single-tdesc behavior.
+    if (batchRank <= 0 || !isa<MemRefType>(op.getSourceType())) {
+      SmallVector<Value> newOps;
+      auto newTdescTy = getUnrolledTypes(tdescTy, *targetShape)[0];
+      auto newOp = xegpu::CreateNdDescOp::create(
+          rewriter, loc, newTdescTy, op.getSource(), op.getMixedSizes(),
+          op.getMixedStrides());
+      newOps.push_back(newOp);
+      Value castOp = unpack(newOps, tdescTy, *targetShape, loc, rewriter);
+      rewriter.replaceOp(op, castOp);
+      return success();
+    }
+
+    // For rank > 2 with memref source: create one tdesc per batch tile via
+    // memref.subview. Each subview slices the batch dimensions, so the
+    // resulting tdesc has the batch offset baked into its base pointer.
+    // The inner dimensions remain full-size for reuse across multiple
+    // load/store operations with different offsets.
+    ArrayRef<int64_t> shape = tdescTy.getShape();
+    SmallVector<int64_t> batchBlockSize(targetShape->begin(),
+                                        targetShape->begin() + batchRank);
+    batchBlockSize.append(shape.begin() + batchRank, shape.end());
+
+    auto newTdescTy =
+        cast<xegpu::TensorDescType>(getUnrolledTypes(tdescTy, *targetShape)[0]);
+
     SmallVector<Value> newOps;
+    for (SmallVector<int64_t> batchOffsets :
+         StaticTileOffsetRange(shape, batchBlockSize)) {
+      SmallVector<OpFoldResult> svOffsets;
+      for (int64_t off : batchOffsets)
+        svOffsets.push_back(rewriter.getIndexAttr(off));
+
+      SmallVector<OpFoldResult> svSizes;
+      for (int64_t d : batchBlockSize)
+        svSizes.push_back(rewriter.getIndexAttr(d));
+
+      SmallVector<OpFoldResult> svStrides(rank, rewriter.getIndexAttr(1));
+
+      auto subview = memref::SubViewOp::create(rewriter, loc, op.getSource(),
+                                               svOffsets, svSizes, svStrides);
+
+      auto newOp = xegpu::CreateNdDescOp::create(
+          rewriter, loc, newTdescTy,
+          cast<TypedValue<MemRefType>>(subview.getResult()));
+      newOps.push_back(newOp);
+    }
 
-    auto newTdescTy = getUnrolledTypes(tdescTy, *targetShape)[0];
-    auto newOp =
-        xegpu::CreateNdDescOp::create(rewriter, loc, newTdescTy, op.getSource(),
-                                      op.getMixedSizes(), op.getMixedStrides());
-    newOps.push_back(newOp);
     Value castOp = unpack(newOps, tdescTy, *targetShape, loc, rewriter);
     rewriter.replaceOp(op, castOp);
-
     return success();
   }
 };
@@ -224,22 +267,58 @@ struct UnrollPrefetchNdOp : public UnrollPattern<xegpu::PrefetchNdOp> {
     if (layout)
       layout = layout.dropInstData();
 
-    SmallVector<Type> convertedTdescTypes =
-        getUnrolledTypes(tdescTy, *targetShape, /*returnSingleType*/ true);
+    int64_t rank = tdescTy.getRank();
+    int64_t batchRank = rank - 2;
 
-    SmallVector<Value> convertedTdesc = pack(
-        op.getTensorDesc(), convertedTdescTypes, *targetShape, loc, rewriter);
+    if (batchRank <= 0) {
+      SmallVector<Type> convertedTdescTypes =
+          getUnrolledTypes(tdescTy, *targetShape);
+      SmallVector<Value> convertedTdesc = pack(
+          op.getTensorDesc(), convertedTdescTypes, *targetShape, loc, rewriter);
 
-    auto createPrefetch = [&](SmallVector<OpFoldResult> offsets) -> Value {
-      xegpu::PrefetchNdOp::create(rewriter, loc, convertedTdesc[0], offsets,
-                                  op.getL1HintAttr(), op.getL2HintAttr(),
-                                  op.getL3HintAttr(), layout);
-      // return dummy Value to satisfy function's signature
-      return nullptr;
-    };
-
-    computeUnrolledOffsets(op.getMixedOffsets(), tdescTy, *targetShape,
-                           createPrefetch, loc, rewriter);
+      auto createPrefetch = [&](SmallVector<OpFoldResult> offsets) -> Value {
+        xegpu::PrefetchNdOp::create(rewriter, loc, convertedTdesc[0], offsets,
+                                    op.getL1HintAttr(), op.getL2HintAttr(),
+                                    op.getL3HintAttr(), layout);
+        return nullptr;
+      };
+      computeUnrolledOffsets(op.getMixedOffsets(), tdescTy, *targetShape,
+                             createPrefetch, loc, rewriter);
+    } else {
+      // Rank > 2: batch tdescs cover [batchTarget..., innerShape...].
+      // Each batch tdesc is reused for multiple inner prefetches via offsets.
+      ArrayRef<int64_t> shape = tdescTy.getShape();
+      SmallVector<int64_t> innerShape(shape.begin() + batchRank, shape.end());
+      SmallVector<int64_t> innerTarget(targetShape->begin() + batchRank,
+                                       targetShape->end());
+
+      SmallVector<Type> batchTdescTypes =
+          getUnrolledTypes(tdescTy, *targetShape);
+      SmallVector<Value> batchTdescs = pack(op.getTensorDesc(), batchTdescTypes,
+                                            *targetShape, loc, rewriter);
+
+      auto innerTdescTy = xegpu::TensorDescType::get(
+          tdescTy.getContext(), innerShape, tdescTy.getElementType(),
+          tdescTy.getEncoding(), /*layout=*/nullptr);
+
+      SmallVector<OpFoldResult> mixedOffsets = op.getMixedOffsets();
+      SmallVector<OpFoldResult> innerOffsets(mixedOffsets.begin() + batchRank,
+                                             mixedOffsets.end());
+
+      for (auto [idx, batchTdesc] : llvm::enumerate(batchTdescs)) {
+        auto createPrefetch = [&](SmallVector<OpFoldResult> offsets) -> Value {
+          SmallVector<OpFoldResult> fullOffsets(batchRank,
+                                                rewriter.getIndexAttr(0));
+          fullOffsets.append(offsets.begin(), offsets.end());
+          xegpu::PrefetchNdOp::create(rewriter, loc, batchTdesc, fullOffsets,
+                                      op.getL1HintAttr(), op.getL2HintAttr(),
+                                      op.getL3HintAttr(), layout);
+          return nullptr;
+        };
+        computeUnrolledOffsets(innerOffsets, innerTdescTy, innerTarget,
+                               createPrefetch, loc, rewriter);
+      }
+    }
 
     rewriter.eraseOp(op);
     return success();
@@ -266,24 +345,65 @@ struct UnrollLoadNdOp : public UnrollPattern<xegpu::LoadNdOp> {
     Type elemTy = tdescTy.getElementType();
     VectorType newValueTy = valueTy.cloneWith(*targetShape, elemTy);
 
-    SmallVector<Type> convertedTdescTypes =
-        getUnrolledTypes(tdescTy, *targetShape, /*returnSingleType*/ true);
-
-    SmallVector<Value> convertedTdescs = pack(
-        op.getTensorDesc(), convertedTdescTypes, *targetShape, loc, rewriter);
+    int64_t rank = tdescTy.getRank();
+    int64_t batchRank = rank - 2;
     SmallVector<Value> newOps;
 
-    auto createLoad = [&](SmallVector<OpFoldResult> offsets) {
-      return xegpu::LoadNdOp::create(
-          rewriter, loc, newValueTy, convertedTdescs[0], offsets,
-          op.getPackedAttr(), op.getTransposeAttr(), op.getL1HintAttr(),
-          op.getL2HintAttr(), op.getL3HintAttr(), layout);
-    };
-    newOps = computeUnrolledOffsets(op.getMixedOffsets(), tdescTy, *targetShape,
-                                    createLoad, loc, rewriter);
+    if (batchRank <= 0) {
+      // Rank <= 2: original behavior with single tdesc.
+      SmallVector<Type> convertedTdescTypes =
+          getUnrolledTypes(tdescTy, *targetShape);
+      SmallVector<Value> convertedTdescs = pack(
+          op.getTensorDesc(), convertedTdescTypes, *targetShape, loc, rewriter);
+
+      auto createLoad = [&](SmallVector<OpFoldResult> offsets) {
+        return xegpu::LoadNdOp::create(
+            rewriter, loc, newValueTy, convertedTdescs[0], offsets,
+            op.getPackedAttr(), op.getTransposeAttr(), op.getL1HintAttr(),
+            op.getL2HintAttr(), op.getL3HintAttr(), layout);
+      };
+      newOps = computeUnrolledOffsets(op.getMixedOffsets(), tdescTy,
+                                      *targetShape, createLoad, loc, rewriter);
+    } else {
+      // Rank > 2: batch tdescs cover [batchTarget..., innerShape...].
+      // Each batch tdesc is reused for multiple inner loads via offsets.
+      ArrayRef<int64_t> shape = tdescTy.getShape();
+      SmallVector<int64_t> innerShape(shape.begin() + batchRank, shape.end());
+      SmallVector<int64_t> innerTarget(targetShape->begin() + batchRank,
+                                       targetShape->end());
+
+      SmallVector<Type> batchTdescTypes =
+          getUnrolledTypes(tdescTy, *targetShape);
+      SmallVector<Value> batchTdescs = pack(op.getTensorDesc(), batchTdescTypes,
+                                            *targetShape, loc, rewriter);
+
+      // For each batch tdesc, pack it down to a single targetShape-sized
+      // tdesc and iterate with inner offsets (reusing the same tdesc).
+      auto innerTdescTy = xegpu::TensorDescType::get(
+          tdescTy.getContext(), innerShape, elemTy, tdescTy.getEncoding(),
+          /*layout=*/nullptr);
+
+      SmallVector<OpFoldResult> mixedOffsets = op.getMixedOffsets();
+      SmallVector<OpFoldResult> innerOffsets(mixedOffsets.begin() + batchRank,
+                                             mixedOffsets.end());
+
+      for (auto [idx, batchTdesc] : llvm::enumerate(batchTdescs)) {
+        auto createLoad = [&](SmallVector<OpFoldResult> offsets) {
+          SmallVector<OpFoldResult> fullOffsets(batchRank,
+                                                rewriter.getIndexAttr(0));
+          fullOffsets.append(offsets.begin(), offsets.end());
+          return xegpu::LoadNdOp::create(
+              rewriter, loc, newValueTy, batchTdesc, fullOffsets,
+              op.getPackedAttr(), op.getTransposeAttr(), op.getL1HintAttr(),
+              op.getL2HintAttr(), op.getL3HintAttr(), layout);
+        };
+        auto batchLoads = computeUnrolledOffsets(
+            innerOffsets, innerTdescTy, innerTarget, createLoad, loc, rewriter);
+        newOps.append(batchLoads.begin(), batchLoads.end());
+      }
+    }
 
     Value castOp = unpack(newOps, op.getType(), *targetShape, loc, rewriter);
-
     rewriter.replaceOp(op, castOp);
     return success();
   }
@@ -307,26 +427,65 @@ struct UnrollStoreNdOp : public UnrollPattern<xegpu::StoreNdOp> {
 
     SmallVector<Type> convertedValTypes =
         getUnrolledTypes(valueTy, *targetShape);
-    SmallVector<Type> convertedTdescTypes =
-        getUnrolledTypes(tdescTy, *targetShape, /*returnSingleType*/ true);
-
-    SmallVector<Value> convertedTdescs = pack(
-        op.getTensorDesc(), convertedTdescTypes, *targetShape, loc, rewriter);
 
     SmallVector<Value> convertedValues =
         pack(op.getValue(), convertedValTypes, *targetShape, loc, rewriter);
 
+    int64_t rank = tdescTy.getRank();
+    int64_t batchRank = rank - 2;
     size_t valueIndex = 0;
-    auto createStore = [&](SmallVector<OpFoldResult> offsets) {
-      xegpu::StoreNdOp::create(rewriter, loc, convertedValues[valueIndex++],
-                               convertedTdescs[0], offsets, op.getL1HintAttr(),
-                               op.getL2HintAttr(), op.getL3HintAttr(), layout);
-      // return dummy Value to satisfy function's signature
-      return nullptr;
-    };
 
-    computeUnrolledOffsets(op.getMixedOffsets(), tdescTy, *targetShape,
-                           createStore, loc, rewriter);
+    if (batchRank <= 0) {
+      SmallVector<Type> convertedTdescTypes =
+          getUnrolledTypes(tdescTy, *targetShape);
+      SmallVector<Value> convertedTdescs = pack(
+          op.getTensorDesc(), convertedTdescTypes, *targetShape, loc, rewriter);
+
+      auto createStore = [&](SmallVector<OpFoldResult> offsets) {
+        xegpu::StoreNdOp::create(rewriter, loc, convertedValues[valueIndex++],
+                                 convertedTdescs[0], offsets,
+                                 op.getL1HintAttr(), op.getL2HintAttr(),
+                                 op.getL3HintAttr(), layout);
+        return (Value) nullptr;
+      };
+      computeUnrolledOffsets(op.getMixedOffsets(), tdescTy, *targetShape,
+                             createStore, loc, rewriter);
+    } else {
+      // Rank > 2: batch tdescs cover [batchTarget..., innerShape...].
+      // Each batch tdesc is reused for multiple inner stores via offsets.
+      ArrayRef<int64_t> shape = tdescTy.getShape();
+      SmallVector<int64_t> innerShape(shape.begin() + batchRank, shape.end());
+      SmallVector<int64_t> innerTarget(targetShape->begin() + batchRank,
+                                       targetShape->end());
+
+      SmallVector<Type> batchTdescTypes =
+          getUnrolledTypes(tdescTy, *targetShape);
+      SmallVector<Value> batchTdescs = pack(op.getTensorDesc(), batchTdescTypes,
+                                            *targetShape, loc, rewriter);
+
+      auto innerTdescTy = xegpu::TensorDescType::get(
+          tdescTy.getContext(), innerShape, tdescTy.getElementType(),
+          tdescTy.getEncoding(), /*layout=*/nullptr);
+
+      SmallVector<OpFoldResult> mixedOffsets = op.getMixedOffsets();
+      SmallVector<OpFoldResult> innerOffsets(mixedOffsets.begin() + batchRank,
+                                             mixedOffsets.end());
+
+      for (auto [idx, batchTdesc] : llvm::enumerate(batchTdescs)) {
+        auto createStore = [&](SmallVector<OpFoldResult> offsets) {
+          SmallVector<OpFoldResult> fullOffsets(batchRank,
+                                                rewriter.getIndexAttr(0));
+          fullOffsets.append(offsets.begin(), offsets.end());
+          xegpu::StoreNdOp::create(rewriter, loc, convertedValues[valueIndex++],
+                                   batchTdesc, fullOffsets, op.getL1HintAttr(),
+                                   op.getL2HintAttr(), op.getL3HintAttr(),
+                                   layout);
+          return (Value) nullptr;
+        };
+        computeUnrolledOffsets(innerOffsets, innerTdescTy, innerTarget,
+                               createStore, loc, rewriter);
+      }
+    }
 
     rewriter.eraseOp(op);
     return success();
@@ -340,15 +499,26 @@ struct UnrollDpasOp : public UnrollPattern<xegpu::DpasOp> {
     Location loc = op.getLoc();
 
     std::optional<SmallVector<int64_t>> targetShape = getTargetShape(op);
-    if (!targetShape || targetShape->size() != 3)
+    if (!targetShape || targetShape->size() < 3)
       return failure();
-    auto M = (*targetShape)[0];
-    auto K = (*targetShape)[1];
-    auto N = (*targetShape)[2];
 
-    int64_t aBlockSize[2] = {M, K};
-    int64_t bBlockSize[2] = {K, N};
-    int64_t cBlockSize[2] = {M, N};
+    // targetShape is [batch..., M, K, N]
+    int64_t tsRank = targetShape->size();
+    auto M = (*targetShape)[tsRank - 3];
+    auto K = (*targetShape)[tsRank - 2];
+    auto N = (*targetShape)[tsRank - 1];
+    ArrayRef<int64_t> batchDims(targetShape->data(), tsRank - 3);
+
+    // Build block sizes including batch dimensions.
+    SmallVector<int64_t> aBlockSize(batchDims);
+    aBlockSize.push_back(M);
+    aBlockSize.push_back(K);
+    SmallVector<int64_t> bBlockSize(batchDims);
+    bBlockSize.push_back(K);
+    bBlockSize.push_back(N);
+    SmallVector<int64_t> cBlockSize(batchDims);
+    cBlockSize.push_back(M);
+    cBlockSize.push_back(N);
 
     auto a = op.getLhs();
     auto b = op.getRhs();
@@ -371,29 +541,40 @@ struct UnrollDpasOp : public UnrollPattern<xegpu::DpasOp> {
 
     auto aShape = a.getType().getShape();
     auto bShape = b.getType().getShape();
-    int64_t mIters = aShape[0] / M;
-    int64_t kIters = aShape[1] / K;
-    int64_t nIters = bShape[1] / N;
+
+    // Compute iteration counts. Batch dims only iterate over M and N (not
+    // K-reduction), so compute batch iterations from the C block size.
+    int64_t batchRank = batchDims.size();
+    int64_t mIters = aShape[batchRank] / M;
+    int64_t kIters = aShape[batchRank + 1] / K;
+    int64_t nIters = bShape[batchRank + 1] / N;
+
+    // Compute batch iterations (product of batch dim ratios).
+    int64_t batchIters = 1;
+    for (int64_t d = 0; d < batchRank; ++d)
+      batchIters *= aShape[d] / batchDims[d];
 
     SmallVector<Value> newOps;
-    for (int64_t i = 0; i < mIters; ++i) {
-      for (int64_t j = 0; j < nIters; ++j) {
-        Value tmpC;
-        if (c)
-          tmpC = cVals[i * nIters + j];
-
-        for (int64_t k = 0; k < kIters; ++k) {
-          Value aVec = aVals[i * kIters + k];
-          Value bVec = bVals[k * nIters + j];
-          SmallVector<Value> operands({aVec, bVec});
-          if (tmpC)
-            operands.push_back(tmpC);
-
-          tmpC =
-              xegpu::DpasOp::create(rewriter, loc, vecTy, operands,
-                                    xegpu::dropInstDataOnAttrs(op->getAttrs()));
+    for (int64_t batch = 0; batch < batchIters; ++batch) {
+      for (int64_t i = 0; i < mIters; ++i) {
+        for (int64_t j = 0; j < nIters; ++j) {
+          Value tmpC;
+          if (c)
+            tmpC = cVals[batch * (mIters * nIters) + i * nIters + j];
+
+          for (int64_t k = 0; k < kIters; ++k) {
+            Value aVec = aVals[batch * (mIters * kIters) + i * kIters + k];
+            Value bVec = bVals[batch * (kIters * nIters) + k * nIters + j];
+            SmallVector<Value> operands({aVec, bVec});
+            if (tmpC)
+              operands.push_back(tmpC);
+
+            tmpC = xegpu::DpasOp::create(
+                rewriter, loc, vecTy, operands,
+                xegpu::dropInstDataOnAttrs(op->getAttrs()));
+          }
+          newOps.push_back(tmpC);
         }
-        newOps.push_back(tmpC);
       }
     }
     Value castOp = unpack(newOps, resultTy, cBlockSize, loc, rewriter);
@@ -409,18 +590,32 @@ struct UnrollDpasMxOp : public UnrollPattern<xegpu::DpasMxOp> {
     Location loc = op.getLoc();
 
     std::optional<SmallVector<int64_t>> targetShape = getTargetShape(op);
-    if (!targetShape || targetShape->size() != 4)
+    if (!targetShape || targetShape->size() < 4)
       return failure();
-    auto M = (*targetShape)[0];
-    auto K = (*targetShape)[1];
-    auto N = (*targetShape)[2];
-    auto S = (*targetShape)[3];
 
-    int64_t aBlockSize[2] = {M, K};
-    int64_t bBlockSize[2] = {K, N};
-    int64_t cBlockSize[2] = {M, N};
-    int64_t aScaleBlockSize[2] = {M, S};
-    int64_t bScaleBlockSize[2] = {S, N};
+    // targetShape is [batch..., M, K, N, S]
+    int64_t tsRank = targetShape->size();
+    auto M = (*targetShape)[tsRank - 4];
+    auto K = (*targetShape)[tsRank - 3];
+    auto N = (*targetShape)[tsRank - 2];
+    auto S = (*targetShape)[tsRank - 1];
+    ArrayRef<int64_t> batchDims(targetShape->data(), tsRank - 4);
+
+    SmallVector<int64_t> aBlockSize(batchDims);
+    aBlockSize.push_back(M);
+    aBlockSize.push_back(K);
+    SmallVector<int64_t> bBlockSize(batchDims);
+    bBlockSize.push_back(K);
+    bBlockSize.push_back(N);
+    SmallVector<int64_t> cBlockSize(batchDims);
+    cBlockSize.push_back(M);
+    cBlockSize.push_back(N);
+    SmallVector<int64_t> aScaleBlockSize(batchDims);
+    aScaleBlockSize.push_back(M);
+    aScaleBlockSize.push_back(S);
+    SmallVector<int64_t> bScaleBlockSize(batchDims);
+    bScaleBlockSize.push_back(S);
+    bScaleBlockSize.push_back(N);
 
     auto a = op.getA();
     auto b = op.getB();
@@ -445,35 +640,44 @@ struct UnrollDpasMxOp : public UnrollPattern<xegpu::DpasMxOp> {
 
     auto aShape = a.getType().getShape();
     auto bShape = b.getType().getShape();
-    int64_t mIters = aShape[0] / M;
-    int64_t kIters = aShape[1] / K;
-    int64_t nIters = bShape[1] / N;
+    int64_t batchRank = batchDims.size();
+    int64_t mIters = aShape[batchRank] / M;
+    int64_t kIters = aShape[batchRank + 1] / K;
+    int64_t nIters = bShape[batchRank + 1] / N;
+
+    int64_t batchIters = 1;
+    for (int64_t d = 0; d < batchRank; ++d)
+      batchIters *= aShape[d] / batchDims[d];
 
     SmallVector<Value> newOps;
     xegpu::DpasMxOp newDpasMxOp;
-    for (int64_t i = 0; i < mIters; ++i) {
-      for (int64_t j = 0; j < nIters; ++j) {
-        Value tmpC;
-        if (c)
-          tmpC = cVals[i * nIters + j];
-
-        for (int64_t k = 0; k < kIters; ++k) {
-          Value aVec = aVals[i * kIters + k];
-          Value bVec = bVals[k * nIters + j];
-          SmallVector<Value> operands({aVec, bVec});
-          if (tmpC)
-            operands.push_back(tmpC);
-          if (ascale)
-            operands.push_back(aScaleVals[i * kIters + k]);
-          if (bscale)
-            operands.push_back(bScaleVals[k * nIters + j]);
-
-          newDpasMxOp = xegpu::DpasMxOp::create(
-              rewriter, loc, vecTy, operands,
-              xegpu::dropInstDataOnAttrs(op->getAttrs()));
-          tmpC = newDpasMxOp.getResult();
+    for (int64_t batch = 0; batch < batchIters; ++batch) {
+      for (int64_t i = 0; i < mIters; ++i) {
+        for (int64_t j = 0; j < nIters; ++j) {
+          Value tmpC;
+          if (c)
+            tmpC = cVals[batch * (mIters * nIters) + i * nIters + j];
+
+          for (int64_t k = 0; k < kIters; ++k) {
+            Value aVec = aVals[batch * (mIters * kIters) + i * kIters + k];
+            Value bVec = bVals[batch * (kIters * nIters) + k * nIters + j];
+            SmallVector<Value> operands({aVec, bVec});
+            if (tmpC)
+              operands.push_back(tmpC);
+            if (ascale)
+              operands.push_back(
+                  aScaleVals[batch * (mIters * kIters) + i * kIters + k]);
+            if (bscale)
+              operands.push_back(
+                  bScaleVals[batch * (kIters * nIters) + k * nIters + j]);
+
+            newDpasMxOp = xegpu::DpasMxOp::create(
+                rewriter, loc, vecTy, operands,
+                xegpu::dropInstDataOnAttrs(op->getAttrs()));
+            tmpC = newDpasMxOp.getResult();
+          }
+          newOps.push_back(newDpasMxOp);
         }
-        newOps.push_back(newDpasMxOp);
       }
     }
     Value castOp = unpack(newOps, resultTy, cBlockSize, loc, rewriter);
diff --git a/mlir/lib/Dialect/XeGPU/Transforms/XeGPUWgToSgDistribute.cpp b/mlir/lib/Dialect/XeGPU/Transforms/XeGPUWgToSgDistribute.cpp
index 1aa03ebc0f376..1463dc96bcf47 100644
--- a/mlir/lib/Dialect/XeGPU/Transforms/XeGPUWgToSgDistribute.cpp
+++ b/mlir/lib/Dialect/XeGPU/Transforms/XeGPUWgToSgDistribute.cpp
@@ -304,7 +304,7 @@ struct WgToSgDpasOp : public OpConversionPattern<xegpu::DpasOp> {
                   ConversionPatternRewriter &rewriter) const override {
     Location loc = op.getLoc();
     VectorType resultTy = op.getResult().getType();
-    if (resultTy.getRank() != 2)
+    if (resultTy.getRank() < 2)
       return failure();
 
     auto layoutCd = op.getLayoutCdAttr();
@@ -328,8 +328,13 @@ struct WgToSgDpasOp : public OpConversionPattern<xegpu::DpasOp> {
             cast<VectorType>(aVec.getType()).getShape();
         ArrayRef<int64_t> bVecShape =
             cast<VectorType>(bVec.getType()).getShape();
-        VectorType resTy = VectorType::get({aVecShape[0], bVecShape[1]},
-                                           resultTy.getElementType());
+        // Build result shape: batch dims from A + [M, N] from last dims of
+        // A and B.
+        SmallVector<int64_t> resShape(aVecShape.drop_back(2));
+        resShape.push_back(aVecShape[aVecShape.size() - 2]);
+        resShape.push_back(bVecShape[bVecShape.size() - 1]);
+        VectorType resTy =
+            VectorType::get(resShape, resultTy.getElementType());
         auto newDpasOp = xegpu::DpasOp::create(rewriter, loc, resTy, operands);
         newDpasOp.setLayoutCdAttr(layoutCd.dropSgLayoutAndData());
         newDpasOp.setLayoutAAttr(layoutA.dropSgLayoutAndData());
@@ -353,7 +358,7 @@ struct WgToSgDpasMxOp : public OpConversionPattern<xegpu::DpasMxOp> {
     Location loc = op.getLoc();
     VectorType resultTy = op.getResult().getType();
 
-    if (resultTy.getRank() != 2)
+    if (resultTy.getRank() < 2)
       return failure();
 
     auto layoutCd = op.getLayoutCdAttr();
@@ -379,8 +384,12 @@ struct WgToSgDpasMxOp : public OpConversionPattern<xegpu::DpasMxOp> {
             cast<VectorType>(aVec.getType()).getShape();
         ArrayRef<int64_t> bVecShape =
             cast<VectorType>(bVec.getType()).getShape();
-        VectorType resTy = VectorType::get({aVecShape[0], bVecShape[1]},
-                                           resultTy.getElementType());
+        // Build result shape: batch dims from A + [M, N]
+        SmallVector<int64_t> resShape(aVecShape.drop_back(2));
+        resShape.push_back(aVecShape[aVecShape.size() - 2]);
+        resShape.push_back(bVecShape[bVecShape.size() - 1]);
+        VectorType resTy =
+            VectorType::get(resShape, resultTy.getElementType());
         auto newDpasMxOp = xegpu::DpasMxOp::create(
             rewriter, loc, resTy, aVec, bVec, accVal, scaleAVal, scaleBVal,
             layoutA.dropSgLayoutAndData(), layoutB.dropSgLayoutAndData(),
diff --git a/mlir/lib/Dialect/XeGPU/Utils/XeGPUUtils.cpp b/mlir/lib/Dialect/XeGPU/Utils/XeGPUUtils.cpp
index 0fb0ac6e3416d..d7fcc6b7bf046 100644
--- a/mlir/lib/Dialect/XeGPU/Utils/XeGPUUtils.cpp
+++ b/mlir/lib/Dialect/XeGPU/Utils/XeGPUUtils.cpp
@@ -73,21 +73,31 @@ FailureOr<VectorType>
 mlir::xegpu::getDistributedVectorType(VectorType originalType,
                                       xegpu::LayoutAttr layout) {
   int64_t rank = originalType.getRank();
-  // Distributed vector type is only supported for 1D, 2D and 3D vectors.
-  if (rank < 1 || rank > 3)
+  if (rank < 1)
     return failure();
   ArrayRef<int64_t> shape = originalType.getShape();
-  // arrayLength is 1 for 1D and 2D vectors, and equal to the first dimension
-  // of the 3D vector.
+  // For rank > 2, leading dimensions are treated as batch/array dimensions.
+  // Drop them and use the product as arrayLength.
   int arrayLength = 1;
-  if (rank == 3) {
-    arrayLength = shape[0];
+  while (shape.size() > 2) {
+    arrayLength *= shape[0];
     shape = shape.drop_front();
   }
+  // Drop matching leading dims from layout if the layout rank exceeds the
+  // remaining shape rank.
+  auto laneLayout = layout.getEffectiveLaneLayoutAsInt();
+  auto laneData = layout.getEffectiveLaneDataAsInt();
+  while (!laneLayout.empty() && laneLayout.size() > shape.size()) {
+    laneLayout.erase(laneLayout.begin());
+    laneData.erase(laneData.begin());
+  }
+  auto trimmedLayout = xegpu::LayoutAttr::get(
+      layout.getContext(), SmallVector<int32_t>(laneLayout.begin(), laneLayout.end()),
+      SmallVector<int32_t>(laneData.begin(), laneData.end()));
   auto helperTdescTy = xegpu::TensorDescType::get(
       shape, originalType.getElementType(), arrayLength,
       /*boundary_check=*/true,
-      /*memory_space=*/xegpu::MemorySpace::Global, layout);
+      /*memory_space=*/xegpu::MemorySpace::Global, trimmedLayout);
   return xegpu::getDistributedVectorType(helperTdescTy);
 }
 
diff --git a/mlir/test/Dialect/XeGPU/invalid.mlir b/mlir/test/Dialect/XeGPU/invalid.mlir
index 341427e7e1231..bca942d0b655e 100644
--- a/mlir/test/Dialect/XeGPU/invalid.mlir
+++ b/mlir/test/Dialect/XeGPU/invalid.mlir
@@ -111,8 +111,8 @@ func.func @load_nd_vc_4(%src: memref<24x32xf32>) {
 // -----
 func.func @subgroup_load_nd_9(%src: memref<4x8x16xf16>) {
   %1 = xegpu.create_nd_tdesc %src : memref<4x8x16xf16> -> !xegpu.tensor_desc<4x8x16xf16>
-  // expected-error at +1 {{Expects a 1D or 2D TensorDesc}}
-  %2 = xegpu.load_nd %1[0, 0, 0] <{l1_hint = #xegpu.cache_hint<cached>, l2_hint = #xegpu.cache_hint<uncached>}> : !xegpu.tensor_desc<4x8x16xf16> -> vector<4x8x16xf16>
+  // expected-error at +1 {{Result shape [4, 8, 8] is not consistent with tensor descriptor}}
+  %2 = xegpu.load_nd %1[0, 0, 0] <{l1_hint = #xegpu.cache_hint<cached>, l2_hint = #xegpu.cache_hint<uncached>}> : !xegpu.tensor_desc<4x8x16xf16> -> vector<4x8x8xf16>
   return
 }
 
@@ -178,10 +178,10 @@ func.func @store_nd_vc_3(%dst: memref<24x32xf16>) {
 
 // -----
 func.func @store_nd_vc_4(%dst: memref<8x24x32xf16>) {
-  %1 = arith.constant dense<1.0>: vector<8x24x32xf16>
+  %1 = arith.constant dense<1.0>: vector<8x24x16xf16>
   %2 = xegpu.create_nd_tdesc %dst : memref<8x24x32xf16> -> !xegpu.tensor_desc<8x24x32xf16>
-  // expected-error at +1 {{Expects a 1D or 2D TensorDesc}}
-  xegpu.store_nd %1, %2[0, 0, 0] <{l1_hint = #xegpu.cache_hint<write_back>, l2_hint = #xegpu.cache_hint<uncached>}>: vector<8x24x32xf16>, !xegpu.tensor_desc<8x24x32xf16>
+  // expected-error at +1 {{Value shape [8, 24, 16] is not consistent with tensor descriptor}}
+  xegpu.store_nd %1, %2[0, 0, 0] <{l1_hint = #xegpu.cache_hint<write_back>, l2_hint = #xegpu.cache_hint<uncached>}>: vector<8x24x16xf16>, !xegpu.tensor_desc<8x24x32xf16>
   return
 }
 
@@ -394,7 +394,7 @@ func.func @dpas_vc_1(%a : vector<8x8xf16>, %b: vector<8x16x2xf16>) {
 
 // -----
 func.func @dpas_vc_2(%a : vector<8x8x2xf16>, %b: vector<8x16x2xf16>) {
-  // expected-error at +1 {{op A operand must be a 2D vector}}
+  // expected-error at +1 {{Batch dimension rank mismatch among A, B, and result}}
   %1 = xegpu.dpas %a, %b : vector<8x8x2xf16>, vector<8x16x2xf16> -> vector<8x16xf32>
   return
 }
@@ -702,21 +702,21 @@ func.func @dpas_mx_acc_result_type_mismatch(%a : vector<8x16xf8E5M2>, %b: vector
 
 // -----
 func.func @dpas_mx_a_not_2d(%a : vector<128xf8E5M2>, %b: vector<16x16xf8E5M2>) {
-  // expected-error at +1 {{A operand must be a 2D vector.}}
+  // expected-error at +1 {{A operand must be at least a 2D vector.}}
   %1 = xegpu.dpas_mx %a, %b : (vector<128xf8E5M2>, vector<16x16xf8E5M2>) -> vector<8x16xf32>
   return
 }
 
 // -----
 func.func @dpas_mx_b_not_2d(%a : vector<8x16xf8E5M2>, %b: vector<256xf8E5M2>) {
-  // expected-error at +1 {{B operand must be a 2D or 3D vector.}}
+  // expected-error at +1 {{B operand must be at least a 2D vector.}}
   %1 = xegpu.dpas_mx %a, %b : (vector<8x16xf8E5M2>, vector<256xf8E5M2>) -> vector<8x16xf32>
   return
 }
 
 // -----
 func.func @dpas_mx_result_not_2d(%a : vector<8x16xf8E5M2>, %b: vector<16x16xf8E5M2>) {
-  // expected-error at +1 {{Result must be a 2D vector.}}
+  // expected-error at +1 {{Result must be at least a 2D vector.}}
   %1 = xegpu.dpas_mx %a, %b : (vector<8x16xf8E5M2>, vector<16x16xf8E5M2>) -> vector<128xf32>
   return
 }
diff --git a/mlir/test/Dialect/XeGPU/xegpu-blocking.mlir b/mlir/test/Dialect/XeGPU/xegpu-blocking.mlir
index c0ea112edc818..f2c25f45f45bc 100644
--- a/mlir/test/Dialect/XeGPU/xegpu-blocking.mlir
+++ b/mlir/test/Dialect/XeGPU/xegpu-blocking.mlir
@@ -746,3 +746,72 @@ gpu.module @test_kernel {
     gpu.return
   }
 }
+
+// -----
+// Test 3D batch dimension unrolling for dpas, load_nd, store_nd
+#a_3d = #xegpu.layout<inst_data = [1, 8, 16]>
+#b_3d = #xegpu.layout<inst_data = [1, 16, 16]>
+#c_3d = #xegpu.layout<inst_data = [1, 8, 16]>
+gpu.module @test_kernel {
+  // CHECK-LABEL: func @gemm_3d_batch_blocking
+  gpu.func @gemm_3d_batch_blocking(%A: memref<4x8x32xf16>, %B: memref<4x32x16xf16>, %C: memref<4x8x16xf32>) {
+    %c0 = arith.constant 0 : index
+    %c_tdesc = xegpu.create_nd_tdesc %C : memref<4x8x16xf32> -> !xegpu.tensor_desc<4x8x16xf32, #c_3d>
+    // CHECK-COUNT-4: xegpu.load_nd {{.*}} -> vector<1x8x16xf32>
+    %c_init = xegpu.load_nd %c_tdesc[0, 0, 0] {layout = #c_3d}: !xegpu.tensor_desc<4x8x16xf32, #c_3d> -> vector<4x8x16xf32>
+    %a_tdesc = xegpu.create_nd_tdesc %A : memref<4x8x32xf16> -> !xegpu.tensor_desc<4x8x32xf16, #a_3d>
+    %b_tdesc = xegpu.create_nd_tdesc %B : memref<4x32x16xf16> -> !xegpu.tensor_desc<4x32x16xf16, #b_3d>
+    // CHECK-COUNT-8: xegpu.load_nd {{.*}} -> vector<1x8x16xf16>
+    %a = xegpu.load_nd %a_tdesc[0, 0, 0] {layout = #a_3d}: !xegpu.tensor_desc<4x8x32xf16, #a_3d> -> vector<4x8x32xf16>
+    // CHECK-COUNT-8: xegpu.load_nd {{.*}} -> vector<1x16x16xf16>
+    %b = xegpu.load_nd %b_tdesc[0, 0, 0] {layout = #b_3d}: !xegpu.tensor_desc<4x32x16xf16, #b_3d> -> vector<4x32x16xf16>
+    // CHECK-COUNT-8: xegpu.dpas {{.*}} : vector<1x8x16xf16>, vector<1x16x16xf16>, vector<1x8x16xf32> -> vector<1x8x16xf32>
+    %d = xegpu.dpas %a, %b, %c_init {layout_a = #a_3d, layout_b = #b_3d, layout_cd = #c_3d}
+      : vector<4x8x32xf16>, vector<4x32x16xf16>, vector<4x8x16xf32> -> vector<4x8x16xf32>
+    // CHECK-COUNT-4: xegpu.store_nd {{.*}} : vector<1x8x16xf32>, !xegpu.tensor_desc<1x8x16xf32>
+    xegpu.store_nd %d, %c_tdesc[0, 0, 0] {layout = #c_3d}: vector<4x8x16xf32>, !xegpu.tensor_desc<4x8x16xf32, #c_3d>
+    gpu.return
+  }
+}
+
+// -----
+// Test 3D batch dimension unrolling for dpas_mx with scale_a and scale_b
+#a_3d_mx = #xegpu.layout<inst_data = [1, 8, 64]>
+#b_3d_mx = #xegpu.layout<inst_data = [1, 64, 16]>
+#c_3d_mx = #xegpu.layout<inst_data = [1, 8, 16]>
+#sa_3d = #xegpu.layout<inst_data = [1, 8, 2]>
+#sb_3d = #xegpu.layout<inst_data = [1, 2, 16]>
+gpu.module @test_kernel {
+  // CHECK-LABEL: func @dpas_mx_3d_batch_blocking
+  gpu.func @dpas_mx_3d_batch_blocking(%A: memref<2x16x64xf4E2M1FN>, %B: memref<2x64x32xf4E2M1FN>, %C: memref<2x16x32xf32>, %SA: memref<2x16x2xf8E8M0FNU>, %SB: memref<2x2x32xf8E8M0FNU>) {
+    %c0 = arith.constant 0 : index
+    %c_tdesc = xegpu.create_nd_tdesc %C : memref<2x16x32xf32> -> !xegpu.tensor_desc<2x16x32xf32, #c_3d_mx>
+    // CHECK-COUNT-8: xegpu.load_nd {{.*}} -> vector<1x8x16xf32>
+    %c_init = xegpu.load_nd %c_tdesc[0, 0, 0] {layout = #c_3d_mx}: !xegpu.tensor_desc<2x16x32xf32, #c_3d_mx> -> vector<2x16x32xf32>
+    %a_tdesc = xegpu.create_nd_tdesc %A : memref<2x16x64xf4E2M1FN> -> !xegpu.tensor_desc<2x16x64xf4E2M1FN, #a_3d_mx>
+    %b_tdesc = xegpu.create_nd_tdesc %B : memref<2x64x32xf4E2M1FN> -> !xegpu.tensor_desc<2x64x32xf4E2M1FN, #b_3d_mx>
+    %sa_tdesc = xegpu.create_nd_tdesc %SA : memref<2x16x2xf8E8M0FNU> -> !xegpu.tensor_desc<2x16x2xf8E8M0FNU, #sa_3d>
+    %sb_tdesc = xegpu.create_nd_tdesc %SB : memref<2x2x32xf8E8M0FNU> -> !xegpu.tensor_desc<2x2x32xf8E8M0FNU, #sb_3d>
+    // CHECK-COUNT-4: xegpu.load_nd {{.*}} -> vector<1x8x64xf4E2M1FN>
+    %a = xegpu.load_nd %a_tdesc[0, 0, 0] {layout = #a_3d_mx}: !xegpu.tensor_desc<2x16x64xf4E2M1FN, #a_3d_mx> -> vector<2x16x64xf4E2M1FN>
+    // CHECK-COUNT-4: xegpu.load_nd {{.*}} -> vector<1x64x16xf4E2M1FN>
+    %b = xegpu.load_nd %b_tdesc[0, 0, 0] {layout = #b_3d_mx}: !xegpu.tensor_desc<2x64x32xf4E2M1FN, #b_3d_mx> -> vector<2x64x32xf4E2M1FN>
+    // CHECK-COUNT-4: xegpu.load_nd {{.*}} -> vector<1x8x2xf8E8M0FNU>
+    %sa = xegpu.load_nd %sa_tdesc[0, 0, 0] {layout = #sa_3d}: !xegpu.tensor_desc<2x16x2xf8E8M0FNU, #sa_3d> -> vector<2x16x2xf8E8M0FNU>
+    // CHECK-COUNT-4: xegpu.load_nd {{.*}} -> vector<1x2x16xf8E8M0FNU>
+    %sb = xegpu.load_nd %sb_tdesc[0, 0, 0] {layout = #sb_3d}: !xegpu.tensor_desc<2x2x32xf8E8M0FNU, #sb_3d> -> vector<2x2x32xf8E8M0FNU>
+    // dpas_mx: [2,16,64] x [2,64,32] -> [2,16,32] with scales [2,16,2] and [2,2,32]
+    // unrolled: batch=2, M=16/8=2, K=64/64=1, N=32/16=2 -> 2*2*2=8 results (with k-reduction)
+    // CHECK-COUNT-8: xegpu.dpas_mx {{.*}} : (vector<1x8x64xf4E2M1FN>, vector<1x64x16xf4E2M1FN>, vector<1x8x16xf32>, vector<1x8x2xf8E8M0FNU>, vector<1x2x16xf8E8M0FNU>) -> vector<1x8x16xf32>
+    %d = xegpu.dpas_mx %a, %b, %c_init scale_a = %sa scale_b = %sb
+          {layout_a = #a_3d_mx, layout_b = #b_3d_mx, layout_cd = #c_3d_mx,
+           layout_a_scale = #sa_3d, layout_b_scale = #sb_3d}
+        : (vector<2x16x64xf4E2M1FN>, vector<2x64x32xf4E2M1FN>,
+           vector<2x16x32xf32>,
+           vector<2x16x2xf8E8M0FNU>, vector<2x2x32xf8E8M0FNU>)
+        -> vector<2x16x32xf32>
+    // CHECK-COUNT-8: xegpu.store_nd {{.*}} : vector<1x8x16xf32>, !xegpu.tensor_desc<1x8x16xf32>
+    xegpu.store_nd %d, %c_tdesc[0, 0, 0] {layout = #c_3d_mx}: vector<2x16x32xf32>, !xegpu.tensor_desc<2x16x32xf32, #c_3d_mx>
+    gpu.return
+  }
+}
diff --git a/mlir/test/Integration/Dialect/XeGPU/WG/simple_3d_gemm.mlir b/mlir/test/Integration/Dialect/XeGPU/WG/simple_3d_gemm.mlir
new file mode 100644
index 0000000000000..e82f586c57527
--- /dev/null
+++ b/mlir/test/Integration/Dialect/XeGPU/WG/simple_3d_gemm.mlir
@@ -0,0 +1,55 @@
+// RUN: mlir-opt %s --gpu-lower-to-xevm-pipeline="xegpu-op-level=workgroup" \
+// RUN: | FileCheck %s
+
+// XFAIL: *
+
+#a = #xegpu.layout<sg_layout = [1, 8, 4], sg_data = [4, 8, 32], inst_data = [1, 8, 16]>
+#b = #xegpu.layout<sg_layout = [1, 8, 4], sg_data = [4, 32, 16], inst_data = [1, 16, 16]>
+#c = #xegpu.layout<sg_layout = [1, 8, 4], sg_data = [4, 8, 16], inst_data = [1, 8, 16]>
+#a_prefetch = #xegpu.layout<sg_layout = [4, 8, 1], sg_data = [1, 8, 32], inst_data = [1, 8, 16]>
+#b_prefetch = #xegpu.layout<sg_layout = [4, 4, 2], sg_data = [1, 8, 32], inst_data = [1, 8, 16]>
+
+gpu.module @test_kernel {
+  gpu.func @test_kernel(%A: memref<4x64x256xf16>, %B: memref<4x256x64xf16>, %C: memref<4x64x64xf32>) kernel {
+    %c0 = arith.constant 0 : index
+    %c32 = arith.constant 32 : index
+    %c64 = arith.constant 64 : index
+    %c96 = arith.constant 96 : index
+    %c256 = arith.constant 256 : index
+    %block_id_x = gpu.block_id x
+    %block_id_y = gpu.block_id y
+    %block_id_z = gpu.block_id z
+    %m = arith.muli %block_id_x, %c64 : index
+    %n = arith.muli %block_id_y, %c64 : index
+    %c_tdesc = xegpu.create_nd_tdesc %C : memref<4x64x64xf32> -> !xegpu.tensor_desc<4x64x64xf32, #c>
+    %c_init_value = xegpu.load_nd %c_tdesc[%block_id_z, %m, %n] {layout = #c} : !xegpu.tensor_desc<4x64x64xf32, #c> -> vector<4x64x64xf32>
+    %a_tdesc = xegpu.create_nd_tdesc %A : memref<4x64x256xf16> -> !xegpu.tensor_desc<4x64x32xf16, #a>
+    %b_tdesc = xegpu.create_nd_tdesc %B : memref<4x256x64xf16> -> !xegpu.tensor_desc<4x32x64xf16, #b>
+    // Prefetch A 3 times.
+    %a_prefetch_tdesc = xegpu.create_nd_tdesc %A : memref<4x64x256xf16> -> !xegpu.tensor_desc<4x64x32xf16, #a_prefetch>
+    xegpu.prefetch_nd %a_prefetch_tdesc[%block_id_z, %m, %c0] {layout = #a_prefetch} : !xegpu.tensor_desc<4x64x32xf16, #a_prefetch>
+    xegpu.prefetch_nd %a_prefetch_tdesc[%block_id_z, %m, %c32] {layout = #a_prefetch} : !xegpu.tensor_desc<4x64x32xf16, #a_prefetch>
+    xegpu.prefetch_nd %a_prefetch_tdesc[%block_id_z, %m, %c64] {layout = #a_prefetch} : !xegpu.tensor_desc<4x64x32xf16, #a_prefetch>
+    // Prefetch B 3 times.
+    %b_prefetch_tdesc = xegpu.create_nd_tdesc %B : memref<4x256x64xf16> -> !xegpu.tensor_desc<4x32x64xf16, #b_prefetch>
+    xegpu.prefetch_nd %b_prefetch_tdesc[%block_id_z, %c0, %n] {layout = #b_prefetch} : !xegpu.tensor_desc<4x32x64xf16, #b_prefetch>
+    xegpu.prefetch_nd %b_prefetch_tdesc[%block_id_z, %c32, %n] {layout = #b_prefetch} : !xegpu.tensor_desc<4x32x64xf16, #b_prefetch>
+    xegpu.prefetch_nd %b_prefetch_tdesc[%block_id_z, %c64, %n] {layout = #b_prefetch} : !xegpu.tensor_desc<4x32x64xf16, #b_prefetch>
+
+    %out = scf.for %k = %c0 to %c256 step %c32
+      iter_args(%c_value = %c_init_value)
+      -> (vector<4x64x64xf32>) {
+      %a_value = xegpu.load_nd %a_tdesc[%block_id_z, %m, %k] {layout = #a} : !xegpu.tensor_desc<4x64x32xf16, #a> -> vector<4x64x32xf16>
+      %b_value = xegpu.load_nd %b_tdesc[%block_id_z, %k, %n] {layout = #b} : !xegpu.tensor_desc<4x32x64xf16, #b> -> vector<4x32x64xf16>
+      // Prefetch next tiles.
+      %prefetch_offset = arith.addi %k, %c96 : index
+      xegpu.prefetch_nd %a_prefetch_tdesc[%block_id_z, %m, %prefetch_offset] {layout = #a_prefetch} : !xegpu.tensor_desc<4x64x32xf16, #a_prefetch>
+      xegpu.prefetch_nd %b_prefetch_tdesc[%block_id_z, %prefetch_offset, %n] {layout = #b_prefetch} : !xegpu.tensor_desc<4x32x64xf16, #b_prefetch>
+      %c_new_value = xegpu.dpas %a_value, %b_value, %c_value {layout_a = #a, layout_b = #b, layout_cd = #c}
+        : vector<4x64x32xf16>, vector<4x32x64xf16>, vector<4x64x64xf32> -> vector<4x64x64xf32>
+      scf.yield %c_new_value : vector<4x64x64xf32>
+    }
+    xegpu.store_nd %out, %c_tdesc[%block_id_z, %m, %n] {layout = #c} : vector<4x64x64xf32>, !xegpu.tensor_desc<4x64x64xf32, #c>
+    gpu.return
+  }
+}
diff --git a/mlir/test/Integration/Dialect/XeGPU/WG/simple_3d_mxfp_gemm.mlir b/mlir/test/Integration/Dialect/XeGPU/WG/simple_3d_mxfp_gemm.mlir
new file mode 100644
index 0000000000000..1801ca866ed8a
--- /dev/null
+++ b/mlir/test/Integration/Dialect/XeGPU/WG/simple_3d_mxfp_gemm.mlir
@@ -0,0 +1,55 @@
+// RUN: mlir-opt %s --gpu-lower-to-xevm-pipeline="xegpu-op-level=workgroup zebin-chip=cri" \
+// RUN: | FileCheck %s
+
+// XFAIL: *
+
+// 3D batched MXFP GEMM: [4, 128, 512] x [4, 512, 128] -> [4, 128, 128]
+// with scale_a [4, 128, 16] and scale_b [4, 16, 128]
+// Batch dim = 4, split first dim of A to [4, M], last dim of B to [4, N].
+
+#a = #xegpu.layout<sg_layout = [1, 4, 2], sg_data = [4, 32, 512], inst_data = [1, 8, 64]>
+#b = #xegpu.layout<sg_layout = [1, 4, 2], sg_data = [4, 512, 64], inst_data = [1, 64, 16]>
+#c = #xegpu.layout<sg_layout = [1, 4, 2], sg_data = [4, 32, 64], inst_data = [1, 8, 16]>
+#a_scale = #xegpu.layout<sg_layout = [1, 4, 2], sg_data = [4, 32, 16], inst_data = [1, 8, 2]>
+#b_scale = #xegpu.layout<sg_layout = [1, 4, 2], sg_data = [4, 16, 64], inst_data = [1, 2, 16]>
+
+gpu.module @test {
+  gpu.func @gemm_3d_mxfp(%arg0: memref<4x128x512xf4E2M1FN>, %arg1: memref<4x512x128xf4E2M1FN>, %arg2: memref<4x128x16xf8E8M0FNU>, %arg3: memref<4x16x128xf8E8M0FNU>, %arg4: memref<4x128x128xf32>) kernel {
+    %c0 = arith.constant 0 : index
+    %c128 = arith.constant 128 : index
+    %block_id_x = gpu.block_id x
+    %block_id_y = gpu.block_id y
+    %block_id_z = gpu.block_id z
+    %m = arith.muli %block_id_x, %c128 : index
+    %n = arith.muli %block_id_y, %c128 : index
+
+    %a_tdesc = xegpu.create_nd_tdesc %arg0 : memref<4x128x512xf4E2M1FN> -> !xegpu.tensor_desc<4x128x512xf4E2M1FN>
+    %a = xegpu.load_nd %a_tdesc[%block_id_z, %m, %c0] {layout = #a} : !xegpu.tensor_desc<4x128x512xf4E2M1FN> -> vector<4x128x512xf4E2M1FN>
+
+    %b_tdesc = xegpu.create_nd_tdesc %arg1 : memref<4x512x128xf4E2M1FN> -> !xegpu.tensor_desc<4x512x128xf4E2M1FN>
+    %b = xegpu.load_nd %b_tdesc[%block_id_z, %c0, %n] {layout = #b} : !xegpu.tensor_desc<4x512x128xf4E2M1FN> -> vector<4x512x128xf4E2M1FN>
+
+    %cd_tdesc = xegpu.create_nd_tdesc %arg4 : memref<4x128x128xf32> -> !xegpu.tensor_desc<4x128x128xf32, #c>
+    %c = xegpu.load_nd %cd_tdesc[%block_id_z, %m, %n] {layout = #c} : !xegpu.tensor_desc<4x128x128xf32, #c> -> vector<4x128x128xf32>
+
+    %a_scale_tdesc = xegpu.create_nd_tdesc %arg2 : memref<4x128x16xf8E8M0FNU> -> !xegpu.tensor_desc<4x128x16xf8E8M0FNU>
+    %scale_a = xegpu.load_nd %a_scale_tdesc[%block_id_z, %m, %c0] {layout = #a_scale} : !xegpu.tensor_desc<4x128x16xf8E8M0FNU> -> vector<4x128x16xf8E8M0FNU>
+
+    %b_scale_tdesc = xegpu.create_nd_tdesc %arg3 : memref<4x16x128xf8E8M0FNU> -> !xegpu.tensor_desc<4x16x128xf8E8M0FNU>
+    %scale_b = xegpu.load_nd %b_scale_tdesc[%block_id_z, %c0, %n] {layout = #b_scale} : !xegpu.tensor_desc<4x16x128xf8E8M0FNU> -> vector<4x16x128xf8E8M0FNU>
+
+    %d = xegpu.dpas_mx %a, %b, %c scale_a = %scale_a scale_b = %scale_b
+          {layout_a = #a,
+           layout_b = #b,
+           layout_cd = #c,
+           layout_a_scale = #a_scale,
+           layout_b_scale = #b_scale}
+        : (vector<4x128x512xf4E2M1FN>, vector<4x512x128xf4E2M1FN>,
+          vector<4x128x128xf32>,
+          vector<4x128x16xf8E8M0FNU>, vector<4x16x128xf8E8M0FNU>)
+        -> vector<4x128x128xf32>
+
+    xegpu.store_nd %d, %cd_tdesc[%block_id_z, %m, %n] {layout = #c} : vector<4x128x128xf32>, !xegpu.tensor_desc<4x128x128xf32, #c>
+    gpu.return
+  }
+}
diff --git a/mlir/test/lib/Dialect/XeGPU/TestXeGPUTransforms.cpp b/mlir/test/lib/Dialect/XeGPU/TestXeGPUTransforms.cpp
index 5c3721630837d..afda640c515a5 100644
--- a/mlir/test/lib/Dialect/XeGPU/TestXeGPUTransforms.cpp
+++ b/mlir/test/lib/Dialect/XeGPU/TestXeGPUTransforms.cpp
@@ -111,8 +111,7 @@ struct TestXeGPUUnrollingPatterns
     });
 
     options.setUnrolledTypesFn(
-        [&](ShapedType type, ArrayRef<int64_t> tileShape,
-            bool returnSingleType = false) -> SmallVector<Type> {
+        [&](ShapedType type, ArrayRef<int64_t> tileShape) -> SmallVector<Type> {
           Type elemTy = type.getElementType();
           Type newTy;
 
@@ -131,13 +130,14 @@ struct TestXeGPUUnrollingPatterns
 
             newTy = xegpu::TensorDescType::get(ctx, tileShape, elemTy, encoding,
                                                layout);
-
-          } else {
-            newTy = type.clone(tileShape, elemTy);
+            // compute the product of batch (higher) dimensions
+            ArrayRef<int64_t> shape = type.getShape();
+            int64_t batchCount =
+                shape.size() > 2 ? computeProduct(shape.drop_back(2)) : 1;
+            return SmallVector<Type>(batchCount, newTy);
           }
 
-          if (returnSingleType)
-            return SmallVector<Type>{newTy};
+          newTy = type.clone(tileShape, elemTy);
           std::optional<SmallVector<int64_t>> ratio =
               computeShapeRatio(type.getShape(), tileShape);
           assert(ratio && "Expecting the ratio to be valid.");

>From 3249d942dc84e095ae95317ea01cdab94268792c Mon Sep 17 00:00:00 2001
From: Jianhui Li <jian.hui.li at intel.com>
Date: Fri, 5 Jun 2026 03:32:08 +0000
Subject: [PATCH 2/4] fix compilation issue

---
 mlir/lib/Dialect/XeGPU/Transforms/XeGPUUnroll.cpp | 9 ++++++---
 1 file changed, 6 insertions(+), 3 deletions(-)

diff --git a/mlir/lib/Dialect/XeGPU/Transforms/XeGPUUnroll.cpp b/mlir/lib/Dialect/XeGPU/Transforms/XeGPUUnroll.cpp
index 95ad1808dac07..88f77834548d3 100644
--- a/mlir/lib/Dialect/XeGPU/Transforms/XeGPUUnroll.cpp
+++ b/mlir/lib/Dialect/XeGPU/Transforms/XeGPUUnroll.cpp
@@ -306,11 +306,12 @@ struct UnrollPrefetchNdOp : public UnrollPattern<xegpu::PrefetchNdOp> {
                                              mixedOffsets.end());
 
       for (auto [idx, batchTdesc] : llvm::enumerate(batchTdescs)) {
+        Value tdesc = batchTdesc;
         auto createPrefetch = [&](SmallVector<OpFoldResult> offsets) -> Value {
           SmallVector<OpFoldResult> fullOffsets(batchRank,
                                                 rewriter.getIndexAttr(0));
           fullOffsets.append(offsets.begin(), offsets.end());
-          xegpu::PrefetchNdOp::create(rewriter, loc, batchTdesc, fullOffsets,
+          xegpu::PrefetchNdOp::create(rewriter, loc, tdesc, fullOffsets,
                                       op.getL1HintAttr(), op.getL2HintAttr(),
                                       op.getL3HintAttr(), layout);
           return nullptr;
@@ -388,12 +389,13 @@ struct UnrollLoadNdOp : public UnrollPattern<xegpu::LoadNdOp> {
                                              mixedOffsets.end());
 
       for (auto [idx, batchTdesc] : llvm::enumerate(batchTdescs)) {
+        Value tdesc = batchTdesc;
         auto createLoad = [&](SmallVector<OpFoldResult> offsets) {
           SmallVector<OpFoldResult> fullOffsets(batchRank,
                                                 rewriter.getIndexAttr(0));
           fullOffsets.append(offsets.begin(), offsets.end());
           return xegpu::LoadNdOp::create(
-              rewriter, loc, newValueTy, batchTdesc, fullOffsets,
+              rewriter, loc, newValueTy, tdesc, fullOffsets,
               op.getPackedAttr(), op.getTransposeAttr(), op.getL1HintAttr(),
               op.getL2HintAttr(), op.getL3HintAttr(), layout);
         };
@@ -472,12 +474,13 @@ struct UnrollStoreNdOp : public UnrollPattern<xegpu::StoreNdOp> {
                                              mixedOffsets.end());
 
       for (auto [idx, batchTdesc] : llvm::enumerate(batchTdescs)) {
+        Value tdesc = batchTdesc;
         auto createStore = [&](SmallVector<OpFoldResult> offsets) {
           SmallVector<OpFoldResult> fullOffsets(batchRank,
                                                 rewriter.getIndexAttr(0));
           fullOffsets.append(offsets.begin(), offsets.end());
           xegpu::StoreNdOp::create(rewriter, loc, convertedValues[valueIndex++],
-                                   batchTdesc, fullOffsets, op.getL1HintAttr(),
+                                   tdesc, fullOffsets, op.getL1HintAttr(),
                                    op.getL2HintAttr(), op.getL3HintAttr(),
                                    layout);
           return (Value) nullptr;

>From 51df0586076340fe31af0ebe48ed88b7b7a31fac Mon Sep 17 00:00:00 2001
From: Jianhui Li <jian.hui.li at intel.com>
Date: Fri, 5 Jun 2026 03:35:57 +0000
Subject: [PATCH 3/4] git clang format

---
 mlir/lib/Dialect/XeGPU/Transforms/XeGPUUnroll.cpp | 6 +++---
 1 file changed, 3 insertions(+), 3 deletions(-)

diff --git a/mlir/lib/Dialect/XeGPU/Transforms/XeGPUUnroll.cpp b/mlir/lib/Dialect/XeGPU/Transforms/XeGPUUnroll.cpp
index 88f77834548d3..c69e306dd0fd0 100644
--- a/mlir/lib/Dialect/XeGPU/Transforms/XeGPUUnroll.cpp
+++ b/mlir/lib/Dialect/XeGPU/Transforms/XeGPUUnroll.cpp
@@ -395,9 +395,9 @@ struct UnrollLoadNdOp : public UnrollPattern<xegpu::LoadNdOp> {
                                                 rewriter.getIndexAttr(0));
           fullOffsets.append(offsets.begin(), offsets.end());
           return xegpu::LoadNdOp::create(
-              rewriter, loc, newValueTy, tdesc, fullOffsets,
-              op.getPackedAttr(), op.getTransposeAttr(), op.getL1HintAttr(),
-              op.getL2HintAttr(), op.getL3HintAttr(), layout);
+              rewriter, loc, newValueTy, tdesc, fullOffsets, op.getPackedAttr(),
+              op.getTransposeAttr(), op.getL1HintAttr(), op.getL2HintAttr(),
+              op.getL3HintAttr(), layout);
         };
         auto batchLoads = computeUnrolledOffsets(
             innerOffsets, innerTdescTy, innerTarget, createLoad, loc, rewriter);

>From 46870bab87e3574ef8db500431a84ddea9792b0c Mon Sep 17 00:00:00 2001
From: Jianhui Li <jian.hui.li at intel.com>
Date: Fri, 5 Jun 2026 17:59:04 +0000
Subject: [PATCH 4/4] git clang format

---
 .../XeGPU/Transforms/XeGPUPropagateLayout.cpp  | 18 +++++++++---------
 .../XeGPU/Transforms/XeGPUWgToSgDistribute.cpp |  6 ++----
 mlir/lib/Dialect/XeGPU/Utils/XeGPUUtils.cpp    |  3 ++-
 3 files changed, 13 insertions(+), 14 deletions(-)

diff --git a/mlir/lib/Dialect/XeGPU/Transforms/XeGPUPropagateLayout.cpp b/mlir/lib/Dialect/XeGPU/Transforms/XeGPUPropagateLayout.cpp
index 8026dddfab16d..439bd5f859bfb 100644
--- a/mlir/lib/Dialect/XeGPU/Transforms/XeGPUPropagateLayout.cpp
+++ b/mlir/lib/Dialect/XeGPU/Transforms/XeGPUPropagateLayout.cpp
@@ -306,8 +306,8 @@ static LayoutInfo getSIMTLayoutInfoBlockIO(Ty ty,
   SmallVector<int32_t> laneData(rank, 1);
   laneLayout[rank - 1] = uArch->getSubgroupSize();
   laneData[rank - 1] = packingFactor;
-  return LayoutInfo(xegpu::LayoutAttr::get(ty.getContext(), laneLayout,
-                                           laneData));
+  return LayoutInfo(
+      xegpu::LayoutAttr::get(ty.getContext(), laneLayout, laneData));
 }
 
 //===----------------------------------------------------------------------===//
@@ -423,18 +423,18 @@ class LayoutInfoPropagation
   visitOperation(Operation *op, ArrayRef<LayoutInfoLattice *> operands,
                  ArrayRef<const LayoutInfoLattice *> results) override;
 
-  void visitBranchOperand(OpOperand &operand) override {};
+  void visitBranchOperand(OpOperand &operand) override{};
 
-  void visitCallOperand(OpOperand &operand) override {};
+  void visitCallOperand(OpOperand &operand) override{};
 
   void
   visitNonControlFlowArguments(RegionSuccessor &successor,
-                               ArrayRef<BlockArgument> arguments) override {};
+                               ArrayRef<BlockArgument> arguments) override{};
 
-  void visitExternalCall(CallOpInterface call,
-                         ArrayRef<LayoutInfoLattice *> operands,
-                         ArrayRef<const LayoutInfoLattice *> results) override {
-  };
+  void
+  visitExternalCall(CallOpInterface call,
+                    ArrayRef<LayoutInfoLattice *> operands,
+                    ArrayRef<const LayoutInfoLattice *> results) override{};
 
   void setToExitState(LayoutInfoLattice *lattice) override {
     (void)lattice->meet(LayoutInfo());
diff --git a/mlir/lib/Dialect/XeGPU/Transforms/XeGPUWgToSgDistribute.cpp b/mlir/lib/Dialect/XeGPU/Transforms/XeGPUWgToSgDistribute.cpp
index 1463dc96bcf47..e3f321a7af543 100644
--- a/mlir/lib/Dialect/XeGPU/Transforms/XeGPUWgToSgDistribute.cpp
+++ b/mlir/lib/Dialect/XeGPU/Transforms/XeGPUWgToSgDistribute.cpp
@@ -333,8 +333,7 @@ struct WgToSgDpasOp : public OpConversionPattern<xegpu::DpasOp> {
         SmallVector<int64_t> resShape(aVecShape.drop_back(2));
         resShape.push_back(aVecShape[aVecShape.size() - 2]);
         resShape.push_back(bVecShape[bVecShape.size() - 1]);
-        VectorType resTy =
-            VectorType::get(resShape, resultTy.getElementType());
+        VectorType resTy = VectorType::get(resShape, resultTy.getElementType());
         auto newDpasOp = xegpu::DpasOp::create(rewriter, loc, resTy, operands);
         newDpasOp.setLayoutCdAttr(layoutCd.dropSgLayoutAndData());
         newDpasOp.setLayoutAAttr(layoutA.dropSgLayoutAndData());
@@ -388,8 +387,7 @@ struct WgToSgDpasMxOp : public OpConversionPattern<xegpu::DpasMxOp> {
         SmallVector<int64_t> resShape(aVecShape.drop_back(2));
         resShape.push_back(aVecShape[aVecShape.size() - 2]);
         resShape.push_back(bVecShape[bVecShape.size() - 1]);
-        VectorType resTy =
-            VectorType::get(resShape, resultTy.getElementType());
+        VectorType resTy = VectorType::get(resShape, resultTy.getElementType());
         auto newDpasMxOp = xegpu::DpasMxOp::create(
             rewriter, loc, resTy, aVec, bVec, accVal, scaleAVal, scaleBVal,
             layoutA.dropSgLayoutAndData(), layoutB.dropSgLayoutAndData(),
diff --git a/mlir/lib/Dialect/XeGPU/Utils/XeGPUUtils.cpp b/mlir/lib/Dialect/XeGPU/Utils/XeGPUUtils.cpp
index d7fcc6b7bf046..e4a085bdde6d3 100644
--- a/mlir/lib/Dialect/XeGPU/Utils/XeGPUUtils.cpp
+++ b/mlir/lib/Dialect/XeGPU/Utils/XeGPUUtils.cpp
@@ -92,7 +92,8 @@ mlir::xegpu::getDistributedVectorType(VectorType originalType,
     laneData.erase(laneData.begin());
   }
   auto trimmedLayout = xegpu::LayoutAttr::get(
-      layout.getContext(), SmallVector<int32_t>(laneLayout.begin(), laneLayout.end()),
+      layout.getContext(),
+      SmallVector<int32_t>(laneLayout.begin(), laneLayout.end()),
       SmallVector<int32_t>(laneData.begin(), laneData.end()));
   auto helperTdescTy = xegpu::TensorDescType::get(
       shape, originalType.getElementType(), arrayLength,



More information about the Mlir-commits mailing list