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