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