"""Compact inference representation of the frozen native-token acoustic model.""" from __future__ import annotations import copy from pathlib import Path import torch from torch import nn from .native_clock_talker import NativeClockTalker from .residual_predictor import Qwen3TTSResidualPredictor class ProjectedTextEmbedding(nn.Module): """Keep projected rows reachable from the complete native byte vocabulary.""" def __init__(self, weight: torch.Tensor, teacher_ids: torch.Tensor, vocabulary_size: int): super().__init__() self.embedding = nn.Embedding.from_pretrained(weight, freeze=True) rows = torch.full((vocabulary_size,), -1, device=weight.device, dtype=torch.long) rows[teacher_ids.to(weight.device)] = torch.arange(len(teacher_ids), device=weight.device) self.register_buffer("teacher_to_row", rows) def forward(self, teacher_ids: torch.Tensor) -> torch.Tensor: return self.embedding(self.teacher_to_row[teacher_ids]) class CompactNativeClockTalker(NativeClockTalker): """Fold frozen text projection and share identical codec input embeddings.""" serialization_buffers = ("prefix", "prefix_english", "prefix_auto", "tts_eos", "tts_pad") tied_weight_keys = { "residual_predictor.q0_embedding.weight": "codec_history_embeddings.0.weight", **{ f"residual_predictor.code_embeddings.{group}.weight": f"codec_history_embeddings.{group + 1}.weight" for group in range(14) }, } def prepare_for_serialization(self) -> "CompactNativeClockTalker": """Save speaker conditions and the runtime precision of rotary frequencies.""" self._non_persistent_buffers_set.difference_update(self.serialization_buffers) for backbone in (self.backbone, self.residual_predictor.backbone): backbone.rotary_emb._non_persistent_buffers_set.difference_update( ("inv_freq", "original_inv_freq") ) return self def get_config(self) -> dict: """Describe the compact model without paths to its initialization assets.""" residual_config = self.residual_predictor.backbone.config.to_dict() residual_config["rope_theta"] = self.residual_predictor.backbone.config.rope_parameters["rope_theta"] return { "backbone_config": self.backbone.config.to_dict(), "residual_predictor_config": residual_config, "attn_implementation": self.backbone.config._attn_implementation, "hidden_size": self.hidden_size, "codebook_size": self.codebook_size, "num_code_groups": len(self.codec_history_embeddings), "projected_num_embeddings": self.text_embedding.embedding.num_embeddings, "teacher_vocab_size": self.text_embedding.teacher_to_row.numel(), "prefix_shapes": { name: list(getattr(self, name).shape) for name in self.serialization_buffers }, "native_teacher_ids": list(self.native_teacher_ids), "native_teacher_offsets": list(self.native_teacher_offsets), } @classmethod def from_config( cls, config: dict, *, attn_implementation: str | None = None, ) -> "CompactNativeClockTalker": """Build the compact modules and shared weights for checkpoint loading.""" from transformers import Qwen3Config, Qwen3Model model = cls() model.hidden_size = int(config["hidden_size"]) model.codebook_size = int(config["codebook_size"]) model.eos_class = model.codebook_size backbone_config = Qwen3Config(**config["backbone_config"]) backbone_config._attn_implementation = attn_implementation or config["attn_implementation"] model.backbone = Qwen3Model(backbone_config) model.backbone.embed_tokens = None model.text_embedding = ProjectedTextEmbedding( torch.empty(int(config["projected_num_embeddings"]), model.hidden_size), torch.empty(0, dtype=torch.long), int(config["teacher_vocab_size"]), ) model.text_projection = nn.Identity() model.codec_history_embeddings = nn.ModuleList( nn.Embedding(model.codebook_size, model.hidden_size) for _ in range(int(config["num_code_groups"])) ) model.codec_bos = nn.Parameter(torch.empty(model.hidden_size)) model.q0_head = nn.Linear(model.hidden_size, model.codebook_size, bias=False) model.codec_eos_head = nn.Linear(model.hidden_size, 1, bias=False) model.residual_predictor = Qwen3TTSResidualPredictor( model.hidden_size, model.codebook_size, config["residual_predictor_config"], ) model.residual_predictor.q0_embedding = model.codec_history_embeddings[0] for group in range(len(model.residual_predictor.code_embeddings)): model.residual_predictor.code_embeddings[group] = model.codec_history_embeddings[group + 1] model.native_teacher_ids = list(config["native_teacher_ids"]) model.native_teacher_offsets = list(config["native_teacher_offsets"]) for name in cls.serialization_buffers: model.register_buffer(name, torch.empty(config["prefix_shapes"][name])) return model.prepare_for_serialization() @classmethod def from_teacher( cls, teacher_path: str | Path, native_tokenizer_path: str | Path, speaker_vector_path: str | Path, *, dtype: torch.dtype = torch.float32, attn_implementation: str = "eager", device: str | torch.device = "cpu", projection_batch_size: int = 1, projected_table_path: str | Path | None = None, ) -> "CompactNativeClockTalker": source = NativeClockTalker.from_teacher( teacher_path, native_tokenizer_path, speaker_vector_path, dtype=dtype, attn_implementation=attn_implementation, ).eval().to(device) return cls.from_native(source, projection_batch_size=projection_batch_size, projected_table_path=projected_table_path) @classmethod @torch.no_grad() def from_native( cls, source: NativeClockTalker, *, projection_batch_size: int = 1, projected_table_path: str | Path | None = None, ) -> "CompactNativeClockTalker": """Reuse acoustic weights; leave the source model's modules unmodified. A batch size of one repeats the streaming projection's matrix shape. Larger batches use GEMM and may change floating-point rounding. """ model = cls() model.hidden_size = source.hidden_size for name in ("backbone", "codec_history_embeddings", "q0_head", "codec_eos_head"): setattr(model, name, getattr(source, name)) model.codec_bos = source.codec_bos model.native_teacher_ids = list(source.native_teacher_ids) model.native_teacher_offsets = list(source.native_teacher_offsets) for name, buffer in source.named_buffers(recurse=False): model.register_buffer(name, buffer, persistent=name not in source._non_persistent_buffers_set) device = source.text_embedding.weight.device teacher_ids = torch.tensor(sorted(set(source.native_teacher_ids)), device=device) table_path = Path(projected_table_path) if projected_table_path else None if table_path is not None and table_path.exists(): from safetensors.torch import load_file table = load_file(str(table_path), device=str(device)) if not torch.equal(table["teacher_ids"], teacher_ids): raise ValueError("Projected table belongs to a different native vocabulary") projected = table["projected"].to(dtype=source.text_embedding.weight.dtype) else: projected = source.q0_head.weight.new_empty((len(teacher_ids), source.hidden_size)) for start in range(0, len(teacher_ids), projection_batch_size): ids = teacher_ids[start:start + projection_batch_size].reshape(-1, 1) projected[start:start + len(ids)] = source.text_projection(source.text_embedding(ids))[:, 0] if table_path is not None: from safetensors.torch import save_file table_path.parent.mkdir(parents=True, exist_ok=True) save_file({"teacher_ids": teacher_ids.cpu(), "projected": projected.cpu()}, str(table_path)) model.text_embedding = ProjectedTextEmbedding(projected, teacher_ids, source.text_embedding.num_embeddings) model.text_projection = nn.Identity() # These fifteen input tables were copied from the same teacher tensors. model.residual_predictor = copy.deepcopy(source.residual_predictor) pairs = [(model.residual_predictor, "q0_embedding", model.codec_history_embeddings[0])] pairs.extend((model.residual_predictor.code_embeddings, str(group), model.codec_history_embeddings[group + 1]) for group in range(len(model.residual_predictor.code_embeddings))) for owner, name, history_embedding in pairs: if not torch.equal(getattr(owner, name).weight, history_embedding.weight): raise ValueError("Codec input embeddings have diverged and cannot be shared") setattr(owner, name, history_embedding) return model.eval()