ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

PyTorch多卡训练报错:Expected all tensors on same device 排查与修复

PyTorch多卡训练报错:Expected all tensors on same device 排查与修复 先说结论这不是显卡坏了也不是 CUDA 装错了而是训练循环里的张量散落在两张卡上。你看到的报错核心是 PyTorch 在告诉你当前参与运算的 Tensor 至少有两个不同的 device一个在cuda:2另一个在cuda:0它没办法自动帮你搬运只能直接中断。这篇文章写给所有正在用多卡训练、或者刚把代码从单卡改到多卡的朋友我会把这条报错的原理、定位方法、常见触发场景和最终修复方案一次讲清楚。我见过很多人第一次看到cuda:2或cuda:0时会下意识去检查驱动、重装 CUDA其实多半没那个必要。驱动只要能跑nvidia-smiCUDA 版本和 PyTorch 能匹配问题基本都出在代码里。接下来我们用一套可复现的排查流程把这条报错连根拔掉。1. 这行报错到底在说什么1.1 先读懂消息里的关键信息完整错误通常长这样RuntimeError: Expected all tensors to be on the same device, but found at least two devices, cuda:2 and cuda:0!个别 PyTorch 版本会把顺序反着写比如cuda:0 and cuda:2意思没有区别。它只是在说某一处操作同时接收了来自cuda:2的张量和来自cuda:0的张量而 PyTorch 不允许这种跨设备直接参与同一个运算。注意这句话并不是在判定“哪一个是错的”。顺序不代表责任方它只是把参与运算的设备名列出来。很多人在排查时误以为“报错顺序在前的是罪魁祸首”结果浪费了很多时间。我建议你只把它当成“现场提示”真正的源头还要结合调用栈去找。另一个容易混淆的点是设备编号。cuda:2指的是当前进程可见 GPU 列表中编号为 2 的显卡而不是物理机上插槽编号为 2 的卡。如果你设置了CUDA_VISIBLE_DEVICES逻辑编号和物理编号不一定一致这个后面会细说。1.2 为什么 PyTorch 不肯“自动帮你搬”很多人第一次遇到这个错都会问数据在cuda:0模型在cuda:2为什么 PyTorch 不自己 copy 一下答案很简单自动搬会带来巨大的性能隐患。深度学习训练中张量经常要被搬运几百次、上千次。如果 PyTorch 在每次运算前都默默加一次设备检查加一次 copy那训练速度会严重退化而且这种隐式拷贝很难被察觉。PyTorch 的设计哲学是“显式优于隐式”尤其是设备分配这种代价很高的操作必须由开发者自己负责。打个比方你想让两台电脑里的文件合并到一个 Excel 表里程序可以提示你“文件不在同一台电脑”但它不该擅自把文件通过网络传过来。谁来传、传到哪、什么时候传都得你自己决定。所以当你看到这条报错真正的问题不是“PyTorch 不智能”而是“代码里某个环节没有把设备统一好”。我们接下来就是要把这个环节找出来。2. 最常见的触发场景模型和设备散落一地2.1model.to(cuda:2)之后模型真的全过去了吗不少人觉得只要对模型调用了.to(cuda:2)整个模型就都在cuda:2上了。这个想法对大多数情况成立但有三个例外第一optimizer 的状态不会跟着 model 迁移。比如 Adam 里的 moment 和 variance 张量在初始化时可能还在 CPU 上。如果 optimizer 是在model.to(device)之前创建的optimizer.state里的张量就一直留在 CPU。训练中一旦用某个自定义逻辑直接读取 optimizer 状态去计算就可能冒出 CPU 和 GPU 混合访问。第二某些 buffer 或临时变量不是通过register_buffer注册的。如果模型类里定义了一个普通self.something torch.zeros(...)然后在前向里让它参与运算这个 tensor 不会随模型迁移。它默认在 CPU最终造成的报错往往是CPU tensor和cuda tensor不一致。第三调用model.to(device)后再重新给模型的某个子模块赋值新张量。例如在外部修改model.head.weight.data new_tensor如果这个new_tensor是在cuda:0上创建的权重就和模型主体分家了。所以模型主体已经迁移不代表所有相关张量都在同一设备上。这也是为什么推荐在训练循环里统一管理所有张量而不是到处散落.to()。2.2 DataLoader 或 Dataset 里提前搬到 GPU这个坑在多卡环境下特别常见。有些人为了让代码“更快”在Dataset.__getitem__里直接返回image.to(cuda:0)或者把 label 预先.cuda()。单卡时看着没问题一旦开多卡问题就来了DataLoader 的 worker 进程可能跑在不同线程/进程里设备上下文不一定是主进程的设备。就算能跑通每个 batch 都要经过额外的 H2D 或 D2D 拷贝性能不升反降。返回的 tensor 被固定在了某个绝对设备号上比如cuda:0而模型可能在cuda:2于是每次前向都会触发这条报错。正确的做法是Dataset 只负责产出 CPU tensor由主进程统一.to(device)。这样既干净又灵活换卡、换机器都不用改 Dataset。2.3 加载 checkpoints 时忽略了map_location这是一个非常隐蔽的坑。当你用torch.load(model.pt)加载一个在cuda:2上保存的模型时如果没有指定map_locationPyTorch 会尝试把张量放到原本保存时的设备上。假如当前环境里可见设备编号发生了变化或者当前进程默认设备是cuda:0就很容易产生“模型一部分在cuda:2一部分在cuda:0”的状态。更麻烦的是这种错误往往不是立刻出现而是等模型跑了一两步之后某个层开始运算时才炸。因为它取决于 checkpoint 里状态的设备分布很难一眼看出来。解决方式很简单后面会专门讲加载时统一map_location再显式.to(device)。3. 排查设备不一致的标准流程3.1 先拿到完整调用栈别只盯着最后一句话遇到这种报错我个人的第一个动作是看 traceback而不是查 CUDA 版本。完整的错误信息会告诉你出错的具体代码行但最后一行往往只是底层算子抛出的。你需要向上翻找到第一个和自己项目代码相关的调用栈。比如如果 traceback 最后停在 loss.cross_entropy那就要去看你在调 loss 之前传给它的logits和labels分别是哪个 device。如果停在某个自定义 Loss 类内部那就去检查这个类里手动创建的张量。不要觉得看调用栈浪费时间这比盲改代码高效得多。通常 80% 的这类报错都能在 5 分钟内从调用栈里定位到。3.2 用 print 和 assert 快速确认现场定位出嫌疑代码行之后在该行之前打印关键张量的设备print(input device:, inputs.device) print(label device:, labels.device) print(model device:, next(model.parameters()).device)如果某个打印结果和你预期不一致那就是问题根源。如果你不想一次一次手动打印可以直接插断言让程序第一时间告诉你assert inputs.device labels.device, finputs: {inputs.device}, labels: {labels.device}这种方式适合在训练脚本里临时加定位到问题后删掉即可。不要嫌麻烦这种报错如果靠猜可能折腾一下午靠断言两三分钟就能水落石出。3.3 用CUDA_VISIBLE_DEVICES缩小范围如果报错只出现在多卡场景中单卡时一切正常那么大概率是代码里写了某个绝对设备号。你可以在启动命令前面强制只暴露一张卡CUDA_VISIBLE_DEVICES0 python train.py如果单卡正常就证明模型和数据是统一迁移的问题出在多卡编号上。如果单卡也报错比如报cuda:2越界那就是代码里硬编码了超出可见范围的设备号。有一点必须强调CUDA_VISIBLE_DEVICES2,3会把你物理上的 2 号卡和 3 号卡重新映射成逻辑编号 0 和 1。也就是说原来的cuda:2、cuda:3在进程里会变成cuda:0、cuda:1。如果你代码里写死的cuda:2不变它实际指向的是逻辑设备 2但逻辑设备只有 0 和 1必然越界。这种重映射机制是很多老手都会踩的坑。所以最佳实践是代码里不要依赖物理绝对编号尽量用rank或环境变量赋值。4. 修复方案统一到同一个 device 上4.1 推荐一套“万能”的to_device工具函数既然报错根源是设备不统一那我们就写一个统一的映射入口保证整个流程只调用一次迁移逻辑。这个函数可以递归处理 Tensor、字典、列表和元组非常实用def to_device(data, device): if isinstance(data, torch.Tensor): return data.to(device, non_blockingTrue) if isinstance(data, dict): return {k: to_device(v, device) for k, v in data.items()} if isinstance(data, list): return [to_device(x, device) for x in data] if isinstance(data, tuple): return tuple(to_device(x, device) for x in data) return data然后在训练循环里统一调用device torch.device(cuda if torch.cuda.is_available() else cpu) model model.to(device) for batch in dataloader: batch to_device(batch, device) inputs batch[image] labels batch[label] outputs model(inputs) loss loss_fn(outputs, labels) ...这样能避免在训练脚本里到处散落.to()也避免你漏掉某个字段。使用non_blockingTrue是想在异步拷贝场景下有一点性能收益单卡时没有副作用多卡时配合pin_memory效果更好。4.2 模型加载时不要忘了map_location我推荐的加载方式是先统一到 CPU再把模型挪到目标 device。代码长这样checkpoint torch.load(model.pt, map_locationcpu) model.load_state_dict(checkpoint[model_state_dict]) model model.to(device)这样写的好处是无论 checkpoint 当初保存在哪个 GPU 上加载时都不会触发设备推断。你先把全部权重拉到 CPU再一次性.to(device)从源头杜绝“一部分在cuda:0、一部分在cuda:2”的情况。如果你希望保留部分层在 CPU 上跑推理比如 CPU 上做预处理那也要清楚知道自己具体想怎么做而不是把设备交给torch.load自动判断。4.3 DataParallel / DDP 的正确姿势先说DataParallel它是单进程多线程的封装通常会这样用model nn.DataParallel(model, device_ids[0, 1])它会自动把输入 batch 切分到device_ids里各张卡上前向计算完后把输出收集到output_device默认是device_ids[0]也就是逻辑cuda:0上。问题出在损失计算如果你的 labels 还是在cuda:1上而模型输出已经回到cuda:0那在计算 loss 时就会碰到cuda:0和cuda:1不一致。最好的做法是不要依赖 DataParallel 自动切分数据而是手动搬到output_device再交给模型。常见的方式是把数据统一导到主卡device torch.device(cuda, 0) inputs inputs.to(device) labels labels.to(device) outputs model(inputs) loss loss_fn(outputs, labels)但现在真的不建议新项目用DataParallel了。它在多进程、多机扩展和性能上都有限制。新项目请直接上DistributedDataParallelDDP。DDP 的正确启动方式核心是每个进程只负责一张卡import os import torch import torch.distributed as dist def main(): rank int(os.environ[RANK]) local_rank int(os.environ[LOCAL_RANK]) torch.cuda.set_device(local_rank) device torch.device(fcuda:{local_rank}) model create_model().to(device) ddp_model torch.nn.parallel.DistributedDataParallel(model, device_ids[local_rank]) for batch in dataloader: batch to_device(batch, device) ...关键就是torch.cuda.set_device(local_rank)。如果不设置torch.cuda.current_device()可能一直默认是cuda:0哪怕你在别的进程里手动把模型.to(cuda:2)后续一些隐式操作还是会在cuda:0上分配张量进而引发各种设备打架。5. 现场问题速查表5.1 按错误信息类型检索我整理了一张速查表你在排查时可以对照一下报错特征常见原因直接解决办法报错提到cuda:0和cuda:2模型在cuda:2输入数据在cuda:0或反过来打印模型参数设备和输入设备用统一to_device搬运报错提到cuda和 CPU某个 tensor 直接在 CPU 上创建忘记迁移检查模型内普通 tensor、optimizer 状态、损失函数内部常量报错提到cuda:2但设备只有cuda:0/1代码里写死了绝对设备号cuda:2但不可见用CUDA_VISIBLE_DEVICES重映射或改用rank加载 checkpoint 后才报错torch.load未指定map_location状态残留在原来设备加载时加map_locationcpu再.to(device)只在 DataParallel 下报错输出在主卡cuda:0labels 还在其它卡把 labels 手动搬到output_device或换成 DDP只在 DDP 多进程下报错每个进程没有正确set_device全局默认还是cuda:0在初始化后执行torch.cuda.set_device(local_rank)这张表不能覆盖所有情况但能覆盖我遇到过的绝大多数情况。只要你能定位到“是哪两个设备”再结合调用栈去看“谁被分配到了错设备”基本就解决了。5.2 容易忽略的隐性来源除了显式的输入输出还有几个隐性来源特别容易漏掉自定义损失函数内部的常量比如torch.tensor(0.5)它默认在 CPU一旦和 CUDA tensor 相乘就会报设备不一致。解决办法是写torch.tensor(0.5, devicedevice)或者直接用 Python 标量。自定义层的权重初始化如果初始化为torch.zeros(...)然后放到一个 CUDA tensor 参与运算也会触发问题。初始化时最好用torch.zeros(..., devicedevice)。评估指标计算在计算 mAP、accuracy 时预测结果可能在cuda:0标签可能在cuda:1一样会报错。统一过to_device后再算。多个 dataloader 的不同返回结构一个返回dict一个返回tuple如果你的to_device只处理了 tuple却忽略了 dict就会漏掉部分张量。分布式通信算子的返回值dist.all_gather返回的是一个张量列表列表里每个张量可能来自不同 rank它仍会保留原 rank 的设备。使用时需要再.to(device)或明确知道目标设备。6. 我踩过的坑和后续建议6.1 别在代码里写死物理设备号这个坑我复盘过很多次。早期我习惯在配置文件里写gpu_ids [2, 3]然后在代码里直接拼cuda:{id}。单机跑没问题一旦换到另一台 GPU 插槽顺序不同的机器或者某张卡被占用需要重新指定可见设备时代码就各种报警。后来我改成“启动脚本指定可见设备训练代码只用相对编号”。比如export CUDA_VISIBLE_DEVICES2,3 python train.py --local_rank 0在训练代码里设备号直接用local_rank或者RANK不要和物理编号扯上关系。这样无论是单机多卡还是多机多卡代码都不用改。6.2 训练脚本最稳的样板一开始就统一 device我现在的习惯是新建训练脚本第一件事就是定义设备入口。之后再写任何涉及张量的代码都从这里拿 deviceimport torch def get_device(): if not torch.cuda.is_available(): return torch.device(cpu) local_rank int(os.environ.get(LOCAL_RANK, 0)) torch.cuda.set_device(local_rank) return torch.device(fcuda:{local_rank}) device get_device()然后模型、数据、损失函数里所有需要构造张量的地方全部从device变量取。像torch.zeros(..., devicedevice)、torch.tensor(..., devicedevice)不直接写死cuda:0。这个方法看起来简单却很有效。因为它把“设备”变成了一个变量而不是散落在代码里的魔法数字。以后换环境、换卡、做分布式都只是一行配置的差别。6.3 最后再分享一个习惯我遇到这种报错时现在的第一反应不是去改代码而是先在出错点上方打印两个设备的实际值。很多时候问题出在“设备编号比你想象的多了一维”你以为所有数据都在cuda:0实际上 dataloader 的某个 worker 已经悄悄返回了cuda:2的张量你以为model.to(device)已经把模型搬过去了但 optimizer 的 state 还留在 CPU。把这些信息打出来看到真实分布解决方案就自己浮出来了。对于正在做多卡训练的朋友我的建议是把设备管理当成和“数据清洗”一样重要的事情所有张量的来源和去向都做到可控、可打印、可断言。这样不仅今天这个cuda:2 / cuda:0的报错会消失以后很多更难查的并发问题也会少一大半。
返回列表