ARTICLE DETAIL

资讯详情

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

PyTorch模型载入实战:GPU与CPU互转的配置骨架与验证清单

PyTorch模型载入实战:GPU与CPU互转的配置骨架与验证清单 1. 为什么模型载入总在设备上翻车你手里大概率有一份训练好的.pth或.pt权重训练时用的是多卡服务器现在要拿到本地单卡甚至纯 CPU 机器上跑推理。直接torch.load然后load_state_dict十有八九会撞上RuntimeError: Attempting to deserialize object on a CUDA device but torch.cuda.is_available() is False或者Expected all tensors to be on the same device。这不是代码写错了而是 PyTorch 的序列化机制默认把张量绑在保存时的设备上。PyTorch 模型载入的本质是两件事一是把磁盘上的字节流反序列化成 Python 对象这一步由torch.load完成二是把权重字典塞进你定义好的网络结构里这一步由load_state_dict完成。GPU 和 CPU 互转的所有坑几乎都出在第一步——反序列化时张量被放到了哪个设备。map_location就是控制这件事的开关它决定了 checkpoint 里的每个张量最终落在 CPU 内存还是某块 GPU 显存上。这篇面向的是需要在本地笔记本、公司服务器、云主机之间来回搬模型的工程师。不管你是做 fine-tuning 还是纯推理下面这套配置骨架和验证清单可以直接抄。我试过在 4 卡 A100 上训完、拷到只有 CPU 的 Mac 上跑通中间踩的坑基本都覆盖到了。2. 载入前的环境确认与 TaoToken 接入准备在动手改代码之前先把运行环境摸清楚。很多人报错是因为根本没确认当前机器有没有 CUDA、有几块卡、PyTorch 是不是 GPU 版本。import torch print(torch version:, torch.__version__) print(cuda available:, torch.cuda.is_available()) print(cuda version:, torch.version.cuda) print(device count:, torch.cuda.device_count()) if torch.cuda.is_available(): print(current device:, torch.cuda.current_device()) print(device name:, torch.cuda.get_device_name(0))如果cuda available是False那所有.cuda()调用都会失败你必须走 CPU 载入路径。如果device count大于 1说明有多卡载入时要考虑DataParallel或DistributedDataParallel的前缀问题。如果你在调试过程中需要快速验证某个模型结构或让模型帮你解释报错可以用 TaoToken 的模型对话能力做辅助排查。它的接入地址是https://taotoken.net/api对话入口在https://taotoken.net/models?utm_sourcetaotoken_aicg_blog_endutm_mediumcsdnutm_campaignrewriteutm_contentmodel_chat。API Key 在控制台生成https://taotoken.net/console/api-keys?utm_sourcetaotoken_aicg_blog_endutm_mediumcsdnutm_campaignrewriteutm_contentapi_keys。接入文档在https://taotoken.net/doc?utm_sourcetaotoken_aicg_blog_endutm_mediumcsdnutm_campaignrewriteutm_contentdoc。这些只是辅助你查文档和验证思路核心的载入逻辑还是靠下面的代码骨架。3. 可复制的载入配置骨架3.1 map_location 的三种写法与选择map_location接受三种形式字符串、torch.device对象、可调用函数。字符串最直观函数最灵活。import torch # 写法一强制全部载入 CPU ckpt torch.load(model.pth, map_locationcpu) # 写法二载入到指定 GPU ckpt torch.load(model.pth, map_locationcuda:0) # 写法三函数式按张量原始位置动态决定 ckpt torch.load( model.pth, map_locationlambda storage, loc: storage.cpu() )函数式写法的好处是你可以根据loc判断原始设备做条件映射。比如原始存在cuda:1你想统一挪到cuda:0ckpt torch.load( model.pth, map_locationlambda storage, loc: storage.cuda(0) if loc.startswith(cuda) else storage )实测下来最稳的策略是先无条件载入 CPU再按需搬到目标设备。这样无论源文件来自几卡、什么版本都不会在反序列化阶段炸掉。3.2 单卡、多卡、CPU 三种目标形态的完整骨架下面这个骨架覆盖了从多卡 checkpoint 到单卡、CPU 的完整迁移路径。假设你有一个ModelArch类定义好了网络结构。import torch import torch.nn as nn class ModelArch(nn.Module): def __init__(self, num_classes10): super().__init__() self.backbone nn.Sequential( nn.Linear(512, 256), nn.ReLU(), nn.Linear(256, num_classes) ) def forward(self, x): return self.backbone(x) def load_checkpoint(path, devicecpu): 统一入口先把权重载到 CPU再按目标设备处理 ckpt torch.load(path, map_locationcpu) # 兼容两种保存格式纯 state_dict 或 包含 state_dict 的字典 if isinstance(ckpt, dict) and state_dict in ckpt: ckpt ckpt[state_dict] return ckpt def strip_dataparallel_prefix(state_dict): 去掉 DataParallel 保存时多出来的 module. 前缀 new_sd {} for k, v in state_dict.items(): if k.startswith(module.): new_sd[k[len(module.):]] v else: new_sd[k] v return new_sd # ---- 场景 A载入为 CPU 模型 ---- def build_cpu_model(path): model ModelArch() sd load_checkpoint(path, devicecpu) sd strip_dataparallel_prefix(sd) model.load_state_dict(sd) model.eval() return model # ---- 场景 B载入为单 GPU 模型 ---- def build_single_gpu_model(path, gpu_id0): device torch.device(fcuda:{gpu_id}) model ModelArch().to(device) sd load_checkpoint(path, devicecpu) sd strip_dataparallel_prefix(sd) model.load_state_dict(sd) model.eval() return model # ---- 场景 C载入为多 GPU 模型 ---- def build_multi_gpu_model(path, device_ids(0, 1)): model ModelArch().cuda(device_ids[0]) model nn.DataParallel(model, device_idslist(device_ids)) sd load_checkpoint(path, devicecpu) # 多卡模式下权重需要带 module. 前缀 if not any(k.startswith(module.) for k in sd.keys()): sd {fmodule.{k}: v for k, v in sd.items()} model.load_state_dict(sd) model.eval() return model关键点在于strip_dataparallel_prefix和反向加前缀这两个操作。多卡训练保存的权重键名带module.单卡或 CPU 网络没有这个前缀直接load_state_dict会报Missing key(s)和Unexpected key(s)。反过来如果你拿单卡权重去喂DataParallel包裹的模型也要手动补前缀。3.3 保存时的设备无关写法载入的对称操作是保存。为了让权重文件在设备间自由迁移保存时最好先搬到 CPUdef save_checkpoint(model, path): # 如果是 DataParallel取 .module 拿到原始网络 if isinstance(model, nn.DataParallel): model model.module state model.state_dict() # 全部转 CPU 再存文件与设备解耦 cpu_state {k: v.cpu() for k, v in state.items()} torch.save(cpu_state, path)注意state_dict()后面必须加括号这是方法调用不是属性。漏掉括号会得到function object has no attribute copy这类错误因为torch.save拿到的是一个函数对象而不是字典。4. 验证请求与成功结果载入完成后别急着跑全量推理先用一个假输入做前向验证确认设备一致、形状正确。def verify_model(model, devicecpu, input_dim512): model.eval() dummy torch.randn(2, input_dim).to(device) with torch.no_grad(): out model(dummy) print(input device:, dummy.device) print(output device:, out.device) print(output shape:, out.shape) assert out.device.type device.split(:)[0], 输出设备与预期不符 print(verify passed) # CPU 验证 cpu_model build_cpu_model(model.pth) verify_model(cpu_model, devicecpu) # 单卡验证 if torch.cuda.is_available(): gpu_model build_single_gpu_model(model.pth, gpu_id0) verify_model(gpu_model, devicecuda:0)成功时你会看到类似输出input device: cpu output device: cpu output shape: torch.Size([2, 10]) verify passed如果走 GPU 路径output device应该显示cuda:0。这一步能同时抓出设备不匹配和维度错误两类问题。另外建议检查一下模型参数的设备分布def check_param_devices(model): devices set() for name, param in model.named_parameters(): devices.add(str(param.device)) print(parameter devices:, devices) return devices正常情况下所有参数应该在同一个设备上。如果出现{cpu, cuda:0}混合说明有部分层没搬过去前向传播时必然报Expected all tensors to be on the same device。5. 本篇常见错误排查5.1 RuntimeError: Attempting to deserialize object on a CUDA device这是最高频的报错。原因是你在一台没有 CUDA 的机器上torch.load了一个 GPU 保存的权重且没指定map_location。解决方式就是加map_locationcpu。注意这个参数要加在torch.load上不是load_state_dict上。5.2 Missing key(s) / Unexpected key(s) 里的 module. 前缀报错信息里如果出现大量module.xxx的 unexpected key说明权重是多卡保存的而你的网络是单卡结构。用第 3.2 节的strip_dataparallel_prefix处理。反过来如果 missing key 全是module.xxx说明你的模型被DataParallel包了但权重是单卡的需要补前缀。5.3 size mismatch for xxx.weight形状对不上通常是网络结构定义和保存时不一致。比如分类数从 10 改成了 5最后一层Linear的权重形状就变了。这种要么改网络定义要么用strictFalse跳过不匹配的层missing, unexpected model.load_state_dict(sd, strictFalse) print(missing:, missing) print(unexpected:, unexpected)但strictFalse是权宜之计生产环境要确认跳过的层是不是你真正想跳的。5.4 显存不足但模型明明不大如果你先把权重载到 GPU 再转 CPU中间会占用一份显存峰值。正确顺序永远是map_locationcpu先落地再.to(device)搬到目标设备。另外多卡载入时storage.cuda(0)只把张量放到 0 号卡如果 0 号卡被占满也会 OOM可以换成空闲卡。5.5 torch.load 的 weights_only 参数较新版本的 PyTorch 对torch.load增加了weights_only参数默认行为在版本间有变化。如果你载入的是纯权重字典可以显式写torch.load(path, map_locationcpu, weights_onlyTrue)来避免反序列化任意对象的安全风险。如果 checkpoint 里还存了优化器状态、epoch 等自定义对象则需要weights_onlyFalse。6. 长期编码与 Agent 场景的接入建议如果你经常需要在不同设备间迁移模型、跑 fine-tuning 或搭推理服务建议把上面这套载入骨架封装成一个独立模块项目里统一调用避免每个脚本各写一套。对于需要长期跑编码任务、Agent 工作流的场景可以了解 TaoToken 的 Coding Planhttps://taotoken.net/coding-plan?utm_sourcetaotoken_aicg_blog_endutm_mediumcsdnutm_campaignrewriteutm_contentcoding_plan。Claude Code 相关的接入说明在https://taotoken.net/doc/claudecode-anthropic?utm_sourcetaotoken_aicg_blog_endutm_mediumcsdnutm_campaignrewriteutm_contentclaudecode。这些入口适合把模型载入、设备切换这类重复劳动交给工具链处理你专注在模型结构本身。最后留一个实用习惯每次保存 checkpoint 时在文件名里带上设备和卡数信息比如model_dp2_cuda.pth或model_cpu.pth。载入前先看一眼文件名能省掉一半的排查时间。
返回列表