本文摘自php中文网,作者不言,侵删。
本篇文章主要介绍了将TensorFlow的网络导出为单个文件的方法,现在分享给大家,也给大家做个参考。一起过来看看吧有时候,我们需要将TensorFlow的模型导出为单个文件(同时包含模型架构定义与权重),方便在其他地方使用(如在c++中部署网络)。利用tf.train.write_graph()默认情况下只导出了网络的定义(没有权重),而利用tf.train.Saver().save()导出的文件graph_def与权重是分离的,因此需要采用别的方法。
我们知道,graph_def文件中没有包含网络中的Variable值(通常情况存储了权重),但是却包含了constant值,所以如果我们能把Variable转换为constant,即可达到使用一个文件同时存储网络架构与权重的目标。
我们可以采用以下方式冻结权重并保存网络:
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 | import tensorflow as tf
from tensorflow.python.framework.graph_util import convert_variables_to_constants
a = tf.Variable([[ 3 ],[ 4 ]], dtype = tf.float32, name = 'a' )
b = tf.Variable( 4 , dtype = tf.float32, name = 'b' )
output = tf.add(a, b, name = 'out' )
with tf.Session() as sess:
sess.run(tf.global_variables_initializer())
graph = convert_variables_to_constants(sess, sess.graph_def, [ "out" ])
tf.train.write_graph(graph, '.' , 'graph.pb' , as_text = False )
|
当恢复网络时,可以使用如下方式:
1 2 3 4 5 6 7 | import tensorflow as tf
with tf.Session() as sess:
with open ( './graph.pb' , 'rb' ) as f:
graph_def = tf.GraphDef()
graph_def.ParseFromString(f.read())
output = tf.import_graph_def(graph_def, return_elements = [ 'out:0' ])
print (sess.run(output))
|
输出结果为:
[array([[ 7.],
[ 8.]], dtype=float32)]
可以看到之前的权重确实保存了下来!!
问题来了,我们的网络需要能有一个输入自定义数据的接口啊!不然这玩意有什么用。。别急,当然有办法。
1 2 3 4 5 6 7 8 9 10 11 | import tensorflow as tf
from tensorflow.python.framework.graph_util import convert_variables_to_constants
a = tf.Variable([[ 3 ],[ 4 ]], dtype = tf.float32, name = 'a' )
b = tf.Variable( 4 , dtype = tf.float32, name = 'b' )
input_tensor = tf.placeholder(tf.float32, name = 'input' )
output = tf.add((a + b), input_tensor, name = 'out' )
with tf.Session() as sess:
sess.run(tf.global_variables_initializer())
graph = convert_variables_to_constants(sess, sess.graph_def, [ "out" ])
tf.train.write_graph(graph, '.' , 'graph.pb' , as_text = False )
|
用上述代码重新保存网络至graph.pb,这次我们有了一个输入placeholder,下面来看看怎么恢复网络并输入自定义数据。
1 2 3 4 5 6 7 8 | import tensorflow as tf
with tf.Session() as sess:
with open ( './graph.pb' , 'rb' ) as f:
graph_def = tf.GraphDef()
graph_def.ParseFromString(f.read())
output = tf.import_graph_def(graph_def, input_map = { 'input:0' : 4. }, return_elements = [ 'out:0' ], name = 'a' )
print (sess.run(output))
|
输出结果为:
[array([[ 11.],
[ 12.]], dtype=float32)]
可以看到结果没有问题,当然在input_map那里可以替换为新的自定义的placeholder,如下所示:
1 2 3 4 5 6 7 8 9 10 | import tensorflow as tf
new_input = tf.placeholder(tf.float32, shape = ())
with tf.Session() as sess:
with open ( './graph.pb' , 'rb' ) as f:
graph_def = tf.GraphDef()
graph_def.ParseFromString(f.read())
output = tf.import_graph_def(graph_def, input_map = { 'input:0' :new_input}, return_elements = [ 'out:0' ], name = 'a' )
print (sess.run(output, feed_dict = {new_input: 4 }))
|
看看输出,同样没有问题。
[array([[ 11.],
[ 12.]], dtype=float32)]
另外需要说明的一点是,在利用tf.train.write_graph写网络架构的时候,如果令as_text=True了,则在导入网络的时候,需要做一点小修改。
1 2 3 4 5 6 7 8 9 10 11 | import tensorflow as tf
from google.protobuf import text_format
with tf.Session() as sess:
with open ( './graph.pb' , 'r' ) as f:
graph_def = tf.GraphDef()
text_format.Merge(f.read(), graph_def)
output = tf.import_graph_def(graph_def, return_elements = [ 'out:0' ])
print (sess.run(output))
|
相关推荐:
TensorFlow安装以及对jupyter notebook配置详解
以上就是将TensorFlow的模型网络导出为单个文件的方法的详细内容,更多文章请关注木庄网络博客!!
相关阅读 >>
Python 的& 表示什么
Python能做游戏吗
Python字典的值可以是字典吗
使用Python时多少有人走过的坑!避险!
Python中numpy是什么
Python中spyder怎么安装
拿下 Python中的文件操作
为什么选择用Python做爬虫
Python怎么找出最大数
用户输入输出和while循环
更多相关阅读请进入《Python》频道 >>
人民邮电出版社
python入门书籍,非常畅销,超高好评,python官方公认好书。
转载请注明出处:木庄网络博客 » 将TensorFlow的模型网络导出为单个文件的方法