#include "mlir/Conversion/LLVMCommon/ConversionTarget.h"
#include "mlir/Conversion/LLVMCommon/Pattern.h"
#include "mlir/Dialect/ArmSVE/IR/ArmSVEDialect.h"
#include "mlir/Dialect/ArmSVE/Transforms/Transforms.h"
#include "mlir/Dialect/LLVMIR/LLVMDialect.h"
#include "mlir/Dialect/Utils/IndexingUtils.h"
#include "mlir/Dialect/Vector/IR/VectorOps.h"
#include "mlir/IR/PatternMatch.h"
using namespace mlir;
using namespace mlir::arm_sve;
using SdotOpLowering = OneToOneConvertToLLVMPattern<SdotOp, SdotIntrOp>;
using SmmlaOpLowering = OneToOneConvertToLLVMPattern<SmmlaOp, SmmlaIntrOp>;
using UdotOpLowering = OneToOneConvertToLLVMPattern<UdotOp, UdotIntrOp>;
using UmmlaOpLowering = OneToOneConvertToLLVMPattern<UmmlaOp, UmmlaIntrOp>;
using UsmmlaOpLowering = OneToOneConvertToLLVMPattern<UsmmlaOp, UsmmlaIntrOp>;
using DupQLaneLowering =
OneToOneConvertToLLVMPattern<DupQLaneOp, DupQLaneIntrOp>;
using ScalableMaskedAddIOpLowering =
OneToOneConvertToLLVMPattern<ScalableMaskedAddIOp,
ScalableMaskedAddIIntrOp>;
using ScalableMaskedAddFOpLowering =
OneToOneConvertToLLVMPattern<ScalableMaskedAddFOp,
ScalableMaskedAddFIntrOp>;
using ScalableMaskedSubIOpLowering =
OneToOneConvertToLLVMPattern<ScalableMaskedSubIOp,
ScalableMaskedSubIIntrOp>;
using ScalableMaskedSubFOpLowering =
OneToOneConvertToLLVMPattern<ScalableMaskedSubFOp,
ScalableMaskedSubFIntrOp>;
using ScalableMaskedMulIOpLowering =
OneToOneConvertToLLVMPattern<ScalableMaskedMulIOp,
ScalableMaskedMulIIntrOp>;
using ScalableMaskedMulFOpLowering =
OneToOneConvertToLLVMPattern<ScalableMaskedMulFOp,
ScalableMaskedMulFIntrOp>;
using ScalableMaskedSDivIOpLowering =
OneToOneConvertToLLVMPattern<ScalableMaskedSDivIOp,
ScalableMaskedSDivIIntrOp>;
using ScalableMaskedUDivIOpLowering =
OneToOneConvertToLLVMPattern<ScalableMaskedUDivIOp,
ScalableMaskedUDivIIntrOp>;
using ScalableMaskedDivFOpLowering =
OneToOneConvertToLLVMPattern<ScalableMaskedDivFOp,
ScalableMaskedDivFIntrOp>;
namespace {
template <typename Op, typename IntrOp>
struct SvboolConversionOpLowering : public ConvertOpToLLVMPattern<Op> {
using ConvertOpToLLVMPattern<Op>::ConvertOpToLLVMPattern;
LogicalResult
matchAndRewrite(Op convertOp, typename Op::Adaptor,
ConversionPatternRewriter &rewriter) const override {
auto loc = convertOp.getLoc();
auto source = convertOp.getSource();
VectorType sourceType = source.getType();
VectorType resultType = convertOp.getResult().getType();
Value result = arith::ConstantOp::create(rewriter, loc, resultType,
rewriter.getZeroAttr(resultType));
SmallVector<int64_t> tileShape(sourceType.getRank(), 1);
tileShape.back() = sourceType.getShape().back();
for (SmallVector<int64_t> index :
StaticTileOffsetRange(sourceType.getShape(), tileShape)) {
auto extractOrInsertPosition = ArrayRef(index).drop_back();
auto sourceVector = vector::ExtractOp::create(rewriter, loc, source,
extractOrInsertPosition);
VectorType convertedType =
VectorType::Builder(llvm::cast<VectorType>(sourceVector.getType()))
.setDim(0, resultType.getShape().back());
auto convertedVector =
IntrOp::create(rewriter, loc, TypeRange{convertedType}, sourceVector);
result = vector::InsertOp::create(rewriter, loc, convertedVector, result,
extractOrInsertPosition);
}
rewriter.replaceOp(convertOp, result);
return success();
}
};
using ConvertToSvboolOpLowering =
SvboolConversionOpLowering<ConvertToSvboolOp, ConvertToSvboolIntrOp>;
using ConvertFromSvboolOpLowering =
SvboolConversionOpLowering<ConvertFromSvboolOp, ConvertFromSvboolIntrOp>;
using ZipX2OpLowering = OneToOneConvertToLLVMPattern<ZipX2Op, ZipX2IntrOp>;
using ZipX4OpLowering = OneToOneConvertToLLVMPattern<ZipX4Op, ZipX4IntrOp>;
struct PselOpLowering : public ConvertOpToLLVMPattern<PselOp> {
using ConvertOpToLLVMPattern::ConvertOpToLLVMPattern;
LogicalResult
matchAndRewrite(PselOp pselOp, PselOp::Adaptor adaptor,
ConversionPatternRewriter &rewriter) const override {
auto svboolType = VectorType::get(16, rewriter.getI1Type(), true);
auto loc = pselOp.getLoc();
auto svboolP1 = ConvertToSvboolIntrOp::create(rewriter, loc, svboolType,
adaptor.getP1());
auto indexI32 = arith::IndexCastOp::create(
rewriter, loc, rewriter.getI32Type(), pselOp.getIndex());
auto pselIntr = PselIntrOp::create(rewriter, loc, svboolType, svboolP1,
pselOp.getP2(), indexI32);
rewriter.replaceOpWithNewOp<ConvertFromSvboolIntrOp>(
pselOp, adaptor.getP1().getType(), pselIntr);
return success();
}
};
struct CreateMaskOpLowering
: public ConvertOpToLLVMPattern<vector::CreateMaskOp> {
using ConvertOpToLLVMPattern::ConvertOpToLLVMPattern;
LogicalResult
matchAndRewrite(vector::CreateMaskOp createMaskOp,
vector::CreateMaskOp::Adaptor adaptor,
ConversionPatternRewriter &rewriter) const override {
auto maskType = createMaskOp.getVectorType();
if (maskType.getRank() != 1 || !maskType.isScalable())
return rewriter.notifyMatchFailure(createMaskOp, "not 1-D and scalable");
auto maskBaseSize = maskType.getDimSize(0);
if (maskBaseSize < 2 || maskBaseSize > 16 ||
!llvm::isPowerOf2_32(uint32_t(maskBaseSize)))
return rewriter.notifyMatchFailure(createMaskOp,
"not SVE predicate-sized");
auto loc = createMaskOp.getLoc();
auto zero = LLVM::ZeroOp::create(rewriter, loc, rewriter.getI64Type());
rewriter.replaceOpWithNewOp<WhileLTIntrOp>(createMaskOp, maskType, zero,
adaptor.getOperands()[0]);
return success();
}
};
}
void mlir::populateArmSVELegalizeForLLVMExportPatterns(
const LLVMTypeConverter &converter, RewritePatternSet &patterns) {
patterns.add<ConvertFromSvboolOpLowering,
ConvertToSvboolOpLowering,
DupQLaneLowering,
PselOpLowering,
ScalableMaskedAddFOpLowering,
ScalableMaskedAddIOpLowering,
ScalableMaskedDivFOpLowering,
ScalableMaskedMulFOpLowering,
ScalableMaskedMulIOpLowering,
ScalableMaskedSDivIOpLowering,
ScalableMaskedSubFOpLowering,
ScalableMaskedSubIOpLowering,
ScalableMaskedUDivIOpLowering,
SmmlaOpLowering,
UdotOpLowering,
UmmlaOpLowering,
UsmmlaOpLowering,
ZipX2OpLowering,
ZipX4OpLowering,
SdotOpLowering>(converter);
patterns.add<CreateMaskOpLowering>(converter, 4096);
}
void mlir::configureArmSVELegalizeForExportTarget(
LLVMConversionTarget &target) {
target.addLegalOp<BfmmlaOp,
ConvertFromSvboolIntrOp,
ConvertToSvboolIntrOp,
DupQLaneIntrOp,
PselIntrOp,
ScalableMaskedAddFIntrOp,
ScalableMaskedAddIIntrOp,
ScalableMaskedDivFIntrOp,
ScalableMaskedMulFIntrOp,
ScalableMaskedMulIIntrOp,
ScalableMaskedSDivIIntrOp,
ScalableMaskedSubFIntrOp,
ScalableMaskedSubIIntrOp,
ScalableMaskedUDivIIntrOp,
SmmlaIntrOp,
UdotIntrOp,
UmmlaIntrOp,
UsmmlaIntrOp,
WhileLTIntrOp,
ZipX2IntrOp,
ZipX4IntrOp,
SdotIntrOp>();
target.addIllegalOp<ConvertFromSvboolOp,
ConvertToSvboolOp,
DupQLaneOp,
PselOp,
ScalableMaskedAddFOp,
ScalableMaskedAddIOp,
ScalableMaskedDivFOp,
ScalableMaskedMulFOp,
ScalableMaskedMulIOp,
ScalableMaskedSDivIOp,
ScalableMaskedSubFOp,
ScalableMaskedSubIOp,
ScalableMaskedUDivIOp,
SmmlaOp,
UdotOp,
UmmlaOp,
UsmmlaOp,
ZipX2Op,
ZipX4Op,
SdotOp>();
}