from __future__ import annotations import json import logging import re import time from itertools import permutations from pathlib import Path from typing import Any from jiwer import cer, wer logger = logging.getLogger(__name__) _NON_SPEAKER_LABELS = { "начало диалога", "конец диалога", "start", "end", "начало", "конец", "noise", "шум", "музыка", "music", } # --------------------------------------------------------------------------- # Label Studio JSON → gold segments # --------------------------------------------------------------------------- def parse_label_studio_json(data: list[dict]) -> dict[str, list[dict]]: """Convert Label Studio JSON export to gold segments grouped by file_id. Handles the standard LS audio annotation format where each task has regions with 'labels' (speaker) and 'textarea' (transcription). Returns: {file_id: [{start, end, speaker_role, text}, ...]} """ result: dict[str, list[dict]] = {} for task in data: file_id = _extract_file_id(task) if not file_id: continue annotations = task.get("annotations", []) if not annotations: continue best_annotation = annotations[-1] regions = _parse_annotation_result(best_annotation.get("result", [])) segments = sorted(regions, key=lambda s: s["start"]) if segments: result[file_id] = segments return result def _extract_file_id(task: dict) -> str | None: """Extract file identifier from a Label Studio task.""" data = task.get("data", {}) for key in ("audio", "audio_url", "file", "url"): val = data.get(key, "") if val: name = Path(val).stem name = re.sub(r"^[a-f0-9]{8}-", "", name) return name task_id = task.get("id") if task_id is not None: return str(task_id) return None def _parse_annotation_result(result: list[dict]) -> list[dict]: """Parse Label Studio annotation result into segments. Matches 'labels' and 'textarea' regions by their region id. """ label_regions: dict[str, dict] = {} text_regions: dict[str, str] = {} for item in result: item_id = item.get("id", "") item_type = item.get("type", "") value = item.get("value", {}) if item_type == "labels": labels = value.get("labels", []) label_regions[item_id] = { "start": value.get("start", 0.0), "end": value.get("end", 0.0), "speaker_role": labels[0] if labels else "unknown", } elif item_type == "textarea": texts = value.get("text", []) text_regions[item_id] = texts[0] if texts else "" segments: list[dict] = [] for region_id, region in label_regions.items(): text = text_regions.get(region_id, "") segments.append({ "start": round(region["start"], 2), "end": round(region["end"], 2), "speaker_role": region["speaker_role"], "text": text, }) if not segments and result: segments = _parse_flat_regions(result) return segments def _parse_flat_regions(result: list[dict]) -> list[dict]: """Fallback: parse regions where labels and text share the same value block.""" segments: list[dict] = [] for item in result: value = item.get("value", {}) start = value.get("start") end = value.get("end") if start is None or end is None: continue labels = value.get("labels", []) texts = value.get("text", []) segments.append({ "start": round(start, 2), "end": round(end, 2), "speaker_role": labels[0] if labels else "unknown", "text": texts[0] if texts else "", }) return segments def gold_segments_to_jsonl(gold_by_file: dict[str, list[dict]]) -> str: """Serialize gold segments to JSONL string.""" lines: list[str] = [] for file_id, segments in sorted(gold_by_file.items()): for seg in segments: row = {"file_id": file_id, **seg} lines.append(json.dumps(row, ensure_ascii=False)) return "\n".join(lines) # --------------------------------------------------------------------------- # Text metrics # --------------------------------------------------------------------------- def compute_wer(reference: str, hypothesis: str) -> float: if not reference.strip(): return 0.0 if not hypothesis.strip() else 1.0 return wer(reference, hypothesis) def compute_cer(reference: str, hypothesis: str) -> float: if not reference.strip(): return 0.0 if not hypothesis.strip() else 1.0 return cer(reference, hypothesis) # --------------------------------------------------------------------------- # DER (Diarization Error Rate) — manual implementation # --------------------------------------------------------------------------- def _discretize_segments( segments: list[dict], speaker_key: str, resolution: float = 0.025 ) -> tuple[dict[int, str], float, float]: """Convert segments to integer-bin→speaker mapping for grid-aligned comparison.""" if not segments: return {}, 0.0, 0.0 min_t = min(s["start"] for s in segments) max_t = max(s["end"] for s in segments) mapping: dict[int, str] = {} for seg in segments: start_bin = int(seg["start"] / resolution) end_bin = int(seg["end"] / resolution) speaker = seg.get(speaker_key, "unknown") for b in range(start_bin, end_bin): mapping[b] = speaker return mapping, min_t, max_t def find_optimal_speaker_mapping( pred_segments: list[dict], gold_segments: list[dict], pred_key: str = "speaker", gold_key: str = "speaker_role", ) -> dict[str, str]: """Find optimal mapping from predicted speaker labels to gold speaker roles. Uses the mapping that maximizes temporal overlap agreement. """ pred_speakers = sorted(set(s.get(pred_key, "") for s in pred_segments if s.get(pred_key))) gold_speakers = sorted(set( s.get(gold_key, "") for s in gold_segments if s.get(gold_key) and s.get(gold_key, "").lower() not in _NON_SPEAKER_LABELS )) if not pred_speakers or not gold_speakers: return {} gold_real = [ s for s in gold_segments if s.get(gold_key, "").lower() not in _NON_SPEAKER_LABELS ] resolution = 0.025 pred_map, _, _ = _discretize_segments(pred_segments, pred_key, resolution) gold_map, _, _ = _discretize_segments(gold_real, gold_key, resolution) common_bins = set(pred_map.keys()) & set(gold_map.keys()) if not common_bins: return {p: g for p, g in zip(pred_speakers, gold_speakers)} best_mapping: dict[str, str] = {} best_score = -1 padded_gold = list(gold_speakers) + ["__unmapped__"] * max(0, len(pred_speakers) - len(gold_speakers)) for perm in permutations(padded_gold, len(pred_speakers)): candidate = dict(zip(pred_speakers, perm)) score = sum( 1 for b in common_bins if candidate.get(pred_map[b]) == gold_map[b] ) if score > best_score: best_score = score best_mapping = candidate return {k: v for k, v in best_mapping.items() if v != "__unmapped__"} def compute_der( pred_segments: list[dict], gold_segments: list[dict], speaker_mapping: dict[str, str] | None = None, collar: float = 0.5, pred_key: str = "speaker", gold_key: str = "speaker_role", ) -> dict[str, float]: """Compute Diarization Error Rate. Returns dict with: der, missed, false_alarm, confusion, total_ref_duration. """ gold_for_der = [ s for s in gold_segments if s.get(gold_key, "").lower() not in _NON_SPEAKER_LABELS ] if not gold_for_der: return {"der": 0.0, "missed": 0.0, "false_alarm": 0.0, "confusion": 0.0, "total_ref_duration": 0.0} if speaker_mapping is None: speaker_mapping = find_optimal_speaker_mapping( pred_segments, gold_for_der, pred_key, gold_key ) resolution = 0.025 gold_map, min_t, max_t = _discretize_segments(gold_for_der, gold_key, resolution) pred_map, _, _ = _discretize_segments(pred_segments, pred_key, resolution) collar_bins = int(collar / resolution) collar_set: set[int] = set() if collar_bins > 0: for seg in gold_segments: seg_start_bin = int(seg["start"] / resolution) seg_end_bin = int(seg["end"] / resolution) for b in range(seg_start_bin, min(seg_start_bin + collar_bins, seg_end_bin)): collar_set.add(b) for b in range(max(seg_end_bin - collar_bins, seg_start_bin), seg_end_bin): collar_set.add(b) all_bins = set(gold_map.keys()) | set(pred_map.keys()) missed = 0.0 false_alarm = 0.0 confusion = 0.0 total_ref = 0.0 for b in all_bins: if b in collar_set: continue in_ref = b in gold_map in_hyp = b in pred_map if in_ref: total_ref += resolution if in_ref and not in_hyp: missed += resolution elif not in_ref and in_hyp: false_alarm += resolution elif in_ref and in_hyp: gold_spk = gold_map[b] pred_spk = pred_map[b] mapped_pred = speaker_mapping.get(pred_spk, pred_spk) if mapped_pred != gold_spk: confusion += resolution der = (missed + false_alarm + confusion) / max(total_ref, 0.001) return { "der": round(der, 4), "missed": round(missed, 2), "false_alarm": round(false_alarm, 2), "confusion": round(confusion, 2), "total_ref_duration": round(total_ref, 2), } # --------------------------------------------------------------------------- # Segment matching # --------------------------------------------------------------------------- def _time_overlap(a_start: float, a_end: float, b_start: float, b_end: float) -> float: return max(0.0, min(a_end, b_end) - max(a_start, b_start)) def match_segments( pred_segments: list[dict], gold_segments: list[dict], speaker_mapping: dict[str, str] | None = None, min_overlap: float = 0.3, ) -> list[dict[str, Any]]: """Match predicted segments to gold segments by time overlap. Returns a list of dicts with keys: pred, gold, overlap, speaker_match, wer, cer. Unmatched gold segments are also included with pred=None. """ if speaker_mapping is None: speaker_mapping = find_optimal_speaker_mapping(pred_segments, gold_segments) matches: list[dict[str, Any]] = [] used_gold: set[int] = set() for p_seg in pred_segments: best_idx = -1 best_overlap = min_overlap for g_idx, g_seg in enumerate(gold_segments): if g_idx in used_gold: continue ov = _time_overlap(p_seg["start"], p_seg["end"], g_seg["start"], g_seg["end"]) if ov > best_overlap: best_overlap = ov best_idx = g_idx if best_idx >= 0: used_gold.add(best_idx) g_seg = gold_segments[best_idx] p_text = p_seg.get("text", "") g_text = g_seg.get("text", "") pred_speaker = p_seg.get("speaker", "") mapped_speaker = speaker_mapping.get(pred_speaker, pred_speaker) gold_role = g_seg.get("speaker_role", "") matches.append({ "pred": p_seg, "gold": g_seg, "overlap": round(best_overlap, 2), "speaker_match": mapped_speaker == gold_role, "mapped_speaker": mapped_speaker, "wer": compute_wer(g_text, p_text), "cer": compute_cer(g_text, p_text), }) else: matches.append({"pred": p_seg, "gold": None, "overlap": 0.0}) for g_idx, g_seg in enumerate(gold_segments): if g_idx not in used_gold: matches.append({"pred": None, "gold": g_seg, "overlap": 0.0}) return matches def summarise_matches( matches: list[dict], der_result: dict[str, float] | None = None, ) -> dict: matched = [m for m in matches if m.get("gold") is not None and m.get("pred") is not None] only_pred = [m for m in matches if m.get("gold") is None] only_gold = [m for m in matches if m.get("pred") is None] total_gold = len(only_gold) + len(matched) if matched: matched_wer = sum(m["wer"] for m in matched) / len(matched) matched_cer = sum(m["cer"] for m in matched) / len(matched) speaker_correct = sum(1 for m in matched if m.get("speaker_match")) / len(matched) else: matched_wer = 1.0 matched_cer = 1.0 speaker_correct = 0.0 if total_gold > 0: recall = len(matched) / total_gold avg_wer = matched_wer * recall + 1.0 * (1.0 - recall) avg_cer = matched_cer * recall + 1.0 * (1.0 - recall) else: avg_wer = matched_wer avg_cer = matched_cer summary = { "total_predicted_segments": len(only_pred) + len(matched), "total_gold_segments": total_gold, "matched_segments": len(matched), "unmatched_pred": len(only_pred), "missed_gold": len(only_gold), "avg_wer": round(avg_wer, 4), "avg_cer": round(avg_cer, 4), "matched_wer": round(matched_wer, 4), "matched_cer": round(matched_cer, 4), "recall": round(recall, 4) if total_gold > 0 else 0.0, "speaker_accuracy": round(speaker_correct, 4), } if der_result: summary.update({ "der": der_result["der"], "der_missed": der_result["missed"], "der_false_alarm": der_result["false_alarm"], "der_confusion": der_result["confusion"], }) return summary # --------------------------------------------------------------------------- # Benchmark: run multiple model/strategy combos and evaluate # --------------------------------------------------------------------------- def evaluate_single_result( result: dict, gold_segments: list[dict], ) -> dict[str, Any]: """Evaluate a single pipeline result against gold annotations. Returns metrics dict with DER, WER, CER, speaker accuracy. """ pred_segments = result.get("transcription", []) speaker_mapping = find_optimal_speaker_mapping(pred_segments, gold_segments) der_result = compute_der(pred_segments, gold_segments, speaker_mapping) matches = match_segments(pred_segments, gold_segments, speaker_mapping) summary = summarise_matches(matches, der_result) summary["speaker_mapping"] = {k: v for k, v in speaker_mapping.items()} return summary def run_benchmark( file_ids: list[str], combos: list[dict], gold_by_file: dict[str, list[dict]], download_audio_fn, process_file_fn, token: str | None = None, progress_fn=None, ) -> list[dict]: """Run benchmark for each (file x combo) and return ranked results. Args: file_ids: list of file IDs to process combos: list of dicts with keys: diar_model, trans_model, strategy, label (optional) gold_by_file: {file_id: [gold segments]} download_audio_fn: callable(file_id) -> audio_path process_file_fn: callable(...) -> result dict (GPU-decorated) token: HF token progress_fn: optional callable(fraction, desc) for progress reporting Returns: List of benchmark row dicts sorted by avg WER. """ total_steps = len(file_ids) * len(combos) step = 0 rows: list[dict] = [] for combo in combos: diar_model = combo["diar_model"] trans_model = combo["trans_model"] strategy = combo.get("strategy", "hybrid") label = combo.get("label", f"{trans_model}/{diar_model}/{strategy}") combo_metrics: list[dict] = [] combo_time = 0.0 for file_id in file_ids: step += 1 if progress_fn: progress_fn(step / total_steps, desc=f"{label}: {file_id} ({step}/{total_steps})") gold = gold_by_file.get(file_id, []) if not gold: logger.warning("No gold annotations for %s, skipping", file_id) continue try: audio_path = download_audio_fn(file_id) t0 = time.time() result = process_file_fn( audio_path, diar_model, trans_model, None, None, min_speakers=2, max_speakers=None, strategy=strategy, ) elapsed = time.time() - t0 except Exception as exc: logger.exception("Benchmark error for %s / %s", file_id, label) combo_metrics.append({ "file_id": file_id, "error": str(exc), "der": 1.0, "avg_wer": 1.0, "avg_cer": 1.0, "speaker_accuracy": 0.0, "time_sec": 0.0, }) continue if "__error__" in result: combo_metrics.append({ "file_id": file_id, "error": result["__error__"][:200], "der": 1.0, "avg_wer": 1.0, "avg_cer": 1.0, "speaker_accuracy": 0.0, "time_sec": 0.0, }) continue metrics = evaluate_single_result(result, gold) metrics["file_id"] = file_id metrics["time_sec"] = round(elapsed, 1) combo_metrics.append(metrics) combo_time += elapsed n = max(len(combo_metrics), 1) avg_der = sum(m.get("der", 1.0) for m in combo_metrics) / n avg_wer_val = sum(m.get("avg_wer", 1.0) for m in combo_metrics) / n avg_cer_val = sum(m.get("avg_cer", 1.0) for m in combo_metrics) / n avg_spk_acc = sum(m.get("speaker_accuracy", 0.0) for m in combo_metrics) / n rows.append({ "label": label, "trans_model": trans_model, "diar_model": diar_model, "strategy": strategy, "avg_der": round(avg_der, 4), "avg_wer": round(avg_wer_val, 4), "avg_cer": round(avg_cer_val, 4), "avg_speaker_accuracy": round(avg_spk_acc, 4), "total_time_sec": round(combo_time, 1), "files_processed": len(combo_metrics), "errors": sum(1 for m in combo_metrics if "error" in m), "per_file": combo_metrics, }) rows.sort(key=lambda r: r["avg_wer"]) return rows # --------------------------------------------------------------------------- # Grid search: one-at-a-time parameter optimization # --------------------------------------------------------------------------- WHISPER_GRID_PARAMS: dict[str, list] = { "beam_size": [1, 3, 5, 10], "no_speech_threshold": [0.3, 0.45, 0.6], "temperature": [[0.0], [0.0, 0.2, 0.4], [0.0, 0.2, 0.4, 0.6, 0.8, 1.0]], "condition_on_previous_text": [True, False], "repetition_penalty": [1.0, 1.1, 1.2], "initial_prompt": [None, "__domain__"], } DIARIZATION_GRID_PARAMS: dict[str, list] = { "min_speakers": [2, 3], "merge_gap": [0.2, 0.3, 0.5, 1.0], "merge_min_dur": [0.2, 0.3, 0.5], } def run_grid_search( file_ids: list[str], trans_model: str, diar_model: str, strategy: str, gold_by_file: dict[str, list[dict]], download_audio_fn, process_file_fn, token: str | None = None, progress_fn=None, search_whisper: bool = True, search_diar: bool = True, ) -> list[dict]: """One-at-a-time grid search over Whisper and diarization parameters. Fix all params at defaults, vary one at a time, pick best, move to next. Returns list of all tried configurations with metrics, sorted by WER. """ best_whisper: dict = {} best_diar: dict = {"min_speakers": 2, "merge_gap": 0.3, "merge_min_dur": 0.3} all_results: list[dict] = [] step = 0 params_to_try = [] if search_whisper: for name, values in WHISPER_GRID_PARAMS.items(): params_to_try.append(("whisper", name, values)) if search_diar: for name, values in DIARIZATION_GRID_PARAMS.items(): params_to_try.append(("diar", name, values)) total_configs = sum(len(vals) for _, _, vals in params_to_try) for param_group, param_name, param_values in params_to_try: best_score = float("inf") best_value = param_values[0] for val in param_values: step += 1 if progress_fn: progress_fn( step / total_configs, desc=f"Grid: {param_name}={val} ({step}/{total_configs})" ) if param_group == "whisper": cur_whisper = {**best_whisper, param_name: val} cur_diar = best_diar.copy() else: cur_whisper = best_whisper.copy() cur_diar = {**best_diar, param_name: val} combo_scores: list[dict] = [] total_time = 0.0 for file_id in file_ids: gold = gold_by_file.get(file_id, []) if not gold: continue try: audio_path = download_audio_fn(file_id) t0 = time.time() result = process_file_fn( audio_path, diar_model, trans_model, None, None, min_speakers=cur_diar.get("min_speakers", 2), max_speakers=None, strategy=strategy, whisper_kwargs=cur_whisper if cur_whisper else None, merge_gap=cur_diar.get("merge_gap", 0.3), merge_min_dur=cur_diar.get("merge_min_dur", 0.3), ) elapsed = time.time() - t0 total_time += elapsed except Exception as exc: logger.warning("Grid search error %s: %s", file_id, exc) combo_scores.append({"avg_wer": 1.0, "der": 1.0, "avg_cer": 1.0}) continue if "__error__" in result: combo_scores.append({"avg_wer": 1.0, "der": 1.0, "avg_cer": 1.0}) continue metrics = evaluate_single_result(result, gold) combo_scores.append(metrics) n = max(len(combo_scores), 1) avg_wer_val = sum(m.get("avg_wer", 1.0) for m in combo_scores) / n avg_der_val = sum(m.get("der", 1.0) for m in combo_scores) / n avg_cer_val = sum(m.get("avg_cer", 1.0) for m in combo_scores) / n row = { "param_group": param_group, "param_name": param_name, "param_value": str(val), "whisper_kwargs": cur_whisper.copy(), "diar_kwargs": cur_diar.copy(), "avg_wer": round(avg_wer_val, 4), "avg_der": round(avg_der_val, 4), "avg_cer": round(avg_cer_val, 4), "total_time_sec": round(total_time, 1), } all_results.append(row) if avg_wer_val < best_score: best_score = avg_wer_val best_value = val if param_group == "whisper": best_whisper[param_name] = best_value else: best_diar[param_name] = best_value logger.info("Grid search: best %s = %s (WER=%.4f)", param_name, best_value, best_score) all_results.sort(key=lambda r: r["avg_wer"]) return all_results