VoiceFocus / offline_pipeline.py
mariesig
VAD in spectrogram
a7c506c
Raw History Blame
8.73 kB
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 _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: str,
normalize: bool = True,
) -> tuple[np.ndarray | None, str, tuple | None, int | None]:
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.")
raise ValueError("Uploaded file exceeds the 5 MB size limit.")
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