Add VibeVoice-ASR

This commit is contained in:
Zhiliang Peng
2026-01-21 22:18:33 +08:00
committed by GitHub
parent 6c7369bb31
commit 56cb11e7b2
14 changed files with 4062 additions and 94 deletions
+104 -1
View File
@@ -240,9 +240,112 @@ class VibeVoiceConfig(PretrainedConfig):
super().__init__(**kwargs)
class VibeVoiceASRConfig(PretrainedConfig):
model_type = "vibevoice"
is_composition = True
sub_configs = {
"acoustic_tokenizer_config": VibeVoiceAcousticTokenizerConfig,
"semantic_tokenizer_config": VibeVoiceSemanticTokenizerConfig,
"decoder_config": Qwen2Config,
}
# keys_to_ignore_at_inference = ["past_key_values"]
# Default tensor parallel plan for base model `Qwen2`
base_model_tp_plan = {
"layers.*.self_attn.q_proj": "colwise",
"layers.*.self_attn.k_proj": "colwise",
"layers.*.self_attn.v_proj": "colwise",
"layers.*.self_attn.o_proj": "rowwise",
"layers.*.mlp.gate_proj": "colwise",
"layers.*.mlp.up_proj": "colwise",
"layers.*.mlp.down_proj": "rowwise",
}
def __init__(
self,
acoustic_tokenizer_config=None,
semantic_tokenizer_config=None,
decoder_config=None,
**kwargs
):
# kwargs["_attn_implementation"] = "flash_attention_2"
kwargs["_attn_implementation_autoset"] = False
if acoustic_tokenizer_config is None:
self.acoustic_tokenizer_config = self.sub_configs["acoustic_tokenizer_config"]()
elif isinstance(acoustic_tokenizer_config, dict):
acoustic_tokenizer_config["model_type"] = "vibevoice_acoustic_tokenizer"
self.acoustic_tokenizer_config = self.sub_configs["acoustic_tokenizer_config"](**acoustic_tokenizer_config)
elif isinstance(acoustic_tokenizer_config, VibeVoiceAcousticTokenizerConfig):
# If an instance of the config class is provided
self.acoustic_tokenizer_config = acoustic_tokenizer_config
if semantic_tokenizer_config is None:
self.semantic_tokenizer_config = self.sub_configs["semantic_tokenizer_config"]()
elif isinstance(semantic_tokenizer_config, dict):
semantic_tokenizer_config["model_type"] = "vibevoice_semantic_tokenizer"
self.semantic_tokenizer_config = self.sub_configs["semantic_tokenizer_config"](**semantic_tokenizer_config)
elif isinstance(semantic_tokenizer_config, VibeVoiceSemanticTokenizerConfig):
# If an instance of the config class is provided
self.semantic_tokenizer_config = semantic_tokenizer_config
if decoder_config is None:
self.decoder_config = self.sub_configs["decoder_config"]()
elif isinstance(decoder_config, dict):
# If a dictionary is provided, instantiate the config class with it
# self.decoder_config = self.sub_configs["decoder_config"](**decoder_config)
if decoder_config.get("model_type", '') == "qwen2":
self.decoder_config = Qwen2Config(**decoder_config)
else:
raise ValueError(f"Unsupported decoder model type: {decoder_config.get('model_type', '')}")
elif isinstance(decoder_config, Qwen2Config):
# If an instance of the config class is provided
self.decoder_config = decoder_config
# other parameters
self.acoustic_vae_dim = getattr(self.acoustic_tokenizer_config, 'vae_dim', 64)
self.semantic_vae_dim = getattr(self.semantic_tokenizer_config, 'vae_dim', 128)
super().__init__(**kwargs)
def get_text_config(self, decoder: bool = False):
"""Return the text (decoder) config for generation."""
return self.decoder_config
@property
def vocab_size(self):
"""Return vocab_size from decoder config for generation compatibility."""
return self.decoder_config.vocab_size
@property
def num_attention_heads(self):
"""Return num_attention_heads from decoder config for Ulysses SP compatibility."""
return self.decoder_config.num_attention_heads
@property
def num_key_value_heads(self):
"""Return num_key_value_heads from decoder config for Ulysses SP compatibility."""
return self.decoder_config.num_key_value_heads
@property
def hidden_size(self):
"""Return hidden_size from decoder config for model compatibility."""
return self.decoder_config.hidden_size
@property
def num_hidden_layers(self):
"""Return num_hidden_layers from decoder config for Ulysses SP compatibility."""
return self.decoder_config.num_hidden_layers
@property
def head_dim(self):
"""Return head_dim from decoder config for Ulysses SP compatibility."""
return getattr(self.decoder_config, 'head_dim', self.hidden_size // self.num_attention_heads)
__all__ = [
"VibeVoiceAcousticTokenizerConfig",
"VibeVoiceSemanticTokenizerConfig",
"VibeVoiceDiffusionHeadConfig",
"VibeVoiceConfig"
"VibeVoiceConfig",
"VibeVoiceASRConfig"
]