You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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函数里也解决不了问题。


问题根源与解决办法

核心原因

  1. 邻接矩阵未被注册为可训练参数:你在__init__里生成的self.adjacency_matrix是拼接subdiagonal_block(这是一个nn.Parameter)和零张量得到的普通张量,虽然加了requires_grad=True,但它并没有被PyTorch模型注册为可训练参数。优化器只会更新nn.Parameter类型的参数,而且拼接操作生成的新张量和原始subdiagonal_block的梯度传播链被切断了——初始化后邻接矩阵就固定了,不再和可训练的子块关联。
  2. 输入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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.29 13:43:19