Add Wav2VecModelLoader to load wav2vec2 from single .safetensors

https://huggingface.co/Kijai/wav2vec2_safetensors/
This commit is contained in:
kijai
2025-08-26 00:09:46 +03:00
parent e836134b90
commit 6fce0e2d3b
3 changed files with 188 additions and 7 deletions
+77 -3
View File
@@ -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"
}
+105
View File
@@ -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
}