Supporting i64 comparison in VM conversion. (#8104)
diff --git a/iree/compiler/Dialect/VM/Conversion/StandardToVM/ConvertStandardToVM.cpp b/iree/compiler/Dialect/VM/Conversion/StandardToVM/ConvertStandardToVM.cpp index 238b277..e5bb31f 100644 --- a/iree/compiler/Dialect/VM/Conversion/StandardToVM/ConvertStandardToVM.cpp +++ b/iree/compiler/Dialect/VM/Conversion/StandardToVM/ConvertStandardToVM.cpp
@@ -65,7 +65,6 @@ class FuncOpConversion : public OpConversionPattern<FuncOp> { using OpConversionPattern::OpConversionPattern; - LogicalResult matchAndRewrite( FuncOp srcOp, OpAdaptor adaptor, ConversionPatternRewriter &rewriter) const override { @@ -136,7 +135,6 @@ class ReturnOpConversion : public OpConversionPattern<mlir::ReturnOp> { using OpConversionPattern::OpConversionPattern; - LogicalResult matchAndRewrite( mlir::ReturnOp srcOp, OpAdaptor adaptor, ConversionPatternRewriter &rewriter) const override { @@ -147,11 +145,9 @@ }; struct ConstantOpConversion : public OpConversionPattern<arith::ConstantOp> { + TypeConverter &typeConverter; ConstantOpConversion(MLIRContext *context, TypeConverter &typeConverter) : OpConversionPattern(context), typeConverter(typeConverter) {} - - TypeConverter &typeConverter; - LogicalResult matchAndRewrite( arith::ConstantOp srcOp, OpAdaptor adaptor, ConversionPatternRewriter &rewriter) const override { @@ -219,12 +215,12 @@ } }; -class CmpIOpConversion : public OpConversionPattern<arith::CmpIOp> { +struct CmpI32OpConversion : public OpConversionPattern<arith::CmpIOp> { using OpConversionPattern::OpConversionPattern; - LogicalResult matchAndRewrite( arith::CmpIOp srcOp, OpAdaptor adaptor, ConversionPatternRewriter &rewriter) const override { + if (!adaptor.getLhs().getType().isInteger(32)) return failure(); auto returnType = rewriter.getIntegerType(32); switch (srcOp.getPredicate()) { case arith::CmpIPredicate::eq: @@ -273,12 +269,66 @@ } }; -class CmpFOpConversion : public OpConversionPattern<arith::CmpFOp> { +struct CmpI64OpConversion : public OpConversionPattern<arith::CmpIOp> { using OpConversionPattern::OpConversionPattern; + LogicalResult matchAndRewrite( + arith::CmpIOp srcOp, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + if (!adaptor.getLhs().getType().isInteger(64)) return failure(); + auto returnType = rewriter.getIntegerType(32); + switch (srcOp.getPredicate()) { + case arith::CmpIPredicate::eq: + rewriter.replaceOpWithNewOp<IREE::VM::CmpEQI64Op>( + srcOp, returnType, adaptor.getLhs(), adaptor.getRhs()); + return success(); + case arith::CmpIPredicate::ne: + rewriter.replaceOpWithNewOp<IREE::VM::CmpNEI64Op>( + srcOp, returnType, adaptor.getLhs(), adaptor.getRhs()); + return success(); + case arith::CmpIPredicate::slt: + rewriter.replaceOpWithNewOp<IREE::VM::CmpLTI64SOp>( + srcOp, returnType, adaptor.getLhs(), adaptor.getRhs()); + return success(); + case arith::CmpIPredicate::sle: + rewriter.replaceOpWithNewOp<IREE::VM::CmpLTEI64SOp>( + srcOp, returnType, adaptor.getLhs(), adaptor.getRhs()); + return success(); + case arith::CmpIPredicate::sgt: + rewriter.replaceOpWithNewOp<IREE::VM::CmpGTI64SOp>( + srcOp, returnType, adaptor.getLhs(), adaptor.getRhs()); + return success(); + case arith::CmpIPredicate::sge: + rewriter.replaceOpWithNewOp<IREE::VM::CmpGTEI64SOp>( + srcOp, returnType, adaptor.getLhs(), adaptor.getRhs()); + return success(); + case arith::CmpIPredicate::ult: + rewriter.replaceOpWithNewOp<IREE::VM::CmpLTI64UOp>( + srcOp, returnType, adaptor.getLhs(), adaptor.getRhs()); + return success(); + case arith::CmpIPredicate::ule: + rewriter.replaceOpWithNewOp<IREE::VM::CmpLTEI64UOp>( + srcOp, returnType, adaptor.getLhs(), adaptor.getRhs()); + return success(); + case arith::CmpIPredicate::ugt: + rewriter.replaceOpWithNewOp<IREE::VM::CmpGTI64UOp>( + srcOp, returnType, adaptor.getLhs(), adaptor.getRhs()); + return success(); + case arith::CmpIPredicate::uge: + rewriter.replaceOpWithNewOp<IREE::VM::CmpGTEI64UOp>( + srcOp, returnType, adaptor.getLhs(), adaptor.getRhs()); + return success(); + default: + return failure(); + } + } +}; +struct CmpF32OpConversion : public OpConversionPattern<arith::CmpFOp> { + using OpConversionPattern::OpConversionPattern; LogicalResult matchAndRewrite( arith::CmpFOp srcOp, OpAdaptor adaptor, ConversionPatternRewriter &rewriter) const override { + if (!adaptor.getLhs().getType().isF32()) return failure(); auto returnType = rewriter.getIntegerType(32); switch (srcOp.getPredicate()) { case arith::CmpFPredicate::AlwaysFalse: // 0 @@ -365,7 +415,6 @@ template <typename SrcOpTy, typename Dst32OpTy, typename Dst64OpTy> class UnaryArithmeticOpConversion : public OpConversionPattern<SrcOpTy> { using OpConversionPattern<SrcOpTy>::OpConversionPattern; - LogicalResult matchAndRewrite( SrcOpTy srcOp, typename SrcOpTy::Adaptor adaptor, ConversionPatternRewriter &rewriter) const override { @@ -388,7 +437,6 @@ template <typename SrcOpTy, typename Dst32OpTy, typename Dst64OpTy> class BinaryArithmeticOpConversion : public OpConversionPattern<SrcOpTy> { using OpConversionPattern<SrcOpTy>::OpConversionPattern; - LogicalResult matchAndRewrite( SrcOpTy srcOp, typename SrcOpTy::Adaptor adaptor, ConversionPatternRewriter &rewriter) const override { @@ -413,7 +461,6 @@ template <typename SrcOpTy, typename Dst32OpTy, typename Dst64OpTy> class ShiftArithmeticOpConversion : public OpConversionPattern<SrcOpTy> { using OpConversionPattern<SrcOpTy>::OpConversionPattern; - LogicalResult matchAndRewrite( SrcOpTy srcOp, typename SrcOpTy::Adaptor adaptor, ConversionPatternRewriter &rewriter) const override { @@ -442,7 +489,6 @@ template <typename StdOp> class CastingOpConversion : public OpConversionPattern<StdOp> { using OpConversionPattern<StdOp>::OpConversionPattern; - LogicalResult matchAndRewrite( StdOp srcOp, typename StdOp::Adaptor adaptor, ConversionPatternRewriter &rewriter) const override { @@ -453,7 +499,6 @@ class IndexCastOpConversion : public OpConversionPattern<arith::IndexCastOp> { using OpConversionPattern::OpConversionPattern; - LogicalResult matchAndRewrite( arith::IndexCastOp srcOp, OpAdaptor adaptor, ConversionPatternRewriter &rewriter) const override { @@ -475,7 +520,6 @@ class ZeroExtendIOpConversion : public OpConversionPattern<arith::ExtUIOp> { using OpConversionPattern::OpConversionPattern; - LogicalResult matchAndRewrite( arith::ExtUIOp srcOp, OpAdaptor adaptor, ConversionPatternRewriter &rewriter) const override { @@ -509,7 +553,6 @@ class SignExtendIOpConversion : public OpConversionPattern<arith::ExtSIOp> { using OpConversionPattern::OpConversionPattern; - LogicalResult matchAndRewrite( arith::ExtSIOp srcOp, OpAdaptor adaptor, ConversionPatternRewriter &rewriter) const override { @@ -535,7 +578,6 @@ class TruncateIOpConversion : public OpConversionPattern<arith::TruncIOp> { using OpConversionPattern::OpConversionPattern; - LogicalResult matchAndRewrite( arith::TruncIOp srcOp, OpAdaptor adaptor, ConversionPatternRewriter &rewriter) const override { @@ -578,7 +620,6 @@ class SIToFPOpConversion : public OpConversionPattern<arith::SIToFPOp> { using OpConversionPattern::OpConversionPattern; - LogicalResult matchAndRewrite( arith::SIToFPOp srcOp, OpAdaptor adaptor, ConversionPatternRewriter &rewriter) const override { @@ -597,7 +638,6 @@ class UIToFPOpConversion : public OpConversionPattern<arith::UIToFPOp> { using OpConversionPattern::OpConversionPattern; - LogicalResult matchAndRewrite( arith::UIToFPOp srcOp, OpAdaptor adaptor, ConversionPatternRewriter &rewriter) const override { @@ -616,7 +656,6 @@ class FPToSIOpConversion : public OpConversionPattern<arith::FPToSIOp> { using OpConversionPattern::OpConversionPattern; - LogicalResult matchAndRewrite( arith::FPToSIOp srcOp, OpAdaptor adaptor, ConversionPatternRewriter &rewriter) const override { @@ -635,7 +674,6 @@ class FPToUIOpConversion : public OpConversionPattern<arith::FPToUIOp> { using OpConversionPattern::OpConversionPattern; - LogicalResult matchAndRewrite( arith::FPToUIOp srcOp, OpAdaptor adaptor, ConversionPatternRewriter &rewriter) const override { @@ -654,7 +692,6 @@ class BitcastOpConversion : public OpConversionPattern<arith::BitcastOp> { using OpConversionPattern::OpConversionPattern; - LogicalResult matchAndRewrite( arith::BitcastOp srcOp, OpAdaptor adaptor, ConversionPatternRewriter &rewriter) const override { @@ -719,7 +756,6 @@ class AssertOpConversion : public OpConversionPattern<AssertOp> { using OpConversionPattern::OpConversionPattern; - LogicalResult matchAndRewrite( AssertOp srcOp, OpAdaptor adaptor, ConversionPatternRewriter &rewriter) const override { @@ -740,7 +776,6 @@ class BranchOpConversion : public OpConversionPattern<BranchOp> { using OpConversionPattern::OpConversionPattern; - LogicalResult matchAndRewrite( BranchOp srcOp, OpAdaptor adaptor, ConversionPatternRewriter &rewriter) const override { @@ -752,7 +787,6 @@ class CondBranchOpConversion : public OpConversionPattern<CondBranchOp> { using OpConversionPattern::OpConversionPattern; - LogicalResult matchAndRewrite( CondBranchOp srcOp, OpAdaptor adaptor, ConversionPatternRewriter &rewriter) const override { @@ -766,7 +800,6 @@ class CallOpConversion : public OpConversionPattern<CallOp> { using OpConversionPattern::OpConversionPattern; - LogicalResult matchAndRewrite( CallOp srcOp, OpAdaptor adaptor, ConversionPatternRewriter &rewriter) const override { @@ -793,9 +826,10 @@ TypeConverter &typeConverter, OwningRewritePatternList &patterns) { patterns.insert<AssertOpConversion, BranchOpConversion, CallOpConversion, - CmpIOpConversion, CmpFOpConversion, CondBranchOpConversion, - ModuleOpConversion, FuncOpConversion, ReturnOpConversion, - SelectOpConversion>(typeConverter, context); + CmpI32OpConversion, CmpI64OpConversion, CmpF32OpConversion, + CondBranchOpConversion, ModuleOpConversion, FuncOpConversion, + ReturnOpConversion, SelectOpConversion>(typeConverter, + context); // TODO(#2878): figure out how to pass the type converter in a supported way. // Right now if we pass the type converter as the first argument - triggering
diff --git a/iree/compiler/Dialect/VM/Conversion/StandardToVM/test/comparison_ops.mlir b/iree/compiler/Dialect/VM/Conversion/StandardToVM/test/comparison_ops.mlir index 7c33590..73626a9 100644 --- a/iree/compiler/Dialect/VM/Conversion/StandardToVM/test/comparison_ops.mlir +++ b/iree/compiler/Dialect/VM/Conversion/StandardToVM/test/comparison_ops.mlir
@@ -169,3 +169,20 @@ } } + +// ----- +// CHECK-LABEL: @t011_cmp_uge_i64 +module @t011_cmp_uge_i64 { + +module { + // CHECK: vm.func private @my_fn + // CHECK-SAME: %[[ARG0:[a-zA-Z0-9$._-]+]] + // CHECK-SAME: %[[ARG1:[a-zA-Z0-9$._-]+]] + func @my_fn(%arg0: i64, %arg1 : i64) -> (i1) { + // CHECK: vm.cmp.gte.i64.u %[[ARG0]], %[[ARG1]] : i64 + %1 = arith.cmpi uge, %arg0, %arg1 : i64 + return %1 : i1 + } +} + +}