• Home
  • Line#
  • Scopes#
  • Navigate#
  • Raw
  • Download
1# Copyright 2015 The TensorFlow Authors. All Rights Reserved.
2#
3# Licensed under the Apache License, Version 2.0 (the "License");
4# you may not use this file except in compliance with the License.
5# You may obtain a copy of the License at
6#
7#     http://www.apache.org/licenses/LICENSE-2.0
8#
9# Unless required by applicable law or agreed to in writing, software
10# distributed under the License is distributed on an "AS IS" BASIS,
11# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12# See the License for the specific language governing permissions and
13# limitations under the License.
14# ==============================================================================
15
16"""Utility functions for reading/writing graphs."""
17import os
18import os.path
19
20from google.protobuf import text_format
21from tensorflow.python.framework import ops
22from tensorflow.python.lib.io import file_io
23from tensorflow.python.util.tf_export import tf_export
24
25
26@tf_export('io.write_graph', v1=['io.write_graph', 'train.write_graph'])
27def write_graph(graph_or_graph_def, logdir, name, as_text=True):
28  """Writes a graph proto to a file.
29
30  The graph is written as a text proto unless `as_text` is `False`.
31
32  ```python
33  v = tf.Variable(0, name='my_variable')
34  sess = tf.compat.v1.Session()
35  tf.io.write_graph(sess.graph_def, '/tmp/my-model', 'train.pbtxt')
36  ```
37
38  or
39
40  ```python
41  v = tf.Variable(0, name='my_variable')
42  sess = tf.compat.v1.Session()
43  tf.io.write_graph(sess.graph, '/tmp/my-model', 'train.pbtxt')
44  ```
45
46  Args:
47    graph_or_graph_def: A `Graph` or a `GraphDef` protocol buffer.
48    logdir: Directory where to write the graph. This can refer to remote
49      filesystems, such as Google Cloud Storage (GCS).
50    name: Filename for the graph.
51    as_text: If `True`, writes the graph as an ASCII proto.
52
53  Returns:
54    The path of the output proto file.
55  """
56  if isinstance(graph_or_graph_def, ops.Graph):
57    graph_def = graph_or_graph_def.as_graph_def()
58  else:
59    graph_def = graph_or_graph_def
60
61  # gcs does not have the concept of directory at the moment.
62  if not logdir.startswith('gs:'):
63    file_io.recursive_create_dir(logdir)
64  path = os.path.join(logdir, name)
65  if as_text:
66    file_io.atomic_write_string_to_file(path,
67                                        text_format.MessageToString(
68                                            graph_def, float_format=''))
69  else:
70    file_io.atomic_write_string_to_file(
71        path, graph_def.SerializeToString(deterministic=True))
72  return path
73