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

图像分割中使用class_weight时model.fit()报错的技术咨询

问题分析与解决

错误原因

图像分割属于像素级分类任务,每个样本(单张图像)包含大量像素点,每个像素对应一个类别标签。而Keras的class_weight参数是为样本级分类任务设计的(每个样本仅对应一个类别),直接传入会导致Keras尝试将每个像素的类别数组转换为单标量,触发TypeError。

解决方案

放弃使用class_weight参数,改为手动生成像素级的样本权重矩阵,并通过model.fit的sample_weight参数传入。同时修正损失函数选择的问题。

步骤1:修正损失函数

你的任务是4类别分割,模型输出为4通道,当前使用的binary_crossentropy适用于二分类/多标签任务,应替换为categorical_crossentropy(适配one-hot编码的标签):

total_loss = 'categorical_crossentropy'

步骤2:生成像素级样本权重

利用已有的类别索引矩阵msk_argmax,结合自定义的类别权重字典,生成每个像素对应的权重矩阵:

# 基于类别索引生成像素级权重
sample_weights = np.zeros_like(msk_argmax, dtype=np.float32)
for cls_idx, weight in class_weights_manual.items():
    sample_weights[msk_argmax == cls_idx] = weight

# 增加通道维度,匹配输入数据格式
sample_weights = sample_weights[..., np.newaxis]

步骤3:修改model.fit调用

移除class_weight参数,传入生成的sample_weights:

history=model.fit(img, msk, epochs=50, verbose=1, validation_split=0.2, shuffle=False,
                  sample_weight=sample_weights)

额外优化:移除冗余代码

你的msk_argmax已经是0-3的类别索引,无需再用LabelEncoder处理,可删除以下冗余代码:

labelencoder = LabelEncoder()
masks_reshaped_encoded = labelencoder.fit_transform(masks_reshaped)

修改后的关键代码片段

# ... 其他代码保持不变 ...

msk = msk.astype(np.uint8)
print(msk.shape)
msk_argmax=np.argmax(msk, axis=3)

class_weights_manual= {0:1, 1:2, 2:4, 3:8}

# 生成像素级样本权重
sample_weights = np.zeros_like(msk_argmax, dtype=np.float32)
for cls_idx, weight in class_weights_manual.items():
    sample_weights[msk_argmax == cls_idx] = weight
sample_weights = sample_weights[..., np.newaxis]

# 修正损失函数并编译模型
total_loss = 'categorical_crossentropy'
metrics = ['accuracy']
optim = 'adam'
model.compile(optimizer = optim, loss=total_loss, metrics= metrics)

# 传入sample_weight而非class_weight
history=model.fit(img, msk, epochs=50, verbose=1, validation_split=0.2, shuffle=False,
                  sample_weight=sample_weights)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.24 14:40:21