PyTorch 1.6.0中(predicted==labels)触发CUDA设备断言错误求助
CUDA RuntimeError: device-side assert triggered 问题排查
先帮你梳理下场景:你运行了自己写的evaluate函数,开启calc_loss=True后,代码走到pass 1、pass 2,然后在right = torch.sum(predicted==labels).item()这行触发了CUDA设备端断言错误,但单独测试这行代码却没问题。这种情况其实很常见——CUDA的异步执行特性会让报错延迟显现,真正的问题大概率出在之前的操作里,只是到这行才爆出来。
下面是几个最可能的原因和排查方向:
1. 模型输出类别数与标签取值范围不匹配
这是触发这类CUDA断言错误最常见的原因。
- 你的模型最后一层输出的logits维度(比如
outputs.shape[1])应该等于数据集的类别数num_classes,而标签的取值必须在[0, num_classes-1]范围内(因为F.cross_entropy默认标签是类别的索引,不是one-hot编码)。 - 举个例子:如果你的模型是5分类(输出维度是5),但标签里出现了5或者负数,
F.cross_entropy内部会触发断言,但因为CUDA操作是异步的,这个错误不会立刻抛出,而是延迟到后面的torch.sum这行才显现。 - 你可以在
outputs = model(inputs)之后加两行打印验证:print("Output shape:", outputs.shape) print("Label range:", labels.min().item(), labels.max().item())
2. 张量设备不匹配
虽然你能执行outputs = model(inputs),说明输入和模型在同一设备,但labels有可能和outputs不在同一个设备上(比如模型在CUDA,标签在CPU)。
- 这种情况下,
F.cross_entropy或者predicted==labels的比较会隐式触发设备同步,进而暴露之前的异步错误。 - 可以打印设备信息确认:
如果不一致,要确保print("Output device:", outputs.device) print("Label device:", labels.device)labels被移到模型所在的设备上(比如labels = labels.to(outputs.device))。
3. loss.detach的写法错误
你的代码里写了losses.append(loss.detach),这里detach是一个方法,应该写成loss.detach()才对——你现在是把方法对象添加到了列表里,而不是张量的数值。虽然这个错误不一定直接触发CUDA断言,但可能会导致后续的内存或计算异常,建议先修正这个写法。
快速验证思路
你可以先把calc_loss=False调用evaluate函数,如果这时候不报错,那基本可以确定问题出在F.cross_entropy(outputs, labels)这一步,也就是上面说的第一个原因(类别数和标签不匹配)。
内容的提问来源于stack exchange,提问作者Bowen Zhang
相关产品推荐
相关产品推荐

