You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何在Jupyter不同笔记本中恢复TensorFlow计算图?

嘿,我瞅了你贴的TensorFlow模型保存和恢复的代码片段,估摸着你大概率在恢复环节踩坑了对吧?毕竟直接写saver.restore()却没先搞定计算图的问题,十有八九会报找不到变量或者操作的错。咱先把这事儿捋明白:

TensorFlow模型保存与恢复的正确姿势

一、先给你的保存代码补补细节

你的保存逻辑其实没大问题,但可以调整得更规范些,避免不必要的资源泄漏:

# 先确保你已经定义好了完整的计算图:比如X/Y占位符、train_step训练操作这些
saver = tf.train.Saver()  # 默认保存所有可训练变量
init_op = tf.global_variables_initializer()

# 用上下文管理器管理Session更稳妥,不用手动关闭
with tf.Session() as sess:
    sess.run(init_op)  # 初始化所有全局变量
    for ep in range(epoch):
        train_step.run(feed_dict={X: X_train, Y: y_train.reshape(-1,1)})
    # 保存后会生成3个文件:.meta(存计算图)、.data-xxx(存变量值)、checkpoint(记录最新模型路径)
    saver.save(sess, 'my_test_model')

小贴士:sess.as_default()不是不能用,但直接用with tf.Session()上下文管理器是更标准的写法,能自动释放资源。

二、恢复模型的两大正确方式(避开坑点)

你现在的恢复代码只写了saver.restore(),但核心问题是:恢复前必须让TensorFlow能找到和训练时完全一致的计算图结构!这里给你两种靠谱的实现方式:

方式1:复刻计算图结构,再恢复变量

如果是在新脚本里恢复,先把训练时的所有占位符、模型层、操作原封不动地重新定义一遍,再初始化Saver加载变量:

# 第一步:1:1复刻训练时的计算图
X = tf.placeholder(tf.float32, shape=[None, 你的特征维度])
Y = tf.placeholder(tf.float32, shape=[None, 1])
# 举个例子:假设你的模型是简单全连接层
W = tf.Variable(tf.random_normal([你的特征维度, 1]), name='weights')
b = tf.Variable(tf.zeros([1]), name='biases')
logits = tf.matmul(X, W) + b
train_step = tf.train.GradientDescentOptimizer(0.01).minimize(tf.reduce_mean(tf.square(logits - Y)))

# 第二步:加载模型变量
with tf.Session() as sess:
    saver = tf.train.Saver()
    saver.restore(sess, 'my_test_model')
    # 现在就能直接跑测试了
    test_result = sess.run(logits, feed_dict={X: X_test})
    print("测试输出:", test_result)

方式2:直接加载保存的计算图(更省心)

要是不想重新写一遍计算图,可以直接从.meta文件加载整个训练时的计算图,省事儿又不容易出错:

with tf.Session() as sess:
    # 从.meta文件加载完整计算图
    saver = tf.train.import_meta_graph('my_test_model.meta')
    # 自动加载最新的变量 checkpoint
    saver.restore(sess, tf.train.latest_checkpoint('./'))
    
    # 获取图里的张量/操作,注意要和训练时的name对应上
    graph = tf.get_default_graph()
    X = graph.get_tensor_by_name("Placeholder:0")  # 比如训练时X的默认name是Placeholder,后缀:0是TensorFlow自动加的
    logits = graph.get_tensor_by_name("add:0")  # 假设模型输出是add操作的结果
    
    # 执行测试
    test_output = sess.run(logits, feed_dict={X: X_test})
    print("测试结果:", test_output)

三、常见踩坑排查

  • 要是报NotFoundError,要么是计算图结构和训练时不一样,要么是获取张量时的name写错了。可以用graph.get_operations()打印所有操作的name核对。
  • 如果只想保存部分变量,给tf.train.Saver()传个变量列表就行,比如saver = tf.train.Saver([W, b])。

内容的提问来源于stack exchange,提问作者Solomiia

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.25 07:22:08