InstructAV2AV / ovi /utils /model_loading_utils.py
suimu's picture
init
e0177dc
Raw
History Blame Contribute Delete
6.79 kB
import torch
import os
import json
from safetensors.torch import load_file
from ovi.modules.fusion import FusionModel
from ovi.modules.t5 import T5EncoderModel
from ovi.modules.vae2_2 import Wan2_2_VAE
from ovi.modules.mmaudio.features_utils import FeaturesUtils
from ovi.distributed_comms.util import get_world_size, get_local_rank, get_global_rank
def init_wan_vae_2_2(ckpt_dir, rank=0):
vae_config = {}
vae_config['device'] = rank
vae_pth = os.path.join(ckpt_dir, "Wan2.2-TI2V-5B/Wan2.2_VAE.pth")
vae_config['vae_pth'] = vae_pth
vae_model = Wan2_2_VAE(**vae_config)
return vae_model
def init_mmaudio_vae(ckpt_dir, rank=0):
vae_config = {}
vae_config['mode'] = '16k'
vae_config['need_vae_encoder'] = True
tod_vae_ckpt = os.path.join(ckpt_dir, "MMAudio/ext_weights/v1-16.pth")
bigvgan_vocoder_ckpt = os.path.join(ckpt_dir, "MMAudio/ext_weights/best_netG.pt")
vae_config['tod_vae_ckpt'] = tod_vae_ckpt
vae_config['bigvgan_vocoder_ckpt'] = bigvgan_vocoder_ckpt
vae = FeaturesUtils(**vae_config).to(rank)
# vae = FeaturesUtils(**vae_config)
return vae
def init_fusion_score_model_ovi(
rank: int = 0,
meta_init=False,
av2av_edit=False,
concat_edit_source_latents=True,
has_video=True,
has_audio=True,
use_siga=False,
):
# import pdb; pdb.set_trace()
if (has_video):
video_config = "ovi/configs/model/dit/video.json"
assert os.path.exists(video_config), f"{video_config} does not exist"
with open(video_config) as f:
video_config = json.load(f)
else:
video_config = None
if (has_audio):
audio_config = "ovi/configs/model/dit/audio.json"
assert os.path.exists(audio_config), f"{audio_config} does not exist"
with open(audio_config) as f:
audio_config = json.load(f)
else:
audio_config = None
# assert os.path.exists(video_config), f"{video_config} does not exist"
# assert os.path.exists(audio_config), f"{audio_config} does not exist"
if meta_init:
with torch.device("meta"):
fusion_model = FusionModel(
video_config,
audio_config,
av2av_edit=av2av_edit,
concat_edit_source_latents=concat_edit_source_latents,
use_siga=use_siga,
)
else:
fusion_model = FusionModel(
video_config,
audio_config,
av2av_edit=av2av_edit,
concat_edit_source_latents=concat_edit_source_latents,
use_siga=use_siga,
)
params_all = sum(p.numel() for p in fusion_model.parameters())
if rank == 0:
print(
f"Score model (Fusion) all parameters:{params_all}"
)
return fusion_model, video_config, audio_config
def init_text_model(ckpt_dir, rank, cpu_offload=False):
wan_dir = os.path.join(ckpt_dir, "Wan2.2-TI2V-5B")
text_encoder_path = os.path.join(wan_dir, "models_t5_umt5-xxl-enc-bf16.pth")
text_tokenizer_path = os.path.join(wan_dir, "google/umt5-xxl")
text_encoder = T5EncoderModel(
text_len=512,
dtype=torch.bfloat16,
device=rank,
checkpoint_path=text_encoder_path,
tokenizer_path=text_tokenizer_path,
cpu_offload=cpu_offload,
shard_fn=None)
return text_encoder
def _maybe_adapt_patch_embedding_weight(key, value, target_shape):
if not (key.endswith("patch_embedding.weight") or key.endswith("patch_embedding.0.weight")):
return None, None
if len(value.shape) != len(target_shape):
return None, None
if value.shape[0] != target_shape[0] or value.shape[2:] != target_shape[2:]:
return None, None
ckpt_in_dim = value.shape[1]
target_in_dim = target_shape[1]
if ckpt_in_dim == target_in_dim * 2:
return value[:, :target_in_dim, ...].contiguous(), "cropped leading input channels"
return None, None
def _adapt_state_dict_for_model(model, state_dict):
model_state_dict = model.state_dict()
adapted_state_dict = {}
adapted_messages = []
skipped_messages = []
for key, value in state_dict.items():
target_value = model_state_dict.get(key)
if target_value is None or not hasattr(value, "shape"):
adapted_state_dict[key] = value
continue
if value.shape == target_value.shape:
adapted_state_dict[key] = value
continue
adapted_value, note = _maybe_adapt_patch_embedding_weight(key, value, target_value.shape)
if adapted_value is not None:
adapted_state_dict[key] = adapted_value
adapted_messages.append(
f"{key}: {tuple(value.shape)} -> {tuple(target_value.shape)} ({note})"
)
continue
skipped_messages.append(
f"{key}: checkpoint {tuple(value.shape)} != model {tuple(target_value.shape)}"
)
return adapted_state_dict, adapted_messages, skipped_messages
def load_fusion_checkpoint(model, checkpoint_path, from_meta=False):
# import pdb; pdb.set_trace()
if checkpoint_path and os.path.exists(checkpoint_path):
if checkpoint_path.endswith(".safetensors"):
df = load_file(checkpoint_path, device="cpu")
elif checkpoint_path.endswith(".pt"):
try:
df = torch.load(checkpoint_path, map_location="cpu", weights_only=False)
df = df['module'] if 'module' in df else df
except Exception as e:
df = torch.load(checkpoint_path, map_location="cpu", weights_only=True)
df = df['app']['model']
else:
raise RuntimeError("We only support .safetensors and .pt checkpoints")
df, adapted_messages, skipped_messages = _adapt_state_dict_for_model(model, df)
missing, unexpected = model.load_state_dict(df, strict=False, assign=from_meta)
if (get_local_rank() == 0):
print("****************************************************")
if adapted_messages:
print("adapted keys:")
for message in adapted_messages:
print(message)
if skipped_messages:
print("skipped mismatched keys:")
for message in skipped_messages:
print(message)
print(f"missing keys: {missing}")
print(f"unexpected keys: {unexpected}")
print("****************************************************")
del df
import gc
gc.collect()
print(f"Successfully loaded fusion checkpoint from {checkpoint_path}")
else:
raise RuntimeError("{checkpoint=} does not exists'")