PyTorch autocast报TypeError:不识别device_type/dtype参数求助
问题原因与解决办法
报错复现
测试代码:
with autocast(device_type='cuda', dtype=torch.float16): a=0.00000000000000001 print(a)
报错信息:
Traceback (most recent call last): File "<stdin>", line 1, in <module> TypeError: __init__() got an unexpected keyword argument 'device_type'
另外在分类模型训练循环中执行with autocast(dtype=self.precision):时,同样触发错误:TypeError: __init()__ got an unexpected keyword argument 'dtype'
环境相关版本:
- pytorch 1.9.0 py3.8_cuda11.1_cudnn8.0.5_0
- python 3.8.15
- cudatoolkit 11.1.74
原因分析
这是PyTorch版本差异导致的问题,和cudatoolkit无关:
- PyTorch 1.9.0仅提供
torch.cuda.amp.autocast接口,这个接口仅支持enabled一个参数,不支持device_type和dtype; device_type参数是PyTorch 1.10+版本中torch.amp.autocast接口新增的,用于同时支持CPU和CUDA的自动混合精度;dtype参数也是后续版本才加入的自定义精度选项;- 你写的代码是基于高版本PyTorch的API,在1.9.0版本中自然无法识别这些参数。
解决办法
方案1:适配PyTorch 1.9.0写法
如果不想升级版本,直接使用CUDA专属的自动混合精度接口,写法如下:
from torch.cuda.amp import autocast # 默认启用自动混合精度,默认用float16 with autocast(): a=0.00000000000000001 print(a) # 也可以通过enabled参数控制开关 with autocast(enabled=True): # 训练逻辑 pass
注意:PyTorch 1.9.0的autocast不支持手动指定dtype,只能使用默认的float16。
方案2:升级PyTorch版本
如果需要使用device_type、dtype等参数,建议升级到PyTorch 1.10及以上版本,执行以下conda命令(适配你的cuda11.1环境):
conda install pytorch>=1.10.0 torchvision torchaudio cudatoolkit=11.1 -c pytorch
内容的提问来源于stack exchange,提问作者smaragda ben
相关产品推荐
相关产品推荐

