(config, args)
| 599 | |
| 600 | |
| 601 | def main(config, args): |
| 602 | |
| 603 | # Process args |
| 604 | debug = args.debug |
| 605 | use_8BitAdam = args.use_8BitAdam |
| 606 | |
| 607 | |
| 608 | # Read Frequently Used config |
| 609 | resume_from_checkpoint = config["resume_from_checkpoint"] |
| 610 | output_folder = config["output_folder"] |
| 611 | experiment_name = config["experiment_name"] |
| 612 | mixed_precision = config["mixed_precision"] |
| 613 | report_to = config["report_to"] |
| 614 | seed = config["seed"] |
| 615 | base_model_path = config["base_model_path"] |
| 616 | pretrained_transformer_path = config["pretrained_transformer_path"] |
| 617 | download_folder_path = config["download_folder_path"] |
| 618 | train_csv_relative_path = config["train_csv_relative_path"] |
| 619 | validation_csv_relative_path = config["validation_csv_relative_path"] |
| 620 | gradient_checkpointing = config["gradient_checkpointing"] |
| 621 | learning_rate = config["learning_rate"] |
| 622 | train_batch_size = config["train_batch_size"] |
| 623 | dataloader_num_workers = config["dataloader_num_workers"] if not debug else 1 # In debug mode, only has 1 worker |
| 624 | gradient_accumulation_steps = config["gradient_accumulation_steps"] |
| 625 | max_train_steps = config["max_train_steps"] |
| 626 | lr_warmup_steps = config["lr_warmup_steps"] |
| 627 | checkpointing_steps = config["checkpointing_steps"] |
| 628 | scale_lr = config["scale_lr"] |
| 629 | checkpoints_total_limit = config["checkpoints_total_limit"] |
| 630 | revision = config["revision"] |
| 631 | lr_scheduler = config["lr_scheduler"] |
| 632 | validation_step = config["validation_step"] |
| 633 | first_iter_validation = config["first_iter_validation"] |
| 634 | num_inference_steps = config["num_inference_steps"] |
| 635 | max_grad_norm = config["max_grad_norm"] |
| 636 | |
| 637 | |
| 638 | # Value Check |
| 639 | if max_train_steps is None: |
| 640 | print("max_train_steps must be set") |
| 641 | os.exit(0) |
| 642 | |
| 643 | if torch.backends.mps.is_available() and mixed_precision == "bf16": |
| 644 | # due to pytorch#99272, MPS does not yet support bfloat16. |
| 645 | raise ValueError( |
| 646 | "Mixed precision training with bfloat16 is not supported on MPS. Please use fp16 (recommended) or fp32 instead." |
| 647 | ) |
| 648 | |
| 649 | output_dir = os.path.join(output_folder, experiment_name) |
| 650 | logging_dir = Path(output_dir, config["logging_name"]) |
| 651 | |
| 652 | accelerator_project_config = ProjectConfiguration(project_dir = output_dir, logging_dir = logging_dir) |
| 653 | ddp_kwargs = DistributedDataParallelKwargs(find_unused_parameters=False) # HACK: his find_unused_parameters is said to increase the computation cost |
| 654 | init_kwargs = InitProcessGroupKwargs(backend="nccl", timeout=timedelta(seconds=config["nccl_timeout"])) |
| 655 | accelerator = Accelerator( |
| 656 | gradient_accumulation_steps = gradient_accumulation_steps, |
| 657 | mixed_precision = mixed_precision, |
| 658 | log_with = report_to, |
no test coverage detected