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

