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

如何在PyTorch中不使用torch.block_diag创建块对角权重张量?

问题原因与解决方案

错误原因

手动构建块对角矩阵的方式存在两个核心问题:

  1. 原地赋值破坏梯度追踪:torch.zeros(n,n)默认创建requires_grad=False的张量,后续切片赋值属于原地操作,PyTorch对这类操作的梯度支持不完善,会导致反向传播时计算图断裂。
  2. 静态构建无法同步参数更新:在__init__中预先构建self.M,会让它成为初始化时的静态张量,后续M1-M4作为可训练参数被优化器更新时,self.M不会自动同步更新,同时这种静态张量的计算图在反向传播时会出现重复访问已释放张量的错误。

而torch.block_diag的版本能暂时运行,是因为它返回的张量直接关联了M1-M4的计算图,但同样存在静态构建无法同步参数更新的隐患——初始化后self.M不会随M1-M4的更新而变化,相当于固定了权重,无法真正训练。

正确实现方式

应该在forward方法中动态构建块对角矩阵,确保每次前向传播都使用最新的参数值,同时保证梯度追踪正常:

class BlockLinear(nn.Module):
    def __init__(self, n):
        super().__init__()
        self.n = n
        self.m = int(np.sqrt(self.n))

        # 定义可训练参数
        self.M1 = nn.Parameter(torch.randn(self.m, self.m))
        self.M2 = nn.Parameter(torch.randn(self.m, self.m))
        self.M3 = nn.Parameter(torch.randn(self.m, self.m))
        self.M4 = nn.Parameter(torch.randn(self.m, self.m))
    
    def forward(self, x):
        # 动态构建块对角矩阵
        M = torch.zeros(self.n, self.n, device=x.device, dtype=x.dtype)
        M[:self.m, :self.m] = self.M1
        M[self.m:2*self.m, self.m:2*self.m] = self.M2
        M[2*self.m:3*self.m, 2*self.m:3*self.m] = self.M3
        M[3*self.m:, 3*self.m:] = self.M4 
        # 执行线性变换
        x = torch.einsum('ij, bi -> bj', M, x)
        return x

更通用的高维度扩展方案

如果要推广到任意数量的块(比如k个m×m块,n=k×m),可以用循环简化构建逻辑,避免硬编码切片:

class BlockLinear(nn.Module):
    def __init__(self, n):
        super().__init__()
        self.n = n
        self.m = int(np.sqrt(self.n))
        self.num_blocks = self.n // self.m  # 计算块的数量
        
        # 用ParameterList管理所有块参数,方便扩展
        self.blocks = nn.ParameterList([
            nn.Parameter(torch.randn(self.m, self.m)) 
            for _ in range(self.num_blocks)
        ])
    
    def forward(self, x):
        M = torch.zeros(self.n, self.n, device=x.device, dtype=x.dtype)
        for i, block in enumerate(self.blocks):
            start = i * self.m
            end = (i+1) * self.m
            M[start:end, start:end] = block
        x = torch.einsum('ij, bi -> bj', M, x)
        return x

额外说明

  • 动态构建时要指定device和dtype与输入x一致,避免设备/类型不匹配的错误。
  • 使用nn.ParameterList管理多个块参数,既符合PyTorch的参数管理规范,也方便后续扩展到任意数量的块。

内容的提问来源于stack exchange,提问作者user572780

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.28 09:43:28