图像分割中使用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
相关产品推荐
相关产品推荐

