Text-to-Speech
English
German
voice-acting
qwen3
moss-audio-tokenizer-v2
audio-generation
File size: 10,801 Bytes
d911efa
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
#!/usr/bin/env python3
"""M1 (Arm P) und M2 (Arm S): Qwen3-0.6B-Backbone + MOSS-Talker-Head.

M1: Head (lokaler Transformer, 12 Audio-Embeddings = 12 Audio-Köpfe, binärer Stop-Kopf) aus $SC/out/sft3/export,
    Hidden 2560, EINGEFROREN; trainierbar: Backbone (1024) + proj_out 1024->2560 (vor dem Head) + proj_in 2560->1024
    (Summe der Audio-Embeddings am Backbone-Eingang).  Der Head wird NIE in ein Backbone gemergt;
    audio_lm_heads.N.weight IST audio_embeddings.N.weight (gewichtsgeteilt).
M2: gleiche Architektur, Head mit Hidden 1024 zufällig initialisiert, alles trainierbar.

`hidden_size != local_hidden_size` ist im Original nicht vorgesehen (Config wirft, _global_hidden_to_local ist Identität,
Audio-Embeddings haben hidden_size) -> Unterklasse mit eigenem __init__ und Projektionen."""
import os, sys, json, importlib.util
import torch, torch.nn as nn
from transformers import Qwen3Config
from transformers.models.gpt2.configuration_gpt2 import GPT2Config
from safetensors import safe_open
from safetensors.torch import load_file

SC = '/e/scratch/reformo/schuhmann1_moss'
SFT3 = f'{SC}/out/sft3/export'
QWEN = f'{SC}/models/qwen3-0.6b'


def _load_export_modules():
    """modeling_moss_tts aus dem SFT3-Export als Paket importieren (relative Imports)."""
    import importlib
    if 'moss_export' in sys.modules:
        return sys.modules['moss_export']
    spec = importlib.util.spec_from_file_location('moss_export', f'{SFT3}/__init__.py',
                                                  submodule_search_locations=[SFT3])
    mod = importlib.util.module_from_spec(spec)
    sys.modules['moss_export'] = mod
    spec.loader.exec_module(mod)
    return mod


def export_classes():
    _load_export_modules()
    import importlib
    cfgm = importlib.import_module('moss_export.configuration_moss_tts')
    modm = importlib.import_module('moss_export.modeling_moss_tts')
    procm = importlib.import_module('moss_export.processing_moss_tts')
    return cfgm.MossTTSLocalConfig, modm.MossTTSLocalModel, procm.MossTTSLocalProcessor


def make_config(kind):
    """kind: 'M1' (local 2560, SFT3-Head) | 'M2' (local 1024, von Null)."""
    MossCfg, _, _ = export_classes()
    base = json.load(open(f'{SFT3}/config.json'))
    q = json.load(open(f'{QWEN}/config.json'))
    for k in ('architectures', 'transformers_version'):
        q.pop(k, None)
    q['use_cache'] = False
    q['gradient_checkpointing_use_reentrant'] = False
    g = dict(base['gpt2_config'])
    if kind == 'M1':
        local = 2560                         # n_embd 2560, n_head 32, n_inner 9728 bleiben (SFT3-Head)
    elif kind == 'M2':
        local = 1024
        g['n_embd'] = 1024; g['n_head'] = 16; g['n_inner'] = 4096
    else:
        raise ValueError(kind)
    g_tmp = dict(g); g_tmp['n_embd'] = int(q['hidden_size'])    # Config-Prüfung umgehen
    kw = {k: base[k] for k in ('audio_codebook_sizes', 'audio_pad_code', 'audio_pad_token_id', 'audio_start_token_id',
                                'audio_end_token_id', 'audio_user_slot_token_id', 'audio_assistant_slot_token_id',
                                'audio_vocab_size', 'im_end_token_id', 'im_start_token_id', 'n_vq', 'pad_token_id',
                                'local_text_head_mode', 'local_transformer_layers', 'sampling_rate',
                                'audio_tokenizer_name_or_path', 'use_static_local_kv_cache')}
    cfg = MossCfg(qwen3_config=q, gpt2_config=g_tmp, attn_implementation='sdpa',
                  local_transformer_attn_implementation='sdpa', **kw)
    cfg.gpt2_config = GPT2Config(**g)
    cfg.local_hidden_size = local
    cfg.small_kind = kind
    cfg.dtype = 'bfloat16'
    return cfg


class MossTTSSmallModel:
    """Fabrik: erzeugt die Unterklasse erst, wenn das Export-Modul geladen ist."""
    _cls = None

    @classmethod
    def cls(klass):
        if klass._cls is not None:
            return klass._cls
        _, MossTTSLocalModel, _ = export_classes()
        import importlib
        gpt2m = importlib.import_module('moss_export.gpt2_decoder')
        qw = importlib.import_module('moss_export.qwen3_decoder')

        class _Small(MossTTSLocalModel):
            def __init__(self, config):
                # Elternklasse überspringen (sie baut Embeddings mit hidden_size); PreTrainedModel-Init direkt.
                super(MossTTSLocalModel, self).__init__(config)
                self._tied_weights_keys = self._build_tied_weights_keys(config)
                config.qwen3_config.pad_token_id = config.pad_token_id
                config.qwen3_config._attn_implementation = config.attn_implementation
                lg = config.gpt2_config.to_dict()
                lg['n_layer'] = int(getattr(config, 'local_transformer_layers', config.gpt2_config.n_layer))
                lg['n_positions'] = int(config.n_vq) + 1
                lg['n_ctx'] = int(config.n_vq) + 1
                lg = GPT2Config(**lg)
                lg.pad_token_id = config.pad_token_id
                lg._attn_implementation = config.local_transformer_attn_implementation
                self.transformer = qw.MossQwen3Model(config.qwen3_config)
                self.local_transformer = gpt2m.MossTTSNanoGPT2Model(lg, attn_implementation=config.local_transformer_attn_implementation)
                self.local_transformer.wte = nn.Identity()
                H = int(config.hidden_size); L = int(config.local_hidden_size)
                self.audio_embeddings = nn.ModuleList([nn.Embedding(int(config.audio_codebook_sizes[i]), L) for i in range(config.n_vq)])
                self.text_lm_head = nn.Linear(H, int(config.vocab_size), bias=False)
                self.audio_lm_heads = nn.ModuleList([nn.Linear(L, int(config.audio_codebook_sizes[i]), bias=False) for i in range(config.n_vq)])
                self.local_text_lm_head = nn.Linear(L, 2, bias=False) if self._use_binary_local_text_head() else None
                if H != L:
                    self.proj_in = nn.Linear(L, H, bias=False)      # Audio-Embedding-Summe -> Backbone
                    self.proj_out = nn.Linear(H, L, bias=False)     # Backbone-Hidden -> Head
                else:
                    self.proj_in = nn.Identity(); self.proj_out = nn.Identity()
                self.post_init()
                self.tie_weights()
                self.initialize_local_text_lm_head_from_text_lm_head()

            def _build_inputs_embeds(self, input_ids):
                if input_ids.ndim != 3 or input_ids.shape[-1] != self.config.n_vq + 1:
                    raise ValueError(f'Expected [B,T,{self.config.n_vq + 1}], got {tuple(input_ids.shape)}')
                text_ids = input_ids[..., 0]
                emb = self.transformer.embed_tokens(text_ids)
                acc = None
                for ci, e in enumerate(self.audio_embeddings):
                    ids = input_ids[..., ci + 1]
                    valid = ids.ne(self.config.audio_pad_token_id)
                    a = e(ids.masked_fill(~valid, 0)) * valid.unsqueeze(-1)
                    acc = a if acc is None else acc + a
                return emb + self.proj_in(acc.to(emb.dtype))

            def _global_hidden_to_local(self, h):
                return self.proj_out(h)

        klass._cls = _Small
        return _Small


def load_qwen_backbone(model, log=print):
    sd = load_file(f'{QWEN}/model.safetensors', device='cpu')
    tgt = model.transformer.state_dict()
    loaded = {}
    for k, v in sd.items():
        if k.startswith('model.'):
            kk = k[len('model.'):]
            if kk in tgt:
                assert tuple(v.shape) == tuple(tgt[kk].shape), (k, v.shape, tgt[kk].shape)
                loaded[kk] = v.to(tgt[kk].dtype)
    missing, unexpected = model.transformer.load_state_dict(loaded, strict=False)
    missing = [m for m in missing if 'rotary' not in m]
    if missing or unexpected:
        raise RuntimeError(f'Qwen3 load: missing={missing[:10]} unexpected={unexpected[:10]}')
    extra = [k for k in sd if not k.startswith('model.')]
    log(f'[qwen3] loaded {len(loaded)} tensors into backbone; non-backbone keys in file: {extra}')
    model.tie_weights()


def load_sft3_head(model, log=print):
    """Lokaler Transformer + 12 Audio-Embeddings (+ Köpfe, gewichtsgeteilt) + binärer Stop-Kopf aus dem SFT3-Export."""
    want_prefix = ('local_transformer.', 'audio_embeddings.', 'local_text_lm_head.')
    loaded = {}
    with safe_open(f'{SFT3}/model.safetensors', framework='pt', device='cpu') as f:
        keys = list(f.keys())
        for k in keys:
            if k.startswith(want_prefix):
                loaded[k] = f.get_tensor(k)
    sd = model.state_dict()
    for k, v in loaded.items():
        assert k in sd, k
        assert tuple(v.shape) == tuple(sd[k].shape), (k, v.shape, sd[k].shape)
    missing, unexpected = model.load_state_dict({k: v.to(sd[k].dtype) for k, v in loaded.items()}, strict=False)
    assert not unexpected, unexpected
    model.tie_weights()
    # Kontrolle: Head-Gewichte sind Embeddings
    for i in range(model.config.n_vq):
        assert model.audio_lm_heads[i].weight.data_ptr() == model.audio_embeddings[i].weight.data_ptr()
    n = sum(v.numel() for v in loaded.values())
    log(f'[sft3-head] loaded {len(loaded)} tensors ({n/1e6:.1f}M params) — local_transformer + audio_embeddings + stop head')
    return loaded.keys()


def freeze_head(model, log=print):
    n = 0
    for mod in (model.local_transformer, model.audio_embeddings, model.local_text_lm_head):
        for p in mod.parameters():
            p.requires_grad_(False); n += p.numel()
    log(f'[freeze] head frozen: {n/1e6:.1f}M params')


def build(kind, log=print, dtype=torch.float32):
    cfg = make_config(kind)
    cls = MossTTSSmallModel.cls()
    model = cls(cfg)
    load_qwen_backbone(model, log)
    if kind == 'M1':
        load_sft3_head(model, log)
        freeze_head(model, log)
    ntr = sum(p.numel() for p in model.parameters() if p.requires_grad)
    ntot = sum(p.numel() for p in model.parameters())
    log(f'[{kind}] params total={ntot/1e6:.1f}M trainable={ntr/1e6:.1f}M hidden={cfg.hidden_size} local={cfg.local_hidden_size}')
    return model.to(dtype), cfg


def param_groups(model, kind, lr_backbone, lr_new):
    """Backbone (vortrainiert) vs. neue Teile (Projektionen / Head von Null)."""
    bb, new = [], []
    for n, p in model.named_parameters():
        if not p.requires_grad:
            continue
        (bb if n.startswith('transformer.') or n.startswith('text_lm_head.') else new).append(p)
    return [{'params': bb, 'lr': lr_backbone, 'name': 'backbone'}, {'params': new, 'lr': lr_new, 'name': 'new'}]


if __name__ == '__main__':
    kind = sys.argv[1] if len(sys.argv) > 1 else 'M2'
    m, cfg = build(kind)
    print(m.__class__.__name__, cfg.small_kind)