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

Charitha Saumya llvmlistbot at llvm.org
Fri Jun 5 15:49:29 PDT 2026


================
@@ -266,24 +346,66 @@ 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());
----------------
charithaintc wrote:

else branch is very similar for prefetch and load and maybe store. consider moving to a helper and reusing the logic rather than repeating logic. 

https://github.com/llvm/llvm-project/pull/201725


More information about the Mlir-commits mailing list