#!/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)