| 2 | |
| 3 | |
| 4 | def parse_args(): |
| 5 | parser = argparse.ArgumentParser(description="Training parameters") |
| 6 | parser.add_argument("--dresses_dataset_base_path", type=str, required=True, help="Base path of the dresses dataset.") |
| 7 | parser.add_argument("--dresses_dataset_metadata_path", type=str, required=True, help="Path to the metadata file of the dresses dataset.") |
| 8 | parser.add_argument("--lower_dataset_base_path", type=str, required=True, help="Base path of the lower body dataset.") |
| 9 | parser.add_argument("--lower_dataset_metadata_path", type=str, required=True, help="Path to the metadata file of the lower body dataset.") |
| 10 | parser.add_argument("--upper_dataset_base_path", type=str, required=True, help="Base path of the upper body dataset.") |
| 11 | parser.add_argument("--upper_dataset_metadata_path", type=str, required=True, help="Path to the metadata file of the upper body dataset.") |
| 12 | parser.add_argument("--height", type=int, required=True, help="Height of images and videos.") |
| 13 | parser.add_argument("--width", type=int, required=True, help="Width of images and videos.") |
| 14 | parser.add_argument("--num_frames", type=int, default=49, help="Number of frames per video. Frames are sampled from the video prefix.") |
| 15 | |
| 16 | parser.add_argument("--vae_model_path", type=str, required=True, help="Path of VAE model.") |
| 17 | parser.add_argument("--text_encoder_model_path", type=str, required=True, help="Path of Text Encoder model.") |
| 18 | parser.add_argument("--dit_model_path", type=str, nargs='+', required=True, help="Paths of DIT model.") |
| 19 | parser.add_argument("--tokenizer_path", type=str, required=True, help="Path of Tokenizer model.") |
| 20 | |
| 21 | parser.add_argument("--lora_base_model", type=str, choices=("dit", "vace"), default="vace", help="Which model LoRA is added to.") |
| 22 | parser.add_argument("--lora_target_modules", type=str, default="q,k,v,o,ffn.0,ffn.2", help="Which layers LoRA is added to.") |
| 23 | parser.add_argument("--lora_rank", type=int, default=32, help="Rank of LoRA.") |
| 24 | |
| 25 | parser.add_argument("--max_timestep_boundary", type=float, default=1.0, help="Max timestep boundary (for mixed models, e.g., Wan-AI/Wan2.2-I2V-A14B).") |
| 26 | parser.add_argument("--min_timestep_boundary", type=float, default=0.0, help="Min timestep boundary (for mixed models, e.g., Wan-AI/Wan2.2-I2V-A14B).") |
| 27 | |
| 28 | parser.add_argument("--output_path", type=str, default="./models", help="Output save path.") |
| 29 | parser.add_argument("--remove_prefix_in_ckpt", type=str, default=None, help="Prefix to remove from saved LoRA keys. Defaults to the selected LoRA base model prefix.") |
| 30 | |
| 31 | parser.add_argument("--learning_rate", type=float, default=1e-4, help="Learning rate.") |
| 32 | parser.add_argument("--weight_decay", type=float, default=0.01, help="Weight decay.") |
| 33 | parser.add_argument("--dataset_num_workers", type=int, default=0, help="Number of workers for data loading.") |
| 34 | parser.add_argument("--save_steps", type=int, default=None, help="Number of checkpoint saving invervals. If None, checkpoints will be saved every epoch.") |
| 35 | parser.add_argument("--num_epochs", type=int, default=1, help="Number of epochs.") |
| 36 | parser.add_argument("--gradient_accumulation_steps", type=int, default=1, help="Gradient accumulation steps.") |
| 37 | parser.add_argument("--find_unused_parameters", default=False, action="store_true", help="Whether to find unused parameters in DDP.") |
| 38 | |
| 39 | return parser |