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

如何将Keras CNN的分类输出映射为可理解的猫狗品种名称?

将CNN预测结果转换为猫狗品种名称的解决方案

1. 建立索引与品种名称的映射关系

你需要先明确训练数据中每个类别索引对应的具体品种,这一步必须和训练时的类别编码逻辑完全一致:

场景1:使用ImageDataGenerator.flow_from_directory训练

如果训练时用该方法加载数据,生成器会自动生成class_indices字典(键为品种名,值为对应索引)。你可以在训练代码中生成并保存逆映射字典:

import json

# 假设训练数据生成器为train_generator
class_label_map = {v: k for k, v in train_generator.class_indices.items()}

# 保存映射到本地文件,方便预测时调用
with open('class_label_map.json', 'w') as f:
    json.dump(class_label_map, f)

场景2:手动编码类别

如果是手动用LabelEncoder或自定义方式编码类别,你需要按训练时的编码顺序,整理出所有37个品种的列表,再生成映射字典:

# 按训练时的编码顺序,依次列出所有猫狗品种
breed_names = ['阿富汗猎犬', '巴哥犬', ..., '英短猫']  # 替换为你的37个品种
class_label_map = {i: breed_names[i] for i in range(len(breed_names))}

# 保存映射文件
import json
with open('class_label_map.json', 'w') as f:
    json.dump(class_label_map, f)

2. 修改预测代码,实现结果转换

在预测代码中加载映射字典,用np.argmax得到的索引去字典中匹配对应的品种名,同时修正图像输入的潜在问题:

import json
import cv2
import numpy as np
import matplotlib.pyplot as plt
import random
import keras as k

# 加载预存的类别映射字典
with open('class_label_map.json', 'r') as f:
    class_label_map = json.load(f)
    # 将json保存的字符串索引转为整数
    class_label_map = {int(idx): breed for idx, breed in class_label_map.items()}

# 加载训练好的模型
model = k.models.load_model(checkpoint_filepath)

i = 1
while i < 15:
    try:
        # 随机读取一张图像
        img = cv2.imread(files[random.randint(1, 150)])
        unfiltered_image = img.copy()
        
        # 关键:cv2读取的图像是BGR格式,需转为训练时使用的RGB格式
        img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)
        # 如果训练时没有对图像做反转,删除np.invert这一步!否则输入数据分布不匹配
        # img = np.invert(img)
        img = np.expand_dims(img, axis=0)  # 增加batch维度
        
        # 执行预测,关闭日志输出
        prediction = model.predict(img, verbose=0)
        pred_index = np.argmax(prediction)
        pred_breed = class_label_map[pred_index]
        # 计算预测置信度
        confidence = prediction[0][pred_index] * 100
        
        # 输出可理解的结果
        print(f"这张图像大概率是:{pred_breed}(置信度:{confidence:.2f}%)")
        # 显示图像时需将BGR转回RGB
        plt.imshow(unfiltered_image[:, :, ::-1])
        plt.show()
        
    finally:
        i += 1

3. 关键注意事项

  • 图像格式一致性:cv2默认读取BGR格式图像,而多数训练流程(如ImageDataGenerator)使用RGB格式,必须转换通道顺序,否则会严重影响预测准确性。
  • 取消不必要的图像反转:如果训练时没有对图像执行np.invert操作,预测时务必删除该步骤,保证输入图像与训练数据的视觉特征一致。
  • 映射字典准确性:这是结果转换的核心,必须确保字典中的索引与训练时的类别编码完全对应,之前用if语句失败大概率是因为对应关系错误。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.21 16:44:57