Download code/packing.py from laion/Humaneness-Voice-Small: direct link, hf CLI and curl.
- Browser
- Download file 8.19 kB
-
https://huggingface.co/laion/Humaneness-Voice-Small/resolve/main/code/packing.py
- Command line
-
hf download hf://laion/Humaneness-Voice-Small/code/packing.py
-
curl -L -o packing.py https://huggingface.co/laion/Humaneness-Voice-Small/resolve/main/code/packing.py
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 | |