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

基于多分类的剪刀石头布预测模型分类结果位置不稳定问题排查

问题分析与解决建议

问题1:类别输出位置随机变化的原因与解决

核心原因

本质是类别到模型输出索引的映射未固定,大概率是以下两种情况:

  1. 代码1的训练/验证集划分是随机拆分,且每次划分后单独生成类别标签映射(比如用flow_from_directory时未指定固定的class_indices),导致每次训练时模型学到的类别索引对应关系随机变化。
  2. 手动生成标签时(代码2)没有固定类别与索引的绑定,而是依赖数据集遍历的随机顺序,导致映射不一致。

解决方法

  • 固定类别映射:在加载数据集时,手动指定class_indices参数,强制类别到索引的对应关系,示例:
    # 假设数据按文件夹分类:rock/、paper/、scissors/
    class_mapping = {'rock': 0, 'paper': 1, 'scissors': 2}
    train_generator = train_datagen.flow_from_directory(
        train_dir,
        target_size=(img_height, img_width),
        batch_size=batch_size,
        class_mode='categorical',
        class_indices=class_mapping  # 强制固定映射
    )
    val_generator = val_datagen.flow_from_directory(
        val_dir,
        target_size=(img_height, img_width),
        batch_size=batch_size,
        class_mode='categorical',
        class_indices=class_mapping  # 验证集用相同映射
    )
    
  • 保存映射文件:训练完成后,将class_mapping保存为JSON文件,预测时加载该文件,确保预测阶段的类别索引和训练完全一致:
    import json
    with open('class_mapping.json', 'w') as f:
        json.dump(class_mapping, f)
    # 预测时加载
    with open('class_mapping.json', 'r') as f:
        class_mapping = json.load(f)
    # 根据预测结果反查类别
    predicted_class = [k for k, v in class_mapping.items() if v == np.argmax(prediction)][0]
    
  • 避免随机生成映射:如果是手动拆分数据集,不要让训练集和验证集各自生成标签映射,统一使用提前定义好的class_mapping来给样本打标签。

问题2:模型输出仅偏向前两个类别的原因与解决

核心原因

  1. 数据分布不平衡:scissors类的样本数量远少于rock和paper,导致模型倾向于预测样本多的类别,只有当scissors的索引恰好是第2位时,模型才偶尔能学到少量特征。
  2. 标签映射混乱:当类别索引随机变化时,scissors的标签可能在训练时被错误分配,导致模型根本没有学到该类的特征,只有当索引固定在第2位时,标签才正确,模型才有输出。
  3. 模型与损失函数不匹配:之前修改class_mode和输出层神经元数量时,可能没有同步调整损失函数,比如用categorical模式却用了sparse_categorical_crossentropy损失,导致模型训练异常。

解决方法

  • 检查数据分布:统计三个类别的样本数量,如果scissors样本不足:
    • 对scissors类做数据增强(旋转、翻转、缩放等),生成更多样本;
    • 采用过采样(复制scissors类样本)或欠采样(减少rock/paper类样本)平衡数据;
    • 在损失函数中加入类别权重,给scissors类更高的权重,示例:
      class_weights = {0: 1.0, 1: 1.0, 2: 3.0}  # 假设scissors是索引2,权重设为3
      model.fit(train_generator, class_weight=class_weights, ...)
      
  • 验证标签一致性:打印训练集和验证集的class_indices,确认scissors对应的索引在训练和验证阶段完全一致,且每个样本的标签都正确对应类别。
  • 匹配模型与损失函数:
    • 若class_mode='categorical',输出层必须是Dense(3, activation='softmax'),损失函数用categorical_crossentropy;
    • 若class_mode='sparse',输出层同样是Dense(3, activation='softmax'),损失函数用sparse_categorical_crossentropy;
    • 确保两者严格匹配,避免训练过程中梯度计算异常。
  • 评估模型性能:用混淆矩阵查看模型对每个类别的预测结果,明确scissors类的预测准确率、召回率,针对性调整模型(比如增加卷积层、调整学习率)。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.04 20:01:14