from __future__ import annotations import json from pathlib import Path import numpy as np import torch from safetensors import safe_open from torch import nn import torch.nn.functional as F from .residual_predictor import Qwen3TTSResidualPredictor def native_teacher_mapping( native_tokenizer_path: str | Path, teacher_path: str | Path, native_vocabulary_size: int | None = None, ) -> tuple[torch.Tensor, torch.Tensor]: """Map exact ByteLevel pieces to teacher IDs, with empty special-token bags.""" native_path = Path(native_tokenizer_path) if native_path.is_dir(): native_path = native_path / "tokenizer.json" native = json.loads(native_path.read_text())["model"]["vocab"] teacher = json.loads((Path(teacher_path) / "vocab.json").read_text()) visible = list(range(33, 127)) + list(range(161, 173)) + list(range(174, 256)) inverse = {chr(byte): byte for byte in visible} inverse.update({chr(256 + index): byte for index, byte in enumerate(byte for byte in range(256) if byte not in visible)}) def raw(piece: str) -> bytes: return bytes(inverse[character] for character in piece) teacher_bytes = {raw(piece): index for piece, index in teacher.items()} by_first_byte: dict[int, dict[bytes, int]] = {} for piece, index in teacher_bytes.items(): by_first_byte.setdefault(piece[0], {})[piece] = index widths = {first: sorted({len(piece) for piece in pieces}, reverse=True) for first, pieces in by_first_byte.items()} native_by_id = {index: raw(piece) for piece, index in native.items()} size = native_vocabulary_size or (max(native_by_id) + 1) flat, offsets = [], [0] for native_id in range(size): piece = native_by_id.get(native_id, b"") if piece in teacher_bytes: flat.append(teacher_bytes[piece]) else: position = 0 while position < len(piece): candidates = by_first_byte[piece[position]] for width in widths[piece[position]]: if width > len(piece) - position: continue candidate = piece[position:position + width] if candidate in candidates: flat.append(candidates[candidate]) position += width break else: raise ValueError(f"Teacher vocabulary cannot represent native token {native_id}") offsets.append(len(flat)) return torch.tensor(flat, dtype=torch.long), torch.tensor(offsets, dtype=torch.long) class TeacherTextProjection(nn.Module): def __init__(self, text_hidden_size: int, hidden_size: int): super().__init__() self.linear_fc1 = nn.Linear(text_hidden_size, text_hidden_size) self.linear_fc2 = nn.Linear(text_hidden_size, hidden_size) def forward(self, hidden: torch.Tensor) -> torch.Tensor: return self.linear_fc2(F.silu(self.linear_fc1(hidden))) class TeacherAcousticModel(nn.Module): """Pretrained text, acoustic and residual modules shared by speech models.""" codebook_size = 2048 num_code_groups = 16 def __init__( self, config: dict, native_teacher_ids: torch.Tensor, native_teacher_offsets: torch.Tensor, speaker_vector: torch.Tensor, *, attn_implementation: str = "eager", ): super().__init__() from transformers import Qwen3Config, Qwen3Model self.config = dict(config) self.hidden_size = int(config["hidden_size"]) backbone_config = Qwen3Config( vocab_size=self.codebook_size, hidden_size=self.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"]), rms_norm_eps=float(config["rms_norm_eps"]), rope_theta=float(config["rope_theta"]), attention_bias=bool(config["attention_bias"]), attention_dropout=float(config.get("attention_dropout", 0.0)), tie_word_embeddings=False, use_cache=True, ) backbone_config._attn_implementation = attn_implementation self.backbone = Qwen3Model(backbone_config) self.backbone.embed_tokens = None self.text_embedding = nn.Embedding(int(config["text_vocab_size"]), int(config["text_hidden_size"])) self.text_projection = TeacherTextProjection(int(config["text_hidden_size"]), self.hidden_size) self.text_embedding.requires_grad_(False) self.text_projection.requires_grad_(False) self.register_buffer("native_teacher_ids", native_teacher_ids, persistent=False) self.register_buffer("native_teacher_offsets", native_teacher_offsets, persistent=False) self.register_buffer("speaker_vector", speaker_vector.reshape(self.hidden_size), persistent=False) self.codec_history_embeddings = nn.ModuleList( [nn.Embedding(self.codebook_size, self.hidden_size) for _ in range(self.num_code_groups)] ) self.codec_bos = nn.Parameter(torch.zeros(self.hidden_size)) self.q0_head = nn.Linear(self.hidden_size, self.codebook_size, bias=False) self.residual_predictor = Qwen3TTSResidualPredictor( self.hidden_size, self.codebook_size, config["code_predictor_config"] ) @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", ) -> "TeacherAcousticModel": teacher_path, native_tokenizer_path = Path(teacher_path), Path(native_tokenizer_path) config = json.loads((teacher_path / "config.json").read_text())["talker_config"] native_root = native_tokenizer_path if native_tokenizer_path.is_dir() else native_tokenizer_path.parent native_config = json.loads((native_root / "config.json").read_text()) native_size = int(native_config.get("text_config", native_config)["vocab_size"]) ids, offsets = native_teacher_mapping(native_tokenizer_path, teacher_path, native_size) speaker = torch.as_tensor(np.load(speaker_vector_path), dtype=torch.float32) model = cls(config, ids, offsets, speaker, attn_implementation=attn_implementation).to(dtype=dtype) with safe_open(teacher_path / "model.safetensors", framework="pt", device="cpu") as weights: with torch.no_grad(): for name, parameter in model.backbone.named_parameters(): parameter.copy_(weights.get_tensor(f"talker.model.{name}")) model.text_embedding.weight.copy_(weights.get_tensor("talker.model.text_embedding.weight")) for name, parameter in model.text_projection.named_parameters(): parameter.copy_(weights.get_tensor(f"talker.text_projection.{name}")) codec = weights.get_tensor("talker.model.codec_embedding.weight") model.codec_history_embeddings[0].weight.copy_(codec[:model.codebook_size]) model.codec_bos.copy_(codec[int(config["codec_bos_id"])]) model.q0_head.weight.copy_(weights.get_tensor("talker.codec_head.weight")[:model.codebook_size]) model.residual_predictor.q0_embedding.weight.copy_(codec[:model.codebook_size]) for name, parameter in model.residual_predictor.backbone.named_parameters(): parameter.copy_(weights.get_tensor(f"talker.code_predictor.model.{name}")) for group in range(15): embedding = weights.get_tensor(f"talker.code_predictor.model.codec_embedding.{group}.weight") model.codec_history_embeddings[group + 1].weight.copy_(embedding) if group < 14: model.residual_predictor.code_embeddings[group].weight.copy_(embedding) model.residual_predictor.heads[group].weight.copy_( weights.get_tensor(f"talker.code_predictor.lm_head.{group}.weight") ) return model