RustyMark/dots.tts
0
1from __future__ import annotations2 3from pathlib import Path4 5import torch6from loguru import logger7 8from dots_tts.data.pipelines.tts_pipeline import TTS_INTERLEAVE_PREFIX9from dots_tts.runtime import DotsTtsRuntime10from dots_tts.utils.util import get_dtype11 12 13class DoubleStreamingSession:14 """Incremental interleave session for text-token to audio-chunk generation."""15 16 def __init__(17 self,18 runtime: DotsTtsRuntime,19 *,20 prompt_audio_path: str | None = None,21 prompt_text: str | None = None,22 ode_method: str = "euler",23 num_steps: int = 10,24 guidance_scale: float = 1.2,25 speaker_scale: float = 1.5,26 eos_threshold: float = 0.8,27 initial_silence_audio_tokens: int = 1,28 ) -> None:29 normalized_prompt_text = runtime._process_prompt_text(prompt_text)30 if normalized_prompt_text:31 raise ValueError("Double streaming does not support prompt_text.")32 33 self.runtime = runtime34 self.model = runtime.model35 self.device = runtime.device36 self.ode_method = ode_method37 self.num_steps = int(num_steps)38 self.guidance_scale = float(guidance_scale)39 self.speaker_scale = float(speaker_scale)40 self.eos_threshold = float(eos_threshold)41 self.max_generate_length = runtime.max_generate_length42 self._initial_silence_audio_tokens = max(43 0,44 min(10, int(initial_silence_audio_tokens or 0)),45 )46 47 self._dtype = get_dtype(runtime.precision)48 self._use_amp = self.device.type == "cuda" and self._dtype in {49 torch.float16,50 torch.bfloat16,51 }52 self._prefix_token_ids = tuple(53 self.model.tokenizer.encode(54 TTS_INTERLEAVE_PREFIX,55 add_special_tokens=False,56 )57 )58 self._state = self.model._allocate_generate_state(59 max_audio_patch_count=self.max_generate_length,60 device=self.device,61 dtype=self._dtype,62 )63 self._vocoder_state = self.model.vocoder.init_stream_state(64 batch_size=1,65 chunk_size=self.model.core.latent_patch_size,66 )67 self._g_cond = None68 self._started = False69 self._text_finished = False70 self._closed = False71 self._decoded_patch_count = 072 73 if prompt_audio_path is not None:74 cache = getattr(self.runtime, "_double_streaming_prompt_g_cond_cache", None)75 if cache is None:76 cache = {}77 setattr(self.runtime, "_double_streaming_prompt_g_cond_cache", cache)78 prompt_cache_key = (79 str(Path(prompt_audio_path).expanduser().resolve()),80 str(self.device),81 str(self._dtype),82 self.speaker_scale,83 )84 cached_g_cond = cache.get(prompt_cache_key)85 if cached_g_cond is None:86 prompt_audio = self.runtime._load_prompt_audio(prompt_audio_path)87 with torch.no_grad():88 with torch.autocast(89 device_type=self.device.type,90 dtype=self._dtype,91 enabled=self._use_amp,92 ):93 prompt_conditioning = self.model._prepare_prompt_conditioning(94 prompt_audio,95 use_prompt_prefill=False,96 speaker_scale=self.speaker_scale,97 )98 cached_g_cond = prompt_conditioning.g_cond.detach()99 cache[prompt_cache_key] = cached_g_cond100 logger.info(101 "Double streaming prompt conditioning cached: path={} device={} "102 "dtype={} speaker_scale={}",103 prompt_cache_key[0],104 self.device,105 self._dtype,106 self.speaker_scale,107 )108 else:109 logger.info(110 "Double streaming prompt conditioning cache hit: path={} device={} "111 "dtype={} speaker_scale={}",112 prompt_cache_key[0],113 self.device,114 self._dtype,115 self.speaker_scale,116 )117 self._g_cond = cached_g_cond118 119 logger.info(120 "Double streaming session started: prefix_token_count={} precision={} "121 "ode_method={} num_steps={} guidance_scale={} speaker_scale={} max_audio_patch_count={} "122 "initial_silence_audio_tokens={} has_ref_audio_only={}",123 len(self._prefix_token_ids),124 runtime.precision,125 self.ode_method,126 self.num_steps,127 self.guidance_scale,128 self.speaker_scale,129 self.max_generate_length,130 self._initial_silence_audio_tokens,131 self._g_cond is not None,132 )133 134 @property135 def is_finished(self) -> bool:136 return self._closed137 138 def push_text_token(self, text_token: int) -> torch.Tensor | None:139 self._ensure_active()140 if self._text_finished:141 raise RuntimeError("Cannot push text tokens after finish_text().")142 if self._state.end_flag:143 raise RuntimeError(144 "Double streaming generation has already reached EOS. "145 "Call finish_text() to flush the remaining audio tail."146 )147 148 token_id = int(text_token)149 if not self._started:150 chunk_token_ids = [*self._prefix_token_ids, token_id]151 self._started = True152 else:153 chunk_token_ids = [token_id]154 155 self._consume_text_chunk(chunk_token_ids)156 return self._decode_audio_chunk()157 158 def finish_text(self):159 self._ensure_active()160 161 if not self._state.end_flag:162 if not self._text_finished:163 text_end_chunk = [self.model.core.text_cond_end_id]164 if not self._started:165 text_end_chunk = [*self._prefix_token_ids, *text_end_chunk]166 self._started = True167 self._consume_text_chunk(text_end_chunk)168 self._text_finished = True169 170 while not self._state.end_flag:171 audio_chunk = self._decode_audio_chunk(continue_audio_span=True)172 if audio_chunk is not None:173 yield audio_chunk174 else:175 self._text_finished = True176 177 final_chunk = self.model.vocoder.stream_flush(self._vocoder_state)178 self._closed = True179 logger.info(180 "Double streaming session finished: decoded_patch_count={}",181 self._decoded_patch_count,182 )183 if final_chunk.size(-1) > 0:184 yield final_chunk185 186 def _ensure_active(self) -> None:187 if self._closed:188 raise RuntimeError("Double streaming session is already closed.")189 190 def _consume_text_chunk(self, token_ids: list[int]) -> None:191 schedule = torch.tensor(192 [token_ids],193 dtype=torch.long,194 device=self.device,195 )196 with torch.no_grad():197 with torch.autocast(198 device_type=self.device.type,199 dtype=self._dtype,200 enabled=self._use_amp,201 ):202 self.model._consume_text_schedule(203 schedule,204 position=0,205 next_audio_position=schedule.size(1),206 state=self._state,207 )208 209 def _get_initial_silence_audio_patch(210 self,211 patch_index: int,212 audio_patch: torch.Tensor,213 ) -> torch.Tensor:214 cache = getattr(self.runtime, "_double_streaming_silence_audio_patch_cache", None)215 if cache is None:216 cache = {}217 setattr(self.runtime, "_double_streaming_silence_audio_patch_cache", cache)218 219 cache_count = 10220 patch_size = int(self.model.core.latent_patch_size)221 key = (222 str(self.device),223 str(self._dtype),224 patch_size,225 int(audio_patch.size(-1)),226 cache_count,227 )228 cached_patches = cache.get(key)229 if cached_patches is None:230 hop_size = int(getattr(self.model.vocoder, "hop_size", 1))231 zero_samples = cache_count * patch_size * hop_size232 zero_audio = torch.zeros(233 (1, 1, zero_samples),234 device=self.device,235 dtype=torch.float32,236 )237 silence_latents = self.model.vocoder.extract_latents(zero_audio)238 silence_latents, _ = torch.split(239 silence_latents,240 int(audio_patch.size(-1)),241 dim=1,242 )243 silence_latents = silence_latents.transpose(1, 2)244 target_frames = cache_count * patch_size245 if silence_latents.size(1) < target_frames:246 silence_latents = torch.cat(247 [248 silence_latents,249 silence_latents.new_zeros(250 (251 silence_latents.size(0),252 target_frames - silence_latents.size(1),253 silence_latents.size(2),254 )255 ),256 ],257 dim=1,258 )259 silence_latents = silence_latents[:, :target_frames, :]260 cached_patches = self.model.core.io_helper.normalize(silence_latents)261 cached_patches = cached_patches.to(device=self.device, dtype=audio_patch.dtype)262 cached_patches = cached_patches.reshape(263 1,264 cache_count,265 patch_size,266 int(audio_patch.size(-1)),267 ).detach()268 cache[key] = cached_patches269 logger.info(270 "Double streaming initial silence cache built: patches={} patch_size={} "271 "hop_size={} device={} dtype={}",272 cache_count,273 patch_size,274 hop_size,275 self.device,276 audio_patch.dtype,277 )278 return cached_patches[:, int(patch_index)].clone()279 280 def _consume_audio_patch(self, audio_patch: torch.Tensor) -> None:281 self.model._consume_audio_patch(self._state, audio_patch=audio_patch)282 283 def _decode_audio_chunk(self, *, continue_audio_span: bool = False) -> torch.Tensor | None:284 if self._decoded_patch_count >= self.max_generate_length:285 raise RuntimeError(286 "Double streaming exceeded max_generate_length before reaching EOS."287 )288 289 with torch.no_grad():290 with torch.autocast(291 device_type=self.device.type,292 dtype=self._dtype,293 enabled=self._use_amp,294 ):295 stop_after_current_audio = self.model._should_stop_after_current_audio(296 self._state,297 eos_threshold=self.eos_threshold,298 )299 audio_patch = self.model._decode_next_audio(300 self._state,301 device=self.device,302 g_cond=self._g_cond,303 ode_method=self.ode_method,304 num_steps=self.num_steps,305 guidance_scale=self.guidance_scale,306 )307 if self._decoded_patch_count < self._initial_silence_audio_tokens:308 audio_patch = self._get_initial_silence_audio_patch(309 self._decoded_patch_count,310 audio_patch,311 )312 self._consume_audio_patch(audio_patch)313 if continue_audio_span:314 self.model._append_hidden_chunk(self._state, self._state.llm_hiddens)315 self._decoded_patch_count += 1316 latent_patch = self.model.core.io_helper.denormalize(audio_patch)317 audio_chunk = self.model.vocoder.stream_step(318 latent_patch.transpose(1, 2),319 self._vocoder_state,320 )321 if stop_after_current_audio:322 self._state.end_flag = True323 324 if audio_chunk.size(-1) == 0:325 return None326 return audio_chunk327 328 329class DotsTtsRuntimeDoubleStreaming(DotsTtsRuntime):330 def start_double_streaming(331 self,332 *,333 prompt_audio_path: str | None = None,334 prompt_text: str | None = None,335 ode_method: str = "euler",336 num_steps: int = 10,337 guidance_scale: float = 1.2,338 speaker_scale: float = 1.5,339 eos_threshold: float = 0.8,340 initial_silence_audio_tokens: int = 1,341 ) -> DoubleStreamingSession:342 return DoubleStreamingSession(343 self,344 prompt_audio_path=prompt_audio_path,345 prompt_text=prompt_text,346 ode_method=ode_method,347 num_steps=num_steps,348 guidance_scale=guidance_scale,349 speaker_scale=speaker_scale,350 eos_threshold=eos_threshold,351 initial_silence_audio_tokens=initial_silence_audio_tokens,352 )353 354 355__all__ = ["DotsTtsRuntimeDoubleStreaming", "DoubleStreamingSession"]356 