EdgeIn-v1 / teacher.py
chenjz24's picture
Upload folder using huggingface_hub
f74eb65 verified
Raw History Blame Contribute Delete
8.59 kB
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