File size: 8,730 Bytes
ba729b8
4e945b9
ba729b8
 
e368e99
03a5049
e368e99
4e945b9
 
 
 
 
 
 
 
 
 
ba729b8
7274b79
 
4e945b9
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
e368e99
 
 
 
 
4e945b9
 
 
 
 
 
7274b79
c62a089
4e945b9
e368e99
7274b79
4e945b9
e368e99
 
4e945b9
 
 
 
 
 
 
 
e368e99
4e945b9
e368e99
4e945b9
 
e368e99
4e945b9
 
 
 
 
 
 
 
 
 
 
 
 
e368e99
4e945b9
 
 
 
 
 
 
 
 
 
 
 
 
e368e99
4e945b9
c62a089
e368e99
 
c62a089
4e945b9
 
 
e368e99
 
4e945b9
e368e99
 
 
4e945b9
 
 
 
e368e99
 
 
4e945b9
 
 
 
a7c506c
4e945b9
 
 
7274b79
4e945b9
 
7274b79
e368e99
 
4e945b9
 
 
 
 
 
 
 
 
 
 
e368e99
4e945b9
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
03a5049
4e945b9
 
e368e99
4e945b9
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
a7c506c
4e945b9
a7c506c
4e945b9
 
 
e368e99
 
 
4e945b9
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
e368e99
 
4e945b9
e368e99
4e945b9
 
e368e99
4e945b9
e368e99
 
 
4e945b9
ba729b8
909a8c1
 
4e945b9
ba729b8
7274b79
4e945b9
4f24eb5
 
 
4e945b9
ba729b8
03a5049
 
 
909a8c1
 
 
 
ba729b8
4e945b9
 
 
 
ba729b8
 
c62a089
4e945b9
ba729b8
4e945b9
eb05480
c62a089
4e945b9
 
 
03a5049
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
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