训练CIBHash模型反复出现NaN问题的原因与解决问询
CIBHash模型复现中损失爆炸引发NaN问题的分析与解决建议
复现CIBHash模型(项目仓库:https://github.com/zexuanqiu/CIBHash)时,使用官方CIFAR-10数据集,每次评估后都会出现损失爆炸进而产生NaN问题,具体复现场景如下:
- 执行命令
python main.py cifar16 --train --dataset cifar10 --encode_length 16 --cuda(默认validate_frequency=20),在epoch=20评估完成后,epoch=21出现损失爆炸; - 执行命令
python main.py cifar16 --train --dataset cifar10 --encode_length 16 --cuda --validate_frequency=3,设置validate_frequency=3后,在epoch=4稳定出现该问题。
样本输出:
模型核心代码
run_training_session函数
def run_training_session(self, run_num, logger): self.train() # Scramble hyperparameters if number of runs is greater than 1. if self.hparams.num_runs > 1: logger.log('RANDOM RUN: %d/%d' % (run_num, self.hparams.num_runs)) for hparam, values in self.get_hparams_grid().items(): assert hasattr(self.hparams, hparam) self.hparams.__dict__[hparam] = random.choice(values) random.seed(self.hparams.seed) torch.manual_seed(self.hparams.seed) self.define_parameters() # if encode_length is 16, then al least 80 epochs! if self.hparams.encode_length == 16: self.hparams.epochs = max(80, self.hparams.epochs) logger.log('hparams: %s' % self.flag_hparams()) device = torch.device('cuda' if self.hparams.cuda else 'cpu') self.to(device) optimizer = self.configure_optimizers() train_loader, val_loader, _, database_loader = self.data.get_loaders( self.hparams.batch_size, self.hparams.num_workers, shuffle_train=True, get_test=False) best_val_perf = float('-inf') best_state_dict = None bad_epochs = 0 try: for epoch in range(1, self.hparams.epochs + 1): forward_sum = {} num_steps = 0 for batch_num, batch in enumerate(train_loader): optimizer.zero_grad() imgi, imgj, _ = batch imgi = imgi.to(device) imgj = imgj.to(device) forward = self.forward(imgi, imgj, device) for key in forward: if key in forward_sum: forward_sum[key] += forward[key] else: forward_sum[key] = forward[key] num_steps += 1 if math.isnan(forward_sum['loss']): logger.log('Stopping epoch because loss is NaN') break forward['loss'].backward() optimizer.step() if math.isnan(forward_sum['loss']): logger.log('Stopping training session because loss is NaN') break logger.log('End of epoch {:3d}'.format(epoch), False) logger.log(' '.join([' | {:s} {:8.4f}'.format( key, forward_sum[key] / num_steps) for key in forward_sum]), True) if epoch % self.hparams.validate_frequency == 0: print('evaluating...') val_perf = self.evaluate(database_loader, val_loader, self.data.topK, device) logger.log(' | val perf {:8.4f}'.format(val_perf), False) if val_perf > best_val_perf: best_val_perf = val_perf bad_epochs = 0 logger.log('\t\t*Best model so far, deep copying*') best_state_dict = deepcopy(self.state_dict()) else: bad_epochs += 1 logger.log('\t\tBad epoch %d' % bad_epochs) if bad_epochs > self.hparams.num_bad_epochs: break except KeyboardInterrupt: logger.log('-' * 89) logger.log('Exiting from training early') return best_state_dict, best_val_perf
forward函数
def forward(self, imgi, imgj, device): imgi = self.vgg.features(imgi) imgi = imgi.view(imgi.size(0), -1) imgi = self.vgg.classifier(imgi) prob_i = torch.sigmoid(self.encoder(imgi)) z_i = hash_layer(prob_i - torch.empty_like(prob_i).uniform_().to(prob_i.device)) imgj = self.vgg.features(imgj) imgj = imgj.view(imgj.size(0), -1) imgj = self.vgg.classifier(imgj) prob_j = torch.sigmoid(self.encoder(imgj)) z_j = hash_layer(prob_j - torch.empty_like(prob_j).uniform_().to(prob_j.device)) kl_loss = (self.compute_kl(prob_i, prob_j) + self.compute_kl(prob_j, prob_i)) / 2 contra_loss = self.criterion(z_i, z_j, device) loss = contra_loss + self.hparams.weight * kl_loss return {'loss': loss, 'contra_loss': contra_loss, 'kl_loss': kl_loss}
已尝试的解决方案包括:
- 按照GitHub issue #6的建议修改
z_i和z_j的计算逻辑; - 使用
gradient_gripping方法;
但均未解决NaN问题,且作者在GitHub issue #7中表示训练时未遇到该问题。
可能的原因及解决方案
1. 评估阶段未切换模型模式
当前代码仅在训练前调用self.train(),但评估阶段未显式切换到eval()模式,导致BatchNorm、Dropout等层在评估时仍保持训练状态,参数被意外更新,后续训练时引发梯度异常。
解决方案:在evaluate函数开头添加模式切换,并在评估结束后恢复训练模式:
def evaluate(self, database_loader, val_loader, topK, device): self.eval() with torch.no_grad(): # 原有评估逻辑 self.train()
2. KL散度计算数值不稳定
当prob_i或prob_j趋近于0或1时,torch.log会产生无穷大值,直接导致KL散度爆炸。
解决方案:在计算KL散度时添加小epsilon值避免数值溢出:
def compute_kl(self, p, q): eps = 1e-8 p = torch.clamp(p, eps, 1-eps) q = torch.clamp(q, eps, 1-eps) return torch.sum(p * torch.log(p / q), dim=1).mean()
3. 优化器学习率过高
评估后模型参数发生变化,若初始学习率过高,后续训练步的梯度更新会导致参数突变,引发损失爆炸。
解决方案:
- 降低初始学习率,例如将默认学习率调整为原来的1/10;
- 使用学习率调度器(如
torch.optim.lr_scheduler.ReduceLROnPlateau),在验证性能下降时自动降低学习率。
4. 哈希层梯度不稳定
hash_layer是不可导的离散操作,直通估计(STE)的梯度存在噪声,累积后易引发数值不稳定。
解决方案:
- 对模型参数的梯度进行裁剪,在
forward['loss'].backward()后添加:
torch.nn.utils.clip_grad_norm_(self.parameters(), max_norm=1.0)
- 暂时移除哈希层的随机采样逻辑,先训练连续的概率值,待模型稳定后再加入离散采样。
5. 损失累加方式存在隐患
当前代码直接对张量进行累加(forward_sum[key] += forward[key]),可能导致梯度累积或数值溢出。
解决方案:改为累加损失的数值而非张量:
forward_sum[key] += forward[key].item()
内容的提问来源于stack exchange,提问作者Fuug4Mo
相关产品推荐
相关产品推荐

