Allow using GGUF Multi/InfiniteTalk models
This commit is contained in:
+7
-15
@@ -12,9 +12,7 @@ class MultiTalkModelLoader:
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"model": (folder_paths.get_filename_list("diffusion_models"), {"tooltip": "These models are loaded from the 'ComfyUI/models/diffusion_models' -folder",}),
|
||||
|
||||
"base_precision": (["fp32", "bf16", "fp16"], {"default": "fp16"}),
|
||||
"model": (folder_paths.get_filename_list("unet_gguf") + folder_paths.get_filename_list("diffusion_models"), {"tooltip": "These models are loaded from the 'ComfyUI/models/diffusion_models' -folder",}),
|
||||
},
|
||||
}
|
||||
|
||||
@@ -23,19 +21,17 @@ class MultiTalkModelLoader:
|
||||
FUNCTION = "loadmodel"
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
|
||||
def loadmodel(self, model, base_precision):
|
||||
def loadmodel(self, model, base_precision=None):
|
||||
from .multitalk import AudioProjModel
|
||||
|
||||
device = mm.get_torch_device()
|
||||
offload_device = mm.unet_offload_device()
|
||||
base_dtype = {"fp8_e4m3fn": torch.float8_e4m3fn, "fp8_e4m3fn_fast": torch.float8_e4m3fn, "bf16": torch.bfloat16, "fp16": torch.float16, "fp16_fast": torch.float16, "fp32": torch.float32}[base_precision]
|
||||
|
||||
model_path = folder_paths.get_full_path_or_raise("diffusion_models", model)
|
||||
if model_path.endswith(".gguf"):
|
||||
from diffusers.models.model_loading_utils import load_gguf_checkpoint
|
||||
sd = load_gguf_checkpoint(model_path)
|
||||
else:
|
||||
sd = load_torch_file(model_path, device=offload_device, safe_load=True)
|
||||
|
||||
audio_proj_keys = [k for k in sd.keys() if "audio_proj" in k]
|
||||
audio_proj_sd = {k.replace("audio_proj.", ""): sd.pop(k) for k in audio_proj_keys}
|
||||
|
||||
audio_window=5
|
||||
intermediate_dim=512
|
||||
output_dim=768
|
||||
@@ -52,19 +48,15 @@ class MultiTalkModelLoader:
|
||||
context_tokens=context_tokens,
|
||||
norm_output_audio=norm_output_audio,
|
||||
)
|
||||
#fantasytalking_proj_model.load_state_dict(sd, strict=False)
|
||||
|
||||
for name, param in multitalk_proj_model.named_parameters():
|
||||
set_module_tensor_to_device(multitalk_proj_model, name, device=offload_device, dtype=base_dtype, value=audio_proj_sd[name])
|
||||
|
||||
multitalk = {
|
||||
"proj_model": multitalk_proj_model,
|
||||
"sd": sd,
|
||||
"is_gguf": model_path.endswith(".gguf")
|
||||
}
|
||||
|
||||
return (multitalk,)
|
||||
|
||||
|
||||
def loudness_norm(audio_array, sr=16000, lufs=-23):
|
||||
try:
|
||||
import pyloudnorm
|
||||
|
||||
@@ -1012,6 +1012,8 @@ class WanVideoModelLoader:
|
||||
sd.update(ip_adapter_sd)
|
||||
|
||||
if multitalk_model is not None:
|
||||
if multitalk_model["is_gguf"] and not gguf:
|
||||
raise ValueError("Multitalk/InfiniteTalk model is a GGUF model, main model also has to be a GGUF model.")
|
||||
# init audio module
|
||||
from .multitalk.multitalk import SingleStreamMultiAttention
|
||||
from .wanvideo.modules.model import WanRMSNorm, WanLayerNorm
|
||||
@@ -1032,10 +1034,9 @@ class WanVideoModelLoader:
|
||||
)
|
||||
block.norm_x = WanLayerNorm(dim, transformer.eps, elementwise_affine=True) if norm_input_visual else nn.Identity()
|
||||
log.info("MultiTalk model detected, patching model...")
|
||||
|
||||
transformer.audio_proj = multitalk_model["proj_model"]
|
||||
sd.update(multitalk_model["sd"])
|
||||
|
||||
|
||||
# Additional cond latents
|
||||
if "add_conv_in.weight" in sd:
|
||||
def zero_module(module):
|
||||
@@ -1236,9 +1237,6 @@ class WanVideoModelLoader:
|
||||
|
||||
del sd
|
||||
|
||||
if multitalk_model is not None:
|
||||
transformer.audio_proj = multitalk_model["proj_model"]
|
||||
|
||||
if vram_management_args is not None:
|
||||
if gguf:
|
||||
raise ValueError("GGUF models don't support vram management")
|
||||
|
||||
Reference in New Issue
Block a user