| 2424 | return true; |
| 2425 | } |
| 2426 | static bool LoadPool(TFNode* tf_node, TFGraph& tf_graph, StaticGraph* graph) |
| 2427 | { |
| 2428 | TFNode* input = tf_node->inputs[0]; |
| 2429 | StaticNode* node = tf_node->static_node; |
| 2430 | |
| 2431 | AddNodeInputTensor(node, input->static_tensor); |
| 2432 | |
| 2433 | PoolParam param = any_cast<PoolParam>(OpManager::GetOpDefParam("Pooling")); |
| 2434 | |
| 2435 | const tensorflow::NodeDef* node_def = tf_node->pb_defs[0]; |
| 2436 | tensorflow::AttrValue value; |
| 2437 | |
| 2438 | if(GetAttrValue(node_def, "ksize", value)) |
| 2439 | { |
| 2440 | param.kernel_h = value.list().i(1); |
| 2441 | param.kernel_w = value.list().i(2); |
| 2442 | } |
| 2443 | |
| 2444 | if(GetAttrValue(node_def, "strides", value)) |
| 2445 | { |
| 2446 | param.stride_h = value.list().i(1); |
| 2447 | param.stride_w = value.list().i(2); |
| 2448 | } |
| 2449 | |
| 2450 | if(GetAttrValue(node_def, "padding", value)) |
| 2451 | { |
| 2452 | if(value.s() == "VALID") |
| 2453 | { |
| 2454 | param.pad_h0 = 0; |
| 2455 | param.pad_h1 = 0; |
| 2456 | param.pad_w0 = 0; |
| 2457 | param.pad_w1 = 0; |
| 2458 | } |
| 2459 | else if(value.s() == "SAME") |
| 2460 | { |
| 2461 | param.pad_h0 = -1; |
| 2462 | param.pad_h1 = -1; |
| 2463 | param.pad_w0 = -1; |
| 2464 | param.pad_w1 = -1; |
| 2465 | } |
| 2466 | } |
| 2467 | |
| 2468 | if(tf_node->op == "AvgPool") |
| 2469 | { |
| 2470 | param.alg = kPoolAvg; |
| 2471 | } |
| 2472 | else if(tf_node->op == "MaxPool") |
| 2473 | { |
| 2474 | param.alg = kPoolMax; |
| 2475 | } |
| 2476 | |
| 2477 | StaticOp* op = CreateStaticOp(graph, "Pooling"); |
| 2478 | SetOperatorParam(op, param); |
| 2479 | SetNodeOp(node, op); |
| 2480 | |
| 2481 | return true; |
| 2482 | } |
| 2483 |
nothing calls this directly
no test coverage detected