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

Function PropagateShapes

tensorflow/compiler/jit/shape_inference.cc:42–142  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

40}
41
42Status 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 }

Callers 2

InferShapesFunction · 0.85
InferStaticallyMethod · 0.85

Calls 15

GetReversePostOrderFunction · 0.85
IsNextIterationMethod · 0.80
IsMergeMethod · 0.80
input_nodeMethod · 0.80
IsIdentityMethod · 0.80
IsSwitchMethod · 0.80
nameMethod · 0.65
outputMethod · 0.65
GetNodeAttrFunction · 0.50

Tested by

no test coverage detected