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

使用functional_call时nn.LSTM无法正常计算梯度的问题

问题:使用torch.nn.utils.stateless.functional_call时LSTM层梯度为None,直接调用模型则正常

问题描述

当调用包含LSTM的模型时,使用functional_call仅能得到线性层的梯度,LSTM层梯度为None;但直接调用模型实例时,所有层(包括LSTM和线性层)都能正常计算梯度。

复现代码

import torch
from torch.nn.utils.stateless import functional_call
import torch.autograd as autograd
import torch.nn as nn

# 定义模型
class Encoder(nn.Module):
    def __init__(self, action_dim, z_dim, skill_length):
        super().__init__()
        print(action_dim)
        self.lin1 = nn.Linear(action_dim, action_dim)
        self.lstm = nn.LSTM(input_size=action_dim, hidden_size=z_dim, batch_first=True)
        self.lin2 = nn.Linear(z_dim, z_dim)

    def forward(self, skill):
        a, b, c = skill.shape
        skill = skill.reshape(-1, skill.shape[-1])
        embed = self.lin1(skill)
        embed = embed.reshape(a, b, c)
        mean, _ = self.lstm(embed)
        mean = mean[:, -1, :]
        mean = self.lin2(mean)

        return mean

# 参数初始化函数
def pars(model):
    params = {}
    for name, param in model.named_parameters():
        if len(param.shape) == 1:
            init = torch.nn.init.constant_(param, 0)
        else:
            init = torch.nn.init.orthogonal_(param)
        params[name] = nn.Parameter(init)
    return params

# 初始化模型和输入
model = Encoder(4, 2, 5)
x = torch.rand(3, 5, 4)
params = pars(model)

# 使用functional_call调用并计算梯度
samp = functional_call(model, params, x)
grad_f = autograd.grad(torch.mean(samp), params.values(), retain_graph=True, allow_unused=True)
print(grad_f)
# 输出中线性层有梯度,LSTM层梯度为None

# 直接调用模型并计算梯度
samp = model(x)
grad = autograd.grad(torch.mean(samp), model.parameters(), retain_graph=True)
print(grad)
# 输出中所有层都有梯度

2022年12月11日更新补充
在模型forward中加入梯度验证:

class Encoder(nn.Module):
    def __init__(self, action_dim, z_dim, skill_length):
        super().__init__()
        print(action_dim)
        self.lin1 = nn.Linear(action_dim, action_dim)
        self.lstm = nn.LSTM(input_size=action_dim, hidden_size=z_dim, batch_first=True)
        self.lin2 = nn.Linear(z_dim, z_dim)

    def forward(self, skill):
        a, b, c = skill.shape
        skill = skill.reshape(-1, skill.shape[-1])
        embed = self.lin1(skill)
        embed = embed.reshape(a, b, c)
        mean, _ = self.lstm(embed)
        import pdb; pdb.set_trace()
        grad1 = autograd.grad(mean.mean(), params.values(), retain_graph=True, allow_unused=True)
        # grad1仅能得到lin1层的梯度,LSTM层为None
        grad2 = autograd.grad(mean.mean(), self.parameters(), retain_graph=True, allow_unused=True)
        # grad2仅能得到LSTM层的梯度,lin1层为None
        mean = mean[:, -1, :]
        mean = self.lin2(mean)

        return mean

补充说明:直接调用模型时,autograd.grad可以同时获取lin1和LSTM层的梯度。

原因分析

问题出在pars函数的参数创建方式上。原代码中使用torch.nn.init.constant_(param, 0)和torch.nn.init.orthogonal_(param)对模型自身的参数进行原地修改,随后将修改后的原参数包装成新的nn.Parameter存入params字典。这导致params中的参数与模型实例自身的参数共享底层张量数据,functional_call无法正确替换LSTM层的参数进行计算——实际计算时LSTM仍在使用模型自身的参数,而非params中的参数,因此对params.values()求梯度时,LSTM层的梯度为None。

解决方案

修改pars函数,创建完全独立于模型自身参数的新参数张量,避免原地修改原模型参数:

def pars(model):
    params = {}
    for name, param in model.named_parameters():
        if len(param.shape) == 1:
            # 创建新的全零参数,不修改原模型参数
            new_param = nn.Parameter(torch.zeros_like(param))
        else:
            # 创建新的正交初始化参数,不修改原模型参数
            new_param = nn.Parameter(torch.nn.init.orthogonal_(torch.empty_like(param)))
        params[name] = new_param
    return params

验证结果

使用修改后的pars函数重新运行代码,functional_call调用后计算梯度,LSTM层和线性层都会产生正常的梯度值。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.17 08:20:26