tiny-turn-detector / scripts /create_dev_subset.py
Yash-V1002's picture
Deploy Tiny Turn Detector
875e4af verified
Raw
History Blame Contribute Delete
6.13 kB
"""
Create a reproducible, stratified development subset of
pipecat-ai/smart-turn-data-v3.2-train, WITHOUT downloading/materializing
the full 41GB dataset.
Status: this script is complete and ready to run, but has NOT been executed
in this sandbox — `datasets`/`huggingface_hub` aren't installed here and
`bash` has no route to huggingface.co (confirmed 403 host_not_allowed via
curl; see docs/INITIAL_ANALYSIS.md and experiments/EXPERIMENTS.md). Run this
in an environment with network access and `pip install datasets`.
Usage:
python scripts/create_dev_subset.py \
--target-n 10000 \
--seed 42 \
--max-scan 200000 \
--out data/processed/dev_subset_manifest.jsonl
Methodology (also written into the output manifest's header for
reproducibility):
1. Stream the dataset (`streaming=True`) so we never download the full
41GB just to build a 10k-example subset.
2. Compute per-clip `duration_sec` and `duration_bucket` from the audio
array's length/sample rate as each streamed record is visited (cheap —
doesn't require decoding beyond what `datasets` already does to expose
the array).
3. Apply stratified reservoir sampling (src/turn_detector/data.py) across
endpoint_bool x language x dataset(source) x synthetic x midfiller x
endfiller x duration_bucket, using a two-pass approach over the first
`--max-scan` streamed records (pass 1: count stratum sizes; pass 2:
per-stratum reservoir sampling). This means `--max-scan` records are
scanned twice; the dataset itself is never materialized beyond that.
4. The chosen `target_n` records' `id` values (not the audio itself) are
written to the manifest, plus the exact sampling parameters, so the
subset is exactly reproducible from the manifest + original dataset.
5. The test set (pipecat-ai/smart-turn-data-v3.2-test) is never touched
by this script — it exists only to build a TRAIN-side dev subset.
"""
from __future__ import annotations
import argparse
import json
import sys
from pathlib import Path
sys.path.insert(0, str(Path(__file__).resolve().parents[1] / "src"))
from turn_detector.data import (
TRAIN_DATASET_ID, load_hf_dataset, stratified_reservoir_sample,
duration_bucket, DevSubsetManifest, DatasetAccessError,
)
def record_to_stratum_record(record: dict) -> dict:
"""Extract only the lightweight metadata fields needed for
stratification + the id needed to re-fetch the audio later — never
keeps the decoded audio array in memory longer than needed to compute
duration.
"""
audio = record.get("audio")
duration_sec = None
if isinstance(audio, dict) and "array" in audio and "sampling_rate" in audio:
duration_sec = len(audio["array"]) / audio["sampling_rate"]
return {
"id": record.get("id"),
"endpoint_bool": record.get("endpoint_bool"),
"language": record.get("language"),
"dataset": record.get("dataset"),
"synthetic": record.get("synthetic"),
"midfiller": record.get("midfiller"),
"endfiller": record.get("endfiller"),
"duration_sec": duration_sec,
}
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--target-n", type=int, default=10_000)
ap.add_argument("--seed", type=int, default=42)
ap.add_argument(
"--max-scan", type=int, default=None,
help="Cap on records scanned per pass for stratum-size counting "
"and reservoir filling. Default: scan the whole stream (safest for "
"correctness, slowest for a 270k-row dataset). Set e.g. 200000 to "
"bound runtime at some accuracy cost in stratum proportion "
"accuracy for very rare strata.",
)
ap.add_argument("--out", type=Path, default=Path("data/processed/dev_subset_manifest.jsonl"))
args = ap.parse_args()
try:
ds = load_hf_dataset(TRAIN_DATASET_ID, split="train", streaming=True)
except DatasetAccessError as e:
print(f"ERROR: could not load dataset: {e}", file=sys.stderr)
print(
"This script cannot run without network access to Hugging Face. "
"See docs/INITIAL_ANALYSIS.md for the current sandbox limitation.",
file=sys.stderr,
)
sys.exit(1)
# Two lightweight passes over the (streamed) dataset. Streaming
# datasets support repeated iteration by re-creating the iterator; we
# rely on `stratified_reservoir_sample`'s two-`for record in records`
# loops, so we materialize just the lightweight metadata (not audio)
# into a list first — this is a deliberate, documented tradeoff: it
# means `--max-scan` records of *metadata only* are held in memory
# (small), while full audio arrays are never retained beyond a single
# record's processing.
print(f"Scanning up to {args.max_scan or 'ALL'} records for stratification metadata...")
stratum_records = []
for i, record in enumerate(ds):
stratum_records.append(record_to_stratum_record(record))
if args.max_scan is not None and i + 1 >= args.max_scan:
break
if (i + 1) % 10_000 == 0:
print(f" scanned {i + 1} records...")
print(f"Scanned {len(stratum_records)} records. Running stratified sampling...")
sample = stratified_reservoir_sample(
stratum_records, target_n=args.target_n, seed=args.seed, max_scan=None,
)
manifest = DevSubsetManifest(
source_dataset_id=TRAIN_DATASET_ID,
source_split="train",
target_n=args.target_n,
actual_n=len(sample),
seed=args.seed,
stratify_columns=("endpoint_bool", "language", "dataset", "synthetic", "midfiller", "endfiller", "duration_bucket"),
max_scan=args.max_scan,
)
args.out.parent.mkdir(parents=True, exist_ok=True)
with open(args.out, "w") as f:
f.write(json.dumps({"_manifest": manifest.to_dict()}) + "\n")
for r in sample:
f.write(json.dumps(r) + "\n")
print(f"Wrote {len(sample)} records + manifest header to {args.out}")
if __name__ == "__main__":
main()