PyTorch自回归生成AI模型预训练CUDA断言错误排查修复
排查与修复CUDA断言错误方案
核心可能原因及对应修复
1. 词表索引越界(最常见触发场景)
CUDA断言错误常隐藏输入/标签中存在超出模型词表大小的索引问题——虽然错误栈指向x.to(device),但实际是后续计算(如交叉熵损失)触发断言,只是栈追踪延迟显示到设备转移步骤。
- 排查:打印输入、标签的最大索引,和模型词表大小对比:
print("输入最大索引:", x.max().item()) print("标签最大索引:", labels.max().item()) print("模型词表大小:", model.vocab_size) - 修复:若索引等于或大于词表大小,修正tokenizer编码逻辑,确保未知词使用
<unk>对应的合法索引(小于词表大小);或调整模型词表参数匹配数据中的最大索引。
2. 输入/标签张量形状不匹配
模型前向传播时,输入与标签的维度、形状不一致会触发底层断言。
- 排查:打印输入和标签的形状:
自回归任务中,两者需保持维度一致(如均为print("输入形状:", x.shape) print("标签形状:", labels.shape)[batch_size, seq_len]),若采用偏移标签逻辑(标签为输入的后移版本),需确认切片操作正确(如input_ids = x[:, :-1],labels = x[:, 1:])。 - 修复:调整数据加载器的输出格式,保证输入与标签形状完全匹配;修正切片逻辑避免维度错位。
3. 设备一致性问题
虽然执行了x.to(device),但模型参数或辅助张量(如注意力掩码)仍留在CPU,会导致后续计算时设备不匹配触发断言。
- 排查:检查模型和输入的设备信息:
print("模型所在设备:", next(model.parameters()).device) print("输入转移前设备:", x.device) - 修复:训练前将模型完整转移到目标设备:
同时确保注意力掩码、位置编码等辅助张量同步转移到对应设备。model = model.to(device)
4. 数据集存在异常样本
数据集中存在长度为0、形状异常的样本,会导致张量转移时触发底层断言。
- 排查:在数据加载的
collate_fn中添加过滤逻辑:
训练循环中跳过返回def collate_fn(batch): # 过滤空样本 batch = [item for item in batch if len(item["input_ids"]) > 0] if not batch: return None return torch.utils.data.dataloader.default_collate(batch)None的批次。
5. CUDA内存异常
内存不足可能导致张量转移时触发隐性断言错误(并非直接OOM报错)。
- 排查:降低
batch_size,关闭非必要的内存占用操作(如梯度检查点外的冗余张量),观察是否仍触发错误。 - 修复:启用混合精度训练减少内存占用:
定期调用scaler = torch.cuda.amp.GradScaler() with torch.cuda.amp.autocast(): outputs = model(x) loss = criterion(outputs, labels) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update()torch.cuda.empty_cache()清理闲置内存。
调试技巧
- 先在CPU环境下跑小批量数据,CPU的错误信息会直接定位到触发点,比CUDA的延迟断言更清晰。
- 设置环境变量
CUDA_LAUNCH_BLOCKING=1,强制CUDA同步执行,让错误在实际触发点抛出,而非延迟到设备转移步骤,获得更准确的栈追踪。
内容的提问来源于stack exchange,提问作者Betu Raja
相关产品推荐
相关产品推荐

