Spaces:
Running on CPU Upgrade
Running on CPU Upgrade
mariesig commited on
Commit ·
1df4f51
1
Parent(s): cdd0f38
initial commit
Browse files- aic_dataset.py +60 -0
- app.py +22 -7
- audio_tools.py +22 -59
- constants.py +11 -40
- intro.md +2 -10
- requirements.txt +2 -1
- sdk.py +0 -4
- stt_streamers/__init__.py +4 -0
- stt_streamers/deepgram_streamer.py +195 -0
- stt_streamers/soniox_streamer.py +141 -0
- transcribe.py +51 -0
- word_error_rate.py +78 -0
aic_dataset.py
ADDED
|
@@ -0,0 +1,60 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import os
|
| 2 |
+
from tkinter import ALL
|
| 3 |
+
import gradio as gr
|
| 4 |
+
from gradio_client import file
|
| 5 |
+
from huggingface_hub import hf_hub_download
|
| 6 |
+
from constants import DATASET_REPO, MIX_DIR, TRANS_DIR, DATASET_METADATA, HF_TOKEN
|
| 7 |
+
|
| 8 |
+
|
| 9 |
+
def _get_base_filenames_from_metadata(metadata_path):
|
| 10 |
+
local_path = hf_hub_download(
|
| 11 |
+
repo_id=DATASET_REPO,
|
| 12 |
+
repo_type="dataset",
|
| 13 |
+
filename=metadata_path,
|
| 14 |
+
token=HF_TOKEN
|
| 15 |
+
)
|
| 16 |
+
with open(local_path, "r", encoding="utf-8") as f:
|
| 17 |
+
lines = f.read().splitlines()
|
| 18 |
+
base_names = []
|
| 19 |
+
for line in lines[1:]: # skip header
|
| 20 |
+
parts = line.split(",")
|
| 21 |
+
if parts and parts[0]:
|
| 22 |
+
# Remove directory and extension
|
| 23 |
+
filename = os.path.splitext(os.path.basename(parts[0]))[0]
|
| 24 |
+
base_names.append(filename)
|
| 25 |
+
return base_names
|
| 26 |
+
|
| 27 |
+
ALL_FILES = _get_base_filenames_from_metadata(DATASET_METADATA)
|
| 28 |
+
|
| 29 |
+
def get_local_mix_path(file_stem: str) -> str:
|
| 30 |
+
if not file_stem:
|
| 31 |
+
return ""
|
| 32 |
+
|
| 33 |
+
mix_path = f"{MIX_DIR}/{file_stem}.wav"
|
| 34 |
+
|
| 35 |
+
# Download selected files into local cache; returns local filesystem paths
|
| 36 |
+
mix_local = hf_hub_download(
|
| 37 |
+
repo_id=DATASET_REPO, repo_type="dataset",
|
| 38 |
+
filename=mix_path, token=HF_TOKEN
|
| 39 |
+
)
|
| 40 |
+
|
| 41 |
+
return mix_local
|
| 42 |
+
|
| 43 |
+
def download_transcript(file_stem: str) -> str:
|
| 44 |
+
"""
|
| 45 |
+
file_stem is the base filename to be downloaded
|
| 46 |
+
"""
|
| 47 |
+
if not file_stem:
|
| 48 |
+
return ""
|
| 49 |
+
|
| 50 |
+
transcript_path = f"{TRANS_DIR}/{file_stem}.txt"
|
| 51 |
+
|
| 52 |
+
|
| 53 |
+
transcript_text = ""
|
| 54 |
+
transcript_local = hf_hub_download(
|
| 55 |
+
repo_id=DATASET_REPO, repo_type="dataset",
|
| 56 |
+
filename=transcript_path, token=HF_TOKEN
|
| 57 |
+
)
|
| 58 |
+
with open(transcript_local, "r", encoding="utf-8", errors="replace") as f:
|
| 59 |
+
transcript_text = f.read()
|
| 60 |
+
return transcript_text
|
app.py
CHANGED
|
@@ -9,10 +9,11 @@ from constants import (
|
|
| 9 |
MINUTES_KEEP,
|
| 10 |
)
|
| 11 |
from sdk import SDKParams, process_file_sdk
|
| 12 |
-
from audio_tools import spec_image
|
| 13 |
import shutil
|
| 14 |
import tempfile
|
| 15 |
-
|
|
|
|
| 16 |
|
| 17 |
# ===============================
|
| 18 |
# Temporary File & Cache Management
|
|
@@ -152,7 +153,7 @@ def start_processing(sample_path: str) -> tuple[Any, str, str]:
|
|
| 152 |
|
| 153 |
|
| 154 |
def enable_new_input():
|
| 155 |
-
return gr.update(
|
| 156 |
|
| 157 |
|
| 158 |
# ===============================
|
|
@@ -169,11 +170,19 @@ with gr.Blocks() as demo:
|
|
| 169 |
)
|
| 170 |
with gr.Row():
|
| 171 |
gr.Markdown(open("intro.md").read())
|
|
|
|
| 172 |
with gr.Row():
|
| 173 |
-
|
| 174 |
-
|
| 175 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
| 176 |
)
|
|
|
|
|
|
|
|
|
|
| 177 |
percent_slider = gr.Slider(
|
| 178 |
minimum=1,
|
| 179 |
maximum=100,
|
|
@@ -192,7 +201,10 @@ with gr.Blocks() as demo:
|
|
| 192 |
noisy_image = gr.Image(label="Noisy spectrogram", format="png", type="filepath")
|
| 193 |
enhanced_audio = gr.Audio(type="filepath", label="Enhanced audio")
|
| 194 |
enhanced_image = gr.Image(label="Enhanced spectrogram", format="png", type="filepath")
|
|
|
|
|
|
|
| 195 |
|
|
|
|
| 196 |
btn.click(cleanup, [input_enhancement, last_audio_file, audio_file], None).then(
|
| 197 |
start_processing,
|
| 198 |
inputs=audio_file,
|
|
@@ -204,7 +216,10 @@ with gr.Blocks() as demo:
|
|
| 204 |
percent_slider,
|
| 205 |
],
|
| 206 |
outputs=[noisy_audio, noisy_image, enhanced_audio, enhanced_image],
|
| 207 |
-
).then(enable_new_input, None, btn)
|
|
|
|
|
|
|
|
|
|
| 208 |
|
| 209 |
cleanup_tmp(minutes_keep=0, filter=[])
|
| 210 |
demo.launch(allowed_paths=["/tmp", "/"])
|
|
|
|
| 9 |
MINUTES_KEEP,
|
| 10 |
)
|
| 11 |
from sdk import SDKParams, process_file_sdk
|
| 12 |
+
from audio_tools import spec_image
|
| 13 |
import shutil
|
| 14 |
import tempfile
|
| 15 |
+
from aic_dataset import ALL_FILES, get_local_mix_path
|
| 16 |
+
from transcribe import transcribe_file
|
| 17 |
|
| 18 |
# ===============================
|
| 19 |
# Temporary File & Cache Management
|
|
|
|
| 153 |
|
| 154 |
|
| 155 |
def enable_new_input():
|
| 156 |
+
return gr.update(interactive=True)
|
| 157 |
|
| 158 |
|
| 159 |
# ===============================
|
|
|
|
| 170 |
)
|
| 171 |
with gr.Row():
|
| 172 |
gr.Markdown(open("intro.md").read())
|
| 173 |
+
|
| 174 |
with gr.Row():
|
| 175 |
+
dataset_dropdown = gr.Dropdown(
|
| 176 |
+
choices=ALL_FILES,
|
| 177 |
+
label="Pick example from our dataset",
|
| 178 |
+
value=None
|
| 179 |
+
)
|
| 180 |
+
audio_file = gr.Audio(
|
| 181 |
+
type="filepath", label="Input", visible=True, sources=["upload"]
|
| 182 |
)
|
| 183 |
+
|
| 184 |
+
with gr.Row():
|
| 185 |
+
with gr.Column():
|
| 186 |
percent_slider = gr.Slider(
|
| 187 |
minimum=1,
|
| 188 |
maximum=100,
|
|
|
|
| 201 |
noisy_image = gr.Image(label="Noisy spectrogram", format="png", type="filepath")
|
| 202 |
enhanced_audio = gr.Audio(type="filepath", label="Enhanced audio")
|
| 203 |
enhanced_image = gr.Image(label="Enhanced spectrogram", format="png", type="filepath")
|
| 204 |
+
transcript_box = gr.Textbox(label="Transcript", lines=8)
|
| 205 |
+
|
| 206 |
|
| 207 |
+
|
| 208 |
btn.click(cleanup, [input_enhancement, last_audio_file, audio_file], None).then(
|
| 209 |
start_processing,
|
| 210 |
inputs=audio_file,
|
|
|
|
| 216 |
percent_slider,
|
| 217 |
],
|
| 218 |
outputs=[noisy_audio, noisy_image, enhanced_audio, enhanced_image],
|
| 219 |
+
).success(transcribe_file, inputs=[enhanced_audio], outputs=[transcript_box]).then(enable_new_input, None, [btn])
|
| 220 |
+
|
| 221 |
+
dataset_dropdown.change(get_local_mix_path, inputs=dataset_dropdown, outputs=[audio_file])
|
| 222 |
+
|
| 223 |
|
| 224 |
cleanup_tmp(minutes_keep=0, filter=[])
|
| 225 |
demo.launch(allowed_paths=["/tmp", "/"])
|
audio_tools.py
CHANGED
|
@@ -4,8 +4,6 @@ import librosa
|
|
| 4 |
from PIL import Image
|
| 5 |
import io
|
| 6 |
import matplotlib.pyplot as plt
|
| 7 |
-
import soundfile as sf
|
| 8 |
-
|
| 9 |
|
| 10 |
def spec_image(
|
| 11 |
audio_wav: str,
|
|
@@ -44,62 +42,27 @@ def spec_image(
|
|
| 44 |
return Image.open(buf).convert("RGB")
|
| 45 |
|
| 46 |
|
| 47 |
-
def
|
| 48 |
-
signal_path: str,
|
| 49 |
-
noise_path: str,
|
| 50 |
-
output_path: str = "output.wav",
|
| 51 |
-
snr_db: float = 10.0,
|
| 52 |
-
rng: Optional[np.random.Generator] = None,
|
| 53 |
-
) -> bool:
|
| 54 |
"""
|
| 55 |
-
|
| 56 |
-
|
| 57 |
-
Args:
|
| 58 |
-
signal_wav: Path to clean/foreground audio (wav/mp3).
|
| 59 |
-
noise_wav: Path to noise audio (wav/mp3).
|
| 60 |
-
snr_db: Desired SNR in dB (signal/noise).
|
| 61 |
-
normalize: If True, peak-normalize the mixture to |x|max=1 after mixing
|
| 62 |
-
(note: can slightly alter achieved SNR).
|
| 63 |
-
rng: Optional numpy Generator for reproducible random cropping.
|
| 64 |
-
|
| 65 |
-
Returns:
|
| 66 |
-
Bool wether clipping occurred.
|
| 67 |
"""
|
| 68 |
-
|
| 69 |
-
|
| 70 |
-
|
| 71 |
-
|
| 72 |
-
|
| 73 |
-
|
| 74 |
-
|
| 75 |
-
|
| 76 |
-
|
| 77 |
-
|
| 78 |
-
|
| 79 |
-
|
| 80 |
-
|
| 81 |
-
|
| 82 |
-
|
| 83 |
-
|
| 84 |
-
|
| 85 |
-
|
| 86 |
-
|
| 87 |
-
|
| 88 |
-
noise_power = float(np.mean(noise**2))
|
| 89 |
-
|
| 90 |
-
if sig_power == 0.0:
|
| 91 |
-
out = noise * 0.0
|
| 92 |
-
elif noise_power == 0.0:
|
| 93 |
-
out = sig.copy()
|
| 94 |
-
else:
|
| 95 |
-
target_noise_power = sig_power / (10.0 ** (snr_db / 10.0))
|
| 96 |
-
scale = np.sqrt(target_noise_power / noise_power)
|
| 97 |
-
noise_scaled = noise * scale
|
| 98 |
-
out = sig + noise_scaled
|
| 99 |
-
|
| 100 |
-
peak = np.max(np.abs(out)) or 1.0
|
| 101 |
-
if peak > 1.0:
|
| 102 |
-
clipped = True
|
| 103 |
-
out = out / peak
|
| 104 |
-
sf.write(output_path, out, sr_s)
|
| 105 |
-
return clipped
|
|
|
|
| 4 |
from PIL import Image
|
| 5 |
import io
|
| 6 |
import matplotlib.pyplot as plt
|
|
|
|
|
|
|
| 7 |
|
| 8 |
def spec_image(
|
| 9 |
audio_wav: str,
|
|
|
|
| 42 |
return Image.open(buf).convert("RGB")
|
| 43 |
|
| 44 |
|
| 45 |
+
def compute_wer(reference: str, hypothesis: str) -> float:
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 46 |
"""
|
| 47 |
+
Compute Word Error Rate (WER) between reference and hypothesis transcripts.
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 48 |
"""
|
| 49 |
+
ref_words = reference.split()
|
| 50 |
+
hyp_words = hypothesis.split()
|
| 51 |
+
d = np.zeros((len(ref_words) + 1, len(hyp_words) + 1), dtype=np.uint8)
|
| 52 |
+
for i in range(len(ref_words) + 1):
|
| 53 |
+
d[i][0] = i
|
| 54 |
+
for j in range(len(hyp_words) + 1):
|
| 55 |
+
d[0][j] = j
|
| 56 |
+
for i in range(1, len(ref_words) + 1):
|
| 57 |
+
for j in range(1, len(hyp_words) + 1):
|
| 58 |
+
if ref_words[i - 1] == hyp_words[j - 1]:
|
| 59 |
+
cost = 0
|
| 60 |
+
else:
|
| 61 |
+
cost = 1
|
| 62 |
+
d[i][j] = min(
|
| 63 |
+
d[i - 1][j] + 1, # Deletion
|
| 64 |
+
d[i][j - 1] + 1, # Insertion
|
| 65 |
+
d[i - 1][j - 1] + cost, # Substitution
|
| 66 |
+
)
|
| 67 |
+
wer = d[len(ref_words)][len(hyp_words)] / max(len(ref_words), 1)
|
| 68 |
+
return wer
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
constants.py
CHANGED
|
@@ -1,48 +1,19 @@
|
|
| 1 |
from typing import Final
|
|
|
|
| 2 |
|
| 3 |
-
|
| 4 |
-
ENHANCEMENT_MODELS: Final = [("Quail", "QUAIL"), ("Finch", "FINCH"), ("Lark 2", "LARK_V2")]
|
| 5 |
-
API_V2_URL: Final = "https://api.ai-coustics.io/v2"
|
| 6 |
CHUNK_SIZE: Final = 1024
|
| 7 |
TIMEOUT_FACTOR_MB: Final = 60
|
| 8 |
BASE_TIMEOUT_SECONDS: Final = 120
|
| 9 |
|
| 10 |
-
|
| 11 |
MINUTES_KEEP: Final = 60
|
| 12 |
|
| 13 |
-
|
| 14 |
-
|
| 15 |
-
|
| 16 |
-
|
| 17 |
-
|
| 18 |
-
|
| 19 |
-
|
| 20 |
-
|
| 21 |
-
|
| 22 |
-
|
| 23 |
-
[
|
| 24 |
-
"assets/samples/input/Background.wav",
|
| 25 |
-
100,
|
| 26 |
-
"FINCH",
|
| 27 |
-
],
|
| 28 |
-
[
|
| 29 |
-
"assets/samples/input/Reverb.wav",
|
| 30 |
-
100,
|
| 31 |
-
"QUAIL",
|
| 32 |
-
],
|
| 33 |
-
[
|
| 34 |
-
"assets/samples/input/Distortion.wav",
|
| 35 |
-
100,
|
| 36 |
-
"LARK_V2",
|
| 37 |
-
],
|
| 38 |
-
[
|
| 39 |
-
"assets/samples/input/Wind.wav",
|
| 40 |
-
100,
|
| 41 |
-
"LARK_V2",
|
| 42 |
-
],
|
| 43 |
-
[
|
| 44 |
-
"assets/samples/input/Music.wav",
|
| 45 |
-
100,
|
| 46 |
-
"LARK_V2",
|
| 47 |
-
],
|
| 48 |
-
]
|
|
|
|
| 1 |
from typing import Final
|
| 2 |
+
import os
|
| 3 |
|
|
|
|
|
|
|
|
|
|
| 4 |
CHUNK_SIZE: Final = 1024
|
| 5 |
TIMEOUT_FACTOR_MB: Final = 60
|
| 6 |
BASE_TIMEOUT_SECONDS: Final = 120
|
| 7 |
|
|
|
|
| 8 |
MINUTES_KEEP: Final = 60
|
| 9 |
|
| 10 |
+
DATASET_REPO: Final = "ai-coustics/leo_butch_voice_focus_open_source_en_vad"
|
| 11 |
+
MIX_DIR: Final = "mix"
|
| 12 |
+
SPEECH_DIR: Final = "speech"
|
| 13 |
+
TRANS_DIR: Final = "transcripts"
|
| 14 |
+
DATASET_METADATA: Final = "metadata.csv"
|
| 15 |
+
|
| 16 |
+
# Private access token from Space secrets:
|
| 17 |
+
HF_TOKEN: Final = os.getenv("HF_TOKEN") # set in HF Space Secrets
|
| 18 |
+
|
| 19 |
+
DEFAULT_SR: Final = 16000
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
intro.md
CHANGED
|
@@ -1,10 +1,2 @@
|
|
| 1 |
-
|
| 2 |
-
|
| 3 |
-
- [**Quail**](https://ai-coustics.com/2025/07/02/meet-quail-the-most-advanced-real-time-speech-enhancement-model/): Our lightest model, created for real-time voice isolation. Try it for free or integrate via our [SDK](https://docs.ai-coustics.com/sdk/overview).
|
| 4 |
-
- [**Finch**](https://ai-coustics.com/2025/06/12/introducing-finch-2-the-new-aicoustics-model-for-studio-quality-speech/): A larger, subtractive model focused on voice isolation and noise removal while preserving the original voice.
|
| 5 |
-
- [**Lark**](https://ai-coustics.com/2025/07/29/lark-2-next-generation-reconstructive-speech-enhancement/): Our most advanced, generative model for universal speech enhancement. Lark not only removes noise but also reconstructs and restores lost details in audio, keeping the speaker’s voice intact.
|
| 6 |
-
|
| 7 |
-
Finch and Lark are available via our [API](https://docs.ai-coustics.com/api-reference/v2/upload-media-file). To get started, request an [API key](https://ai-coustics.com/api/). You’ll receive 120 minutes for free. For additional usage, see our [pricing details](https://ai-coustics.com/pricing/).
|
| 8 |
-
|
| 9 |
-
**How to use**
|
| 10 |
-
Upload or record a noisy audio sample, optionally add background noise, and adjust the enhancement level to your preference. Select the model you want to use. Note that Quail is typically the fastest option. You can also try the example clips below to explore to preview the capabilities of each model.
|
|
|
|
| 1 |
+
This demo showcases **Voice Focus**, our real-time speech enhancement model, combined with **speech-to-text (STT)** transcription.
|
| 2 |
+
Enhance speech, view transcripts instantly, and evaluate recognition quality — either offline or live.
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
requirements.txt
CHANGED
|
@@ -6,4 +6,5 @@ librosa>=0.10.1,<0.11
|
|
| 6 |
loguru~=0.7
|
| 7 |
aic-sdk>=2.0.0
|
| 8 |
dotenv
|
| 9 |
-
resampy
|
|
|
|
|
|
| 6 |
loguru~=0.7
|
| 7 |
aic-sdk>=2.0.0
|
| 8 |
dotenv
|
| 9 |
+
resampy
|
| 10 |
+
whisper-normalizer
|
sdk.py
CHANGED
|
@@ -9,10 +9,6 @@ import soundfile as sf
|
|
| 9 |
load_dotenv()
|
| 10 |
|
| 11 |
|
| 12 |
-
|
| 13 |
-
|
| 14 |
-
|
| 15 |
-
|
| 16 |
class SDKParams:
|
| 17 |
def __init__(self, enhancement_level: float, sdk_key: str, model_id: str = "quail-vf-1.1-l-16khz"):
|
| 18 |
self.enhancement_level = enhancement_level
|
|
|
|
| 9 |
load_dotenv()
|
| 10 |
|
| 11 |
|
|
|
|
|
|
|
|
|
|
|
|
|
| 12 |
class SDKParams:
|
| 13 |
def __init__(self, enhancement_level: float, sdk_key: str, model_id: str = "quail-vf-1.1-l-16khz"):
|
| 14 |
self.enhancement_level = enhancement_level
|
stt_streamers/__init__.py
ADDED
|
@@ -0,0 +1,4 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from .deepgram_streamer import DeepgramStreamer
|
| 2 |
+
from .soniox_streamer import SonioxStreamer
|
| 3 |
+
|
| 4 |
+
__all__ = ["DeepgramStreamer", "SonioxStreamer"]
|
stt_streamers/deepgram_streamer.py
ADDED
|
@@ -0,0 +1,195 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import json
|
| 2 |
+
import os
|
| 3 |
+
import threading
|
| 4 |
+
import urllib.parse
|
| 5 |
+
import numpy as np
|
| 6 |
+
from websockets.sync.client import connect
|
| 7 |
+
from websockets.exceptions import ConnectionClosedOK, ConnectionClosedError
|
| 8 |
+
|
| 9 |
+
DEEPGRAM_WEBSOCKET_URL = "wss://api.deepgram.com/v1/listen"
|
| 10 |
+
|
| 11 |
+
|
| 12 |
+
class DeepgramStreamer:
|
| 13 |
+
def __init__(self, fs_hz: int, stream_name: str, on_update=None) -> None:
|
| 14 |
+
api_key = os.environ.get("DEEPGRAM_API_KEY")
|
| 15 |
+
if not api_key:
|
| 16 |
+
raise RuntimeError("Missing DEEPGRAM_API_KEY.")
|
| 17 |
+
|
| 18 |
+
self.stream_name = stream_name
|
| 19 |
+
self.api_name = "Deepgram V1 Nova-3"
|
| 20 |
+
self.on_update = on_update
|
| 21 |
+
self.final_tokens: list[dict] = []
|
| 22 |
+
self.lock = threading.Lock()
|
| 23 |
+
self.finished_event = threading.Event()
|
| 24 |
+
|
| 25 |
+
# 1. Build the Deepgram URL with query parameters
|
| 26 |
+
config = self.get_config(fs_hz)
|
| 27 |
+
query_string = urllib.parse.urlencode(config)
|
| 28 |
+
url_with_params = f"{DEEPGRAM_WEBSOCKET_URL}?{query_string}"
|
| 29 |
+
|
| 30 |
+
print(f"Connecting {stream_name} to Deepgram...")
|
| 31 |
+
|
| 32 |
+
# 2. Connect with Authorization header
|
| 33 |
+
# Deepgram requires the API key in the headers
|
| 34 |
+
headers = {"Authorization": f"Token {api_key}"}
|
| 35 |
+
self.ws = connect(url_with_params, additional_headers=headers)
|
| 36 |
+
|
| 37 |
+
# 3. Start the receiving thread
|
| 38 |
+
self.thread = threading.Thread(target=self._receive_loop, daemon=True)
|
| 39 |
+
self.thread.start()
|
| 40 |
+
|
| 41 |
+
def stream_array(self, pcm: np.ndarray, fs_hz: int) -> str:
|
| 42 |
+
"""
|
| 43 |
+
Streams audio chunks to Deepgram and waits for the final result.
|
| 44 |
+
"""
|
| 45 |
+
chunk_size = 160 # Keeping the same chunk size as the reference
|
| 46 |
+
num_chunks = int(np.ceil(len(pcm) / chunk_size))
|
| 47 |
+
print(f"Streaming {self.stream_name} audio to Deepgram...")
|
| 48 |
+
|
| 49 |
+
for i in range(num_chunks):
|
| 50 |
+
start_idx = i * chunk_size
|
| 51 |
+
end_idx = min((i + 1) * chunk_size, len(pcm))
|
| 52 |
+
chunk = pcm[start_idx:end_idx]
|
| 53 |
+
self.process_chunk(chunk)
|
| 54 |
+
|
| 55 |
+
# Signal the end of the stream to Deepgram
|
| 56 |
+
self.close()
|
| 57 |
+
|
| 58 |
+
# Wait for the 'finished' signal (Metadata) from the receive loop
|
| 59 |
+
self.finished_event.wait()
|
| 60 |
+
|
| 61 |
+
# Ensure connection is closed and return final text
|
| 62 |
+
self._ensure_closed()
|
| 63 |
+
with self.lock:
|
| 64 |
+
return self.render_tokens(self.final_tokens, [])
|
| 65 |
+
|
| 66 |
+
def get_config(self, fs_hz: int) -> dict:
|
| 67 |
+
"""
|
| 68 |
+
Returns parameters for the Deepgram V1 URL query string.
|
| 69 |
+
"""
|
| 70 |
+
assert fs_hz == 16000, "Only 16 kHz audio is supported."
|
| 71 |
+
|
| 72 |
+
return {
|
| 73 |
+
"model": "nova-3", # Recommended general model
|
| 74 |
+
"encoding": "linear16", # Corresponds to pcm_s16le
|
| 75 |
+
"sample_rate": 16000,
|
| 76 |
+
"channels": 1,
|
| 77 |
+
"smart_format": "true", # handling punctuation/formatting
|
| 78 |
+
"interim_results": "true", # required for non-final updates
|
| 79 |
+
"endpointing": "500", # ms silence to trigger finalization
|
| 80 |
+
}
|
| 81 |
+
|
| 82 |
+
def process_chunk(self, chunk: np.ndarray) -> None:
|
| 83 |
+
"""
|
| 84 |
+
Converts float32 numpy array to int16 bytes and sends to WebSocket.
|
| 85 |
+
"""
|
| 86 |
+
chunk = np.clip(chunk, -1.0, 1.0)
|
| 87 |
+
chunk_int16 = (chunk * 32767).astype(np.int16)
|
| 88 |
+
if len(chunk_int16) > 0:
|
| 89 |
+
try:
|
| 90 |
+
self.ws.send(chunk_int16.tobytes())
|
| 91 |
+
except Exception:
|
| 92 |
+
pass
|
| 93 |
+
|
| 94 |
+
def render_tokens(
|
| 95 |
+
self, final_tokens: list[dict], non_final_tokens: list[dict]
|
| 96 |
+
) -> str:
|
| 97 |
+
"""
|
| 98 |
+
Renders the list of token dicts into a string.
|
| 99 |
+
Matches Soniox logic: treats certain tokens as punctuation triggers.
|
| 100 |
+
"""
|
| 101 |
+
text_parts = []
|
| 102 |
+
for token in final_tokens + non_final_tokens:
|
| 103 |
+
text = token["text"]
|
| 104 |
+
text_parts.append(text)
|
| 105 |
+
# Add newline if the text chunk looks like end-of-sentence punctuation
|
| 106 |
+
# Note: Deepgram 'smart_format' usually attaches punctuation to the word.
|
| 107 |
+
if text.strip() in [".", "?", "!"]:
|
| 108 |
+
text_parts.append("\n")
|
| 109 |
+
return "".join(text_parts)
|
| 110 |
+
|
| 111 |
+
def _receive_loop(self):
|
| 112 |
+
"""
|
| 113 |
+
Background loop to handle incoming JSON messages from Deepgram.
|
| 114 |
+
"""
|
| 115 |
+
try:
|
| 116 |
+
while True:
|
| 117 |
+
message = self.ws.recv()
|
| 118 |
+
res = json.loads(message)
|
| 119 |
+
|
| 120 |
+
# Check for metadata indicating stream end
|
| 121 |
+
if res.get("type") == "Metadata":
|
| 122 |
+
self.finished_event.set()
|
| 123 |
+
break
|
| 124 |
+
|
| 125 |
+
# Deepgram error handling
|
| 126 |
+
if "error" in res:
|
| 127 |
+
print(f"Deepgram Error: {res['error']}")
|
| 128 |
+
break
|
| 129 |
+
|
| 130 |
+
# Process Transcripts
|
| 131 |
+
# Deepgram V1 structure: result -> channel -> alternatives -> [0] -> transcript
|
| 132 |
+
if "channel" in res:
|
| 133 |
+
is_final = res.get("is_final", False)
|
| 134 |
+
alternatives = res["channel"].get("alternatives", [])
|
| 135 |
+
|
| 136 |
+
if alternatives:
|
| 137 |
+
transcript = alternatives[0].get("transcript", "")
|
| 138 |
+
|
| 139 |
+
if transcript:
|
| 140 |
+
# Wrap the transcript in a dict to match the
|
| 141 |
+
# 'render_tokens' expectation of a list[dict]
|
| 142 |
+
token_data = {
|
| 143 |
+
"text": transcript + " ", # Add space for readability
|
| 144 |
+
"is_final": is_final,
|
| 145 |
+
}
|
| 146 |
+
|
| 147 |
+
non_final_tokens = []
|
| 148 |
+
with self.lock:
|
| 149 |
+
if is_final:
|
| 150 |
+
self.final_tokens.append(token_data)
|
| 151 |
+
else:
|
| 152 |
+
non_final_tokens.append(token_data)
|
| 153 |
+
|
| 154 |
+
current_finals = list(self.final_tokens)
|
| 155 |
+
|
| 156 |
+
# Trigger the callback
|
| 157 |
+
text = self.render_tokens(current_finals, non_final_tokens)
|
| 158 |
+
if self.on_update:
|
| 159 |
+
self.on_update(text)
|
| 160 |
+
|
| 161 |
+
except (ConnectionClosedOK, ConnectionClosedError):
|
| 162 |
+
pass
|
| 163 |
+
except Exception as e:
|
| 164 |
+
print(f"Deepgram receive loop error: {e}")
|
| 165 |
+
finally:
|
| 166 |
+
self.finished_event.set()
|
| 167 |
+
|
| 168 |
+
def close(self) -> str:
|
| 169 |
+
"""
|
| 170 |
+
Sends the specific JSON message Deepgram expects to close the stream.
|
| 171 |
+
"""
|
| 172 |
+
if hasattr(self, "ws"):
|
| 173 |
+
try:
|
| 174 |
+
# Deepgram V1 expects this specific JSON to close the stream
|
| 175 |
+
self.ws.send(json.dumps({"type": "CloseStream"}))
|
| 176 |
+
except Exception:
|
| 177 |
+
pass
|
| 178 |
+
return self.render_tokens(self.final_tokens, [])
|
| 179 |
+
|
| 180 |
+
def _ensure_closed(self) -> None:
|
| 181 |
+
"""
|
| 182 |
+
Physical socket closure and thread cleanup.
|
| 183 |
+
"""
|
| 184 |
+
if hasattr(self, "ws"):
|
| 185 |
+
try:
|
| 186 |
+
self.ws.close()
|
| 187 |
+
except Exception:
|
| 188 |
+
pass
|
| 189 |
+
|
| 190 |
+
if (
|
| 191 |
+
hasattr(self, "thread")
|
| 192 |
+
and self.thread.is_alive()
|
| 193 |
+
and self.thread != threading.current_thread()
|
| 194 |
+
):
|
| 195 |
+
self.thread.join(timeout=1.0)
|
stt_streamers/soniox_streamer.py
ADDED
|
@@ -0,0 +1,141 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import json
|
| 2 |
+
import os
|
| 3 |
+
import threading
|
| 4 |
+
import numpy as np
|
| 5 |
+
from websockets.sync.client import connect
|
| 6 |
+
from websockets.exceptions import ConnectionClosedOK
|
| 7 |
+
|
| 8 |
+
SONIOX_WEBSOCKET_URL = "wss://stt-rt.soniox.com/transcribe-websocket"
|
| 9 |
+
|
| 10 |
+
|
| 11 |
+
class SonioxStreamer:
|
| 12 |
+
def __init__(self, fs_hz: int, stream_name: str, on_update=None) -> None:
|
| 13 |
+
api_key = os.environ.get("SONIOX_API_KEY")
|
| 14 |
+
if not api_key:
|
| 15 |
+
raise RuntimeError("Missing SONIOX_API_KEY.")
|
| 16 |
+
self.stream_name = stream_name
|
| 17 |
+
self.api_name = "Soniox RT"
|
| 18 |
+
self.on_update = on_update
|
| 19 |
+
self.final_tokens: list[dict] = []
|
| 20 |
+
self.lock = threading.Lock()
|
| 21 |
+
self.finished_event = threading.Event()
|
| 22 |
+
config = self.get_config(api_key, fs_hz)
|
| 23 |
+
print(f"Connecting {stream_name} to Soniox...")
|
| 24 |
+
self.ws = connect(SONIOX_WEBSOCKET_URL)
|
| 25 |
+
self.ws.send(json.dumps(config))
|
| 26 |
+
|
| 27 |
+
self.thread = threading.Thread(target=self._receive_loop, daemon=True)
|
| 28 |
+
self.thread.start()
|
| 29 |
+
|
| 30 |
+
def stream_array(self, pcm: np.ndarray, fs_hz: int) -> str:
|
| 31 |
+
chunk_size = 160
|
| 32 |
+
num_chunks = int(np.ceil(len(pcm) / chunk_size))
|
| 33 |
+
print(f"Streaming {self.stream_name} audio to Soniox...")
|
| 34 |
+
for i in range(num_chunks):
|
| 35 |
+
start_idx = i * chunk_size
|
| 36 |
+
end_idx = min((i + 1) * chunk_size, len(pcm))
|
| 37 |
+
chunk = pcm[start_idx:end_idx]
|
| 38 |
+
self.process_chunk(chunk)
|
| 39 |
+
try:
|
| 40 |
+
self.ws.send("")
|
| 41 |
+
except Exception:
|
| 42 |
+
pass
|
| 43 |
+
|
| 44 |
+
# 3. Wait for the 'finished' message from the receive loop
|
| 45 |
+
self.finished_event.wait()
|
| 46 |
+
self.close()
|
| 47 |
+
with self.lock:
|
| 48 |
+
return self.render_tokens(self.final_tokens, [])
|
| 49 |
+
|
| 50 |
+
def get_config(self, api_key: str, fs_hz: int) -> dict:
|
| 51 |
+
config = {
|
| 52 |
+
"api_key": api_key,
|
| 53 |
+
"model": "stt-rt-v3",
|
| 54 |
+
"language_hints": ["en", "de"],
|
| 55 |
+
"language_hints_strict": True,
|
| 56 |
+
"enable_language_identification": True,
|
| 57 |
+
"enable_speaker_diarization": False,
|
| 58 |
+
"enable_endpoint_detection": True,
|
| 59 |
+
}
|
| 60 |
+
assert fs_hz == 16000, "Only 16 kHz audio is supported."
|
| 61 |
+
config["audio_format"] = "pcm_s16le"
|
| 62 |
+
config["sample_rate"] = 16000
|
| 63 |
+
config["num_channels"] = 1
|
| 64 |
+
return config
|
| 65 |
+
|
| 66 |
+
def process_chunk(self, chunk: np.ndarray) -> None:
|
| 67 |
+
chunk = np.clip(chunk, -1.0, 1.0)
|
| 68 |
+
chunk_int16 = (chunk * 32767).astype(np.int16)
|
| 69 |
+
if len(chunk_int16) > 0:
|
| 70 |
+
try:
|
| 71 |
+
self.ws.send(chunk_int16.tobytes())
|
| 72 |
+
except Exception:
|
| 73 |
+
pass
|
| 74 |
+
|
| 75 |
+
def render_tokens(
|
| 76 |
+
self, final_tokens: list[dict], non_final_tokens: list[dict]
|
| 77 |
+
) -> str:
|
| 78 |
+
text_parts = []
|
| 79 |
+
for token in final_tokens + non_final_tokens:
|
| 80 |
+
text = token["text"]
|
| 81 |
+
text_parts.append(text)
|
| 82 |
+
if text.strip() in [".", "?", "!"]:
|
| 83 |
+
text_parts.append("\n")
|
| 84 |
+
return "".join(text_parts)
|
| 85 |
+
|
| 86 |
+
def _receive_loop(self):
|
| 87 |
+
try:
|
| 88 |
+
while True:
|
| 89 |
+
message = self.ws.recv()
|
| 90 |
+
res = json.loads(message)
|
| 91 |
+
|
| 92 |
+
if res.get("error_code") is not None:
|
| 93 |
+
break
|
| 94 |
+
|
| 95 |
+
non_final_tokens: list[dict] = []
|
| 96 |
+
|
| 97 |
+
with self.lock:
|
| 98 |
+
for token in res.get("tokens", []):
|
| 99 |
+
if token.get("text"):
|
| 100 |
+
if token.get("is_final"):
|
| 101 |
+
self.final_tokens.append(token)
|
| 102 |
+
else:
|
| 103 |
+
non_final_tokens.append(token)
|
| 104 |
+
|
| 105 |
+
current_finals = list(self.final_tokens)
|
| 106 |
+
|
| 107 |
+
text = self.render_tokens(current_finals, non_final_tokens)
|
| 108 |
+
|
| 109 |
+
if self.on_update:
|
| 110 |
+
self.on_update(text)
|
| 111 |
+
|
| 112 |
+
if res.get("finished"):
|
| 113 |
+
# Signal stream_array to stop waiting
|
| 114 |
+
self.finished_event.set()
|
| 115 |
+
break
|
| 116 |
+
|
| 117 |
+
except ConnectionClosedOK:
|
| 118 |
+
pass
|
| 119 |
+
except Exception:
|
| 120 |
+
pass
|
| 121 |
+
finally:
|
| 122 |
+
self.finished_event.set()
|
| 123 |
+
|
| 124 |
+
def close(self) -> str:
|
| 125 |
+
"""Closes the connection."""
|
| 126 |
+
if hasattr(self, "ws"):
|
| 127 |
+
try:
|
| 128 |
+
self.ws.send("")
|
| 129 |
+
except Exception:
|
| 130 |
+
pass
|
| 131 |
+
self.ws.close()
|
| 132 |
+
|
| 133 |
+
if (
|
| 134 |
+
hasattr(self, "thread")
|
| 135 |
+
and self.thread.is_alive()
|
| 136 |
+
and self.thread != threading.current_thread()
|
| 137 |
+
):
|
| 138 |
+
self.thread.join(timeout=1.0)
|
| 139 |
+
|
| 140 |
+
with self.lock:
|
| 141 |
+
return self.render_tokens(self.final_tokens, [])
|
transcribe.py
ADDED
|
@@ -0,0 +1,51 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from pathlib import Path
|
| 2 |
+
from regex import D
|
| 3 |
+
import resampy
|
| 4 |
+
import soundfile as sf
|
| 5 |
+
from stt_streamers import DeepgramStreamer, SonioxStreamer
|
| 6 |
+
from constants import DEFAULT_SR
|
| 7 |
+
|
| 8 |
+
def transcribe_file(audio_file_path: str, streamer_type: str = "deepgram", mode: str = "RAW"):
|
| 9 |
+
"""
|
| 10 |
+
Transcribe an audio file using the specified STT streamer.
|
| 11 |
+
|
| 12 |
+
Args:
|
| 13 |
+
streamer_type (str): "soniox" or "deepgram"
|
| 14 |
+
audio_file_path (str): Path to WAV file
|
| 15 |
+
mode (str): Optional label for streamer instance ("RAW", "ENHANCED", etc.)
|
| 16 |
+
|
| 17 |
+
Returns:
|
| 18 |
+
str: Transcript text
|
| 19 |
+
"""
|
| 20 |
+
|
| 21 |
+
# Validate file
|
| 22 |
+
audio_path = Path(audio_file_path)
|
| 23 |
+
if not audio_path.exists():
|
| 24 |
+
raise FileNotFoundError(f"Audio file not found: {audio_file_path}")
|
| 25 |
+
|
| 26 |
+
# Load audio
|
| 27 |
+
pcm, fs_hz = sf.read(audio_file_path, dtype="float32")
|
| 28 |
+
if fs_hz != DEFAULT_SR:
|
| 29 |
+
pcm = resampy.resample(pcm, fs_hz, DEFAULT_SR)
|
| 30 |
+
fs_hz = DEFAULT_SR
|
| 31 |
+
|
| 32 |
+
# Select streamer
|
| 33 |
+
api_map = {
|
| 34 |
+
"soniox": SonioxStreamer,
|
| 35 |
+
"deepgram": DeepgramStreamer,
|
| 36 |
+
}
|
| 37 |
+
|
| 38 |
+
streamer_key = streamer_type.lower()
|
| 39 |
+
if streamer_key not in api_map:
|
| 40 |
+
raise ValueError(
|
| 41 |
+
f"Invalid streamer_type '{streamer_type}'. "
|
| 42 |
+
f"Choose from: {', '.join(api_map.keys())}"
|
| 43 |
+
)
|
| 44 |
+
|
| 45 |
+
StreamerClass = api_map[streamer_key]
|
| 46 |
+
streamer = StreamerClass(fs_hz, mode)
|
| 47 |
+
|
| 48 |
+
# Transcribe
|
| 49 |
+
transcript = streamer.stream_array(pcm, fs_hz)
|
| 50 |
+
|
| 51 |
+
return transcript
|
word_error_rate.py
ADDED
|
@@ -0,0 +1,78 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import numpy as np
|
| 2 |
+
from whisper_normalizer.english import EnglishTextNormalizer
|
| 3 |
+
|
| 4 |
+
def tokenize_text(text: str) -> list[str]:
|
| 5 |
+
"""Normalize + tokenize into words."""
|
| 6 |
+
normalizer = EnglishTextNormalizer()
|
| 7 |
+
normalized = normalizer(text)
|
| 8 |
+
# simple whitespace word tokens; adjust if you need punctuation handling
|
| 9 |
+
return normalized.split()
|
| 10 |
+
|
| 11 |
+
def compute_metric(y_true: str, y_pred: str, detailed: bool = False)-> float | tuple[float, float, float, float]:
|
| 12 |
+
"""
|
| 13 |
+
Returns:
|
| 14 |
+
- if detailed=False: wer
|
| 15 |
+
- if detailed=True: (wer, ins_rate, del_rate, sub_rate)
|
| 16 |
+
"""
|
| 17 |
+
ref = tokenize_text(y_true)
|
| 18 |
+
hyp = tokenize_text(y_pred)
|
| 19 |
+
|
| 20 |
+
total_err, ins, dels, subs = levenshtein_word_errors(ref, hyp)
|
| 21 |
+
ref_len = max(1, len(ref))
|
| 22 |
+
|
| 23 |
+
if detailed:
|
| 24 |
+
return (
|
| 25 |
+
total_err / ref_len,
|
| 26 |
+
ins / ref_len,
|
| 27 |
+
dels / ref_len,
|
| 28 |
+
subs / ref_len,
|
| 29 |
+
)
|
| 30 |
+
return total_err / ref_len
|
| 31 |
+
|
| 32 |
+
|
| 33 |
+
def levenshtein_word_errors(reference: list[str], prediction: list[str]):
|
| 34 |
+
"""Edit distance with insertion/deletion/substitution breakdown."""
|
| 35 |
+
n, m = len(reference), len(prediction)
|
| 36 |
+
|
| 37 |
+
# distance + counts
|
| 38 |
+
d = np.zeros((n + 1, m + 1), dtype=np.int32)
|
| 39 |
+
ins = np.zeros((n + 1, m + 1), dtype=np.int32)
|
| 40 |
+
dels = np.zeros((n + 1, m + 1), dtype=np.int32)
|
| 41 |
+
subs = np.zeros((n + 1, m + 1), dtype=np.int32)
|
| 42 |
+
|
| 43 |
+
# init borders
|
| 44 |
+
for j in range(1, m + 1):
|
| 45 |
+
d[0, j] = j
|
| 46 |
+
ins[0, j] = j
|
| 47 |
+
for i in range(1, n + 1):
|
| 48 |
+
d[i, 0] = i
|
| 49 |
+
dels[i, 0] = i
|
| 50 |
+
|
| 51 |
+
# DP
|
| 52 |
+
for i in range(1, n + 1):
|
| 53 |
+
for j in range(1, m + 1):
|
| 54 |
+
is_sub = 0 if reference[i - 1] == prediction[j - 1] else 1
|
| 55 |
+
|
| 56 |
+
# candidates: delete, insert, substitute/match
|
| 57 |
+
del_cost = d[i - 1, j] + 1
|
| 58 |
+
ins_cost = d[i, j - 1] + 1
|
| 59 |
+
sub_cost = d[i - 1, j - 1] + is_sub
|
| 60 |
+
|
| 61 |
+
# choose best (tie-breaking: prefer sub/match, then del, then ins)
|
| 62 |
+
if sub_cost <= del_cost and sub_cost <= ins_cost:
|
| 63 |
+
d[i, j] = sub_cost
|
| 64 |
+
ins[i, j] = ins[i - 1, j - 1]
|
| 65 |
+
dels[i, j] = dels[i - 1, j - 1]
|
| 66 |
+
subs[i, j] = subs[i - 1, j - 1] + is_sub
|
| 67 |
+
elif del_cost <= ins_cost:
|
| 68 |
+
d[i, j] = del_cost
|
| 69 |
+
ins[i, j] = ins[i - 1, j]
|
| 70 |
+
dels[i, j] = dels[i - 1, j] + 1
|
| 71 |
+
subs[i, j] = subs[i - 1, j]
|
| 72 |
+
else:
|
| 73 |
+
d[i, j] = ins_cost
|
| 74 |
+
ins[i, j] = ins[i, j - 1] + 1
|
| 75 |
+
dels[i, j] = dels[i, j - 1]
|
| 76 |
+
subs[i, j] = subs[i, j - 1]
|
| 77 |
+
|
| 78 |
+
return int(d[n, m]), int(ins[n, m]), int(dels[n, m]), int(subs[n, m])
|