From b3107a8eabcbf8121c7af424b4511f02e5c7e45a Mon Sep 17 00:00:00 2001 From: 51616 Date: Mon, 23 Dec 2024 10:44:26 +0000 Subject: [PATCH] fix too long generation --- hyperlora/training_utils.py | 24 ++++++++++++------------ 1 file changed, 12 insertions(+), 12 deletions(-) diff --git a/hyperlora/training_utils.py b/hyperlora/training_utils.py index a0f8022..c0f6930 100644 --- a/hyperlora/training_utils.py +++ b/hyperlora/training_utils.py @@ -112,20 +112,16 @@ def train_model( ############## Evaluation + # TODO: generalize gen_kwargs for validation + gen_kwargs = dict(do_sample=False, max_length=2**13, max_new_tokens=100) + # pad_token_id=tokenizer.pad_token_id, + # eos_token_id=? + # NOTE: could also set kv_cache implementation here eval_trainer_args = Seq2SeqTrainingArguments( predict_with_generate=True, generation_max_length=2**13, - generation_config=GenerationConfig( - do_sample=False, - # just a placeholder, will be overridden with `max_new_tokens` - max_length=2**13, - max_new_tokens=100, - # pad_token_id=tokenizer.pad_token_id, - # eos_token_id=? - ), - **training_args.to_dict(), - # compute_metrics=compute_metrics, + generation_config=GenerationConfig(**gen_kwargs), ) eval_trainer_args.eval_strategy = "no" @@ -151,7 +147,11 @@ def train_model( 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") + 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( @@ -161,7 +161,7 @@ def train_model( ) if test_dataset is not None: - test_result = eval_trainer.predict(test_dataset) + 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)