如何修改含原地操作的PyTorch自定义LSTM模块以支持微分?
解决原地操作导致的梯度计算错误问题
这个错误的根源确实是你代码里的原地赋值操作(比如o[it, t] = ...、hx[it, t+1] = ...)——PyTorch的自动微分依赖追踪张量的完整创建路径,原地修改会直接切断这个路径,导致梯度计算时找不到所需的中间变量。
要在保留“重复10次取均值”逻辑的前提下解决这个问题,核心思路是避免提前初始化张量然后原地修改,改用列表收集每一步的计算结果,最后再拼接成张量。这样所有操作都是非原地的,能完整保留计算图。
下面是修改后的完整代码,我会标注关键改动点:
import torch from torch import nn class StochasticLSTM(nn.Module): def __init__(self, input_size: int, hidden_size: int, dropout_rate: float): """ Args: - dropout_rate: should be between 0 and 1 """ super(StochasticLSTM, self).__init__() self.iter = 10 self.input_size = input_size self.hidden_size = hidden_size if not 0 <= dropout_rate <= 1: raise Exception("Dropout rate should be between 0 and 1") self.dropout = dropout_rate self.bernoulli_x = torch.distributions.Bernoulli( torch.full((self.input_size,), 1 - self.dropout) ) self.bernoulli_h = torch.distributions.Bernoulli( torch.full((hidden_size,), 1 - self.dropout) ) self.Wi = nn.Linear(self.input_size, self.hidden_size) self.Ui = nn.Linear(self.hidden_size, self.hidden_size) self.Wf = nn.Linear(self.input_size, self.hidden_size) self.Uf = nn.Linear(self.hidden_size, self.hidden_size) self.Wo = nn.Linear(self.input_size, self.hidden_size) self.Uo = nn.Linear(self.hidden_size, self.hidden_size) self.Wg = nn.Linear(self.input_size, self.hidden_size) self.Ug = nn.Linear(self.hidden_size, self.hidden_size) def forward(self, input, hx=None): """ input shape (sequence, batch, input dimension) output shape (sequence, batch, output dimension) return output, (hidden_state, cell_state) """ T, B, _ = input.shape # 用列表收集每一次迭代的结果,替代原地修改的张量 iter_outputs = [] iter_hiddens = [] iter_cells = [] for it in range(self.iter): # 每个迭代维护独立的h和c状态,而不是共享大张量 if hx is None: current_h = torch.zeros((B, self.hidden_size), dtype=input.dtype, device=input.device) current_c = torch.zeros((B, self.hidden_size), dtype=input.dtype, device=input.device) else: current_h, current_c = hx # 注意:这里不需要repeat,每个迭代用相同的初始h/c即可 # 收集当前迭代每个时间步的输出、h、c step_outputs = [] step_hiddens = [] step_cells = [] # 采样当前迭代的dropout掩码,确保和输入在同一设备 zx = self.bernoulli_x.sample().to(input.device) zh = self.bernoulli_h.sample().to(input.device) for t in range(T): x = input[t] * zx h = current_h * zh i = torch.sigmoid(self.Ui(h) + self.Wi(x)) f = torch.sigmoid(self.Uf(h) + self.Wf(x)) o_t = torch.sigmoid(self.Uo(h) + self.Wo(x)) g = torch.tanh(self.Ug(h) + self.Wg(x)) current_c = f * current_c + i * g current_h = o_t * torch.tanh(current_c) # 把当前时间步的结果加入列表 step_outputs.append(o_t) step_hiddens.append(current_h) step_cells.append(current_c) # 把当前迭代的所有时间步结果堆叠成张量,加入迭代列表 iter_outputs.append(torch.stack(step_outputs, dim=0)) # shape (T, B, hidden_size) iter_hiddens.append(torch.stack(step_hiddens, dim=0)) # shape (T, B, hidden_size) iter_cells.append(torch.stack(step_cells, dim=0)) # shape (T, B, hidden_size) # 把所有迭代的结果堆叠,然后计算均值 o = torch.mean(torch.stack(iter_outputs, dim=0), dim=0) # shape (T, B, hidden_size) hx_mean = torch.mean(torch.stack(iter_hiddens, dim=0), dim=0) # shape (T, B, hidden_size) c_mean = torch.mean(torch.stack(iter_cells, dim=0), dim=0) # shape (T, B, hidden_size) return o, (hx_mean, c_mean)
关键改动说明:
- 用列表替代原地张量:不再提前初始化
o、hx、c这些大张量,而是用iter_outputs、step_outputs等列表收集每一步结果,最后再堆叠成张量,完全避免原地操作。 - 每个迭代独立维护状态:原来的代码试图用一个大张量保存所有迭代的h/c状态并原地修改,现在改成每个迭代单独维护
current_h和current_c,逻辑更清晰,也不会出现跨迭代的原地修改问题。 - 设备对齐:添加了
.to(input.device)确保dropout掩码和输入在同一设备上(比如GPU),避免潜在的设备不匹配问题。 - 初始状态处理:原来的
hx重复10次的操作其实没必要,每个迭代用相同的初始h/c即可,这样更符合“重复传递10次”的逻辑。
这样修改后,所有张量的创建都是通过非原地操作完成的,PyTorch可以正常追踪计算图,梯度计算就不会再报错了,同时也保留了“重复10次取均值”的核心逻辑。
内容的提问来源于stack exchange,提问作者Khanetor
相关产品推荐
相关产品推荐

