[Mlir-commits] [mlir] [mlir][xegpu] Support batched matmul in VectorToXeGPU ContractionLowering (PR #211947)
Jianhui Li
llvmlistbot at llvm.org
Mon Aug 3 22:01:21 PDT 2026
================
@@ -935,21 +992,24 @@ struct ContractionLowering : public OpRewritePattern<vector::ContractionOp> {
return rewriter.notifyMatchFailure(contractOp,
"Expects add combining kind");
+ TypedValue<VectorType> lhs = contractOp.getLhs();
+ TypedValue<VectorType> rhs = contractOp.getRhs();
TypedValue<Type> acc = contractOp.getAcc();
VectorType accType = dyn_cast<VectorType>(acc.getType());
- if (!accType || accType.getRank() != 2)
- return rewriter.notifyMatchFailure(contractOp, "Expects acc 2D vector");
+ if (!accType)
+ return rewriter.notifyMatchFailure(contractOp, "Expects vector acc");
- // Accept only plain 2D data layout.
- // VNNI packing is applied to DPAS as a separate lowering step.
- TypedValue<VectorType> lhs = contractOp.getLhs();
- TypedValue<VectorType> rhs = contractOp.getRhs();
- if (lhs.getType().getRank() != 2 || rhs.getType().getRank() != 2)
- return rewriter.notifyMatchFailure(contractOp,
- "Expects lhs and rhs 2D vectors");
+ if (!getRowMajorMatmulBatchRank(contractOp.getIndexingMapsAttr()))
+ return rewriter.notifyMatchFailure(
+ contractOp,
+ "Expects a (batched) row-major matmul: leading dims must "
+ "be batch dims shared by lhs, rhs, and acc; innermost two "
+ "dims must be (M, K), (K, N), and (M, N)");
- if (!isRowMajorMatmul(contractOp.getIndexingMapsAttr()))
- return rewriter.notifyMatchFailure(contractOp, "Invalid indexing maps");
+ // xegpu.dpas operands are limited to rank 4 (2 batch + 2 core dims).
+ if (accType.getRank() > 4)
----------------
Jianhui-Li wrote:
Good call. Switched to using the returned batch rank — the check now reads if (*batchRank > 2) (max 2 batch dims)
https://github.com/llvm/llvm-project/pull/211947
More information about the Mlir-commits
mailing list