最新版PyTorch设置张量默认使用指定设备的最佳实践与注意事项
PyTorch 全局设置默认张量设备的实现方案
最新稳定版PyTorch(2.x全系列,1.13及以上版本均支持)提供了官方原生API实现全局默认设备配置,不需要逐行给张量创建接口加device参数或者.cuda(),程序初始化阶段跑一行torch.set_default_device()就够了。
示例代码:
import torch # 程序启动时配置全局默认设备,支持cuda、mps、cpu等所有合法设备标识 torch.set_default_device("cuda:0") # 后续所有未显式指定device参数的张量工厂函数,都会直接在目标设备创建张量 x = torch.randn(3) print(x.device) # 输出: cuda:0 y = x + 5 print(y.device) # 输出: cuda:0
该API支持传入设备字符串或者torch.device对象,比如苹果硅设备可以传入"mps",多卡场景可以指定具体卡编号如"cuda:1"。如果张量创建时显式传入了device参数,会优先使用显式指定的设备,不受全局配置影响。
别用网上流传的猴子补丁修改torch.randn、torch.tensor这类工厂函数的野路子,这类写法完全没有兼容性保障,PyTorch版本一更新很容易直接崩,甚至出静默错误算错结果都发现不了。
全局设置默认设备的弊端与不适用场景
这种全局配置虽然能简化单设备场景的代码编写,但存在不少明显局限,以下场景不建议使用:
- 第三方库兼容性差:很多基于PyTorch开发的第三方工具库内部默认按CPU逻辑编写张量创建、运算流程,全局修改默认设备很容易触发库内部的设备不匹配报错,甚至出现无感知的跨设备数据拷贝,大幅拖慢运行速度。
- 小张量运算性能反降:GPU执行极小张量(比如标量、维度小于10的向量)的创建、简单运算时,PCIe传输、kernel调度的开销远高于运算本身,这类操作默认放在GPU上运行速度反而不如CPU。
- 显存管控难度上升:全局默认GPU的情况下,调试用的临时张量、逻辑分支里的冗余张量都会默认占用显存,很容易出现无意义的显存占用,甚至触发显存不足(OOM),排查显存问题的成本远高于手动指定设备的写法。
- 多设备/分布式场景适配性差:如果代码需要同时使用CPU做数据预处理、多块GPU做模型并行/分布式训练,全局默认单设备会打乱跨设备张量流转的逻辑,很容易出现隐式的设备不匹配问题,调试成本极高。
实际使用建议:如果只是自己写单卡训练、推理的独立小脚本,全局设置默认设备确实能省很多重复传参的功夫;但如果是写供他人调用的工具库、或者涉及多设备/分布式的复杂项目,老老实实显式传递
device参数才是最稳妥的做法。
内容的提问来源于stack exchange,提问作者chausies
相关产品推荐
相关产品推荐

