CoolFace
Apppublic

RustyMark/dots.tts

sourceHugging Faceapache-2.0updated 3mo agoView on Hugging Face
0likes
runtime_double_streaming.py356 linesDownload Raw Back to dots_tts
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