Download code/moss_small.py from laion/Humaneness-Voice-Small: direct link, hf CLI and curl.
- Browser
- Download file 10.8 kB
-
https://huggingface.co/laion/Humaneness-Voice-Small/resolve/main/code/moss_small.py
- Command line
-
hf download hf://laion/Humaneness-Voice-Small/code/moss_small.py
-
curl -L -o moss_small.py https://huggingface.co/laion/Humaneness-Voice-Small/resolve/main/code/moss_small.py
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 | |
| 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) | |