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

如何在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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.05 17:42:09