brintos

brintos / llvm-project-archived public Read only

0
0
Text · 3.2 KiB · c97441a Raw
60 lines · plain
1// RUN: mlir-opt %s -inline | FileCheck %s2 3// These tests verify that regions with operations from TOSA dialect4// can be inlined.5 6// CHECK-LABEL: func @inlined_if_fn7// Check that both the calls and the functions are eliminated after inlining:8// CHECK-NOT: @add9// CHECK-NOT: @sub10func.func @inlined_if_fn(%arg0: tensor<f32>, %arg1: tensor<f32>, %arg2: tensor<i1>) -> tensor<f32> {11  %0 = "tosa.cond_if"(%arg2, %arg0, %arg1) ({12  ^bb0(%arg3: tensor<f32>, %arg4: tensor<f32>):13    %1 = call @add(%arg3, %arg4) : (tensor<f32>, tensor<f32>) -> tensor<f32>14    "tosa.yield"(%1) : (tensor<f32>) -> ()15  },  {16  ^bb0(%arg3: tensor<f32>, %arg4: tensor<f32>):17    %1 = call @sub(%arg3, %arg4) : (tensor<f32>, tensor<f32>) -> tensor<f32>18    "tosa.yield"(%1) : (tensor<f32>) -> ()19  }) : (tensor<i1>, tensor<f32>, tensor<f32>) -> tensor<f32>20  return %0 : tensor<f32>21}22func.func private @add(%arg0: tensor<f32>, %arg1: tensor<f32>) -> tensor<f32> {23  %0 = "tosa.add"(%arg0, %arg1) : (tensor<f32>, tensor<f32>) -> tensor<f32>24  return %0 : tensor<f32>25}26func.func private @sub(%arg0: tensor<f32>, %arg1: tensor<f32>) -> tensor<f32> {27  %0 = "tosa.sub"(%arg0, %arg1) : (tensor<f32>, tensor<f32>) -> tensor<f32>28  return %0 : tensor<f32>29}30 31// -----32 33// CHECK-LABEL: func @inlined_while_fn34func.func @inlined_while_fn(%arg0: tensor<i32>, %arg1: tensor<i32>, %arg2: tensor<i32>, %arg3: tensor<10xi32>) -> tensor<10xi32> {35  // Check that calls are inlined and functions eliminated:36  // CHECK-NOT: @while37  %1:4 = "tosa.while_loop"(%arg0, %arg1, %arg2, %arg3) ({38  ^bb0(%arg4: tensor<i32>, %arg5: tensor<i32>, %arg6: tensor<i32>, %arg7: tensor<10xi32>):39    %2 = call @while_cond_40(%arg4, %arg5, %arg6, %arg7) : (tensor<i32>, tensor<i32>, tensor<i32>, tensor<10xi32>) -> tensor<i1>40    "tosa.yield"(%2) : (tensor<i1>) -> ()41  },  {42  ^bb0(%arg4: tensor<i32>, %arg5: tensor<i32>, %arg6: tensor<i32>, %arg7: tensor<10xi32>):43    %2:4 = call @while_body_50(%arg4, %arg5, %arg6, %arg7) : (tensor<i32>, tensor<i32>, tensor<i32>, tensor<10xi32>) -> (tensor<i32>, tensor<i32>, tensor<i32>, tensor<10xi32>)44    "tosa.yield"(%2#0, %2#1, %2#2, %2#3) : (tensor<i32>, tensor<i32>, tensor<i32>, tensor<10xi32>) -> ()45  }) : (tensor<i32>, tensor<i32>, tensor<i32>, tensor<10xi32>) -> (tensor<i32>, tensor<i32>, tensor<i32>, tensor<10xi32>)46  return %1#3 : tensor<10xi32>47}48func.func private @while_body_50(%arg0: tensor<i32>, %arg1: tensor<i32>, %arg2: tensor<i32>, %arg3: tensor<10xi32>) -> (tensor<i32>, tensor<i32>, tensor<i32>, tensor<10xi32>) {49  %1 = "tosa.add"(%arg0, %arg1) : (tensor<i32>, tensor<i32>) -> tensor<i32>50  %4 = "tosa.const_shape"() {values = dense<1> : tensor<1xindex>} : () -> !tosa.shape<1>51  %3 = "tosa.reshape"(%1, %4) : (tensor<i32>, !tosa.shape<1>) -> tensor<1xi32>52  %2 = "tosa.add"(%arg3, %3) : (tensor<10xi32>, tensor<1xi32>) -> tensor<10xi32>53  return %1, %arg1, %arg2, %2: tensor<i32>, tensor<i32>, tensor<i32>, tensor<10xi32>54}55func.func private @while_cond_40(%arg0: tensor<i32>, %arg1: tensor<i32>, %arg2: tensor<i32>, %arg3: tensor<10xi32>) -> tensor<i1> {56  %0 = "tosa.greater_equal"(%arg0, %arg1) : (tensor<i32>, tensor<i32>) -> tensor<i1>57  %1 = "tosa.logical_not"(%0) : (tensor<i1>) -> tensor<i1>58  return %1 : tensor<i1>59}60