Text-to-Speech
English
German
voice-acting
qwen3
moss-audio-tokenizer-v2
audio-generation
Humaneness-Voice-Small / code /conditioning.py
ChristophSchuhmann's picture
Document architecture, prompts, code, and full run statistics
d911efa verified
Raw History Blame Contribute Delete
7.5 kB
"""Projected score tokens. Missing measurements are masked, never imputed scores.
The schema file pins names, units and measurement scopes. A tokenizer reserves
six textual marker IDs inside M2's existing padded vocabulary; its original
embedding rows and all original model weights remain loadable without resizing.
"""
from contextlib import contextmanager
import json
from pathlib import Path
import sys
import torch
from torch import nn
SMALL = Path('/e/scratch/reformo/schuhmann1_moss/code/small_tts')
sys.path.insert(0, str(SMALL))
MARKERS = ('<|score_global_start|>', '<|score_global_slot|>', '<|score_global_end|>',
'<|score_sentence_start|>', '<|score_sentence_slot|>', '<|score_sentence_end|>')
GLOBAL_TOKENS = 5
LOCAL_TOKENS = 10
def install_markers(tokenizer, vocab_size):
import inspect
parameters = inspect.signature(tokenizer.add_special_tokens).parameters
option = ('replace_extra_special_tokens' if 'replace_extra_special_tokens' in parameters
else 'replace_additional_special_tokens')
tokenizer.add_special_tokens({'additional_special_tokens': list(MARKERS)}, **{option: False})
ids = dict(zip(MARKERS, tokenizer.convert_tokens_to_ids(list(MARKERS))))
assert len(set(ids.values())) == 6 and all(0 <= x < vocab_size for x in ids.values()), ids
for marker, token in ids.items():
assert tokenizer.encode(marker, add_special_tokens=False) == [token]
return ids
def feature_vector(values, dimensions):
"""Return values scaled by documented bounds and an explicit validity mask.
Missing values become numerical zero only alongside mask=0. Normalization
bounds are scale anchors, not clipping bounds: actual regression outputs can
exceed nominal score ranges. Only non-finite measurements (or violations of
explicitly configured hard bounds) are rejected.
"""
import math
scaled, mask = [], []
for dim in dimensions:
value = values.get(dim['name'])
valid = value is not None
if valid:
value = float(value)
low, high = map(float, dim['normalization_bounds'])
legal = dim.get('hard_bounds')
if not math.isfinite(value) or high <= low or (legal is not None and not legal[0] <= value <= legal[1]):
raise ValueError(f'Invalid measured feature {dim["name"]}={value}')
value = 2 * (value - low) / (high - low) - 1
scaled.append(value if valid else 0.0)
mask.append(float(valid))
return scaled + mask
class ScoreConditioner(nn.Module):
def __init__(self, global_dims, local_dims, hidden=1024):
super().__init__()
self.global_dims, self.local_dims = global_dims, local_dims
self.hidden = hidden
self.global_projection = nn.Sequential(nn.Linear(2 * global_dims, 256), nn.SiLU(),
nn.Linear(256, GLOBAL_TOKENS * hidden))
self.local_projection = nn.Sequential(nn.Linear(2 * local_dims, 512), nn.SiLU(),
nn.Linear(512, LOCAL_TOKENS * hidden))
self.boundaries = nn.Embedding(4, hidden)
for module in self.modules():
if isinstance(module, (nn.Linear, nn.Embedding)):
nn.init.normal_(module.weight, mean=0.0, std=0.02)
if isinstance(module, nn.Linear):
nn.init.zeros_(module.bias)
def forward(self, global_values, local_values):
# At least one padded sentence keeps the same DDP graph in every mode.
B, S, _ = local_values.shape
g = self.global_projection(global_values).reshape(B, GLOBAL_TOKENS, self.hidden)
local = self.local_projection(local_values).reshape(B, S, LOCAL_TOKENS, self.hidden)
boundary = self.boundaries.weight.to(g.dtype)
gs = boundary[0].reshape(1, 1, -1).expand(B, 1, -1)
ge = boundary[1].reshape(1, 1, -1).expand(B, 1, -1)
ls = boundary[2].reshape(1, 1, 1, -1).expand(B, S, 1, -1)
le = boundary[3].reshape(1, 1, 1, -1).expand(B, S, 1, -1)
local = torch.cat([ls, local, le], dim=2).reshape(B, S * (LOCAL_TOKENS + 2), self.hidden)
# Bank index zero is not a score token. Gradients still traverse the
# complete bank with zero contribution when a batch has no score mode.
zero = torch.zeros(B, 1, self.hidden, device=g.device, dtype=g.dtype)
return torch.cat([zero, gs, g, ge, local], dim=1)
def scored_class():
import moss_small
parent = moss_small.MossTTSSmallModel.cls()
class ScoreM2(parent):
def __init__(self, config, schema):
super().__init__(config)
self.score_conditioner = ScoreConditioner(len(schema['global_dimensions']), len(schema['local_dimensions']),
hidden=int(config.hidden_size))
self._generation_conditioning = None
def score_embeddings(self, input_ids, conditioning):
embeds = super()._build_inputs_embeds(input_ids)
if conditioning is None:
return embeds
values, local, index = conditioning
bank = self.score_conditioner(values, local).to(embeds.dtype)
if index.shape[1] != input_ids.shape[1]:
# Cached generation after prefill has only new audio rows.
if input_ids.shape[1] != 1:
raise ValueError('Conditioning positions do not match the prompt; KV-cache generation is required')
return embeds
scores = bank.gather(1, index.unsqueeze(-1).expand(-1, -1, bank.shape[-1]))
return torch.where(index.unsqueeze(-1).gt(0), scores, embeds)
def _build_inputs_embeds(self, input_ids):
return self.score_embeddings(input_ids, self._generation_conditioning)
def forward(self, input_ids=None, attention_mask=None, score_conditioning=None, **kwargs):
if score_conditioning is not None:
kwargs['inputs_embeds'] = self.score_embeddings(input_ids, score_conditioning)
input_ids = None
return super().forward(input_ids=input_ids, attention_mask=attention_mask, **kwargs)
@contextmanager
def generation_scores(self, conditioning):
if self._generation_conditioning is not None:
raise RuntimeError('Overlapping generation on one score-conditioned model is unsupported')
self._generation_conditioning = conditioning
try:
yield
finally:
self._generation_conditioning = None
return ScoreM2
def load_start_model(schema_path, checkpoint, log=print):
import moss_small
schema = json.loads(Path(schema_path).read_text())
config = moss_small.make_config('M2')
model = scored_class()(config, schema)
state = torch.load(Path(checkpoint) / 'model_bf16.pt', map_location='cpu', weights_only=True)
missing, unexpected = model.load_state_dict(state, strict=False)
expected = {'score_conditioner.' + name for name in model.score_conditioner.state_dict()}
if set(missing) != expected or unexpected:
raise RuntimeError(f'Initial M2 checkpoint mismatch: missing={missing}, unexpected={unexpected}')
del state
model.tie_weights()
log(f'Loaded final M2 checkpoint {checkpoint}; new score parameters={sum(p.numel() for p in model.score_conditioner.parameters())}')
return model, config, schema