From 6648f3c655b00f4500a2efaab736d9840a7197df Mon Sep 17 00:00:00 2001 From: Hawk Lee Date: Tue, 30 Dec 2025 23:25:39 +0800 Subject: [PATCH] Optimize 0.5B integration: Cleanup debug prints, update docs, mark Preset Maker as experimental --- README.md | 17 +- aiia_vibevoice_preset_maker.py | 125 +++++++++---- aiia_vibevoice_realtime_tts.py | 170 ++++++++++++++++-- .../modeling_vibevoice_streaming_inference.py | 44 ++++- 4 files changed, 292 insertions(+), 64 deletions(-) diff --git a/README.md b/README.md index a6ca26a..ad55965 100644 --- a/README.md +++ b/README.md @@ -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 克隆。 + - **仅供研究**: 此节点保留给开发者进行研究调试,普通用户**不推荐**使用。 **手动下载命令**: diff --git a/aiia_vibevoice_preset_maker.py b/aiia_vibevoice_preset_maker.py index 9314fd9..638cb5b 100644 --- a/aiia_vibevoice_preset_maker.py +++ b/aiia_vibevoice_preset_maker.py @@ -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) } } diff --git a/aiia_vibevoice_realtime_tts.py b/aiia_vibevoice_realtime_tts.py index 53ac4b5..c9d6c4d 100644 --- a/aiia_vibevoice_realtime_tts.py +++ b/aiia_vibevoice_realtime_tts.py @@ -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 diff --git a/vibevoice_core/modular/modeling_vibevoice_streaming_inference.py b/vibevoice_core/modular/modeling_vibevoice_streaming_inference.py index 6d45a63..64e718c 100644 --- a/vibevoice_core/modular/modeling_vibevoice_streaming_inference.py +++ b/vibevoice_core/modular/modeling_vibevoice_streaming_inference.py @@ -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]