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

训练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稳定出现该问题。

样本输出:
sample output

模型核心代码

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.23 18:52:05