diff --git a/multitalk/nodes.py b/multitalk/nodes.py index 0ecea7d..8b3016d 100644 --- a/multitalk/nodes.py +++ b/multitalk/nodes.py @@ -3,9 +3,81 @@ from comfy import model_management as mm from comfy.utils import load_torch_file, common_upscale from accelerate import init_empty_weights import torch -from ..utils import log +from ..utils import log, set_module_tensor_to_device +import os +import json +script_directory = os.path.dirname(os.path.abspath(__file__)) +folder_paths.add_model_folder_path("wav2vec", os.path.join(folder_paths.models_dir, "wav2vec")) +class Wav2VecModelLoader: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "model": (folder_paths.get_filename_list("wav2vec"), {"tooltip": "These models are loaded from the 'ComfyUI/models/wav2vec' -folder",}), + "base_precision": (["fp32", "bf16", "fp16"], {"default": "fp16"}), + "load_device": (["main_device", "offload_device"], {"default": "main_device", "tooltip": "Initial device to load the model to, NOT recommended with the larger models unless you have 48GB+ VRAM"}), + }, + + } + + + RETURN_TYPES = ("WAV2VECMODEL",) + RETURN_NAMES = ("wav2vec_model", ) + FUNCTION = "loadmodel" + CATEGORY = "WanVideoWrapper" + + def loadmodel(self, model, base_precision, load_device): + from transformers import Wav2Vec2Config, Wav2Vec2FeatureExtractor + from ..multitalk.wav2vec2 import Wav2Vec2Model as MultiTalkWav2Vec2Model + + 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] + device = mm.get_torch_device() + offload_device = mm.unet_offload_device() + + if load_device == "offload_device": + transfomer_load_device = offload_device + else: + transfomer_load_device = device + + config_path = os.path.join(script_directory, "wav2vec_config.json") + wav2vec_config = Wav2Vec2Config(**json.load(open(config_path))) + + with init_empty_weights(): + wav2vec = MultiTalkWav2Vec2Model(wav2vec_config).eval() + + feature_extractor_config = { + "do_normalize": False, + "feature_size": 1, + "padding_side": "right", + "padding_value": 0.0, + "return_attention_mask": False, + "sampling_rate": 16000 + } + wav2vec_feature_extractor = Wav2Vec2FeatureExtractor(**feature_extractor_config) + + model_path = folder_paths.get_full_path_or_raise("wav2vec", model) + sd = load_torch_file(model_path, device=transfomer_load_device, safe_load=True) + + for name, param in wav2vec.named_parameters(): + key = "wav2vec2." + name + if "original0" in name: + key = "wav2vec2.encoder.pos_conv_embed.conv.weight_g" + elif "original1" in name: + key = "wav2vec2.encoder.pos_conv_embed.conv.weight_v" + value=sd[key] + set_module_tensor_to_device(wav2vec, name, device=offload_device, dtype=base_dtype, value=value) + + wav2vec_processor_model = { + "feature_extractor": wav2vec_feature_extractor, + "model": wav2vec, + "dtype": base_dtype, + "model_type": "tencent", + } + + return (wav2vec_processor_model,) + class MultiTalkModelLoader: @classmethod def INPUT_TYPES(s): @@ -336,11 +408,13 @@ class WanVideoImageToVideoMultiTalk: NODE_CLASS_MAPPINGS = { "MultiTalkModelLoader": MultiTalkModelLoader, "MultiTalkWav2VecEmbeds": MultiTalkWav2VecEmbeds, - "WanVideoImageToVideoMultiTalk": WanVideoImageToVideoMultiTalk + "WanVideoImageToVideoMultiTalk": WanVideoImageToVideoMultiTalk, + "Wav2VecModelLoader": Wav2VecModelLoader } NODE_DISPLAY_NAME_MAPPINGS = { "MultiTalkModelLoader": "Multi/InfiniteTalk Model Loader", "MultiTalkWav2VecEmbeds": "Multi/InfiniteTalk Wav2Vec Embeds", - "WanVideoImageToVideoMultiTalk": "WanVideo Long I2V Multi/InfiniteTalk" + "WanVideoImageToVideoMultiTalk": "WanVideo Long I2V Multi/InfiniteTalk", + "Wav2VecModelLoader": "Wav2Vec Model Loader" } \ No newline at end of file diff --git a/multitalk/wav2vec_config.json b/multitalk/wav2vec_config.json new file mode 100644 index 0000000..74fe676 --- /dev/null +++ b/multitalk/wav2vec_config.json @@ -0,0 +1,105 @@ +{ + "activation_dropout": 0.1, + "adapter_kernel_size": 3, + "adapter_stride": 2, + "add_adapter": false, + "apply_spec_augment": true, + "architectures": [ + "Wav2Vec2ForPreTraining" + ], + "attention_dropout": 0.1, + "bos_token_id": 1, + "classifier_proj_size": 256, + "codevector_dim": 256, + "contrastive_logits_temperature": 0.1, + "conv_bias": false, + "conv_dim": [ + 512, + 512, + 512, + 512, + 512, + 512, + 512 + ], + "conv_kernel": [ + 10, + 3, + 3, + 3, + 3, + 2, + 2 + ], + "conv_stride": [ + 5, + 2, + 2, + 2, + 2, + 2, + 2 + ], + "ctc_loss_reduction": "sum", + "ctc_zero_infinity": false, + "diversity_loss_weight": 0.1, + "do_stable_layer_norm": false, + "eos_token_id": 2, + "feat_extract_activation": "gelu", + "feat_extract_norm": "group", + "feat_proj_dropout": 0.0, + "feat_quantizer_dropout": 0.0, + "final_dropout": 0.1, + "hidden_act": "gelu", + "hidden_dropout": 0.1, + "hidden_size": 768, + "initializer_range": 0.02, + "intermediate_size": 3072, + "layer_norm_eps": 1e-05, + "layerdrop": 0.1, + "mask_feature_length": 10, + "mask_feature_min_masks": 0, + "mask_feature_prob": 0.0, + "mask_time_length": 10, + "mask_time_min_masks": 2, + "mask_time_prob": 0.05, + "model_type": "wav2vec2", + "num_adapter_layers": 3, + "num_attention_heads": 12, + "num_codevector_groups": 2, + "num_codevectors_per_group": 320, + "num_conv_pos_embedding_groups": 16, + "num_conv_pos_embeddings": 128, + "num_feat_extract_layers": 7, + "num_hidden_layers": 12, + "num_negatives": 100, + "output_hidden_size": 768, + "pad_token_id": 0, + "proj_codevector_dim": 256, + "tdnn_dilation": [ + 1, + 2, + 3, + 1, + 1 + ], + "tdnn_dim": [ + 512, + 512, + 512, + 512, + 1500 + ], + "tdnn_kernel": [ + 5, + 3, + 3, + 1, + 1 + ], + "torch_dtype": "float32", + "transformers_version": "4.16.2", + "use_weighted_layer_sum": false, + "vocab_size": 32, + "xvector_output_dim": 512 +} diff --git a/nodes.py b/nodes.py index 7c355ae..87e5126 100644 --- a/nodes.py +++ b/nodes.py @@ -2013,9 +2013,6 @@ class WanVideoSampler: minimax_latents = minimax_latents.to(device, dtype) minimax_mask_latents = minimax_mask_latents.to(device, dtype) - # Stand-In - standin_input = image_embeds.get("standin_input", None) - # Context windows is_looped = False context_reference_latent = None @@ -2310,7 +2307,12 @@ class WanVideoSampler: if isinstance(rope_function, dict): ntk_alphas = rope_function["ntk_scale_f"], rope_function["ntk_scale_h"], rope_function["ntk_scale_w"] rope_function = rope_function["rope_function"] - + + # Stand-In + standin_input = image_embeds.get("standin_input", None) + if standin_input is not None: + rope_function = "comfy" # only works with this currently + freqs = None transformer.rope_embedder.k = None transformer.rope_embedder.num_frames = None