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

PyTorch中剪枝模型的deepcopy报错问题求解

解决PyTorch剪枝模型无法deepcopy的问题

问题根源

PyTorch的剪枝工具(torch.nn.utils.prune)会对目标层的参数做包装:它把原权重替换成一个带掩码的计算张量(非用户显式创建的叶子节点),这类张量带有grad_fn但未标记为requires_grad=True,而Python的deepcopy只支持拷贝用户显式创建的叶子张量,因此触发报错。

可行解决方案

方案1:临时移除剪枝再恢复(简单直接)

思路是在拷贝前先移除剪枝逻辑,完成拷贝后再重新应用剪枝掩码,同时保留原有的梯度控制:

import torch
from torch import nn
from copy import deepcopy
import torch.nn.utils.prune as prune

device = torch.device("cpu")

# 初始化模型并剪枝
model = nn.Sequential(
    nn.Linear(1,5),
    nn.Linear(5,1)
)
mask = torch.tensor([1,0,0,1,0]).reshape(-1,1)
prune.custom_from_mask(model[0], name='weight', mask=mask)

# 步骤1:保存剪枝的掩码和原权重
pruned_layer = model[0]
original_weight = pruned_layer.weight_orig.clone()
prune_mask = pruned_layer.weight_mask.clone()

# 步骤2:移除剪枝,恢复原权重
prune.remove(pruned_layer, 'weight')

# 步骤3:deepcopy模型
new_model = deepcopy(model)

# 步骤4:对拷贝后的模型重新应用剪枝,并冻结被掩码的权重
new_pruned_layer = new_model[0]
prune.custom_from_mask(new_pruned_layer, name='weight', mask=prune_mask)
# 冻结被剪枝的权重(可选,根据需求)
new_pruned_layer.weight.requires_grad = True
new_pruned_layer.weight_orig.requires_grad = True
# 手动设置被掩码部分的权重不参与梯度更新
with torch.no_grad():
    new_pruned_layer.weight_orig[prune_mask == 0] = 0.0
new_pruned_layer.weight_orig.grad_mask = prune_mask.bool()

方案2:自定义模块的深拷贝逻辑(更优雅)

如果需要频繁拷贝模型,可以给剪枝后的模块添加自定义的__deepcopy__方法,自动处理剪枝参数的拷贝:

def pruned_module_deepcopy(self, memo):
    # 创建模块的浅拷贝
    cls = self.__class__
    new_module = cls.__new__(cls)
    memo[id(self)] = new_module
    # 拷贝模块的属性
    for k, v in self.__dict__.items():
        if k.endswith('_orig') or k.endswith('_mask'):
            # 直接拷贝原权重和掩码张量
            setattr(new_module, k, deepcopy(v, memo))
        else:
            setattr(new_module, k, deepcopy(v, memo))
    # 重新应用剪枝
    for name in prune.get_pruned_parameters(self):
        prune.custom_from_mask(new_module, name=name.split('.')[-1], mask=getattr(self, f"{name.split('.')[-1]}_mask"))
    return new_module

# 给被剪枝的Linear层绑定自定义深拷贝方法
model[0].__deepcopy__ = pruned_module_deepcopy.__get__(model[0], nn.Linear)

# 现在可以直接deepcopy了
new_model = deepcopy(model)

验证效果

运行上述代码后,new_model会和原模型保持相同的剪枝状态,且能正常参与训练,不会触发deepcopy的报错。同时,被掩码的权重不会被更新(如果设置了冻结逻辑)。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.28 00:52:47