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