Niko-NN's picture
fix: 3 critical issues — GPU decorator, gold labels, grid search params
faa18b6
Raw
History Blame Contribute Delete
36.4 kB
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"])