MCPcopy Create free account
hub / github.com/FusionBrainLab/SONAR-LLM / parse_args

Function parse_args

train_sonarllm.py:40–73  ·  view source on GitHub ↗
()

Source from the content-addressed store, hash-verified

38# 1) Parse command line arguments
39# ------------------------------------------------------------------------
40def parse_args():
41 parser = argparse.ArgumentParser()
42 parser.add_argument("--epochs", type=int, default=2, help="Number of epochs")
43 parser.add_argument("--batch_size", type=int, default=4, help="Per-GPU batch size")
44 parser.add_argument("--lr", type=float, default=5e-4, help="Learning rate")
45 parser.add_argument("--weight_decay", type=float, default=1e-3, help="Weight decay")
46 parser.add_argument("--warmup_steps", type=int, default=350, help="Warmup steps for the LR scheduler")
47 parser.add_argument("--logging_steps", type=int, default=10)
48 parser.add_argument("--save_steps", type=int, default=500)
49 parser.add_argument("--eval_steps", type=int, default=100)
50 parser.add_argument("--local_rank", type=int, default=-1, help="Local rank for DDP")
51 parser.add_argument("--max_train_samples", type=int, default=-1,
52 help="Use fewer samples for quick debugging. -1 for all.")
53 parser.add_argument("--max_val_samples", type=int, default=2048,
54 help="Use fewer samples for quick debugging. -1 for all.")
55 parser.add_argument("--output_dir", type=str, default="./ddp_checkpoints_2")
56
57 parser.add_argument("--use_wandb", action="store_true", help="Enable wandb logging.")
58 parser.add_argument("--wandb_project", type=str, default="my_project", help="Wandb project name.")
59 parser.add_argument("--wandb_run_name", type=str, default="my_ddp_run", help="Wandb run name.")
60
61 parser.add_argument("--use_mixed_precision", action="store_true",
62 help="Use mixed precision training (fp16 autocast).")
63 parser.add_argument("--grad_accum_steps", type=int, default=1,
64 help="Number of gradient accumulation steps.")
65
66 parser.add_argument("--gradient_checkpointing", action="store_true",
67 help="Enable gradient checkpointing (use_cache=False).")
68
69 parser.add_argument("--start_from", type=str, default=None,
70 help="Path to a full checkpoint from which to resume training (including mid-epoch).")
71
72 args = parser.parse_args()
73 return args
74
75# ------------------------------------------------------------------------
76# 2) Dataset classes and data collator

Callers 1

mainFunction · 0.70

Calls

no outgoing calls

Tested by

no test coverage detected