如何在避免命名冲突且不导入完整PyTorch模块的情况下,以torch.device形式引用该对象?
解决PyTorch device导入冲突且保持torch.device调用风格的方法
首先得明确:Python里as后面只能跟简单标识符,不能像torch.device这种带点的路径,所以你之前尝试的from torch import device as torch.device才会报语法错误。不过有几个巧妙的方法能实现你想要的效果,既不用导入整个PyTorch模块,又能继续用torch.device的方式引用,还避开名称冲突:
方法一:用模块对象模拟torch命名空间
你可以手动创建一个小型的模块对象,把导入的device挂载到它的device属性上,完美模拟原有的调用方式:
import types # 创建一个名为"torch"的空模块对象 torch = types.ModuleType('torch') # 只导入需要的device,挂载到模块对象上 from torch import device torch.device = device
这样之后,你就能在代码里正常用torch.device("cuda")这类写法了,而且完全没有导入整个PyTorch模块,只加载了device相关的部分。
方法二:用类作为命名空间
如果觉得创建模块对象有点麻烦,用一个简单的类来做命名空间也可以:
class TorchNamespace: # 在类内部直接导入device,相当于把它变成类属性 from torch import device # 实例化类(或者直接用类本身也可以) torch = TorchNamespace()
之后同样可以通过torch.device来访问,效果和方法一一致,代码更简洁。
这两种方法都比torch_device这类别名更贴近你原本的使用习惯,而且严格控制了导入的内容,避免了加载整个庞大的PyTorch模块。
内容的提问来源于stack exchange,提问作者Jonas De Schouwer
相关产品推荐
相关产品推荐

