Download app.py from wq2012/tec: direct link, hf CLI and curl.
- Browser
- Download file 9.44 kB
-
https://huggingface.co/spaces/wq2012/tec/resolve/main/app.py
- Command line
-
hf download hf://spaces/wq2012/tec/app.py
-
curl -L -o app.py https://huggingface.co/spaces/wq2012/tec/resolve/main/app.py
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() | |