PyTorch中实现带parameters()属性的参数分组的方法及最佳实践
你现在的代码里self.base()和self.classifier()是普通的实例方法,不是PyTorch的子模块对象,所以自然没法调用model.base.parameters()——PyTorch只会把直接注册为模型属性的nn.Module子类实例(比如nn.Linear、nn.Sequential)的参数纳入管理,普通方法不会被识别为参数组的载体。
怎么修改才能实现model.base.parameters()?
根本思路是把base和classifier对应的层封装成PyTorch子模块,这样它们就自带parameters()方法,同时PyTorch会自动帮你管理这些参数的注册、设备同步等操作。下面是两种常用的实现方式:
方式1:用nn.Sequential快速封装子模块
这是最简洁的方式,适合层结构线性串联的场景:
import torch.nn as nn class MyModel(nn.Module): def __init__(self): super(MyModel, self).__init__() # 把base的层封装成Sequential子模块 self.base = nn.Sequential( nn.Linear(1, 512), nn.ReLU(), nn.Linear(512, 264), nn.ReLU() ) # 把classifier的层封装成Sequential子模块 self.classifier = nn.Sequential( nn.Linear(264, 128), nn.ReLU(), nn.Linear(128, 964) ) def forward(self, y0): # 直接调用子模块的forward方法,逻辑更简洁 y2 = self.base(y0) return self.classifier(y2)
这样修改后,self.base和self.classifier都是nn.Sequential(属于nn.Module子类)的实例,你完全可以用model.base.parameters()来获取base部分的所有参数,然后在优化器里分组设置:
import torch.optim as optim optimizer = optim.SGD( [ {'params': model.base.parameters()}, # 使用默认的lr=1e-2 {'params': model.classifier.parameters(), 'lr': 1e-3} ], lr=1e-2, momentum=0.9 )
方式2:自定义子模块类(适合复杂结构)
如果你的base或classifier有更复杂的分支逻辑,用nn.Sequential没法满足,可以自定义独立的子模块类:
import torch.nn as nn # 定义base子模块 class BaseModule(nn.Module): def __init__(self): super().__init__() self.fc1 = nn.Linear(1, 512) self.fc2 = nn.Linear(512, 264) self.relu = nn.ReLU() def forward(self, x): x = self.relu(self.fc1(x)) return self.relu(self.fc2(x)) # 定义classifier子模块 class ClassifierModule(nn.Module): def __init__(self): super().__init__() self.fc3 = nn.Linear(264, 128) self.fc4 = nn.Linear(128, 964) self.relu = nn.ReLU() def forward(self, x): x = self.relu(self.fc3(x)) return self.fc4(x) # 主模型 class MyModel(nn.Module): def __init__(self): super(MyModel, self).__init__() self.base = BaseModule() self.classifier = ClassifierModule() def forward(self, y0): y2 = self.base(y0) return self.classifier(y2)
这种方式的好处是把不同模块的逻辑完全解耦,代码可读性更强,同样支持model.base.parameters()的调用。
关于nn.ParameterList的疑问
你提到的nn.ParameterList确实可以手动收集参数,但这绝对不是最佳实践——手动维护参数列表不仅繁琐,还容易遗漏参数(比如偏置项),而且PyTorch的模块系统提供的自动参数管理(比如.cuda()、.eval()时的行为)也无法生效。只有在极少数特殊场景下(比如动态生成参数)才需要用到nn.ParameterList,你现在的场景完全没必要。
总结最佳实践
- 把需要分组优化的层封装成子模块(
nn.Sequential或自定义nn.Module子类) - 在主模型中把这些子模块作为实例属性注册
- 直接通过
子模块.parameters()获取对应组的参数,传入优化器即可
内容的提问来源于stack exchange,提问作者Blade

