manual generation eval

This commit is contained in:
51616 2024-12-24 11:39:43 +00:00
parent ca38148acd
commit e25ca3d95f
2 changed files with 30 additions and 39 deletions

View file

@ -109,36 +109,6 @@ def compute_metrics(eval_pred: EvalPrediction) -> dict:
)
def compute_rouge(pred_texts, label_texts):
out = defaultdict(list)
scorer = rouge_scorer.RougeScorer(["rouge1", "rougeL"], use_stemmer=False)
for pred_text, label_text in zip(pred_texts, label_texts):
scores = scorer.score(pred_text, label_text)
for k, v in scores.items():
out[f"{k}.f1"].append(v.fmeasure)
for k in out:
out[k] = np.mean(out[k])
return out
def compute_generation_based_metrics(
eval_pred: EvalPrediction,
tokenizer: AutoTokenizer,
) -> dict:
pred_toks, labels = eval_pred.predictions, eval_pred.label_ids
start_indices = np.argmax(labels != -100, axis=1)
gen_toks = [x[start_indices[i] :] for i, x in enumerate(pred_toks)]
label_toks = [x[start_indices[i] :] for i, x in enumerate(labels)]
# labels are padded with -100, so we need to replace them with the pad token id
label_toks = [np.where(x == -100, tokenizer.pad_token_id, x) for x in label_toks]
gen_text = tokenizer.batch_decode(gen_toks, skip_special_tokens=True)
label_text = tokenizer.batch_decode(label_toks, skip_special_tokens=True)
rouge = compute_rouge(gen_text, label_text)
return {**rouge}
def main(output_dir: str):
############ Argument parsing
parser = ArgumentParser(
@ -377,7 +347,6 @@ def main(output_dir: str):
partial(train_collator, tokenizer=tokenizer),
partial(generation_collator, tokenizer=tokenizer),
compute_metrics,
partial(compute_generation_based_metrics, tokenizer=tokenizer),
)

View file

@ -1,9 +1,11 @@
from collections import defaultdict
import json
import logging
from dataclasses import fields
from enum import Enum
import numpy as np
from rouge_score import rouge_scorer
from transformers import (
GenerationConfig,
Seq2SeqTrainer,
@ -17,6 +19,18 @@ TRAINING_TASK = Enum("TRAINING_TASK", ["CAUSAL_LM", "COMPLETION"])
logger = logging.getLogger()
def compute_rouge(pred_texts, label_texts):
out = defaultdict(list)
scorer = rouge_scorer.RougeScorer(["rouge1", "rougeL"], use_stemmer=False)
for pred_text, label_text in zip(pred_texts, label_texts):
scores = scorer.score(pred_text, label_text)
for k, v in scores.items():
out[f"{k}.f1"].append(v.fmeasure)
for k in out:
out[k] = np.mean(out[k])
return out
def save_generated_text(samples, output_dir, split):
with open(f"{output_dir}/{split}_generated_text.jsonl", "w") as f:
for sample in samples:
@ -24,24 +38,26 @@ def save_generated_text(samples, output_dir, split):
def decode_test_result(test_dataset, test_result, tokenizer):
out = dict()
out = []
for sample, pred_toks in zip(test_dataset, test_result.predictions):
d = dict()
if "labels" in sample:
start_idx = np.argmax(sample["labels"] != -100)
label_toks = sample["labels"][start_idx:]
# labels are padded with -100, so we need to replace them with the pad token id
label_toks = np.where(label_toks == -100, tokenizer.pad_token_id, label_toks)
label_text = tokenizer.decode(label_toks, skip_special_tokens=True)
out["label"] = label_text
d["label"] = label_text
# HACK: remove the label part
input_toks = sample["input_ids"][:start_idx]
gen_toks = pred_toks[len(input_toks) :]
out["input"] = tokenizer.decode(input_toks, skip_special_tokens=True)
out["generated"] = tokenizer.decode(gen_toks, skip_special_tokens=True)
d["input"] = tokenizer.decode(input_toks, skip_special_tokens=True)
d["generated"] = tokenizer.decode(gen_toks, skip_special_tokens=True)
out.append(d)
yield out
return out
def eval_generation(eval_trainer, tokenizer, dataset, split, gen_kwargs):
@ -50,10 +66,18 @@ def eval_generation(eval_trainer, tokenizer, dataset, split, gen_kwargs):
metric_key_prefix=split,
**gen_kwargs,
)
decoded_txts = decode_test_result(dataset, eval_result, tokenizer)
rouge_metrics = compute_rouge(
[txt["generated"] for txt in decoded_txts],
[txt["label"] for txt in decoded_txts],
)
eval_result.metrics.update(rouge_metrics)
eval_trainer.log_metrics("eval" if split == "val" else split, eval_result.metrics)
eval_trainer.save_metrics("eval" if split == "val" else split, eval_result.metrics)
save_generated_text(
decode_test_result(dataset, eval_result, tokenizer),
decoded_txts,
split=split,
output_dir=eval_trainer.args.output_dir,
)
@ -69,7 +93,6 @@ def train_model(
train_collator=None,
generation_collator=None,
compute_metrics=None,
compute_generation_based_metrics=None,
):
# last_checkpoint = None
@ -176,7 +199,6 @@ def train_model(
# w/ left padding?
# removing label part from input_ids
data_collator=generation_collator,
compute_metrics=compute_generation_based_metrics,
)
if val_dataset is not None:
if isinstance(val_dataset, dict):