| 36 | using util::IsBinaryOp; |
| 37 | |
| 38 | ::tensorflow::Status IdentifyHardSwish::Run(Model* model, std::size_t op_index, |
| 39 | bool* modified) { |
| 40 | *modified = false; |
| 41 | const auto add_with_relu6_op_it = (model->operators.begin() + op_index); |
| 42 | const auto add_with_relu6_op = add_with_relu6_op_it->get(); |
| 43 | if (!util::IsBinaryOp(add_with_relu6_op, OperatorType::kAdd, |
| 44 | FusedActivationFunctionType::kRelu6)) { |
| 45 | return ::tensorflow::Status::OK(); |
| 46 | } |
| 47 | std::vector<const Operator*> ops; |
| 48 | ops.push_back(add_with_relu6_op); |
| 49 | const auto* mul_op = GetOpWithInput(*model, add_with_relu6_op->outputs[0]); |
| 50 | ops.push_back(mul_op); |
| 51 | |
| 52 | if (mul_op->type == OperatorType::kFakeQuant) { |
| 53 | mul_op = GetOpWithInput(*model, mul_op->outputs[0]); |
| 54 | ops.push_back(mul_op); |
| 55 | } |
| 56 | if (!IsBinaryOp(mul_op, OperatorType::kMul)) { |
| 57 | return ::tensorflow::Status::OK(); |
| 58 | } |
| 59 | |
| 60 | const auto* output_op = GetOpWithInput(*model, mul_op->outputs[0]); |
| 61 | ops.push_back(output_op); |
| 62 | if (output_op->type == OperatorType::kFakeQuant) { |
| 63 | output_op = GetOpWithInput(*model, output_op->outputs[0]); |
| 64 | ops.push_back(output_op); |
| 65 | } |
| 66 | if (!IsBinaryOp(output_op, OperatorType::kMul)) { |
| 67 | return ::tensorflow::Status::OK(); |
| 68 | } |
| 69 | const auto add_3_tensor = |
| 70 | util::GetSingleScalarInputIndexOfBinaryOp(model, add_with_relu6_op, 3.0f); |
| 71 | if (add_3_tensor < 0) { |
| 72 | // Expected 3.0f got something else.; |
| 73 | return ::tensorflow::Status::OK(); |
| 74 | } |
| 75 | const auto input_tensor_name = add_with_relu6_op->inputs[1 - add_3_tensor]; |
| 76 | |
| 77 | // Now we verify that the 3 mul arguments are respectively: |
| 78 | // 1. non-constant input of add_with_relu6_op |
| 79 | // 2. 1/6 |
| 80 | // 3. (and add_with_relu6_op[0].outputs[0] - which we already know!) |
| 81 | std::vector<string> mul_inputs = mul_op->inputs; |
| 82 | mul_inputs.insert(mul_inputs.end(), output_op->inputs.begin(), |
| 83 | output_op->inputs.end()); |
| 84 | |
| 85 | // 1. Check that we have the input tensor as one of the multiplicants |
| 86 | if (std::find(mul_inputs.begin(), mul_inputs.end(), input_tensor_name) == |
| 87 | mul_inputs.end()) { |
| 88 | // Input tensor not found! << input_tensor_name << std::endl; |
| 89 | return ::tensorflow::Status::OK(); |
| 90 | } |
| 91 | // 2. Find 1/6 |
| 92 | bool found = false; |
| 93 | for (const auto& input : mul_inputs) { |
| 94 | found |= util::CheckArrayIsScalarFloat(model, input, 1.f / 6.f); |
| 95 | } |
nothing calls this directly
no test coverage detected