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

1""" 

2Wrapper for VibeVoice model. 

3""" 

4import os 

5import base64 

6import logging 

7import asyncio 

8import tempfile 

9 

10import torch 

11from torch import inference_mode 

12 

13from typing import override 

14from typing import Optional 

15from typing import Dict 

16from typing import Union 

17from typing import Any 

18 

19# Copy from transformers PR: https://github.com/huggingface/transformers/pull/40546/files 

20from vibevoice_processor import VibeVoiceProcessor 

21from modeling_vibevoice_inference import VibeVoiceForConditionalGenerationInference 

22 

23from wrapper_model import ModelGeneration 

24 

25 

26class VoiceMapper: 

27 """Maps speaker names to voice file paths""" 

28 

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) 

39 

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 

48 

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))] 

52 

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 

58 

59 self.voice_presets = dict(sorted(self.voice_presets.items())) 

60 

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 } 

66 

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())}") 

69 

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 

84 

85 

86class VibeVoiceGeneration(ModelGeneration): 

87 """Handle audio generation using the VibeVoice model. 

88 https://github.com/microsoft/VibeVoice""" 

89 

90 HF_MODEL_NAME = "microsoft/VibeVoice-1.5B" 

91 

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 

103 

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 

109 

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") 

117 

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") 

125 

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") 

137 

138 def init_parallelism(self) -> None: 

139 self.load_timer.start("torch_dist") 

140 

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)) 

144 

145 self.device_id = self.local_rank 

146 

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 

153 

154 self.device = torch.device(f"cuda:{self.device_id}") 

155 

156 torch.cuda.set_device(self.local_rank) 

157 

158 if self.world_size > 1: 

159 logging.warning("VibeVoice does not support multi-GPU setups (yet).") 

160 

161 def init_model_parallelism(self) -> None: 

162 if self.world_size > 1: 

163 logging.warning("VibeVoice does not support multi-GPU setups (yet).") 

164 

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 

171 

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") 

178 

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") 

187 

188 @inference_mode() 

189 async def warmup(self) -> None: 

190 logging.info("Warmup for VibeVoice generation.") 

191 await self.generate(text="Warmup") 

192 

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. 

195 

196 The caller is responsible for deleting the file when done. 

197 

198 Args: 

199 voice_sample: Base64-encoded WAV audio. 

200 

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 

210 

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. 

213 

214 Silently ignores errors so that cleanup never raises inside a finally block. 

215 

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}") 

224 

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. 

237 

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) 

248 

249 self._assert_model_init() 

250 

251 self.running = True 

252 

253 _tmp_voice_path: Optional[str] = None 

254 try: 

255 if not self.voice_mapper: 

256 raise ValueError("VoiceMapper not initialized") 

257 

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] 

284 

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") 

301 

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") 

325 

326 return output_path 

327 finally: 

328 self.running = False 

329 gen_timer.end("total") 

330 self._cleanup_tmp_voice_file(_tmp_voice_path) 

331 

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") 

339 

340 job_id = data_json.get("job_id", None) 

341 

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 }