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

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)

遇到的问题

  1. 测试自制手绘PNG图像时,用matplotlib.image.imread或cv2.imread加载后转为numpy数组,维度不符合模型要求,导致卷积层报错;
  2. 尝试用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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.22 19:07:09