使用torchmetrics时如何处理NaN值?PyTorch Lightning项目求助
解决PyTorch Lightning+TorchMetrics中指标NaN问题
针对部分样本导致批次指标NaN、进而epoch聚合指标失效的问题,给你三个可行方案,全部基于你当前使用的工具栈实现:
方案1:计算前过滤无效样本
如果能提前判断哪些样本不符合指标计算条件(比如标签缺失、输入异常),直接在批次里过滤掉这些样本,从根源避免批次指标出现NaN。
示例代码(以多分类Accuracy为例):
import torch from torchmetrics import Accuracy import pytorch_lightning as pl class MyModel(pl.LightningModule): def __init__(self): super().__init__() self.train_acc = Accuracy(task="multiclass", num_classes=10) def calculate_loss(self, x, y): # 替换成你自己的损失计算逻辑 preds = self(x) return torch.nn.functional.cross_entropy(preds, y) def training_step(self, batch, batch_idx): x, y = batch # 假设y=-1的样本是无效样本,生成过滤掩码 valid_mask = y != -1 # 若当前批次全是无效样本,直接跳过指标更新 if valid_mask.sum() == 0: loss = self.calculate_loss(x, y) self.log("train_loss", loss) return loss # 只保留有效样本计算预测和指标 x_valid = x[valid_mask] y_valid = y[valid_mask] preds_valid = self(x_valid) loss = self.calculate_loss(x_valid, y_valid) self.train_acc(preds_valid, y_valid) # 日志epoch级聚合指标 self.log("train_acc", self.train_acc, on_epoch=True, prog_bar=True) self.log("train_loss", loss) return loss
方案2:用MeanMetric聚合有效批次
如果没法提前过滤样本,只能接受部分批次指标为NaN,可以用MeanMetric来聚合所有非NaN的批次结果,自动忽略NaN值。
示例代码:
import torch from torchmetrics import MeanMetric, Accuracy import pytorch_lightning as pl class MyModel(pl.LightningModule): def __init__(self): super().__init__() # 用于计算单批次指标 self.batch_acc = Accuracy(task="multiclass", num_classes=10) # 用MeanMetric聚合所有非NaN的批次结果 self.train_acc = MeanMetric() def calculate_loss(self, x, y): preds = self(x) return torch.nn.functional.cross_entropy(preds, y) def training_step(self, batch, batch_idx): x, y = batch preds = self(x) loss = self.calculate_loss(x, y) # 计算当前批次的指标值 batch_acc_val = self.batch_acc(preds, y) # 仅当批次指标非NaN时,更新聚合器 if not torch.isnan(batch_acc_val): self.train_acc.update(batch_acc_val) # 日志epoch级聚合后的指标 self.log("train_acc", self.train_acc, on_epoch=True, prog_bar=True) self.log("train_loss", loss) return loss def on_epoch_end(self): # 重置批次指标和聚合器,为下一轮epoch做准备 self.batch_acc.reset() self.train_acc.reset()
方案3:自定义支持NaN跳过的指标
如果上述方案都不满足需求,可以自定义一个torchmetrics指标类,在内部自动过滤含NaN的样本或批次。
示例代码:
import torch from torchmetrics import Metric from torchmetrics.utilities.data import dim_zero_cat class AccWithNaNHandling(Metric): def __init__(self, num_classes): super().__init__() self.num_classes = num_classes # 定义状态变量,用于累计正确数和总样本数 self.add_state("correct", default=torch.tensor(0), dist_reduce_fx="sum") self.add_state("total", default=torch.tensor(0), dist_reduce_fx="sum") def update(self, preds, target): # 过滤掉target为NaN,或者preds中存在NaN的样本 valid_mask = ~torch.isnan(target) & ~torch.isnan(preds).any(dim=-1) # 若当前批次无有效样本,直接返回 if valid_mask.sum() == 0: return preds_valid = preds[valid_mask] target_valid = target[valid_mask] # 计算正确数并累计 preds_argmax = preds_valid.argmax(dim=-1) self.correct += (preds_argmax == target_valid).sum() self.total += target_valid.numel() def compute(self): # 若累计总样本数为0,返回默认值(可根据需求调整) if self.total == 0: return torch.tensor(0.0) return self.correct.float() / self.total
在LightningModule中使用这个自定义指标:
class MyModel(pl.LightningModule): def __init__(self): super().__init__() self.train_acc = AccWithNaNHandling(num_classes=10) def calculate_loss(self, x, y): preds = self(x) return torch.nn.functional.cross_entropy(preds, y) def training_step(self, batch, batch_idx): x, y = batch preds = self(x) loss = self.calculate_loss(x, y) # 直接更新自定义指标,内部会自动处理NaN样本 self.train_acc(preds, y) self.log("train_acc", self.train_acc, on_epoch=True, prog_bar=True) return loss
内容的提问来源于stack exchange,提问作者Nadav Schweiger
相关产品推荐
相关产品推荐

