如何在PyTorch中仅训练线性层权重矩阵的对角权重?
实现仅对角线元素可训练的缩放层
方法1:自定义缩放层(最直接高效)
直接跳过标准nn.Linear,自定义一个仅维护对角线缩放参数的层,本质就是对输入每个特征维度做独立缩放,完全匹配你要的"仅前序输出缩放"逻辑,训练时只有对角线参数参与更新,没有冗余计算。
import torch import torch.nn as nn class DiagonalScalingLayer(nn.Module): def __init__(self, in_features): super().__init__() # 初始化对角线缩放参数,默认设为1,也可改用其他初始化策略(比如xavier) self.diag_weights = nn.Parameter(torch.ones(in_features)) def forward(self, x): # 逐特征维度相乘,等价于对角线矩阵乘法 return x * self.diag_weights
这个层的参数只有一维张量diag_weights,训练时只会更新这些参数,完全满足高效调优的需求。
方法2:改造现有Linear层(用钩子约束参数)
如果必须基于nn.Linear层改造,可以通过注册钩子,强制约束非对角线元素始终为0,同时冻结其梯度:
import torch import torch.nn as nn # 初始化Linear层(输入输出维度要一致,否则无法形成对角线矩阵) linear_layer = nn.Linear(5, 5) # 先一次性把非对角线元素置0,初始化对角线为1 with torch.no_grad(): linear_layer.weight.fill_(0) torch.diagonal(linear_layer.weight).fill_(1) # 注册前向钩子:每次前向传播前,强制把非对角线元素重置为0 def enforce_diagonal(module, input): with torch.no_grad(): # 先保存当前对角线元素,再清空权重矩阵,最后恢复对角线 diag_vals = torch.diagonal(module.weight) module.weight.fill_(0) torch.diagonal(module.weight).copy_(diag_vals) linear_layer.register_forward_pre_hook(enforce_diagonal) # 注册反向钩子:把非对角线元素的梯度置0,确保优化器只更新对角线参数 def zero_off_diag_grad(module, grad_input, grad_output): grad_weight = grad_input[0] if isinstance(grad_input, tuple) else grad_input with torch.no_grad(): # 只保留对角线的梯度,其余置0 grad_weight.copy_(torch.diag(torch.diagonal(grad_weight))) return (grad_weight,) + grad_input[1:] if isinstance(grad_input, tuple) else grad_weight linear_layer.register_backward_hook(zero_off_diag_grad)
- 前向钩子保证每次前向计算都是纯缩放操作;
- 反向钩子确保非对角线参数不会被更新;
- 所有参数修改都用
torch.no_grad()包裹,避免干扰计算图。
关于requires_grad的说明
requires_grad是张量级别的属性,无法单独设置张量的部分元素是否需要梯度,所以要么只维护需要训练的参数(方法1),要么在计算流程中约束不需要训练的部分(方法2),这两种方式都能实现你要的训练时动态保持对角线矩阵的需求。
内容的提问来源于stack exchange,提问作者SlimeyGuy123
相关产品推荐
相关产品推荐

