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