Text-to-Speech
English
German
voice-acting
qwen3
moss-audio-tokenizer-v2
audio-generation
Humaneness-Voice-Small / code /build_caption_stage.py
ChristophSchuhmann's picture
Document architecture, prompts, code, and full run statistics
5de8603 verified
Raw History Blame Contribute Delete
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()