Optimize 0.5B integration: Cleanup debug prints, update docs, mark Preset Maker as experimental

This commit is contained in:
Hawk Lee
2025-12-30 23:25:39 +08:00
parent 6614ac07cf
commit 6648f3c655
4 changed files with 292 additions and 64 deletions
+13 -4
View File
@@ -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 克隆。
- **仅供研究**: 此节点保留给开发者进行研究调试,普通用户**不推荐**使用。
**手动下载命令**:
+88 -37
View File
@@ -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
View File
@@ -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]