Faster-RCNN添加数据增强后训练损失变为NaN问题排查
Faster R-CNN 加入Albumentations增强后训练损失NaN问题排查
问题背景
基于PyTorch官方预训练的Faster R-CNN(MobileNetV3-Large FPN backbone),在roboflow国际象棋棋子数据集上开展训练,模型初始化代码如下:
def get_model(n_classes): model = models.detection.fasterrcnn_mobilenet_v3_large_fpn(pretrained=True) in_features = model.roi_heads.box_predictor.cls_score.in_features model.roi_heads.box_predictor = models.detection.faster_rcnn.FastRCNNPredictor(in_features, n_classes) return model
未加入数据增强时,数据集类__getitem__方法实现如下,训练流程可正常运行,训练10轮后预测边界框效果良好,mAP指标可达0.4~0.8:
def __getitem__(self, index): id = self.ids[index] image = self._load_image(id) target = copy.deepcopy(self._load_target(id)) boxes = torch.tensor([t["bbox"] for t in target]) new_boxes = torch.add(boxes[:,:2],boxes[:,2:]) boxes = torch.cat((boxes[:,:2],new_boxes),1) labels = torch.tensor([t["category_id"] for t in target], dtype=torch.int64) image = torch.from_numpy(image).permute(2,0,1) targ = {} targ['boxes'] = boxes targ['labels'] = labels targ['image_id'] = torch.tensor(index) targ['area'] = (boxes[:,2]-boxes[:,0]) * (boxes[:,3]-boxes[:,1]) targ['iscrowd'] = torch.tensor([t["iscrowd"] for t in target], dtype=torch.int64) return image, targ
后续尝试加入基于Albumentations的数据增强逻辑,初始版本仅保留张量转换操作:
def get_transforms(train=False): if train: transform = A.Compose([ ToTensorV2() ], bbox_params=A.BboxParams(format='pascal_voc',label_fields=["labels"])) else: transform = A.Compose([ ToTensorV2() ], bbox_params=A.BboxParams(format='pascal_voc',label_fields=["labels"])) return transform
对应修改__getitem__方法加入增强调用逻辑后,训练过程中出现损失为NaN的异常:
def __getitem__(self, index): id = self.ids[index] image = self._load_image(id) target = copy.deepcopy(self._load_target(id)) boxes = torch.tensor([t["bbox"] for t in target]) new_boxes = torch.add(boxes[:,:2],boxes[:,2:]) boxes = torch.cat((boxes[:,:2],new_boxes),1) labels = torch.tensor([t["category_id"] for t in target], dtype=torch.int64) if self.transforms is not None: transformed = self.transforms(image=image, bboxes=boxes, labels=labels) image = transformed['image'] boxes = torch.tensor(transformed['bboxes']).view(len(transformed["bboxes"]),4) labels = torch.tensor(transformed["labels"],dtype=torch.int64) else: image = torch.from_numpy(image).permute(2,0,1) targ = {} targ['boxes'] = boxes targ['labels'] = labels targ['image_id'] = torch.tensor(index) targ['area'] = (boxes[:,2]-boxes[:,0]) * (boxes[:,3]-boxes[:,1]) targ['iscrowd'] = torch.tensor([t["iscrowd"] for t in target], dtype=torch.int64) return image, targ
batch_size设为10时,训练中断前的最后一段日志输出如下:
Epoch: [0] [10/18] eta: 0:02:41 lr: 0.003237 loss: 2.3237 (2.6498) loss_classifier: 1.4347 (1.8002) loss_box_reg: 0.7538 (0.7682) loss_objectness: 0.0441 (0.0595) loss_rpn_box_reg: 0.0221 (0.0220) time: 20.2499 data: 0.1298 Loss is nan, stopping training {'loss_classifier': tensor(nan, grad_fn=<NllLossBackward0>), 'loss_box_reg': tensor(nan, grad_fn=<DivBackward0>), 'loss_objectness': tensor(nan, grad_fn=<BinaryCrossEntropyWithLogitsBackward0>), 'loss_rpn_box_reg': tensor(nan, dtype=torch.float64, grad_fn=<DivBackward0>)}
补充观测:训练时采用图像切片输入,部分训练样本为空样本(不含任何检测目标)。处理空样本时,损失值旁括号内的数值会大幅升高(注:日志中括号外为当前batch损失,括号内为训练启动以来的滑动平均损失),分别尝试Adam、SGD优化器均会复现NaN问题:
# 处理空样本前 Epoch: [0] [17/26] eta: 0:00:14 lr: 0.003601 loss: 2.4854 (3.9266) loss_classifier: 1.1224 (2.2893) loss_box_reg: 0.7182 (1.2226) loss_objectness: 0.0497 (0.3413) loss_rpn_box_reg: 0.0116 (0.0735) time: 1.6587 data: 0.0102 # 处理空样本后 Epoch: [0] [18/26] eta: 0:00:12 lr: 0.003801 loss: 2.8132 (61.1689) loss_classifier: 1.5675 (28.8652) loss_box_reg: 0.7563 (29.8348) loss_objectness: 0.1070 (2.2412) loss_rpn_box_reg: 0.0145 (0.2278) time: 1.6240 data: 0.0098
1. 训练过程损失变为NaN的核心原因
核心bug是空标注样本和Albumentations转换逻辑的兼容性问题,触发链路非常明确:
- 当样本没有任何标注目标时,初始构造的
boxes是shape为[0,4]的空张量 - 低版本Albumentations处理零长度bbox输入时存在已知缺陷:传入空bbox列表/张量后,转换返回的bboxes不是预期的空列表,而是包含无效值——可能是NaN、坐标超出图像范围、宽/高为0的退化框
- 代码未对增强后的bbox做合法性校验,直接转成张量传入Faster R-CNN,会直接触发两类数值错误:
- 计算框面积、回归损失时遇到0值/负值做除法,除以0直接产生NaN
- 非法坐标传入RoI Align层时会产生数值溢出,反向传播时梯度爆炸,几轮迭代后NaN会扩散到所有损失项
- 观测到的空样本处理后平均损失突然飙升,就是NaN出现的前兆:退化框先产生极大的损失值拉爆滑动平均,下一个迭代步就会出现全损失NaN。
- 代码里还有一个隐性兼容问题:直接把torch张量格式的bbox传入Albumentations,部分版本对torch张量类型的空bbox处理逻辑有问题,会直接返回未定义值。
2. 根因定位步骤
按优先级从易到难排查:
- 第一步:校验增强后输出的所有bbox合法性
传入Albumentations时先把bbox和label转成Python原生列表,避免torch张量的兼容性问题,同时过滤所有无效框(坐标越界、宽高小于等于0),参考修正逻辑:
加完过滤后如果NaN消失,即可确认是无效bbox导致的问题。if self.transforms is not None: transformed = self.transforms( image=image, bboxes=boxes.numpy().tolist(), labels=labels.numpy().tolist() ) image = transformed['image'] valid_bboxes = [] valid_labels = [] img_h, img_w = image.shape[1], image.shape[2] for bbox, label in zip(transformed['bboxes'], transformed['labels']): x1, y1, x2, y2 = bbox if 0 <= x1 < x2 <= img_w and 0 <= y1 < y2 <= img_h: valid_bboxes.append(bbox) valid_labels.append(label) boxes = torch.tensor(valid_bboxes, dtype=torch.float32).reshape(-1, 4) labels = torch.tensor(valid_labels, dtype=torch.int64) - 第二步:单独校验空样本输出
把所有无标注的空样本单独提取出来跑增强逻辑,打印转换后的bboxes返回值,绝大多数情况下可以观测到空样本输入时,Albumentations返回了非空的无效值。 - 第三步:校验张量数值类型
从日志可看到loss_rpn_box_reg为float64类型,说明输入数据存在double类型,和模型float32的参数类型不匹配会放大数值不稳定问题,需确保image、boxes统一为torch.float32类型,labels为torch.int64类型,禁止保留float64精度的张量输入模型。 - 第四步:空样本对照验证
临时在数据加载时过滤所有空样本,如果训练全程不再出现NaN,即可100%确认问题是空样本与增强逻辑的交互bug。
内容的提问来源于stack exchange,提问作者Alican Kartal
相关产品推荐
相关产品推荐

