MCPcopy Create free account
hub / github.com/DeepRec-AI/DeepRec / Run

Method Run

tensorflow/lite/toco/graph_transformations/identify_hardswish.cc:38–114  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

36using 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 }

Callers

nothing calls this directly

Calls 15

GetOpWithInputFunction · 0.85
CheckArrayIsScalarFloatFunction · 0.85
LogNameFunction · 0.85
DeleteOpAndArraysFunction · 0.85
pop_backMethod · 0.80
IsBinaryOpFunction · 0.70
beginMethod · 0.45
getMethod · 0.45
push_backMethod · 0.45
insertMethod · 0.45
endMethod · 0.45

Tested by

no test coverage detected