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

在forward()中修改PyTorch预训练模型参数导致训练持续变慢

问题背景与现象

我需要用固定张量A和可训练参数B替换预训练模型的全部权重,即模型权重W = A + B,其中A是不可训练的固定张量,B是可训练参数,目标是在预训练模型结构下仅训练B。

我实现的代码如下:

Net类定义

class Net(nn.Module):
    
    def __init__(self, pre_model, B):
    
        super(Net, self).__init__()
        self.B = B
        self.pre_model = copy.deepcopy(pre_model)
        for params in self.pre_model.parameters(): 
            params.requires_grad = False
    
    def forward(self, x, A):
        for i, params in enumerate(self.pre_model.parameters()):
            params.copy_(A[i].detach().clone()) # 不需要训练A所以detach
            params.add_(self.B[i])      # 需要训练B所以不detach
            params.retain_grad()

        x = self.pre_model(x)
        return x

A和B的初始化代码

b = []
A = []
for params in list(pre_model.parameters()):
    A.append(torch.rand_like(params))
    b_temp = nn.Parameter(torch.rand_like(params))
    b.append(b_temp.detach().clone())
B = nn.ParameterList(b)

训练过程中B确实在被训练,但每一轮迭代的训练速度持续变慢,训练日志如下:

Epoch 1:
24%|██▍ | 47/196 [00:05<00:23, 6.44it/s]
57%|█████▋ | 111/196 [00:18<00:19, 4.28it/s]
96%|█████████▋| 189/196 [00:41<00:02, 2.90it/s]
Epoch 2:
6%|▌ | 11/196 [00:04<01:14, 2.50it/s]

我认为已经正确分离所有不需要训练的参数,但不确定速度变慢的原因。可通过以下主循环代码复现该问题:

from torch.cuda import synchronize
device = 'cuda'
pre_model = models.resnet18().to(device)
b = []
A = []
for params in list(pre_model.parameters()):
    A.append(torch.rand_like(params))
    b_temp = nn.Parameter(torch.rand_like(params))
    b.append(b_temp.detach().clone())
B = nn.ParameterList(b)

modelwithAB = Net(pre_model, B)
optimizer = torch.optim.Adam(modelwithAB.parameters(), lr=1e-3)

image = torch.randn(2, 3, 224, 224).to(device)
print(torch.cuda.memory_allocated()/1024**2)

for i in tqdm(range(300)):
    optimizer.zero_grad()
    out = modelwithAB(image, A)
    start = time.time()
    out.mean().backward()
    torch.cuda.synchronize()
    optimizer.step()
    if i%40==0:
        print("-", torch.cuda.memory_allocated()/1024**2, "-", time.time()-start)

问题原因分析

核心问题:params.retain_grad()的累积效应

每次前向传播时,你都对预训练模型的参数调用了params.retain_grad(),但没有在反向传播后清除这些梯度。PyTorch中,调用retain_grad()会让张量的梯度在反向传播后被保留,而非自动释放。随着迭代次数增加,越来越多的梯度张量被存储在内存中,导致显存占用持续上升,进而拖慢训练速度(显存不足时会触发CPU-GPU数据交换,进一步降低速度)。

其他潜在问题

  • 前向传播中修改参数的低效性:每次前向都循环修改预训练模型的参数,这种操作会破坏PyTorch的计算图优化,且每次复制和相加操作的额外开销累积后,也会导致速度下降。
  • B的初始化错误:初始化B时使用b.append(b_temp.detach().clone()),会把nn.Parameter转换成普通张量,后续用nn.ParameterList包装时存在不必要的张量操作,增加了额外开销。

修正方案

方案1:移除params.retain_grad()并调整权重计算方式

不需要在预训练模型参数上保留梯度,因为我们只需要训练B。正确做法是直接将权重设置为A + B,而非修改预训练模型参数本身:

修改Net类:

class Net(nn.Module):
    def __init__(self, pre_model, B, A):
        super(Net, self).__init__()
        self.B = B
        self.A = A
        self.pre_model = pre_model
        # 冻结预训练模型的所有参数
        for params in self.pre_model.parameters():
            params.requires_grad = False
    
    def forward(self, x):
        # 遍历参数,将权重设置为A+B
        for idx, param in enumerate(self.pre_model.parameters()):
            param.data = self.A[idx].data + self.B[idx].data
        return self.pre_model(x)

修正B的初始化:

b = []
A = []
for params in list(pre_model.parameters()):
    A.append(torch.rand_like(params))
    # 直接存储nn.Parameter,无需detach
    b.append(nn.Parameter(torch.rand_like(params)))
B = nn.ParameterList(b)

方案2:更高效的参数替换方式(推荐)

自定义模型参数,直接在计算时使用A + B,避免修改预训练模型参数:

class Net(nn.Module):
    def __init__(self, pre_model, A):
        super(Net, self).__init__()
        self.A = A
        # 初始化可训练参数B,与A形状一致
        self.B = nn.ParameterList([nn.Parameter(torch.rand_like(a)) for a in A])
        self.pre_model = pre_model
        # 冻结预训练模型参数
        for param in self.pre_model.parameters():
            param.requires_grad = False
    
    def forward(self, x):
        # 将每个层的权重设置为A+B
        for a, b, param in zip(self.A, self.B, self.pre_model.parameters()):
            param.data = a.data + b.data
        return self.pre_model(x)

这样修改后,既避免了retain_grad()导致的显存泄漏,又简化了初始化逻辑,训练速度会恢复正常。


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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.06 04:01:05