Optimize 0.5B integration: Cleanup debug prints, update docs, mark Preset Maker as experimental
This commit is contained in:
@@ -520,11 +520,20 @@ git clone https://github.com/havvk/ComfyUI_AIIA.git
|
||||
- **必选参数**: `voice_preset` (音色预设) - **必须选择**。
|
||||
- **功能**: 极速实时生成。基于预计算的 `.pt` 缓存文件生成语音。
|
||||
- **不支持**: `reference_audio` (直接克隆)。
|
||||
- **如何获取中文预设?**: 请使用配套的 `🎤 VibeVoice Preset Maker` 节点自行制作。
|
||||
- **特点**:
|
||||
- **极低延迟**: 首包延迟极低,适合即时交互。
|
||||
- **BF16 加速**: 自动使用 Bloat16 精度进行推理(如果硬件支持),大幅提升速度。
|
||||
- **多语言支持**: 官方预设涵盖英、日、韩、法、德等。
|
||||
|
||||
#### 3. 🎤 VibeVoice Preset Maker (0.5B)
|
||||
- **用途**: 制作 0.5B 模型专用的 `.pt` 音色预设。
|
||||
- **流程**: 连接参考音频 -> 运行节点 ->生成预设 -> 重启 ComfyUI -> 在 `Realtime 0.5B` 节点中使用。
|
||||
#### 3. 🎤 VibeVoice Preset Maker (0.5B) (Experimental ⚠️)
|
||||
|
||||
- **用途**: 尝试制作 0.5B 模型专用的 `.pt` 音色预设。
|
||||
- **现状**: **极不稳定**。
|
||||
- **原因**: 社区反馈和测试表明,VibeVoice-Realtime-0.5B 模型的权重似乎对自定义音色进行了限制或未进行充分的 Zero-Shot 泛化训练。即使使用长达 1 分钟的高质量音频,生成时也极易出现**死循环、胡言乱语或噪音**。
|
||||
- **建议**:
|
||||
- **首选**: 请直接下载并在 `Realtime 0.5B` 节点中使用 **微软官方提供的预设** (Carter, Emma 等)。
|
||||
- **尝试**: 如果您一定要克隆音色,请使用 **VibeVoice 1.5B / 7B (Standard)** 节点,它们原生支持完美的 Zero-Shot 克隆。
|
||||
- **仅供研究**: 此节点保留给开发者进行研究调试,普通用户**不推荐**使用。
|
||||
|
||||
|
||||
**手动下载命令**:
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
import os
|
||||
|
||||
class AIIA_VibeVoice_Preset_Maker:
|
||||
@classmethod
|
||||
@@ -47,48 +48,91 @@ class AIIA_VibeVoice_Preset_Maker:
|
||||
# `_create_voice_prompt` calls `self.audio_feature_extractor`.
|
||||
# Taking a cue from `AIIA_VibeVoice_TTS`, raw audio is fine if we pass it correctly.
|
||||
|
||||
|
||||
# We need numpy array for processor
|
||||
if waveform.ndim == 3: # [B, C, T]
|
||||
waveform = waveform.mean(dim=1) # Mix to mono
|
||||
# ComfyUI AUDIO is typically [Batch, Samples, Channels] or [Batch, Samples]
|
||||
if waveform.ndim == 3:
|
||||
# Check heuristic: if dim 2 is small (channels) and dim 1 is large (time) -> [B, T, C]
|
||||
if waveform.shape[2] < waveform.shape[1]:
|
||||
# Input is [B, T, C]. Mix query channels (dim 2) to mono
|
||||
waveform = waveform.mean(dim=2) # -> [B, T]
|
||||
else:
|
||||
# Input is [B, C, T]. Mix query channels (dim 1) to mono
|
||||
waveform = waveform.mean(dim=1) # -> [B, T]
|
||||
|
||||
if waveform.ndim == 2: # [B, T] -> [T] (take first batch)
|
||||
audio_np = waveform[0].cpu().numpy()
|
||||
else:
|
||||
audio_np = waveform.cpu().numpy()
|
||||
|
||||
# Resample if needed (processor expects specific SR? Usually handled by feature extractor, but let's assume raw is okay if passed properly)
|
||||
# Actually, in `AIIA_VibeVoice_TTS`, we verified 0.5B doesn't use `reference_audio` which is why we are here.
|
||||
# But the processor CAN handle it.
|
||||
waveform = waveform[0]
|
||||
|
||||
# Resample to Processor's expected SR
|
||||
target_sr = 22050 # Default for VibeVoice / Encodec
|
||||
if hasattr(processor, "feature_extractor") and hasattr(processor.feature_extractor, "sampling_rate"):
|
||||
target_sr = processor.feature_extractor.sampling_rate
|
||||
|
||||
if sample_rate != target_sr:
|
||||
print(f"[AIIA] Resampling audio from {sample_rate} to {target_sr}")
|
||||
resampler = torchaudio.transforms.Resample(orig_freq=sample_rate, new_freq=target_sr)
|
||||
waveform = waveform.cpu() # ensure cpu for torchaudio transforms usually
|
||||
if waveform.dim() == 1: waveform = waveform.unsqueeze(0) # [1, T]
|
||||
waveform = resampler(waveform)
|
||||
waveform = waveform.squeeze(0) # [T]
|
||||
|
||||
audio_np = waveform.cpu().numpy()
|
||||
|
||||
# KEY FIX: Normalize audio explicitly for consistency with processor
|
||||
# _create_voice_prompt normalizes internally, but prepare_speech_inputs DOES NOT.
|
||||
# We must normalize here so both share the same volume level (-25dB).
|
||||
if hasattr(processor, "audio_normalizer") and processor.audio_normalizer:
|
||||
print("[AIIA] Normalizing audio volume to -25dB...")
|
||||
audio_np = processor.audio_normalizer(audio_np)
|
||||
|
||||
# 3. Construct Prompt (Imitating `_process_single`)
|
||||
# System Prompt
|
||||
system_prompt = " Transform the text provided by various speakers into speech output, utilizing the distinct voice of each respective speaker.\n"
|
||||
system_tokens = tokenizer.encode(system_prompt, add_special_tokens=False) # Qwen tokenizer adds BOS? `_process_single` doesn't explicitly add=False, likely default True?
|
||||
# Wait, `tokenizer.encode(system_prompt)` in `_process_single` line 226.
|
||||
# `_process_single` line 360 calls `tokenizer.encode(..., add_special_tokens=False)`.
|
||||
# Let's assume standard encode for system prompt uses defaults (likely adds BOS).
|
||||
|
||||
system_tokens = tokenizer.encode(system_prompt, add_special_tokens=False)
|
||||
|
||||
# Voice Prompt
|
||||
# We need to call internal `_create_voice_prompt`
|
||||
# It expects `voice_samples` as list of strings (paths) or arrays.
|
||||
voice_tokens, voice_speech_inputs, voice_speech_masks = processor._create_voice_prompt([audio_np])
|
||||
|
||||
# Debug Voice Tokens
|
||||
if len(voice_tokens) > 0:
|
||||
print(f"[AIIA] Generated {len(voice_tokens)} Voice Tokens. Range: [{min(voice_tokens)}, {max(voice_tokens)}]")
|
||||
else:
|
||||
print("[AIIA] WARNING: Generated 0 Voice Tokens! Input audio might be silent or too short.")
|
||||
|
||||
# Header
|
||||
header_text = ' Text input:\n'
|
||||
header_tokens = tokenizer.encode(header_text, add_special_tokens=False)
|
||||
|
||||
# Combine for MAIN Prompt (Input IDs)
|
||||
# System + Voice + Header
|
||||
prompt_tokens = system_tokens + voice_tokens + header_tokens
|
||||
|
||||
# Masks
|
||||
# System (text) + Voice (speech) + Header (text)
|
||||
# Text parts are False in speech_input_mask
|
||||
speech_input_masks = [False] * len(system_tokens) + voice_speech_masks + [False] * len(header_tokens)
|
||||
# Add " Speaker 0:" suffix (Crucial for prompt continuity!)
|
||||
prefix_text = " Speaker 0:"
|
||||
prefix_tokens = tokenizer.encode(prefix_text, add_special_tokens=False)
|
||||
|
||||
# KEY FIX: Separate inputs for LM and TTS_LM
|
||||
# LM (Semantic) sees ONLY text. TTS_LM sees Text + Speech.
|
||||
# This explains the cache size discrepancy (108 vs 316) in official presets.
|
||||
|
||||
# LM Input: System + Header + Prefix
|
||||
lm_tokens = system_tokens + header_tokens + prefix_tokens
|
||||
|
||||
# TTS LM Input: System + Voice + Header + Prefix
|
||||
tts_lm_tokens = system_tokens + voice_tokens + header_tokens + prefix_tokens
|
||||
|
||||
# Masks (aligned with tts_lm_tokens)
|
||||
# System (text) + Voice (speech) + Header (text) + Prefix (text)
|
||||
# speech_input_masks: Text=False, Speech=True
|
||||
# tts_text_masks: Text=1, Speech=0
|
||||
speech_input_masks = [False] * len(system_tokens) + voice_speech_masks + [False] * len(header_tokens) + [False] * len(prefix_tokens)
|
||||
tts_text_masks_list = [1] * len(system_tokens) + [0] * len(voice_tokens) + [1] * len(header_tokens) + [1] * len(prefix_tokens)
|
||||
|
||||
# 4. Prepare Batch for Forward Pass
|
||||
# We need to wrap this into batch encoding format
|
||||
input_ids = torch.tensor([prompt_tokens], device=device, dtype=torch.long)
|
||||
input_ids = torch.tensor([lm_tokens], device=device, dtype=torch.long) # For LM
|
||||
tts_lm_input_ids = torch.tensor([tts_lm_tokens], device=device, dtype=torch.long) # For TTS LM
|
||||
|
||||
speech_input_mask_tensor = torch.tensor([speech_input_masks], device=device, dtype=torch.bool)
|
||||
tts_text_masks_tensor = torch.tensor([tts_text_masks_list], device=device, dtype=torch.long)
|
||||
|
||||
# Prepare speech tensors
|
||||
speech_dict = processor.prepare_speech_inputs([audio_np], return_tensors="pt", device=device)
|
||||
@@ -114,29 +158,29 @@ class AIIA_VibeVoice_Preset_Maker:
|
||||
|
||||
# TTS LM Forward
|
||||
# We need to pass `speech_tensors`, `speech_masks`, `speech_input_mask`
|
||||
# `tts_text_masks`? prompt doesn't have TTS text yet.
|
||||
# In streaming processor `__call__`, it calls `_batch_encode` which sets up inputs.
|
||||
# Let's look at `model.forward_tts_lm` signature or usage.
|
||||
# It usually takes `input_ids` (same as LM), `lm_last_hidden_state`, `speech_tensors`...
|
||||
|
||||
# Ensure speech_tensors match model dtype (Half/Float16)
|
||||
if speech_tensors is not None:
|
||||
speech_tensors = speech_tensors.to(dtype=model.dtype, device=model.device)
|
||||
|
||||
tts_lm_out = model.forward_tts_lm(
|
||||
input_ids=input_ids,
|
||||
attention_mask=torch.ones_like(input_ids),
|
||||
input_ids=tts_lm_input_ids, # Use TTS_LM specific input (includes speech tokens)
|
||||
attention_mask=torch.ones_like(tts_lm_input_ids),
|
||||
lm_last_hidden_state=lm_last_hidden,
|
||||
speech_tensors=speech_tensors,
|
||||
speech_masks=speech_masks,
|
||||
speech_input_mask=speech_input_mask_tensor,
|
||||
# tts_text_masks=None, # Prompt doesn't contain TTS text yet
|
||||
tts_text_masks=tts_text_masks_tensor, # Prompt mixed mask (System/Header=1, Voice=0)
|
||||
use_cache=True,
|
||||
return_dict=True
|
||||
)
|
||||
tts_lm_cache = tts_lm_out.past_key_values
|
||||
tts_lm_last_hidden = tts_lm_out.last_hidden_state
|
||||
|
||||
# 6. Negative Cache (Standard "<|image_pad|>" for unconditioned)
|
||||
neg_input_id = tokenizer.convert_tokens_to_ids("<|image_pad|>")
|
||||
# Length must match Positive Cache to align positions!
|
||||
seq_len = input_ids.shape[1]
|
||||
neg_input_ids = torch.full((1, seq_len), neg_input_id, device=device, dtype=torch.long)
|
||||
# FIX: Negative Cache should be length 1 (as seen in official presets), not prompt length!
|
||||
neg_input_ids = torch.full((1, 1), neg_input_id, device=device, dtype=torch.long)
|
||||
|
||||
with torch.no_grad():
|
||||
neg_lm_o = model.forward_lm(
|
||||
@@ -146,6 +190,7 @@ class AIIA_VibeVoice_Preset_Maker:
|
||||
return_dict=True
|
||||
)
|
||||
neg_lm_cache = neg_lm_o.past_key_values
|
||||
neg_lm_last_hidden = neg_lm_o.last_hidden_state
|
||||
|
||||
neg_tts_lm_o = model.forward_tts_lm(
|
||||
input_ids=neg_input_ids,
|
||||
@@ -153,9 +198,12 @@ class AIIA_VibeVoice_Preset_Maker:
|
||||
use_cache=True,
|
||||
return_dict=True,
|
||||
lm_last_hidden_state=neg_lm_o.last_hidden_state,
|
||||
tts_text_masks=torch.ones_like(neg_input_ids) # Treat all as "text" (mask=1) for negative?
|
||||
# FIX: Negative Cache (Padding) is TEXT (Type 1), not Speech (0).
|
||||
# Marking it as Speech prevents Scatter from working (0!=1) and adds wrong Type Embedding.
|
||||
tts_text_masks=torch.ones_like(neg_input_ids)
|
||||
)
|
||||
neg_tts_lm_cache = neg_tts_lm_o.past_key_values
|
||||
neg_tts_lm_last_hidden = neg_tts_lm_o.last_hidden_state
|
||||
|
||||
# 7. Save to PT
|
||||
# Convert to CPU before saving
|
||||
@@ -181,13 +229,16 @@ class AIIA_VibeVoice_Preset_Maker:
|
||||
'last_hidden_state': recursive_to_cpu(lm_last_hidden)
|
||||
},
|
||||
'tts_lm': {
|
||||
'past_key_values': recursive_to_cpu(tts_lm_cache)
|
||||
'past_key_values': recursive_to_cpu(tts_lm_cache),
|
||||
'last_hidden_state': recursive_to_cpu(tts_lm_last_hidden)
|
||||
},
|
||||
'neg_lm': {
|
||||
'past_key_values': recursive_to_cpu(neg_lm_cache)
|
||||
'past_key_values': recursive_to_cpu(neg_lm_cache),
|
||||
'last_hidden_state': recursive_to_cpu(neg_lm_last_hidden)
|
||||
},
|
||||
'neg_tts_lm': {
|
||||
'past_key_values': recursive_to_cpu(neg_tts_lm_cache)
|
||||
'past_key_values': recursive_to_cpu(neg_tts_lm_cache),
|
||||
'last_hidden_state': recursive_to_cpu(neg_tts_lm_last_hidden)
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
+151
-19
@@ -1,3 +1,6 @@
|
||||
import os
|
||||
import torch
|
||||
import torchaudio
|
||||
|
||||
class AIIA_VibeVoice_Realtime_TTS:
|
||||
@classmethod
|
||||
@@ -35,7 +38,7 @@ class AIIA_VibeVoice_Realtime_TTS:
|
||||
CATEGORY = "AIIA/VibeVoice"
|
||||
|
||||
def generate(self, vibevoice_model, text, voice_preset, ddpm_steps, speed, normalize_text,
|
||||
do_sample, temperature, top_k, top_p, cfg_scale, voice_preset_input=None): # cfg_scale added based on standard
|
||||
do_sample, temperature, top_k, top_p, cfg_scale, voice_preset_input=None):
|
||||
model = vibevoice_model["model"]
|
||||
tokenizer = vibevoice_model["tokenizer"]
|
||||
processor = vibevoice_model.get("processor")
|
||||
@@ -71,38 +74,55 @@ class AIIA_VibeVoice_Realtime_TTS:
|
||||
print(f"[AIIA] Using 0.5B Streaming Inference with preset: {voice_preset_name}")
|
||||
|
||||
# Load Preset
|
||||
# Preset path already resolved
|
||||
if not os.path.exists(preset_path):
|
||||
raise FileNotFoundError(f"Preset file not found at: {preset_path}")
|
||||
|
||||
preset_data = torch.load(preset_path, map_location=device)
|
||||
|
||||
# Helper to get attributes safe and move to device
|
||||
# Helper to get attributes safe and move to device + cast to dtype (Recursive for tuples/lists)
|
||||
def get_tensor(obj, key):
|
||||
val = obj.get(key) if isinstance(obj, dict) else getattr(obj, key, None)
|
||||
if isinstance(val, torch.Tensor): return val.to(device)
|
||||
return val
|
||||
|
||||
# Extract Caches
|
||||
# Extract Caches & Hidden States
|
||||
lm_cache = get_tensor(preset_data.get('lm'), 'past_key_values')
|
||||
tts_lm_cache = get_tensor(preset_data.get('tts_lm'), 'past_key_values')
|
||||
lm_last_hidden = get_tensor(preset_data.get('lm'), 'last_hidden_state')
|
||||
|
||||
tts_lm_cache = get_tensor(preset_data.get('tts_lm'), 'past_key_values')
|
||||
tts_lm_last_hidden = get_tensor(preset_data.get('tts_lm'), 'last_hidden_state')
|
||||
|
||||
# Extract Speech Tensors (Critical for Cross-Attention)
|
||||
speech_data = preset_data.get('speech')
|
||||
if speech_data:
|
||||
speech_tensors = get_tensor(speech_data, 'speech_tensors')
|
||||
speech_masks = get_tensor(speech_data, 'speech_masks')
|
||||
speech_input_mask = get_tensor(speech_data, 'speech_input_mask')
|
||||
else:
|
||||
# Fallback for old presets (might fail if model needs them)
|
||||
speech_tensors, speech_masks, speech_input_mask = None, None, None
|
||||
|
||||
if lm_cache is None or tts_lm_cache is None:
|
||||
raise ValueError(f"Invalid preset file: {voice_preset}. Missing cache data.")
|
||||
|
||||
# Extract Negative Caches (if present) or Create Dummy
|
||||
if 'neg_lm' in preset_data:
|
||||
neg_lm_cache = get_tensor(preset_data.get('neg_lm'), 'past_key_values')
|
||||
neg_lm_last_hidden = get_tensor(preset_data.get('neg_lm'), 'last_hidden_state')
|
||||
|
||||
neg_tts_lm_cache = get_tensor(preset_data.get('neg_tts_lm'), 'past_key_values')
|
||||
neg_tts_lm_last_hidden = get_tensor(preset_data.get('neg_tts_lm'), 'last_hidden_state')
|
||||
else:
|
||||
print("[AIIA] Preset missing negative cache, generating on fly...")
|
||||
seq_len = lm_cache[0][0].shape[2]
|
||||
neg_input_id = tokenizer.convert_tokens_to_ids("<|image_pad|>")
|
||||
neg_input_ids = torch.full((1, seq_len), neg_input_id, device=device)
|
||||
|
||||
# Note: This logic assumes model.forward_lm/tts_lm handles everything
|
||||
# But for robustness we should use try/except block if model methods fail?
|
||||
# Assuming original code was functional for this part.
|
||||
neg_lm_o = model.forward_lm(input_ids=neg_input_ids, attention_mask=torch.ones_like(neg_input_ids), use_cache=True, return_dict=True)
|
||||
neg_lm_cache = neg_lm_o.past_key_values
|
||||
neg_lm_last_hidden = neg_lm_o.last_hidden_state
|
||||
|
||||
neg_tts_lm_o = model.forward_tts_lm(
|
||||
input_ids=neg_input_ids,
|
||||
@@ -110,17 +130,94 @@ class AIIA_VibeVoice_Realtime_TTS:
|
||||
use_cache=True,
|
||||
return_dict=True,
|
||||
lm_last_hidden_state=neg_lm_o.last_hidden_state,
|
||||
tts_text_masks=torch.ones_like(neg_input_ids)
|
||||
tts_text_masks=torch.zeros_like(neg_input_ids)
|
||||
)
|
||||
neg_tts_lm_cache = neg_tts_lm_o.past_key_values
|
||||
neg_tts_lm_last_hidden = neg_tts_lm_o.last_hidden_state
|
||||
|
||||
# Universal casting helper
|
||||
def cast_recursive(item, dtype, device):
|
||||
if isinstance(item, torch.Tensor):
|
||||
if item.is_floating_point():
|
||||
return item.to(device=device, dtype=dtype)
|
||||
return item.to(device)
|
||||
|
||||
# Handle DynamicCache by casting internals IN PLACE (preserving object type)
|
||||
try:
|
||||
from transformers.cache_utils import DynamicCache
|
||||
if isinstance(item, DynamicCache):
|
||||
# Qwen2 Strictness: Must be a Cache object. Convert data inside.
|
||||
if hasattr(item, 'key_cache'):
|
||||
item.key_cache = [cast_recursive(k, dtype, device) for k in item.key_cache]
|
||||
if hasattr(item, 'value_cache'):
|
||||
item.value_cache = [cast_recursive(v, dtype, device) for v in item.value_cache]
|
||||
return item
|
||||
except:
|
||||
pass
|
||||
|
||||
if isinstance(item, (list, tuple)):
|
||||
return type(item)(cast_recursive(x, dtype, device) for x in item)
|
||||
return item
|
||||
|
||||
# Robustly get model dtype
|
||||
def get_deep_dtype(model_obj):
|
||||
try:
|
||||
return model_obj.model.language_model.model.layers[0].self_attn.q_proj.weight.dtype
|
||||
except:
|
||||
pass
|
||||
try:
|
||||
return model_obj.model.language_model.dtype
|
||||
except:
|
||||
pass
|
||||
try:
|
||||
return next(model_obj.parameters()).dtype
|
||||
except:
|
||||
return model_obj.dtype
|
||||
|
||||
target_dtype = get_deep_dtype(model)
|
||||
|
||||
# Force Cast ALL caches to model dtype (Crucial for FP16/Mixed Precision)
|
||||
lm_cache = cast_recursive(lm_cache, target_dtype, device)
|
||||
tts_lm_cache = cast_recursive(tts_lm_cache, target_dtype, device)
|
||||
|
||||
# Handle Negative Cache Variables
|
||||
if 'neg_lm_cache' in locals():
|
||||
neg_lm_cache = cast_recursive(neg_lm_cache, target_dtype, device)
|
||||
if 'neg_tts_lm_cache' in locals():
|
||||
neg_tts_lm_cache = cast_recursive(neg_tts_lm_cache, target_dtype, device)
|
||||
|
||||
# Cast Hidden States
|
||||
lm_last_hidden = cast_recursive(lm_last_hidden, target_dtype, device)
|
||||
if 'tts_lm_last_hidden' in locals() and tts_lm_last_hidden is not None:
|
||||
tts_lm_last_hidden = cast_recursive(tts_lm_last_hidden, target_dtype, device)
|
||||
if 'neg_lm_last_hidden' in locals() and neg_lm_last_hidden is not None:
|
||||
neg_lm_last_hidden = cast_recursive(neg_lm_last_hidden, target_dtype, device)
|
||||
if 'neg_tts_lm_last_hidden' in locals() and neg_tts_lm_last_hidden is not None:
|
||||
neg_tts_lm_last_hidden = cast_recursive(neg_tts_lm_last_hidden, target_dtype, device)
|
||||
|
||||
# Cast Speech Tensors
|
||||
if 'speech_tensors' in locals() and speech_tensors is not None:
|
||||
speech_tensors = cast_recursive(speech_tensors, target_dtype, device)
|
||||
if 'speech_masks' in locals() and speech_masks is not None:
|
||||
speech_masks = cast_recursive(speech_masks, torch.bool, device)
|
||||
if 'speech_input_mask' in locals() and speech_input_mask is not None:
|
||||
speech_input_mask = cast_recursive(speech_input_mask, torch.bool, device)
|
||||
|
||||
# Use ModelOutput (fixes BOTH 'not iterable' and 'has no attribute' errors)
|
||||
from transformers.modeling_outputs import ModelOutput
|
||||
|
||||
|
||||
# Helper to create ModelOutput safe
|
||||
def create_output(cache, hidden):
|
||||
if hidden is not None:
|
||||
return ModelOutput(past_key_values=cache, last_hidden_state=hidden)
|
||||
return ModelOutput(past_key_values=cache)
|
||||
|
||||
# Wrap in SimpleNamespace
|
||||
from types import SimpleNamespace
|
||||
all_prefilled = {
|
||||
"lm": SimpleNamespace(past_key_values=lm_cache, last_hidden_state=lm_last_hidden),
|
||||
"tts_lm": SimpleNamespace(past_key_values=tts_lm_cache),
|
||||
"neg_lm": SimpleNamespace(past_key_values=neg_lm_cache),
|
||||
"neg_tts_lm": SimpleNamespace(past_key_values=neg_tts_lm_cache)
|
||||
"lm": create_output(lm_cache, lm_last_hidden),
|
||||
"tts_lm": create_output(tts_lm_cache, tts_lm_last_hidden),
|
||||
"neg_lm": create_output(neg_lm_cache, neg_lm_last_hidden),
|
||||
"neg_tts_lm": create_output(neg_tts_lm_cache, neg_tts_lm_last_hidden)
|
||||
}
|
||||
|
||||
# Target tokens
|
||||
@@ -128,17 +225,45 @@ class AIIA_VibeVoice_Realtime_TTS:
|
||||
target_tokens = tokenizer.encode(target_text.strip() + "\n", add_special_tokens=False, return_tensors="pt").to(device)
|
||||
|
||||
# Dummy inputs
|
||||
cache_len = lm_cache[0][0].shape[2]
|
||||
dummy_ids = torch.zeros((1, cache_len), dtype=torch.long, device=device)
|
||||
# Fix: LM and TTS_LM have DIFFERENT cache lengths. Must track separately.
|
||||
|
||||
# Helper to get cache length regardless of tuple or DynamicCache
|
||||
def get_cache_len(cache_item):
|
||||
if isinstance(cache_item, (list, tuple)):
|
||||
return cache_item[0][0].shape[2]
|
||||
elif hasattr(cache_item, 'key_cache'): # DynamicCache
|
||||
return cache_item.key_cache[0].shape[2]
|
||||
return 0
|
||||
|
||||
lm_cache_len = get_cache_len(lm_cache)
|
||||
tts_cache_len = get_cache_len(tts_lm_cache)
|
||||
|
||||
# Main LM inputs (matches lm_cache)
|
||||
lm_dummy_ids = torch.zeros((1, lm_cache_len), dtype=torch.long, device=device)
|
||||
lm_dummy_mask = torch.ones((1, lm_cache_len), dtype=torch.long, device=device)
|
||||
|
||||
# TTS LM inputs (matches tts_lm_cache)
|
||||
tts_dummy_ids = torch.zeros((1, tts_cache_len), dtype=torch.long, device=device)
|
||||
tts_dummy_mask = torch.ones((1, tts_cache_len), dtype=torch.long, device=device)
|
||||
|
||||
# Resolution for sampling
|
||||
f_do_sample = getattr(model.generation_config, "do_sample", False) if do_sample == "auto" else (do_sample == "false")
|
||||
f_do_sample = getattr(model.generation_config, "do_sample", False) if do_sample == "auto" else (do_sample == "true")
|
||||
|
||||
# ComfyUI Progress Bar
|
||||
from comfy.utils import ProgressBar
|
||||
total_steps = len(text) * 20 # Rough estimate
|
||||
pbar = ProgressBar(total_steps)
|
||||
|
||||
def progress_callback(step_increment):
|
||||
pbar.update(step_increment)
|
||||
|
||||
output = model.generate(
|
||||
all_prefilled_outputs=all_prefilled,
|
||||
tts_text_ids=target_tokens,
|
||||
tts_lm_input_ids=dummy_ids,
|
||||
input_ids=dummy_ids,
|
||||
tts_lm_input_ids=tts_dummy_ids,
|
||||
tts_lm_attention_mask=tts_dummy_mask,
|
||||
input_ids=lm_dummy_ids,
|
||||
attention_mask=lm_dummy_mask,
|
||||
max_new_tokens=4000,
|
||||
cfg_scale=cfg_scale,
|
||||
do_sample=f_do_sample,
|
||||
@@ -147,7 +272,14 @@ class AIIA_VibeVoice_Realtime_TTS:
|
||||
top_p=top_p,
|
||||
expected_steps=len(text)*20,
|
||||
max_length_times=10.0,
|
||||
show_progress_bar=True
|
||||
show_progress_bar=True,
|
||||
tokenizer=tokenizer,
|
||||
progress_callback=progress_callback,
|
||||
|
||||
# Pass Acoustic Conditioning - REMOVED (Official presets don't use this during gen)
|
||||
# speech_tensors=speech_tensors,
|
||||
# speech_masks=speech_masks,
|
||||
# speech_input_mask=speech_input_mask
|
||||
)
|
||||
|
||||
# Format Audio Output
|
||||
|
||||
@@ -265,10 +265,44 @@ class VibeVoiceStreamingForConditionalGenerationInference(VibeVoiceStreamingPreT
|
||||
# Will be replaced with lm_last_hidden_state
|
||||
inputs_embeds = self.model.get_input_embeddings()(input_ids)
|
||||
|
||||
# Replace the last part of inputs_embeds with lm_last_hidden_state
|
||||
start_idx = inputs_embeds.shape[1] - lm_last_hidden_state.shape[1]
|
||||
inputs_embeds[:, start_idx:, :] = lm_last_hidden_state
|
||||
|
||||
# Replace the parts of inputs_embeds marked as text (or aligned speech) with lm_last_hidden_state
|
||||
# Logic:
|
||||
# 1. Total Count Match: If len(hidden) == len(mask), overwrite ALL (covers Text-Only or Single Speech Step).
|
||||
# 2. Text Count Match: If len(hidden) == sum(mask), scatter into Text positions (covers Sandwich Prompt).
|
||||
# 3. Fallback: Tail Overwrite.
|
||||
if lm_last_hidden_state is not None and tts_text_masks is not None:
|
||||
mask = tts_text_masks.bool() # (B, S) - True for Text
|
||||
|
||||
num_lm_tokens = lm_last_hidden_state.numel() // lm_last_hidden_state.shape[-1]
|
||||
num_total_tokens = mask.numel()
|
||||
num_text_tokens = mask.sum()
|
||||
|
||||
if num_total_tokens == num_lm_tokens:
|
||||
# Case 1: Perfect Alignment (e.g. Single Speech Step or Full Text)
|
||||
inputs_embeds = lm_last_hidden_state.view(inputs_embeds.shape)
|
||||
if self.config.use_return_dict and inputs_embeds.numel() < 1000: # Light debug
|
||||
print(f"[AIIA Debug] Full Injection! Shape {inputs_embeds.shape} Mean {inputs_embeds.mean().item():.3f}")
|
||||
elif num_text_tokens == num_lm_tokens:
|
||||
# Case 2: Sparse Text Alignment (Sandwich Prompt)
|
||||
# Flatten for scatter (B*S, H) handling is complex, iterating per batch is safer if B>1
|
||||
for b in range(inputs_embeds.shape[0]):
|
||||
b_mask = mask[b] # (S)
|
||||
b_lm = lm_last_hidden_state[b] # (S_text, H)
|
||||
if b_mask.sum() == b_lm.shape[0]:
|
||||
inputs_embeds[b, b_mask, :] = b_lm
|
||||
print(f"[AIIA Debug] Text Injection! Batch {b} Count {b_mask.sum()}")
|
||||
else:
|
||||
# Fallback for batch mismatch
|
||||
start_idx = inputs_embeds.shape[1] - lm_last_hidden_state.shape[1]
|
||||
if start_idx >= 0:
|
||||
inputs_embeds[b, start_idx:, :] = b_lm
|
||||
print(f"[AIIA Debug] Batch Fallback Injection!")
|
||||
else:
|
||||
# Case 3: Fallback (Tail Overwrite)
|
||||
start_idx = inputs_embeds.shape[1] - lm_last_hidden_state.shape[1]
|
||||
if start_idx >= 0:
|
||||
inputs_embeds[:, start_idx:, :] = lm_last_hidden_state
|
||||
|
||||
# Adds type embedding via `tts_text_masks`.
|
||||
inputs_embeds = inputs_embeds + self.model.tts_input_types(tts_text_masks.long())
|
||||
|
||||
@@ -492,6 +526,8 @@ class VibeVoiceStreamingForConditionalGenerationInference(VibeVoiceStreamingPreT
|
||||
|
||||
# Initialize audio chunks storage for each sample
|
||||
audio_chunks = [[] for _ in range(batch_size)]
|
||||
|
||||
|
||||
tts_text_window_index = 0
|
||||
reach_max_step_sample = torch.zeros(batch_size, dtype=torch.bool, device=device)
|
||||
first_text_window_size = TTS_TEXT_WINDOW_SIZE if tts_text_ids.shape[1] >= TTS_TEXT_WINDOW_SIZE else tts_text_ids.shape[1]
|
||||
|
||||
Reference in New Issue
Block a user