TensorFlow跨类命名空间问题:复用模型图而非新建实例
解决TensorFlow多实例复用模型计算图的问题
先给你点透问题根源:在TensorFlow 1.x环境下,每个Test实例初始化时默认会创建独立的计算图,而你保存的模型是和第一个实例的计算图绑定的。第二个实例用自己的新图加载模型时,找不到原模型对应的张量、操作节点,自然会抛出异常。
下面给你几个实用的解决思路,按需选用:
思路1:让所有Test实例共享默认计算图
最简单的方式是强制所有实例复用全局默认图。你只需要在实例化类时,把代码放在默认图的上下文管理器里:
import tensorflow as tf import numpy as np import matplotlib.pyplot as plt # 获取全局默认计算图,所有实例共用它 default_graph = tf.get_default_graph() X = np.random.random((10000, 2)) y = (2*X[:, 0] + 3*X[:, 1]).reshape(-1, 1) inpShape = (None, 2) outShape = (None, 1) layers = [7, 1] activations = [tf.sigmoid, None] # 第一个实例:在默认图上下文里构建模型 with default_graph.as_default(): t = Test(inpShape, outShape, layers, activations) t.fit(X, y, 10000) yHat = t.predict(X, restorePoint=t.restorePoints[-1]) plt.plot(yHat, y, '.', label='original') # 第二个实例:同样用默认图加载模型 with default_graph.as_default(): t1 = Test(inpShape, outShape, layers, activations) yHat1 = t1.predict(X, restorePoint=t.restorePoints[-1]) plt.plot(yHat1, y, '.', label='copied') plt.legend() plt.show()
思路2:把计算图作为参数传入Test类
更灵活的方式是给Test类加一个graph参数,显式指定要使用的计算图,完全掌控图的复用逻辑。
先修改Test类的构造函数:
class Test: def __init__(self, inpShape, outShape, layers, activations, graph=None): # 优先用传入的图,没有则用默认图 self.graph = graph or tf.get_default_graph() # 所有模型构建逻辑都放在指定图的上下文里 with self.graph.as_default(): self._build_model(inpShape, outShape, layers, activations) # ... 其他初始化代码(比如定义损失、优化器) def _build_model(self, inpShape, outShape, layers, activations): # 把原来的模型层构建代码移到这里,确保在指定图中执行 self.inputs = tf.placeholder(tf.float32, shape=inpShape, name='inputs') # ... 后续的隐藏层、输出层定义逻辑
然后测试代码可以这样写:
# 创建一个共享计算图 shared_graph = tf.Graph() # 第一个实例绑定共享图 t = Test(inpShape, outShape, layers, activations, graph=shared_graph) t.fit(X, y, 10000) yHat = t.predict(X, restorePoint=t.restorePoints[-1]) # 第二个实例同样用这个共享图 t1 = Test(inpShape, outShape, layers, activations, graph=shared_graph) yHat1 = t1.predict(X, restorePoint=t.restorePoints[-1])
思路3:类内部维护共享图(进阶)
如果想让Test类自动处理图的复用,可以在类里定义一个类级别的共享图,所有实例默认复用它:
class Test: # 类级别的共享计算图,所有实例共用 _shared_graph = tf.Graph() def __init__(self, inpShape, outShape, layers, activations): with Test._shared_graph.as_default(): self._build_model(inpShape, outShape, layers, activations) # ... 其他初始化逻辑
这样不管你实例化多少个Test对象,都会共用同一个计算图,彻底避免图不匹配的问题。
额外提示
如果你的项目允许切换到TensorFlow 2.x,这个问题会自动消失——TF2.x默认使用即时执行模式,没有计算图隔离的限制,模型保存和加载会更简洁直观。
内容的提问来源于stack exchange,提问作者ssm
相关产品推荐
相关产品推荐

