Text-to-Speech
English
German
voice-acting
qwen3
moss-audio-tokenizer-v2
audio-generation
ChristophSchuhmann's picture
Document architecture, prompts, code, and full run statistics
d911efa verified
Raw History Blame Contribute Delete
8.19 kB
"""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