基于TensorFlow的MNIST预训练模型恢复后预测方法咨询
用预训练的MNIST深度神经网络做预测的完整方案
我来帮你搞定预训练模型预测的问题!先把你没写完的模型定义补全,再一步步教你保存模型、恢复模型并执行实际预测。
第一步:补全深度神经网络模型定义
先把你的DeepNN函数补充完整,这是模型的核心结构,和你之前的代码完全衔接:
import tensorflow as tf from tensorflow.examples.tutorials.mnist import input_data mnist = input_data.read_data_sets('tmp/data/', one_hot=True) n_nodes_hl1 = n_nodes_hl2 = n_nodes_hl3 = 500 n_classes, batch_size = 10, 100 x = tf.placeholder(tf.float32, shape=(None, 784)) y = tf.placeholder(tf.float32) def DeepNN(data): # 定义各层权重与偏置变量 hidden_1_layer = {'weights': tf.Variable(tf.random_normal([784, n_nodes_hl1])), 'biases': tf.Variable(tf.random_normal([n_nodes_hl1]))} hidden_2_layer = {'weights': tf.Variable(tf.random_normal([n_nodes_hl1, n_nodes_hl2])), 'biases': tf.Variable(tf.random_normal([n_nodes_hl2]))} hidden_3_layer = {'weights': tf.Variable(tf.random_normal([n_nodes_hl2, n_nodes_hl3])), 'biases': tf.Variable(tf.random_normal([n_nodes_hl3]))} output_layer = {'weights': tf.Variable(tf.random_normal([n_nodes_hl3, n_classes])), 'biases': tf.Variable(tf.random_normal([n_classes]))} # 计算各层输出(带ReLU激活) l1 = tf.add(tf.matmul(data, hidden_1_layer['weights']), hidden_1_layer['biases']) l1 = tf.nn.relu(l1) l2 = tf.add(tf.matmul(l1, hidden_2_layer['weights']), hidden_2_layer['biases']) l2 = tf.nn.relu(l2) l3 = tf.add(tf.matmul(l2, hidden_3_layer['weights']), hidden_3_layer['biases']) l3 = tf.nn.relu(l3) output = tf.add(tf.matmul(l3, output_layer['weights']), output_layer['biases']) return output
第二步:训练模型并保存预训练参数
在训练代码里加入模型保存逻辑,用tf.train.Saver()持久化训练好的参数:
# 初始化模型、损失函数与优化器 prediction = DeepNN(x) cost = tf.reduce_mean(tf.nn.softmax_cross_entropy_with_logits(logits=prediction, labels=y)) optimizer = tf.train.AdamOptimizer().minimize(cost) # 初始化Saver对象(用于保存/恢复模型) saver = tf.train.Saver() # 定义模型保存路径 save_path = "./mnist_deepnn_model" # 启动训练流程 epochs = 10 with tf.Session() as sess: sess.run(tf.global_variables_initializer()) for epoch in range(epochs): epoch_loss = 0 for _ in range(int(mnist.train.num_examples/batch_size)): epoch_x, epoch_y = mnist.train.next_batch(batch_size) _, c = sess.run([optimizer, cost], feed_dict={x: epoch_x, y: epoch_y}) epoch_loss += c print(f'Epoch {epoch+1} 完成,总损失: {epoch_loss:.2f}') # 训练结束后保存模型 saved_path = saver.save(sess, save_path) print(f"模型已保存至: {saved_path}")
第三步:恢复预训练模型并执行预测
现在可以单独写一段代码,加载保存的模型,对MNIST图片做预测:
# 注意:这里必须重新定义和训练时完全一致的模型结构 def DeepNN(data): hidden_1_layer = {'weights': tf.Variable(tf.random_normal([784, n_nodes_hl1])), 'biases': tf.Variable(tf.random_normal([n_nodes_hl1]))} hidden_2_layer = {'weights': tf.Variable(tf.random_normal([n_nodes_hl1, n_nodes_hl2])), 'biases': tf.Variable(tf.random_normal([n_nodes_hl2]))} hidden_3_layer = {'weights': tf.Variable(tf.random_normal([n_nodes_hl2, n_nodes_hl3])), 'biases': tf.Variable(tf.random_normal([n_nodes_hl3]))} output_layer = {'weights': tf.Variable(tf.random_normal([n_nodes_hl3, n_classes])), 'biases': tf.Variable(tf.random_normal([n_classes]))} l1 = tf.add(tf.matmul(data, hidden_1_layer['weights']), hidden_1_layer['biases']) l1 = tf.nn.relu(l1) l2 = tf.add(tf.matmul(l1, hidden_2_layer['weights']), hidden_2_layer['biases']) l2 = tf.nn.relu(l2) l3 = tf.add(tf.matmul(l2, hidden_3_layer['weights']), hidden_3_layer['biases']) l3 = tf.nn.relu(l3) output = tf.add(tf.matmul(l3, output_layer['weights']), output_layer['biases']) return output # 初始化预测相关节点 x = tf.placeholder(tf.float32, shape=(None, 784)) prediction = DeepNN(x) saver = tf.train.Saver() # 取一张测试集图片作为预测示例 test_image = mnist.test.images[0].reshape(1, 784) # 单张图片要调整形状为(1,784) true_label = mnist.test.labels[0] with tf.Session() as sess: # 恢复预训练模型参数 saver.restore(sess, "./mnist_deepnn_model") print("模型恢复成功!") # 执行预测:获取输出层概率,再取最大值对应的类别 pred_prob = sess.run(prediction, feed_dict={x: test_image}) pred_label = tf.argmax(pred_prob, 1).eval() print(f"真实标签: {tf.argmax(true_label, 0).eval()}") print(f"预测标签: {pred_label[0]}")
关键注意事项
- 恢复模型时,模型结构必须和训练时完全一致,包括层数、节点数、激活函数等,否则会出现参数不匹配的报错。
tf.train.Saver()默认保存所有变量,如果你只想保存特定层的参数,可以在初始化时指定var_list参数。- 预测时输入数据的形状要和训练时对齐:单张图片需reshape为
(1, 784),多张则为(N, 784)(N为图片数量)。
内容的提问来源于stack exchange,提问作者Pranjal Verma
相关产品推荐
相关产品推荐

