From 36b2c22e07f1663d7d7388568ab979b35ecceb94 Mon Sep 17 00:00:00 2001 From: Corentin Ferry Date: Wed, 31 Jul 2024 10:41:18 +0200 Subject: [PATCH] [mlir][emitc] Lower arith.divui, remui (#99313) This commit lowers `arith.divui` and `arith.remui` to EmitC by wrapping those operations with type conversions. --- .../Conversion/ArithToEmitC/ArithToEmitC.cpp | 34 +++++++++++++++++++ .../arith-to-emitc-unsupported.mlir | 16 +++++++++ .../ArithToEmitC/arith-to-emitc.mlir | 18 ++++++++++ 3 files changed, 68 insertions(+) diff --git a/mlir/lib/Conversion/ArithToEmitC/ArithToEmitC.cpp b/mlir/lib/Conversion/ArithToEmitC/ArithToEmitC.cpp index db93f186fbd5..50384d9a08e5 100644 --- a/mlir/lib/Conversion/ArithToEmitC/ArithToEmitC.cpp +++ b/mlir/lib/Conversion/ArithToEmitC/ArithToEmitC.cpp @@ -421,6 +421,38 @@ public: } }; +template +class BinaryUIOpConversion final : public OpConversionPattern { +public: + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(ArithOp uiBinOp, typename ArithOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + Type newRetTy = this->getTypeConverter()->convertType(uiBinOp.getType()); + if (!newRetTy) + return rewriter.notifyMatchFailure(uiBinOp, + "converting result type failed"); + if (!isa(newRetTy)) { + return rewriter.notifyMatchFailure(uiBinOp, "expected integer type"); + } + Type unsignedType = + adaptIntegralTypeSignedness(newRetTy, /*needsUnsigned=*/true); + if (!unsignedType) + return rewriter.notifyMatchFailure(uiBinOp, + "converting result type failed"); + Value lhsAdapted = adaptValueType(uiBinOp.getLhs(), rewriter, unsignedType); + Value rhsAdapted = adaptValueType(uiBinOp.getRhs(), rewriter, unsignedType); + + auto newDivOp = + rewriter.create(uiBinOp.getLoc(), unsignedType, + ArrayRef{lhsAdapted, rhsAdapted}); + Value resultAdapted = adaptValueType(newDivOp, rewriter, newRetTy); + rewriter.replaceOp(uiBinOp, resultAdapted); + return success(); + } +}; + template class IntegerOpConversion final : public OpConversionPattern { public: @@ -722,6 +754,8 @@ void mlir::populateArithToEmitCPatterns(TypeConverter &typeConverter, ArithOpConversion, ArithOpConversion, ArithOpConversion, + BinaryUIOpConversion, + BinaryUIOpConversion, IntegerOpConversion, IntegerOpConversion, IntegerOpConversion, diff --git a/mlir/test/Conversion/ArithToEmitC/arith-to-emitc-unsupported.mlir b/mlir/test/Conversion/ArithToEmitC/arith-to-emitc-unsupported.mlir index 766ad4039335..ef0e71ee8673 100644 --- a/mlir/test/Conversion/ArithToEmitC/arith-to-emitc-unsupported.mlir +++ b/mlir/test/Conversion/ArithToEmitC/arith-to-emitc-unsupported.mlir @@ -134,3 +134,19 @@ func.func @arith_shrui_i1(%arg0: i1, %arg1: i1) { %shrui = arith.shrui %arg0, %arg1 : i1 return } + +// ----- + +func.func @arith_divui_vector(%arg0: vector<5xi32>, %arg1: vector<5xi32>) -> vector<5xi32> { + // expected-error @+1 {{failed to legalize operation 'arith.divui'}} + %divui = arith.divui %arg0, %arg1 : vector<5xi32> + return %divui: vector<5xi32> +} + +// ----- + +func.func @arith_remui_vector(%arg0: vector<5xi32>, %arg1: vector<5xi32>) -> vector<5xi32> { + // expected-error @+1 {{failed to legalize operation 'arith.remui'}} + %divui = arith.remui %arg0, %arg1 : vector<5xi32> + return %divui: vector<5xi32> +} diff --git a/mlir/test/Conversion/ArithToEmitC/arith-to-emitc.mlir b/mlir/test/Conversion/ArithToEmitC/arith-to-emitc.mlir index 858ccd117144..afd1198ede0f 100644 --- a/mlir/test/Conversion/ArithToEmitC/arith-to-emitc.mlir +++ b/mlir/test/Conversion/ArithToEmitC/arith-to-emitc.mlir @@ -717,3 +717,21 @@ func.func @arith_index_castui(%arg0: i32) -> i32 { return %int : i32 } + +// ----- + +func.func @arith_divui_remui(%arg0: i32, %arg1: i32) -> i32 { + // CHECK-LABEL: arith_divui_remui + // CHECK-SAME: (%[[Arg0:[^ ]*]]: i32, %[[Arg1:[^ ]*]]: i32) + // CHECK: %[[Conv0:.*]] = emitc.cast %[[Arg0]] : i32 to ui32 + // CHECK: %[[Conv1:.*]] = emitc.cast %[[Arg1]] : i32 to ui32 + // CHECK: %[[Div:.*]] = emitc.div %[[Conv0]], %[[Conv1]] : (ui32, ui32) -> ui32 + %div = arith.divui %arg0, %arg1 : i32 + + // CHECK: %[[Conv2:.*]] = emitc.cast %[[Arg0]] : i32 to ui32 + // CHECK: %[[Conv3:.*]] = emitc.cast %[[Arg1]] : i32 to ui32 + // CHECK: %[[Rem:.*]] = emitc.rem %[[Conv2]], %[[Conv3]] : (ui32, ui32) -> ui32 + %rem = arith.remui %arg0, %arg1 : i32 + + return %div : i32 +}