使用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
相关产品推荐
相关产品推荐

