WGAN损失仅5个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

