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

原地剪枝nn.Linear权重引发异常:原因与替代方案解析

《Temporal Neuron Variance Pruning》剪枝中的PyTorch原地修改权重异常解析

问题背景

在实现《Temporal Neuron Variance Pruning》论文的模型剪枝逻辑时,原地修改nn.Linear层的权重形状后,第二次反向传播触发形状不匹配的RuntimeError,需通过特殊手段规避,以下是问题复现、可行方案及核心疑问的解析。

执行失败的代码

import torch
import torch.nn as nn

def test1():  
  layer = nn.Linear(100, 10)
  x = 5 - torch.sum(layer(torch.ones(100)))
  x.backward()
  layer.weight.data = layer.weight.data[:, :90]
  layer.weight.grad.data = layer.weight.grad.data[:, :90]
  x = 5 - torch.sum(layer(torch.ones(90)))
  x.backward()
test1()

报错信息

---------------------------------------------------------------------------
RuntimeError                              Traceback (most recent call last)
<ipython-input-3-bb36a010bd86> in <cell line: 10>()
      8     x = 5 - torch.sum(layer(torch.ones(90)))
      9     x.backward()
---> 10 test1()
     11 # and this works as well
     12 

2 frames
/usr/local/lib/python3.10/dist-packages/torch/autograd/__init__.py in backward(tensors, grad_tensors, retain_graph, create_graph, grad_variables, inputs)
    249     # some Python versions print out the first line of a multi-line function
    250     # calls in the traceback and some print out the last line
---> 251     Variable._execution_engine.run_backward(  # Calls into the C++ engine to run the backward pass
    252         tensors,
    253         grad_tensors_,

RuntimeError: Function TBackward0 returned an invalid gradient at index 0 - got [10, 90] but expected shape compatible with [10, 100]

可正常执行的代码1

import torch
import torch.nn as nn

def test2():  
  layer = torch.nn.Linear(100, 10)
  x = 5 - torch.sum(layer(torch.ones(100)))
  x.backward()
  del x    #核心修改
  layer.weight.data = layer.weight.data[:, :90]
  layer.weight.grad.data = layer.weight.grad.data[:, :90]
  x = 5 - torch.sum(layer(torch.ones(90)))
  x.backward()
test2()

可正常执行的代码2

import torch
import torch.nn as nn

def test3():  
  layer = torch.nn.Linear(100, 10)
  x = 5 - torch.sum(layer(torch.ones(100)))
  x.backward()
  layer.weight.data = layer.weight.data[:, :90]
  layer.weight.grad.data = layer.weight.grad.data[:, :90]
  layer.weight = torch.nn.Parameter(layer.weight)   #核心修改
  x = 5 - torch.sum(layer(torch.ones(90)))
  x.backward()
test3()

核心疑问解答

  1. 什么是TBackward0?
    TBackward0是PyTorch自动生成的转置操作反向传播算子,对应nn.Linear前向传播中y = x @ weight.T + bias的转置矩阵乘法的反向计算逻辑。后缀0是该算子的版本标识,由自动微分框架根据前向操作动态生成。

  2. TBackward0的定义位置?
    它是PyTorch C++后端的ATen核心库中自动生成的反向算子,没有直接的Python层定义。相关逻辑由自动微分系统根据前向张量运算动态构建,负责转置乘法的梯度回传计算。

  3. 该RuntimeError的触发位置?
    触发于PyTorch自动微分引擎的梯度校验阶段。第二次调用backward()时,新计算图的权重形状已变为[10,90],但第一次反向传播后残留的计算图节点(如TBackward0)仍持有原权重[10,100]的形状元信息,梯度张量形状不匹配时被引擎检测抛出异常。

  4. 为何仍要求与原权重形状兼容?我已正确修改梯度张量
    第一次反向传播后,原计算图未被完全释放:变量x(第一次前向输出)仍持有对原计算图的引用,导致TBackward0等反向算子保留了旧的形状校验逻辑。即使修改了权重和梯度的.data属性,算子内部存储的形状元数据并未更新。

    • 案例2中del x销毁了原计算图的根节点,触发垃圾回收后原计算图被完全释放,新反向传播仅使用新构建的计算图;
    • 案例3中重新包装权重为nn.Parameter,相当于替换了计算图中引用的张量对象,切断了与原计算图的关联,新计算图使用更新后的形状元数据。
  5. 除上述两种可行方案外,是否有其他解决方法?
    还有以下几种替代方案:

    • 用torch.no_grad()包裹权重修改操作,同时调用torch.autograd.graph.clear()(PyTorch 1.13+支持)强制清除残留计算图;
    • 创建新的nn.Linear层替换原层:layer = nn.Linear(90, 10, bias=layer.bias is not None),再将剪枝后的权重和梯度赋值给新层;
    • 直接替换权重参数:layer.weight = nn.Parameter(layer.weight[:, :90])(本质与案例3一致,但更简洁);
    • 确保第一次反向传播后,所有引用原计算图的张量(如原x)被彻底释放,避免残留计算图节点。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.30 08:14:52