#!/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