VoiceFocus / offline_pipeline.py
mariesig
handle various input types e.g across different gradio versions
0e1fd79
Raw History Blame
9.28 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 _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