TensorFlow:如何在两个独立图中保持变量名称一致?
解决TensorFlow多图变量名自动加后缀的问题
这问题我之前踩过坑!TensorFlow默认会在同一个图内给重复的变量名自动加_1/_2后缀避免冲突,但咱们要实现多图变量名完全一致,核心是利用不同图的独立命名空间特性,配合显式指定变量名,具体有两个靠谱方案:
方案1:在独立图上下文里显式指定变量name
每个计算图都是独立的命名空间,只要你在创建变量时显式指定name参数,并且把变量定义严格包裹在对应图的as_default()上下文里,就不会出现自动加后缀的情况。
示例代码:
import tensorflow as tf # 构建第一个图 g1 = tf.Graph() with g1.as_default(): # 显式指定name="sample_data",不要依赖TF自动生成 sample_data = tf.Variable([0, 0, 0, 0], name="sample_data") # 生成初始化操作 init_op = tf.global_variables_initializer() # 构建第二个图,完全复用同名变量 g2 = tf.Graph() with g2.as_default(): # 同样指定name="sample_data",因为在独立图中,不会和g1的变量冲突 sample_data = tf.Variable([0, 0, 0, 0], name="sample_data") init_op = tf.global_variables_initializer() # 测试通过变量名初始化 with tf.Session(graph=g1) as sess: sess.run(init_op, feed_dict={g1.get_tensor_by_name("sample_data:0"): [1,2,3,4]}) print("g1的sample_data值:", sess.run(sample_data)) # 输出 [1 2 3 4] with tf.Session(graph=g2) as sess: sess.run(init_op, feed_dict={g2.get_tensor_by_name("sample_data:0"): [5,6,7,8]}) print("g2的sample_data值:", sess.run(sample_data)) # 输出 [5 6 7 8]
方案2:用tf.get_variable()配合独立图作用域
如果你的代码习惯用tf.get_variable()来创建变量,同样要确保每个变量定义在对应图的上下文里,显式指定变量名即可,完全不用担心跨图冲突:
import tensorflow as tf g1 = tf.Graph() with g1.as_default(): # 可选:用variable_scope做分组,不影响跨图同名 with tf.variable_scope("my_scope"): sample_data = tf.get_variable( name="sample_data", shape=[4], initializer=tf.zeros_initializer() ) init_op = tf.global_variables_initializer() g2 = tf.Graph() with g2.as_default(): with tf.variable_scope("my_scope"): sample_data = tf.get_variable( name="sample_data", shape=[4], initializer=tf.zeros_initializer() ) init_op = tf.global_variables_initializer() # 验证变量名一致 print("g1的变量名:", g1.get_tensor_by_name("my_scope/sample_data:0").name) print("g2的变量名:", g2.get_tensor_by_name("my_scope/sample_data:0").name) # 输出都是 my_scope/sample_data:0
避坑提醒
千万不要犯这个低级错误:把多个图的变量定义都放在默认图的上下文里(也就是没加with graph.as_default()),这种情况下TF会把所有变量都塞到默认图里,自然会出现_1后缀。必须严格给每个图的定义代码加上专属的上下文包裹。
内容的提问来源于stack exchange,提问作者Hoeze
相关产品推荐
相关产品推荐

