Yash-V1002's picture
Deploy Tiny Turn Detector
875e4af verified
Raw
History Blame Contribute Delete
3.06 kB
"""
Split strategy: since there are no speaker/session IDs (confirmed —
docs/INITIAL_ANALYSIS.md §4), this module implements two comparable split
strategies over the DEVELOPMENT subset only (never the official held-out
test set, which is used as-is for final reporting):
- split_random: plain random split, ignoring `dataset` (source) grouping.
- split_source_aware: group-disjoint split by the `dataset` column, so no
single source/voice-engine appears in both train and validation.
The Phase 2 brief explicitly asks us to compare these two, to quantify
whether source-specific acoustic characteristics make random evaluation
overly optimistic. This module produces the splits; experiments/EXPERIMENTS.md
records the actual comparison once real data is available.
IMPORTANT: this is a source-DISJOINT split, not a speaker-disjoint split.
The two are not equivalent, and this module's naming/docstrings are
deliberately explicit about that so no downstream document accidentally
implies otherwise.
"""
from __future__ import annotations
import numpy as np
def split_random(
records: list[dict],
val_frac: float = 0.2,
seed: int = 42,
) -> tuple[list[dict], list[dict]]:
rng = np.random.default_rng(seed)
idx = np.arange(len(records))
rng.shuffle(idx)
n_val = int(round(len(records) * val_frac))
val_idx = set(idx[:n_val].tolist())
train, val = [], []
for i, r in enumerate(records):
(val if i in val_idx else train).append(r)
return train, val
def split_source_aware(
records: list[dict],
val_frac: float = 0.2,
seed: int = 42,
source_key: str = "dataset",
) -> tuple[list[dict], list[dict]]:
"""Group-disjoint split by `source_key` (the `dataset` / voice-engine
column). Whole sources are assigned entirely to train or entirely to
val, chosen to get val_frac of the *records* (not sources) into val,
via a greedy bin-packing over source sizes (deterministic given seed
for tie-breaking order).
"""
rng = np.random.default_rng(seed)
by_source: dict[str, list[dict]] = {}
for r in records:
by_source.setdefault(str(r.get(source_key)), []).append(r)
sources = list(by_source.keys())
rng.shuffle(sources) # randomize assignment order for tie-breaking
target_val_n = int(round(len(records) * val_frac))
val, train = [], []
val_n = 0
for s in sources:
group = by_source[s]
if val_n < target_val_n:
val.extend(group)
val_n += len(group)
else:
train.extend(group)
return train, val
def source_overlap(train: list[dict], val: list[dict], source_key: str = "dataset") -> set:
"""Returns the set of source values present in BOTH train and val.
Should be empty for split_source_aware; will typically be non-empty for
split_random (that's the point of the comparison).
"""
train_sources = {str(r.get(source_key)) for r in train}
val_sources = {str(r.get(source_key)) for r in val}
return train_sources & val_sources