from __future__ import annotations import sys import types # torchmetrics eagerly imports onnxruntime for DNSMOS, but its C extension # is broken on ZeroGPU Spaces. We don't need it — inject a lightweight stub. try: import onnxruntime as _ort_test # noqa: F401 except Exception: import importlib class _FakeInferenceSession: def __init__(self, *a, **kw): raise RuntimeError("onnxruntime stub — not available") def _make_stub(name): m = types.ModuleType(name) m.__spec__ = importlib.machinery.ModuleSpec(name, None) m.__path__ = [] m.__file__ = __file__ return m _ort_stub = _make_stub("onnxruntime") _ort_stub.__version__ = "0.0.0" _ort_stub.InferenceSession = _FakeInferenceSession _ort_stub.SessionOptions = type("SessionOptions", (), {}) _ort_stub.GraphOptimizationLevel = type("GraphOptimizationLevel", (), {"ORT_ENABLE_ALL": 99}) _ort_stub.ExecutionMode = type("ExecutionMode", (), {"ORT_SEQUENTIAL": 0}) for _sub in [ "onnxruntime.capi", "onnxruntime.capi._pybind_state", "onnxruntime.capi.onnxruntime_pybind11_state", ]: sys.modules[_sub] = _make_stub(_sub) sys.modules["onnxruntime"] = _ort_stub import json import logging import os import tempfile from pathlib import Path from typing import Any import gradio as gr from huggingface_hub import hf_hub_download, list_repo_files, snapshot_download from evaluation import ( compute_der, evaluate_single_result, find_optimal_speaker_mapping, match_segments, parse_label_studio_json, run_benchmark, run_grid_search, summarise_matches, ) from models_config import ( DIARIZATION_MODELS, TRANSCRIPTION_MODELS, diarization_choices, transcription_choices, ) from pipeline import ( TRANSCRIPTION_STRATEGIES, load_audio, parse_time_str, prefetch_model_weights, process_file, save_audio_tmp, unload_models, ) def _noop_gpu(duration: int = 300): # noqa: N802 def _decorator(fn): return fn return _decorator try: import spaces _is_zerogpu = hasattr(spaces, "GPU") and os.environ.get("SPACE_RUNTIME_STATELESS_GPU") GPU = spaces.GPU if _is_zerogpu else _noop_gpu except ImportError: GPU = _noop_gpu logging.basicConfig(level=logging.INFO) logger = logging.getLogger(__name__) APP_TITLE = "Store Dialogs — Diarization Experiments" DATASET_REPO = os.environ.get("STORE_DIALOGS_DATASET_REPO", "Niko-NN/gold-store-dialogs") BENCHMARK_VERSION = os.environ.get("STORE_DIALOGS_BENCHMARK_VERSION", "v1") # --------------------------------------------------------------------------- # Data helpers # --------------------------------------------------------------------------- _data_cache: dict[str, Any] = {} def _token() -> str | None: return os.environ.get("HF_TOKEN") or None def _get_dataset_dir() -> Path: if "dataset_dir" in _data_cache: return _data_cache["dataset_dir"] local = os.environ.get("STORE_DIALOGS_DATA_DIR", "") if local and Path(local).exists(): _data_cache["dataset_dir"] = Path(local) return _data_cache["dataset_dir"] tmp = Path(tempfile.mkdtemp(prefix="store_dialogs_")) local_dir = tmp / "dataset" snapshot_download( repo_id=DATASET_REPO, repo_type="dataset", token=_token(), local_dir=str(local_dir), local_dir_use_symlinks=False, allow_patterns=["gold/*", "reports/*", "predictions/*", "metadata/*", "audio/*"], ) _data_cache["dataset_dir"] = local_dir return local_dir def _list_audio_files() -> list[str]: """List audio files available on the Hub (audio/ folder).""" try: files = list_repo_files(DATASET_REPO, repo_type="dataset", token=_token()) return sorted( Path(f).stem for f in files if f.startswith("audio/") and not f.endswith("/") ) except Exception as exc: logger.warning("Cannot list audio files from Hub: %s", exc) return [] def _download_audio(file_id: str) -> str: """Download a single audio file from the dataset repo and return the local path.""" extensions = [".mp3", ".wav", ".ogg", ".flac", ".mp4"] for ext in extensions: remote = f"audio/{file_id}{ext}" try: return hf_hub_download( repo_id=DATASET_REPO, repo_type="dataset", filename=remote, token=_token(), ) except Exception: continue raise FileNotFoundError(f"Audio file for '{file_id}' not found in {DATASET_REPO}/audio/") def _load_jsonl(path: Path) -> list[dict]: rows: list[dict] = [] with path.open("r", encoding="utf-8") as f: for line in f: line = line.strip() if line: rows.append(json.loads(line)) return rows def _load_gold_turns() -> dict[str, list[dict]]: """Load gold turns grouped by file_id.""" data_dir = _get_dataset_dir() turns_path = data_dir / "gold" / BENCHMARK_VERSION / "turns.jsonl" if not turns_path.exists(): return {} rows = _load_jsonl(turns_path) by_file: dict[str, list[dict]] = {} for r in rows: by_file.setdefault(r["file_id"], []).append(r) return by_file # --------------------------------------------------------------------------- # GPU-decorated processing functions # --------------------------------------------------------------------------- @GPU(duration=120) def _run_single( audio_path: str, diar_model_id: str, trans_model_id: str, start_sec: float | None, end_sec: float | None, min_speakers: int | None = 2, max_speakers: int | None = None, strategy: str = "hybrid", whisper_kwargs: dict | None = None, merge_gap: float = 0.3, merge_min_dur: float = 0.3, ) -> dict: import traceback as _tb try: token = _token() return process_file( audio_path, diar_model_id, trans_model_id, token, start_sec, end_sec, min_speakers=min_speakers, max_speakers=max_speakers, strategy=strategy, whisper_kwargs=whisper_kwargs, merge_gap=merge_gap, merge_min_dur=merge_min_dur, ) except Exception: return {"__error__": _tb.format_exc()} # --------------------------------------------------------------------------- # Tab 1 — Single file processing # --------------------------------------------------------------------------- def handle_single_file( hub_file_id: str | None, uploaded_file: str | None, diar_model_id: str, trans_model_id: str, start_time_str: str, end_time_str: str, min_speakers: int, max_speakers: int, strategy: str, ): if not diar_model_id or not trans_model_id: return "Выберите модели диаризации и транскрипции.", None, [], "{}" start_sec = parse_time_str(start_time_str) end_sec = parse_time_str(end_time_str) min_sp = min_speakers if min_speakers and min_speakers > 0 else None max_sp = max_speakers if max_speakers and max_speakers > 0 else None try: if uploaded_file: audio_path = uploaded_file elif hub_file_id: audio_path = _download_audio(hub_file_id) else: return "Выберите файл из датасета или загрузите свой.", None, [], "{}" prefetch_model_weights(trans_model_id) result = _run_single( audio_path, diar_model_id, trans_model_id, start_sec, end_sec, min_speakers=min_sp, max_speakers=max_sp, strategy=strategy, ) except Exception as exc: import traceback tb = traceback.format_exc() logger.exception("Processing failed") return f"Ошибка: {type(exc).__name__}: {exc}\n\n```\n{tb}\n```", None, [], "{}" if "__error__" in result: return f"Ошибка в GPU-воркере:\n\n```\n{result['__error__']}\n```", None, [], "{}" summary_lines = [ f"### {result['file']}", f"- Диаризация: **{DIARIZATION_MODELS[diar_model_id]['name']}**", f"- Транскрипция: **{TRANSCRIPTION_MODELS[trans_model_id]['name']}**", f"- Спикеров найдено: **{result['num_speakers']}** ({', '.join(result['speakers'])})", f"- Сегментов транскрипции: **{len(result['transcription'])}**", ] summary_md = "\n".join(summary_lines) table_rows = [ [ seg.get("speaker", ""), seg.get("start", ""), seg.get("end", ""), seg.get("text", ""), ] for seg in result["transcription"] ] audio_for_player = uploaded_file or audio_path if audio_for_player and not audio_for_player.startswith(tempfile.gettempdir()): import shutil tmp_copy = os.path.join(tempfile.gettempdir(), os.path.basename(audio_for_player)) shutil.copy2(audio_for_player, tmp_copy) audio_for_player = tmp_copy return summary_md, audio_for_player, table_rows, json.dumps(result, ensure_ascii=False, indent=2) # --------------------------------------------------------------------------- # Tab 2 — Batch processing # --------------------------------------------------------------------------- def handle_batch( file_ids_text: str, diar_model_id: str, trans_model_id: str, progress: gr.Progress = gr.Progress(), ): if not diar_model_id or not trans_model_id: return "Выберите модели.", [], "{}" file_ids = [fid.strip() for fid in file_ids_text.split("\n") if fid.strip()] if not file_ids: return "Укажите ID файлов (по одному на строку).", [], "{}" all_results: list[dict] = [] errors: list[str] = [] for i, fid in enumerate(file_ids): progress((i) / len(file_ids), desc=f"Обработка {fid} ({i+1}/{len(file_ids)})") try: audio_path = _download_audio(fid) result = _run_single(audio_path, diar_model_id, trans_model_id, None, None) all_results.append(result) except Exception as exc: logger.exception("Batch error for %s", fid) errors.append(f"{fid}: {exc}") progress(1.0, desc="Готово") table_rows = [ [ r["file"], r["num_speakers"], len(r["transcription"]), " | ".join( f'{s.get("speaker", "?")}: {s.get("text", "")[:60]}' for s in r["transcription"][:3] ), ] for r in all_results ] status_parts = [f"Обработано файлов: **{len(all_results)}** из {len(file_ids)}"] if errors: status_parts.append(f"\n\nОшибки ({len(errors)}):\n" + "\n".join(f"- {e}" for e in errors)) return "\n".join(status_parts), table_rows, json.dumps(all_results, ensure_ascii=False, indent=2) # --------------------------------------------------------------------------- # Tab 3 — Compare with gold # --------------------------------------------------------------------------- _gold_cache: dict[str, list[dict]] = {} def handle_upload_gold(ls_file): """Parse Label Studio JSON export and cache gold annotations.""" if ls_file is None: return "Загрузите JSON-экспорт из Label Studio." try: with open(ls_file, "r", encoding="utf-8") as f: data = json.load(f) if isinstance(data, dict): data = [data] gold_by_file = parse_label_studio_json(data) _gold_cache.clear() _gold_cache.update(gold_by_file) files_info = [] for fid, segs in sorted(gold_by_file.items()): speakers = sorted(set(s["speaker_role"] for s in segs)) files_info.append(f"- **{fid}**: {len(segs)} сегм., спикеры: {', '.join(speakers)}") return ( f"### Загружено gold-аннотаций для {len(gold_by_file)} файлов\n\n" + "\n".join(files_info) ) except Exception as exc: return f"Ошибка парсинга: {type(exc).__name__}: {exc}" def handle_compare(results_json: str): if not results_json or results_json == "{}": return "Сначала обработайте файл на вкладке 'Обработка файла'.", [], "{}" try: result = json.loads(results_json) except json.JSONDecodeError: return "Невалидный JSON результата.", [], "{}" file_id = result.get("file", "") gold_turns = _gold_cache.get(file_id, []) if not gold_turns: hub_gold = _load_gold_turns() gold_turns = hub_gold.get(file_id, []) if not gold_turns: available = list(_gold_cache.keys()) hint = f"\n\nДоступные файлы в gold: {', '.join(available)}" if available else "" return f"Gold-разметка для `{file_id}` не найдена.{hint}", [], "{}" pred_segments = result.get("transcription", []) speaker_mapping = find_optimal_speaker_mapping(pred_segments, gold_turns) der_result = compute_der(pred_segments, gold_turns, speaker_mapping) matches = match_segments(pred_segments, gold_turns, speaker_mapping) summary = summarise_matches(matches, der_result) mapping_str = ", ".join(f"{k}→{v}" for k, v in speaker_mapping.items()) summary_lines = [ f"### Сравнение: {file_id}", f"#### Маппинг спикеров: {mapping_str}", "", "| Метрика | Значение |", "|---------|----------|", f"| **DER** (Diarization Error Rate) | **{der_result['der']:.1%}** |", f"| — пропущено (missed) | {der_result['missed']:.1f}с |", f"| — ложное срабатывание (false alarm) | {der_result['false_alarm']:.1f}с |", f"| — путаница спикеров (confusion) | {der_result['confusion']:.1f}с |", f"| **WER** (средний) | **{summary['avg_wer']:.1%}** |", f"| **CER** (средний) | **{summary['avg_cer']:.1%}** |", f"| Точность спикера | **{summary['speaker_accuracy']:.1%}** |", f"| Совпавших сегментов | {summary['matched_segments']} / {summary['total_gold_segments']} |", f"| Лишних предсказаний | {summary['unmatched_pred']} |", f"| Пропущенных gold | {summary['missed_gold']} |", ] match_rows = [] for m in matches: g = m.get("gold") p = m.get("pred") if p and g: mapped = m.get("mapped_speaker", p.get("speaker", "")) spk_icon = "✓" if m.get("speaker_match") else "✗" match_rows.append([ f'{p.get("start", ""):.1f}-{p.get("end", ""):.1f}', f'{p.get("speaker", "")} ({mapped})', p.get("text", "")[:80], g.get("speaker_role", ""), g.get("text", "")[:80], f'{spk_icon}', f'{m.get("wer", 0):.0%}', ]) elif p and not g: match_rows.append([ f'{p.get("start", ""):.1f}-{p.get("end", ""):.1f}', p.get("speaker", ""), p.get("text", "")[:80], "—", "— (лишний)", "—", "—", ]) elif g and not p: match_rows.append([ f'{g.get("start", ""):.1f}-{g.get("end", ""):.1f}', "—", "— (пропущен)", g.get("speaker_role", ""), g.get("text", "")[:80], "—", "—", ]) return "\n".join(summary_lines), match_rows, json.dumps(summary, ensure_ascii=False, indent=2) # --------------------------------------------------------------------------- # Tab 4 — Benchmark handler # --------------------------------------------------------------------------- def _handle_benchmark( trans_model_ids: list[str], diar_model_ids: list[str], strategy_ids: list[str], max_files: int, progress: gr.Progress = gr.Progress(), ): if not _gold_cache: return "Сначала загрузите gold-разметку (Label Studio JSON).", [], "{}" if not trans_model_ids or not diar_model_ids or not strategy_ids: return "Выберите хотя бы одну модель и стратегию.", [], "{}" file_ids = sorted(_gold_cache.keys()) if max_files and max_files > 0: file_ids = file_ids[:int(max_files)] for tm in trans_model_ids: prefetch_model_weights(tm) combos: list[dict] = [] for dm in diar_model_ids: for tm in trans_model_ids: for st in strategy_ids: t_name = TRANSCRIPTION_MODELS.get(tm, {}).get("name", tm) d_name = DIARIZATION_MODELS.get(dm, {}).get("name", dm) combos.append({ "diar_model": dm, "trans_model": tm, "strategy": st, "label": f"{t_name} + {d_name} / {st}", }) def _progress_fn(frac, desc=""): progress(frac, desc=desc) results = run_benchmark( file_ids=file_ids, combos=combos, gold_by_file=_gold_cache, download_audio_fn=_download_audio, process_file_fn=_run_single, token=_token(), progress_fn=_progress_fn, ) table_rows = [ [ r["label"], r["strategy"], f"{r['avg_der']:.1%}", f"{r['avg_wer']:.1%}", f"{r['avg_cer']:.1%}", f"{r['avg_speaker_accuracy']:.1%}", f"{r['total_time_sec']:.0f}", r["files_processed"], r["errors"], ] for r in results ] best = results[0] if results else None summary_parts = [ f"### Бенчмарк завершён", f"- Файлов: **{len(file_ids)}**, комбинаций: **{len(combos)}**", ] if best: summary_parts.append( f"- Лучшая комбинация: **{best['label']}** " f"(WER={best['avg_wer']:.1%}, DER={best['avg_der']:.1%}, CER={best['avg_cer']:.1%})" ) raw_clean = [{k: v for k, v in r.items() if k != "per_file"} for r in results] return ( "\n".join(summary_parts), table_rows, json.dumps(results, ensure_ascii=False, indent=2, default=str), ) # --------------------------------------------------------------------------- # Tab 5 — Grid Search handler # --------------------------------------------------------------------------- def _handle_grid_search( trans_model_id: str, diar_model_id: str, strategy_id: str, max_files: int, search_whisper: bool, search_diar: bool, progress: gr.Progress = gr.Progress(), ): if not _gold_cache: return "Сначала загрузите gold-разметку (Label Studio JSON).", [], "{}" if not trans_model_id or not diar_model_id: return "Выберите модели транскрипции и диаризации.", [], "{}" prefetch_model_weights(trans_model_id) file_ids = sorted(_gold_cache.keys()) if max_files and max_files > 0: file_ids = file_ids[:int(max_files)] def _progress_fn(frac, desc=""): progress(frac, desc=desc) results = run_grid_search( file_ids=file_ids, trans_model=trans_model_id, diar_model=diar_model_id, strategy=strategy_id, gold_by_file=_gold_cache, download_audio_fn=_download_audio, process_file_fn=_run_single, token=_token(), progress_fn=_progress_fn, search_whisper=search_whisper, search_diar=search_diar, ) table_rows = [ [ r["param_group"], r["param_name"], r["param_value"], f"{r['avg_wer']:.1%}", f"{r['avg_der']:.1%}", f"{r['avg_cer']:.1%}", f"{r['total_time_sec']:.0f}", ] for r in results ] best = results[0] if results else None summary_parts = [ f"### Grid Search завершён", f"- Протестировано конфигураций: **{len(results)}**", f"- Файлов: **{len(file_ids)}**", ] if best: summary_parts.append( f"- Лучший результат: **{best['param_name']}={best['param_value']}** " f"(WER={best['avg_wer']:.1%}, DER={best['avg_der']:.1%})" ) summary_parts.append(f"\nЛучшие Whisper-параметры: `{best['whisper_kwargs']}`") summary_parts.append(f"Лучшие Diar-параметры: `{best['diar_kwargs']}`") return ( "\n".join(summary_parts), table_rows, json.dumps(results, ensure_ascii=False, indent=2, default=str), ) # --------------------------------------------------------------------------- # Build Gradio UI # --------------------------------------------------------------------------- def _refresh_file_list(): """Lazy-load audio file list (called on button click, not at startup).""" files = _list_audio_files() return gr.Dropdown(choices=files) def build_demo() -> gr.Blocks: with gr.Blocks(title=APP_TITLE, theme=gr.themes.Soft()) as demo: gr.Markdown(f"# {APP_TITLE}") last_result_json = gr.State("{}") # ---- Tab 1: Single file ------------------------------------------- with gr.Tab("Обработка файла"): with gr.Row(): hub_file_dd = gr.Dropdown( choices=[], label="Файл из датасета", interactive=True, allow_custom_value=True, ) refresh_files_btn = gr.Button("🔄", variant="secondary", scale=0, min_width=40) uploaded_audio = gr.Audio( label="Или загрузите аудио", type="filepath", ) with gr.Row(): diar_dd = gr.Dropdown( choices=diarization_choices(), label="Модель диаризации", value=diarization_choices()[0][1] if diarization_choices() else None, interactive=True, ) trans_dd = gr.Dropdown( choices=transcription_choices(), label="Модель транскрипции", value=transcription_choices()[0][1] if transcription_choices() else None, interactive=True, ) with gr.Row(): start_time = gr.Textbox(label="Начало (мм:сс)", placeholder="0:00") end_time = gr.Textbox(label="Конец (мм:сс)", placeholder="15:00") with gr.Row(): min_speakers_num = gr.Number( label="Мин. спикеров", value=2, minimum=1, maximum=20, precision=0, interactive=True, ) max_speakers_num = gr.Number( label="Макс. спикеров (0 = авто)", value=0, minimum=0, maximum=20, precision=0, interactive=True, ) with gr.Row(): strategy_dd = gr.Dropdown( choices=[(v, k) for k, v in TRANSCRIPTION_STRATEGIES.items()], value="hybrid", label="Стратегия транскрипции", interactive=True, ) run_btn = gr.Button("Запустить", variant="primary") single_summary = gr.Markdown("Результаты появятся здесь.") audio_player = gr.Audio(label="Аудио", interactive=False) single_table = gr.Dataframe( headers=["speaker", "start", "end", "text"], label="Сегменты", wrap=True, interactive=False, ) single_raw = gr.Code(label="Raw JSON", language="json", interactive=False) refresh_files_btn.click(_refresh_file_list, outputs=[hub_file_dd]) run_btn.click( handle_single_file, inputs=[hub_file_dd, uploaded_audio, diar_dd, trans_dd, start_time, end_time, min_speakers_num, max_speakers_num, strategy_dd], outputs=[single_summary, audio_player, single_table, single_raw], ).then(lambda raw: raw, inputs=[single_raw], outputs=[last_result_json]) # ---- Tab 2: Batch ------------------------------------------------- with gr.Tab("Batch-режим"): gr.Markdown( "Введите ID файлов (по одному на строку). " "Формат: имя файла без расширения, как в датасете." ) batch_file_ids = gr.Textbox( label="ID файлов", lines=10, placeholder="2_2026-01-15_11-13-50_587\n2_2026-01-15_12-28-57_896\n...", ) with gr.Row(): batch_diar_dd = gr.Dropdown( choices=diarization_choices(), label="Модель диаризации", value=diarization_choices()[0][1] if diarization_choices() else None, interactive=True, ) batch_trans_dd = gr.Dropdown( choices=transcription_choices(), label="Модель транскрипции", value=transcription_choices()[0][1] if transcription_choices() else None, interactive=True, ) batch_btn = gr.Button("Запустить batch", variant="primary") batch_summary = gr.Markdown() batch_table = gr.Dataframe( headers=["file", "speakers", "segments", "preview"], label="Результаты", wrap=True, interactive=False, ) batch_raw = gr.Code(label="Raw JSON", language="json", interactive=False) batch_btn.click( handle_batch, inputs=[batch_file_ids, batch_diar_dd, batch_trans_dd], outputs=[batch_summary, batch_table, batch_raw], ) # ---- Tab 3: Compare ----------------------------------------------- with gr.Tab("Сравнение с gold"): gr.Markdown( "**Шаг 1:** Загрузите JSON-экспорт из Label Studio (gold-разметка).\n\n" "**Шаг 2:** Обработайте файл на вкладке 'Обработка файла'.\n\n" "**Шаг 3:** Нажмите 'Сравнить' для автоматической оценки." ) with gr.Row(): gold_upload = gr.File( label="Label Studio JSON экспорт", file_types=[".json"], type="filepath", ) gold_status = gr.Markdown("Gold-разметка не загружена.") gold_upload.change(handle_upload_gold, inputs=[gold_upload], outputs=[gold_status]) cmp_btn = gr.Button("Сравнить с результатом", variant="primary") cmp_summary = gr.Markdown() cmp_table = gr.Dataframe( headers=[ "время", "pred_спикер", "pred_текст", "gold_спикер", "gold_текст", "спикер", "WER", ], label="Сегменты: предсказание vs gold", wrap=True, interactive=False, ) cmp_raw = gr.Code(label="Метрики JSON", language="json", interactive=False) cmp_btn.click( handle_compare, inputs=[last_result_json], outputs=[cmp_summary, cmp_table, cmp_raw], ) # ---- Tab 4: Benchmark ---------------------------------------------- with gr.Tab("Бенчмарк"): gr.Markdown( "Автоматическое сравнение комбинаций **модель × стратегия** на тестовых файлах.\n\n" "1. Загрузите gold-разметку (Label Studio JSON)\n" "2. Выберите модели и стратегии\n" "3. Нажмите «Запустить бенчмарк»" ) with gr.Row(): bench_gold_upload = gr.File( label="Label Studio JSON (gold)", file_types=[".json"], type="filepath", ) bench_gold_status = gr.Markdown("Gold не загружен.") bench_gold_upload.change( handle_upload_gold, inputs=[bench_gold_upload], outputs=[bench_gold_status], ) with gr.Row(): bench_trans_models = gr.CheckboxGroup( choices=[(v["name"], k) for k, v in TRANSCRIPTION_MODELS.items()], label="Модели транскрипции", value=list(TRANSCRIPTION_MODELS.keys())[:2], ) bench_diar_models = gr.CheckboxGroup( choices=[(v["name"], k) for k, v in DIARIZATION_MODELS.items()], label="Модели диаризации", value=list(DIARIZATION_MODELS.keys()), ) with gr.Row(): bench_strategies = gr.CheckboxGroup( choices=[(v, k) for k, v in TRANSCRIPTION_STRATEGIES.items()], label="Стратегии", value=["hybrid"], ) bench_max_files = gr.Number( label="Макс. файлов (0 = все)", value=3, minimum=0, maximum=100, precision=0, ) bench_run_btn = gr.Button("Запустить бенчмарк", variant="primary") bench_summary = gr.Markdown() bench_table = gr.Dataframe( headers=["Модель", "Стратегия", "DER", "WER", "CER", "Спикер%", "Время(с)", "Файлов", "Ошибок"], label="Рейтинг комбинаций", wrap=True, interactive=False, ) bench_raw = gr.Code(label="Raw JSON", language="json", interactive=False) bench_run_btn.click( _handle_benchmark, inputs=[bench_trans_models, bench_diar_models, bench_strategies, bench_max_files], outputs=[bench_summary, bench_table, bench_raw], ) # ---- Tab 5: Grid Search ------------------------------------------- with gr.Tab("Grid Search"): gr.Markdown( "Подбор оптимальных параметров для выбранной модели.\n\n" "Стратегия **один-параметр-за-раз**: фиксируем остальные, варьируем один, " "выбираем лучший, переходим к следующему." ) with gr.Row(): grid_gold_upload = gr.File( label="Label Studio JSON (gold)", file_types=[".json"], type="filepath", ) grid_gold_status = gr.Markdown("Gold не загружен.") grid_gold_upload.change( handle_upload_gold, inputs=[grid_gold_upload], outputs=[grid_gold_status], ) with gr.Row(): grid_trans_dd = gr.Dropdown( choices=transcription_choices(), label="Модель транскрипции", value=transcription_choices()[0][1] if transcription_choices() else None, interactive=True, ) grid_diar_dd = gr.Dropdown( choices=diarization_choices(), label="Модель диаризации", value=diarization_choices()[0][1] if diarization_choices() else None, interactive=True, ) with gr.Row(): grid_strategy_dd = gr.Dropdown( choices=[(v, k) for k, v in TRANSCRIPTION_STRATEGIES.items()], value="hybrid", label="Стратегия", interactive=True, ) grid_max_files = gr.Number( label="Макс. файлов (0 = все)", value=2, minimum=0, maximum=50, precision=0, ) with gr.Row(): grid_search_whisper = gr.Checkbox(label="Whisper параметры", value=True) grid_search_diar = gr.Checkbox(label="Diar параметры", value=True) grid_run_btn = gr.Button("Запустить Grid Search", variant="primary") grid_summary = gr.Markdown() grid_table = gr.Dataframe( headers=["Группа", "Параметр", "Значение", "WER", "DER", "CER", "Время(с)"], label="Результаты grid search", wrap=True, interactive=False, ) grid_raw = gr.Code(label="Raw JSON", language="json", interactive=False) grid_run_btn.click( _handle_grid_search, inputs=[ grid_trans_dd, grid_diar_dd, grid_strategy_dd, grid_max_files, grid_search_whisper, grid_search_diar, ], outputs=[grid_summary, grid_table, grid_raw], ) # ---- Utility buttons ---------------------------------------------- with gr.Row(): diag_btn = gr.Button("Диагностика импортов", variant="secondary", size="sm") unload_btn = gr.Button("Выгрузить модели из памяти", variant="stop", size="sm") diag_output = gr.Code(label="Диагностика", language=None, interactive=False) def _run_diagnostics(): import traceback as _tb lines = [] for name in ["torch", "torchaudio", "librosa", "soundfile", "pyannote.core", "pyannote.audio", "faster_whisper", "gigaam", "jiwer", "spaces"]: try: mod = __import__(name) ver = getattr(mod, "__version__", "?") lines.append(f"OK {name} ({ver})") except Exception: lines.append(f"ERR {name}:\n{_tb.format_exc()}") import torch lines.append(f"\ntorch.cuda.is_available(): {torch.cuda.is_available()}") lines.append(f"torch version: {torch.__version__}") return "\n".join(lines) diag_btn.click(_run_diagnostics, outputs=[diag_output]) unload_btn.click(unload_models, inputs=None, outputs=None) return demo # --------------------------------------------------------------------------- # Entrypoint # --------------------------------------------------------------------------- demo = build_demo() _hf_cache = os.path.join(os.path.expanduser("~"), ".cache", "huggingface") demo.launch(allowed_paths=[_hf_cache, "/tmp"])