[Mlir-commits] [mlir] [MLIR][XeGPU][XeVM] XeGPU to XeVM: Add lowering for xegpu.dpas_mx (PR #196981)
Artem Kroviakov
llvmlistbot at llvm.org
Wed May 13 02:31:16 PDT 2026
================
@@ -1066,6 +1076,78 @@ class AtomicRMWToXeVMPattern : public OpConversionPattern<xegpu::AtomicRMWOp> {
}
};
+class DpasMxToXeVMPattern : public OpConversionPattern<xegpu::DpasMxOp> {
+ using OpConversionPattern::OpConversionPattern;
+ LogicalResult
+ matchAndRewrite(xegpu::DpasMxOp op, xegpu::DpasMxOp::Adaptor adaptor,
+ ConversionPatternRewriter &rewriter) const override {
+ auto loc = op.getLoc();
+ auto ctxt = rewriter.getContext();
+ auto aTy = op.getA().getType();
+ auto bTy = op.getB().getType();
+ auto resultType = cast<VectorType>(op.getType());
+
+ auto chipStr = xegpu::getChipStr(op);
+ if (!chipStr)
+ return rewriter.notifyMatchFailure(op, "cannot determine target chip");
+
+ const auto *uArch = xegpu::uArch::getUArch(*chipStr);
+ if (!uArch)
+ return rewriter.notifyMatchFailure(op, "unsupported target uArch");
+
+ // TODO: Add supported shape check
+
+ xevm::ElemType precATy = encodePrecision(aTy.getElementType());
+ xevm::ElemType precBTy = encodePrecision(bTy.getElementType());
+ Value c = op.getAcc();
+ if (!c) {
+ auto elementTy = resultType.getElementType();
+ Attribute initValueAttr;
+ if (isa<FloatType>(elementTy))
+ initValueAttr = FloatAttr::get(elementTy, 0.0);
+ else
+ initValueAttr = IntegerAttr::get(elementTy, 0);
+ c = arith::ConstantOp::create(
+ rewriter, loc, DenseElementsAttr::get(resultType, initValueAttr));
+ }
+
+ Value aVec = adaptor.getA();
+ Value bVec = adaptor.getB();
+ auto aVecTy = cast<VectorType>(aVec.getType());
+ auto bVecTy = cast<VectorType>(bVec.getType());
+ if (aVecTy.getElementTypeBitWidth() == 4)
+ aVec = vector::BitCastOp::create(
+ rewriter, loc,
+ VectorType::get(aVecTy.getNumElements() / 2, rewriter.getI8Type()),
+ aVec);
+ if (bVecTy.getElementTypeBitWidth() == 4)
+ bVec = vector::BitCastOp::create(
+ rewriter, loc,
+ VectorType::get(bVecTy.getNumElements() / 2, rewriter.getI8Type()),
+ bVec);
+ auto cvecty = cast<VectorType>(c.getType());
+ xevm::ElemType precCTy = encodePrecision(cvecty.getElementType());
+ xevm::ElemType precDTy = encodePrecision(resultType.getElementType());
+ VectorType cNty =
+ VectorType::get(cvecty.getNumElements(), cvecty.getElementType());
+ if (cvecty != cNty)
+ c = vector::ShapeCastOp::create(rewriter, loc, cNty, c);
+ Value scaleA = adaptor.getScaleA();
+ Value scaleB = adaptor.getScaleB();
+ Value dpasMxRes = xevm::MMAMxOp::create(
+ rewriter, loc, cNty, aVec, bVec, scaleA, scaleB, c,
+ xevm::MMAShapeAttr::get(ctxt, cvecty.getNumElements(), executionSize,
+ systolicDepth *
+ getNumOperandsPerDword(precATy)),
+ xevm::MMATypesAttr::get(ctxt, precDTy, precATy, precBTy, precCTy));
+ if (cvecty != cNty)
+ dpasMxRes =
+ vector::ShapeCastOp::create(rewriter, loc, resultType, dpasMxRes);
----------------
akroviakov wrote:
Shouldn't the source materialization do this automatically?
https://github.com/llvm/llvm-project/pull/196981
More information about the Mlir-commits
mailing list