mirror of
https://github.com/SakanaAI/doc-to-lora.git
synced 2026-07-23 17:01:04 +02:00
fix unseen split + wandb tags
This commit is contained in:
parent
1afdb35150
commit
b1bf124c8b
2 changed files with 12 additions and 6 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue