是否可实现上下文管理器自动处理PyTorch张量转移到CUDA设备?
PyTorch设置默认设备/上下文自动指定设备方案
方案1:全局设置默认设备(PyTorch 1.12及以上版本支持)
直接通过torch.set_default_device()全局指定所有新创建张量的默认设备,后续所有张量工厂函数(randn/zeros/ones/tensor等)生成的张量都会自动放在指定设备上,无需逐一声明:
# 检查CUDA可用后设置默认设备 if torch.cuda.is_available(): torch.set_default_device('cuda') # 后续生成的张量默认在cuda上 A = torch.randn(1, 10) B = torch.randn(10, 20) C = A @ B # 运算自动在cuda上执行
如果需要切回CPU默认,执行torch.set_default_device('cpu')即可。
方案2:上下文管理器局部生效(完全符合需求的用法)
PyTorch原生支持torch.device作为上下文管理器,仅对作用域内新创建的张量生效,不会修改全局默认配置:
cuda_device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') # 仅该代码块内的张量自动使用cuda设备 with torch.device(cuda_device): A = torch.randn(1, 10) B = torch.randn(10, 20) C = A @ B # 运算在cuda上执行,结果也存储在cuda上
低版本PyTorch兼容方案(1.12以下版本)
如果使用的是旧版本PyTorch,可以通过设置默认张量类型实现类似效果,该方法仅对浮点类型张量生效:
if torch.cuda.is_available(): torch.set_default_tensor_type('torch.cuda.FloatTensor') # 后续创建的浮点张量默认在cuda上 A = torch.randn(1, 10) B = torch.randn(10, 20)
注意事项
- 以上方案均不影响你使用
torch.nn.functional.conv2d的自定义层逻辑,只要输入张量已经在CUDA上,运算会自动调度到GPU执行 - 如果是从外部导入的numpy数组转PyTorch张量,
torch.tensor(numpy_array)也会自动遵循默认设备设置,无需额外调用to()方法
内容的提问来源于stack exchange,提问作者dom
相关产品推荐
相关产品推荐

