| 38 | # 1) Parse command line arguments |
| 39 | # ------------------------------------------------------------------------ |
| 40 | def 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 |