import os from typing import Any import gradio as gr import numpy as np import librosa from constants import APP_TMP_DIR, STREAMER_CLASSES from hf_dataset_utils import get_audio, get_transcript from sdk import SDKParams, SDKWrapper from utils import ( compute_wer, get_vad_labels, normalize_lufs, spec_image, to_gradio_audio, ) SDK_OFFLINE = SDKWrapper() def _extract_uploaded_path(sample_input: Any) -> str | None: if sample_input is None: return None if isinstance(sample_input, str): return sample_input for attr in ("path", "name"): value = getattr(sample_input, attr, None) if isinstance(value, str): return value if isinstance(sample_input, dict): for key in ("path", "name"): value = sample_input.get(key) if isinstance(value, str): return value return None def _safe_progress(progress: gr.Progress, value: float, desc: str) -> None: progress(max(0.0, min(1.0, value)), desc=desc) def _empty_pipeline_result(sample_id: str) -> tuple[Any, str, str, str, str, str, str]: return ( None, "", "", "Unavailable", "Unavailable", "Unavailable", sample_id, ) def _finalize_stream_transcript(streamer) -> str: if hasattr(streamer, "close_stream"): streamer.close_stream() else: streamer.close() streamer.finished_event.wait() with streamer.lock: return streamer.render_tokens(streamer.final_tokens, []) def _init_sdk(sample_rate: int, enhancement_level: int) -> int: sdk_params = SDKParams( sample_rate=sample_rate, enhancement_level=enhancement_level / 100.0, ) SDK_OFFLINE.init_processor(sdk_params) return SDK_OFFLINE.num_frames def _init_streamers( sample_rate: int, stt_model: str, sample_id: str, progress: gr.Progress, ): if stt_model not in STREAMER_CLASSES: raise ValueError(f"Unknown STT model: {stt_model}") streamer_class = STREAMER_CLASSES[stt_model] _safe_progress(progress, 0.12, f"Initializing {stt_model} stream 1/2...") streamer_noisy = streamer_class(sample_rate, f"{sample_id}_noisy") _safe_progress(progress, 0.18, f"Initializing {stt_model} stream 2/2...") streamer_enhanced = streamer_class(sample_rate, f"{sample_id}_enhanced") return streamer_noisy, streamer_enhanced def _attach_wer( original_transcript: str, noisy_transcript: str, enhanced_transcript: str, ) -> tuple[str, str]: wer_enhanced = compute_wer(original_transcript, enhanced_transcript) wer_noisy = compute_wer(original_transcript, noisy_transcript) noisy_transcript = f"{noisy_transcript} (WER: {wer_noisy * 100:.2f}%)" enhanced_transcript = f"{enhanced_transcript} (WER: {wer_enhanced * 100:.2f}%)" return noisy_transcript, enhanced_transcript def _process_audio_chunks( sample: np.ndarray, sample_rate: int, chunk_size: int, streamer_noisy, streamer_enhanced, progress: gr.Progress, ) -> tuple[np.ndarray, list[list[float]]]: accumulated_enhanced: list[np.ndarray] = [] vad_timestamps: list[list[float]] = [] n = len(sample) for i in range(0, n, chunk_size): raw_chunk = sample[i : i + chunk_size] original_chunk_len = raw_chunk.size if original_chunk_len < chunk_size: raw_chunk = np.pad( raw_chunk, (0, chunk_size - original_chunk_len), mode="constant", constant_values=0.0, ) enhanced_chunk = SDK_OFFLINE.process_chunk(raw_chunk.reshape(1, -1)) enhanced_1d = np.asarray(enhanced_chunk, dtype=np.float32).flatten() streamer_noisy.process_chunk(raw_chunk) streamer_enhanced.process_chunk(enhanced_1d) accumulated_enhanced.append(enhanced_1d) loop_progress = (i + original_chunk_len) / n if n > 0 else 1.0 _safe_progress( progress, 0.20 + 0.50 * loop_progress, "Enhancing audio...", ) if SDK_OFFLINE.vad_context.is_speech_detected(): start_in_sec = i / sample_rate end_in_sec = min(i + original_chunk_len, n) / sample_rate vad_timestamps.append([start_in_sec, end_in_sec]) enhanced_array = np.concatenate(accumulated_enhanced).astype(np.float32) return enhanced_array, vad_timestamps def _save_spectrograms( sample: np.ndarray, enhanced_array: np.ndarray, sample_rate: int, sample_id: str, vad_timestamps: list[list[float]], ) -> tuple[str, str]: os.makedirs(APP_TMP_DIR, exist_ok=True) enhanced_spec_path = os.path.join(APP_TMP_DIR, f"{sample_id}_enhanced_spectrogram.png") noisy_spec_path = os.path.join(APP_TMP_DIR, f"{sample_id}_noisy_spectrogram.png") spec_image(enhanced_array, sr=sample_rate, vad_timestamps=vad_timestamps).save(enhanced_spec_path) spec_image(sample, sr=sample_rate, vad_timestamps=vad_timestamps).save(noisy_spec_path) return enhanced_spec_path, noisy_spec_path def run_offline_pipeline( sample: np.ndarray, sample_rate: int, enhancement_level: int, stt_model: str, sample_id: str, progress=gr.Progress(), ) -> tuple[Any, str, str, str, str, str, str]: _safe_progress(progress, 0.00, "Starting...") if sample is None or len(sample) == 0: gr.Warning("No audio to enhance. Please upload a file first.") return _empty_pipeline_result(sample_id) _safe_progress(progress, 0.05, "Initializing enhancement...") chunk_size = _init_sdk(sample_rate, enhancement_level) try: streamer_noisy, streamer_enhanced = _init_streamers( sample_rate=sample_rate, stt_model=stt_model, sample_id=sample_id, progress=progress, ) except Exception as e: raise RuntimeError(f"Failed to initialize STT streaming: {e}") from e enhanced_array, vad_timestamps = _process_audio_chunks( sample=sample, sample_rate=sample_rate, chunk_size=chunk_size, streamer_noisy=streamer_noisy, streamer_enhanced=streamer_enhanced, progress=progress, ) _safe_progress(progress, 0.72, "Finalizing transcripts...") noisy_transcript = _finalize_stream_transcript(streamer_noisy) _safe_progress(progress, 0.80, "Finalizing transcripts...") enhanced_transcript = _finalize_stream_transcript(streamer_enhanced) _safe_progress(progress, 0.94, "Loading reference transcript...") try: original_transcript = get_transcript(sample_id) except Exception: original_transcript = "Unavailable" if original_transcript != "Unavailable": _safe_progress(progress, 0.96, "Computing WER...") noisy_transcript, enhanced_transcript = _attach_wer( original_transcript=original_transcript, noisy_transcript=noisy_transcript, enhanced_transcript=enhanced_transcript, ) _safe_progress(progress, 0.99, "Generating outputs...") gradio_enhanced_audio = to_gradio_audio(enhanced_array, sample_rate) enhanced_spec_path, noisy_spec_path = _save_spectrograms( sample=sample, enhanced_array=enhanced_array, sample_rate=sample_rate, sample_id=sample_id, vad_timestamps=vad_timestamps ) vad_labels = get_vad_labels( vad_timestamps, length=len(sample) / sample_rate, ) _safe_progress(progress, 1.00, "Done.") return ( gr.update(value=gradio_enhanced_audio, subtitles=vad_labels), enhanced_spec_path, noisy_spec_path, original_transcript, noisy_transcript, enhanced_transcript, sample_id, ) def load_local_file( sample_path: Any, normalize: bool = True, ) -> tuple[np.ndarray | None, str, tuple | None, int | None]: sample_path = _extract_uploaded_path(sample_path) if not sample_path or not os.path.exists(sample_path): return None, "", None, None if os.path.getsize(sample_path) > 5 * 1024 * 1024: gr.Warning("File size exceeds 5 MB limit. Please upload a smaller file.") return None, "", None, None new_sample_stem = os.path.splitext(os.path.basename(sample_path))[0] y, sample_rate = librosa.load(sample_path, sr=None, mono=True) sample_rate = int(sample_rate) y = np.asarray(y, dtype=np.float32) if normalize: y = normalize_lufs(y, sample_rate) gradio_audio = to_gradio_audio(y, sample_rate) return y, new_sample_stem, gradio_audio, sample_rate def load_file_from_dataset( sample_id: str, ) -> tuple[tuple | None, np.ndarray | None, str, int | None]: if not sample_id: gr.Warning("Please select a sample from the dropdown.") return None, None, "", None new_sample_stem = sample_id try: y, sample_rate = get_audio(sample_id, prefix="mix") except Exception as e: gr.Warning(str(e)) raise y = np.asarray(y, dtype=np.float32) if y.ndim > 1: y = np.mean(y, axis=0) gradio_audio = to_gradio_audio(y, sample_rate) return gradio_audio, y, new_sample_stem, sample_rate