PyTorch调用to(device)时嵌套子模块的参数会同步迁移到目标设备吗
结论
inside类中的self.weight_h和self.weight_v会自动迁移到GPU2,不需要单独为inside类的参数指定迁移操作- 该逻辑适用于所有PyTorch>=1.0版本,你当前使用的1.8.1版本完全兼容
原理说明
PyTorch的nn.Module内置了递归管理子模块、参数、缓冲区的机制:
- 当你在
main_class的__init__方法中将inside()实例赋值给self.in_class时,inside实例会被自动注册为main_class的子模块 - 调用模块的
.to(device)方法时,PyTorch会递归遍历当前模块的所有参数、缓冲区,以及所有已注册子模块的参数、缓冲区,统一迁移到指定设备上 - 你的代码里
self.weight_h是nn.Parameter类型,self.weight_v是nn.Conv2d(属于nn.Module子类),都会被inside模块自动注册为自身的参数/子模块,自然会被上层的.to()操作递归处理
快速验证方法
你可以在执行object = object.to(device)后添加以下代码确认参数所在设备:
# 输出weight_h的设备 print(object.in_class.weight_h.device) # 输出weight_v卷积层参数的设备 print(object.in_class.weight_v.weight.device)
输出结果会显示cuda:2(假设你的GPU2对应的CUDA设备编号为2),即可确认参数已正确迁移到目标GPU。
注意事项
只有被注册为模块属性的子模块/参数才会被自动处理,如果你将子模块存放在普通Python列表/字典中,需要改用nn.ModuleList/nn.ModuleDict包裹才能触发自动注册和设备迁移逻辑。
小建议:代码中尽量不要使用
object作为变量名,该名称是Python内置的基类标识符,重名可能引发不可预期的问题。
内容的提问来源于stack exchange,提问作者Mohit Lamba
相关产品推荐
相关产品推荐

