From 8f8c85693d43a881dc32673c4a54aa2e112a34f6 Mon Sep 17 00:00:00 2001 From: Leojc Date: Wed, 27 Aug 2025 15:28:01 +0800 Subject: [PATCH] Fix VAE and transformer loading for 5B model (#297) --- scripts/wan2.2_fun/train.py | 16 +++++++++++----- scripts/wan2.2_fun/train_control.py | 14 ++++++++++---- scripts/wan2.2_fun/train_control_lora.py | 14 ++++++++++---- scripts/wan2.2_fun/train_lora.py | 14 ++++++++++---- 4 files changed, 41 insertions(+), 17 deletions(-) diff --git a/scripts/wan2.2_fun/train.py b/scripts/wan2.2_fun/train.py index 3f813ae..0c9ca9e 100644 --- a/scripts/wan2.2_fun/train.py +++ b/scripts/wan2.2_fun/train.py @@ -73,7 +73,7 @@ from videox_fun.data.bucket_sampler import (ASPECT_RATIO_512, from videox_fun.data.dataset_image_video import (ImageVideoDataset, ImageVideoSampler, get_random_mask) -from videox_fun.models import (AutoencoderKLWan, CLIPModel, WanT5EncoderModel, +from videox_fun.models import (AutoencoderKLWan, AutoencoderKLWan3_8, CLIPModel, WanT5EncoderModel, Wan2_2Transformer3DModel) from videox_fun.pipeline import WanFunInpaintPipeline, WanFunPipeline from videox_fun.utils.discrete_sampler import DiscreteSampling @@ -862,15 +862,21 @@ 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']), ) vae.eval() - + # Get Transformer - sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer') \ - if args.boundary_type == "low" else config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer') + if args.boundary_type == "low" or args.boundary_type == "full": + sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer') + else: + sub_path = config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer') transformer3d = 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']), diff --git a/scripts/wan2.2_fun/train_control.py b/scripts/wan2.2_fun/train_control.py index dd8e387..c302d4d 100644 --- a/scripts/wan2.2_fun/train_control.py +++ b/scripts/wan2.2_fun/train_control.py @@ -73,7 +73,7 @@ from videox_fun.data.dataset_image_video import (ImageVideoControlDataset, get_random_mask, process_pose_file, process_pose_params) -from videox_fun.models import (AutoencoderKLWan, CLIPModel, WanT5EncoderModel, +from videox_fun.models import (AutoencoderKLWan, AutoencoderKLWan3_8, CLIPModel, WanT5EncoderModel, Wan2_2Transformer3DModel) from videox_fun.pipeline import WanFunControlPipeline from videox_fun.utils.discrete_sampler import DiscreteSampling @@ -837,15 +837,21 @@ 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']), ) vae.eval() # Get Transformer - sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer') \ - if args.boundary_type == "low" else config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer') + if args.boundary_type == "low" or args.boundary_type == "full": + sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer') + else: + sub_path = config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer') transformer3d = 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']), diff --git a/scripts/wan2.2_fun/train_control_lora.py b/scripts/wan2.2_fun/train_control_lora.py index 611454d..c378e0f 100644 --- a/scripts/wan2.2_fun/train_control_lora.py +++ b/scripts/wan2.2_fun/train_control_lora.py @@ -73,7 +73,7 @@ from videox_fun.data.dataset_image_video import (ImageVideoControlDataset, get_random_mask, process_pose_file, process_pose_params) -from videox_fun.models import (AutoencoderKLWan, CLIPModel, WanT5EncoderModel, +from videox_fun.models import (AutoencoderKLWan, AutoencoderKLWan3_8, CLIPModel, WanT5EncoderModel, Wan2_2Transformer3DModel) from videox_fun.pipeline import WanFunControlPipeline from videox_fun.utils.discrete_sampler import DiscreteSampling @@ -832,15 +832,21 @@ 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']), ) vae.eval() # Get Transformer - sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer') \ - if args.boundary_type == "low" else config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer') + if args.boundary_type == "low" or args.boundary_type == "full": + sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer') + else: + sub_path = config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer') transformer3d = 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']), diff --git a/scripts/wan2.2_fun/train_lora.py b/scripts/wan2.2_fun/train_lora.py index 4b31342..04f7655 100644 --- a/scripts/wan2.2_fun/train_lora.py +++ b/scripts/wan2.2_fun/train_lora.py @@ -69,7 +69,7 @@ from videox_fun.data.bucket_sampler import (ASPECT_RATIO_512, from videox_fun.data.dataset_image_video import (ImageVideoDataset, ImageVideoSampler, get_random_mask) -from videox_fun.models import (AutoencoderKLWan, WanT5EncoderModel, +from videox_fun.models import (AutoencoderKLWan, AutoencoderKLWan3_8, WanT5EncoderModel, Wan2_2Transformer3DModel) from videox_fun.pipeline import WanFunInpaintPipeline, WanFunPipeline from videox_fun.utils.discrete_sampler import DiscreteSampling @@ -852,15 +852,21 @@ 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']), ) vae.eval() # Get Transformer - sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer') \ - if args.boundary_type == "low" else config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer') + if args.boundary_type == "low" or args.boundary_type == "full": + sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer') + else: + sub_path = config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer') transformer3d = 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']),