TensorFlow已训练CNN模型推理代码求助:报错排查与操作指导
搞定TensorFlow+TensorLayer CNN推理的那些坑
嘿,看你已经搞定了CNN的训练,现在卡在推理环节了,我帮你一步步梳理问题,把代码给你修正过来:
首先回答你最关心的:推理阶段需要用minibatch吗?
不是必须,但非常推荐。如果测试数据量很大,一次性喂进去会占满显存导致OOM;就算数据量小,用minibatch也能提高推理效率。你训练时用的TensorLayer tl.iterate.minibatches在推理时照样能用,只要把shuffle设成False就行,避免打乱测试数据的顺序。
然后是你的核心问题:加载模型+修复推理代码
你的现有推理代码有几个明显的问题,我给你拆解并修正:
1. 加载模型的正确姿势
你现在的代码直接调用tf.train.Saver().restore(),但连会话sess都没创建,而且加载前必须先定义和训练时完全一模一样的网络结构——TensorFlow需要知道变量的结构才能正确加载参数。
2. 推理时的关键参数设置
训练时用的is_training和keep_prob,推理时必须切换:
is_training设为False:BatchNorm、Dropout这些层在训练和推理时的行为完全不同,推理时要切换到评估模式。keep_prob设为1.0:推理时不需要Dropout,直接关闭它。
3. 数据预处理要和训练一致
你的数据加载部分基本没问题,但一定要确保:训练时如果对数据做了归一化(比如除以255)、标准化,推理时必须做完全相同的操作,否则预测结果会完全跑偏。
修复后的完整推理代码
结合你的训练代码,我把推理代码修正如下(注意要把MyCNN替换成你自己的网络实现):
import tensorflow as tf import tensorlayer as tl import scipy.io as sio import numpy as np import os import time # -------------------------- # 重点!!!必须和训练时的MyCNN完全一致 # 复制你训练代码里的MyCNN实现到这里 # -------------------------- def MyCNN(net_in, is_training, keep_prob): # 示例结构(替换成你自己的): # net = tl.layers.Conv2d(net_in, 32, (3,3), act=tf.nn.relu, name='conv1') # net = tl.layers.MaxPool2d(net, (2,2), name='pool1') # ... 剩下的网络层 ... # 确保层的数量、参数、名称和训练时完全相同 pass # -------------------------- # 加载测试数据 # -------------------------- print("\n\nPreparing testing data........................") test_data = sio.loadmat('MyTest.mat') Z0 = test_data['Real_testing1'] img_num_test = Z0.shape[0] # 这里必须和训练时的预处理完全一致!!! # 比如训练时做了 Z0 = Z0 / 255.0,这里也要加 X_test = np.empty([img_num_test, 128, 128, 1], dtype=np.float32) X_test[:,:,:,0] = Z0 print("\tTesting X shape: {0}".format(X_test.shape)) # -------------------------- # 初始化会话并加载模型 # -------------------------- print("\n\nRestore the network ...") save_dir = "checkpoints/" epoch = 1000 model_name = save_dir + str(epoch) + '_model' # 定义和训练时完全一致的输入占位符 x = tf.placeholder(tf.float32, shape=[None, 128, 128, 1], name='x') keep_prob = tf.placeholder(tf.float32, name='keep_prob') is_training = tf.placeholder(tf.bool, name='is_training') # 构建网络 net_out = MyCNN(x, is_training, keep_prob) y = net_out y_op = tf.argmax(tf.nn.softmax(y), 1) # 得到预测的类别 # 创建会话并加载模型 sess = tf.Session() saver = tf.train.Saver() # 加载模型,这里不需要初始化变量,restore会直接加载训练好的参数 saver.restore(sess, save_path=model_name) print("Model loaded successfully!") # -------------------------- # 推理环节(两种方式可选) # -------------------------- start_time_begin = time.time() print("\n\nRunning network...") # 方式一:用minibatch推理(推荐,显存友好) batch_size = 32 # 根据你的显卡显存调整大小 predictions = [] # 用tl的minibatch拆分测试数据,shuffle设为False保持顺序 for X_test_batch in tl.iterate.minibatches(X_test, batch_size=batch_size, shuffle=False): feed_dict = { x: X_test_batch, is_training: False, # 推理时关闭训练模式 keep_prob: 1.0 # 推理时不用Dropout } batch_pred = sess.run(y_op, feed_dict=feed_dict) predictions.extend(batch_pred) # 方式二:单张图片推理(适合小数据量测试) # single_pred = sess.run(y_op, feed_dict={ # x: X_test[0:1, :, :, :], # 取第一张图,注意维度要匹配[1,128,128,1] # is_training: False, # keep_prob: 1.0 # }) # print("Single image prediction:", single_pred) print("First 9 predictions:", predictions[:9]) print("Inference completed! Total time: {:.2f} seconds".format(time.time() - start_time_begin)) sess.close()
最后再划几个重点
- 网络结构绝对不能改:推理时的
MyCNN必须和训练时的代码完全一致,包括每一层的名称、参数数量,否则TensorFlow找不到对应的变量,加载模型会报错。 - 别乱初始化变量:加载模型前不要调用
tf.global_variables_initializer(),否则会把训练好的参数覆盖掉。 - 数据预处理要对齐:训练时怎么处理数据,推理时就怎么处理,比如归一化、裁剪、通道顺序这些细节都不能错。
内容的提问来源于stack exchange,提问作者fireman
相关产品推荐
相关产品推荐

