Download code/build_caption_stage.py from laion/Humaneness-Voice-Small: direct link, hf CLI and curl.
- Browser
- Download file 11.5 kB
-
https://huggingface.co/laion/Humaneness-Voice-Small/resolve/main/code/build_caption_stage.py
- Command line
-
hf download hf://laion/Humaneness-Voice-Small/code/build_caption_stage.py
-
curl -L -o build_caption_stage.py https://huggingface.co/laion/Humaneness-Voice-Small/resolve/main/code/build_caption_stage.py
11.5 kB
| #!/usr/bin/env python3 | |
| """Build one compact virtual-presentation index from an immutable base manifest.""" | |
| from __future__ import annotations | |
| from collections import Counter | |
| from datetime import datetime | |
| import hashlib | |
| import json | |
| import math | |
| import os | |
| from pathlib import Path | |
| import sys | |
| import numpy as np | |
| from caption_curriculum_dataset import DESC_DTYPE, MISSING | |
| SC = Path("/e/scratch/reformo/schuhmann1_moss") | |
| ROOT = SC / "out/m2_600m_ladder" | |
| OUT = ROOT / "caption_curriculum_v1/manifests" | |
| ASSETS = ROOT / "caption_curriculum_v1/assets/manifest.json" | |
| STAGES = [ | |
| ("S1", "x100", ROOT / "manifests/S1/manifest.json"), | |
| ("S2", "x50", ROOT / "manifests_v2/S2/manifest.json"), | |
| ("S3", "tier7", ROOT / "manifests/S2/manifest.json"), | |
| ("S4", "tier6", ROOT / "manifests/S3/manifest.json"), | |
| ("S5", "tier5", ROOT / "manifests/S4/manifest.json"), | |
| ("S6", "tier4", ROOT / "manifests/S5/manifest.json"), | |
| ("S7", "tier3", ROOT / "manifests/S5b/manifest.json"), | |
| ("S8", "tier2", ROOT / "manifests/S6/manifest.json"), | |
| ("S9", "tier1", ROOT / "manifests/S5c/manifest.json"), | |
| ("S10", "tier0", ROOT / "manifests/S7/manifest.json"), | |
| ] | |
| MASK64 = (1 << 64) - 1 | |
| def sha(path: Path) -> str: | |
| h = hashlib.sha256() | |
| with path.open("rb") as stream: | |
| for block in iter(lambda: stream.read(8 << 20), b""): | |
| h.update(block) | |
| return h.hexdigest() | |
| def splitmix(value: int) -> int: | |
| value = (value + 0x9E3779B97F4A7C15) & MASK64 | |
| value = ((value ^ (value >> 30)) * 0xBF58476D1CE4E5B9) & MASK64 | |
| value = ((value ^ (value >> 27)) * 0x94D049BB133111EB) & MASK64 | |
| return value ^ (value >> 31) | |
| def base_hash(stage: str, uid: str) -> int: | |
| return int.from_bytes(hashlib.blake2b(f"20260923|{stage}|{uid}".encode(), digest_size=8).digest(), "little") | |
| def chosen(value: int, seed: int, probability: float) -> bool: | |
| return splitmix(value ^ seed) / 2**64 < probability | |
| def lookup(keys: np.ndarray, values: np.ndarray, uids: list[str]): | |
| query = np.asarray([uid.encode() for uid in uids], dtype=keys.dtype) | |
| positions = np.searchsorted(keys, query) | |
| safe = np.minimum(positions, len(keys) - 1) | |
| found = (positions < len(keys)) & (keys[safe] == query) | |
| result = np.full(len(query), MISSING, dtype=np.uint32) | |
| result[found] = values[safe[found]] | |
| return found, result, safe | |
| def chunks(path: Path, size: int = 100_000): | |
| with path.open("rb") as stream: | |
| rows, physical = [], 0 | |
| for line in stream: | |
| row = json.loads(line) | |
| if row.get("form", "A") == "A": | |
| rows.append((physical, row)) | |
| if len(rows) >= size: | |
| yield rows; rows = [] | |
| physical += 1 | |
| if rows: | |
| yield rows | |
| def main(): | |
| stage_index = int(sys.argv[1]) | |
| stage, tier, base_manifest_path = STAGES[stage_index] | |
| dest = OUT / stage | |
| manifest_path = dest / "manifest.json" | |
| if manifest_path.exists() and json.loads(manifest_path.read_text()).get("status") == "complete": | |
| print(manifest_path); return | |
| dest.mkdir(parents=True, exist_ok=True) | |
| for stale in dest.glob("*.partial"): | |
| stale.unlink() | |
| assets = json.loads(ASSETS.read_text()); assert assets["status"] == "complete" | |
| asset_root = ASSETS.parent | |
| ann_keys = np.load(asset_root / assets["annotations"]["keys"], mmap_mode="r") | |
| ann_values = np.load(asset_root / assets["annotations"]["values"], mmap_mode="r") | |
| ann_variants = np.load(asset_root / assets["annotations"]["variants"], mmap_mode="r") | |
| ann_timbre = np.load(asset_root / assets["annotations"]["timbre_has_tags"], mmap_mode="r") | |
| ref_keys = np.load(asset_root / assets["references"]["keys"], mmap_mode="r") | |
| ref_values = np.load(asset_root / assets["references"]["values"], mmap_mode="r") | |
| variant_ids = assets["annotations"]["variant_ids"] | |
| ref_variant_codes = {variant_ids[x] for x in ("REF_PERF_SENT", "REF_PERF_GLOBAL", "REF_PERF_SPARSE")} | |
| base = json.loads(base_manifest_path.read_text()) | |
| base_records = base_manifest_path.parent / base["records"] | |
| counts = Counter() | |
| # Pass one establishes validated reference coverage. Main presentations | |
| # use a deterministic 50/50 draw among reference-eligible targets; targets | |
| # without a validated different-UID reference remain instruction-only. | |
| for batch in chunks(base_records): | |
| rows = [row for _, row in batch]; uids = [row["uid"] for row in rows] | |
| ann_found, _, ann_pos = lookup(ann_keys, ann_values, uids) | |
| ref_found, _, _ = lookup(ref_keys, ref_values, uids) | |
| old = np.asarray([bool(row.get("bude_caption")) for row in rows]) | |
| counts["unique"] += len(rows); counts["unique_ref"] += int(ref_found.sum()) | |
| counts["old"] += int(old.sum()); counts["old_ref"] += int((old & ref_found).sum()) | |
| counts["ann"] += int(ann_found.sum()); counts["ann_ref"] += int((ann_found & ref_found).sum()) | |
| if ann_found.any(): | |
| codes = ann_variants[ann_pos] | |
| gemma_ref = ann_found & np.isin(codes, list(ref_variant_codes)) | |
| gemma_nr = ann_found & ~gemma_ref | |
| counts["gemma_fixed_ref"] += int((gemma_ref & ref_found).sum()) | |
| counts["gemma_ref_missing"] += int((gemma_ref & ~ref_found).sum()) | |
| counts["gemma_nr_ref_eligible"] += int((gemma_nr & ref_found).sum()) | |
| if counts["unique"] % 2_000_000 < len(rows): | |
| print(stage, "audit", dict(counts), flush=True) | |
| assert counts["unique"] == base["presentations"] - base.get("b_form", 0) | |
| assert counts["gemma_ref_missing"] == 0, counts | |
| p_general = p_old = p_ann = 0.5 | |
| # REF_* Gemma instructions intrinsically require a reference. Balance the | |
| # eligible NR_* remainder so the expected eligible Gemma split is 50/50. | |
| gemma_extra = max(0.0, 0.5 * counts["ann_ref"] - counts["gemma_fixed_ref"]) | |
| p_gemma_nr = gemma_extra / max(1, counts["gemma_nr_ref_eligible"]) | |
| assert 0 <= p_gemma_nr <= 1 | |
| p_transcript = 0.10 * counts["unique"] / counts["unique_ref"] | |
| assert p_transcript <= 1 | |
| probabilities = {"general": p_general, "old": p_old, "annotation": p_ann, | |
| "gemma_nr": p_gemma_nr, "transcript": p_transcript} | |
| raw_path = dest / "descriptors.raw.partial" | |
| family_counts, family_refs = Counter(), Counter() | |
| hours = 0.0; emitted = 0 | |
| with raw_path.open("wb", buffering=16 << 20) as output: | |
| for batch in chunks(base_records): | |
| physical = [index for index, _ in batch] | |
| rows = [row for _, row in batch]; uids = [row["uid"] for row in rows] | |
| ann_found, ann_index, ann_pos = lookup(ann_keys, ann_values, uids) | |
| ref_found, ref_index, _ = lookup(ref_keys, ref_values, uids) | |
| records = [] | |
| for i, row in enumerate(rows): | |
| uid = uids[i]; h = base_hash(stage, uid); duration = float(row["dur_s"]) | |
| def emit(family, surface=0, use_reference=False, annotation=MISSING): | |
| nonlocal emitted, hours | |
| reference = int(ref_index[i]) if use_reference else int(MISSING) | |
| assert not use_reference or ref_found[i] | |
| records.append((physical[i], reference, int(annotation), family, surface, | |
| int(use_reference), 0)) | |
| emitted += 1; hours += duration / 3600 | |
| family_counts[family] += 1; family_refs[family] += int(use_reference) | |
| emit(0, use_reference=bool(ref_found[i] and chosen(h, 0x10, p_general))) | |
| if row.get("bude_caption"): | |
| emit(1, use_reference=bool(ref_found[i] and chosen(h, 0x11, p_old))) | |
| if ann_found[i]: | |
| ai = int(ann_index[i]) | |
| use_b = bool(ref_found[i] and chosen(h, 0x12, p_ann)) | |
| use_t = bool(ref_found[i] and chosen(h, 0x13, p_ann)) | |
| use_v = bool(ref_found[i] and chosen(h, 0x14, p_ann)) | |
| emit(2, use_reference=use_b, annotation=ai) | |
| t_surface = splitmix(h ^ 0x23) % 3 if ann_timbre[ann_pos[i]] else 1 | |
| emit(3, surface=int(t_surface), use_reference=use_t, annotation=ai) | |
| emit(4, surface=int(splitmix(h ^ 0x24) % 3), use_reference=use_v, annotation=ai) | |
| variant = int(ann_variants[ann_pos[i]]) | |
| if variant in ref_variant_codes: | |
| use_g = True | |
| else: | |
| use_g = bool(ref_found[i] and chosen(h, 0x15, p_gemma_nr)) | |
| emit(5, use_reference=use_g, annotation=ai) | |
| if ref_found[i] and chosen(h, 0x16, p_transcript): | |
| emit(6, use_reference=True) | |
| np.asarray(records, dtype=DESC_DTYPE).tofile(output) | |
| if emitted // 2_000_000 > (emitted - len(records)) // 2_000_000: | |
| print(stage, "emit", f"{emitted:,}", flush=True) | |
| raw = np.memmap(raw_path, mode="r", dtype=DESC_DTYPE, shape=(emitted,)) | |
| descriptors_name = "descriptors.npy" | |
| descriptors_tmp = dest / (descriptors_name + ".partial") | |
| descriptors = np.lib.format.open_memmap(descriptors_tmp, mode="w+", dtype=DESC_DTYPE, shape=(emitted,)) | |
| for start in range(0, emitted, 2_000_000): | |
| descriptors[start:start + 2_000_000] = raw[start:start + 2_000_000] | |
| descriptors.flush(); del descriptors, raw | |
| os.replace(descriptors_tmp, dest / descriptors_name); raw_path.unlink() | |
| order_name = "order.npy"; order_tmp = dest / (order_name + ".partial") | |
| order = np.lib.format.open_memmap(order_tmp, mode="w+", dtype=np.uint32, shape=(emitted,)) | |
| order[:] = np.arange(emitted, dtype=np.uint32) | |
| np.random.RandomState(20260923 + stage_index).shuffle(order) | |
| order.flush(); del order | |
| os.replace(order_tmp, dest / order_name) | |
| payload = { | |
| "status": "complete", "format": "m2-caption-curriculum-v1", "stage": stage, "tier": tier, | |
| "created_at": datetime.now().astimezone().isoformat(), "seed": 20260923, | |
| "base_manifest": str(base_manifest_path), "base_manifest_sha256": sha(base_manifest_path), | |
| "assets_manifest": str(ASSETS), "assets_manifest_sha256": sha(ASSETS), | |
| "descriptors": descriptors_name, "descriptors_sha256": sha(dest / descriptors_name), | |
| "order": order_name, "order_sha256": sha(dest / order_name), | |
| "presentations": emitted, "hours": round(hours, 3), | |
| "unique_audio_occurrences": counts["unique"], "annotated_occurrences": counts["ann"], | |
| "historical_bude_occurrences": counts["old"], "family_counts": dict(family_counts), | |
| "family_reference_counts": dict(family_refs), "reference_probabilities": probabilities, | |
| "family_reference_eligible_counts": { | |
| "0": counts["unique_ref"], "1": counts["old_ref"], | |
| "2": counts["ann_ref"], "3": counts["ann_ref"], | |
| "4": counts["ann_ref"], "5": counts["ann_ref"], | |
| "6": counts["unique_ref"], | |
| }, | |
| "reference_policy": "50% deterministic draw among validated reference-eligible targets; ineligible targets remain instruction-only; transcript-only extra is 10% of all unique targets and always referenced", | |
| } | |
| manifest_tmp = dest / "manifest.json.partial" | |
| manifest_tmp.write_text(json.dumps(payload, indent=2) + "\n") | |
| os.replace(manifest_tmp, manifest_path) | |
| print(json.dumps(payload, indent=2), flush=True) | |
| if __name__ == "__main__": | |
| main() | |