/
TurGG
/
IRA
Обзор
Документация
Войти
/
TurGG
/
IRA
Код
Запросы
1
Задачи
Вики
Пакеты
0
Релизы
0
CI/CD
Аналитика
Безопасность
master
pages/tests.py
471 строка
22 KB
Твоё Имя
final version
25 дек 2025, 15:57
25 дек 2025, 15:57
acec5f3
Код
Авторство
О чём код?
import io import os import queue import re import tempfile import time from dataclasses import dataclass from pathlib import Path from typing import Dict, List, Tuple import streamlit as st from difflib import SequenceMatcher from pydub import AudioSegment from streamlit_webrtc import WebRtcMode, webrtc_streamer from transcribe import transcribe_audio # Ensure ffmpeg/ffprobe are visible to pydub (local bundle under ffmpeg_tmp). _FFMPEG_BIN = Path(__file__).resolve().parent.parent / "ffmpeg_tmp" / "ffmpeg" / "ffmpeg-8.0.1-essentials_build" / "bin" if _FFMPEG_BIN.exists(): os.environ["PATH"] = f"{_FFMPEG_BIN}{os.pathsep}{os.environ.get('PATH', '')}" _ffmpeg_exe = _FFMPEG_BIN / "ffmpeg.exe" _ffprobe_exe = _FFMPEG_BIN / "ffprobe.exe" if _ffmpeg_exe.exists(): AudioSegment.converter = str(_ffmpeg_exe) AudioSegment.ffmpeg = str(_ffmpeg_exe) if _ffprobe_exe.exists(): AudioSegment.ffprobe = str(_ffprobe_exe) def normalize_words(text: str) -> List[str]: sanitized = re.sub(r"[^\w\s]", " ", text.lower()) return [token for token in sanitized.split() if token] def calculate_wer(reference: str, hypothesis: str) -> float: ref_words = normalize_words(reference) hyp_words = normalize_words(hypothesis) if not ref_words: return 0.0 if not hyp_words else 1.0 distances = [[i + j if j == 0 else 0 for j in range(len(hyp_words) + 1)] for i in range(len(ref_words) + 1)] for i in range(len(ref_words) + 1): distances[i][0] = i for j in range(len(hyp_words) + 1): distances[0][j] = j for i in range(1, len(ref_words) + 1): for j in range(1, len(hyp_words) + 1): if ref_words[i - 1] == hyp_words[j - 1]: distances[i][j] = distances[i - 1][j - 1] else: substitution = distances[i - 1][j - 1] + 1 insertion = distances[i][j - 1] + 1 deletion = distances[i - 1][j] + 1 distances[i][j] = min(substitution, insertion, deletion) wer_value = distances[-1][-1] / len(ref_words) return round(wer_value, 3) def analyze_structure(text: str) -> float: sentences = [s for s in re.split(r"[.!?]+", text) if s and s.strip()] paragraphs = max(1, text.count("\n\n") + 1) punctuation_marks = sum(text.count(p) for p in ".!?") sentence_score = min(1.0, len(sentences) / max(1, paragraphs * 2)) punctuation_score = min(1.0, punctuation_marks / max(1, len(sentences))) return round(0.55 * sentence_score + 0.45 * punctuation_score, 3) def coverage_of_terms(hypothesis: str, terms: List[str]) -> float: if not terms: return 0.0 hyp_lower = hypothesis.lower() hits = sum(1 for term in terms if term and term.lower() in hyp_lower) return round(hits / len(terms), 3) def run_local_transcription( audio_bytes: bytes, suffix: str, use_ai_enhancement: bool ) -> Tuple[str, float]: with tempfile.NamedTemporaryFile(delete=False, suffix=suffix) as tmp: tmp.write(audio_bytes) tmp.flush() path = Path(tmp.name) start = time.perf_counter() result = transcribe_audio(str(path), use_ai_enhancement=use_ai_enhancement) duration = time.perf_counter() - start path.unlink(missing_ok=True) text = ( result.get("dialog_text") or result.get("formatted_text") or result.get("cleaned_text") or result.get("raw_text") or "" ) return text.strip(), round(duration, 2) def collect_audio_frames(ctx) -> List: """Drain audio frames from a streamlit-webrtc context (if present).""" if not ctx or not getattr(ctx, "audio_receiver", None): return [] frames = [] while True: try: frame = ctx.audio_receiver.get_frame(timeout=0.1) except queue.Empty: break except Exception: break else: frames.append(frame) return frames def frames_to_wav_bytes(frames: List) -> bytes: """Convert a list of AudioFrame objects to WAV bytes using pydub.""" if not frames: return b"" first = frames[0] sample_rate = int(getattr(first, "sample_rate", 16000)) sample_width = getattr(first, "format", None) sample_width_bytes = getattr(sample_width, "bytes", 2) if sample_width else 2 channels = 1 layout = getattr(first, "layout", None) if layout is not None and hasattr(layout, "channels"): try: channels = len(layout.channels) except Exception: channels = 1 raw_audio = b"".join(frame.to_ndarray().tobytes() for frame in frames) segment = AudioSegment( data=raw_audio, sample_width=sample_width_bytes, frame_rate=sample_rate, channels=channels, ) buffer = io.BytesIO() segment.export(buffer, format="wav") return buffer.getvalue() @dataclass class ModelResult: label: str transcript: str duration: float wer: float structure: float term_coverage: float def evaluate_models( reference: str, transcripts: Dict[str, Tuple[str, float]], terms: List[str] ) -> List[ModelResult]: results = [] for label, (text, duration) in transcripts.items(): result = ModelResult( label=label, transcript=text, duration=duration, wer=calculate_wer(reference, text), structure=analyze_structure(text), term_coverage=coverage_of_terms(text, terms), ) results.append(result) return results def summarize(results: List[ModelResult]) -> Tuple[str, float]: if not results: return "", 0.0 best = min(results, key=lambda item: item.wer + (1 - item.structure) + (1 - item.term_coverage)) score = round(best.wer * 0.7 + (1 - best.structure) * 0.2 + (1 - best.term_coverage) * 0.1, 3) return best.label, score def render_criteria() -> None: st.markdown( """ **Что мы оцениваем:** 1. **Точность распознавания речи (WER)** — это доля вставок, удалений и замен по сравнению с эталонным текстом. Чем меньше показатель, тем точнее нейросеть. - Критичность терминов: замена «запрещать» на «предлагать» может полностью исказить требование. - Потеря «не», «только», «всегда» меняет смысл гипотезы. - Высокий WER превращает результат в черновик, который потребуется переписывать вручную. 2. **Распознавание бизнес-терминологии и сокращений** — задача модели повторить узкоспециализированные слова, названия компаний и внутренние акронимы. - Ошибки в ключевых именах подрывают доверие к документу. - Неправильно расшифрованный акроним делает требование бессмысленным. 3. **Форматирование, пунктуация и структура** — точки, запятые, абзацы и списки определяют логическую границу мысли. - Без знаков препинания текст превращается в «простыню», где трудно ориентироваться. - Отформатированный текст проще разобрать автоматически. 4. **Обработка шумов, разной акустики и динамики** — проверяем устойчивость модели к внешним шумам, качеству микрофона и перекрывающимся голосам. - Интервью пишутся не в студии, поэтому устойчивость к шуму критична для практики. 5. **Диаризация говорящих** — важно понимать, кто задаёт вопрос, а кто отвечает, чтобы отделить требования. - Без явного разграничения теряется контекст: кто упомянул идею и на какие спросили уточнения. """ ) def render_step_guidance() -> None: st.markdown( """ #### Пошагово 1. **Шаг 1. Подготовьте данные.** - Загрузите аудио в формате WAV/MP3/M4A. - Вставьте эталонный текст, который вам кажется наиболее точным. - Добавьте ключевые термины (по одному на строку), которые хотим увидеть в каждом ответе. 2. **Шаг 2. Запустите тест.** - Нажмите кнопку «Запустить тест», чтобы последовательно прогнали GigaAM и Whisper. - Подождите, пока модели отработают — замеряем время и качество. - Сравните результаты по таблице, выделенным ошибкам и отсутствующим словам. 3. **Шаг 3. Анализируйте текстовые карточки.** - В карточках сразу видно разметку ошибок, пропущенные слова и полный текст. - Красные и оранжевые подсветки укажут на замены и вставки, а блок пропущенных слов покажет пропуски. """ ) def highlight_transcript(reference: str, hypothesis: str) -> str: ref_words = [w.lower() for w in re.findall(r"\b\w+\b", reference)] hyp_words = re.findall(r"\b\w+\b", hypothesis) hyp_lower = [w.lower() for w in hyp_words] matcher = SequenceMatcher(None, ref_words, hyp_lower) statuses = ["equal"] * len(hyp_words) for tag, i1, i2, j1, j2 in matcher.get_opcodes(): if tag in ("replace", "insert"): for idx in range(j1, j2): statuses[idx] = tag parts = re.split(r"(\b\w+\b)", hypothesis) word_idx = 0 html_parts = [] for part in parts: if re.fullmatch(r"\b\w+\b", part): status = statuses[word_idx] if status != "equal": cls = "diff-replace" if status == "replace" else "diff-insert" html_parts.append(f"<span class='diff-word {cls}'>{part}</span>") else: html_parts.append(part) word_idx += 1 else: html_parts.append(part) return "".join(html_parts) def find_missing_words(reference: str, hypothesis: str) -> List[str]: ref_words = [w.lower() for w in re.findall(r"\b\w+\b", reference)] hyp_words = [w.lower() for w in re.findall(r"\b\w+\b", hypothesis)] matcher = SequenceMatcher(None, ref_words, hyp_words) missing = [] seen = set() for tag, i1, i2, _, _ in matcher.get_opcodes(): if tag == "delete": for word in ref_words[i1:i2]: if word not in seen: missing.append(word) seen.add(word) return missing def main() -> None: st.set_page_config(page_title="Тесты транскрипции", layout="wide") st.title("Сравнительный тест транскрипторов") st.caption( "Загрузите стих и эталонный текст, затем запускайте GigaAM и Whisper." ) render_criteria() render_step_guidance() for key, value in ( ("recording_audio", False), ("recorded_audio", None), ("recorded_audio_mime", "audio/wav"), ("last_live_recorder_payload", None), ("uploaded_audio_bytes", None), ("uploaded_filename", None), ("live_webrtc_ctx", None), ("audio_source_choice", None), ): st.session_state.setdefault(key, value) st.markdown( """ <style> .diff-word {padding:2px 4px; border-radius:4px;} .diff-replace {background:#fde2e2; color:#b91c1c; font-weight:600;} .diff-insert {background:#fff4d9; color:#78350f; font-weight:600;} .transcript-card {background:#0f111a; color:#f5f7fb; border:1px solid rgba(255,255,255,0.12); padding:1rem; border-radius:10px; font-family:monospace; white-space:pre-wrap; box-shadow:0 4px 12px rgba(0,0,0,0.3);} .missing-words {color:#f87171; font-weight:600; margin-top:0.35rem; font-size:0.95rem;} .missing-words span {font-size:0.95rem;} </style> """, unsafe_allow_html=True, ) with st.expander("Шаг 1. Подготовьте входные данные", expanded=True): st.markdown( "Загрузите исходное аудио, вставьте эталонную расшифровку и перечислите ключевые термины, " "которые должны встретиться в результатах. Это позволит оценивать как точность, так и полноту." ) uploaded = st.file_uploader( "Аудиофайл (wav / mp3 / m4a)", type=["wav", "mp3", "m4a"], help="Загрузка поддерживает стандартные аудиоформаты. Размер до ~50 МБ.", ) reference_text = st.text_area( "Эталонный текст (то, что должен повторить человек)", value=( "Новая система быстрой выплаты должна полностью устранить задержки более чем на 72 часа, " "чтобы разработчики успевали масштабировать поток финансирования." ), height=150, ) term_input = st.text_area( "Ключевые термины / бизнес-слова (по одному на строку)", value="выручка\nмаржа\nконтроль качества\nавтоматизация", height=90, ) terms = [term.strip() for term in term_input.splitlines() if term.strip()] if uploaded: st.session_state.uploaded_audio_bytes = uploaded.read() st.session_state.uploaded_filename = uploaded.name st.session_state.audio_source_choice = "Файл" st.markdown("#### Живая запись с микрофона") col_start, col_stop = st.columns(2) if st.session_state.recording_audio: col_start.button("Идёт запись…", key="recording_placeholder", disabled=True) else: if col_start.button("Начать запись", key="start_live_recording"): st.session_state.recording_audio = True st.session_state.recorded_audio = None with col_stop: if st.session_state.recording_audio: stop_request = col_stop.button("Остановить запись", key="stop_live_recording") else: col_stop.empty() stop_request = False if st.session_state.recording_audio: ctx = webrtc_streamer( key="live_audio_recording", mode=WebRtcMode.SENDONLY, media_stream_constraints={"audio": True, "video": False}, async_processing=True, ) st.session_state.live_webrtc_ctx = ctx else: ctx = None if stop_request and st.session_state.live_webrtc_ctx: frames = collect_audio_frames(st.session_state.live_webrtc_ctx) wav_bytes = frames_to_wav_bytes(frames) st.session_state.recorded_audio = wav_bytes st.session_state.live_webrtc_ctx.stop() st.session_state.live_webrtc_ctx = None st.session_state.recording_audio = False if wav_bytes: st.success("Запись сохранена и готова к тесту.") st.session_state.audio_source_choice = "Запись" else: st.warning("Не удалось захватить звук, попробуй ещё раз.") recorded_audio = st.session_state.get("recorded_audio") if recorded_audio: st.audio(recorded_audio, format="audio/wav") st.caption("Последняя запись будет автоматически предложена как источник.") uploaded_audio_bytes = st.session_state.get("uploaded_audio_bytes") audio_choices: List[str] = [] if recorded_audio: audio_choices.append("Запись") if uploaded_audio_bytes: audio_choices.append("Файл") if "audio_source_choice" not in st.session_state: st.session_state.audio_source_choice = audio_choices[0] if audio_choices else None audio_source_choice = st.session_state.get("audio_source_choice") if audio_source_choice not in audio_choices: audio_source_choice = audio_choices[0] if audio_choices else None st.session_state.audio_source_choice = audio_source_choice if audio_choices: chosen_index = max(0, audio_choices.index(audio_source_choice)) if audio_source_choice else 0 audio_source_choice = st.radio( "Источник аудио для теста", audio_choices, index=chosen_index, key="audio_source_choice", ) selected_audio_bytes = None if audio_source_choice == "Запись": selected_audio_bytes = recorded_audio elif audio_source_choice == "Файл": selected_audio_bytes = uploaded_audio_bytes if audio_source_choice: msg = ( "Будет использована последняя запись с микрофона." if audio_source_choice == "Запись" else f"Будет использован файл {st.session_state.get('uploaded_filename') or 'из загрузки'}" ) st.info(msg) can_run = bool(selected_audio_bytes) and bool(reference_text.strip()) if can_run: st.info( "Шаг 2. Готово! Нажмите кнопку, чтобы последовательно прогнать GigaAM и Whisper. " "Мы запоминаем время отклика и качество каждой модели.", ) if st.button("Запустить тест"): if not can_run: st.warning("Добавь аудио и эталонный текст перед запуском.") st.stop() audio_bytes = selected_audio_bytes if audio_source_choice == "Запись": suffix = ".wav" else: suffix = Path(st.session_state.get("uploaded_filename", "audio.wav")).suffix or ".wav" transcripts: Dict[str, Tuple[str, float]] = {} with st.spinner("GigaAM (контрольный запуск)..."): text, duration = run_local_transcription(audio_bytes, suffix, use_ai_enhancement=False) transcripts["GigaAM (контроль)"] = (text, duration) with st.spinner("Whisper..."): text, duration = run_local_transcription(audio_bytes, suffix, use_ai_enhancement=True) transcripts["Whisper"] = (text, duration) results = evaluate_models(reference_text, transcripts, terms) st.session_state.tests_results = results winner_label, winner_score = summarize(results) st.session_state.last_winner = winner_label st.session_state.last_score = winner_score results = st.session_state.get("tests_results", []) if results: st.subheader("Результаты сравнения") rows = [ { "Модель": item.label, "WER": f"{item.wer:.3f}", "Структура": f"{item.structure:.2f}", "Термы": f"{item.term_coverage:.2f}", "Время, с": f"{item.duration:.2f}", } for item in results ] st.table(rows) winner = st.session_state.get("last_winner") score = st.session_state.get("last_score") if winner: st.success(f"Победитель: {winner} (показатель {score:.3f})") for idx, item in enumerate(results): st.markdown(f"**{item.label}** — текст:", unsafe_allow_html=True) highlighted = highlight_transcript(reference_text, item.transcript) st.markdown(f"<div class='transcript-card'>{highlighted}</div>", unsafe_allow_html=True) missing_words = find_missing_words(reference_text, item.transcript) if missing_words: formatted = ", ".join(missing_words) st.markdown( f"<div class='missing-words'>Пропущенные слова: {formatted}</div>", unsafe_allow_html=True, ) with st.expander("Полный текст расшифровки", expanded=False): st.text_area( f"Текст {item.label}", value=item.transcript, height=150, key=f"full_transcript_{idx}", ) else: st.info( "Результатов пока нет. Пройдите первый шаг — загрузите аудио, эталон и термины — затем запустите сравнение.", icon="ℹ️", ) if __name__ == "__main__": main()