如何在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
相关产品推荐
相关产品推荐

