为何StyleGAN3中torch.device('cuda',0)报错而torch.cuda.device(0)正常
PyTorch设备API差异及解决方案
报错差异的核心原因
两个API的设计逻辑完全不同,触发校验的时机不一样:
torch.device('cuda', 0)是设备描述符对象,创建阶段就会主动校验目标GPU的CUDA算力是否匹配当前PyTorch版本的最低要求。你遇到报错是因为当前使用的GPU算力低于官方编译的PyTorch版本支持的最低值(比如Kepler架构显卡算力3.x,PyTorch 1.10+版本已经停止支持该算力区间),因此初始化时直接抛出错误。torch.cuda.device(0)是CUDA上下文切换工具,仅负责将当前进程的默认CUDA设备切换到索引0的显卡,初始化阶段不会主动校验GPU算力,只有当实际在该设备上分配张量、运行CUDA算子时才会触发校验,所以你暂时没有遇到报错,本质是当前运行的代码逻辑还没有触碰到算力校验的环节。
是否需要全局替换代码
不建议直接全局替换,原因如下:
- 临时不报错的状态不稳定:后续如果代码涉及调用StyleGAN3内置的自定义CUDA算子、半精度计算等对算力有要求的操作,还是会触发运行时报错,严重时甚至会出现计算结果异常但无报错的静默故障。
- 两个API的使用场景不通用:
torch.device对象是PyTorch全栈支持的设备参数类型,比如创建张量时的device参数、模型迁移时的to()方法入参,都要求传入torch.device类型,直接替换为torch.cuda.device对象会触发类型错误。
更稳妥的修复方案
你可以根据自己的使用场景选择以下方案,不需要修改原有业务代码:
- 方案1:降级PyTorch和CUDA版本到匹配你显卡算力的版本,比如算力3.x的显卡可以降级到PyTorch 1.9 + CUDA 11.1的组合,原有
torch.device代码可以直接正常运行。 - 方案2:保留现有PyTorch版本,在运行代码前新增环境变量绕过算力校验:
该配置会让PyTorch跳过初始化阶段的算力校验,兼容旧显卡运行。# 把7.5替换为你自己显卡的算力值,比如算力3.5就填3.5 export TORCH_CUDA_ARCH_LIST="7.5+PTX" - 方案3:如果只是临时验证代码逻辑,不需要GPU加速,可以直接将设备修改为CPU运行,
torch.device('cpu')即可规避GPU相关的所有校验问题。
内容的提问来源于stack exchange,提问作者Laura Alvarez
相关产品推荐
相关产品推荐

