PyTorch自定义网络权重矩阵无法更新问题求助
问题:自定义邻接矩阵线性网络无法更新参数
我为MNIST数据集写了一个简单的线性分类网络,用全局邻接矩阵做计算,矩阵只有左下角块非零,其余都是0。模型结构很基础,只有两层,没有非线性激活,但训练时邻接矩阵完全不更新,模型根本学不了。我已经验证过训练循环在标准架构上能正常跑(用SGD优化器和交叉熵损失),所以问题肯定出在网络类的定义上。对我来说,必须用这个全局邻接矩阵做计算,想知道问题根源和解决办法。
代码如下:
class Simple_Direct_Network_Adjacency_Matrix_Implementation_Dim2(nn.Module): def __init__(self, input_dim , middle_dim, output_dim): super().__init__() self.input_dim = input_dim _ = middle_dim #This is an hack: we want dim 2 now, so this input to the class gets ignored self.output_dim = output_dim self.total_dim = self.input_dim + self.output_dim self.subdiagonal_block = nn.Parameter(torch.empty(self.output_dim, self.input_dim)) nn.init.normal_(self.subdiagonal_block , mean=0 , std=0.1) self.adjacency_matrix = self.make_subdiagonal_matrix().requires_grad_(requires_grad=True) def make_subdiagonal_matrix(self): over_block = torch.zeros(self.input_dim, self.input_dim) side_block = torch.zeros(self.total_dim, self.output_dim) matrix = torch.cat((over_block , self.subdiagonal_block), 0) matrix = torch.cat((matrix, side_block), 1) return matrix def forward(self, batch_of_inputs): # Flatten the batch of input images flat_inputs = batch_of_inputs.view(-1 , batch_of_inputs.size(0)) # Append zeros to match flat_inputs_total = torch.cat((flat_inputs, torch.zeros(self.output_dim , flat_inputs.size(1))), dim=0) # Perform matrix multiplication y_total_final = torch.mm(self.adjacency_matrix , flat_inputs_total) # Extract logits logits = y_total_final[-self.output_dim: , :].t() return logits
注:我试过省略requires_grad,没用;用nn.Parameter定义的参数矩阵也没更新;把邻接矩阵的构建移到forward函数里也解决不了问题。
问题根源与解决办法
核心原因
- 邻接矩阵未被注册为可训练参数:你在
__init__里生成的self.adjacency_matrix是拼接subdiagonal_block(这是一个nn.Parameter)和零张量得到的普通张量,虽然加了requires_grad=True,但它并没有被PyTorch模型注册为可训练参数。优化器只会更新nn.Parameter类型的参数,而且拼接操作生成的新张量和原始subdiagonal_block的梯度传播链被切断了——初始化后邻接矩阵就固定了,不再和可训练的子块关联。 - 输入flatten维度错误:
batch_of_inputs.view(-1, batch_of_inputs.size(0))把样本数和特征数的顺序搞反了,会导致矩阵乘法维度不匹配,间接影响梯度计算。
修复方案
class Simple_Direct_Network_Adjacency_Matrix_Implementation_Dim2(nn.Module): def __init__(self, input_dim , middle_dim, output_dim): super().__init__() self.input_dim = input_dim _ = middle_dim # 忽略该参数,当前用两层结构 self.output_dim = output_dim self.total_dim = self.input_dim + self.output_dim # 仅将可训练的子块注册为模型参数 self.subdiagonal_block = nn.Parameter(torch.empty(self.output_dim, self.input_dim)) nn.init.normal_(self.subdiagonal_block, mean=0, std=0.1) def make_subdiagonal_matrix(self): # 让零张量和参数同设备,避免GPU/CPU不匹配问题 over_block = torch.zeros(self.input_dim, self.input_dim, device=self.subdiagonal_block.device) side_block = torch.zeros(self.total_dim, self.output_dim, device=self.subdiagonal_block.device) matrix = torch.cat((over_block, self.subdiagonal_block), 0) matrix = torch.cat((matrix, side_block), 1) return matrix def forward(self, batch_of_inputs): # 正确flatten:先转为(样本数, 特征数),再转置适配矩阵乘法维度 flat_inputs = batch_of_inputs.view(batch_of_inputs.size(0), -1).T # 拼接零向量,形状变为(total_dim, 样本数) flat_inputs_total = torch.cat( (flat_inputs, torch.zeros(self.output_dim, flat_inputs.size(1), device=flat_inputs.device)), dim=0 ) # 动态构建邻接矩阵,确保梯度能传递到subdiagonal_block adjacency_matrix = self.make_subdiagonal_matrix() y_total_final = torch.mm(adjacency_matrix, flat_inputs_total) # 提取输出并转置回(样本数, 输出维度) logits = y_total_final[-self.output_dim:, :].T return logits
关键改动说明
- 移除
__init__里的self.adjacency_matrix,改为在forward中动态生成邻接矩阵,这样每次前向传播都会基于当前的subdiagonal_block构建,梯度可以正常传递到可训练参数。 - 修正输入flatten的维度顺序,确保矩阵乘法维度匹配。
- 给零张量添加设备参数,避免模型在GPU运行时出现设备不匹配错误。
修改后,优化器就能正常更新subdiagonal_block,邻接矩阵会随参数动态变化,模型即可正常学习。
内容的提问来源于stack exchange,提问作者Noumeno
相关产品推荐
相关产品推荐

