(config, args)
| 675 | |
| 676 | |
| 677 | def main(config, args): |
| 678 | |
| 679 | # Process args |
| 680 | debug = args.debug |
| 681 | use_8BitAdam = args.use_8BitAdam |
| 682 | |
| 683 | |
| 684 | # Read Frequently Used config |
| 685 | resume_from_checkpoint = config["resume_from_checkpoint"] |
| 686 | output_folder = config["output_folder"] |
| 687 | experiment_name = config["experiment_name"] |
| 688 | mixed_precision = config["mixed_precision"] |
| 689 | report_to = config["report_to"] |
| 690 | seed = config["seed"] |
| 691 | base_model_path = config["base_model_path"] |
| 692 | pretrained_transformer_path = config["pretrained_transformer_path"] |
| 693 | download_folder_path = config["download_folder_path"] |
| 694 | train_csv_relative_path = config["train_csv_relative_path"] |
| 695 | validation_csv_relative_path = config["validation_csv_relative_path"] |
| 696 | gradient_checkpointing = config["gradient_checkpointing"] |
| 697 | learning_rate = config["learning_rate"] |
| 698 | train_batch_size = config["train_batch_size"] |
| 699 | dataloader_num_workers = config["dataloader_num_workers"] if not debug else 1 # In debug mode, only has 1 worker |
| 700 | gradient_accumulation_steps = config["gradient_accumulation_steps"] |
| 701 | max_train_steps = config["max_train_steps"] |
| 702 | lr_warmup_steps = config["lr_warmup_steps"] |
| 703 | checkpointing_steps = config["checkpointing_steps"] |
| 704 | scale_lr = config["scale_lr"] |
| 705 | checkpoints_total_limit = config["checkpoints_total_limit"] |
| 706 | revision = config["revision"] |
| 707 | lr_scheduler = config["lr_scheduler"] |
| 708 | validation_step = config["validation_step"] |
| 709 | first_iter_validation = config["first_iter_validation"] |
| 710 | num_inference_steps = config["num_inference_steps"] |
| 711 | max_grad_norm = config["max_grad_norm"] |
| 712 | |
| 713 | |
| 714 | # Organize |
| 715 | use_FrameIn = True |
| 716 | |
| 717 | |
| 718 | # Value Check |
| 719 | if max_train_steps is None: |
| 720 | print("max_train_steps must be set") |
| 721 | os.exit(0) |
| 722 | |
| 723 | if torch.backends.mps.is_available() and mixed_precision == "bf16": |
| 724 | # due to pytorch#99272, MPS does not yet support bfloat16. |
| 725 | raise ValueError( |
| 726 | "Mixed precision training with bfloat16 is not supported on MPS. Please use fp16 (recommended) or fp32 instead." |
| 727 | ) |
| 728 | |
| 729 | output_dir = os.path.join(output_folder, experiment_name) |
| 730 | logging_dir = Path(output_dir, config["logging_name"]) |
| 731 | |
| 732 | accelerator_project_config = ProjectConfiguration(project_dir = output_dir, logging_dir = logging_dir) |
| 733 | ddp_kwargs = DistributedDataParallelKwargs(find_unused_parameters=False) # HACK: his find_unused_parameters is said to increase the computation cost |
| 734 | init_kwargs = InitProcessGroupKwargs(backend="nccl", timeout=timedelta(seconds=config["nccl_timeout"])) |
no test coverage detected