#!/usr/bin/env python3 """Build two bounded random-access containers for captions and references.""" from __future__ import annotations from array import array from datetime import datetime import hashlib import json import os from pathlib import Path import re import numpy as np import pyarrow.parquet as pq SC = Path("/e/scratch/reformo/schuhmann1_moss") SOURCE = SC / "out/m2_s3_caption_enrichment/production_v1" REFS = SC / "out/refassign/reference_assignments_x100.parquet" OUT = SC / "out/m2_600m_ladder/caption_curriculum_v1/assets" PARTS = 256 KEY_DTYPE = "S128" # audited maxima: annotations 102 bytes, references 78 bytes KEY_BYTES = np.dtype(KEY_DTYPE).itemsize 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 gemma_caption(instruction: str) -> str: # Quote-safe production validation proves every straight-double-quoted span # is spoken text. CAPTION/TRANSCRIPT training must carry that text once. body = re.sub(r'"[^"\n]*"', " ", instruction) body = " ".join(body.replace('"', "'").split()).strip(" ;") assert body return body def voice_prose(tags: list[str]) -> str: words = [" ".join(str(tag).replace("_", " ").replace("-", " ").split()) for tag in tags] words = [word for word in words if word] assert words if len(words) == 1: joined = words[0] else: joined = ", ".join(words[:-1]) + ", and " + words[-1] return "Use a voice characterized by " + joined + "." def write_npy(path: Path, values, dtype=None): temporary = Path(str(path) + ".partial") with temporary.open("wb") as stream: np.save(stream, np.asarray(values, dtype=dtype), allow_pickle=False) os.replace(temporary, path) def sorted_lookup(raw_keys: Path, count: int, prefix: str, variants=None, tag_flags=None): keys = np.memmap(raw_keys, mode="r", dtype=KEY_DTYPE, shape=(count,)) order = np.argsort(keys, kind="stable") sorted_keys = np.asarray(keys[order]) assert count == 0 or not np.any(sorted_keys[1:] == sorted_keys[:-1]) write_npy(OUT / f"{prefix}.keys.npy", sorted_keys, dtype=KEY_DTYPE) write_npy(OUT / f"{prefix}.values.npy", order, dtype=np.uint32) if variants is not None: write_npy(OUT / f"{prefix}.variants.npy", np.asarray(variants, dtype=np.uint8)[order]) if tag_flags is not None: write_npy(OUT / f"{prefix}.timbre_has_tags.npy", np.asarray(tag_flags, dtype=np.uint8)[order]) del sorted_keys, order, keys raw_keys.unlink() def build_annotations() -> dict: records_tmp = OUT / "annotations.records.jsonl.partial" raw_keys = OUT / "annotations.keys.raw.partial" offsets = array("q", [0]) variants, tag_flags = array("B"), array("B") variant_ids = {name: i for i, name in enumerate(( "NR_TIMBRE_SENT", "NR_ARCHETYPE_SENT", "NR_TIMBRE_GLOBAL", "REF_PERF_SENT", "REF_PERF_GLOBAL", "REF_PERF_SPARSE"))} counts = {name: 0 for name in variant_ids} with records_tmp.open("wb", buffering=8 << 20) as records, raw_keys.open("wb", buffering=8 << 20) as keys: for part in range(PARTS): wp = SOURCE / "whisper" / f"whisper.part-{part:05d}.parquet" gp = SOURCE / "gemma" / f"gemma.part-{part:05d}.parquet" w = pq.read_table(wp, columns=["uid", "bude_caption", "timbre_tags_json", "timbre_prose", "voice_tags_json"]).to_pylist() g = pq.read_table(gp, columns=["uid", "variant", "instruction", "pass"]).to_pylist() assert len(w) == len(g) for wr, gr in zip(w, g): assert wr["uid"] == gr["uid"] and gr["pass"] timbre_tags = json.loads(wr["timbre_tags_json"] or "[]") voice_tags = json.loads(wr["voice_tags_json"] or "[]") assert wr["bude_caption"] and wr["timbre_prose"] and voice_tags row = { "uid": wr["uid"], "bude": wr["bude_caption"], "timbre_tags": timbre_tags, "timbre_prose": wr["timbre_prose"], "voice_tags": voice_tags, "voice_prose": voice_prose(voice_tags), "gemma_variant": gr["variant"], "gemma_caption": gemma_caption(gr["instruction"]), } blob = json.dumps(row, ensure_ascii=False, separators=(",", ":")).encode() + b"\n" records.write(blob); offsets.append(offsets[-1] + len(blob)) key = wr["uid"].encode(); assert len(key) <= KEY_BYTES keys.write(key.ljust(KEY_BYTES, b"\0")) variants.append(variant_ids[gr["variant"]]); tag_flags.append(bool(timbre_tags)) counts[gr["variant"]] += 1 print(f"annotations part {part:03d}: {len(offsets)-1:,}", flush=True) os.replace(records_tmp, OUT / "annotations.records.jsonl") write_npy(OUT / "annotations.offsets.npy", np.frombuffer(offsets, dtype=np.int64)) count = len(offsets) - 1 sorted_lookup(raw_keys, count, "annotations", variants, tag_flags) return {"rows": count, "records": "annotations.records.jsonl", "offsets": "annotations.offsets.npy", "keys": "annotations.keys.npy", "values": "annotations.values.npy", "variants": "annotations.variants.npy", "timbre_has_tags": "annotations.timbre_has_tags.npy", "variant_ids": variant_ids, "variant_counts": counts} def build_references() -> dict: records_tmp = OUT / "references.records.jsonl.partial" raw_keys = OUT / "references.keys.raw.partial" offsets = array("q", [0]) path_exists = {} skipped_path = 0 with records_tmp.open("wb", buffering=8 << 20) as records, raw_keys.open("wb", buffering=8 << 20) as keys: pf = pq.ParquetFile(REFS) for batch_i, batch in enumerate(pf.iter_batches(batch_size=250_000, columns=["uid", "reference_uid", "reference_locator_json", "reference_eligible"])): for source in batch.to_pylist(): if not source["reference_eligible"] or not source["reference_locator_json"]: continue locator = json.loads(source["reference_locator_json"]) path = locator["path"] exists = path_exists.get(path) if exists is None: exists = Path(path).exists(); path_exists[path] = exists if not exists: skipped_path += 1 continue row = {"uid": source["uid"], "reference_uid": source["reference_uid"], "locator": locator} blob = json.dumps(row, separators=(",", ":")).encode() + b"\n" records.write(blob); offsets.append(offsets[-1] + len(blob)) key = source["uid"].encode(); assert len(key) <= KEY_BYTES keys.write(key.ljust(KEY_BYTES, b"\0")) print(f"references batch {batch_i:03d}: {len(offsets)-1:,}", flush=True) os.replace(records_tmp, OUT / "references.records.jsonl") write_npy(OUT / "references.offsets.npy", np.frombuffer(offsets, dtype=np.int64)) count = len(offsets) - 1 sorted_lookup(raw_keys, count, "references") return {"rows": count, "records": "references.records.jsonl", "offsets": "references.offsets.npy", "keys": "references.keys.npy", "values": "references.values.npy", "skipped_missing_path": skipped_path} def main(): OUT.mkdir(parents=True, exist_ok=True) manifest = OUT / "manifest.json" if manifest.exists() and json.loads(manifest.read_text()).get("status") == "complete": print(manifest); return for stale in OUT.glob("*.partial"): stale.unlink() annotations = build_annotations() assert annotations["rows"] == 3_682_649 references = build_references() payload = {"status": "complete", "format": "m2-caption-assets-v1", "created_at": datetime.now().astimezone().isoformat(), "annotations": annotations, "references": references} for group in (annotations, references): for key in ("records", "offsets", "keys", "values"): group[key + "_sha256"] = sha(OUT / group[key]) temporary = OUT / "manifest.json.partial" temporary.write_text(json.dumps(payload, indent=2) + "\n") os.replace(temporary, manifest) print(manifest) if __name__ == "__main__": main()