torch.nn.Linear类工作原理及self.linear(x)传参逻辑疑问
关于torch.nn.Linear传参逻辑的解答
你提到的「self.linear没有定义接收参数的逻辑」是对PyTorch nn.Module封装机制的误解,具体原理如下:
class LR(nn.Module): def __init__(self, input_size, output_size): super().__init__() self.linear = Linear(input_size, output_size) def forward(self, x): pred = self.linear(x) return pred
- 代码里的
self.linear是torch.nn.Linear类的实例,所有继承自nn.Module的类都实现了Python的__call__魔法方法,支持像调用函数一样直接调用类实例,你传入的x首先会被传递到Linear实例的__call__方法中 nn.Module的__call__是PyTorch框架预先实现的通用逻辑,内部除了调用类的forward方法外,还会自动处理计算图记录、前后钩子触发等操作,不需要用户自行实现torch.nn.Linear官方已经内置了forward方法的实现,核心逻辑就是对输入做线性变换:output = x @ 权重矩阵.T + 偏置项,你传入的x最终就是在这里参与计算,得到的结果再赋值给pred
提示:不要直接手动调用
self.linear.forward(x),直接调用实例的写法会自动处理所有框架层的必要操作,能避免反向传播失效等异常问题
内容的提问来源于stack exchange,提问作者Vayun Ekbote
相关产品推荐
相关产品推荐

