如何在不同Graph中使用另一Graph中的embedding张量?
解决TensorFlow跨Graph使用张量的问题
这个坑我之前踩过!TensorFlow里每个张量都是牢牢绑定在它被创建的Graph上的,直接跨图调用肯定会报那个“tensor must be from the same graph”的错误。下面给你几个实用的解决办法,按需选择:
方法1:将原Graph的结构导入目标Graph
如果你的embedding张量还依赖graph1里的其他运算(不是单纯的预初始化变量),这个方法最合适——把graph1的图结构整体导入到graph2中,就能拿到属于graph2的对应张量了:
# 假设你已经构建好graph1和其中的embedding张量 with graph1.as_default(): embedding = tf.get_variable("embedding", shape=[100, 50]) # 导出graph1的图定义 graph1_def = graph1.as_graph_def() # 在graph2中导入graph1的定义并获取embedding with graph2.as_default(): # 导入时指定要返回的张量,注意名称要加":0"(TensorFlow中张量的默认输出后缀) imported_embedding = tf.import_graph_def( graph1_def, return_elements=["embedding:0"], name="" # 为空表示不添加前缀,避免名称混乱 )[0] # 现在imported_embedding属于graph2,随便用就行 print(imported_embedding.graph is graph2) # 会输出True
方法2:提取张量值后在目标Graph中复用
如果只需要embedding的数值,不需要保留原计算图的依赖关系,直接把值取出来再放到graph2里更简单:
import tensorflow as tf # 构建graph1 graph1 = tf.Graph() with graph1.as_default(): embedding = tf.get_variable("embedding", shape=[100, 50], initializer=tf.random_normal_initializer()) init_op = tf.global_variables_initializer() # 在graph1的会话中计算出embedding的具体值 with tf.Session(graph=graph1) as sess: sess.run(init_op) embedding_np = sess.run(embedding) # 构建graph2并使用这个值 graph2 = tf.Graph() with graph2.as_default(): # 方式A:用变量直接初始化,后续可以当作普通变量用 embedding_in_graph2 = tf.get_variable( "embedding_in_graph2", shape=[100, 50], initializer=tf.constant_initializer(embedding_np) ) # 方式B:用占位符,运行时喂入值(适合需要动态更新的场景) embedding_ph = tf.placeholder(tf.float32, shape=[100, 50]) demo_op = tf.reduce_sum(embedding_ph) # 在graph2的会话中验证 with tf.Session(graph=graph2) as sess: sess.run(tf.global_variables_initializer()) print(sess.run(embedding_in_graph2[0][:5])) # 输出和graph1中embedding的前5个值一致 print(sess.run(demo_op, feed_dict={embedding_ph: embedding_np}))
额外提示
- 用
get_variable的时候要注意,两个Graph里的变量名称最好不要重复,避免不必要的冲突; - 如果你的TensorFlow版本比较新(2.x以上),其实默认是即时执行模式,Graph的概念没那么强,但如果是兼容1.x的代码,上面的方法依然适用。
内容的提问来源于stack exchange,提问作者xiyan
相关产品推荐
相关产品推荐

