diff --git a/multitalk/nodes.py b/multitalk/nodes.py index b4fcb04..e4ee204 100644 --- a/multitalk/nodes.py +++ b/multitalk/nodes.py @@ -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,18 +21,16 @@ 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) - 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} + 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_window=5 intermediate_dim=512 @@ -52,18 +48,14 @@ 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: diff --git a/nodes_model_loading.py b/nodes_model_loading.py index 9cc68f3..1c7ce4c 100644 --- a/nodes_model_loading.py +++ b/nodes_model_loading.py @@ -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,9 +1034,8 @@ 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: @@ -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")