PyTorch中使用自定义损失函数更新VAE潜在向量时的梯度错误及不收敛问题求助
PyTorch中使用自定义损失函数更新VAE潜在向量时的梯度错误及不收敛问题求助
看起来你遇到的问题核心是潜在向量没有被正确设置为可训练张量,导致梯度无法反向传播到它身上,进而出现梯度错误或者不收敛的情况。我来帮你拆解问题并给出修复方案:
问题根源分析
梯度错误的原因
你在with torch.no_grad()上下文里生成了z_latent_vect,这会直接切断它的计算图,导致这个张量的requires_grad为False且没有grad_fn。当你把它传给优化器时,PyTorch自然会报错——优化器根本无法对一个没有梯度信息的张量进行更新。强行设置loss的requires_grad后不收敛的原因
你解开注释的loss = Variable(loss, requires_grad=True)只是强行给损失张量加上了梯度标记,但潜在向量z_latent_vect本身还是没有梯度连接。反向传播时梯度根本流不到z_latent_vect那里,所以它完全不会被更新,损失自然也不会下降。
具体修复方案
我们需要让z_latent_vect成为一个独立的可训练张量,同时保证编码器和解码器的权重不被更新。下面是修改后的完整代码,我会标注关键修改点:
class VAE_GD_Loss(nn.Module): def __init__(self): super(VAE_GD_Loss, self).__init__() def forward(self, bad_seg, recons_mask, vector): # l2 normed squared and the soft dice loss are calculated loss = torch.sum(vector**2) + Soft_Dice_Loss(recons_mask, bad_seg) return loss # 先确保Soft_Dice_Loss的实现是可导的(示例实现) class Soft_Dice_Loss(nn.Module): def __init__(self): super().__init__() def forward(self, pred, target): pred = pred.float() target = target.float() intersection = torch.sum(pred * target) union = torch.sum(pred) + torch.sum(target) # 加小epsilon避免除零 dice = (2. * intersection + 1e-6) / (union + 1e-6) return 1 - dice def optimize_latent_vector(model, inp__, num_epochs=50, learning_rate=0.01): inp__ = inp__.to(device) # 1. 移除torch.no_grad(),获取初始潜在向量后切断与编码器的连接 mu, log_var = model.encoder(inp__) z_latent_vect = model.reparameterize(mu, log_var) # 关键:detach切断与编码器的计算图连接,手动开启requires_grad让它成为可训练张量 z_latent_vect = z_latent_vect.detach().requires_grad_(True) # 2. 只把潜在向量传给优化器 optimizer_lat = torch.optim.Adam([z_latent_vect], lr=learning_rate) dec_only = model.decoder # 3. 强制冻结解码器所有参数,确保权重不会被更新 for param in dec_only.parameters(): param.requires_grad = False # 4. 提前创建损失函数实例,避免循环内重复初始化 vg_loss = VAE_GD_Loss() for epoch in range(num_epochs): optimizer_lat.zero_grad() dec_only.eval() # 从潜在向量解码 recons_mask = dec_only(z_latent_vect) # 计算损失 loss = vg_loss(inp__, recons_mask, z_latent_vect) # 反向传播+更新潜在向量 loss.backward() optimizer_lat.step() print(f"Epoch {epoch+1}/{num_epochs}: Loss = {loss.item():.4f}") return z_latent_vect
额外注意事项
- 如果你希望编码器的权重也完全不更新,在调用这个函数之前可以设置
model.encoder.eval(),并给编码器的所有参数设置requires_grad=False。 - 确保你的
Soft_Dice_Loss实现是可导的,不要使用任何会切断计算图的操作(比如torch.no_grad()、.detach()除非必要)。 - 可以尝试调整学习率和 epoch 数,比如把学习率调到0.001试试,有时候过大的学习率会导致损失震荡不收敛。
这样修改后,潜在向量就能正常被优化器更新,损失也会按照你的预期逐渐下降了。
备注:内容来源于stack exchange,提问作者Jimut123
相关产品推荐
相关产品推荐

