#!/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()