diff --git a/multitalk/nodes.py b/multitalk/nodes.py index ed8aba5..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",}), }, } @@ -25,16 +23,14 @@ class MultiTalkModelLoader: 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 1b39dd0..be893ab 100644 --- a/nodes_model_loading.py +++ b/nodes_model_loading.py @@ -1209,6 +1209,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 @@ -1231,7 +1233,6 @@ class WanVideoModelLoader: 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: @@ -1312,6 +1313,9 @@ 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")