save generated val/test text

This commit is contained in:
51616 2024-12-23 10:31:13 +00:00
parent 8c71161071
commit c19e9d37d9
2 changed files with 48 additions and 13 deletions

View file

@ -131,14 +131,9 @@ def compute_generation_based_metrics(
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)]
# print(pred_toks[indices])
# print(labels[indices])
# breakpoint()
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)
# acc = (pred_toks[indices] == labels[indices]).mean()
# "gen_text": gen_text, "label_text": label_text
return {**rouge}

View file

@ -1,5 +1,6 @@
from enum import Enum
import json
import logging
import numpy as np
from transformers import (
GenerationConfig,
@ -12,6 +13,29 @@ from transformers.trainer_utils import get_last_checkpoint
TRAINING_TASK = Enum("TRAINING_TASK", ["CAUSAL_LM", "COMPLETION"])
logger = logging.getLogger()
def save_generated_text(samples, output_dir, split):
with open(f"{output_dir}/{split}_generated_text.jsonl", "w") as f:
for sample in samples:
f.write(json.dumps(sample) + "\n")
def decode_test_result(test_dataset, test_result, tokenizer):
for sample, pred_toks, labels in zip(
test_dataset, test_result.predictions, test_result.label_ids
):
start_idx = np.argmax(labels != -100, axis=0)
input_toks = sample["input_ids"][:start_idx]
gen_toks = pred_toks[start_idx:]
label_toks = labels[start_idx:]
input_text = tokenizer.decode(input_toks, skip_special_tokens=True)
gen_text = tokenizer.decode(gen_toks, skip_special_tokens=True)
label_text = tokenizer.decode(label_toks, skip_special_tokens=True)
yield {"input": input_text, "generated": gen_text, "label": label_text}
def train_model(
model,
@ -103,6 +127,7 @@ def train_model(
**training_args.to_dict(),
# compute_metrics=compute_metrics,
)
eval_trainer_args.eval_strategy = "no"
# Seq2SeqTrainer is actually just the same as Trainer
# (although it uses a different data collator, i.e., explicit prompt/answer separation)
@ -110,24 +135,39 @@ def train_model(
# allowing us to compute metrics on the generated outputs
# no clue why they call this seq2seq...
logger.info("=" * 80 + "\n" + "Evaluating model..." + "\n" + "=" * 80)
model.eval()
eval_trainer = Seq2SeqTrainer(
model=model,
args=eval_trainer_args,
# train_dataset=train_dataset,
eval_dataset=val_dataset,
# eval_dataset=val_dataset,
# TODO: use a different collator for test, e.g., more max_len
data_collator=data_collator,
compute_metrics=compute_generation_based_metrics,
)
if val_dataset is not None:
eval_result = eval_trainer.evaluate()
eval_trainer.log_metrics("eval", eval_result)
eval_trainer.save_metrics("eval", eval_result)
if isinstance(val_dataset, dict):
val_dataset = val_dataset["val"]
eval_result = eval_trainer.predict(val_dataset, metric_key_prefix="val")
eval_trainer.log_metrics("eval", eval_result.metrics)
eval_trainer.save_metrics("eval", eval_result.metrics)
save_generated_text(
decode_test_result(val_dataset, eval_result, tokenizer),
split="val",
output_dir=training_args.output_dir,
)
if test_dataset is not None:
test_result = eval_trainer.predict(test_dataset)
print(test_result)
# eval_trainer.log_metrics("test", test_result)
# eval_trainer.save_metrics("test", test_result)
eval_trainer.log_metrics("test", test_result.metrics)
eval_trainer.save_metrics("test", test_result.metrics)
save_generated_text(
decode_test_result(test_dataset, test_result, tokenizer),
split="test",
output_dir=training_args.output_dir,
)