mirror of
https://github.com/SakanaAI/doc-to-lora.git
synced 2026-07-23 17:01:04 +02:00
manual generation eval
This commit is contained in:
parent
ca38148acd
commit
e25ca3d95f
2 changed files with 30 additions and 39 deletions
|
|
@ -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),
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue