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

Function main

train_code/train_cogvideox_motion_FrameINO.py:549–1218  ·  view source on GitHub ↗
(config, args)

Source from the content-addressed store, hash-verified

547
548
549def 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

Calls 15

MixedBatchSamplerClass · 0.90
img_tensor_to_vae_latentFunction · 0.85
enable_slicingMethod · 0.80
enable_tilingMethod · 0.80
backwardMethod · 0.80
updateMethod · 0.80
get_optimizerFunction · 0.70

Tested by

no test coverage detected