Merge branch 'main' into dev
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",}),
|
||||
},
|
||||
}
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user