(config, args)
| 547 | |
| 548 | |
| 549 | def main(config, args): |
| 550 | |
| 551 | # Process args |
| 552 | debug = args.debug |
| 553 | use_8BitAdam = args.use_8BitAdam |
| 554 | |
| 555 | |
| 556 | # Read Frequently Used config |
| 557 | resume_from_checkpoint = config["resume_from_checkpoint"] |
| 558 | output_folder = config["output_folder"] |
| 559 | experiment_name = config["experiment_name"] |
| 560 | mixed_precision = config["mixed_precision"] |
| 561 | report_to = config["report_to"] |
| 562 | seed = config["seed"] |
| 563 | base_model_path = config["base_model_path"] |
| 564 | pretrained_transformer_path = config["pretrained_transformer_path"] |
| 565 | download_folder_path = config["download_folder_path"] |
| 566 | train_csv_relative_path = config["train_csv_relative_path"] |
| 567 | validation_csv_relative_path = config["validation_csv_relative_path"] |
| 568 | gradient_checkpointing = config["gradient_checkpointing"] |
| 569 | enable_slicing = config["enable_slicing"] |
| 570 | enable_tiling = config["enable_tiling"] |
| 571 | learning_rate = config["learning_rate"] |
| 572 | train_batch_size = config["train_batch_size"] |
| 573 | dataloader_num_workers = config["dataloader_num_workers"] if not debug else 1 # In debug mode, only has 1 worker |
| 574 | gradient_accumulation_steps = config["gradient_accumulation_steps"] |
| 575 | max_train_steps = config["max_train_steps"] |
| 576 | lr_warmup_steps = config["lr_warmup_steps"] |
| 577 | checkpointing_steps = config["checkpointing_steps"] |
| 578 | scale_lr = config["scale_lr"] |
| 579 | checkpoints_total_limit = config["checkpoints_total_limit"] |
| 580 | revision = config["revision"] |
| 581 | variant = config["variant"] |
| 582 | lr_scheduler = config["lr_scheduler"] |
| 583 | use_rotary_positional_embeddings = config["use_rotary_positional_embeddings"], |
| 584 | use_learned_positional_embeddings = config["use_learned_positional_embeddings"] |
| 585 | validation_step = config["validation_step"] |
| 586 | first_iter_validation = config["first_iter_validation"] |
| 587 | num_inference_steps = config["num_inference_steps"] |
| 588 | |
| 589 | # Organize |
| 590 | use_FrameIn = True |
| 591 | |
| 592 | |
| 593 | # Value Check |
| 594 | if max_train_steps is None: |
| 595 | print("max_train_steps must be set") |
| 596 | os.exit(0) |
| 597 | |
| 598 | if torch.backends.mps.is_available() and mixed_precision == "bf16": |
| 599 | # due to pytorch#99272, MPS does not yet support bfloat16. |
| 600 | raise ValueError( |
| 601 | "Mixed precision training with bfloat16 is not supported on MPS. Please use fp16 (recommended) or fp32 instead." |
| 602 | ) |
| 603 | |
| 604 | output_dir = os.path.join(output_folder, experiment_name) |
| 605 | logging_dir = Path(output_dir, config["logging_name"]) |
| 606 |
no test coverage detected