Coverage for wrapper/vibevoice/wrapper_vibevoice.py: 73%
196 statements
« prev ^ index » next coverage.py v7.15.4, created at 2026-08-09 04:47 +0000
« prev ^ index » next coverage.py v7.15.4, created at 2026-08-09 04:47 +0000
1"""
2Wrapper for VibeVoice model.
3"""
4import os
5import base64
6import logging
7import asyncio
8import tempfile
10import torch
11from torch import inference_mode
13from typing import override
14from typing import Optional
15from typing import Dict
16from typing import Union
17from typing import Any
19# Copy from transformers PR: https://github.com/huggingface/transformers/pull/40546/files
20from vibevoice_processor import VibeVoiceProcessor
21from modeling_vibevoice_inference import VibeVoiceForConditionalGenerationInference
23from wrapper_model import ModelGeneration
26class VoiceMapper:
27 """Maps speaker names to voice file paths"""
29 def __init__(self) -> None:
30 self.setup_voice_presets()
31 new_dict = {}
32 for name, path in self.voice_presets.items():
33 if '_' in name:
34 name = name.split('_')[0]
35 if '-' in name:
36 name = name.split('-')[-1]
37 new_dict[name] = path
38 self.voice_presets.update(new_dict)
40 def setup_voice_presets(self) -> None:
41 """Setup voice presets by scanning the voices directory."""
42 voices_dir = os.path.join(os.path.dirname(__file__), "voices")
43 if not os.path.exists(voices_dir):
44 logging.warning(f"Voices directory not found at {voices_dir}")
45 self.voice_presets: dict[str, str] = {}
46 self.available_voices: dict[str, str] = {}
47 return
49 wav_files = [
50 f for f in os.listdir(voices_dir)
51 if f.lower().endswith('.wav') and os.path.isfile(os.path.join(voices_dir, f))]
53 self.voice_presets = {}
54 for wav_file in wav_files:
55 name = os.path.splitext(wav_file)[0]
56 full_path = os.path.join(voices_dir, wav_file)
57 self.voice_presets[name] = full_path
59 self.voice_presets = dict(sorted(self.voice_presets.items()))
61 # Filter out voices that don't exist (this is now redundant but kept for safety)
62 self.available_voices = {
63 name: path for name, path in self.voice_presets.items()
64 if os.path.exists(path)
65 }
67 logging.info(f"Found {len(self.available_voices)} voice files in {voices_dir}")
68 logging.info(f"Available voices: {', '.join(self.available_voices.keys())}")
70 def get_voice_path(self, speaker_name: str) -> str:
71 """Get voice file path for a given speaker name"""
72 if speaker_name in self.voice_presets:
73 return self.voice_presets[speaker_name]
74 speaker_lower = speaker_name.lower()
75 for preset_name, path in self.voice_presets.items():
76 if preset_name.lower() in speaker_lower or speaker_lower in preset_name.lower():
77 return path
78 voices_list = list(self.voice_presets.values())
79 if not voices_list:
80 raise ValueError("No voice presets available.")
81 default_voice = voices_list[0]
82 logging.warning(f"No voice preset found for '{speaker_name}', using default voice: {default_voice}")
83 return default_voice
86class VibeVoiceGeneration(ModelGeneration):
87 """Handle audio generation using the VibeVoice model.
88 https://github.com/microsoft/VibeVoice"""
90 HF_MODEL_NAME = "microsoft/VibeVoice-1.5B"
92 def __init__(
93 self,
94 model_name: str = "vibevoice",
95 param_dtype: torch.dtype = torch.bfloat16,
96 ) -> None:
97 super().__init__(model_name)
98 self.param_dtype = param_dtype
99 # https://github.com/microsoft/VibeVoice/tree/main/demo/voices
100 self.voice_mapper: Optional[VoiceMapper] = None
101 self.processor: Optional[VibeVoiceProcessor] = None
102 self.vibevoice: Optional[VibeVoiceForConditionalGenerationInference] = None
104 def __del__(self) -> None:
105 if self.processor is not None:
106 self.processor = None
107 if self.vibevoice is not None:
108 self.vibevoice = None
110 def load_model(self) -> None:
111 self.load_timer.start("voice_mapper")
112 self.voice_mapper = VoiceMapper()
113 logging.info("Available voices:")
114 for voice_name, voice_path in self.voice_mapper.available_voices.items():
115 logging.info(f" {voice_name}: {voice_path}")
116 self.load_timer.end("voice_mapper")
118 self.load_timer.start("processor")
119 self.processor = VibeVoiceProcessor.from_pretrained(self.HF_MODEL_NAME)
120 # TODO can we move this to the GPU device?
121 # self.processor.to(self.device)
122 # self.processor.tokenizer = self.processor.tokenizer.to(self.device)
123 # self.audio_processor = self.processor.audio_processor.to(self.device)
124 self.load_timer.end("processor")
126 self.load_timer.start("vibevoice")
127 self.vibevoice = VibeVoiceForConditionalGenerationInference.from_pretrained(
128 self.HF_MODEL_NAME,
129 torch_dtype=self.param_dtype,
130 attn_implementation="flash_attention_2")
131 if not self.vibevoice:
132 raise ValueError("Failed to load VibeVoice model")
133 self.vibevoice.eval()
134 self.vibevoice.set_ddpm_inference_steps(num_steps=10)
135 self.vibevoice.to(self.device)
136 self.load_timer.end("vibevoice")
138 def init_parallelism(self) -> None:
139 self.load_timer.start("torch_dist")
141 self.rank = int(os.getenv("RANK", 0))
142 self.local_rank = int(os.getenv("LOCAL_RANK", 0))
143 self.world_size = int(os.getenv("WORLD_SIZE", 1))
145 self.device_id = self.local_rank
147 if not torch.cuda.is_available():
148 self.device_id = 0
149 self.device = torch.device("cpu")
150 logging.warning("CUDA is not available. Running on CPU.")
151 self.load_timer.end("torch_dist")
152 return # CPU mode, no parallelism needed
154 self.device = torch.device(f"cuda:{self.device_id}")
156 torch.cuda.set_device(self.local_rank)
158 if self.world_size > 1:
159 logging.warning("VibeVoice does not support multi-GPU setups (yet).")
161 def init_model_parallelism(self) -> None:
162 if self.world_size > 1:
163 logging.warning("VibeVoice does not support multi-GPU setups (yet).")
165 def model_compile(self) -> None:
166 """Compile the model components for optimized performance.
167 self.processor cannot be compiled due to dynamic input shapes.
168 """
169 if not self.torch_compile:
170 return
172 self.load_timer.start("vibevoice_compile")
173 self.vibevoice = torch.compile(
174 self.vibevoice,
175 mode="max-autotune-no-cudagraphs",
176 )
177 self.load_timer.end("vibevoice_compile")
179 def _assert_model_init(self) -> None:
180 super()._assert_model_init()
181 if not self.voice_mapper:
182 raise ValueError("VoiceMapper not initialized")
183 if not self.processor:
184 raise ValueError("Processor not initialized")
185 if not self.vibevoice:
186 raise ValueError("VibeVoice model not initialized")
188 @inference_mode()
189 async def warmup(self) -> None:
190 logging.info("Warmup for VibeVoice generation.")
191 await self.generate(text="Warmup")
193 def _decode_voice_sample_to_tmp_file(self, voice_sample: str) -> str:
194 """Decode a base64-encoded WAV voice sample to a temporary file.
196 The caller is responsible for deleting the file when done.
198 Args:
199 voice_sample: Base64-encoded WAV audio.
201 Returns:
202 Path to the temporary WAV file.
203 """
204 audio_bytes = base64.b64decode(voice_sample)
205 with tempfile.NamedTemporaryFile(suffix=".wav", delete=False) as tmp_f:
206 tmp_f.write(audio_bytes)
207 tmp_path = tmp_f.name
208 logging.info(f"Decoded voice_sample to temporary file: {tmp_path} ({len(audio_bytes)} bytes).")
209 return tmp_path
211 def _cleanup_tmp_voice_file(self, tmp_path: Optional[str]) -> None:
212 """Remove a temporary voice file created by _decode_voice_sample_to_tmp_file.
214 Silently ignores errors so that cleanup never raises inside a finally block.
216 Args:
217 tmp_path: Path returned by _decode_voice_sample_to_tmp_file, or None (no-op).
218 """
219 if tmp_path is not None:
220 try:
221 os.unlink(tmp_path)
222 except OSError as e:
223 logging.warning(f"Could not remove temporary voice file {tmp_path}: {e}")
225 @override
226 @inference_mode()
227 async def generate(
228 self,
229 text: str,
230 voice: str = "woman_000",
231 voice_sample: Optional[str] = None,
232 cfg_scale: float = 1.3,
233 job_id: Optional[str] = None,
234 output_type: str = "audio_path",
235 ) -> str:
236 """Generate speech audio from text.
238 Args:
239 text: The text to synthesise.
240 voice: Name of a built-in voice preset (used when *voice_sample* is not provided).
241 voice_sample: Base64-encoded WAV audio to clone the voice from. When supplied
242 the model uses this audio as the reference speaker instead of a preset.
243 cfg_scale: Classifier-free guidance scale.
244 job_id: Optional job identifier used for output file naming.
245 output_type: Output format selector (currently only "audio_path" is supported).
246 """
247 gen_timer = self._new_gen_timer(job_id)
249 self._assert_model_init()
251 self.running = True
253 _tmp_voice_path: Optional[str] = None
254 try:
255 if not self.voice_mapper:
256 raise ValueError("VoiceMapper not initialized")
258 if voice_sample is not None:
259 if job_id is not None:
260 # Save the input audio for debugging, mirroring how other wrappers
261 # save their inputs (e.g. /tmp/{job_id}.png for images).
262 # Use basename to strip any path-traversal characters from job_id.
263 safe_job_id = os.path.basename(job_id)
264 debug_voice_path = f"/tmp/{safe_job_id}_voice_sample.wav"
265 audio_bytes = base64.b64decode(voice_sample)
266 with open(debug_voice_path, "wb") as out_f:
267 out_f.write(audio_bytes)
268 logging.info(
269 f"Saved voice_sample to {debug_voice_path} ({len(audio_bytes)} bytes)."
270 )
271 voice_path = debug_voice_path
272 # _tmp_voice_path stays None: debug file is not cleaned up so it
273 # remains available for post-mortem inspection.
274 else:
275 # No job_id available (e.g. during warmup): fall back to a temp file
276 # that is cleaned up in the finally block.
277 _tmp_voice_path = self._decode_voice_sample_to_tmp_file(voice_sample)
278 voice_path = _tmp_voice_path
279 logging.info("Using cloned voice from provided voice_sample.")
280 else:
281 voice_path = self.voice_mapper.get_voice_path(voice)
282 logging.info(f"Using voice: {voice} -> {voice_path}")
283 voice_samples = [voice_path]
285 # https://github.com/microsoft/VibeVoice/blob/main/demo/inference_from_file.py
286 if not self.processor:
287 raise ValueError("Processor not initialized")
288 gen_timer.start("processor")
289 inputs = self.processor(
290 text=["Speaker 0:" + text + "\n"], # Wrap in list for batch processing
291 voice_samples=[voice_samples], # Wrap in list for batch processing
292 padding=True,
293 return_tensors="pt",
294 return_attention_mask=True,
295 )
296 inputs = {
297 k: v.to(self.device) if torch.is_tensor(v) else v
298 for k, v in inputs.items()
299 }
300 gen_timer.end("processor")
302 gen_timer.start("vibevoice")
303 if self.vibevoice is None:
304 raise RuntimeError("VibeVoice model not loaded")
305 outputs = await asyncio.to_thread(
306 self.vibevoice.generate,
307 **inputs,
308 max_new_tokens=None,
309 cfg_scale=cfg_scale,
310 tokenizer=self.processor.tokenizer,
311 generation_config={'do_sample': False},
312 verbose=True,
313 )
314 output_path = "/tmp/file.wav"
315 if job_id is not None:
316 output_path = f"/tmp/file_{job_id}.wav"
317 if outputs.speech_outputs and outputs.speech_outputs[0] is not None:
318 speech_output = outputs.speech_outputs[0]
319 await asyncio.to_thread(
320 self.processor.save_audio,
321 speech_output,
322 output_path=output_path,
323 )
324 gen_timer.end("vibevoice")
326 return output_path
327 finally:
328 self.running = False
329 gen_timer.end("total")
330 self._cleanup_tmp_voice_file(_tmp_voice_path)
332 async def get_rest_args(
333 self,
334 data_json: Dict[str, Union[str, int, float]]
335 ) -> Dict[str, Any]:
336 """Get REST argguments."""
337 if data_json is None:
338 raise ValueError("Missing JSON body")
340 job_id = data_json.get("job_id", None)
342 text = data_json.get("text", None)
343 if text is None:
344 raise ValueError("Missing 'text' parameter")
345 voice = data_json.get("voice", "af_heart")
346 voice_sample = data_json.get("voice_sample", None)
347 args: Dict[str, Any] = {
348 "job_id": job_id,
349 "text": text,
350 "voice": voice,
351 }
352 if voice_sample is not None:
353 args["voice_sample"] = voice_sample
354 return {
355 "task": self.model_name,
356 "args": args,
357 }