From 6bff905c37f84690c486aab55c3aa6ec61ddbbb7 Mon Sep 17 00:00:00 2001 From: 51616 Date: Mon, 23 Dec 2024 14:58:05 +0000 Subject: [PATCH] refactor + fix generated results --- hyperlora/intx_sft.py | 4 ++- hyperlora/training_utils.py | 68 +++++++++++++++++++++---------------- 2 files changed, 41 insertions(+), 31 deletions(-) diff --git a/hyperlora/intx_sft.py b/hyperlora/intx_sft.py index 384baf2..7f5bb25 100644 --- a/hyperlora/intx_sft.py +++ b/hyperlora/intx_sft.py @@ -130,7 +130,8 @@ 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)] - + # 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) @@ -259,6 +260,7 @@ def main(output_dir: str): # https://huggingface.co/blog/packing-with-FA2 # data_collator = DataCollatorForSeq2Seq(tokenizer, model, pad_to_multiple_of=8) + # TODO: have to add truncation here for longer inputs def collator(inp_list, tokenizer): # input is a list of tokenized sequences padding_kwargs = dict(padding=True, pad_to_multiple_of=8, return_tensors="pt") diff --git a/hyperlora/training_utils.py b/hyperlora/training_utils.py index c0f6930..566e415 100644 --- a/hyperlora/training_utils.py +++ b/hyperlora/training_utils.py @@ -1,6 +1,8 @@ -from enum import Enum import json import logging +from dataclasses import fields +from enum import Enum + import numpy as np from transformers import ( GenerationConfig, @@ -30,6 +32,8 @@ def decode_test_result(test_dataset, test_result, tokenizer): input_toks = sample["input_ids"][:start_idx] gen_toks = pred_toks[start_idx:] label_toks = 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) input_text = tokenizer.decode(input_toks, skip_special_tokens=True) gen_text = tokenizer.decode(gen_toks, skip_special_tokens=True) @@ -37,6 +41,21 @@ def decode_test_result(test_dataset, test_result, tokenizer): yield {"input": input_text, "generated": gen_text, "label": label_text} +def eval_generation(eval_trainer, tokenizer, dataset, split, gen_kwargs): + eval_result = eval_trainer.predict( + dataset, + metric_key_prefix=split, + **gen_kwargs, + ) + 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), + split=split, + output_dir=eval_trainer.args.output_dir, + ) + + def train_model( model, tokenizer, @@ -111,19 +130,30 @@ def train_model( trainer.save_model() ############## Evaluation - # TODO: generalize gen_kwargs for validation - gen_kwargs = dict(do_sample=False, max_length=2**13, max_new_tokens=100) + max_new_tokens = 100 + # max_input_len=2**13 # for input truncation + + gen_kwargs = dict(do_sample=False, max_new_tokens=max_new_tokens) # pad_token_id=tokenizer.pad_token_id, # eos_token_id=? + eval_trainer_args = {} + + # Copy only necessary attributes from training_args to eval_trainer_args + seq2seq_training_args_fields = {f.name for f in fields(Seq2SeqTrainingArguments)} + for attr, value in training_args.to_dict().items(): + if attr in seq2seq_training_args_fields: + eval_trainer_args[attr] = value + eval_trainer_args["eval_strategy"] = "no" + eval_trainer_args["overwrite_output_dir"] = True + # NOTE: could also set kv_cache implementation here eval_trainer_args = Seq2SeqTrainingArguments( + **eval_trainer_args, predict_with_generate=True, - generation_max_length=2**13, generation_config=GenerationConfig(**gen_kwargs), ) - 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) @@ -138,36 +168,14 @@ def train_model( eval_trainer = Seq2SeqTrainer( model=model, args=eval_trainer_args, - # train_dataset=train_dataset, - # eval_dataset=val_dataset, - # TODO: use a different collator for test, e.g., more max_len + # TODO: use a different collator for test, e.g., more max_len truncation data_collator=data_collator, compute_metrics=compute_generation_based_metrics, ) if val_dataset is not None: if isinstance(val_dataset, dict): val_dataset = val_dataset["val"] - eval_result = eval_trainer.predict( - val_dataset, - metric_key_prefix="val", - **gen_kwargs, - ) - 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, - ) + eval_generation(eval_trainer, tokenizer, val_dataset, "val", gen_kwargs) if test_dataset is not None: - test_result = eval_trainer.predict(test_dataset, **gen_kwargs) - - 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, - ) + eval_generation(eval_trainer, tokenizer, test_dataset, "test", gen_kwargs)