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()