You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

子类模块重写nn.Module的.to方法在父类调用时无效的问题

原因分析
  • PyTorch的nn.Module.to()方法在递归处理子模块时,内部调用的是子模块的._apply()方法,而非直接调用.to()。你重写的.to()只有在直接调用时才会触发,父类递归走的是底层_apply流程,所以不会执行自定义逻辑。
  • 虽然.to()最终会调用_apply(),但你重写的to()并没有改变_apply的行为,所以父类的递归迁移不会触发你的代码。
解决方法

要让父类调用.to()时也触发子类的自定义逻辑,需要重写_apply()方法——这是PyTorch模块设备迁移的核心入口。修改Child类如下:

import torch
import torch.nn as nn


class Child(nn.Module):
    def __init__(self, device):
        super().__init__()
        self.device = device
        self.param = nn.Parameter(torch.randn((3,2)))

    def forward(self,x):
        print("nothing again")
        return x
    
    def _apply(self, fn):
        # 先调用父类的_apply完成参数和子模块的设备迁移
        super()._apply(fn)
        # 执行自定义逻辑
        print("SPECIAL")
        # 从参数中获取目标设备
        self.device = next(self.parameters()).device
        return self
    
    def to(self,*args,**kwargs) :
        super().to(*args,**kwargs)
        print("SPECIAL")
        if("device" in kwargs) :
            self.device=kwargs["device"]
        return self

class Dad(nn.Module) :
    def __init__(self):
        super().__init__()
        self.child = Child('cuda')
    
    def forward(self,x):
        x=self.child(x)
        print("nothing")
        return x


daron = Dad()
daron.to('cpu')
print("Param is now : ",daron.child.param.device)
print("Child device is now : ",daron.child.device)
daron.child.to('cuda')
print("Param is now : ",daron.child.param.device)
print("Child device is now : ",daron.child.device)

说明

  • 重写_apply()方法后,无论是父类递归调用还是直接调用子模块的.to(),都会触发自定义逻辑:父类调用时走_apply流程,直接调用.to()时则会先执行父类的to()(内部调用_apply),再执行你重写的to()里的自定义代码。
  • 通过next(self.parameters()).device获取目标设备,是因为_apply接收的是一个设备转换函数,无法直接拿到kwargs里的设备参数,这种方式能准确获取当前模块参数所在的设备。

内容的提问来源于stack exchange,提问作者Frotaur

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.01 06:16:08