File size: 8,590 Bytes
f74eb65
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
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