tec / app.py
wq2012's picture
Publish interactive Textual Echo Cancellation (TEC) + ASR demo Space
00dd564 verified
Raw History Blame Contribute Delete
9.44 kB
#!/usr/bin/env python3
"""Hugging Face Space Gradio demo for Textual Echo Cancellation (TEC)."""
import os
import tempfile
from typing import Dict, Tuple
os.environ.setdefault("OMP_NUM_THREADS", "4")
os.environ.setdefault("TF_CPP_MIN_LOG_LEVEL", "2")
from faster_whisper import WhisperModel # noqa: E402
import gradio as gr # noqa: E402
from huggingface_hub import snapshot_download # noqa: E402
from tec import configs # noqa: E402
from tec import data_prep # noqa: E402
from tec import inference # noqa: E402
MODEL_REPOS: Dict[str, Tuple[str, type]] = {
"TecSingleInterfering (wq2012/tec_single_interfering)": (
"wq2012/tec_single_interfering",
configs.TecSingleInterfering,
),
"TecMultiInterfering (wq2012/tec_multi_interfering)": (
"wq2012/tec_multi_interfering",
configs.TecMultiInterfering,
),
}
_TEC_RUNNERS: Dict[str, inference.TecInferenceRunner] = {}
_ASR_MODEL = None
def get_tec_runner(model_choice: str) -> inference.TecInferenceRunner:
"""Loads and caches the selected pretrained TEC model from Hugging Face."""
if model_choice not in MODEL_REPOS:
model_choice = list(MODEL_REPOS.keys())[0]
if model_choice not in _TEC_RUNNERS:
repo_id, cfg_cls = MODEL_REPOS[model_choice]
model_dir = snapshot_download(repo_id=repo_id)
ckpt_path = os.path.join(model_dir, "best.ckpt")
_TEC_RUNNERS[model_choice] = inference.TecInferenceRunner(
model_config_cls=cfg_cls,
checkpoint_path=ckpt_path,
decode_max_output_frames=None,
)
return _TEC_RUNNERS[model_choice]
def get_asr_model() -> WhisperModel:
"""Loads and caches the CPU Whisper ASR model."""
global _ASR_MODEL
if _ASR_MODEL is None:
_ASR_MODEL = WhisperModel(
"base.en", device="cpu", compute_type="int8", cpu_threads=4
)
return _ASR_MODEL
def transcribe_wav(wav_path: str) -> str:
"""Runs ASR on a WAV file and returns the recognized transcript."""
asr = get_asr_model()
segments, _ = asr.transcribe(wav_path, beam_size=5)
text = " ".join(seg.text.strip() for seg in segments).strip()
return text if text else "(no speech recognized)"
def run_tec_and_asr(
audio_path: str,
interfering_text: str,
model_choice: str,
) -> Tuple[str, str, str, str]:
"""Runs TEC enhancement and compares ASR before and after TEC."""
if not audio_path:
raise gr.Error("Please upload an audio file or select an example below.")
interfering_text = (interfering_text or "").strip()
# 1. Read and normalize uploaded audio to 24 kHz mono WAV.
mixed_wav, sr = data_prep.read_wav_file(audio_path, target_sample_rate=24000)
tmp_dir = tempfile.mkdtemp(prefix="tec_space_")
norm_input_wav = os.path.join(tmp_dir, "uploaded_input_24k.wav")
enhanced_wav = os.path.join(tmp_dir, "tec_enhanced_24k.wav")
data_prep.write_wav_file(norm_input_wav, mixed_wav, sample_rate=sr)
# 2. Run ASR on the user uploaded audio file (Before TEC).
asr_before = transcribe_wav(norm_input_wav)
# 3. Run Textual Echo Cancellation (TEC) using the interfering text.
runner = get_tec_runner(model_choice)
outputs = runner.predict(
mixed_waveforms=[mixed_wav],
interfering_texts=[interfering_text],
)
pred_wav = outputs["predicted_waveforms"][0]
pred_len = int(outputs["predicted_waveform_lengths"][0])
if pred_len > 0:
pred_wav = pred_wav[:pred_len]
data_prep.write_wav_file(enhanced_wav, pred_wav, sample_rate=sr)
# 4. Run ASR on the textual echo cancelled file (After TEC).
asr_after = transcribe_wav(enhanced_wav)
# 5. Format comparison summary.
side_bytes = len(interfering_text.encode("utf-8"))
audio_bytes = len(mixed_wav) * 2
summary_md = (
"### ASR Recognition Comparison\n\n"
"| Approach | Audio Source | Side Input Payload | ASR Recognition Result |\n"
"| :--- | :--- | :--- | :--- |\n"
f"| **1. Baseline (No Echo Cancellation)** | User Uploaded Audio | `0 B` | `{asr_before}` |\n"
f"| **2. Textual Echo Cancellation (TEC)** | TEC Enhanced Audio | `{side_bytes} B` (`{side_bytes / 1000.0:.3f} KB` text vs. `{audio_bytes / 1000.0:.1f} KB` audio) | **`{asr_after}`** |\n\n"
f"**Interfering TTS Playback Text Cancelled:** *\"{interfering_text or '(empty)'}\"*"
)
return enhanced_wav, asr_before, asr_after, summary_md
EXAMPLES = [
[
os.path.join(os.path.dirname(__file__), "examples", "sample_1_mixed.wav"),
"a table showing the figures for the year ending Michaelmas eighteen oh two.",
"TecSingleInterfering (wq2012/tec_single_interfering)",
],
[
os.path.join(os.path.dirname(__file__), "examples", "sample_2_mixed.wav"),
"and to approve or disapprove the public policy written into these laws.",
"TecSingleInterfering (wq2012/tec_single_interfering)",
],
[
os.path.join(os.path.dirname(__file__), "examples", "sample_3_mixed.wav"),
"was living in the city while the walls were still standing, though in a ruinous condition.",
"TecSingleInterfering (wq2012/tec_single_interfering)",
],
[
os.path.join(os.path.dirname(__file__), "examples", "sample_4_mixed.wav"),
"a subsequent bullet, which was lethal, shattered the right side of his skull.",
"TecSingleInterfering (wq2012/tec_single_interfering)",
],
]
DESCRIPTION_MD = """
# πŸŽ™οΈ Textual Echo Cancellation (TEC) Demo
When a user speaks to a smart speaker or voice assistant while the device is playing back a Text-to-Speech (TTS) response, the microphone records an overlapping mixture of the **user's speech** and the **device's reverberant TTS playback echo**.
**Textual Echo Cancellation (TEC)** cancels the TTS playback echo using only the **source text of the TTS playback** (< 0.1 KB side input) via a multi-source attention sequence-to-sequence neural network, dramatically improving downstream Automatic Speech Recognition (ASR) without streaming reference audio waveforms.
- πŸ“„ **Paper**: [Textual Echo Cancellation (IEEE SLT 2021 / arXiv:2008.06006)](https://arxiv.org/abs/2008.06006)
- πŸ’» **GitHub**: [https://github.com/wq2012/tec](https://github.com/wq2012/tec) | πŸ“¦ **PyPI**: [`pip install textual-echo-cancellation`](https://pypi.org/project/textual-echo-cancellation/)
- πŸ€— **Pretrained Models**: [`wq2012/tec_single_interfering`](https://huggingface.co/wq2012/tec_single_interfering) | [`wq2012/tec_multi_interfering`](https://huggingface.co/wq2012/tec_multi_interfering)
"""
def build_demo() -> gr.Blocks:
"""Constructs the Gradio Blocks interface."""
with gr.Blocks(title="Textual Echo Cancellation (TEC) Demo") as demo:
gr.Markdown(DESCRIPTION_MD)
with gr.Row():
with gr.Column(scale=1):
gr.Markdown("### Inputs")
input_audio = gr.Audio(
sources=["upload", "microphone"],
type="filepath",
label="1. Upload Audio File (Microphone Mixture: User Speech + TTS Echo)",
)
input_text = gr.Textbox(
label="2. Interfering TTS Playback Text (Source Text to Cancel)",
placeholder=(
"Enter the source text spoken by the interfering TTS voice..."
),
lines=3,
)
model_dropdown = gr.Dropdown(
choices=list(MODEL_REPOS.keys()),
value=list(MODEL_REPOS.keys())[0],
label="Pretrained TEC Model",
)
run_btn = gr.Button(
"Run Textual Echo Cancellation & Compare ASR",
variant="primary",
)
with gr.Column(scale=1):
gr.Markdown("### Outputs")
output_audio = gr.Audio(
type="filepath",
label="Textual Echo Cancelled Audio (Enhanced User Speech)",
)
asr_before_box = gr.Textbox(
label="1. ASR Recognition on Uploaded Audio (Before TEC)",
lines=2,
interactive=False,
)
asr_after_box = gr.Textbox(
label="2. ASR Recognition on Textual Echo Cancelled Audio (After TEC)",
lines=2,
interactive=False,
)
summary_markdown = gr.Markdown(
label="Comparison Summary",
)
run_btn.click(
fn=run_tec_and_asr,
inputs=[input_audio, input_text, model_dropdown],
outputs=[output_audio, asr_before_box, asr_after_box, summary_markdown],
)
gr.Markdown(
"### Built-in Examples (LibriTTS User Query + Reverberant LJ Speech TTS Echo at 0 dB SNR)\n"
"*Click any row below to populate the inputs, then click **Run Textual Echo Cancellation & Compare ASR**:*\n"
"- **Example 1** β€” Ground-Truth User Query: *\"I can't see you at all, anywhere.\"*\n"
"- **Example 2** β€” Ground-Truth User Query: *\"Because the thing had been such a scare?\"*\n"
"- **Example 3** β€” Ground-Truth User Query: *\"I must know about you.\"*\n"
"- **Example 4** β€” Ground-Truth User Query: *\"I should much prefer that you called in the aid of the police.\"*"
)
gr.Examples(
examples=EXAMPLES,
inputs=[input_audio, input_text, model_dropdown],
outputs=[output_audio, asr_before_box, asr_after_box, summary_markdown],
fn=run_tec_and_asr,
cache_examples=False,
)
return demo
if __name__ == "__main__":
get_asr_model()
get_tec_runner(list(MODEL_REPOS.keys())[0])
demo_app = build_demo()
demo_app.launch()