[Mlir-commits] [mlir] [MLIR][XeGPU] Add unrolling/blocking support for 3D+ batched operations (PR #201725)
Artem Kroviakov
llvmlistbot at llvm.org
Mon Jun 8 06:59:55 PDT 2026
================
@@ -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;
----------------
akroviakov wrote:
```suggestion
Type newTy = VectorType::get(tileShape, elemTy);
```
https://github.com/llvm/llvm-project/pull/201725
More information about the Mlir-commits
mailing list