Audio-Text-to-Text
Transformers
Safetensors
edgeinstant
feature-extraction
audio
text-to-speech
custom_code
Instructions to use chenjz24/EdgeIn-v3 with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use chenjz24/EdgeIn-v3 with Transformers:
# pip install -U transformers accelerate # Load model directly from transformers import AutoModel model = AutoModel.from_pretrained("chenjz24/EdgeIn-v3", trust_remote_code=True, device_map="auto") - Notebooks
- Google Colab
- Kaggle
File size: 9,385 Bytes
d81a465 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 | """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()
|