CoolFace
Apppublic

Mike0021/zonos2

sourceHugging Faceupdated 4mo agoView on Hugging Face
3likes
core.py219 linesDownload Raw Back to zonos2
1from __future__ import annotations2 3from contextlib import contextmanager4from dataclasses import dataclass, field5from typing import TYPE_CHECKING, List, Literal6 7import torch8 9if TYPE_CHECKING:10    from zonos2.attention import BaseAttnBackend, BaseAttnMetadata11    from zonos2.kvcache import BaseCacheHandle12 13 14@dataclass15class Context:16    page_size: int17    attn_backend: BaseAttnBackend18    _batch: TTSBatch | None = field(default=None, init=False)19 20    @property21    def batch(self) -> TTSBatch:22        assert self._batch is not None, "No active batch in context"23        return self._batch24 25    @contextmanager26    def forward_batch(self, batch: TTSBatch):27        assert self._batch is None, "Nested forward_batch is not allowed"28        try:29            self._batch = batch30            yield31        finally:32            self._batch = None33 34 35_GLOBAL_CTX: Context | None = None36 37 38def set_global_ctx(ctx: Context):39    global _GLOBAL_CTX40    assert _GLOBAL_CTX is None, "Global context is already set"41    _GLOBAL_CTX = ctx42 43 44def get_global_ctx() -> Context:45    assert _GLOBAL_CTX is not None, "Global context is not set"46    return _GLOBAL_CTX47 48 49# =============================================================================50# TTS-specific data structures51# =============================================================================52 53 54@dataclass55class TTSSamplingParams:56    """Sampling parameters for TTS generation."""57 58    temperature: float = 1.1559    topk: int = 10660    top_p: float = 0.061    min_p: float = 0.1862    max_tokens: int = 102463    ignore_eos: bool = False64    repetition_window: int = 5065    repetition_penalty: float = 1.266    repetition_codebooks: int = 867    seed: int | None = None68 69 70@dataclass(eq=False)71class TTSReq:72    """Request class for TTS generation with 2D token format.73 74    Tokens are in unpacked format: [cb0, cb1, ..., cb8, text_token] per frame.75    """76 77    input_ids: torch.Tensor  # 2D CPU tensor (seq_len, frame_width)78    table_idx: int79    cached_len: int80    output_len: int81    uid: int82    sampling_params: TTSSamplingParams83    cache_handle: BaseCacheHandle84    n_codebooks: int = 985    eoa_id: int = 102486    eos_frame: int = -1  # Aligned frame where EOS first appeared (-1 = not seen)87    eos_countdown: int = -1  # Steps remaining after EOS (-1 = not in countdown)88    total_generated: int = 0  # Total frames generated (for logging)89    rng: torch.Generator | None = None  # Per-request RNG for deterministic sampling90    speaker_embedding: torch.Tensor | None = None  # 1D CPU float32 tensor91    speaker_token_position: int = -1  # Injection position within the prompt sequence92 93    def __post_init__(self) -> None:94        assert self.input_ids.is_cpu95        assert self.input_ids.dim() == 2, "TTS input_ids must be 2D (seq_len, frame_width)"96        self.device_len = len(self.input_ids)97        self.max_device_len = len(self.input_ids) + self.output_len98        assert 0 <= self.cached_len < self.device_len <= self.max_device_len99 100        if self.speaker_embedding is not None:101            emb = self.speaker_embedding102            if emb.dim() == 2 and emb.shape[0] == 1:103                emb = emb.squeeze(0)104            if emb.dim() != 1:105                raise ValueError(106                    f"speaker_embedding must be 1D or (1, D), got shape {tuple(emb.shape)}"107                )108            self.speaker_embedding = emb.to(dtype=torch.float32, device="cpu")109 110        if self.speaker_token_position < 0:111            # Training convention: reserved speaker slot is at prompt position 0.112            self.speaker_token_position = 0113        if self.speaker_token_position >= self.device_len:114            self.speaker_token_position = 0115 116    @property117    def frame_width(self) -> int:118        """Number of elements per frame (n_codebooks + extras)."""119        return self.input_ids.shape[-1]120 121    @property122    def remain_len(self) -> int:123        return self.max_device_len - self.device_len124 125    @property126    def extend_len(self) -> int:127        return self.device_len - self.cached_len128 129    @property130    def num_completion_tokens(self) -> int:131        """Number of generated tokens (frames)."""132        return self.device_len - self.cached_len133 134    def complete_one(self) -> None:135        self.cached_len = self.device_len136        self.device_len += 1137        self.total_generated += 1138 139    def append_host(self, next_token: torch.Tensor) -> None:140        """Append a single frame (unpacked token) to input_ids."""141        assert next_token.dim() == 1, "next_token must be 1D (frame_width,)"142        self.input_ids = torch.cat([self.input_ids, next_token.unsqueeze(0)], dim=0)143 144    def can_decode(self) -> bool:145        return self.remain_len > 0 and self.eos_countdown != 0146 147    def check_eos(self, audio_codes: List[int]) -> bool:148        """Check for EOS and update countdown state.149 150        Args:151            audio_codes: List of audio codebook values for one frame152 153        Returns:154            True if sequence is finished (countdown reached 0)155        """156        if self.sampling_params.ignore_eos:157            return False158 159        # Match Zonos2 reference inference: any sampled EOA codebook starts the160        # delayed stop countdown. The aligned frame is shifted back by the161        # highest EOA codebook index and clamped at zero.162        # Use total_generated because this request only sees one decode frame at a time.163        if self.eos_frame < 0:164            step = self.total_generated - 1165            eos_cols = [c == self.eoa_id for c in audio_codes[: self.n_codebooks]]166            if any(eos_cols):167                # First EOS: compute aligned frame168                max_eos_cb = max(i for i, is_eos in enumerate(eos_cols) if is_eos)169                self.eos_frame = max(0, step - max_eos_cb)170                self.eos_countdown = self.n_codebooks + 1171 172        # Decrement countdown173        if self.eos_countdown > 0:174            self.eos_countdown -= 1175            if self.eos_countdown == 0:176                return True177 178        return False179 180    def __repr__(self) -> str:181        return (182            f"{type(self).__name__}(table_idx={self.table_idx}, "183            f"cached_len={self.cached_len}, device_len={self.device_len}, "184            f"max_device_len={self.max_device_len}, eos_frame={self.eos_frame})"185        )186 187 188@dataclass189class TTSBatch:190    """Batch of TTS requests with 2D token format."""191 192    reqs: List[TTSReq]193    phase: Literal["prefill", "decode"]194    # these fields should be set by scheduler195    input_ids: torch.Tensor = field(init=False)  # (total_tokens, frame_width)196    out_loc: torch.Tensor = field(init=False)197    padded_reqs: List[TTSReq] = field(init=False)198    # this field should be set by attention backend199    attn_metadata: BaseAttnMetadata = field(init=False)200    # Optional per-batch speaker conditioning data (set by TTS scheduler).201    speaker_emb_values: torch.Tensor | None = field(default=None, init=False)202    speaker_token_positions: torch.Tensor | None = field(default=None, init=False)203 204    @property205    def is_prefill(self) -> bool:206        return self.phase == "prefill"207 208    @property209    def is_decode(self) -> bool:210        return self.phase == "decode"211 212    @property213    def size(self) -> int:214        return len(self.reqs)215 216    @property217    def padded_size(self) -> int:218        return len(self.padded_reqs)219