From b1bf124c8bae7e08ae79b70a79a737d0b35e426f Mon Sep 17 00:00:00 2001 From: 51616 Date: Tue, 7 Jan 2025 10:46:03 +0000 Subject: [PATCH] fix unseen split + wandb tags --- hyperlora/intx_sft.py | 16 +++++++++++----- hyperlora/training_utils.py | 2 +- 2 files changed, 12 insertions(+), 6 deletions(-) diff --git a/hyperlora/intx_sft.py b/hyperlora/intx_sft.py index 85d9630..e938bf9 100644 --- a/hyperlora/intx_sft.py +++ b/hyperlora/intx_sft.py @@ -296,8 +296,9 @@ def main(output_dir): else: # take some samples from train_ds n_val_samples = data_args.max_val_samples_per_ds - train_ds = train_ds.skip(n_val_samples) val_ds["train_unseen"] = train_ds.take(n_val_samples) + train_ds = train_ds.skip(n_val_samples) + val_train_indices = np.random.permutation(len(train_ds))[:500] val_ds["train"] = train_ds.select(val_train_indices) test_ds = tokenized_ds.get("test", None) @@ -393,10 +394,11 @@ def main(output_dir): # might improve/decrease training speed w/ longer inputs wandb.init( - project=os.environ["WANDB_PROJECT"], + project=os.getenv("WANDB_PROJECT"), name=run_name, group=run_name, config=args, + tags=os.getenv("WANDB_TAGS").split(","), notes=ctx_args.notes, ) @@ -422,15 +424,19 @@ def main(output_dir): if __name__ == "__main__": os.environ["TOKENIZERS_PARALLELISM"] = "true" os.environ["WANDB_PROJECT"] = "ctx_to_lora" - os.environ["WANDB_WATCH"] = "all" + os.environ["WANDB_WATCH"] = "" # "all" os.environ["WANDB_CONSOLE"] = "off" os.environ["PYTORCH_CUDA_ALLOC_CONF"] = "expandable_segments:True" run_name = get_run_name() output_dir = f"train_outputs/{run_name}" - setup_logging(output_dir, debug=os.environ.get("DEBUG", False)) + setup_logging(output_dir, debug=os.getenv("DEBUG", False)) logger.debug(f"CMD: {' '.join(os.sys.argv)}") - save_yaml(extract_cli_args(os.sys.argv), f"{output_dir}/cli_args.yaml") + cli_args = extract_cli_args(os.sys.argv) + save_yaml(cli_args, f"{output_dir}/cli_args.yaml") + if "config" in cli_args: + config_name = os.path.basename(cli_args["config"]).split(".yaml")[0] + os.environ["WANDB_TAGS"] = config_name # disable_caching() main(output_dir) diff --git a/hyperlora/training_utils.py b/hyperlora/training_utils.py index 9e22245..bd7ced8 100644 --- a/hyperlora/training_utils.py +++ b/hyperlora/training_utils.py @@ -182,11 +182,11 @@ def train_model( # is done when load_best_model_at_end=True (our default) train_result = trainer.train(resume_from_checkpoint=checkpoint) trainer.log_metrics("train", train_result.metrics) + trainer.save_model() clear_gpu() metrics = trainer.evaluate(dict(**val_dataset, test=test_dataset)) trainer.log_metrics("eval", metrics) trainer.save_metrics("eval", metrics) - trainer.save_model() clear_gpu()