#!/usr/bin/env python3 """Lazy container-backed loader for the multi-caption S1->S10 curriculum.""" from __future__ import annotations import hashlib import json import os from pathlib import Path import random import numpy as np from timed_prompt import repair_prompt TAR_MEMBER_MAP = "/e/scratch/reformo/schuhmann1_moss/out/m2_600m_ladder/tar_member_map/manifest.json" FAMILIES = ("general_script", "bude_v1", "bude_v1_1", "timbre", "voice_tagging", "gemma", "transcript_only") DESC_DTYPE = np.dtype([ ("base", " dict: start, end = int(self.offsets[index]), int(self.offsets[index + 1]) return json.loads(os.pread(self.fd, end - start, start)) class PhysicalBase(JsonlOffsets): def __init__(self, manifest: str | Path): path = Path(manifest) spec = json.loads(path.read_text()) root = path.parent super().__init__(root / spec["records"], root / spec["offsets"]) def shuffled(tags: list[str], uid: str, stage: str, presentation: int, family: str) -> list[str]: result = list(tags) seed = int.from_bytes(hashlib.blake2b( f"20260923|{stage}|{uid}|{family}|{presentation}".encode(), digest_size=8).digest(), "little") random.Random(seed).shuffle(result) return result def caption_prompt(caption: str, transcript: str) -> str: caption = " ".join(str(caption).replace('"', "'").split()) assert caption return f'CAPTION: {caption}\nTRANSCRIPT: "{transcript}"' class CaptionCurriculumRecords: def __init__(self, manifest: str | Path): self.manifest_path = Path(manifest) self.spec = json.loads(self.manifest_path.read_text()) assert self.spec["status"] == "complete" assert self.spec["format"] == "m2-caption-curriculum-v1" root = self.manifest_path.parent self.phase = self.spec["stage"] self.descriptors = np.load(root / self.spec["descriptors"], mmap_mode="r", allow_pickle=False) self.order = np.load(root / self.spec["order"], mmap_mode="r", allow_pickle=False) assert self.descriptors.dtype == DESC_DTYPE assert len(self.descriptors) == len(self.order) == self.spec["presentations"] self.base = PhysicalBase(self.spec["base_manifest"]) assets = json.loads(Path(self.spec["assets_manifest"]).read_text()) assert assets["status"] == "complete" asset_root = Path(self.spec["assets_manifest"]).parent self.annotations = JsonlOffsets(asset_root / assets["annotations"]["records"], asset_root / assets["annotations"]["offsets"]) self.references = JsonlOffsets(asset_root / assets["references"]["records"], asset_root / assets["references"]["offsets"]) 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): 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 example(self, index: int, packer): descriptor = self.descriptors[int(self.order[index])] base = self.base.row(int(descriptor["base"])) uid = base["uid"] family = FAMILIES[int(descriptor["family"])] mode = "reference" if int(descriptor["mode"]) else "instruction" meta = dict(self.reader.record(base["record_locator"])) assert meta["uid"] == uid meta["frames"] = int(meta.get("moss_frames", base["target_locator"]["full_frames"])) target = self.reader.codes(base["target_locator"]) assert len(target) == meta["frames"] transcript = str(meta["text"]) annotation = None if int(descriptor["annotation"]) != MISSING: annotation = self.annotations.row(int(descriptor["annotation"])) assert annotation["uid"] == uid prompt_info = {"timed": False, "eligible": False, "duration_tags": 0, "format_id": "m2-multiformat-caption-v1-20260923"} form = "B" if family == "general_script": meta["prompt"], prompt_info = repair_prompt( meta, caption=None, form="A", mode=mode, presentation=int(index), phase=self.phase) form = "A" elif family == "bude_v1": assert base.get("bude_caption") meta["prompt"] = caption_prompt(base["bude_caption"], transcript) elif family == "bude_v1_1": meta["prompt"] = caption_prompt(annotation["bude"], transcript) elif family in ("timbre", "voice_tagging"): prefix = "timbre" if family == "timbre" else "voice" tags = shuffled(annotation[prefix + "_tags"], uid, self.phase, int(index), family) prose = annotation[prefix + "_prose"] surface = int(descriptor["surface"]) if surface == 0: assert tags body = ", ".join(tags) elif surface == 1: body = prose else: assert tags and prose body = "Tags: " + ", ".join(tags) + ". Caption: " + prose meta["prompt"] = caption_prompt(body, transcript) elif family == "gemma": meta["prompt"] = caption_prompt(annotation["gemma_caption"], transcript) elif family == "transcript_only": meta["prompt"] = f'TRANSCRIPT: "{transcript}"' form = "A" else: raise AssertionError(family) reference = None if mode == "reference": assert int(descriptor["reference"]) != MISSING reference_row = self.references.row(int(descriptor["reference"])) assert reference_row["uid"] == uid and reference_row["reference_uid"] != uid reference = self.reader.codes(reference_row["locator"]) assert 0 < len(reference) <= 37 result = packer.pack_mode(meta, target, mode, reference) result["accounting"] = { "uid": uid, "mode": mode, "form": form, "family": family, "frames": len(target), "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"] if family == "general_script" else "m2-multiformat-caption-v1-20260923"), } return result