PyTorch如何自动识别模型参数?自定义MLP实例答疑
PyTorch自动识别模型参数的原理
自定义MLP分类器代码
class MLPClassifier(nn.Module): """ A basic multi-layer perceptron classifier with 3 layers. """ def __init__(self, input_size, hidden_size, num_classes): """ The constructor for the MLPClassifier class. """ super(MLPClassifier, self).__init__() self.fc1 = nn.Linear(input_size, hidden_size) # weights & biases for the input-to-hidden layer self.ac1 = nn.ReLU() # non-linear activation for the input-to-hidden layer self.fc2 = nn.Linear(hidden_size, num_classes) # weights & biases for the hidden-to-output layer self.ac2 = nn.Softmax(dim=1) # non-linear activation for the hidden-to-output layer
运行脚本及输出
hyper_param_input_size = 4 hyper_param_hidden_size = 64 hyper_param_num_classes = 3 model = MLPClassifier(hyper_param_input_size, hyper_param_hidden_size, hyper_param_num_classes) for p in model.parameters(): print(p.shape)
输出结果:
torch.Size([64, 4]) torch.Size([64]) torch.Size([3, 64]) torch.Size([3])
自动识别参数的原理
你猜的方向基本对,但细节更明确:
- 当你的类继承
nn.Module,并且在__init__里调用了super().__init__()后,nn.Module的父类逻辑会自动维护一个可学习参数的注册表。 - 当你把
nn.Linear这种nn.Module子类实例(比如self.fc1)赋值给模型属性时,nn.Module的属性赋值逻辑会自动检测:如果该属性是nn.Module子类对象,就会把它加入模型的子模块列表;接着递归遍历所有子模块,把里面的nn.Parameter实例(比如Linear层的weight和bias)全部收集到全局参数列表中。 - 像ReLU、Softmax这类激活层属于无参数Module,它们内部没有定义可学习的
nn.Parameter,所以不会出现在model.parameters()的输出里。
内容的提问来源于stack exchange,提问作者AlanSTACK
相关产品推荐
相关产品推荐

