Text-to-Speech
English
German
voice-acting
qwen3
moss-audio-tokenizer-v2
audio-generation
File size: 7,501 Bytes
d911efa
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
"""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