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

如何在PyTorch Lightning训练阶段对Faster RCNN的N个训练批次计算mAP?

解决方案:训练过程中计算训练批次的Mean AP指标

可以通过val_check_interval结合自定义逻辑实现需求,但不建议用on_validation_epoch_end,更适合用validation_step或on_train_batch_end来完成训练批次的AP计算,以下是两种具体实现方式:


方式一:利用验证间隔触发训练批次AP计算

配置val_check_interval=N让训练每N个批次触发一次验证流程,在验证步骤中采样训练批次计算指标:

import pytorch_lightning as pl
import torch
from torchmetrics.detection import MeanAveragePrecision

class FasterRCNNModule(pl.LightningModule):
    def __init__(self, model, num_train_ap_batches=5):
        super().__init__()
        self.model = model
        self.train_map = MeanAveragePrecision()
        self.num_train_batches = num_train_ap_batches
        self.train_iter = None

    def validation_epoch_start(self):
        # 初始化训练集迭代器,重置指标
        if self.train_iter is None:
            self.train_iter = iter(self.trainer.datamodule.train_dataloader())
        self.train_map.reset()

    def validation_step(self, batch, batch_idx):
        # 仅在指定批次内采样训练数据计算AP
        if batch_idx < self.num_train_batches:
            try:
                images, targets = next(self.train_iter)
            except StopIteration:
                # 迭代器耗尽时重新初始化
                self.train_iter = iter(self.trainer.datamodule.train_dataloader())
                images, targets = next(self.train_iter)
            
            # 切换模型模式获取预测结果
            self.model.eval()
            with torch.no_grad():
                preds = self.model(images, targets)  # eval模式返回预测结果
            self.model.train()  # 切回训练模式
            
            # 更新并计算指标
            self.train_map.update(preds, targets)
            if batch_idx == self.num_train_batches - 1:
                map_result = self.train_map.compute()
                self.log("train/mAP", map_result["map"], prog_bar=True, logger=True)

    def configure_optimizers(self):
        return torch.optim.SGD(self.model.parameters(), lr=0.001)

训练时只需设置:

trainer = pl.Trainer(val_check_interval=N)  # N为触发计算的训练批次间隔

方式二:用训练批次结束回调直接计算

在on_train_batch_end中判断批次间隔,触发训练批次的AP计算:

import pytorch_lightning as pl
import torch
from torchmetrics.detection import MeanAveragePrecision

class FasterRCNNModule(pl.LightningModule):
    def __init__(self, model, num_train_ap_batches=5, check_every=10):
        super().__init__()
        self.model = model
        self.train_map = MeanAveragePrecision()
        self.num_train_batches = num_train_ap_batches
        self.check_every = check_every
        self.train_iter = None

    def on_train_batch_end(self, outputs, batch, batch_idx):
        # 每check_every个训练批次触发一次AP计算
        if (batch_idx + 1) % self.check_every == 0:
            if self.train_iter is None:
                self.train_iter = iter(self.trainer.datamodule.train_dataloader())
            
            self.train_map.reset()
            for _ in range(self.num_train_batches):
                try:
                    images, targets = next(self.train_iter)
                except StopIteration:
                    self.train_iter = iter(self.trainer.datamodule.train_dataloader())
                    images, targets = next(self.train_iter)
                
                # 转移到当前设备
                images = images.to(self.device)
                targets = [{k: v.to(self.device) for k, v in t.items()} for t in targets]
                
                # 切换模式计算预测
                self.model.eval()
                with torch.no_grad():
                    preds = self.model(images, targets)
                self.model.train()
                
                self.train_map.update(preds, targets)
            
            map_result = self.train_map.compute()
            self.log("train/mAP", map_result["map"], prog_bar=True, logger=True)

    def configure_optimizers(self):
        return torch.optim.SGD(self.model.parameters(), lr=0.001)

关键注意事项

  • 模型模式切换:计算AP前必须切到eval()模式,计算完成后切回train(),避免BatchNorm、Dropout等训练层干扰结果。
  • 迭代器处理:训练集迭代器耗尽时要重新初始化,防止报错。
  • 日志记录:用self.log并设置prog_bar=True,可以在训练进度条实时查看训练AP结果。

内容的提问来源于stack exchange,提问作者Michael D

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.17 16:02:43