PyTorch RNN字符预测时梯度变量被inplace操作修改报错如何解决?
报错原因
- 就地修改变量报错:你将参数更新操作
optimizer.step()放在了单个字符的时间步循环内,同时开启了retain_graph=True保留了旧计算图。前序时间步更新参数后,旧计算图依赖的参数值已经被修改,后续时间步反向传播时无法找到原始参数计算梯度,就触发了版本不匹配的错误。 - 二次反向传播报错:移除
retain_graph=True后,第一次时间步的反向传播会释放计算图,而你累加的loss还依赖前面已经被释放的计算图节点,后续反向传播时就会提示计算图已经被销毁。
修复方法
你的训练逻辑完全错位了:字符级RNN的标准训练流程是先跑完整个序列的所有时间步,累加完所有损失后,再统一做梯度清空、反向传播、参数更新,不需要每个时间步都执行这三个操作。
修改后的训练函数如下:
def train(train_seq, target_seq): hidden = rnn.initHidden() # 整个序列训练开始前清空梯度 optimizer.zero_grad() total_loss = 0 for i in range(len(train_seq)): output, hidden = rnn(train_seq[i].unsqueeze(0), hidden) target_class = (target_seq[i] == 1).nonzero(as_tuple=True)[0] total_loss += criterion(output, target_class) print(f"done {i} loop") # 所有时间步计算完成后统一反向传播、更新参数 total_loss.backward() optimizer.step() return output, total_loss.item() / train_seq.size(0)
如果你训练的序列过长,担心显存占用过高,可以采用截断反向传播(BPTT)策略,每固定步数更新一次参数并断开隐藏状态的计算图:
def train_long_seq(train_seq, target_seq, trunc_step=32): hidden = rnn.initHidden() optimizer.zero_grad() total_loss = 0 avg_loss = 0 for i in range(len(train_seq)): output, hidden = rnn(train_seq[i].unsqueeze(0), hidden) target_class = (target_seq[i] == 1).nonzero(as_tuple=True)[0] step_loss = criterion(output, target_class) total_loss += step_loss avg_loss += step_loss.item() print(f"done {i} loop") # 每trunc_step步更新一次参数 if (i + 1) % trunc_step == 0: total_loss.backward() optimizer.step() optimizer.zero_grad() # 断开隐藏状态的计算图,避免反向传播到已处理的时间步 hidden = hidden.detach() total_loss = 0 return output, avg_loss / train_seq.size(0)
内容的提问来源于stack exchange,提问作者Joseph Harvey
相关产品推荐
相关产品推荐

