PyTorch复现MoCo V1遇二次反向传播及原地操作错误求助
问题背景
作为PyTorch初学者,在复现MoCo V1模型时遭遇梯度相关错误。已单独验证encoder和momentum_encoder的训练逻辑正常,最初认为设置retain_graph=True无必要,但出现如下报错:
RuntimeError: Trying to backward through the graph a second time (or directly access saved tensors after they have already been freed). Saved intermediate values of the graph are freed when you call .backward() or autograd.grad(). Specify retain_graph=True if you need to backward through the graph a second time or if you need to access saved tensors after calling backward.
尝试改为loss.backward(retain_graph=True)后,又触发新错误:
RuntimeError: one of the variables needed for gradient computation has been modified by an inplace operation: [torch.cuda.FloatTensor [128, 32]], which is output 0 of AsStridedBackward0, is at version 3; expected version 2 instead. Hint: the backtrace further above shows the operation that failed to compute its gradient. The variable in question was changed in there or anywhere later. Good luck!
训练代码如下:
def train(train_loader, encoder, momentum_encoder, device, epochs, m): torch.autograd.set_detect_anomaly(True) encoder_optimizer = optim.Adam(encoder.parameters(), lr=0.001) encoder.apply(xavier_init_weights) encoder.to(device) momentum_encoder.apply(xavier_init_weights) momentum_encoder.to(device) queue = Memory_queue(4) criterion = nn.CrossEntropyLoss() for batch_idx, (datas, _) in enumerate(train_loader, 1): encoder.eval() momentum_encoder.eval() datas = datas.to(device) q = encoder(datas) queue.put(q) for epoch in range(epochs): encoder.train() for batch_idx, (datas, _) in enumerate(train_loader, 1): # 调整起始值为1 encoder_optimizer.zero_grad() encoder.zero_grad() datas = add_noise(datas).to(device) datas_key1 = add_noise(datas).to(device) q = encoder(datas) k = momentum_encoder(datas_key1) k = k.detach() N, C = q.shape l_pos = torch.squeeze(torch.bmm(q.view(N, 1, C), k.view(N, C, 1)).to(device), dim=2) temp_queue = torch.cat(queue.get_all(), dim=0) K, _ = temp_queue.shape l_neg = torch.mm(q.view(N, C), temp_queue.view(C, K)).to(device) logits = torch.cat([l_pos, l_neg], dim=1).to(device) labels = torch.zeros(N).long().to(device) loss = criterion(logits, labels) loss.backward() encoder_optimizer.step() update_model_ema(encoder, momentum_encoder, 0.5) queue.put(q)
问题根源与修复方案
1. 初始反向传播错误原因
epoch循环外的预填充队列阶段,调用encoder(datas)时未禁用梯度计算,导致残留计算图。后续训练时的反向传播尝试复用已释放的图结构,触发第一次报错。
2. 原地操作错误原因
设置retain_graph=True后,队列中存储的q仍带有计算图,后续队列的原地修改操作(如覆盖旧元素)会破坏梯度计算所需的张量版本,引发第二次错误。
具体修复步骤
- 预填充队列禁用梯度:用
torch.no_grad()包裹预填充循环,避免存储不必要的计算图:with torch.no_grad(): for batch_idx, (datas, _) in enumerate(train_loader, 1): encoder.eval() momentum_encoder.eval() datas = datas.to(device) q = encoder(datas) queue.put(q) - 存入队列前剥离梯度:
q是encoder的输出,带有计算图,存入队列前用detach()剥离梯度,避免后续队列操作干扰梯度计算:queue.put(q.detach()) - 移除冗余梯度清零:
encoder_optimizer.zero_grad()已能清除encoder参数的梯度,无需额外调用encoder.zero_grad()。 - 修正动量编码器使用方式:动量编码器始终处于
eval模式,前向传播时禁用梯度;MoCo的动量系数m应接近1(如0.999)而非0.5,且update_model_ema需用原地更新方式,避免影响计算图:def update_model_ema(encoder, momentum_encoder, m): for enc_param, momentum_param in zip(encoder.parameters(), momentum_encoder.parameters()): momentum_param.data = m * momentum_param.data + (1 - m) * enc_param.data - 修正负样本计算的转置方式:用
t()替代view(C, K),避免可能的张量形状错误:l_neg = torch.mm(q.view(N, C), temp_queue.t())
修复后的完整训练函数
def train(train_loader, encoder, momentum_encoder, device, epochs, m): torch.autograd.set_detect_anomaly(True) encoder_optimizer = optim.Adam(encoder.parameters(), lr=0.001) encoder.apply(xavier_init_weights) encoder.to(device) momentum_encoder.apply(xavier_init_weights) momentum_encoder.to(device) queue = Memory_queue(4) criterion = nn.CrossEntropyLoss() # 预填充队列,禁用梯度 with torch.no_grad(): for batch_idx, (datas, _) in enumerate(train_loader, 1): encoder.eval() momentum_encoder.eval() datas = datas.to(device) q = encoder(datas) queue.put(q) for epoch in range(epochs): encoder.train() momentum_encoder.eval() # 动量编码器保持eval模式 for batch_idx, (datas, _) in enumerate(train_loader, 1): encoder_optimizer.zero_grad() datas = add_noise(datas).to(device) datas_key1 = add_noise(datas).to(device) q = encoder(datas) # 动量编码器前向传播禁用梯度 with torch.no_grad(): k = momentum_encoder(datas_key1) k = k.detach() N, C = q.shape l_pos = torch.squeeze(torch.bmm(q.view(N, 1, C), k.view(N, C, 1)), dim=2) temp_queue = torch.cat(queue.get_all(), dim=0) l_neg = torch.mm(q.view(N, C), temp_queue.t()) logits = torch.cat([l_pos, l_neg], dim=1) labels = torch.zeros(N).long().to(device) loss = criterion(logits, labels) loss.backward() encoder_optimizer.step() # 更新动量编码器,禁用梯度 with torch.no_grad(): update_model_ema(encoder, momentum_encoder, m) # 存入队列前剥离梯度 queue.put(q.detach())
内容的提问来源于stack exchange,提问作者LPP

