Text-to-Speech
English
German
voice-acting
qwen3
moss-audio-tokenizer-v2
audio-generation
File size: 8,193 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
155
156
157
158
159
160
161
162
163
"""Preserve the inherited prompt skeleton and inject measured score slots.

Only the global caption and the leading sentence-cue parentheses are replaced
in score mode. The inherited template repeats the SCRIPT in Instruction and
Text; both copies map to the same measured local feature bank.
"""
import re
import sys
from pathlib import Path
import numpy as np
import torch
from conditioning import MARKERS, feature_vector, install_markers

sys.path.insert(0, '/e/scratch/reformo/schuhmann1_moss/code/small_tts')
from pack_moss import MossPacker
from prompt_fmt import user_content

GLOBAL_BLOCK = MARKERS[0] + MARKERS[1] * 5 + MARKERS[2]
LOCAL_BLOCK = MARKERS[3] + MARKERS[4] * 10 + MARKERS[5]


def score_prompt(prompt, sentence_count):
    if not prompt.startswith('GENERAL: ') or '\nSCRIPT:\n' not in prompt:
        raise ValueError('Score mode requires the verified GENERAL/SCRIPT prompt format')
    _, script = prompt.split('\nSCRIPT:\n', 1)
    lines = script.splitlines()
    if len(lines) != sentence_count or not lines:
        raise ValueError(f'Prompt cues and measured sentence rows disagree: {len(lines)} vs {sentence_count}')
    edited, bodies = [], []
    for line in lines:
        match = re.match(r'^\([^)]*\)(.*)$', line)
        if not match:
            raise ValueError('Expected one leading parenthesized sentence cue')
        body = match.group(1)
        bodies.append(body)
        edited.append(LOCAL_BLOCK + body)
    changed = 'GENERAL: ' + GLOBAL_BLOCK + '\nSCRIPT:\n' + '\n'.join(edited)
    # Everything following the cue stays byte-identical, including burst tags,
    # pause spellings, numbers and punctuation. Durations live in the template.
    assert all(new[len(LOCAL_BLOCK):] == old for new, old in zip(edited, bodies))
    return changed


def labels_for(ids, assist, aend, audio_pad):
    labels = torch.full_like(ids, -100)
    labels[:-1] = ids[1:]
    targets = labels[:, 0]
    supervised = targets.eq(assist) | targets.eq(aend)
    first = ids[:, 0].eq(assist).nonzero()
    if not first.numel():
        raise ValueError('No assistant target frames in the packed training example')
    supervised[:max(0, int(first[0]) - 1)] = False
    labels[~supervised] = -100
    audio = labels[:, 1:]
    audio[(audio >= audio_pad) | (audio < 0)] = -100
    return labels


def generated_audio(result, config):
    """MOSS generate returns BOA + 13-channel rows, not a [frames,12] array.

    Follow the already validated M1/M2 evaluator's extraction. Any unexpected
    remaining code value is an error rather than silently clipping corruption.
    """
    start_length, sequence = result[0]
    assert start_length == 0, 'This continuation path requires an empty target prefix'
    assert sequence.ndim == 2 and sequence.shape[1] == int(config.n_vq) + 1
    assert int(sequence[0, 0]) == int(config.audio_start_token_id)
    audio = sequence[1:, 1:]
    keep = (audio[:, 0] != int(config.audio_pad_token_id)) & (audio[:, 0] >= 0)
    audio = audio[keep]
    if audio.numel():
        assert int(audio.min()) >= 0 and int(audio.max()) < int(config.audio_pad_token_id)
    return audio


class ScorePacker(MossPacker):
    def __init__(self, proc, config, schema):
        super().__init__(proc, config)
        self.schema = schema
        self.markers = install_markers(proc.tokenizer, int(config.vocab_size))

    def _prepare(self, meta, mode, ref_codes):
        if mode not in ('instruction', 'reference', 'score'):
            raise ValueError(mode)
        if (ref_codes is not None) != (mode == 'reference'):
            raise ValueError('Reference mode must have a genuine non-target reference')
        record = dict(meta)
        global_values = [0.] * (2 * len(self.schema['global_dimensions']))
        local_values = []
        if mode == 'score':
            measurements = meta['conditioning_measurements']
            if not measurements.get('validated'):
                raise ValueError('Score mode requires independently validated measurement provenance')
            local = measurements['sentences']
            record['prompt'] = score_prompt(meta['prompt'], len(local))
            global_values = feature_vector(measurements['global'], self.schema['global_dimensions'])
            for row in local:
                if row['scope'] not in ('sentence_audio', 'whole_target_single_cue_rescored', 'windowed_sentence_audio') or not row.get('provenance'):
                    raise ValueError('Whole-clip or unknown measurements cannot populate local score slots')
                local_values.append(feature_vector(row['values'], self.schema['local_dimensions']))
            if not local_values:
                raise ValueError('No local measurements for score mode')
        return record, global_values, local_values

    def _index(self, ids, mode, sentence_count):
        index = torch.zeros(ids.shape[0], dtype=torch.long)
        text = ids[:, 0].tolist()
        global_start = self.markers[MARKERS[0]]
        local_start = self.markers[MARKERS[3]]
        gs = [i for i, token in enumerate(text) if token == global_start]
        ls = [i for i, token in enumerate(text) if token == local_start]
        if mode != 'score':
            if gs or ls:
                raise ValueError('Score markers unexpectedly present outside score mode')
            return index
        if len(gs) != 1 or len(ls) != 2 * sentence_count:
            raise ValueError(f'Unexpected template score-block coverage: global={len(gs)}, sentence={len(ls)}')
        expect_global = [global_start] + [self.markers[MARKERS[1]]] * 5 + [self.markers[MARKERS[2]]]
        expect_local = [local_start] + [self.markers[MARKERS[4]]] * 10 + [self.markers[MARKERS[5]]]
        i = gs[0]
        assert text[i:i + 7] == expect_global
        index[i:i + 7] = torch.arange(1, 8)
        for repeat, i in enumerate(ls):
            assert text[i:i + 12] == expect_local
            base = 8 + (repeat % sentence_count) * 12
            index[i:i + 12] = torch.arange(base, base + 12)
        return index

    def pack_mode(self, meta, codes, mode, ref_codes=None, generation=False):
        record, glob, local = self._prepare(meta, mode, ref_codes)
        content = user_content(record, mode == 'reference')
        user = {'role': 'user', 'content': content,
                'audio_codes_list': [torch.as_tensor(np.asarray(ref_codes).copy(), dtype=torch.long)] if ref_codes is not None else []}
        conversation = [user]
        if not generation:
            conversation.append(self.proc.build_assistant_message(audio_codes_list=[torch.as_tensor(np.asarray(codes), dtype=torch.long)]))
        packed = self.proc([conversation], mode='generation' if generation else 'computing_loss')
        ids = packed['input_ids'][0].long()
        index = self._index(ids, mode, len(local))
        labels = torch.full_like(ids, -100) if generation else labels_for(ids, self.assist, self.aend, self.audio_pad)
        if not generation:
            assert int((labels[:, 1] >= 0).sum()) == len(codes), 'Target frame supervision changed'
            assert int((labels[:, 0] == self.aend).sum()) == 1, 'Reference end leaked into the loss'
        return {'input_ids': ids, 'labels': labels, 'score_index': index,
                'global_values': glob, 'local_values': local, 'mode': mode}

    def collate(self, examples):
        from pack_moss import collate
        batch = collate(examples, self.text_pad, self.audio_pad)
        B, T = batch['attention_mask'].shape
        S = max(1, max(len(row['local_values']) for row in examples))
        glob = torch.tensor([row['global_values'] for row in examples], dtype=torch.float32)
        local = torch.zeros(B, S, 2 * len(self.schema['local_dimensions']), dtype=torch.float32)
        index = torch.zeros(B, T, dtype=torch.long)
        for i, row in enumerate(examples):
            index[i, :len(row['score_index'])] = row['score_index']
            if row['local_values']:
                local[i, :len(row['local_values'])] = torch.tensor(row['local_values'])
        batch['score_conditioning'] = (glob, local, index)
        batch['modes'] = [row['mode'] for row in examples]
        return batch