from __future__ import annotations from typing import Any import torch from torch import nn from .sampling import sample_token class Qwen3TTSResidualPredictor(nn.Module): """The complete pretrained five-layer Qwen3-TTS code predictor.""" def __init__(self, hidden_size: int, codebook_size: int, config: dict[str, Any]): super().__init__() from transformers import Qwen3Config, Qwen3Model if hidden_size != int(config["hidden_size"]) or codebook_size != int( config["vocab_size"] ): raise ValueError("Qwen3-TTS residual predictor dimensions do not match") backbone_config = Qwen3Config( vocab_size=codebook_size, hidden_size=hidden_size, intermediate_size=int(config["intermediate_size"]), num_hidden_layers=int(config["num_hidden_layers"]), num_attention_heads=int(config["num_attention_heads"]), num_key_value_heads=int(config["num_key_value_heads"]), head_dim=int(config["head_dim"]), hidden_act=config["hidden_act"], max_position_embeddings=int(config["max_position_embeddings"]), initializer_range=float(config["initializer_range"]), rms_norm_eps=float(config["rms_norm_eps"]), use_cache=True, tie_word_embeddings=False, rope_theta=float(config["rope_theta"]), attention_bias=bool(config["attention_bias"]), attention_dropout=float(config["attention_dropout"]), ) backbone_config._attn_implementation = "eager" self.backbone = Qwen3Model(backbone_config) self.backbone.embed_tokens = None self.q0_embedding = nn.Embedding(codebook_size, hidden_size) self.code_embeddings = nn.ModuleList( [nn.Embedding(codebook_size, hidden_size) for _ in range(14)] ) self.heads = nn.ModuleList( [nn.Linear(hidden_size, codebook_size, bias=False) for _ in range(15)] ) self._greedy_graph = None self._sampling_graph = None @torch.no_grad() def prepare_cuda_graph(self) -> None: """Capture single-session greedy inference after loading the checkpoint.""" self._greedy_graph = self._capture_generation_graph(do_sample=False) @torch.no_grad() def prepare_sampling_cuda_graph(self) -> None: """Capture top-k 50, top-p 1, temperature 0.9 with the default CUDA RNG.""" self._sampling_graph = self._capture_generation_graph(do_sample=True) def _capture_generation_graph(self, *, do_sample: bool) -> tuple: weight = self.q0_embedding.weight if self.training or weight.device.type != "cuda": raise ValueError("Residual CUDA graph requires a CUDA model in eval mode") with torch.cuda.device(weight.device): hidden = weight.new_zeros((1, weight.shape[1])) q0 = torch.zeros(1, device=weight.device, dtype=torch.long) masks = [weight.new_zeros((1, 1, 2, 2))] masks[0][0, 0, 0, 1] = torch.finfo(weight.dtype).min masks.extend(weight.new_zeros((1, 1, 1, n)) for n in range(3, 17)) stream = torch.cuda.Stream(device=weight.device) stream.wait_stream(torch.cuda.current_stream()) with torch.cuda.stream(stream): for _ in range(3): self._generate(hidden, q0, do_sample=do_sample, attention_masks=masks) torch.cuda.current_stream().wait_stream(stream) graph = torch.cuda.CUDAGraph() with torch.cuda.graph(graph): output = self._generate(hidden, q0, do_sample=do_sample, attention_masks=masks) # Keep all captured inputs alive for subsequent replays. return graph, hidden, q0, output, masks def forward(self, hidden: torch.Tensor, frame_prefix: torch.Tensor) -> torch.Tensor: if frame_prefix.ndim != 2 or frame_prefix.shape[1] != 16: raise ValueError("frame_prefix must be [N, 16]") inputs = [hidden.unsqueeze(1), self.q0_embedding(frame_prefix[:, :1])] inputs.extend( embedding(frame_prefix[:, group : group + 1]) for group, embedding in enumerate(self.code_embeddings, start=1) ) outputs = self.backbone( inputs_embeds=torch.cat(inputs, dim=1), use_cache=False ).last_hidden_state return torch.stack( [head(outputs[:, group + 1]) for group, head in enumerate(self.heads)], dim=1, ) @torch.no_grad() def generate( self, hidden: torch.Tensor, q0: torch.Tensor, *, do_sample: bool = False, top_k: int = 50, top_p: float = 1.0, temperature: float = 0.9, generator: torch.Generator | None = None, ) -> torch.Tensor: captured = self._greedy_graph if not do_sample else ( self._sampling_graph if generator is None and top_k == 50 and top_p == 1.0 and temperature == 0.9 else None ) if captured is not None and hidden.shape[0] == 1: graph, static_hidden, static_q0, output, _ = captured static_hidden.copy_(hidden) static_q0.copy_(q0) graph.replay() return output.clone() return self._generate( hidden, q0, do_sample=do_sample, top_k=top_k, top_p=top_p, temperature=temperature, generator=generator, ) def _generate( self, hidden: torch.Tensor, q0: torch.Tensor, *, do_sample: bool = False, top_k: int = 50, top_p: float = 1.0, temperature: float = 0.9, generator: torch.Generator | None = None, attention_masks: list[torch.Tensor] | None = None, ) -> torch.Tensor: inputs = torch.cat( [hidden.unsqueeze(1), self.q0_embedding(q0).unsqueeze(1)], dim=1 ) output = self.backbone( inputs_embeds=inputs, use_cache=True, attention_mask={"full_attention": attention_masks[0]} if attention_masks else None, ) cache = output.past_key_values codes = [q0] for group, head in enumerate(self.heads): code = sample_token( head(output.last_hidden_state[:, -1]), do_sample=do_sample, top_k=top_k, top_p=top_p, temperature=temperature, generator=generator, ) codes.append(code) if group < len(self.code_embeddings): output = self.backbone( inputs_embeds=self.code_embeddings[group](code).unsqueeze(1), past_key_values=cache, use_cache=True, attention_mask={"full_attention": attention_masks[group + 1]} if attention_masks else None, ) cache = output.past_key_values return torch.stack(codes, dim=-1)