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

TensorFlow训练模型对训练集图像预测始终错误,求技术排查

宝可梦图像分类模型测试错误问题排查

问题描述

我正在学习TensorFlow,实践是最好的学习方式。最近用Kaggle的宝可梦数据集练手,训练完模型后,用训练过的图像测试,结果全错,求帮忙找问题。

代码

import os
import pandas as pd
import cv2
import numpy as np
from tensorflow.keras.models import Sequential
from tensorflow.keras.layers import Flatten, Dense, Softmax, Conv2D, MaxPooling2D
import tensorflow.keras as keras

root_dir = "./images"
read_files = os.path.join(root_dir)
file_names = os.listdir(read_files)
data_pok = {}

def build_image_labels():
    data = pd.read_csv("./pokemon.csv")
    data.head()
    for index, row in data.iterrows():
        name = row["Name"]
        type_one = row["Type1"]
        type_two = row["Type2"]
        data_pok[index] = {
            "index": index + 1,
            "name": name,
            "type_one": type_one,
            "type_two": type_two
        }

    type_one = data["Type1"].unique()
    ids = [0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16, 17]
    types_indices = dict(zip(type_one, ids))

    final_images = []
    final_pokemon = []
    count = 0

    for file in file_names:
        pokemon = data_pok[count]
        count += 1
        img = cv2.imread(os.path.join(root_dir, file))
        img = cv2.resize(img, (120, 120))
        final_images.append(np.array(img))
        final_pokemon.append(np.array(pokemon['index']-1))

    final_images = np.array(final_images, dtype=np.float32)/255.0
    final_pokemon = np.array(final_pokemon, dtype=np.float32).reshape(809, 1)

    print("After", final_images.shape, final_pokemon.shape)

    return final_images, final_pokemon, type_one

def build_model():
    model = Sequential([
        Flatten(input_shape=(120,120,3)),
        Dense(100, activation='relu'),
        Dense(100, activation='relu'),
        Dense(100, activation='relu'),
        Dense(809),
    ])
    model.summary()

    model.compile(
        optimizer='Adam',
        loss= keras.losses.SparseCategoricalCrossentropy(from_logits=True),
        metrics=['accuracy']
    )

    (images, pokemons, types) = build_image_labels()

    history = model.fit(images, pokemons, epochs=50)
    probability_model = Sequential([model, Softmax()])

    pokemon = cv2.imread('./images/mewtwo.png')
    pokemon = cv2.resize(pokemon, (120,120))
    pokemon_arr = np.array(pokemon, dtype=np.float32) / 255.0
    pokemon_arr = np.expand_dims(pokemon, axis=0)

    cv2.imshow("Pokemon", pokemon)
    cv2.waitKey(0)

    predictions = probability_model.predict(pokemon)

    id = np.argmax(predictions[0])
    print("id", id)
    print("pokemon index", data_pok[id])
    print("accuracy of the model", history.history['accuracy'][-1])

build_model()

问题原因及解决方法

1. 图像与标签不匹配(核心问题)

os.listdir()返回的文件名顺序和csv中宝可梦的行号顺序不一致,导致训练时标签完全错误,测试自然无法正确识别。

解决:
通过文件名提取编号匹配标签(数据集图片文件名是001.png到809.png,对应csv第0到808行),修改build_image_labels里的循环:

for file in file_names:
    # 从文件名提取编号,转成整数后减1得到csv里的index
    file_id = int(file.split('.')[0]) - 1
    pokemon = data_pok[file_id]
    img = cv2.imread(os.path.join(root_dir, file))
    img = cv2.resize(img, (120, 120))
    final_images.append(np.array(img))
    final_pokemon.append(np.array(pokemon['index']-1))

2. 测试数据预处理错误

测试时扩展维度用了未归一化的原始图像,和训练时的输入数据分布不一致,导致预测错误。

解决:
修改测试部分代码:

pokemon = cv2.imread('./images/mewtwo.png')
pokemon = cv2.resize(pokemon, (120,120))
pokemon_arr = np.array(pokemon, dtype=np.float32) / 255.0
pokemon_arr = np.expand_dims(pokemon_arr, axis=0)  # 用归一化后的数组扩展维度

# 预测时传入处理后的数组
predictions = probability_model.predict(pokemon_arr)

3. 模型结构不适合图像分类

全连接层无法有效提取图像空间特征,809类的个体分类任务难度极大,简单全连接层难以胜任。

解决:
换成CNN结构:

model = Sequential([
    Conv2D(32, (3,3), activation='relu', input_shape=(120,120,3)),
    MaxPooling2D((2,2)),
    Conv2D(64, (3,3), activation='relu'),
    MaxPooling2D((2,2)),
    Conv2D(128, (3,3), activation='relu'),
    MaxPooling2D((2,2)),
    Flatten(),
    Dense(128, activation='relu'),
    Dense(809)  # 个体分类保持809,属性分类改为18
])

4. 任务目标混淆

代码中定义了属性映射字典types_indices但未使用,若原本想做宝可梦属性分类(18类Type1),需调整标签和输出层:

  • 标签改为types_indices[pokemon['type_one']]
  • 最后一层Dense改为Dense(18)

5. 无验证集监控训练效果

训练时未划分验证集,无法判断模型是否过拟合,也无法确认训练的真实效果。

解决:
用train_test_split划分数据集:

from sklearn.model_selection import train_test_split

(images, pokemons, types) = build_image_labels()
train_images, val_images, train_pokemons, val_pokemons = train_test_split(images, pokemons, test_size=0.2, random_state=42)

history = model.fit(train_images, train_pokemons, epochs=50, validation_data=(val_images, val_pokemons))

内容的提问来源于stack exchange,提问作者Ruben Mim

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.13 08:19:58