[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 =
----------------
akroviakov wrote:
How is it different from
`cast<VectorType>(adaptor.getAcc().getType())` or `getTypeConverter()->convertType(resultType)`?
https://github.com/llvm/llvm-project/pull/196981
More information about the Mlir-commits
mailing list