trial is only used when we are sweeping hyperparameters.
(args: argparse.Namespace, config: dict, trial: optuna.Trial = None)
| 20 | |
| 21 | |
| 22 | def train(args: argparse.Namespace, config: dict, trial: optuna.Trial = None): |
| 23 | """ |
| 24 | trial is only used when we are sweeping hyperparameters. |
| 25 | """ |
| 26 | |
| 27 | # accelerator |
| 28 | accelerator = Accelerator() |
| 29 | current_gpu = int(torch.cuda.current_device()) |
| 30 | device = accelerator.device # "cuda" |
| 31 | n_gpus = torch.cuda.device_count() |
| 32 | print( |
| 33 | f"---Number of GPUs: {n_gpus}", |
| 34 | "---current_gpu:", |
| 35 | current_gpu, |
| 36 | "---device:", |
| 37 | device, |
| 38 | ) |
| 39 | |
| 40 | # prepare log |
| 41 | logs_folder = config["train"]["logs_folder"] |
| 42 | if accelerator.is_main_process: |
| 43 | writer = SummaryWriter(log_dir=logs_folder) |
| 44 | |
| 45 | print("---Load LLM Model...") |
| 46 | model = load_model(config, args.checkpoint_path, device) |
| 47 | |
| 48 | trainloader, valloader = create_dataloaders( |
| 49 | train_datasets=config["dataset"]["train_dataset_dir"], |
| 50 | validation_datasets=config["dataset"]["valid_dataset_dir"], |
| 51 | batch_size=config["train"]["batch_size"], |
| 52 | device=device, |
| 53 | infinite_train=False, |
| 54 | num_workers=8, |
| 55 | ) |
| 56 | |
| 57 | eff_batch_size = config["train"]["batch_size"] * config["train"]["accumulate_num"] |
| 58 | |
| 59 | total_steps = (config["train"]["n_epochs"] * len(trainloader)) // config["train"][ |
| 60 | "accumulate_num" |
| 61 | ] |
| 62 | print("---total_steps:", total_steps) |
| 63 | |
| 64 | optimizer = torch.optim.AdamW( |
| 65 | model.parameters(), |
| 66 | lr=config["train"]["lr"], |
| 67 | weight_decay=config["train"]["weight_decay"], |
| 68 | ) |
| 69 | scheduler = WarmupDecayLR( |
| 70 | optimizer, |
| 71 | config["train"]["warmup_steps"], |
| 72 | total_steps, |
| 73 | config["train"]["lr_decay"], |
| 74 | ) |
| 75 | |
| 76 | state = { |
| 77 | "model": model.state_dict(), |
| 78 | "optimizer": optimizer.state_dict(), |
| 79 | "scheduler": scheduler.state_dict(), |
no test coverage detected