Text-to-Speech
English
German
voice-acting
qwen3
moss-audio-tokenizer-v2
audio-generation
Humaneness-Voice-Small / code /moss_small.py
ChristophSchuhmann's picture
Document architecture, prompts, code, and full run statistics
d911efa verified
Raw History Blame Contribute Delete
10.8 kB
#!/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)