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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.15 06:56:26