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

如何在PyTorch Lightning验证步骤中计算mT5模型的F1分数

问题:PyTorch Lightning微调mT5做QA任务,如何基于F1分数选择最优检查点

我正在使用PyTorch Lightning微调mT5模型以完成QA任务,当前我的验证步骤代码如下:

def validation_step(self, batch, batch_idx):
    input_ids = batch['input_ids']
    attention_mask = batch['attention_mask']
    labels = batch['labels']
    loss, outputs = self(input_ids, attention_mask, labels)
    self.log('val_loss', loss, prog_bar=True, logger=True)
    return loss

我希望不再记录val_loss并以此选择最优检查点,而是计算F1分数并基于最佳整体F1选择最优检查点。我已成功在训练循环结束后计算出预测结果:

predictions = []
ground_truths = []
for batch in tqdm(data_module.test_dataloader(), desc="Evaluating"):
    input_ids = batch['input_ids'].to(DEVICE)
    attention_mask = batch['attention_mask'].to(DEVICE)
    labels = batch['labels'].to(DEVICE)
    with torch.no_grad():
        generated_ids = trained_model.model.generate(
            input_ids=input_ids,
            attention_mask=attention_mask,
        )

    predicted_answers = tokenizer.batch_decode(generated_ids, skip_special_tokens=True)
    predictions.extend(predicted_answers)

    mask = labels != -100
    labels = torch.masked_select(labels, mask)
    true_text = tokenizer.decode(labels, skip_special_tokens=True, clean_up_tokenization_spaces=True)
    ground_truths.append(true_text)

请问如何在验证步骤内计算预测值与真实值,从而在该步骤中计算F1分数?


解决方案

要实现基于F1分数选择最优检查点,需分三步修改:在验证步骤收集预测与真实标签、在验证epoch结束时计算整体F1、配置检查点回调跟踪F1分数,具体实现如下:

1. 初始化F1计算工具

推荐使用datasets库的squad指标(适配QA任务的抽取式F1计算),在模型初始化时加载:

from datasets import load_metric

class MT5QA(pl.LightningModule):
    def __init__(self, model_name, tokenizer, ...):
        super().__init__()
        self.model = T5ForConditionalGeneration.from_pretrained(model_name)
        self.tokenizer = tokenizer
        # 加载QA专用F1指标
        self.metric = load_metric("squad")
        # 其他初始化逻辑...

2. 修改validation_step收集预测与真实值

在验证步骤中生成预测,解码并逐个处理样本的真实标签(避免批量mask导致标签拼接错误):

def validation_step(self, batch, batch_idx):
    input_ids = batch['input_ids']
    attention_mask = batch['attention_mask']
    labels = batch['labels']
    
    # 可选:保留loss计算(仅用于日志,不影响检查点选择)
    loss, outputs = self(input_ids, attention_mask, labels)
    
    # 生成预测结果(可根据任务需求添加max_length、num_beams等参数)
    generated_ids = self.model.generate(
        input_ids=input_ids,
        attention_mask=attention_mask,
        max_length=128,
        num_beams=4,
        early_stopping=True
    )
    
    # 解码预测文本
    predicted_answers = self.tokenizer.batch_decode(generated_ids, skip_special_tokens=True)
    
    # 解码真实标签:逐个过滤-100的padding值
    ground_truths = []
    for label in labels:
        mask = label != -100
        cleaned_label = label[mask]
        true_text = self.tokenizer.decode(
            cleaned_label, 
            skip_special_tokens=True, 
            clean_up_tokenization_spaces=True
        )
        ground_truths.append(true_text)
    
    # 记录loss(可选)并返回预测与真实值
    self.log('val_loss', loss, prog_bar=True, logger=True)
    return {"predictions": predicted_answers, "ground_truths": ground_truths}

3. 在validation_epoch_end计算并日志F1分数

合并所有batch的结果,计算整体F1并记录到日志:

def validation_epoch_end(self, outputs):
    # 合并所有batch的预测和真实标签
    all_predictions = []
    all_ground_truths = []
    for output in outputs:
        all_predictions.extend(output["predictions"])
        all_ground_truths.extend(output["ground_truths"])
    
    # 适配squad指标的输入格式
    predictions = [
        {"prediction_text": pred, "id": str(i)} 
        for i, pred in enumerate(all_predictions)
    ]
    references = [
        {"answers": {"text": [gt], "answer_start": [0]}, "id": str(i)} 
        for i, gt in enumerate(all_ground_truths)
    ]
    
    # 计算F1分数
    results = self.metric.compute(predictions=predictions, references=references)
    val_f1 = results["f1"]
    
    # 日志F1分数,sync_dist用于多GPU训练场景
    self.log('val_f1', val_f1, prog_bar=True, logger=True, sync_dist=True)

4. 配置检查点回调跟踪F1

初始化ModelCheckpoint回调,设置基于val_f1保存最优模型:

from pytorch_lightning.callbacks import ModelCheckpoint

# 初始化检查点回调
checkpoint_callback = ModelCheckpoint(
    monitor='val_f1',          # 跟踪的核心指标
    mode='max',                # F1分数越高越好
    save_top_k=1,              # 保存排名前1的模型
    dirpath='./checkpoints/',  # 模型保存路径
    filename='best-mt5-qa-{val_f1:.2f}'  # 文件名格式
)

# 初始化Trainer并传入回调
trainer = pl.Trainer(
    max_epochs=10,
    accelerator='gpu',
    devices=1,
    callbacks=[checkpoint_callback]
)

# 启动训练
trainer.fit(model, datamodule=data_module)

注意事项

  • 若为生成式QA任务,可替换为rouge指标或自定义分词级F1逻辑(基于分词后的交集计算)。
  • generate参数需与训练时的生成逻辑一致,保证预测结果合理性。
  • 多GPU训练时必须设置sync_dist=True,确保指标在各设备间同步。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.13 03:35:42