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
相关产品推荐
相关产品推荐

