PyTorch线性回归构造函数中为何要完整继承nn.Module初始化?
为什么自定义PyTorch模型必须调用
nn.Module父类初始化方法 你参考的自定义线性回归模型代码如下:
class linearRegression(torch.nn.Module): def __init__(self, inputSize, outputSize): super(linearRegression, self).__init__() self.linear = torch.nn.Linear(inputSize, outputSize) def forward(self, x): out = self.linear(x) return out
其中super(linearRegression, self).__init__()这行不是冗余代码,哪怕你后续定义的层都是直接调用官方自带的nn.Linear这类现成模块,这步也不能省。父类初始化的核心作用是给你的自定义类搭好整个Module体系的运行基础,具体承担这几部分关键工作:
- 初始化模块核心内部存储容器
父类__init__会提前创建_parameters、_modules、_buffers三个核心内部字典。你后续写self.linear = torch.nn.Linear(...)的时候,不是简单给实例挂了个普通属性——PyTorch重写了Module类的__setattr__逻辑,只要你赋值的对象是nn.Module子类、Parameter或者Buffer类型,就会自动把对应的子模块、可训练参数、常驻缓存(比如BN层的滑动均值)存到前面的三个字典里。如果没跑父类初始化,这几个容器根本不存在,你赋值子层的时候要么直接抛属性错误,要么后续所有参数相关的功能全部失效。 - 初始化训练/推理全流程基础属性
父类初始化会默认给实例打上training=True的状态标记,同时初始化好钩子注册、设备迁移、数值精度切换相关的内部配置。你后续调用.cuda()/.to()把模型挪到GPU/其他设备、调用.eval()切换推理模式、注册前向/反向传播钩子、设置权重是否需要梯度这些常用功能,全依赖这些提前初始化好的属性,跳过父类初始化调用这些方法会直接报属性不存在的错误。 - 打通模块树的递归遍历逻辑
PyTorch很多核心功能,比如.parameters()递归扫描所有层的可训练权重、state_dict()导出全模型权重、load_state_dict()加载训练好的权重,都是靠遍历每个Module下_modules存储的子模块递归实现的。如果跳过父类初始化,你定义的self.linear不会被注册到子模块列表里,这些遍历逻辑根本找不到你写的线性层,相当于这个层成了游离在PyTorch模型管理体系外的普通Python属性,完全不会被识别为模型的一部分。
最直观的反例:你可以把这行super调用注释掉再跑代码,实例化模型之后调用
model.parameters()会直接返回空迭代器,优化器拿不到线性层的权重,训练的时候参数完全不会更新,而且代码不会抛显性的初始化错误,排查问题要花很多时间。
你可能会疑惑:self.linear本身就是现成的nn.Linear实例,自带所有相关方法,为什么还要外层做初始化?这里要明确:nn.Linear作为独立的Module子类,它自己的初始化流程只会完成它自己这个子模块的内部状态搭建,和你当前写的外层linearRegression类的初始化是完全独立的两件事。外层模型作为一个独立的Module实例,必须自己完成父类初始化流程,才能把自己和内部嵌套的所有子模块整合成一个能被PyTorch正常识别、调度的完整模型。
内容的提问来源于stack exchange,提问作者prarthik krishnan
相关产品推荐
相关产品推荐

