如何在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
相关产品推荐
相关产品推荐

