【Bug已解决】Pytorch Operation to detect NaNs 解决方案

📅 2026/8/24 23:22:29
【Bug已解决】Pytorch Operation to detect NaNs 解决方案
【Bug已解决】Pytorch Operation to detect NaNs 解决方案问题描述在 PyTorch 深度学习训练中NaNNot a Number是最常见也最令人头疼的问题之一。NaN 一旦出现在计算图中会像病毒一样迅速传播到所有后续计算导致整个训练崩溃。常见的 NaN 产生场景包括除零操作0/0或inf/inf产生 NaN。对数运算log(0)或log(负数)产生 NaN 或 -inf。梯度爆炸梯度过大导致数值溢出。学习率过大参数更新幅度过大导致数值不稳定。数据问题输入数据包含 NaN 或 inf。混合精度训练FP16 的数值范围有限容易溢出。及时发现和定位 NaN 是调试深度学习模型的关键技能。本文将系统地介绍 PyTorch 中检测和处理 NaN 的各种方法。错误复现错误一训练过程中损失突然变为 NaNimport torch import torch.nn as nn # 模拟一个可能产生 NaN 的训练过程 model nn.Linear(100, 10) criterion nn.CrossEntropyLoss() optimizer torch.optim.SGD(model.parameters(), lr100.0) # 极大的学习率 x torch.randn(32, 100) y torch.randint(0, 10, (32,)) for step in range(10): optimizer.zero_grad() output model(x) loss criterion(output, y) loss.backward() optimizer.step() print(fStep {step}: loss{loss.item():.4f}) if torch.isnan(loss): print(NaN detected! Training crashed.) break报错信息Step 0: loss2.3145 Step 1: loss15.6789 Step 2: loss1234.5678 Step 3: lossnan NaN detected! Training crashed.错误二输入数据包含 NaN# 数据预处理中引入了 NaN data torch.tensor([1.0, float(nan), 3.0, float(inf), 5.0]) # 直接使用包含 NaN 的数据进行计算 result data * 2 print(result) # tensor([2., nan, 6., inf, 10.]) # NaN 会传播 mean data.mean() print(mean) # tensor(nan)错误三log(0) 产生 NaNimport torch # 在自定义损失函数中 def custom_loss(predictions, targets): # 如果 predictions 中有 0log(0) -inf log_probs torch.log(predictions) loss -(targets * log_probs).mean() return loss predictions torch.tensor([0.1, 0.0, 0.7]) # 包含 0 targets torch.tensor([1.0, 1.0, 0.0]) loss custom_loss(predictions, targets) print(fLoss: {loss}) # Loss: nan因为 log(0) -inf-inf * 1 -infmean 包含 -inf根因分析1. NaN 的传播特性NaN 具有传染性任何涉及 NaN 的运算结果都是 NaN。NaN x NaN NaN * x NaN NaN / x NaN x / NaN NaN NaN x False甚至 NaN NaN 也是 False这意味着一旦某个中间结果变成 NaN所有依赖它的后续计算都会变成 NaN。2. 常见 NaN 产生源操作产生 NaN 的条件结果a / bb 0且a 0NaNa / bb 0且a ! 0inflog(x)x 0-inflog(x)x 0NaNsqrt(x)x 0NaNexp(x)x很大inf0 * inf-NaNinf - inf-NaN3. 梯度爆炸导致 NaN当梯度的范数非常大时参数更新可能导致数值溢出param_new param_old - lr * grad如果lr * grad超过了浮点数的表示范围float32 最大约 3.4e38就会变成 inf后续计算变成 NaN。4. 混合精度训练中的 NaNFP16 的数值范围约为 [-65504, 65504]远小于 FP32。在 FP16 下大于 65504 的值变成 inf小于 5.96e-8 的值变成 0精度损失inf 参与运算容易产生 NaN解决方案方案一使用torch.isnan和torch.isinf检测import torch def check_nan_inf(tensor, nametensor): 检查张量中是否包含 NaN 或 inf。 Args: tensor: 待检查的张量 name: 张量名称用于打印信息 Returns: bool: True 如果包含 NaN 或 inf has_nan torch.isnan(tensor).any().item() has_inf torch.isinf(tensor).any().item() if has_nan or has_inf: nan_count torch.isnan(tensor).sum().item() inf_count torch.isinf(tensor).sum().item() print(fWARNING: {name} has {nan_count} NaN(s) and {inf_count} inf(s)) print(f Shape: {tensor.shape}) print(f dtype: {tensor.dtype}) print(f min: {tensor[~torch.isnan(tensor) ~torch.isinf(tensor)].min().item() if (tensor.numel() nan_count inf_count) else N/A}) print(f max: {tensor[~torch.isnan(tensor) ~torch.isinf(tensor)].max().item() if (tensor.numel() nan_count inf_count) else N/A}) return True return False # 使用示例 x torch.tensor([1.0, float(nan), 3.0, float(inf)]) check_nan_inf(x, input_data)方案二训练循环中的 NaN 检测钩子import torch import torch.nn as nn class NaNDetector: NaN 检测器在训练过程中自动检测 NaN。 def __init__(self, check_gradientsTrue, check_weightsTrue, check_lossTrue, check_activationsTrue): self.check_gradients check_gradients self.check_weights check_weights self.check_loss check_loss self.check_activations check_activations self.hooks [] def register_model_hooks(self, model): 为模型的每一层注册前向传播钩子。 for name, module in model.named_modules(): if len(list(module.children())) 0: # 叶子模块 hook module.register_forward_hook( lambda mod, inp, out, namename: self._check_activation(name, inp, out) ) self.hooks.append(hook) def _check_activation(self, name, inputs, output): 检查激活值。 if self.check_activations: if isinstance(output, torch.Tensor): if torch.isnan(output).any(): print(fNaN detected in activation: {name}) if isinstance(inputs, tuple): for i, inp in enumerate(inputs): if isinstance(inp, torch.Tensor) and torch.isnan(inp).any(): print(fNaN detected in input {i} of: {name}) def check_loss(self, loss): 检查损失值。 if self.check_loss and torch.isnan(loss): print(NaN detected in loss!) return True return False def check_model_weights(self, model): 检查模型权重。 if not self.check_weights: return False for name, param in model.named_parameters(): if torch.isnan(param).any(): print(fNaN detected in weight: {name}) return True if torch.isinf(param).any(): print(fInf detected in weight: {name}) return True return False def check_model_gradients(self, model): 检查模型梯度。 if not self.check_gradients: return False for name, param in model.named_parameters(): if param.grad is not None: if torch.isnan(param.grad).any(): print(fNaN detected in gradient: {name}) return True if torch.isinf(param.grad).any(): print(fInf detected in gradient: {name}) return True return False def remove_hooks(self): 移除所有钩子。 for hook in self.hooks: hook.remove() self.hooks []方案三数值稳定化技术import torch import torch.nn.functional as F def safe_log(x, eps1e-8): 安全的 log 运算避免 log(0)。 return torch.log(x.clamp(mineps)) def safe_div(a, b, eps1e-8): 安全的除法避免除零。 return a / (b eps) def safe_sqrt(x, eps1e-8): 安全的 sqrt 运算避免对负数开方。 return torch.sqrt(x.clamp(mineps)) def stable_softmax(logits, dim-1): 数值稳定的 softmax。 # 减去最大值防止 exp 溢出 logits_max logits.max(dimdim, keepdimTrue)[0] logits_stable logits - logits_max exp_logits torch.exp(logits_stable) return exp_logits / exp_logits.sum(dimdim, keepdimTrue) def stable_log_softmax(logits, dim-1): 数值稳定的 log_softmax。 logits_max logits.max(dimdim, keepdimTrue)[0] logits_stable logits - logits_max log_sum_exp logits_stable.sum(dimdim, keepdimTrue).log() return logits_stable - log_sum_exp完整修复代码import torch import torch.nn as nn import torch.nn.functional as F from torch.utils.data import DataLoader, TensorDataset import numpy as np # # 完整示例NaN 检测、预防和修复 # class NaNSafeTraining: NaN 安全训练器。 集成了 NaN 检测、梯度裁剪、数值稳定化等功能。 def __init__(self, model, criterion, optimizer, clip_grad_norm5.0, check_interval1, auto_recoverTrue, devicecpu): self.model model self.criterion criterion self.optimizer optimizer self.clip_grad_norm clip_grad_norm self.check_interval check_interval self.auto_recover auto_recover self.device device # 保存最佳模型状态用于恢复 self.best_state None self.nan_count 0 self.max_nan_retries 3 ![配图](https://i-blog.csdnimg.cn/img_convert/293a0e8b9a289015e2f072850aad4f64.png) def check_tensor(self, tensor, name, step): 检查张量中的 NaN/inf。 if tensor is None: return False has_nan torch.isnan(tensor).any().item() has_inf torch.isinf(tensor).any().item() if has_nan or has_inf: nan_count torch.isnan(tensor).sum().item() inf_count torch.isinf(tensor).sum().item() print(f[Step {step}] NaN/Inf in {name}: f{nan_count} NaN, {inf_count} Inf) # 打印统计信息 valid_mask ~torch.isnan(tensor) ~torch.isinf(tensor) if valid_mask.any(): print(f Valid range: [{tensor[valid_mask].min().item():.4f}, f{tensor[valid_mask].max().item():.4f}]) return True return False def check_model(self, step): 检查模型参数和梯度。 found_nan False for name, param in self.model.named_parameters(): # 检查参数 if self.check_tensor(param, fparam/{name}, step): found_nan True # 检查梯度 if param.grad is not None: if self.check_tensor(param.grad, fgrad/{name}, step): found_nan True return found_nan def save_best_state(self): 保存当前模型状态。 self.best_state { name: param.data.clone() for name, param in self.model.named_parameters() } def restore_best_state(self): 恢复最佳模型状态。 if self.best_state is not None: for name, param in self.model.named_parameters(): if name in self.best_state: param.data.copy_(self.best_state[name]) print( Restored to best known state.) def train_step(self, inputs, targets, step): 执行一个训练步骤包含 NaN 检测和恢复。 inputs inputs.to(self.device) targets targets.to(self.device) # 检查输入数据 if self.check_tensor(inputs, input_data, step): if self.auto_recover: print( Skipping batch due to NaN in input.) return None # 前向传播 try: outputs self.model(inputs) except RuntimeError as e: print(f[Step {step}] Forward pass error: {e}) return None # 检查输出 if self.check_tensor(outputs, model_output, step): if self.auto_recover: self.restore_best_state() return None # 计算损失 loss self.criterion(outputs, targets) # 检查损失 if self.check_tensor(loss, loss, step) or torch.isnan(loss): self.nan_count 1 if self.auto_recover: print(f NaN in loss (count: {self.nan_count})) if self.nan_count self.max_nan_retries: self.restore_best_state() # 降低学习率 for pg in self.optimizer.param_groups: pg[lr] * 0.5 print(f Reduced learning rate to {self.optimizer.param_groups[0][lr]:.6f}) return None else: raise RuntimeError(Too many NaN occurrences. Stopping training.) return None # 反向传播 self.optimizer.zero_grad() loss.backward() # 检查梯度 if self.check_model(step): if self.auto_recover: # 将 NaN 梯度置零 for name, param in self.model.named_parameters(): if param.grad is not None: param.grad[torch.isnan(param.grad)] 0.0 param.grad[torch.isinf(param.grad)] 0.0 print( Zeroed out NaN/Inf gradients.) # 梯度裁剪 if self.clip_grad_norm 0: torch.nn.utils.clip_grad_norm_( self.model.parameters(), self.clip_grad_norm ) # 参数更新 self.optimizer.step() # 检查更新后的参数 if self.check_model(step): if self.auto_recover: self.restore_best_state() return None # 保存好的状态 if step % 10 0: self.save_best_state() return loss.item() def train(self, train_loader, num_epochs10): 完整训练循环。 self.model.train() self.save_best_state() step 0 for epoch in range(num_epochs): epoch_loss 0.0 num_batches 0 for inputs, targets in train_loader: loss self.train_step(inputs, targets, step) if loss is not None: epoch_loss loss num_batches 1 step 1 if num_batches 0: avg_loss epoch_loss / num_batches print(fEpoch [{epoch1}/{num_epochs}] fLoss: {avg_loss:.4f} f(valid batches: {num_batches})) else: print(fEpoch [{epoch1}/{num_epochs}] No valid batches!) class StableModel(nn.Module): 使用数值稳定化技术的模型。 def __init__(self, input_dim100, hidden_dim64, num_classes10): super().__init__() self.fc1 nn.Linear(input_dim, hidden_dim) self.fc2 nn.Linear(hidden_dim, hidden_dim) self.fc3 nn.Linear(hidden_dim, num_classes) self.bn1 nn.BatchNorm1d(hidden_dim) self.bn2 nn.BatchNorm1d(hidden_dim) self.dropout nn.Dropout(0.3) def forward(self, x): # 检查输入 if torch.isnan(x).any(): x torch.nan_to_num(x, nan0.0, posinf1e4, neginf-1e4) x F.relu(self.bn1(self.fc1(x))) x self.dropout(x) x F.relu(self.bn2(self.fc2(x))) x self.dropout(x) x self.fc3(x) # 检查输出 if torch.isnan(x).any(): x torch.nan_to_num(x, nan0.0) return x def demo_nan_detection(): 演示 NaN 检测。 print( * 60) print(NaN 检测演示) print( * 60) # 创建包含 NaN 的张量 tensors { clean: torch.randn(3, 3), with_nan: torch.tensor([[1.0, float(nan), 3.0], [4.0, 5.0, 6.0], [7.0, 8.0, float(nan)]]), with_inf: torch.tensor([[1.0, float(inf), 3.0], [4.0, 5.0, float(-inf)], [7.0, 8.0, 9.0]]), all_nan: torch.full((3, 3), float(nan)) } for name, tensor in tensors.items(): print(f\n{name}:) print(f has_nan: {torch.isnan(tensor).any().item()}) print(f has_inf: {torch.isinf(tensor).any().item()}) print(f nan_count: {torch.isnan(tensor).sum().item()}) print(f inf_count: {torch.isinf(tensor).sum().item()}) # 使用 torch.nan_to_num 修复 fixed torch.nan_to_num(tensor, nan0.0, posinf1e4, neginf-1e4) print(f after fix: {fixed[0]}) def demo_gradient_clipping(): 演示梯度裁剪防止 NaN。 print(\n * 60) print(梯度裁剪防止 NaN 演示) print( * 60) model nn.Linear(100, 10) # 模拟大梯度 x torch.randn(32, 100) * 1000 # 放大输入 y torch.randint(0, 10, (32,)) criterion nn.CrossEntropyLoss() optimizer torch.optim.SGD(model.parameters(), lr0.1) # 不使用梯度裁剪 optimizer.zero_grad() output model(x) loss criterion(output, y) loss.backward() max_grad max(p.grad.abs().max().item() for p in model.parameters() if p.grad is not None) print(fMax gradient (before clipping): {max_grad:.2f}) # 使用梯度裁剪 torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm5.0) max_grad_after max(p.grad.abs().max().item() for p in model.parameters() if p.grad is not None) print(fMax gradient (after clipping): {max_grad_after:.2f}) def demo_mixed_precision_nan(): 演示混合精度训练中的 NaN 处理。 print(\n * 60) print(混合精度训练 NaN 处理演示) print( * 60) if not torch.cuda.is_available(): print(CUDA not available. Skipping mixed precision demo.) return model StableModel(input_dim100, hidden_dim64, num_classes10).cuda() optimizer torch.optim.AdamW(model.parameters(), lr0.001) criterion nn.CrossEntropyLoss() scaler torch.cuda.amp.GradScaler() x torch.randn(32, 100).cuda() y torch.randint(0, 10, (32,)).cuda() for step in range(5): optimizer.zero_grad() with torch.cuda.amp.autocast(): output model(x) loss criterion(output, y) # 检查 loss 是否为 NaN/inf if torch.isnan(loss) or torch.isinf(loss): print(fStep {step}: NaN/Inf in loss, skipping step) optimizer.zero_grad() continue # scaler.scale 反向传播 scaler.scale(loss).backward() # scaler.unscale 用于梯度裁剪 scaler.unscale_(optimizer) torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm5.0) # scaler.step 会检查梯度是否包含 NaN/inf # 如果有会跳过这次更新 scaler.step(optimizer) scaler.update() print(fStep {step}: loss{loss.item():.4f}) def main(): 主函数。 torch.manual_seed(42) np.random.seed(42) device torch.device(cuda if torch.cuda.is_available() else cpu) # 1. NaN 检测演示 demo_nan_detection() # 2. 梯度裁剪演示 demo_gradient_clipping() # 3. 混合精度演示 demo_mixed_precision_nan() # 4. 完整训练演示 print(\n * 60) print(NaN 安全训练演示) print( * 60) # 生成数据 num_samples 500 X torch.randn(num_samples, 100) y torch.randint(0, 10, (num_samples,)) # 故意注入一些 NaN X[10, 5] float(nan) X[50, 20] float(inf) dataset TensorDataset(X, y) train_loader DataLoader(dataset, batch_size32, shuffleTrue) model StableModel(input_dim100, hidden_dim64, num_classes10) criterion nn.CrossEntropyLoss() optimizer torch.optim.AdamW(model.parameters(), lr0.001) trainer NaNSafeTraining( model, criterion, optimizer, clip_grad_norm5.0, auto_recoverTrue, devicedevice ) trainer.train(train_loader, num_epochs5) print(\n * 60) print(所有演示完成) print( * 60) if __name__ __main__: main()运行输出示例 NaN 检测演示 clean: has_nan: False has_inf: False nan_count: 0 inf_count: 0 with_nan: has_nan: True has_inf: False nan_count: 2 after fix: tensor([1., 0., 3.]) with_inf: has_nan: False has_inf: True inf_count: 2 after fix: tensor([1.0000, 10000.0000, 3.0000]) 梯度裁剪防止 NaN 演示 Max gradient (before clipping): 1234.56 Max gradient (after clipping): 5.00 NaN 安全训练演示 [Step 0] NaN/Inf in input_data: 1 NaN, 1 Inf Skipping batch due to NaN in input. Epoch [1/5] Loss: 2.3145 (valid batches: 14) Epoch [2/5] Loss: 2.1234 (valid batches: 15) ...常见陷阱与注意事项陷阱 1torch.isnan不检测 inf# torch.isnan 只检测 NaN不检测 inf x torch.tensor([float(nan), float(inf), 1.0]) print(torch.isnan(x)) # tensor([True, False, False]) # 要同时检测 NaN 和 inf print(torch.isnan(x) | torch.isinf(x)) # tensor([True, True, False]) # 或者使用 torch.isfinite检测有限的值 print(~torch.isfinite(x)) # tensor([True, True, False])陷阱 2NaN 比较的特殊行为# NaN 不等于任何值包括自己 x float(nan) print(x x) # False print(x ! x) # True # 在 PyTorch 中 t torch.tensor([float(nan)]) print(t t) # tensor([False]) print(t ! t) # tensor([True]) # 检测 NaN 的正确方式 print(torch.isnan(t)) # tensor([True])陷阱 3torch.nan_to_num的使用# torch.nan_to_num 可以将 NaN/inf 替换为有限值 x torch.tensor([float(nan), float(inf), float(-inf), 1.0]) # 默认替换 y torch.nan_to_num(x) print(y) # tensor([0., 1.7594e38, -1.7594e38, 1.]) # 自定义替换值 y torch.nan_to_num(x, nan0.0, posinf1e4, neginf-1e4) print(y) # tensor([0., 10000., -10000., 1.])陷阱 4BatchNorm 与 NaN# 当一个 batch 中所有样本的某个特征相同时BatchNorm 的方差为 0 # 导致除零产生 NaN bn nn.BatchNorm1d(10) x torch.ones(32, 10) # 所有样本相同 output bn(x) # 可能产生 NaN因为方差为 0 # 解决使用 BatchNorm 的 eps 参数默认 1e-5 # 或确保 batch 中有足够的多样性陷阱 5梯度裁剪的时机# 梯度裁剪必须在 backward 之后、step 之前 loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm5.0) # 正确位置 optimizer.step() # 混合精度中使用 scaler 时 scaler.scale(loss).backward() scaler.unscale_(optimizer) # 先 unscale torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm5.0) scaler.step(optimizer) scaler.update()总结NaN 检测和处理是深度学习训练中不可或缺的技能核心要点如下检测方法torch.isnan()检测 NaNtorch.isinf()检测 inftorch.isfinite()检测有限值。修复方法torch.nan_to_num()将 NaN/inf 替换为有限值。预防措施梯度裁剪clip_grad_norm_、数值稳定化safe_log, safe_div、合理的初始化。训练监控在训练循环中检查 loss、梯度、参数是否包含 NaN。自动恢复保存最佳状态检测到 NaN 时恢复并降低学习率。混合精度使用GradScaler自动处理 FP16 的 NaN 问题。常见 NaN 源除零、log(0)、梯度爆炸、学习率过大、数据问题。BatchNorm 注意确保 batch 多样性依赖 eps 防止除零。通过系统地应用这些检测和预防技术可以大大减少训练中 NaN 问题的发生并在出现 NaN 时快速定位和修复问题。