[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