fix unseen split + wandb tags

This commit is contained in:
51616 2025-01-07 10:46:03 +00:00
parent 1afdb35150
commit b1bf124c8b
2 changed files with 12 additions and 6 deletions

View file

@ -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)

View file

@ -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()