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

PyTorch实现多维Seq2Seq LSTM时原地操作报错如何解决

问题1:为什么报错的偏偏是s?

s是所有张量中唯一在计算当前步值时直接依赖自身历史步值的张量:你计算s[idx]的表达式右侧引用了s[prev(idx,k)](也就是之前迭代步写入的s的切片值)。而i/o/c这类张量的计算仅依赖历史步的h,后续迭代也不会再引用它们之前的写入值,就算是原地修改也不会影响计算图的梯度依赖链,因此不会触发版本校验错误。
而h虽然也被后续迭代引用,但h[idx]的计算发生在s[idx]赋值之后,且h的历史值不会被自身的计算逻辑引用,所以版本冲突没有首先触发在h上。


问题2:有没有办法告知PyTorch我只是在填充数组?

没有直接的开关可以让PyTorch忽略这类版本校验,因为这种校验本身是为了避免梯度计算出现逻辑错误:PyTorch的自动微分机制需要记录计算图中每个张量的版本号,一旦张量被原地修改,版本号就会更新,如果反向传播时发现依赖的张量版本和记录的不一致,就会抛出错误防止计算出错误的梯度。
你认为的“只是初始化填充”在动态计算图的逻辑里并不成立:你每一步写入张量的操作都会被追踪,且后续的计算依赖了之前写入的值,原地修改相当于覆盖了计算图中已经记录的依赖节点的值,自然会触发错误。


问题3:如何改写代码避免原地操作?

你之前用torch.cat报错大概率是因为你仍在复用同一个张量变量、或者拼接逻辑仍然依赖了被原地修改的张量,正确的改写思路如下:
放弃预先分配大张量的写法,改用字典存储每个idx对应的计算结果,避免所有原地切片赋值操作:

def forward(self, x):
    """ Note: x is of shape (d1, ..., dn, batch_size, input_size). """
    dimensions = x.shape[:-2]
    batch_size = x.shape[-2]
    # 用字典存储每个idx的结果,不预先分配大张量
    f_dict = {}
    i_dict = {}
    o_dict = {}
    c_dict = {}
    s_dict = {}
    h_dict = {}
    for idx in self.iter_idx(dimensions):
        # 1/ Forget, input, output and cell activation gates.
        f_dict[idx] = torch.empty(self.dim_in, batch_size, self.size_out)
        for l in range(self.dim_in):
            prev_h = sum(torch.mul(h_dict[prev(idx,k)], self.uf[l][k]) for k in np.nonzero(idx)[0])
            f_dict[idx][l] = torch.sigmoid(self.biasf[l] + torch.mm(x[idx], self.wf[l]) + prev_h)
        prev_h_i = sum(torch.mul(h_dict[prev(idx,k)], self.ui[k]) for k in np.nonzero(idx)[0])
        i_dict[idx] = torch.sigmoid(self.biasi + torch.mm(x[idx], self.wi) + prev_h_i)
        prev_h_o = sum(torch.mul(h_dict[prev(idx,k)], self.uo[k]) for k in np.nonzero(idx)[0])
        o_dict[idx] = torch.sigmoid(self.biaso + torch.mm(x[idx], self.wo) + prev_h_o)
        prev_h_c = sum(torch.mul(h_dict[prev(idx,k)], self.uc[k]) for k in np.nonzero(idx)[0])
        c_dict[idx] = torch.sigmoid(self.biasc + torch.mm(x[idx], self.wc) + prev_h_c)
        # 2/ cell state
        prev_s = sum(torch.mul(f_dict[idx][k], s_dict[prev(idx,k)]) for k in np.nonzero(idx)[0])
        s_dict[idx] = torch.tanh(torch.mul(i_dict[idx], c_dict[idx]) + prev_s)
        # 3/ Final output
        h_dict[idx] = torch.mul(o_dict[idx], s_dict[idx])
    # 所有计算完成后拼装输出
    # 如果迭代顺序是维度展开的flatten顺序,可直接stack后reshape:
    h_list = [h_dict[idx] for idx in self.iter_idx(dimensions)]
    h = torch.stack(h_list, dim=0).reshape(*dimensions, batch_size, self.size_out)
    return h

这种写法的核心是:所有参与计算图构建的中间结果都是独立的张量,不存在对同一个张量的原地修改,最后拼接得到的输出张量自动保留完整的梯度链,不会出现版本冲突问题。


内容的提问来源于stack exchange,提问作者Amélie

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.01 19:06:02