| 40 | } |
| 41 | |
| 42 | Status PropagateShapes(const Graph& graph, |
| 43 | const std::map<int, InferredShape>& arg_shapes, |
| 44 | const std::vector<BackEdgeHelper::BackEdge>& back_edges, |
| 45 | ShapeRefiner* shape_refiner) { |
| 46 | std::map<const Node*, const Node*> merge_to_next_iteration; |
| 47 | for (const auto& e : back_edges) { |
| 48 | if (e.src->IsNextIteration() && e.dst->IsMerge()) { |
| 49 | merge_to_next_iteration[e.dst] = e.src; |
| 50 | } |
| 51 | } |
| 52 | |
| 53 | // Visits the nodes in topological order (reverse post-order), inferring |
| 54 | // shapes. |
| 55 | // TODO(phawkins): handle cyclic graphs. |
| 56 | std::vector<Node*> order; |
| 57 | GetReversePostOrder(graph, &order); |
| 58 | |
| 59 | for (Node* n : order) { |
| 60 | // Ignore the status returned by the shape_refiner. We want the best effort |
| 61 | // shapes, even if no shape function is registered for a node. |
| 62 | Status status = shape_refiner->AddNode(n); |
| 63 | if (!status.ok()) { |
| 64 | VLOG(1) << "Shape inference failed for node " << n->name() << ": " |
| 65 | << status; |
| 66 | } else { |
| 67 | shape_inference::InferenceContext* context = shape_refiner->GetContext(n); |
| 68 | for (int i = 0; i < n->num_outputs(); i++) { |
| 69 | shape_inference::ShapeHandle handle = context->output(i); |
| 70 | VLOG(4) << "Output " << i << " for node " << n->name() << ": " |
| 71 | << context->DebugString(handle); |
| 72 | } |
| 73 | } |
| 74 | |
| 75 | if (n->type_string() == "_Arg") { |
| 76 | int index; |
| 77 | TF_RETURN_IF_ERROR(GetNodeAttr(n->attrs(), "index", &index)); |
| 78 | auto it = arg_shapes.find(index); |
| 79 | if (it != arg_shapes.end()) { |
| 80 | const InferredShape& arg_shape = it->second; |
| 81 | shape_inference::InferenceContext* context = |
| 82 | shape_refiner->GetContext(n); |
| 83 | |
| 84 | if (arg_shape.handle_type != DT_INVALID) { |
| 85 | shape_inference::ShapeHandle handle; |
| 86 | TF_RETURN_IF_ERROR(context->MakeShapeFromPartialTensorShape( |
| 87 | arg_shape.handle_shape, &handle)); |
| 88 | |
| 89 | // Sets the shape and type of the variable's value. |
| 90 | context->set_output_handle_shapes_and_types( |
| 91 | 0, std::vector<shape_inference::ShapeAndType>{ |
| 92 | {handle, arg_shape.handle_type}}); |
| 93 | } |
| 94 | |
| 95 | shape_inference::ShapeHandle handle; |
| 96 | TF_RETURN_IF_ERROR( |
| 97 | context->MakeShapeFromPartialTensorShape(arg_shape.shape, &handle)); |
| 98 | TF_RETURN_IF_ERROR(shape_refiner->SetShape(n, 0, handle)); |
| 99 | } |
no test coverage detected