mirror of
https://github.com/SakanaAI/doc-to-lora.git
synced 2026-07-23 17:01:04 +02:00
save generated val/test text
This commit is contained in:
parent
8c71161071
commit
c19e9d37d9
2 changed files with 48 additions and 13 deletions
|
|
@ -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}
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue