suvradeepp's picture
Publish Tiny Hinglish Turn Detector development preview
35d483e verified
Raw
History Blame Contribute Delete
7.75 kB
"""Leakage-aware grouping using exact audio and metadata linkages."""
from __future__ import annotations
from dataclasses import dataclass
import hashlib
import json
import re
import unicodedata
from typing import Any, Iterable, Mapping, MutableMapping, Sequence
@dataclass(frozen=True)
class GroupingConfig:
"""Controls which stable identifiers link examples into components."""
include_record_id: bool = True
include_audio_path: bool = True
include_text: bool = True
minimum_text_characters: int = 12
minimum_text_tokens: int = 3
_FIELD_FAMILIES: dict[str, tuple[str, ...]] = {
"conversation": ("conversation_id", "conversation", "dialogue_id", "dialog_id", "session_id", "call_id"),
"speaker": ("speaker_id", "speaker", "user_id", "actor_id"),
"voice": ("voice_id", "voice", "tts_voice", "speaker_voice"),
"recording": ("recording_id", "source_recording_id", "audio_id", "clip_id"),
"prompt": ("prompt_id", "source_prompt_id", "template_id", "script_id", "parent_id"),
}
def _canonical_value(value: Any) -> str | None:
if value is None:
return None
if isinstance(value, str):
normalized = unicodedata.normalize("NFKC", value).strip().casefold()
return normalized or None
if isinstance(value, (bool, int, float)):
return json.dumps(value, sort_keys=True, allow_nan=False)
try:
return json.dumps(value, sort_keys=True, separators=(",", ":"), ensure_ascii=False, allow_nan=False)
except (TypeError, ValueError):
normalized = str(value).strip().casefold()
return normalized or None
def _private_key(family: str, namespace: str, value: Any) -> str | None:
canonical = _canonical_value(value)
if canonical is None:
return None
digest = hashlib.sha256(f"{namespace}\0{canonical}".encode("utf-8")).hexdigest()[:32]
return f"{family}:{digest}"
def _lookup(record: Mapping[str, Any], field: str) -> Any:
if field in record:
return record[field]
metadata = record.get("metadata")
if isinstance(metadata, Mapping) and field in metadata:
return metadata[field]
return None
def _normalize_text(value: str) -> str:
value = unicodedata.normalize("NFKC", value).casefold()
return " ".join(re.findall(r"[\w']+", value, flags=re.UNICODE))
def derive_group_keys(
record: Mapping[str, Any],
*,
audio_sha256: str | None = None,
config: GroupingConfig | None = None,
) -> list[str]:
"""Derive privacy-preserving exact-audio and metadata linkage keys.
Every shared key induces an edge. :func:`attach_group_ids` computes the
transitive connected components, so an audio duplicate linked to a shared
speaker cannot be split indirectly.
"""
config = config or GroupingConfig()
keys: set[str] = set()
audio_hash = audio_sha256 or _lookup(record, "audio_sha256")
if isinstance(audio_hash, str) and audio_hash:
keys.add(f"audio:{audio_hash.lower()}")
dataset = _canonical_value(
_lookup(record, "dataset") or _lookup(record, "source_dataset") or _lookup(record, "data_source")
)
namespace = dataset or "unknown-dataset"
for family, fields in _FIELD_FAMILIES.items():
for field in fields:
value = _lookup(record, field)
key = _private_key(family, namespace, value)
if key:
keys.add(key)
if config.include_record_id:
explicit_id = _lookup(record, "id")
if explicit_id is None:
# Audit manifests canonicalize the upstream id as record_id.
explicit_id = _lookup(record, "record_id")
key = _private_key("record", namespace, explicit_id)
if key:
keys.add(key)
if config.include_audio_path and not audio_hash:
audio_path = _lookup(record, "audio_path")
if audio_path is None:
audio = _lookup(record, "audio")
if isinstance(audio, Mapping):
audio_path = audio.get("path")
if audio_path:
# Basename is portable across machines and catches copied manifests.
basename = str(audio_path).replace("\\", "/").rsplit("/", 1)[-1]
key = _private_key("audio_path", namespace, basename)
if key:
keys.add(key)
if config.include_text:
text = _lookup(record, "spoken_text") or _lookup(record, "transcript") or _lookup(record, "text")
if isinstance(text, str):
normalized_text = _normalize_text(text)
if (
len(normalized_text) >= config.minimum_text_characters
and len(normalized_text.split()) >= config.minimum_text_tokens
):
key = _private_key("text", "global", normalized_text)
if key:
keys.add(key)
return sorted(keys)
class _UnionFind:
def __init__(self, size: int) -> None:
self.parent = list(range(size))
self.rank = [0] * size
def find(self, item: int) -> int:
while self.parent[item] != item:
self.parent[item] = self.parent[self.parent[item]]
item = self.parent[item]
return item
def union(self, left: int, right: int) -> None:
left_root = self.find(left)
right_root = self.find(right)
if left_root == right_root:
return
if self.rank[left_root] < self.rank[right_root]:
left_root, right_root = right_root, left_root
self.parent[right_root] = left_root
if self.rank[left_root] == self.rank[right_root]:
self.rank[left_root] += 1
def attach_group_ids(
rows: Sequence[MutableMapping[str, Any]],
*,
group_keys_field: str = "group_keys",
) -> Sequence[MutableMapping[str, Any]]:
"""Attach deterministic component ``group_id`` values to manifest rows.
Rows are modified in place to avoid duplicating a full 270k-row manifest.
The resulting IDs do not depend on row order.
"""
union_find = _UnionFind(len(rows))
first_row_for_key: dict[str, int] = {}
normalized_keys: list[list[str]] = []
for index, row in enumerate(rows):
raw_keys = row.get(group_keys_field) or []
if isinstance(raw_keys, str):
raw_keys = [raw_keys]
keys = sorted({str(key) for key in raw_keys if key})
if not keys:
fallback = hashlib.sha256(
f"{row.get('record_id', '')}\0{row.get('source_file', '')}\0{row.get('source_row', index)}".encode(
"utf-8"
)
).hexdigest()
keys = [f"row:{fallback}"]
row[group_keys_field] = keys
normalized_keys.append(keys)
for key in keys:
previous = first_row_for_key.setdefault(key, index)
union_find.union(index, previous)
component_keys: dict[int, set[str]] = {}
for index, keys in enumerate(normalized_keys):
root = union_find.find(index)
component_keys.setdefault(root, set()).update(keys)
component_ids = {
root: "grp_" + hashlib.sha256("\n".join(sorted(keys)).encode("utf-8")).hexdigest()[:24]
for root, keys in component_keys.items()
}
for index, row in enumerate(rows):
row[group_keys_field] = normalized_keys[index]
row["group_id"] = component_ids[union_find.find(index)]
return rows
def group_sizes(rows: Iterable[Mapping[str, Any]]) -> dict[str, int]:
"""Return the number of rows in every leakage component."""
counts: dict[str, int] = {}
for row in rows:
group_id = str(row.get("group_id") or "")
if group_id:
counts[group_id] = counts.get(group_id, 0) + 1
return counts