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

使用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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.04 01:00:53