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)