如何恢复并使用自行创建的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
相关产品推荐
相关产品推荐

