训练Fast R-CNN触发CUDA设备端断言错误的排查问询
解决Fast R-CNN训练时的CUDA设备端断言错误
问题根源
报错触发在fastrcnn_loss函数的torch.where(labels > 0)[0]处,且后续调用torch.manual_seed(1)也报错,核心原因是数据标签不合法导致CUDA上下文异常,进而引发后续操作连锁出错。
具体解决步骤
检查数据集标签合法性
- Fast R-CNN中0是背景类,目标类标签必须从1开始,且不能超过模型定义的
num_classes-1(num_classes包含背景类)。排查所有标注的类别ID,确保没有负数、非背景类的0或超出范围的数值。 - 可在数据加载器中添加临时代码,输出每个batch的标签值定位异常样本:
for images, targets in train_data_loader: print([t['labels'] for t in targets]) break
- Fast R-CNN中0是背景类,目标类标签必须从1开始,且不能超过模型定义的
重置CUDA环境
- 先清理CUDA缓存:
import torch torch.cuda.empty_cache() - 若清理后仍报错,直接重启Python运行环境(Notebook内核、终端进程),彻底重置损坏的CUDA上下文。
- 先清理CUDA缓存:
启用同步调试定位精确错误
- 设置环境变量强制CUDA同步执行,获取准确的错误堆栈:
- 终端运行训练脚本时:
CUDA_LAUNCH_BLOCKING=1 python your_train_script.py - Notebook中先执行:
import os os.environ['CUDA_LAUNCH_BLOCKING'] = '1'
- 终端运行训练脚本时:
- 设置环境变量强制CUDA同步执行,获取准确的错误堆栈:
验证模型与数据的类别数匹配
- 确认模型初始化时的
num_classes参数正确:如果数据集有K个目标类别,num_classes应设为K+1(包含背景类),比如目标类是1~5,num_classes=6。
- 确认模型初始化时的
内容的提问来源于stack exchange,提问作者tearis
相关产品推荐
相关产品推荐

