Mike0021/zonos2
3
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 