MCPcopy Create free account
hub / github.com/UVA-Computer-Vision-Lab/FrameINO / main

Function main

train_code/train_wan_motion_FrameINO.py:677–1333  ·  view source on GitHub ↗
(config, args)

Source from the content-addressed store, hash-verified

675
676
677def 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"]))

Callers 1

Calls 15

MixedBatchSamplerClass · 0.90
DiscreteSamplingClass · 0.90
ID_tensor_to_vae_latentFunction · 0.85
backwardMethod · 0.80
updateMethod · 0.80
filter_kwargsFunction · 0.70
get_optimizerFunction · 0.70
get_sigmasFunction · 0.70

Tested by

no test coverage detected