| 11 | { |
| 12 | |
| 13 | void ArgMaxFP32Test(int axisValue) |
| 14 | { |
| 15 | // Set input data |
| 16 | std::vector<int32_t> inputShape { 1, 3, 2, 4 }; |
| 17 | std::vector<int32_t> outputShape { 1, 3, 4 }; |
| 18 | std::vector<int32_t> axisShape { 1 }; |
| 19 | |
| 20 | std::vector<float> inputValues = { 1.0f, 2.0f, 3.0f, 4.0f, |
| 21 | 5.0f, 6.0f, 7.0f, 8.0f, |
| 22 | |
| 23 | 10.0f, 20.0f, 30.0f, 40.0f, |
| 24 | 50.0f, 60.0f, 70.0f, 80.0f, |
| 25 | |
| 26 | 100.0f, 200.0f, 300.0f, 400.0f, |
| 27 | 500.0f, 600.0f, 700.0f, 800.0f }; |
| 28 | |
| 29 | std::vector<int32_t> expectedOutputValues = { 1, 1, 1, 1, |
| 30 | 1, 1, 1, 1, |
| 31 | 1, 1, 1, 1 }; |
| 32 | |
| 33 | ArgMinMaxTest<float, int32_t>(tflite::BuiltinOperator_ARG_MAX, |
| 34 | ::tflite::TensorType_FLOAT32, |
| 35 | inputShape, |
| 36 | axisShape, |
| 37 | outputShape, |
| 38 | inputValues, |
| 39 | expectedOutputValues, |
| 40 | axisValue, |
| 41 | ::tflite::TensorType_INT32); |
| 42 | } |
| 43 | |
| 44 | void ArgMinFP32Test(int axisValue) |
| 45 | { |