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

如何恢复并使用自行创建的TensorFlow简单模型?

如何恢复并使用你的TensorFlow模型

嘿,我看你已经搭建了一个基础的线性回归TensorFlow计算图,不过目前训练代码还没写完。要恢复模型的前提是先完成训练并把模型保存下来,之后再按步骤恢复使用。下面分阶段给你说明:

一、先补全训练代码并保存模型

你的代码里optimizer = tf.train.GradientDescentOptimiz...没写完,我先帮你补全训练逻辑,并且加上模型保存的步骤——这是后续恢复的基础:

import tensorflow as tf
tf.reset_default_graph()
x_data = [1,2,3]
y_data = [3,4,5]
X = tf.placeholder(tf.float32, name="X")
Y = tf.placeholder(tf.float32, name="Y")
W = tf.Variable(tf.random_uniform([1], -1.0, 1.0), name='W')
b = tf.Variable(tf.random_uniform([1], 0.0, 2.0), name='b')
hypothesis = tf.add(b, tf.multiply(X,W), name="op_restore")
saver = tf.train.Saver()  # 初始化Saver,用来保存/加载模型
cost = tf.reduce_mean(tf.square(hypothesis - Y))
# 补全优化器定义
optimizer = tf.train.GradientDescentOptimizer(learning_rate=0.01)
train_op = optimizer.minimize(cost)

# 启动会话开始训练
with tf.Session() as sess:
    sess.run(tf.global_variables_initializer())
    # 训练1000轮,每100轮打印一次损失值
    for step in range(1000):
        _, current_cost = sess.run([train_op, cost], feed_dict={X: x_data, Y: y_data})
        if step % 100 == 0:
            print(f"Step {step}, 当前损失: {current_cost:.4f}")
    # 训练完成后保存模型,会生成几个相关文件(.ckpt.data、.ckpt.index、.ckpt.meta、checkpoint)
    save_path = saver.save(sess, "./my_model.ckpt")
    print(f"模型已保存到路径: {save_path}")

二、恢复模型并进行预测

恢复模型有两种常用方式,你可以根据自己的需求选择:

方式1:重新构建相同计算图,再加载参数

这种方式需要你重新定义和训练时完全一致的计算图结构(包括所有张量的name不能变),然后加载保存好的参数:

import tensorflow as tf
tf.reset_default_graph()

# 严格复刻训练时的计算图结构
X = tf.placeholder(tf.float32, name="X")
W = tf.Variable(tf.random_uniform([1], -1.0, 1.0), name='W')
b = tf.Variable(tf.random_uniform([1], 0.0, 2.0), name='b')
hypothesis = tf.add(b, tf.multiply(X,W), name="op_restore")

saver = tf.train.Saver()

with tf.Session() as sess:
    # 加载保存的模型参数
    saver.restore(sess, "./my_model.ckpt")
    print("模型恢复成功!")
    # 用测试数据验证模型
    test_x = [4, 5, 6]
    predictions = sess.run(hypothesis, feed_dict={X: test_x})
    for x, pred in zip(test_x, predictions):
        print(f"输入{x},预测结果: {pred:.2f}")

方式2:直接加载保存的计算图(无需重新构建)

如果你不想重复写计算图代码,可以直接加载训练时保存的图结构,然后通过张量的name来获取对应的操作:

import tensorflow as tf

with tf.Session() as sess:
    # 先加载计算图的元数据(.meta文件)
    saver = tf.train.import_meta_graph('./my_model.ckpt.meta')
    # 再加载模型参数
    saver.restore(sess, "./my_model.ckpt")
    
    # 通过name获取图中的输入张量X和预测操作hypothesis
    X = sess.graph.get_tensor_by_name("X:0")
    hypothesis = sess.graph.get_tensor_by_name("op_restore:0")
    
    # 进行预测
    test_x = [4, 5, 6]
    predictions = sess.run(hypothesis, feed_dict={X: test_x})
    for x, pred in zip(test_x, predictions):
        print(f"输入{x},预测结果: {pred:.2f}")

一些注意点

  • 确保恢复模型时的文件路径正确,所有模型相关文件要放在指定路径下,否则会加载失败。
  • 如果用方式1恢复,计算图的结构和张量name必须和训练时完全一致,不然会找不到对应的变量。
  • 默认情况下saver.save()会生成.meta文件,如果你不想保存,可以初始化Saver时设置write_meta_graph=False,但这种情况下只能用方式1恢复模型。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.25 06:20:00