TensorFlow自定义卷积模型无法测试自制手绘图像求助
问题:手绘图案分类模型测试时的维度与预测失效问题
问题背景
作为TensorFlow新手,采用子类化方式构建卷积模型,用于分类10种手绘图案,训练使用Quick Draw数据集,模型输入尺寸为(1, 28, 28, 1)。模型编译配置:
- loss:
SparseCategoricalCrossentropy - optimizer:
SGD - 评估指标:
SparseCategoricalAccuracy
模型代码
@tf.keras.saving.register_keras_serializable() class ConvModel(tf.keras.Model): def __init__(self): super(ConvModel, self).__init__() self.conv1 = tf.keras.layers.Conv2D(32, 4, activation = 'relu', name = 'conv1', input_shape = (28, 28, 1)) self.conv2 = tf.keras.layers.Conv2D(32, 3, activation = 'relu', name = 'conv2') self.conv3 = tf.keras.layers.Conv2D(32, 3, activation = 'relu', name = 'conv3') self.flatten = tf.keras.layers.Flatten(name = 'flatten') self.d1 = tf.keras.layers.Dense(32, activation = 'relu', name = 'd1') self.d2 = tf.keras.layers.Dense(16, activation = 'relu', name = 'd2') self.out = tf.keras.layers.Dense(10, activation = 'softmax', name = 'out') def call(self, image): conv1 = self.conv1(image) conv2 = self.conv2(conv1) conv3 = self.conv3(conv2) flatten = self.flatten(conv3) d1 = self.d1(flatten) d2 = self.d2(d1) output = self.out(d2) return output
训练数据加载与归一化
数据加载代码
files = [name for name in os.listdir() if ".npy" in name] max_size_per_cl = 1500 draw_class = [] # 计算数据集总大小 size = 0 for name in files: draws = np.load(name) draws = draws[:max_size_per_cl] size += draws.shape[0] images = np.zeros((size, 28, 28)) targets = np.zeros((size,)) it = 0 t = 0 for name in files: # 记录类别名称 draw_class.append(name.replace("full_numpy_bitmap_", "").replace(".npy", "")) draws = np.load(name) draws = draws[:max_size_per_cl] # 反转图像(将Quick Draw黑底白笔转为白底黑笔) images[it:it+draws.shape[0]] = np.invert(draws.reshape(-1, 28, 28)) targets[it:it+draws.shape[0]] = t it += draws.shape[0] t += 1 images = images.astype(np.float32) # 打乱数据集 indexes = np.arange(size) np.random.shuffle(indexes) images = images[indexes] targets = targets[indexes] # 划分训练/验证集 images, images_valid, targets, targets_valid = train_test_split(images, targets, test_size=0.33)
归一化代码
scaler = StandardScaler() scaled_images = scaler.fit_transform(images.reshape(-1, 28 * 28)) scaled_images_valid = scaler.transform(images_valid.reshape(-1, 28 * 28)) # 调整为模型输入维度 images = scaled_images.reshape(-1, 28, 28, 1) images_valid = scaled_images_valid.reshape(-1, 28, 28, 1)
遇到的问题
- 测试自制手绘PNG图像时,用
matplotlib.image.imread或cv2.imread加载后转为numpy数组,维度不符合模型要求,导致卷积层报错; - 尝试用
np.resize或cv2.resize调整维度后,模型预测效果极差,softmax输出的10个类别概率接近(如array([[0.1089, 0.0981, ..., 0.0956]], dtype=float32))。
解决方案
一、解决测试图像维度不匹配问题
模型要求输入维度为(batch_size, 28, 28, 1),单张测试图需调整为(1, 28, 28, 1),同时要确保图像为单通道灰度图,处理步骤如下:
标准化测试图像预处理函数
import cv2 import numpy as np def preprocess_test_image(image_path, scaler): # 1. 加载为灰度图 img = cv2.imread(image_path, 0) # 若用matplotlib加载: # import matplotlib.image as mpimg # img = mpimg.imread(image_path) # if img.ndim == 3: # 彩色图转灰度 # img = np.dot(img[...,:3], [0.2989, 0.5870, 0.1140]) # 2. 调整尺寸为28x28(用INTER_AREA插值避免模糊) img = cv2.resize(img, (28, 28), interpolation=cv2.INTER_AREA) # 3. 反转图像(和训练数据明暗对齐:白底黑笔) img = np.invert(img) # 4. 转为float32类型 img = img.astype(np.float32) # 5. 用训练时的scaler做归一化 img_flat = img.reshape(1, 28*28) img_scaled = scaler.transform(img_flat) # 6. 调整为模型要求的输入维度 img_input = img_scaled.reshape(1, 28, 28, 1) return img_input
二、解决预测效果差的核心问题
softmax输出接近的本质是测试数据预处理流程与训练数据完全不一致,需确保以下三点:
- 明暗一致:必须对齐训练时的
np.invert操作,保证测试图是白底黑笔(和训练数据风格统一); - 归一化一致:必须复用训练时拟合好的
scaler,不能重新拟合或跳过归一化; - 尺寸与通道一致:确保测试图是28x28的单通道灰度图,不能是彩色图或其他尺寸。
三、测试验证流程
调用预处理函数后执行预测:
# 若训练时保存了scaler,需先加载(如用joblib) # import joblib # scaler = joblib.load("scaler.pkl") test_img_path = "你的手绘图案路径.png" processed_img = preprocess_test_image(test_img_path, scaler) # 执行预测 pred_probs = model.predict(processed_img) pred_class_idx = np.argmax(pred_probs, axis=1)[0] pred_class_name = draw_class[pred_class_idx] print(f"预测类别概率:{pred_probs}") print(f"预测类别:{pred_class_name}")
额外注意事项
- 训练时需保存
scaler,可使用joblib.dump(scaler, "scaler.pkl"),测试时用joblib.load("scaler.pkl")加载; - 若手绘图像素值范围为0-1(如matplotlib加载的float类型),需先转为0-255的uint8格式再处理;
- 手绘图案尽量和Quick Draw数据集风格对齐:笔画居中、粗细适中。
内容的提问来源于stack exchange,提问作者Lambda7214
相关产品推荐
相关产品推荐

