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