Download code/build_caption_assets.py from laion/Humaneness-Voice-Small: direct link, hf CLI and curl.
- Browser
- Download file 8.52 kB
-
https://huggingface.co/laion/Humaneness-Voice-Small/resolve/main/code/build_caption_assets.py
- Command line
-
hf download hf://laion/Humaneness-Voice-Small/code/build_caption_assets.py
-
curl -L -o build_caption_assets.py https://huggingface.co/laion/Humaneness-Voice-Small/resolve/main/code/build_caption_assets.py
8.52 kB
| #!/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() | |