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

如何将CNN输出的(359,3)预测结果转为(359,)类别标签?

三类超声图像CNN模型预测问题排查与解决方法

核心问题分析

  • 数据集类别严重不平衡:训练集中类别2仅164个样本,远少于类别0(632个)和类别1(804个),模型无法充分学习到类别2的特征,导致预测时无法识别该类。
  • 输出层激活函数错误:多分类任务使用了sigmoid激活函数,该函数适用于多标签分类场景,无法输出合理的单标签多分类概率分布,导致argmax无法正确获取类别2。
  • 训练轮数不足:仅训练3轮,模型未充分收敛,尤其是少数类的特征未被有效学习。

解决步骤

1. 修正模型输出层激活函数

将最后一层的sigmoid替换为softmax,配合categorical_crossentropy损失,实现单标签多分类的概率输出:

import tensorflow as tf
from tensorflow.keras import layers

model = tf.keras.Sequential([
    layers.Conv2D(32, (3, 3), activation='relu', input_shape=X_train.shape[1:]),
    layers.MaxPooling2D((2, 2)),
    layers.Conv2D(64, (3, 3), activation='relu'),
    layers.MaxPooling2D((2, 2)),
    layers.Conv2D(128, (3, 3), activation='relu'),
    layers.MaxPooling2D((2, 2)),
    layers.Conv2D(256, (3, 3), activation='relu'),
    layers.MaxPooling2D((2, 2)),
    layers.Flatten(),
    layers.Dense(256, activation='relu'),
    layers.Dense(3, activation='softmax')  # 替换为softmax
])

model.compile(optimizer='adam',
              loss='categorical_crossentropy',
              metrics=['accuracy'])

2. 处理类别不平衡问题

通过类别权重让模型更关注少数类,避免多数类主导训练:

# 计算类别权重:总样本数/(类别数量*该类样本数)
total_samples = len(y_train)
class_counts = [632, 804, 164]
class_weights = {i: total_samples/(3*count) for i, count in enumerate(class_counts)}

# 训练时传入类别权重参数
history = model.fit(X_train, y_train, epochs=20,
                    validation_data=(X_test, y_test),
                    class_weight=class_weights)

3. 增加训练轮数并加入早停机制

增加训练轮数让模型充分学习,同时用早停防止过拟合:

from tensorflow.keras.callbacks import EarlyStopping

# 当验证集损失连续3轮不下降时停止训练,恢复最优权重
early_stop = EarlyStopping(monitor='val_loss', patience=3, restore_best_weights=True)

history = model.fit(X_train, y_train, epochs=20,
                    validation_data=(X_test, y_test),
                    class_weight=class_weights,
                    callbacks=[early_stop])

4. 获取正确的类别预测结果

修正模型后,使用argmax(axis=1)即可得到形状为(359,)的类别数组:

predictions = model.predict(X_test)
predicted_classes = predictions.argmax(axis=1)

print(predicted_classes.shape)  # 输出:(359,)
print(predicted_classes)        # 输出包含0、1、2的类别数组

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.28 00:55:06