Spaces:
Paused
Paused
| 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 | |
| # --------------------------------------------------------------------------- | |
| 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"]) | |