PyTorch自定义maxout类forward方法未通过__init__传参为何能获取输入x
问题解答
核心逻辑说明
你当前的困惑主要是对PyTorch nn.Module 的运行规则、以及「固定超参数」和「动态输入」的传递逻辑有误解:
__init__方法的作用是初始化模块的固定配置、可学习参数等不会随前向输入变化的内容,比如你maxout类的num_pieces属于模型结构固定超参数,所以放在__init__传入并保存为实例属性;而forward的入参x是每次前向传播都会变化的输入张量,本来就不需要在初始化阶段传入。- 所有继承自
torch.nn.Module的子类,PyTorch都内置封装了__call__魔法方法:当你直接调用模块实例(比如你实例化的maxout(5))时,会自动触发forward方法执行,同时把调用时传入的参数直接透传给forward作为入参。举个简单的测试例子:
import torch test_x = torch.randn(2, 240) # 模拟Linear层输出的240维张量 maxout_inst = maxout(5) res = maxout_inst(test_x) print(res.shape) # 输出为torch.Size([2, 48]),和你网络中的计算结果一致
上述代码中maxout_inst(test_x)本质等价于maxout_inst.forward(test_x),不需要提前把test_x传给__init__。
nn.Sequential的输入传递规则
你用到的nn.Sequential是PyTorch的顺序容器,会自动按你定义的顺序执行内部各个模块的前向计算,自动完成输入传递:
- 你给
self.fcn传入输入inputs时,首先会送给第一个Linear层,得到240维的输出张量 - Sequential自动把这个Linear层的输出作为调用参数,传给下一个模块也就是
maxout(5)的实例,相当于自动执行maxout_inst(linear_output),自然就把Linear的输出作为x传给了maxout的forward方法 - 后续maxout输出的48维张量,会再自动传给第二个Linear层完成剩余计算
内容的提问来源于stack exchange,提问作者Zeeshan Ali
相关产品推荐
相关产品推荐

