Fix VAE and transformer loading for 5B model (#297)
This commit is contained in:
@@ -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']),
|
||||
|
||||
@@ -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']),
|
||||
|
||||
@@ -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']),
|
||||
|
||||
@@ -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']),
|
||||
|
||||
Reference in New Issue
Block a user