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