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