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

如何修改含原地操作的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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.06 21:17:29