python - google - ¿Cómo convertir el gráfico protobuf al formato de cable binario?
protobuf python install (2)
Tengo un método para convertir formato de cable binario a formato legible para humanos, pero no puedo hacer lo contrario de esto
import tensorflow as tf
from tensorflow.python.platform import gfile
def converter(filename):
with gfile.FastGFile(filename,''rb'') as f:
graph_def = tf.GraphDef()
graph_def.ParseFromString(f.read())
tf.import_graph_def(graph_def, name='''')
tf.train.write_graph(graph_def, ''pbtxt/'', ''protobuf.pb'', as_text=True)
return
Solo tengo que escribir el nombre del archivo para esto y funciona. Pero al hacer lo opuesto consigo
File "pb_to_pbtxt.py", line 16, in <module>
converter(''protobuf.pb'') # here you can write the name of the file to be converted
File "pb_to_pbtxt.py", line 11, in converter
graph_def.ParseFromString(f.read())
File "/usr/local/lib/python2.7/dist-packages/google/protobuf/message.py", line 185, in ParseFromString
self.MergeFromString(serialized)
File "/usr/local/lib/python2.7/dist-packages/google/protobuf/internal/python_message.py", line 1008, in MergeFromString
if self._InternalParse(serialized, 0, length) != length:
File "/usr/local/lib/python2.7/dist-packages/google/protobuf/internal/python_message.py", line 1034, in InternalParse
new_pos = local_SkipField(buffer, new_pos, end, tag_bytes)
File "/usr/local/lib/python2.7/dist-packages/google/protobuf/internal/decoder.py", line 868, in SkipField
return WIRETYPE_TO_SKIPPER[wire_type](buffer, pos, end)
File "/usr/local/lib/python2.7/dist-packages/google/protobuf/internal/decoder.py", line 838, in _RaiseInvalidWireType
raise _DecodeError(''Tag had invalid wire type.'')
Puede realizar la traducción inversa utilizando el módulo google.protobuf.text_format
:
import tensorflow as tf
from google.protobuf import text_format
def convert_pbtxt_to_graphdef(filename):
"""Returns a `tf.GraphDef` proto representing the data in the given pbtxt file.
Args:
filename: The name of a file containing a GraphDef pbtxt (text-formatted
`tf.GraphDef` protocol buffer data).
Returns:
A `tf.GraphDef` protocol buffer.
"""
with tf.gfile.FastGFile(filename, ''r'') as f:
graph_def = tf.GraphDef()
file_content = f.read()
# Merges the human-readable string in `file_content` into `graph_def`.
text_format.Merge(file_content, graph_def)
return graph_def
Puede usar tf.Graph.as_graph_def()
y luego SerializeToString()
Protobuf como sigue:
proto_graph = # obtained by calling tf.Graph.as_graph_def()
with open("my_graph.bin", "wb") as f:
f.write(proto_graph.SerializeToString())
Si solo quieres escribir el archivo y no te importa la codificación, también puedes usar tf.train.write_graph()
v = tf.Variable(0, name=''my_variable'')
sess = tf.Session()
tf.train.write_graph(sess.graph_def, ''/tmp/my-model'', ''train.pbtxt'')
Nota: Probado en TF 0.10, no estoy seguro acerca de las versiones anteriores.