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

Function main

train_code/train_wan_motion.py:601–1229  ·  view source on GitHub ↗
(config, args)

Source from the content-addressed store, hash-verified

599
600
601def 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,

Callers 1

Calls 15

VideoDataset_MotionClass · 0.90
MixedBatchSamplerClass · 0.90
DiscreteSamplingClass · 0.90
backwardMethod · 0.80
updateMethod · 0.80
zero_moduleFunction · 0.70
filter_kwargsFunction · 0.70
get_optimizerFunction · 0.70
get_sigmasFunction · 0.70

Tested by

no test coverage detected