TensorFlow:不同图初始化模型却关联错误图的修正方法咨询
解决TensorFlow多图上下文下模型属性绑定错误的问题
我来帮你分析这个问题并给出解决方案:
问题重现
你编写了以下TensorFlow代码,期望在不同图上下文中创建的模型实例,其input_batch张量属于对应的图,但实际结果却不符合预期:
import tensorflow as tf import numpy as np class SimpleModel(): pass def declare_placeholders(self): self.input_batch = tf.placeholder(dtype=tf.int32, shape=[None, None], name='input_batch') SimpleModel.__declare_placeholders = classmethod(declare_placeholders) def init_model(self): self.__declare_placeholders() SimpleModel.__init__ = classmethod(init_model) g_1 = tf.Graph() with g_1.as_default(): model1 = SimpleModel() g_2 = tf.Graph() with g_2.as_default(): model2 = SimpleModel()
出现的异常现象:
- 预期
assert model1.input_batch.graph is g_1不会触发断言错误,但实际触发了 - 反而
assert model1.input_batch.graph is g_2成立,明明model1是在g_1的默认图上下文中初始化的
问题根源
问题出在你错误地使用了classmethod来定义__init__和__declare_placeholders方法:
__init__是Python类的实例初始化方法,本质上应该接收实例对象self作为第一个参数,而不是类对象。当你把它设为classmethod后,init_model里的self其实是SimpleModel类本身,而不是刚创建的实例。- 同样,
__declare_placeholders作为classmethod时,你给self(也就是类)绑定了input_batch属性——这意味着所有实例都会共享这个类属性。当你在g_2上下文创建model2时,会覆盖类上的input_batch,所以model1.input_batch实际上指向的是最后一次创建的、属于g_2的张量。
修正后的代码
只需要把初始化和占位符声明方法改回普通的实例方法即可:
import tensorflow as tf import numpy as np class SimpleModel(): def __init__(self): # 实例初始化时调用占位符声明方法 self.declare_placeholders() def declare_placeholders(self): # 给当前实例绑定input_batch属性,而非类 self.input_batch = tf.placeholder(dtype=tf.int32, shape=[None, None], name='input_batch') g_1 = tf.Graph() with g_1.as_default(): model1 = SimpleModel() g_2 = tf.Graph() with g_2.as_default(): model2 = SimpleModel() # 现在两个断言都会正常通过 assert model1.input_batch.graph is g_1 assert model2.input_batch.graph is g_2
修正说明
- 将
__init__恢复为普通实例方法,负责初始化当前实例并调用占位符声明逻辑 declare_placeholders改为实例方法,给当前创建的实例绑定独立的input_batch属性- 这样每个模型实例都会拥有自己专属的张量,创建时所在的图上下文会正确关联到该张量上,不会出现跨图混淆的问题
内容的提问来源于stack exchange,提问作者evgenii ershenko
相关产品推荐
相关产品推荐

