基于BDD10k训练SegFormer时损失值为NaN的问题排查
SegFormer在BDD10k训练损失NaN问题排查与修复
问题背景
使用预训练nvidia/mit-b0的SegFormer模型在BDD10k数据集(10k张带掩码图像,掩码0-18对应标签,255为忽略值)上执行语义分割任务,训练时损失始终为NaN,预测掩码全为NaN。先后尝试图像归一化、降低学习率、调整训练轮数、更换预训练模型,以及两种预处理流程(原生tf.data构建、HuggingFace Dataset+AutoImageProcessor),问题均未解决。
现有实现的潜在问题
1. 掩码处理的关键疏漏
- 第一种tf.data流程:用
tf.image.decode_jpeg加载掩码,但BDD10k的语义分割掩码多为单通道PNG格式,解码后可能出现像素值异常;未明确处理255忽略值的类型转换,导致模型计算损失时出现无效值。 - 第二种Dataset流程:
preprocess函数直接将标签传入image_processor,未确保掩码为整数类型(尤其是255),且未对掩码做数值范围校验,可能引发数据类型或数值异常。
2. 模型配置缺失忽略值参数
第一种流程初始化模型时,未设置semantic_loss_ignore_index=255,SegFormer内置交叉熵损失会将255当作有效标签计算,超出0-18的标签范围会导致softmax计算异常,进而产生NaN。
3. 数据类型不匹配
- 图像预处理后为float32,但掩码若为float类型(如resize后未强制转int),损失函数计算会出现数值不稳定;
- 第二种流程中
img_to_array返回float32,但掩码需保持int32/int64类型,否则会引发损失计算的无效操作。
排查与修复步骤
步骤1:修正掩码处理逻辑
针对tf.data流程
修改load_and_preprocess函数,确保掩码类型与数值合规:
def load_and_preprocess(image_path, mask_path): # 加载图像与掩码(改用decode_png适配BDD10k掩码格式) image = tf.image.decode_jpeg(tf.io.read_file(image_path), channels=3) mask = tf.image.decode_png(tf.io.read_file(mask_path), channels=1) # 预处理图像与掩码 image = tf.image.resize(image, (height, width)) mask = tf.image.resize(mask, (height, width), method='nearest') # 强制掩码为int32类型,修正异常值为255 mask = tf.cast(mask, tf.int32) mask = tf.squeeze(mask, axis=-1) mask = tf.where(tf.logical_or(mask < 0, mask > 18), 255, mask) image = normalize(image) image = tf.transpose(image, perm=(2, 0, 1)) return {'pixel_values': image, 'labels': mask}
同时初始化模型时添加忽略值配置:
model = TFSegformerForSemanticSegmentation.from_pretrained( 'nvidia/mit-b0', num_labels=num_labels, id2label=id2label, label2id=label2id, ignore_mismatched_sizes=True, semantic_loss_ignore_index=255 # 指定忽略值 )
针对Dataset流程
修改preprocess函数,确保掩码数值与类型合规:
def preprocess(example_batch): images = [transforms(x.convert('RGB')) for x in example_batch['pixel_values']] # 处理掩码:转int32类型,修正异常值为255,确保单通道 labels = [np.array(x).astype(np.int32) for x in example_batch['label']] labels = [np.where(np.logical_or(lbl < 0, lbl > 18), 255, lbl) for lbl in labels] labels = [lbl.squeeze(-1) if lbl.ndim ==3 else lbl for lbl in labels] inputs = image_processor(images, labels=labels) return inputs
步骤2:添加数据校验环节
训练前随机抽取样本,验证数据有效性:
# 针对tf.data流程 for batch in train_dataset.take(1): print("图像数值范围:", tf.reduce_min(batch['pixel_values']), tf.reduce_max(batch['pixel_values'])) print("掩码数值分布:", tf.unique(tf.reshape(batch['labels'], [-1]))) print("图像类型:", batch['pixel_values'].dtype) print("掩码类型:", batch['labels'].dtype)
需确认图像无NaN/inf值,掩码仅包含0-18和255,数据类型符合要求。
步骤3:手动验证损失计算逻辑
临时替换内置损失,用自定义交叉熵损失排查问题:
loss_fn = tf.keras.losses.SparseCategoricalCrossentropy(from_logits=True, ignore_class=255) def custom_loss(y_true, y_pred): y_pred = tf.transpose(y_pred, perm=(0,2,3,1)) # 转换为(batch, H, W, num_classes)格式 return loss_fn(y_true, y_pred) model.compile(optimizer=tf.keras.optimizers.Adam(learning_rate=1e-5), loss=custom_loss)
若手动计算损失正常,说明是模型内置损失的配置问题;若仍为NaN,需重点排查数据本身。
步骤4:添加梯度裁剪抑制数值爆炸
初始梯度过大可能导致NaN,添加梯度裁剪:
optimizer = tf.keras.optimizers.Adam(learning_rate=1e-5, clipnorm=1.0) model.compile(optimizer=optimizer)
步骤5:排查异常样本
遍历数据集,找出加载失败或数值异常的样本:
# 针对tf.data流程 for idx, (image_path, mask_path) in enumerate(dataset): try: sample = load_and_preprocess(image_path, mask_path) if tf.reduce_any(tf.math.is_nan(sample['pixel_values'])) or tf.reduce_any(tf.math.is_nan(sample['labels'])): print(f"异常样本:{image_path.numpy().decode()}, {mask_path.numpy().decode()}") except Exception as e: print(f"加载样本失败:{idx}, {image_path}, 错误:{e}")
内容的提问来源于stack exchange,提问作者Jimenemex
相关产品推荐
相关产品推荐

