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

WGAN损失仅5个batch内暴跌至负无穷,请求排查代码问题

WGAN损失几轮batch内暴跌至负无穷的问题排查

问题描述

我的WGAN损失仅在5个batch内就暴跌至负无穷,损失数值如下:

Epoch:0 batch_num:0 wgan_loss:-16.413176
Epoch:0 batch_num:1 wgan_loss:14472.721 
Epoch:0 batch_num:2 wgan_loss:-10957247.0
Epoch:0 batch_num:3 wgan_loss:-455000130.0 
Epoch:0 batch_num:4 wgan_loss:-3773285000.0

已尝试降低学习率、将batch size限制为20,但问题依旧。判别器集成了Gradient Reversal Layer(GRL),代码如下:

def grl_hook(coeff):
    def fun1(grad):
        return -coeff*grad.clone()
    return fun1

def calc_coeff(iter_num, high=1.0, low=0.0, alpha=2.0, max_iter=50.0):
    return np.float(2.0 * (high - low) / (1.0 + np.exp(-alpha * iter_num / max_iter)) - (high - low) + low)

class DiscriminatorforWGAN(nn.Module):
    def __init__(self, in_feature, hidden_size):
        super(AdversarialNetworkforCDAN, self).__init__()
        self.ad_layer1 = nn.Linear(in_feature, hidden_size)
        self.ad_layer2 = nn.Linear(hidden_size, hidden_size)
        self.ad_layer3 = nn.Linear(hidden_size, 1)
        self.relu1 = nn.ReLU()
        self.relu2 = nn.ReLU()
        self.dropout1 = nn.Dropout(0.2)
        self.dropout2 = nn.Dropout(0.2)
        self.iter_num = -1
        self.alpha = 1.0
        self.low = 0.0
        self.high = 1.0
        self.max_iter = 15.0
        self.coeff = np.float(0.02)
    def forward(self, x):
        if self.training:
            self.iter_num += 1
        if self.iter_num >= self.max_iter:
            self.iter_num = self.max_iter
        coeff = calc_coeff(self.iter_num, self.high, self.low, self.alpha, self.max_iter)
        self.coeff = coeff
        x = x * 1.0
        x.register_hook(grl_hook(coeff))
        x = self.ad_layer1(x)
        x = self.relu1(x)
        x = self.dropout1(x)
        x = self.ad_layer2(x)
        x = self.relu2(x)
        x = self.dropout2(x)
        y = self.ad_layer3(x)
        return y

生成器为简单CNN网络,已遵循WGAN参数钳制、RMSprop优化器等要求,损失函数如下:

def wgan_loss(values_from_target_side, values_from_source_side):
    W_loss = -torch.mean(values_from_target_side) + torch.mean(values_from_source_side)
    return W_loss

问题排查与解决建议

1. 判别器继承类初始化错误

代码中DiscriminatorforWGAN类的__init__方法调用了super(AdversarialNetworkforCDAN, self).__init__(),但当前类名是DiscriminatorforWGAN,这会导致模型参数初始化异常,直接破坏训练稳定性。

修复:
将super的第一个参数改为当前类名:

super(DiscriminatorforWGAN, self).__init__()

2. GRL应用时机错误

当前代码在输入特征x上直接注册梯度反转钩子,随后才经过线性层和激活层。过早应用GRL会打乱梯度传播逻辑,引发梯度爆炸或消失。

修复:
调整GRL到最后一层线性层前应用:

def forward(self, x):
    if self.training:
        self.iter_num += 1
    if self.iter_num >= self.max_iter:
        self.iter_num = self.max_iter
    coeff = calc_coeff(self.iter_num, self.high, self.low, self.alpha, self.max_iter)
    self.coeff = coeff
    
    x = self.ad_layer1(x)
    x = self.relu1(x)
    x = self.dropout1(x)
    x = self.ad_layer2(x)
    x = self.relu2(x)
    x = self.dropout2(x)
    
    # 在最后一层前应用GRL
    x.register_hook(grl_hook(coeff))
    y = self.ad_layer3(x)
    return y

3. 损失函数符号可能混淆

WGAN标准损失逻辑:

  • 判别器损失:E[D(真实样本)] - E[D(生成样本)]
  • 生成器损失:-E[D(生成样本)]

你的损失函数需要确认values_from_target_side和values_from_source_side对应的样本类型。如果对应关系搞反,会让模型训练方向完全错误,导致损失数值失控。

检查:
若values_from_source_side是真实样本的判别输出,values_from_target_side是生成样本的判别输出,判别器损失应为torch.mean(values_from_source_side) - torch.mean(values_from_target_side),当前损失函数符号相反,会导致判别器优化方向错误。

4. 参数钳制执行不规范

WGAN要求每次判别器更新后,将参数钳制到[-0.01, 0.01]范围。需确认是否在每个batch的判别器优化步骤后都执行了钳制操作,若钳制未生效,判别器参数会无限制增长,引发输出数值爆炸。

检查与修复:
在判别器优化器step()后添加参数钳制代码:

for p in discriminator.parameters():
    p.data.clamp_(-0.01, 0.01)

5. GRL系数增长过快

当前calc_coeff函数中,max_iter=15且alpha=1.0,迭代到第7次时系数就接近1.0,快速增长的梯度反转强度会破坏训练稳定性。

调整:
降低alpha值或增大max_iter,让系数增长更平缓:

self.alpha = 0.5
self.max_iter = 100.0

内容的提问来源于stack exchange,提问作者BaeHann

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.24 15:23:10