Download code/caption_curriculum_dataset.py from laion/Humaneness-Voice-Small: direct link, hf CLI and curl.
- Browser
- Download file 7.59 kB
-
https://huggingface.co/laion/Humaneness-Voice-Small/resolve/main/code/caption_curriculum_dataset.py
- Command line
-
hf download hf://laion/Humaneness-Voice-Small/code/caption_curriculum_dataset.py
-
curl -L -o caption_curriculum_dataset.py https://huggingface.co/laion/Humaneness-Voice-Small/resolve/main/code/caption_curriculum_dataset.py
7.59 kB
| #!/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", "<u4"), ("reference", "<u4"), ("annotation", "<u4"), | |
| ("family", "u1"), ("surface", "u1"), ("mode", "u1"), ("reserved", "u1"), | |
| ]) | |
| MISSING = np.iinfo(np.uint32).max | |
| class JsonlOffsets: | |
| def __init__(self, records: str | Path, offsets: str | Path): | |
| self.path = Path(records) | |
| self.offsets = np.load(offsets, mmap_mode="r", allow_pickle=False) | |
| self.fd = os.open(self.path, os.O_RDONLY) | |
| assert self.offsets.dtype == np.int64 and int(self.offsets[0]) == 0 | |
| assert int(self.offsets[-1]) == self.path.stat().st_size | |
| def row(self, index: int) -> 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 | |