diff --git a/scripts/cogvideox_fun/train.py b/scripts/cogvideox_fun/train.py index 32d8538..7401b61 100755 --- a/scripts/cogvideox_fun/train.py +++ b/scripts/cogvideox_fun/train.py @@ -72,7 +72,7 @@ from videox_fun.models import (AutoencoderKLCogVideoX, from videox_fun.pipeline import (CogVideoXFunPipeline, CogVideoXFunControlPipeline, CogVideoXFunInpaintPipeline) -from videox_fun.pipeline.pipeline_CogVideoXFuninpaint import ( +from videox_fun.pipeline.pipeline_cogvideox_fun_inpaint import ( add_noise_to_reference_video, get_3d_rotary_pos_embed, get_resize_crop_region_for_grid) from videox_fun.utils.discrete_sampler import DiscreteSampling @@ -1244,9 +1244,10 @@ def main(): # The trackers initializes automatically on the main process. if accelerator.is_main_process: tracker_config = dict(vars(args)) - tracker_config.pop("validation_prompts") - tracker_config.pop("trainable_modules") - tracker_config.pop("trainable_modules_low_learning_rate") + keys_to_pop = [k for k, v in tracker_config.items() if isinstance(v, list)] + for k in keys_to_pop: + tracker_config.pop(k) + print(f"Removed tracker_config['{k}']") accelerator.init_trackers(args.tracker_project_name, tracker_config) # Function for unwrapping if model was compiled with `torch.compile`. @@ -1366,7 +1367,7 @@ def main(): mask_pixel_values = batch["mask_pixel_values"].to(weight_dtype) mask = batch["mask"].to(weight_dtype) # Increase the batch size when the length of the latent sequence of the current sample is small - if args.training_with_video_token_length: + if args.auto_tile_batch_size and args.training_with_video_token_length: if args.video_sample_n_frames * args.token_sample_size * args.token_sample_size // 16 >= pixel_values.size()[1] * pixel_values.size()[3] * pixel_values.size()[4]: mask_pixel_values = torch.tile(mask_pixel_values, (4, 1, 1, 1, 1)) mask = torch.tile(mask, (4, 1, 1, 1, 1)) diff --git a/scripts/cogvideox_fun/train_control.py b/scripts/cogvideox_fun/train_control.py index 0405db2..a446138 100755 --- a/scripts/cogvideox_fun/train_control.py +++ b/scripts/cogvideox_fun/train_control.py @@ -70,7 +70,7 @@ from videox_fun.models import (AutoencoderKLCogVideoX, from videox_fun.pipeline import (CogVideoXFunPipeline, CogVideoXFunControlPipeline, CogVideoXFunInpaintPipeline) -from videox_fun.pipeline.pipeline_CogVideoXFuninpaint import ( +from videox_fun.pipeline.pipeline_cogvideox_fun_inpaint import ( add_noise_to_reference_video, get_3d_rotary_pos_embed, get_resize_crop_region_for_grid) from videox_fun.utils.discrete_sampler import DiscreteSampling @@ -1179,10 +1179,10 @@ def main(): # The trackers initializes automatically on the main process. if accelerator.is_main_process: tracker_config = dict(vars(args)) - tracker_config.pop("validation_prompts") - tracker_config.pop("validation_paths") - tracker_config.pop("trainable_modules") - tracker_config.pop("trainable_modules_low_learning_rate") + keys_to_pop = [k for k, v in tracker_config.items() if isinstance(v, list)] + for k in keys_to_pop: + tracker_config.pop(k) + print(f"Removed tracker_config['{k}']") accelerator.init_trackers(args.tracker_project_name, tracker_config) # Function for unwrapping if model was compiled with `torch.compile`. diff --git a/scripts/cogvideox_fun/train_lora.py b/scripts/cogvideox_fun/train_lora.py index 22862b2..7a43b22 100755 --- a/scripts/cogvideox_fun/train_lora.py +++ b/scripts/cogvideox_fun/train_lora.py @@ -71,7 +71,7 @@ from videox_fun.models import (AutoencoderKLCogVideoX, from videox_fun.pipeline import (CogVideoXFunPipeline, CogVideoXFunControlPipeline, CogVideoXFunInpaintPipeline) -from videox_fun.pipeline.pipeline_CogVideoXFuninpaint import ( +from videox_fun.pipeline.pipeline_cogvideox_fun_inpaint import ( add_noise_to_reference_video, get_3d_rotary_pos_embed, get_resize_crop_region_for_grid) from videox_fun.utils.discrete_sampler import DiscreteSampling @@ -1180,7 +1180,10 @@ def main(): # The trackers initializes automatically on the main process. if accelerator.is_main_process: tracker_config = dict(vars(args)) - tracker_config.pop("validation_prompts") + keys_to_pop = [k for k, v in tracker_config.items() if isinstance(v, list)] + for k in keys_to_pop: + tracker_config.pop(k) + print(f"Removed tracker_config['{k}']") accelerator.init_trackers(args.tracker_project_name, tracker_config) # Function for unwrapping if model was compiled with `torch.compile`. @@ -1361,7 +1364,7 @@ def main(): mask_pixel_values = batch["mask_pixel_values"].to(weight_dtype) mask = batch["mask"].to(weight_dtype) # Increase the batch size when the length of the latent sequence of the current sample is small - if args.training_with_video_token_length: + if args.auto_tile_batch_size and args.training_with_video_token_length: if args.video_sample_n_frames * args.token_sample_size * args.token_sample_size // 16 >= pixel_values.size()[1] * pixel_values.size()[3] * pixel_values.size()[4]: mask_pixel_values = torch.tile(mask_pixel_values, (4, 1, 1, 1, 1)) mask = torch.tile(mask, (4, 1, 1, 1, 1)) diff --git a/scripts/cogvideox_fun/train_reward_lora.py b/scripts/cogvideox_fun/train_reward_lora.py index 31bf7ce..81e78fa 100755 --- a/scripts/cogvideox_fun/train_reward_lora.py +++ b/scripts/cogvideox_fun/train_reward_lora.py @@ -1049,7 +1049,10 @@ def main(): # The trackers initializes automatically on the main process. if accelerator.is_main_process: tracker_config = dict(vars(args)) - tracker_config.pop("validation_prompts") + keys_to_pop = [k for k, v in tracker_config.items() if isinstance(v, list)] + for k in keys_to_pop: + tracker_config.pop(k) + print(f"Removed tracker_config['{k}']") accelerator.init_trackers(args.tracker_project_name, tracker_config) # Function for unwrapping if model was compiled with `torch.compile`. diff --git a/scripts/flux/train.py b/scripts/flux/train.py index a3ed21e..3b0be2c 100644 --- a/scripts/flux/train.py +++ b/scripts/flux/train.py @@ -1362,10 +1362,10 @@ def main(): # The trackers initializes automatically on the main process. if accelerator.is_main_process: tracker_config = dict(vars(args)) - tracker_config.pop("validation_prompts") - tracker_config.pop("trainable_modules") - tracker_config.pop("trainable_modules_low_learning_rate") - tracker_config.pop("fix_sample_size") + keys_to_pop = [k for k, v in tracker_config.items() if isinstance(v, list)] + for k in keys_to_pop: + tracker_config.pop(k) + print(f"Removed tracker_config['{k}']") accelerator.init_trackers(args.tracker_project_name, tracker_config) # Function for unwrapping if model was compiled with `torch.compile`. diff --git a/scripts/flux/train_lora.py b/scripts/flux/train_lora.py index 00fa6f8..babf2d4 100644 --- a/scripts/flux/train_lora.py +++ b/scripts/flux/train_lora.py @@ -1297,8 +1297,10 @@ def main(): # The trackers initializes automatically on the main process. if accelerator.is_main_process: tracker_config = dict(vars(args)) - tracker_config.pop("validation_prompts") - tracker_config.pop("fix_sample_size") + keys_to_pop = [k for k, v in tracker_config.items() if isinstance(v, list)] + for k in keys_to_pop: + tracker_config.pop(k) + print(f"Removed tracker_config['{k}']") accelerator.init_trackers(args.tracker_project_name, tracker_config) # Function for unwrapping if model was compiled with `torch.compile`. diff --git a/scripts/qwenimage/train.py b/scripts/qwenimage/train.py index d809240..d9e97f5 100644 --- a/scripts/qwenimage/train.py +++ b/scripts/qwenimage/train.py @@ -1228,10 +1228,10 @@ def main(): # The trackers initializes automatically on the main process. if accelerator.is_main_process: tracker_config = dict(vars(args)) - tracker_config.pop("validation_prompts") - tracker_config.pop("trainable_modules") - tracker_config.pop("trainable_modules_low_learning_rate") - tracker_config.pop("fix_sample_size") + keys_to_pop = [k for k, v in tracker_config.items() if isinstance(v, list)] + for k in keys_to_pop: + tracker_config.pop(k) + print(f"Removed tracker_config['{k}']") accelerator.init_trackers(args.tracker_project_name, tracker_config) # Function for unwrapping if model was compiled with `torch.compile`. @@ -1472,7 +1472,7 @@ def main(): masked_loss = masked_loss * weighting final_loss = masked_loss.mean() return final_loss - + weighting = compute_loss_weighting_for_sd3(weighting_scheme=args.weighting_scheme, sigmas=sigmas) loss = custom_mse_loss(noise_pred.float(), target.float(), weighting.float()) loss = loss.mean() diff --git a/scripts/qwenimage/train_edit.py b/scripts/qwenimage/train_edit.py index b30f540..6364cd5 100644 --- a/scripts/qwenimage/train_edit.py +++ b/scripts/qwenimage/train_edit.py @@ -76,11 +76,11 @@ from videox_fun.data.dataset_image import ImageEditDataset from videox_fun.models import (AutoencoderKLQwenImage, Qwen2VLProcessor, Qwen2_5_VLForConditionalGeneration, Qwen2Tokenizer, QwenImageTransformer2DModel) -from videox_fun.pipeline import QwenImageEditPipeline +from videox_fun.pipeline import QwenImageEditPipeline, QwenImageEditPlusPipeline from videox_fun.pipeline.pipeline_qwenimage_edit import PREFERRED_QWENIMAGE_RESOLUTIONS, calculate_dimensions from videox_fun.pipeline.pipeline_qwenimage_edit_plus import CONDITION_IMAGE_SIZE, VAE_IMAGE_SIZE from videox_fun.utils.discrete_sampler import DiscreteSampling -from videox_fun.utils.utils import get_image_to_video_latent, save_videos_grid +from videox_fun.utils.utils import get_image_to_video_latent, save_videos_grid, get_image if is_wandb_available(): import wandb @@ -150,13 +150,22 @@ def log_validation(vae, text_encoder, tokenizer, transformer3d, args, accelerato subfolder="scheduler" ) transformer3d = transformer3d.to("cpu") - pipeline = QwenImageEditPipeline( - vae=accelerator.unwrap_model(vae).to(weight_dtype), - text_encoder=accelerator.unwrap_model(text_encoder), - tokenizer=tokenizer, - transformer=transformer3d_val, - scheduler=scheduler, - ) + if args.train_mode == "qwen_image_edit": + pipeline = QwenImageEditPipeline( + vae=accelerator.unwrap_model(vae).to(weight_dtype), + text_encoder=accelerator.unwrap_model(text_encoder), + tokenizer=tokenizer, + transformer=transformer3d_val, + scheduler=scheduler, + ) + else: + pipeline = QwenImageEditPlusPipeline( + vae=accelerator.unwrap_model(vae).to(weight_dtype), + text_encoder=accelerator.unwrap_model(text_encoder), + tokenizer=tokenizer, + transformer=transformer3d_val, + scheduler=scheduler, + ) pipeline = pipeline.to(accelerator.device) if args.seed is None: @@ -166,12 +175,17 @@ def log_validation(vae, text_encoder, tokenizer, transformer3d, args, accelerato for i in range(len(args.validation_prompts)): with torch.no_grad(): + if args.train_mode == "qwen_image_edit": + image = get_image(args.validation_image_paths[i]) + else: + image = [get_image(args.validation_image_paths[i])] sample = pipeline( args.validation_prompts[i], negative_prompt = "bad detailed", height = args.image_sample_size, width = args.image_sample_size, - generator = generator + generator = generator, + image = image ).images os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) image = sample[0].save(os.path.join(args.output_dir, f"sample/sample-{global_step}-image-{i}.gif")) @@ -246,6 +260,13 @@ def parse_args(): nargs="+", help=("A set of prompts evaluated every `--validation_epochs` and logged to `--report_to`."), ) + parser.add_argument( + "--validation_image_paths", + type=str, + default=None, + nargs="+", + help=("A set of images evaluated every `--validation_epochs` and logged to `--report_to`."), + ) parser.add_argument( "--output_dir", type=str, @@ -1250,10 +1271,10 @@ def main(): # The trackers initializes automatically on the main process. if accelerator.is_main_process: tracker_config = dict(vars(args)) - tracker_config.pop("validation_prompts") - tracker_config.pop("trainable_modules") - tracker_config.pop("trainable_modules_low_learning_rate") - tracker_config.pop("fix_sample_size") + keys_to_pop = [k for k, v in tracker_config.items() if isinstance(v, list)] + for k in keys_to_pop: + tracker_config.pop(k) + print(f"Removed tracker_config['{k}']") accelerator.init_trackers(args.tracker_project_name, tracker_config) # Function for unwrapping if model was compiled with `torch.compile`. diff --git a/scripts/qwenimage/train_edit_lora.py b/scripts/qwenimage/train_edit_lora.py index 5aa9c86..cc77480 100644 --- a/scripts/qwenimage/train_edit_lora.py +++ b/scripts/qwenimage/train_edit_lora.py @@ -77,7 +77,7 @@ from videox_fun.models import (AutoencoderKLQwenImage, AutoencoderKLWan, Qwen2_5_VLForConditionalGeneration, Qwen2Tokenizer, Qwen2VLProcessor, QwenImageTransformer2DModel) -from videox_fun.pipeline import QwenImageEditPipeline, QwenImagePipeline +from videox_fun.pipeline import QwenImageEditPipeline, QwenImageEditPlusPipeline from videox_fun.pipeline.pipeline_qwenimage_edit import ( PREFERRED_QWENIMAGE_RESOLUTIONS, calculate_dimensions) from videox_fun.pipeline.pipeline_qwenimage_edit_plus import ( @@ -85,7 +85,7 @@ from videox_fun.pipeline.pipeline_qwenimage_edit_plus import ( from videox_fun.utils.discrete_sampler import DiscreteSampling from videox_fun.utils.lora_utils import (create_network, merge_lora, unmerge_lora) -from videox_fun.utils.utils import get_image_to_video_latent, save_videos_grid +from videox_fun.utils.utils import get_image_to_video_latent, save_videos_grid, get_image if is_wandb_available(): import wandb @@ -155,13 +155,22 @@ def log_validation(vae, text_encoder, tokenizer, transformer3d, network, args, a subfolder="scheduler" ) transformer3d = transformer3d.to("cpu") - pipeline = QwenImageEditPipeline( - vae=accelerator.unwrap_model(vae).to(weight_dtype), - text_encoder=accelerator.unwrap_model(text_encoder), - tokenizer=tokenizer, - transformer=transformer3d_val, - scheduler=scheduler, - ) + if args.train_mode == "qwen_image_edit": + pipeline = QwenImageEditPipeline( + vae=accelerator.unwrap_model(vae).to(weight_dtype), + text_encoder=accelerator.unwrap_model(text_encoder), + tokenizer=tokenizer, + transformer=transformer3d_val, + scheduler=scheduler, + ) + else: + pipeline = QwenImageEditPlusPipeline( + vae=accelerator.unwrap_model(vae).to(weight_dtype), + text_encoder=accelerator.unwrap_model(text_encoder), + tokenizer=tokenizer, + transformer=transformer3d_val, + scheduler=scheduler, + ) pipeline = pipeline.to(accelerator.device) pipeline = merge_lora( @@ -175,12 +184,17 @@ def log_validation(vae, text_encoder, tokenizer, transformer3d, network, args, a for i in range(len(args.validation_prompts)): with torch.no_grad(): + if args.train_mode == "qwen_image_edit": + image = get_image(args.validation_image_paths[i]) + else: + image = [get_image(args.validation_image_paths[i])] sample = pipeline( args.validation_prompts[i], negative_prompt = "bad detailed", height = args.image_sample_size, width = args.image_sample_size, - generator = generator + generator = generator, + image = image ).images os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) image = sample[0].save(os.path.join(args.output_dir, f"sample/sample-{global_step}-image-{i}.gif")) @@ -255,6 +269,13 @@ def parse_args(): nargs="+", help=("A set of prompts evaluated every `--validation_epochs` and logged to `--report_to`."), ) + parser.add_argument( + "--validation_image_paths", + type=str, + default=None, + nargs="+", + help=("A set of images evaluated every `--validation_epochs` and logged to `--report_to`."), + ) parser.add_argument( "--output_dir", type=str, @@ -1202,8 +1223,10 @@ def main(): # The trackers initializes automatically on the main process. if accelerator.is_main_process: tracker_config = dict(vars(args)) - tracker_config.pop("validation_prompts") - tracker_config.pop("fix_sample_size") + keys_to_pop = [k for k, v in tracker_config.items() if isinstance(v, list)] + for k in keys_to_pop: + tracker_config.pop(k) + print(f"Removed tracker_config['{k}']") accelerator.init_trackers(args.tracker_project_name, tracker_config) # Function for unwrapping if model was compiled with `torch.compile`. diff --git a/scripts/qwenimage/train_lora.py b/scripts/qwenimage/train_lora.py index f834569..b3dfc5f 100644 --- a/scripts/qwenimage/train_lora.py +++ b/scripts/qwenimage/train_lora.py @@ -1173,8 +1173,10 @@ def main(): # The trackers initializes automatically on the main process. if accelerator.is_main_process: tracker_config = dict(vars(args)) - tracker_config.pop("validation_prompts") - tracker_config.pop("fix_sample_size") + keys_to_pop = [k for k, v in tracker_config.items() if isinstance(v, list)] + for k in keys_to_pop: + tracker_config.pop(k) + print(f"Removed tracker_config['{k}']") accelerator.init_trackers(args.tracker_project_name, tracker_config) # Function for unwrapping if model was compiled with `torch.compile`. diff --git a/scripts/wan2.1/train.py b/scripts/wan2.1/train.py index 1a1b11b..2d40717 100755 --- a/scripts/wan2.1/train.py +++ b/scripts/wan2.1/train.py @@ -1420,10 +1420,10 @@ def main(): # The trackers initializes automatically on the main process. if accelerator.is_main_process: tracker_config = dict(vars(args)) - tracker_config.pop("validation_prompts") - tracker_config.pop("trainable_modules") - tracker_config.pop("trainable_modules_low_learning_rate") - tracker_config.pop("fix_sample_size") + keys_to_pop = [k for k, v in tracker_config.items() if isinstance(v, list)] + for k in keys_to_pop: + tracker_config.pop(k) + print(f"Removed tracker_config['{k}']") accelerator.init_trackers(args.tracker_project_name, tracker_config) # Function for unwrapping if model was compiled with `torch.compile`. diff --git a/scripts/wan2.1/train_lora.py b/scripts/wan2.1/train_lora.py index 21a3e37..fa4092a 100755 --- a/scripts/wan2.1/train_lora.py +++ b/scripts/wan2.1/train_lora.py @@ -1355,8 +1355,10 @@ def main(): # The trackers initializes automatically on the main process. if accelerator.is_main_process: tracker_config = dict(vars(args)) - tracker_config.pop("validation_prompts") - tracker_config.pop("fix_sample_size") + keys_to_pop = [k for k, v in tracker_config.items() if isinstance(v, list)] + for k in keys_to_pop: + tracker_config.pop(k) + print(f"Removed tracker_config['{k}']") accelerator.init_trackers(args.tracker_project_name, tracker_config) # Function for unwrapping if model was compiled with `torch.compile`. diff --git a/scripts/wan2.1/train_reward_lora.py b/scripts/wan2.1/train_reward_lora.py index 7cca334..7671267 100755 --- a/scripts/wan2.1/train_reward_lora.py +++ b/scripts/wan2.1/train_reward_lora.py @@ -1054,8 +1054,10 @@ def main(): # The trackers initializes automatically on the main process. if accelerator.is_main_process: tracker_config = dict(vars(args)) - tracker_config.pop("validation_prompts") - tracker_config.pop("backprop_step_list", None) + keys_to_pop = [k for k, v in tracker_config.items() if isinstance(v, list)] + for k in keys_to_pop: + tracker_config.pop(k) + print(f"Removed tracker_config['{k}']") accelerator.init_trackers(args.tracker_project_name, tracker_config) # Train! diff --git a/scripts/wan2.1_fun/train.py b/scripts/wan2.1_fun/train.py index 3e9a72a..8fd8e4a 100755 --- a/scripts/wan2.1_fun/train.py +++ b/scripts/wan2.1_fun/train.py @@ -1417,10 +1417,10 @@ def main(): # The trackers initializes automatically on the main process. if accelerator.is_main_process: tracker_config = dict(vars(args)) - tracker_config.pop("validation_prompts") - tracker_config.pop("trainable_modules") - tracker_config.pop("trainable_modules_low_learning_rate") - tracker_config.pop("fix_sample_size") + keys_to_pop = [k for k, v in tracker_config.items() if isinstance(v, list)] + for k in keys_to_pop: + tracker_config.pop(k) + print(f"Removed tracker_config['{k}']") accelerator.init_trackers(args.tracker_project_name, tracker_config) # Function for unwrapping if model was compiled with `torch.compile`. diff --git a/scripts/wan2.1_fun/train_control.py b/scripts/wan2.1_fun/train_control.py index d08227f..8d627fc 100755 --- a/scripts/wan2.1_fun/train_control.py +++ b/scripts/wan2.1_fun/train_control.py @@ -230,6 +230,13 @@ def parse_args(): nargs="+", help=("A set of prompts evaluated every `--validation_epochs` and logged to `--report_to`."), ) + parser.add_argument( + "--validation_paths", + type=str, + default=None, + nargs="+", + help=("A set of control videos evaluated every `--validation_epochs` and logged to `--report_to`."), + ) parser.add_argument( "--output_dir", type=str, @@ -1415,10 +1422,10 @@ def main(): # The trackers initializes automatically on the main process. if accelerator.is_main_process: tracker_config = dict(vars(args)) - tracker_config.pop("validation_prompts") - tracker_config.pop("trainable_modules") - tracker_config.pop("trainable_modules_low_learning_rate") - tracker_config.pop("fix_sample_size") + keys_to_pop = [k for k, v in tracker_config.items() if isinstance(v, list)] + for k in keys_to_pop: + tracker_config.pop(k) + print(f"Removed tracker_config['{k}']") accelerator.init_trackers(args.tracker_project_name, tracker_config) # Function for unwrapping if model was compiled with `torch.compile`. diff --git a/scripts/wan2.1_fun/train_control_lora.py b/scripts/wan2.1_fun/train_control_lora.py index d81e98d..62b5ee4 100755 --- a/scripts/wan2.1_fun/train_control_lora.py +++ b/scripts/wan2.1_fun/train_control_lora.py @@ -234,6 +234,13 @@ def parse_args(): nargs="+", help=("A set of prompts evaluated every `--validation_epochs` and logged to `--report_to`."), ) + parser.add_argument( + "--validation_paths", + type=str, + default=None, + nargs="+", + help=("A set of control videos evaluated every `--validation_epochs` and logged to `--report_to`."), + ) parser.add_argument( "--output_dir", type=str, @@ -1360,8 +1367,10 @@ def main(): # The trackers initializes automatically on the main process. if accelerator.is_main_process: tracker_config = dict(vars(args)) - tracker_config.pop("validation_prompts") - tracker_config.pop("fix_sample_size") + keys_to_pop = [k for k, v in tracker_config.items() if isinstance(v, list)] + for k in keys_to_pop: + tracker_config.pop(k) + print(f"Removed tracker_config['{k}']") accelerator.init_trackers(args.tracker_project_name, tracker_config) # Function for unwrapping if model was compiled with `torch.compile`. diff --git a/scripts/wan2.1_fun/train_lora.py b/scripts/wan2.1_fun/train_lora.py index a0dd27c..ddf8bbc 100755 --- a/scripts/wan2.1_fun/train_lora.py +++ b/scripts/wan2.1_fun/train_lora.py @@ -1356,8 +1356,10 @@ def main(): # The trackers initializes automatically on the main process. if accelerator.is_main_process: tracker_config = dict(vars(args)) - tracker_config.pop("validation_prompts") - tracker_config.pop("fix_sample_size") + keys_to_pop = [k for k, v in tracker_config.items() if isinstance(v, list)] + for k in keys_to_pop: + tracker_config.pop(k) + print(f"Removed tracker_config['{k}']") accelerator.init_trackers(args.tracker_project_name, tracker_config) # Function for unwrapping if model was compiled with `torch.compile`. diff --git a/scripts/wan2.1_fun/train_reward_lora.py b/scripts/wan2.1_fun/train_reward_lora.py index ac56147..3c0b2dd 100755 --- a/scripts/wan2.1_fun/train_reward_lora.py +++ b/scripts/wan2.1_fun/train_reward_lora.py @@ -1067,8 +1067,10 @@ def main(): # The trackers initializes automatically on the main process. if accelerator.is_main_process: tracker_config = dict(vars(args)) - tracker_config.pop("validation_prompts") - tracker_config.pop("backprop_step_list", None) + keys_to_pop = [k for k, v in tracker_config.items() if isinstance(v, list)] + for k in keys_to_pop: + tracker_config.pop(k) + print(f"Removed tracker_config['{k}']") accelerator.init_trackers(args.tracker_project_name, tracker_config) # Train! diff --git a/scripts/wan2.1_vace/train.py b/scripts/wan2.1_vace/train.py index 975dbd6..2c6a0d4 100644 --- a/scripts/wan2.1_vace/train.py +++ b/scripts/wan2.1_vace/train.py @@ -146,7 +146,8 @@ def log_validation(vae, text_encoder, tokenizer, clip_image_encoder, transformer with torch.no_grad(): with torch.autocast("cuda", dtype=weight_dtype): video_length = int(args.video_sample_n_frames // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if args.video_sample_n_frames != 1 else 1 - input_video, input_video_mask, ref_image, clip_image = get_video_to_video_latent(args.validation_paths[i], video_length=video_length, sample_size=[args.video_sample_size, args.video_sample_size]) + inpaint_video, inpaint_video_mask, clip_image = get_image_to_video_latent(None, None, video_length=video_length, sample_size=[args.video_sample_size, args.video_sample_size]) + control_video, _, _, _ = get_video_to_video_latent(args.validation_paths[i], video_length=video_length, sample_size=[args.video_sample_size, args.video_sample_size]) sample = pipeline( args.validation_prompts[i], num_frames = video_length, @@ -155,7 +156,11 @@ def log_validation(vae, text_encoder, tokenizer, clip_image_encoder, transformer width = args.video_sample_size, generator = generator, - control_video = input_video, + video = inpaint_video, + mask_video = inpaint_video_mask, + control_video = control_video, + subject_ref_images = None, + vace_context_scale = 1, ).videos os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-{i}.gif")) @@ -231,6 +236,13 @@ def parse_args(): nargs="+", help=("A set of prompts evaluated every `--validation_epochs` and logged to `--report_to`."), ) + parser.add_argument( + "--validation_paths", + type=str, + default=None, + nargs="+", + help=("A set of control videos evaluated every `--validation_epochs` and logged to `--report_to`."), + ) parser.add_argument( "--output_dir", type=str, @@ -1403,10 +1415,10 @@ def main(): # The trackers initializes automatically on the main process. if accelerator.is_main_process: tracker_config = dict(vars(args)) - tracker_config.pop("validation_prompts") - tracker_config.pop("trainable_modules") - tracker_config.pop("trainable_modules_low_learning_rate") - tracker_config.pop("fix_sample_size") + keys_to_pop = [k for k, v in tracker_config.items() if isinstance(v, list)] + for k in keys_to_pop: + tracker_config.pop(k) + print(f"Removed tracker_config['{k}']") accelerator.init_trackers(args.tracker_project_name, tracker_config) # Function for unwrapping if model was compiled with `torch.compile`. diff --git a/scripts/wan2.2/train.py b/scripts/wan2.2/train.py index e2dcd92..932aeeb 100644 --- a/scripts/wan2.2/train.py +++ b/scripts/wan2.2/train.py @@ -73,7 +73,7 @@ from videox_fun.data.dataset_image_video import (ImageVideoDataset, get_random_mask) from videox_fun.models import (AutoencoderKLWan, AutoencoderKLWan3_8, WanT5EncoderModel, Wan2_2Transformer3DModel) -from videox_fun.pipeline import WanPipeline, WanI2VPipeline +from videox_fun.pipeline import Wan2_2Pipeline, Wan2_2I2VPipeline from videox_fun.utils.discrete_sampler import DiscreteSampling from videox_fun.utils.utils import get_image_to_video_latent, save_videos_grid @@ -165,29 +165,66 @@ def log_validation(vae, text_encoder, tokenizer, transformer3d, args, config, ac try: logger.info("Running validation... ") - transformer3d_val = Wan2_2Transformer3DModel.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, config['transformer_additional_kwargs'].get('transformer_subpath', 'transformer')), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - ).to(weight_dtype) - transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) + if args.boundary_type == "full": + sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer') + + transformer3d_val = Wan2_2Transformer3DModel.from_pretrained( + os.path.join(args.pretrained_model_name_or_path, sub_path), + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), + ).to(weight_dtype) + transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) + + transformer3d_2_val = None + else: + if args.boundary_type == "low": + sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer') + + transformer3d_val = Wan2_2Transformer3DModel.from_pretrained( + os.path.join(args.pretrained_model_name_or_path, sub_path), + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), + ).to(weight_dtype) + transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) + + sub_path = config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer') + transformer3d_2_val = Wan2_2Transformer3DModel.from_pretrained( + os.path.join(args.pretrained_model_name_or_path, sub_path), + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), + ).to(weight_dtype) + else: + sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer') + + transformer3d_val = Wan2_2Transformer3DModel.from_pretrained( + os.path.join(args.pretrained_model_name_or_path, sub_path), + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), + ).to(weight_dtype) + + sub_path = config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer') + transformer3d_2_val = Wan2_2Transformer3DModel.from_pretrained( + os.path.join(args.pretrained_model_name_or_path, sub_path), + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), + ).to(weight_dtype) + transformer3d_2_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) + scheduler = FlowMatchEulerDiscreteScheduler( **filter_kwargs(FlowMatchEulerDiscreteScheduler, OmegaConf.to_container(config['scheduler_kwargs'])) ) - + if args.train_mode != "normal": - pipeline = WanI2VPipeline( + pipeline = Wan2_2I2VPipeline( vae=accelerator.unwrap_model(vae).to(weight_dtype), text_encoder=accelerator.unwrap_model(text_encoder), tokenizer=tokenizer, transformer=transformer3d_val, + transformer_2=transformer3d_2_val, scheduler=scheduler, ) else: - pipeline = WanPipeline( + pipeline = Wan2_2Pipeline( vae=accelerator.unwrap_model(vae).to(weight_dtype), text_encoder=accelerator.unwrap_model(text_encoder), tokenizer=tokenizer, transformer=transformer3d_val, + transformer_2=transformer3d_2_val, scheduler=scheduler, ) pipeline = pipeline.to(accelerator.device) @@ -1415,10 +1452,10 @@ def main(): # The trackers initializes automatically on the main process. if accelerator.is_main_process: tracker_config = dict(vars(args)) - tracker_config.pop("validation_prompts") - tracker_config.pop("trainable_modules") - tracker_config.pop("trainable_modules_low_learning_rate") - tracker_config.pop("fix_sample_size") + keys_to_pop = [k for k, v in tracker_config.items() if isinstance(v, list)] + for k in keys_to_pop: + tracker_config.pop(k) + print(f"Removed tracker_config['{k}']") accelerator.init_trackers(args.tracker_project_name, tracker_config) # Function for unwrapping if model was compiled with `torch.compile`. diff --git a/scripts/wan2.2/train_lora.py b/scripts/wan2.2/train_lora.py index 24f8fff..3753094 100755 --- a/scripts/wan2.2/train_lora.py +++ b/scripts/wan2.2/train_lora.py @@ -163,11 +163,46 @@ def log_validation(vae, text_encoder, tokenizer, transformer3d, network, config, try: logger.info("Running validation... ") - transformer3d_val = Wan2_2Transformer3DModel.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, config['transformer_additional_kwargs'].get('transformer_subpath', 'transformer')), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - ).to(weight_dtype) - transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) + if args.boundary_type == "full": + sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer') + + transformer3d_val = Wan2_2Transformer3DModel.from_pretrained( + os.path.join(args.pretrained_model_name_or_path, sub_path), + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), + ).to(weight_dtype) + transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) + + transformer3d_2_val = None + else: + if args.boundary_type == "low": + sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer') + + transformer3d_val = Wan2_2Transformer3DModel.from_pretrained( + os.path.join(args.pretrained_model_name_or_path, sub_path), + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), + ).to(weight_dtype) + transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) + + sub_path = config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer') + transformer3d_2_val = Wan2_2Transformer3DModel.from_pretrained( + os.path.join(args.pretrained_model_name_or_path, sub_path), + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), + ).to(weight_dtype) + else: + sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer') + + transformer3d_val = Wan2_2Transformer3DModel.from_pretrained( + os.path.join(args.pretrained_model_name_or_path, sub_path), + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), + ).to(weight_dtype) + + sub_path = config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer') + transformer3d_2_val = Wan2_2Transformer3DModel.from_pretrained( + os.path.join(args.pretrained_model_name_or_path, sub_path), + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), + ).to(weight_dtype) + transformer3d_2_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) + scheduler = FlowMatchEulerDiscreteScheduler( **filter_kwargs(FlowMatchEulerDiscreteScheduler, OmegaConf.to_container(config['scheduler_kwargs'])) ) @@ -178,6 +213,7 @@ def log_validation(vae, text_encoder, tokenizer, transformer3d, network, config, text_encoder=accelerator.unwrap_model(text_encoder), tokenizer=tokenizer, transformer=transformer3d_val, + transformer_2=transformer3d_2_val, scheduler=scheduler, ) else: @@ -186,6 +222,7 @@ def log_validation(vae, text_encoder, tokenizer, transformer3d, network, config, text_encoder=accelerator.unwrap_model(text_encoder), tokenizer=tokenizer, transformer=transformer3d_val, + transformer_2=transformer3d_2_val, scheduler=scheduler, ) pipeline = pipeline.to(accelerator.device) @@ -1362,8 +1399,10 @@ def main(): # The trackers initializes automatically on the main process. if accelerator.is_main_process: tracker_config = dict(vars(args)) - tracker_config.pop("validation_prompts") - tracker_config.pop("fix_sample_size") + keys_to_pop = [k for k, v in tracker_config.items() if isinstance(v, list)] + for k in keys_to_pop: + tracker_config.pop(k) + print(f"Removed tracker_config['{k}']") accelerator.init_trackers(args.tracker_project_name, tracker_config) # Function for unwrapping if model was compiled with `torch.compile`. diff --git a/scripts/wan2.2_fun/train.py b/scripts/wan2.2_fun/train.py index bb88de7..67eedcf 100644 --- a/scripts/wan2.2_fun/train.py +++ b/scripts/wan2.2_fun/train.py @@ -75,7 +75,7 @@ from videox_fun.data.dataset_image_video import (ImageVideoDataset, get_random_mask) from videox_fun.models import (AutoencoderKLWan, AutoencoderKLWan3_8, CLIPModel, WanT5EncoderModel, Wan2_2Transformer3DModel) -from videox_fun.pipeline import WanFunInpaintPipeline, WanFunPipeline +from videox_fun.pipeline import Wan2_2Pipeline, Wan2_2I2VPipeline from videox_fun.utils.discrete_sampler import DiscreteSampling from videox_fun.utils.utils import get_image_to_video_latent, save_videos_grid @@ -157,20 +157,66 @@ def log_validation(vae, text_encoder, tokenizer, transformer3d, args, config, ac **filter_kwargs(FlowMatchEulerDiscreteScheduler, OmegaConf.to_container(config['scheduler_kwargs'])) ) + if args.boundary_type == "full": + sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer') + + transformer3d_val = Wan2_2Transformer3DModel.from_pretrained( + os.path.join(args.pretrained_model_name_or_path, sub_path), + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), + ).to(weight_dtype) + transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) + + transformer3d_2_val = None + else: + if args.boundary_type == "low": + sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer') + + transformer3d_val = Wan2_2Transformer3DModel.from_pretrained( + os.path.join(args.pretrained_model_name_or_path, sub_path), + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), + ).to(weight_dtype) + transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) + + sub_path = config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer') + transformer3d_2_val = Wan2_2Transformer3DModel.from_pretrained( + os.path.join(args.pretrained_model_name_or_path, sub_path), + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), + ).to(weight_dtype) + else: + sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer') + + transformer3d_val = Wan2_2Transformer3DModel.from_pretrained( + os.path.join(args.pretrained_model_name_or_path, sub_path), + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), + ).to(weight_dtype) + + sub_path = config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer') + transformer3d_2_val = Wan2_2Transformer3DModel.from_pretrained( + os.path.join(args.pretrained_model_name_or_path, sub_path), + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), + ).to(weight_dtype) + transformer3d_2_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) + + scheduler = FlowMatchEulerDiscreteScheduler( + **filter_kwargs(FlowMatchEulerDiscreteScheduler, OmegaConf.to_container(config['scheduler_kwargs'])) + ) + if args.train_mode != "normal": - pipeline = WanFunInpaintPipeline( + pipeline = Wan2_2I2VPipeline( vae=accelerator.unwrap_model(vae).to(weight_dtype), text_encoder=accelerator.unwrap_model(text_encoder), tokenizer=tokenizer, transformer=transformer3d_val, + transformer_2=transformer3d_2_val, scheduler=scheduler, ) else: - pipeline = WanFunPipeline( + pipeline = Wan2_2Pipeline( vae=accelerator.unwrap_model(vae).to(weight_dtype), text_encoder=accelerator.unwrap_model(text_encoder), tokenizer=tokenizer, transformer=transformer3d_val, + transformer_2=transformer3d_2_val, scheduler=scheduler, ) pipeline = pipeline.to(accelerator.device) @@ -1423,10 +1469,10 @@ def main(): # The trackers initializes automatically on the main process. if accelerator.is_main_process: tracker_config = dict(vars(args)) - tracker_config.pop("validation_prompts") - tracker_config.pop("trainable_modules") - tracker_config.pop("trainable_modules_low_learning_rate") - tracker_config.pop("fix_sample_size") + keys_to_pop = [k for k, v in tracker_config.items() if isinstance(v, list)] + for k in keys_to_pop: + tracker_config.pop(k) + print(f"Removed tracker_config['{k}']") accelerator.init_trackers(args.tracker_project_name, tracker_config) # Function for unwrapping if model was compiled with `torch.compile`. diff --git a/scripts/wan2.2_fun/train_control.py b/scripts/wan2.2_fun/train_control.py index 0a34e22..d951000 100644 --- a/scripts/wan2.2_fun/train_control.py +++ b/scripts/wan2.2_fun/train_control.py @@ -75,7 +75,7 @@ from videox_fun.data.dataset_image_video import (ImageVideoControlDataset, process_pose_params) from videox_fun.models import (AutoencoderKLWan, AutoencoderKLWan3_8, CLIPModel, WanT5EncoderModel, Wan2_2Transformer3DModel) -from videox_fun.pipeline import WanFunControlPipeline +from videox_fun.pipeline import Wan2_2FunControlPipeline from videox_fun.utils.discrete_sampler import DiscreteSampling from videox_fun.utils.lora_utils import (create_network, merge_lora, unmerge_lora) @@ -152,20 +152,55 @@ def log_validation(vae, text_encoder, tokenizer, transformer3d, args, config, ac try: logger.info("Running validation... ") - transformer3d_val = Wan2_2Transformer3DModel.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, config['transformer_additional_kwargs'].get('transformer_subpath', 'transformer')), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - ).to(weight_dtype) - transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) + if args.boundary_type == "full": + sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer') + + transformer3d_val = Wan2_2Transformer3DModel.from_pretrained( + os.path.join(args.pretrained_model_name_or_path, sub_path), + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), + ).to(weight_dtype) + transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) + + transformer3d_2_val = None + else: + if args.boundary_type == "low": + sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer') + + transformer3d_val = Wan2_2Transformer3DModel.from_pretrained( + os.path.join(args.pretrained_model_name_or_path, sub_path), + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), + ).to(weight_dtype) + transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) + + sub_path = config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer') + transformer3d_2_val = Wan2_2Transformer3DModel.from_pretrained( + os.path.join(args.pretrained_model_name_or_path, sub_path), + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), + ).to(weight_dtype) + else: + sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer') + + transformer3d_val = Wan2_2Transformer3DModel.from_pretrained( + os.path.join(args.pretrained_model_name_or_path, sub_path), + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), + ).to(weight_dtype) + + sub_path = config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer') + transformer3d_2_val = Wan2_2Transformer3DModel.from_pretrained( + os.path.join(args.pretrained_model_name_or_path, sub_path), + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), + ).to(weight_dtype) + transformer3d_2_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) + scheduler = FlowMatchEulerDiscreteScheduler( **filter_kwargs(FlowMatchEulerDiscreteScheduler, OmegaConf.to_container(config['scheduler_kwargs'])) ) - - pipeline = WanFunControlPipeline( + pipeline = Wan2_2FunControlPipeline( vae=accelerator.unwrap_model(vae).to(weight_dtype), text_encoder=accelerator.unwrap_model(text_encoder), tokenizer=tokenizer, transformer=transformer3d_val, + transformer_2=transformer3d_2_val, scheduler=scheduler, ) pipeline = pipeline.to(accelerator.device) @@ -265,6 +300,13 @@ def parse_args(): nargs="+", help=("A set of prompts evaluated every `--validation_epochs` and logged to `--report_to`."), ) + parser.add_argument( + "--validation_paths", + type=str, + default=None, + nargs="+", + help=("A set of control videos evaluated every `--validation_epochs` and logged to `--report_to`."), + ) parser.add_argument( "--output_dir", type=str, @@ -1484,10 +1526,10 @@ def main(): # The trackers initializes automatically on the main process. if accelerator.is_main_process: tracker_config = dict(vars(args)) - tracker_config.pop("validation_prompts") - tracker_config.pop("trainable_modules") - tracker_config.pop("trainable_modules_low_learning_rate") - tracker_config.pop("fix_sample_size") + keys_to_pop = [k for k, v in tracker_config.items() if isinstance(v, list)] + for k in keys_to_pop: + tracker_config.pop(k) + print(f"Removed tracker_config['{k}']") accelerator.init_trackers(args.tracker_project_name, tracker_config) # Function for unwrapping if model was compiled with `torch.compile`. diff --git a/scripts/wan2.2_fun/train_control.sh b/scripts/wan2.2_fun/train_control.sh index 89abfbd..5fd33c6 100644 --- a/scripts/wan2.2_fun/train_control.sh +++ b/scripts/wan2.2_fun/train_control.sh @@ -37,6 +37,7 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.2_fun/train_control.py \ --training_with_video_token_length \ --enable_bucket \ --uniform_sampling \ + --boundary_type="low" \ --train_mode="control_ref" \ --control_ref_image="random" \ --add_inpaint_info \ diff --git a/scripts/wan2.2_fun/train_control_lora.py b/scripts/wan2.2_fun/train_control_lora.py index b41635d..3a97da3 100644 --- a/scripts/wan2.2_fun/train_control_lora.py +++ b/scripts/wan2.2_fun/train_control_lora.py @@ -75,7 +75,7 @@ from videox_fun.data.dataset_image_video import (ImageVideoControlDataset, process_pose_params) from videox_fun.models import (AutoencoderKLWan, AutoencoderKLWan3_8, CLIPModel, WanT5EncoderModel, Wan2_2Transformer3DModel) -from videox_fun.pipeline import WanFunControlPipeline +from videox_fun.pipeline import Wan2_2FunControlPipeline from videox_fun.utils.discrete_sampler import DiscreteSampling from videox_fun.utils.lora_utils import (create_network, merge_lora, unmerge_lora) @@ -152,20 +152,56 @@ def log_validation(vae, text_encoder, tokenizer, transformer3d, network, config, try: logger.info("Running validation... ") - transformer3d_val = Wan2_2Transformer3DModel.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, config['transformer_additional_kwargs'].get('transformer_subpath', 'transformer')), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - ).to(weight_dtype) - transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) + if args.boundary_type == "full": + sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer') + + transformer3d_val = Wan2_2Transformer3DModel.from_pretrained( + os.path.join(args.pretrained_model_name_or_path, sub_path), + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), + ).to(weight_dtype) + transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) + + transformer3d_2_val = None + else: + if args.boundary_type == "low": + sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer') + + transformer3d_val = Wan2_2Transformer3DModel.from_pretrained( + os.path.join(args.pretrained_model_name_or_path, sub_path), + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), + ).to(weight_dtype) + transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) + + sub_path = config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer') + transformer3d_2_val = Wan2_2Transformer3DModel.from_pretrained( + os.path.join(args.pretrained_model_name_or_path, sub_path), + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), + ).to(weight_dtype) + else: + sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer') + + transformer3d_val = Wan2_2Transformer3DModel.from_pretrained( + os.path.join(args.pretrained_model_name_or_path, sub_path), + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), + ).to(weight_dtype) + + sub_path = config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer') + transformer3d_2_val = Wan2_2Transformer3DModel.from_pretrained( + os.path.join(args.pretrained_model_name_or_path, sub_path), + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), + ).to(weight_dtype) + transformer3d_2_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) + scheduler = FlowMatchEulerDiscreteScheduler( **filter_kwargs(FlowMatchEulerDiscreteScheduler, OmegaConf.to_container(config['scheduler_kwargs'])) ) - pipeline = WanFunControlPipeline( + pipeline = Wan2_2FunControlPipeline( vae=accelerator.unwrap_model(vae).to(weight_dtype), text_encoder=accelerator.unwrap_model(text_encoder), tokenizer=tokenizer, transformer=transformer3d_val, + transformer_2=transformer3d_2_val, scheduler=scheduler, ) pipeline = pipeline.to(accelerator.device) @@ -269,6 +305,13 @@ def parse_args(): nargs="+", help=("A set of prompts evaluated every `--validation_epochs` and logged to `--report_to`."), ) + parser.add_argument( + "--validation_paths", + type=str, + default=None, + nargs="+", + help=("A set of control videos evaluated every `--validation_epochs` and logged to `--report_to`."), + ) parser.add_argument( "--output_dir", type=str, @@ -1431,8 +1474,10 @@ def main(): # The trackers initializes automatically on the main process. if accelerator.is_main_process: tracker_config = dict(vars(args)) - tracker_config.pop("validation_prompts") - tracker_config.pop("fix_sample_size") + keys_to_pop = [k for k, v in tracker_config.items() if isinstance(v, list)] + for k in keys_to_pop: + tracker_config.pop(k) + print(f"Removed tracker_config['{k}']") accelerator.init_trackers(args.tracker_project_name, tracker_config) # Function for unwrapping if model was compiled with `torch.compile`. diff --git a/scripts/wan2.2_fun/train_lora.py b/scripts/wan2.2_fun/train_lora.py index 2185309..d5a0e01 100644 --- a/scripts/wan2.2_fun/train_lora.py +++ b/scripts/wan2.2_fun/train_lora.py @@ -71,7 +71,7 @@ from videox_fun.data.dataset_image_video import (ImageVideoDataset, get_random_mask) from videox_fun.models import (AutoencoderKLWan, AutoencoderKLWan3_8, WanT5EncoderModel, Wan2_2Transformer3DModel) -from videox_fun.pipeline import WanFunInpaintPipeline, WanFunPipeline +from videox_fun.pipeline import Wan2_2Pipeline, Wan2_2I2VPipeline from videox_fun.utils.discrete_sampler import DiscreteSampling from videox_fun.utils.lora_utils import (create_network, merge_lora, unmerge_lora) @@ -146,29 +146,66 @@ def log_validation(vae, text_encoder, tokenizer, transformer3d, network, config, try: logger.info("Running validation... ") - transformer3d_val = Wan2_2Transformer3DModel.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, config['transformer_additional_kwargs'].get('transformer_subpath', 'transformer')), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - ).to(weight_dtype) - transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) + if args.boundary_type == "full": + sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer') + + transformer3d_val = Wan2_2Transformer3DModel.from_pretrained( + os.path.join(args.pretrained_model_name_or_path, sub_path), + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), + ).to(weight_dtype) + transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) + + transformer3d_2_val = None + else: + if args.boundary_type == "low": + sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer') + + transformer3d_val = Wan2_2Transformer3DModel.from_pretrained( + os.path.join(args.pretrained_model_name_or_path, sub_path), + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), + ).to(weight_dtype) + transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) + + sub_path = config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer') + transformer3d_2_val = Wan2_2Transformer3DModel.from_pretrained( + os.path.join(args.pretrained_model_name_or_path, sub_path), + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), + ).to(weight_dtype) + else: + sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer') + + transformer3d_val = Wan2_2Transformer3DModel.from_pretrained( + os.path.join(args.pretrained_model_name_or_path, sub_path), + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), + ).to(weight_dtype) + + sub_path = config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer') + transformer3d_2_val = Wan2_2Transformer3DModel.from_pretrained( + os.path.join(args.pretrained_model_name_or_path, sub_path), + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), + ).to(weight_dtype) + transformer3d_2_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) + scheduler = FlowMatchEulerDiscreteScheduler( **filter_kwargs(FlowMatchEulerDiscreteScheduler, OmegaConf.to_container(config['scheduler_kwargs'])) ) if args.train_mode != "normal": - pipeline = WanFunInpaintPipeline( + pipeline = Wan2_2I2VPipeline( vae=accelerator.unwrap_model(vae).to(weight_dtype), text_encoder=accelerator.unwrap_model(text_encoder), tokenizer=tokenizer, transformer=transformer3d_val, + transformer_2=transformer3d_2_val, scheduler=scheduler, ) else: - pipeline = WanFunPipeline( + pipeline = Wan2_2Pipeline( vae=accelerator.unwrap_model(vae).to(weight_dtype), text_encoder=accelerator.unwrap_model(text_encoder), tokenizer=tokenizer, transformer=transformer3d_val, + transformer_2=transformer3d_2_val, scheduler=scheduler, ) pipeline = pipeline.to(accelerator.device) @@ -1363,8 +1400,10 @@ def main(): # The trackers initializes automatically on the main process. if accelerator.is_main_process: tracker_config = dict(vars(args)) - tracker_config.pop("validation_prompts") - tracker_config.pop("fix_sample_size") + keys_to_pop = [k for k, v in tracker_config.items() if isinstance(v, list)] + for k in keys_to_pop: + tracker_config.pop(k) + print(f"Removed tracker_config['{k}']") accelerator.init_trackers(args.tracker_project_name, tracker_config) # Function for unwrapping if model was compiled with `torch.compile`. diff --git a/scripts/wan2.2_fun/train_lora.sh b/scripts/wan2.2_fun/train_lora.sh index 5c8464e..fc3e06f 100644 --- a/scripts/wan2.2_fun/train_lora.sh +++ b/scripts/wan2.2_fun/train_lora.sh @@ -38,5 +38,4 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.2_fun/train_lora.py \ --train_mode="inpaint" \ --boundary_type="low" \ --lora_skip_name="ffn" \ - --boundary_type="low" \ --low_vram diff --git a/scripts/wan2.2_fun/train_reward_lora.py b/scripts/wan2.2_fun/train_reward_lora.py index 3a4dab0..e460665 100644 --- a/scripts/wan2.2_fun/train_reward_lora.py +++ b/scripts/wan2.2_fun/train_reward_lora.py @@ -1230,8 +1230,10 @@ def main(): # The trackers initializes automatically on the main process. if accelerator.is_main_process: tracker_config = dict(vars(args)) - tracker_config.pop("validation_prompts") - tracker_config.pop("backprop_step_list", None) + keys_to_pop = [k for k, v in tracker_config.items() if isinstance(v, list)] + for k in keys_to_pop: + tracker_config.pop(k) + print(f"Removed tracker_config['{k}']") accelerator.init_trackers(args.tracker_project_name, tracker_config) # Train! diff --git a/scripts/wan2.2_vace_fun/train.py b/scripts/wan2.2_vace_fun/train.py index b9faa5a..abf7c5b 100644 --- a/scripts/wan2.2_vace_fun/train.py +++ b/scripts/wan2.2_vace_fun/train.py @@ -74,7 +74,7 @@ from videox_fun.data.dataset_image_video import (ImageVideoControlDataset, padding_image, process_pose_file, process_pose_params) -from videox_fun.models import (AutoencoderKLWan, CLIPModel, +from videox_fun.models import (AutoencoderKLWan, CLIPModel, AutoencoderKLWan3_8, VaceWanTransformer3DModel, WanT5EncoderModel) from videox_fun.pipeline import Wan2_2VaceFunPipeline from videox_fun.utils.discrete_sampler import DiscreteSampling @@ -117,11 +117,46 @@ def log_validation(vae, text_encoder, tokenizer, clip_image_encoder, transformer try: logger.info("Running validation... ") - transformer3d_val = VaceWanTransformer3DModel.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, config['transformer_additional_kwargs'].get('transformer_subpath', 'transformer')), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - ).to(weight_dtype) - transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) + if args.boundary_type == "full": + sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer') + + transformer3d_val = Wan2_2Transformer3DModel.from_pretrained( + os.path.join(args.pretrained_model_name_or_path, sub_path), + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), + ).to(weight_dtype) + transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) + + transformer3d_2_val = None + else: + if args.boundary_type == "low": + sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer') + + transformer3d_val = Wan2_2Transformer3DModel.from_pretrained( + os.path.join(args.pretrained_model_name_or_path, sub_path), + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), + ).to(weight_dtype) + transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) + + sub_path = config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer') + transformer3d_2_val = Wan2_2Transformer3DModel.from_pretrained( + os.path.join(args.pretrained_model_name_or_path, sub_path), + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), + ).to(weight_dtype) + else: + sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer') + + transformer3d_val = Wan2_2Transformer3DModel.from_pretrained( + os.path.join(args.pretrained_model_name_or_path, sub_path), + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), + ).to(weight_dtype) + + sub_path = config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer') + transformer3d_2_val = Wan2_2Transformer3DModel.from_pretrained( + os.path.join(args.pretrained_model_name_or_path, sub_path), + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), + ).to(weight_dtype) + transformer3d_2_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) + scheduler = FlowMatchEulerDiscreteScheduler( **filter_kwargs(FlowMatchEulerDiscreteScheduler, OmegaConf.to_container(config['scheduler_kwargs'])) ) @@ -131,6 +166,7 @@ def log_validation(vae, text_encoder, tokenizer, clip_image_encoder, transformer text_encoder=accelerator.unwrap_model(text_encoder), tokenizer=tokenizer, transformer=transformer3d_val, + transformer_2=transformer3d_2_val, scheduler=scheduler, clip_image_encoder=clip_image_encoder, ) @@ -146,7 +182,8 @@ def log_validation(vae, text_encoder, tokenizer, clip_image_encoder, transformer with torch.no_grad(): with torch.autocast("cuda", dtype=weight_dtype): video_length = int(args.video_sample_n_frames // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if args.video_sample_n_frames != 1 else 1 - input_video, input_video_mask, ref_image, clip_image = get_video_to_video_latent(args.validation_paths[i], video_length=video_length, sample_size=[args.video_sample_size, args.video_sample_size]) + inpaint_video, inpaint_video_mask, clip_image = get_image_to_video_latent(None, None, video_length=video_length, sample_size=[args.video_sample_size, args.video_sample_size]) + control_video, _, _, _ = get_video_to_video_latent(args.validation_paths[i], video_length=video_length, sample_size=[args.video_sample_size, args.video_sample_size]) sample = pipeline( args.validation_prompts[i], num_frames = video_length, @@ -155,7 +192,11 @@ def log_validation(vae, text_encoder, tokenizer, clip_image_encoder, transformer width = args.video_sample_size, generator = generator, - control_video = input_video, + video = inpaint_video, + mask_video = inpaint_video_mask, + control_video = control_video, + subject_ref_images = None, + vace_context_scale = 1, ).videos os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-{i}.gif")) @@ -231,6 +272,13 @@ def parse_args(): nargs="+", help=("A set of prompts evaluated every `--validation_epochs` and logged to `--report_to`."), ) + parser.add_argument( + "--validation_paths", + type=str, + default=None, + nargs="+", + help=("A set of control videos evaluated every `--validation_epochs` and logged to `--report_to`."), + ) parser.add_argument( "--output_dir", type=str, @@ -780,7 +828,11 @@ def main(): ) text_encoder = text_encoder.eval() # Get Vae - vae = AutoencoderKLWan.from_pretrained( + Chosen_AutoencoderKL = { + "AutoencoderKLWan": AutoencoderKLWan, + "AutoencoderKLWan3_8": AutoencoderKLWan3_8 + }[config['vae_kwargs'].get('vae_type', 'AutoencoderKLWan')] + vae = Chosen_AutoencoderKL.from_pretrained( os.path.join(args.pretrained_model_name_or_path, config['vae_kwargs'].get('vae_subpath', 'vae')), additional_kwargs=OmegaConf.to_container(config['vae_kwargs']), ) @@ -1023,7 +1075,8 @@ def main(): # Get the training dataset sample_n_frames_bucket_interval = vae.config.temporal_compression_ratio - + spatial_compression_ratio = vae.config.spatial_compression_ratio + if args.fix_sample_size is not None and args.enable_bucket: args.video_sample_size = max(max(args.fix_sample_size), args.video_sample_size) args.image_sample_size = max(max(args.fix_sample_size), args.image_sample_size) @@ -1186,7 +1239,7 @@ def main(): aspect_ratio_random_crop_sample_size = {key : [x / 512 * args.video_sample_size / random_downsample_ratio for x in ASPECT_RATIO_RANDOM_CROP_512[key]] for key in ASPECT_RATIO_RANDOM_CROP_512.keys()} if args.fix_sample_size is not None: - fix_sample_size = [int(x / 16) * 16 for x in args.fix_sample_size] + fix_sample_size = [int(x / spatial_compression_ratio / 2) * spatial_compression_ratio * 2 for x in args.fix_sample_size] elif args.random_ratio_crop: if rng is None: random_sample_size = aspect_ratio_random_crop_sample_size[ @@ -1196,10 +1249,10 @@ def main(): random_sample_size = aspect_ratio_random_crop_sample_size[ rng.choice(list(aspect_ratio_random_crop_sample_size.keys()), p = ASPECT_RATIO_RANDOM_CROP_PROB) ] - random_sample_size = [int(x / 16) * 16 for x in random_sample_size] + random_sample_size = [int(x / spatial_compression_ratio / 2) * spatial_compression_ratio * 2 for x in random_sample_size] else: closest_size, closest_ratio = get_closest_ratio(h, w, ratios=aspect_ratio_sample_size) - closest_size = [int(x / 16) * 16 for x in closest_size] + closest_size = [int(x / spatial_compression_ratio / 2) * spatial_compression_ratio * 2 for x in closest_size] for example in examples: # To 0~1 @@ -1808,7 +1861,7 @@ def main(): vace_latents = vace_encode_frames(control_pixel_values, subject_ref_images, mask) mask = torch.ones_like(mask) - mask_latents = vace_encode_masks(mask, subject_ref_images) + mask_latents = vace_encode_masks(mask, subject_ref_images, vae_stride=[4, spatial_compression_ratio, spatial_compression_ratio]) vace_context = torch.stack(vace_latent(vace_latents, mask_latents)) if subject_ref_images is not None: diff --git a/videox_fun/models/flux_transformer2d.py b/videox_fun/models/flux_transformer2d.py index 8d42d16..b8020b6 100644 --- a/videox_fun/models/flux_transformer2d.py +++ b/videox_fun/models/flux_transformer2d.py @@ -33,8 +33,8 @@ from diffusers.models.modeling_utils import ModelMixin from diffusers.models.normalization import (AdaLayerNormContinuous, AdaLayerNormZero, AdaLayerNormZeroSingle) -from diffusers.utils import (USE_PEFT_BACKEND, logging, scale_lora_layers, - unscale_lora_layers) +from diffusers.utils import (USE_PEFT_BACKEND, is_torch_version, logging, + scale_lora_layers, unscale_lora_layers) from diffusers.utils.torch_utils import maybe_allow_in_graph from ..dist import (FluxMultiGPUsAttnProcessor2_0, get_sequence_parallel_rank, @@ -695,6 +695,14 @@ class FluxTransformer2DModel( self.sp_world_size = 1 self.sp_world_rank = 0 + def _set_gradient_checkpointing(self, *args, **kwargs): + if "value" in kwargs: + self.gradient_checkpointing = kwargs["value"] + elif "enable" in kwargs: + self.gradient_checkpointing = kwargs["enable"] + else: + raise ValueError("Invalid set gradient checkpointing") + def enable_multi_gpus_inference(self,): self.sp_world_size = get_sequence_parallel_world_size() self.sp_world_rank = get_sequence_parallel_rank() @@ -868,13 +876,20 @@ class FluxTransformer2DModel( for index_block, block in enumerate(self.transformer_blocks): if torch.is_grad_enabled() and self.gradient_checkpointing: - encoder_hidden_states, hidden_states = self._gradient_checkpointing_func( - block, + def create_custom_forward(module): + def custom_forward(*inputs): + return module(*inputs) + + return custom_forward + ckpt_kwargs: Dict[str, Any] = {"use_reentrant": False} if is_torch_version(">=", "1.11.0") else {} + encoder_hidden_states, hidden_states = torch.utils.checkpoint.checkpoint( + create_custom_forward(block), hidden_states, encoder_hidden_states, temb, image_rotary_emb, joint_attention_kwargs, + **ckpt_kwargs, ) else: @@ -900,13 +915,20 @@ class FluxTransformer2DModel( for index_block, block in enumerate(self.single_transformer_blocks): if torch.is_grad_enabled() and self.gradient_checkpointing: - encoder_hidden_states, hidden_states = self._gradient_checkpointing_func( - block, + def create_custom_forward(module): + def custom_forward(*inputs): + return module(*inputs) + + return custom_forward + ckpt_kwargs: Dict[str, Any] = {"use_reentrant": False} if is_torch_version(">=", "1.11.0") else {} + encoder_hidden_states, hidden_states = torch.utils.checkpoint.checkpoint( + create_custom_forward(block), hidden_states, encoder_hidden_states, temb, image_rotary_emb, joint_attention_kwargs, + **ckpt_kwargs, ) else: