Text-to-Speech
English
German
voice-acting
qwen3
moss-audio-tokenizer-v2
audio-generation
Humaneness-Voice-Small / code /cascade_dataset.py
ChristophSchuhmann's picture
Document architecture, prompts, code, and full run statistics
d911efa verified
Raw History Blame Contribute Delete
5.3 kB
#!/usr/bin/env python3
"""M2-600M cascade adapter with deterministic timed-prompt repair.
Audio/code targets and manifests remain immutable. Every presentation passes
through ``timed_prompt.repair_prompt``; repair facts are returned in accounting
so the trainer can prove the observed timed/untimed mixture.
"""
import json
import os
import numpy as np
from timed_prompt import repair_prompt
MODES = ('instruction', 'reference')
TAR_MEMBER_MAP = '/e/scratch/reformo/schuhmann1_moss/out/m2_600m_ladder/tar_member_map/manifest.json'
class CascadeRecords:
def __init__(self, manifest):
from pathlib import Path
self.manifest_path = Path(manifest)
spec = json.loads(self.manifest_path.read_text())
assert spec['status'] == 'complete' and spec['presentations'] > 0
root = self.manifest_path.parent
self.spec = spec
self.phase = str(spec.get('phase') or spec.get('ladder_stage') or self.manifest_path.parent.name)
self.offsets = np.load(root / spec['offsets'], mmap_mode='r', allow_pickle=False)
assert len(self.offsets) == spec['presentations'] + 1
assert self.offsets.dtype == np.int64 and int(self.offsets[0]) == 0
self.order = np.load(root / spec['order'], mmap_mode='r', allow_pickle=False)
assert len(self.order) == spec['presentations']
assert self.order.dtype in (np.dtype('uint32'), np.dtype('int64'))
# Full uniqueness is established once by the manifest verifier. Never
# materialize 48M Python ints/sets independently on all 128 DDP ranks.
assert int(self.order.min()) == 0 and int(self.order.max()) == spec['presentations'] - 1
self.path = root / spec['records']
self.fd = os.open(self.path, os.O_RDONLY)
assert int(self.offsets[-1]) == self.path.stat().st_size
if 'records_bytes' in spec:
assert self.path.stat().st_size == spec['records_bytes']
assert (root / spec['offsets']).stat().st_size == spec['offsets_bytes']
assert (root / spec['order']).stat().st_size == spec['order_bytes']
import sys
sys.path.insert(0, '/e/scratch/reformo/schuhmann1_moss/code/m2_20k_score')
from dataset import ShardReader
from tar_member_map import TarMemberMap
class MappedShardReader(ShardReader):
def __init__(self):
# ShardReader uses the same bounded capacity for Parquet
# reference tables as for TAR member dictionaries. TAR
# dictionaries are overridden below, but Parquet tables still
# need a non-zero LRU (capacity=0 caused an immediate eviction
# followed by move_to_end(KeyError) in S1 smoke 1892952).
super().__init__(capacity=16)
self.packed_members = TarMemberMap(TAR_MEMBER_MAP)
def member(self, path, name):
offset, size = self.packed_members.locate(path, name)
return self.bytes_at(path, offset, size)
self.reader = MappedShardReader()
def __len__(self):
return len(self.order)
def row(self, index):
j = int(self.order[index])
start, end = int(self.offsets[j]), int(self.offsets[j + 1])
return json.loads(os.pread(self.fd, end - start, start))
def example(self, index, packer):
row = self.row(index)
mode = row['mode']
form = row.get('form', 'A')
assert mode in MODES, row['uid']
assert form in ('A', 'B'), row['uid']
meta = self.reader.record(row['record_locator'])
assert meta['uid'] == row['uid']
if 'moss_frames' in meta:
meta['frames'] = int(meta['moss_frames'])
else:
# Sidecar element JSONs carry no moss_frames key (pre-upload state);
# frame count comes from the manifest locator (validated in codes()).
meta['frames'] = int(row['target_locator']['full_frames'])
codes = self.reader.codes(row['target_locator'])
assert len(codes) == meta['frames']
meta = dict(meta)
meta['prompt'], prompt_info = repair_prompt(
meta, caption=row.get('bude_caption'), form=form, mode=mode,
presentation=int(index), phase=self.phase)
reference = None
if mode == 'reference':
assert row.get('reference_uid') and row['reference_uid'] != row['uid']
assert row.get('reference_locator')
reference = self.reader.codes(row['reference_locator'])
assert 0 < len(reference) <= 37
result = packer.pack_mode(meta, codes, mode, reference)
result['accounting'] = {'uid': row['uid'], 'mode': mode, 'form': form,
'frames': len(codes),
'target_audio_hours': float(meta['dur_s']) / 3600,
'reference_frames': len(reference) if reference is not None else 0,
'prompt_timed': bool(prompt_info['timed']),
'prompt_eligible': bool(prompt_info['eligible']),
'duration_tags': int(prompt_info['duration_tags']),
'prompt_format_id': prompt_info['format_id']}
return result