Add Wav2VecModelLoader to load wav2vec2 from single .safetensors
https://huggingface.co/Kijai/wav2vec2_safetensors/
This commit is contained in:
+77
-3
@@ -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"
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
Reference in New Issue
Block a user