import io import math import tempfile from dataclasses import dataclass from pathlib import Path from typing import Dict, Optional, Tuple import gradio as gr import librosa import matplotlib.pyplot as plt from matplotlib.colors import LinearSegmentedColormap import numpy as np import onnxruntime as ort import soundfile as sf from PIL import Image # ----------------------------- # Configuration # ----------------------------- MAX_SECONDS = 10.0 ONNX_DIR = Path("./onnx") ATTN_LIMIT_NOISY_FRAME_OFFSET = 4 # Spectrogram palette mirrors the UI: ink/slate shadows, orange energy, # and warm cream highlights. Keeping it here makes the visual theme explicit # and gives both before/after plots the same stable color scale. SPECTROGRAM_CMAP = LinearSegmentedColormap.from_list( "dpdfnet_ui", [ "#17191c", # ink "#34383d", # slate "#6b321d", # burnt umber "#a64a16", # deep accent "#d76522", # UI orange "#f1cdb6", # accent soft "#fffaf7", # warm highlight ], N=256, ) SPECTROGRAM_DB_MIN = -80.0 SPECTROGRAM_DB_MAX = 0.0 @dataclass(frozen=True) class ModelSpec: name: str sr: int onnx_path: str # ----------------------------- # Model discovery and metadata # ----------------------------- def _infer_model_meta(model_name: str) -> int: normalized = model_name.lower().replace("-", "_") if "48khz" in normalized or "48k" in normalized or "48hr" in normalized: return 48000 if "8khz" in normalized or normalized.endswith("_8k") or "_8k_" in normalized: return 8000 # Fallback for unknown 16 kHz DPDFNet variants return 16000 def _display_label(spec: ModelSpec) -> str: khz = int(spec.sr // 1000) return f"{spec.name} ({khz} kHz)" def discover_model_presets() -> Dict[str, ModelSpec]: ordered_names = [ "baseline", "dpdfnet2", "dpdfnet4", "dpdfnet8", "dpdfnet2_8khz", "dpdfnet8_8khz", "dpdfnet2_48khz_hr", "dpdfnet8_48khz_hr", ] found_paths = {p.stem: p for p in ONNX_DIR.glob("*.onnx") if p.is_file()} presets: Dict[str, ModelSpec] = {} for name in ordered_names: p = found_paths.get(name) if p is None: continue sr = _infer_model_meta(name) spec = ModelSpec( name=name, sr=sr, onnx_path=str(p), ) presets[_display_label(spec)] = spec # Include any additional ONNX files not in the canonical order list. for name, p in sorted(found_paths.items()): if name in ordered_names: continue sr = _infer_model_meta(name) spec = ModelSpec( name=name, sr=sr, onnx_path=str(p), ) presets[_display_label(spec)] = spec return presets MODEL_PRESETS = discover_model_presets() def _model_choice_label(spec: ModelSpec) -> str: """Return an icon-free model name; sample rate is selected separately.""" friendly_names = { "baseline": "Baseline", "dpdfnet2": "DPDFNet2", "dpdfnet4": "DPDFNet4", "dpdfnet8": "DPDFNet8", "dpdfnet2_8khz": "DPDFNet2", "dpdfnet8_8khz": "DPDFNet8", "dpdfnet2_48khz_hr": "DPDFNet2", "dpdfnet8_48khz_hr": "DPDFNet8", } return friendly_names.get( spec.name, spec.name.replace("_", " ").replace("-", " ").title(), ) def _build_model_choices_by_sr() -> Dict[int, list]: """Group discovered ONNX models by sample rate in a stable order.""" preferred_order = { 8000: [ "dpdfnet2_8khz", "dpdfnet8_8khz", ], 16000: [ "baseline", "dpdfnet2", "dpdfnet4", "dpdfnet8", ], 48000: [ "dpdfnet2_48khz_hr", "dpdfnet8_48khz_hr", ], } grouped: Dict[int, list] = {} for sr in sorted({spec.sr for spec in MODEL_PRESETS.values()}): keys_at_sr = [ key for key, spec in MODEL_PRESETS.items() if spec.sr == sr ] key_by_name = { MODEL_PRESETS[key].name: key for key in keys_at_sr } ordered_keys = [ key_by_name[name] for name in preferred_order.get(sr, []) if name in key_by_name ] ordered_keys.extend( key for key in sorted( keys_at_sr, key=lambda item: MODEL_PRESETS[item].name, ) if key not in ordered_keys ) grouped[sr] = [ (_model_choice_label(MODEL_PRESETS[key]), key) for key in ordered_keys ] return grouped MODEL_CHOICES_BY_SR = _build_model_choices_by_sr() # Keep the familiar rates first, while still supporting any additional models. AVAILABLE_SAMPLE_RATES = [ sr for sr in (8000, 16000, 48000) if MODEL_CHOICES_BY_SR.get(sr) ] AVAILABLE_SAMPLE_RATES.extend( sr for sr in sorted(MODEL_CHOICES_BY_SR) if sr not in AVAILABLE_SAMPLE_RATES ) SAMPLE_RATE_CHOICES = [ (f"{sr // 1000} kHz", str(sr)) for sr in AVAILABLE_SAMPLE_RATES ] DEFAULT_SAMPLE_RATE = ( 16000 if 16000 in AVAILABLE_SAMPLE_RATES else (AVAILABLE_SAMPLE_RATES[0] if AVAILABLE_SAMPLE_RATES else None) ) def _default_model_key_for_sr(sr: int) -> Optional[str]: """Prefer the DPDFNet2 variant for every sample-rate group.""" preferred_name = { 8000: "dpdfnet2_8khz", 16000: "dpdfnet2", 48000: "dpdfnet2_48khz_hr", }.get(sr) if preferred_name is not None: for key, spec in MODEL_PRESETS.items(): if spec.sr == sr and spec.name == preferred_name: return key for key, spec in MODEL_PRESETS.items(): if spec.sr == sr and spec.name.lower().startswith("dpdfnet2"): return key choices = MODEL_CHOICES_BY_SR.get(sr, []) return choices[0][1] if choices else None DEFAULT_MODEL_CHOICES = ( MODEL_CHOICES_BY_SR.get(DEFAULT_SAMPLE_RATE, []) if DEFAULT_SAMPLE_RATE is not None else [] ) DEFAULT_MODEL_KEY = ( _default_model_key_for_sr(DEFAULT_SAMPLE_RATE) if DEFAULT_SAMPLE_RATE is not None else None ) ATTN_RADIO_CHOICES = [ ("Gentle · 6 dB", "6"), ("Balanced · 12 dB", "12"), ("Strong · 24 dB", "24"), ("Unlimited", "unlimited"), ] DEFAULT_ATTN_CHOICE = "unlimited" # ----------------------------- # ONNX Runtime + frontend cache # ----------------------------- _SESSIONS: Dict[str, ort.InferenceSession] = {} _INIT_STATES: Dict[str, np.ndarray] = {} def resolve_model_path(local_path: str) -> str: p = Path(local_path) if p.exists(): return str(p) raise gr.Error( f"ONNX model not found at: {local_path}. " "Expected local models under ./onnx/." ) def get_ort_session(model_key: str) -> ort.InferenceSession: if model_key in _SESSIONS: return _SESSIONS[model_key] spec = MODEL_PRESETS[model_key] onnx_path = resolve_model_path(spec.onnx_path) options = ort.SessionOptions() options.intra_op_num_threads = 1 options.inter_op_num_threads = 1 sess = ort.InferenceSession( onnx_path, sess_options=options, providers=["CPUExecutionProvider"], ) _SESSIONS[model_key] = sess return sess def _load_initial_state(model_key: str, session: ort.InferenceSession) -> np.ndarray: if model_key in _INIT_STATES: return _INIT_STATES[model_key] if len(session.get_inputs()) < 2: raise gr.Error("Expected streaming ONNX model with two inputs: (spec, state).") meta = session.get_modelmeta().custom_metadata_map try: state_size = int(meta["state_size"]) erb_norm_state_size = int(meta["erb_norm_state_size"]) spec_norm_state_size = int(meta["spec_norm_state_size"]) erb_norm_init = np.array( [float(x) for x in meta["erb_norm_init"].split(",")], dtype=np.float32 ) spec_norm_init = np.array( [float(x) for x in meta["spec_norm_init"].split(",")], dtype=np.float32 ) except KeyError as exc: raise gr.Error( f"ONNX model is missing required metadata key: {exc}. " "Re-export the model to embed state initialisation metadata." ) init_state = np.zeros(state_size, dtype=np.float32) init_state[0:erb_norm_state_size] = erb_norm_init init_state[erb_norm_state_size:erb_norm_state_size + spec_norm_state_size] = spec_norm_init init_state = np.ascontiguousarray(init_state) _INIT_STATES[model_key] = init_state return init_state # ----------------------------- # STFT/iSTFT (module-free) # ----------------------------- def vorbis_window(window_len: int) -> np.ndarray: window_size_h = window_len / 2 indices = np.arange(window_len) sin = np.sin(0.5 * np.pi * (indices + 0.5) / window_size_h) window = np.sin(0.5 * np.pi * sin * sin) return window.astype(np.float32) def _infer_stft_params(model_key: str, session: ort.InferenceSession) -> Tuple[int, int, np.ndarray]: # ONNX spec input is [B, T, F, 2] (or dynamic variants). spec_shape = session.get_inputs()[0].shape freq_bins = spec_shape[-2] if len(spec_shape) >= 2 else None if isinstance(freq_bins, int) and freq_bins > 1: win_len = int((freq_bins - 1) * 2) else: # 20 ms windows for DPDFNet family. sr = MODEL_PRESETS[model_key].sr win_len = int(round(sr * 0.02)) hop = win_len // 2 win = vorbis_window(win_len) return win_len, hop, win def _preprocess_waveform(waveform: np.ndarray, win_len: int, hop: int, win: np.ndarray) -> np.ndarray: audio = np.asarray(waveform, dtype=np.float32).reshape(-1) audio_pad = np.pad(audio, (0, win_len), mode="constant") spec = librosa.stft( y=audio_pad, n_fft=win_len, hop_length=hop, win_length=win_len, window=win, center=True, pad_mode="reflect", ) spec = spec.T.astype(np.complex64, copy=False) # [T, F] spec_ri = np.stack([spec.real, spec.imag], axis=-1).astype(np.float32, copy=False) # [T, F, 2] return np.ascontiguousarray(spec_ri[None, ...], dtype=np.float32) # [1, T, F, 2] def _postprocess_spec(spec_e: np.ndarray, win_len: int, hop: int, win: np.ndarray) -> np.ndarray: spec_c = np.asarray(spec_e[0], dtype=np.float32) # [T, F, 2] spec = (spec_c[..., 0] + 1j * spec_c[..., 1]).T.astype(np.complex64, copy=False) # [F, T] waveform_e = librosa.istft( spec, hop_length=hop, win_length=win_len, window=win, center=True, length=None, ).astype(np.float32, copy=False) return np.concatenate( [waveform_e[win_len * 2 :], np.zeros(win_len * 2, dtype=np.float32)], axis=0, ) def _validate_attn_limit_db(attn_limit_db: Optional[float]) -> Optional[float]: if attn_limit_db is None: return None value = float(attn_limit_db) if np.isnan(value) or value < 0.0: raise gr.Error("Attenuation limit must be zero or greater.") return value def _apply_attn_limit( spec_noisy: np.ndarray, spec_enh: np.ndarray, attn_limit_db: Optional[float], ) -> np.ndarray: value = _validate_attn_limit_db(attn_limit_db) enhanced = np.asarray(spec_enh, dtype=np.float32) if value is None: return enhanced noisy = np.asarray(spec_noisy, dtype=np.float32) if noisy.shape != enhanced.shape: raise gr.Error( "The noisy and enhanced spectra do not have matching shapes." ) aligned_noisy = np.zeros_like(noisy, dtype=np.float32) if noisy.shape[1] > ATTN_LIMIT_NOISY_FRAME_OFFSET: aligned_noisy[:, ATTN_LIMIT_NOISY_FRAME_OFFSET:, :, :] = noisy[ :, :-ATTN_LIMIT_NOISY_FRAME_OFFSET, :, : ] alpha = float(10.0 ** (-value / 20.0)) limited = alpha * aligned_noisy + (1.0 - alpha) * enhanced return np.ascontiguousarray(limited, dtype=np.float32) # ----------------------------- # ONNX inference (non-streaming pre/post, streaming ONNX state loop) # ----------------------------- def enhance_audio_onnx( audio_mono: np.ndarray, model_key: str, attn_limit_db: Optional[float] = None, ) -> np.ndarray: sess = get_ort_session(model_key) inputs = sess.get_inputs() outputs = sess.get_outputs() if len(inputs) < 2 or len(outputs) < 2: raise gr.Error( "Expected streaming ONNX signature with 2 inputs (spec, state) and 2 outputs (spec_e, state_out)." ) in_spec_name = inputs[0].name in_state_name = inputs[1].name out_spec_name = outputs[0].name out_state_name = outputs[1].name waveform = np.asarray(audio_mono, dtype=np.float32).reshape(-1) win_len, hop, win = _infer_stft_params(model_key, sess) spec_r_np = _preprocess_waveform(waveform, win_len=win_len, hop=hop, win=win) state = _load_initial_state(model_key, sess).copy() spec_e_frames = [] num_frames = int(spec_r_np.shape[1]) for t in range(num_frames): spec_t = np.ascontiguousarray(spec_r_np[:, t : t + 1, :, :], dtype=np.float32) spec_e_t, state = sess.run( [out_spec_name, out_state_name], {in_spec_name: spec_t, in_state_name: state}, ) spec_e_frames.append(np.ascontiguousarray(spec_e_t, dtype=np.float32)) if not spec_e_frames: return waveform spec_e_np = np.concatenate(spec_e_frames, axis=1) spec_e_np = _apply_attn_limit(spec_r_np, spec_e_np, attn_limit_db) waveform_e = _postprocess_spec(spec_e_np, win_len=win_len, hop=hop, win=win) return np.asarray(waveform_e, dtype=np.float32).reshape(-1) # ----------------------------- # Audio utilities # ----------------------------- def _load_wav_from_gradio_path(path: str) -> Tuple[np.ndarray, int]: data, sr = sf.read(path, always_2d=True) data = data.astype(np.float32, copy=False) return data, int(sr) def _to_mono(x: np.ndarray) -> Tuple[np.ndarray, int]: if x.ndim == 1: return x.astype(np.float32, copy=False), 1 if x.shape[1] == 1: return x[:, 0], 1 return x.mean(axis=1), int(x.shape[1]) def _resample(y: np.ndarray, sr_in: int, sr_out: int) -> np.ndarray: if sr_in == sr_out: return y return librosa.resample(y, orig_sr=sr_in, target_sr=sr_out).astype(np.float32, copy=False) def _match_length(y: np.ndarray, target_len: int) -> np.ndarray: if len(y) == target_len: return y if len(y) > target_len: return y[:target_len] out = np.zeros((target_len,), dtype=y.dtype) out[: len(y)] = y return out def _save_wav(y: np.ndarray, sr: int, prefix: str) -> str: tmp = tempfile.NamedTemporaryFile(prefix=prefix, suffix=".wav", delete=False) tmp.close() sf.write(tmp.name, y, sr) return tmp.name def _spectrogram_image(y: np.ndarray, sr: int) -> Image.Image: win_length = max(256, int(0.032 * sr)) hop_length = max(64, int(0.008 * sr)) n_fft = 1 << (int(math.ceil(math.log2(win_length)))) S = librosa.stft(y, n_fft=n_fft, hop_length=hop_length, win_length=win_length, center=False) S_db = librosa.amplitude_to_db(np.abs(S) + 1e-10, ref=np.max) fig, ax = plt.subplots(figsize=(8.4, 3.2)) fig.patch.set_facecolor("#17191c") ax.set_facecolor("#17191c") ax.imshow( S_db, origin="lower", aspect="auto", cmap=SPECTROGRAM_CMAP, vmin=SPECTROGRAM_DB_MIN, vmax=SPECTROGRAM_DB_MAX, interpolation="nearest", ) ax.set_axis_off() fig.subplots_adjust(left=0, right=1, top=1, bottom=0) buf = io.BytesIO() fig.savefig(buf, format="png", dpi=160) plt.close(fig) buf.seek(0) return Image.open(buf) # ----------------------------- # Main pipeline # ----------------------------- def _resolve_attn_choice(attn_choice: Optional[str]) -> Optional[float]: if not attn_choice or attn_choice == "unlimited": return None return _validate_attn_limit_db(float(attn_choice)) def run_enhancement( source: str, mic_path: Optional[str], file_path: Optional[str], model_key: str, attn_choice: Optional[str], ): if not MODEL_PRESETS: raise gr.Error("No ONNX models found under ./onnx/. Add models and retry.") chosen_path = mic_path if source == "Microphone" else file_path if not chosen_path: raise gr.Error("Please provide audio either from the microphone or by uploading a file.") x, sr_orig = _load_wav_from_gradio_path(chosen_path) y_mono, n_ch = _to_mono(x) max_samples = int(MAX_SECONDS * sr_orig) was_trimmed = len(y_mono) > max_samples if was_trimmed: y_mono = y_mono[:max_samples] dur = len(y_mono) / float(sr_orig) spec = MODEL_PRESETS[model_key] sr_model = spec.sr attn_limit_db = _resolve_attn_choice(attn_choice) y_model = _resample(y_mono, sr_orig, sr_model) y_enh_model = enhance_audio_onnx(y_model, model_key, attn_limit_db) y_enh = _resample(y_enh_model, sr_model, sr_orig) y_enh = _match_length(y_enh, len(y_mono)) noisy_out = _save_wav(y_mono, sr_orig, prefix="noisy_mono_") enh_out = _save_wav(y_enh, sr_orig, prefix="enhanced_") noisy_img = _spectrogram_image(y_mono, sr_orig) enh_img = _spectrogram_image(y_enh, sr_orig) details = [spec.name, f"{attn_limit_db:g} dB" if attn_limit_db is not None else "unlimited"] if was_trimmed: details.append("first 10s") status = "✓ Enhanced · " + " · ".join(details) return noisy_out, enh_out, noisy_img, enh_img, status def set_source_visibility(source: str): return ( gr.update(visible=(source == "Microphone")), gr.update(visible=(source == "Upload")), ) def set_model_choices(sample_rate: str): """Show only models that match the selected sample rate.""" try: sr = int(sample_rate) except (TypeError, ValueError): sr = DEFAULT_SAMPLE_RATE if sr is None: return gr.update(choices=[], value=None) choices = MODEL_CHOICES_BY_SR.get(sr, []) return gr.update( choices=choices, value=_default_model_key_for_sr(sr), ) # ----------------------------- # UI — Wide guided flow # ----------------------------- THEME = gr.themes.Base( primary_hue="orange", secondary_hue="slate", neutral_hue="slate", font=["Aptos", "Segoe UI Variable", "Segoe UI", "Helvetica Neue", "sans-serif"], ).set( body_background_fill="#f4f4f1", body_background_fill_dark="#f4f4f1", body_text_color="#17191c", body_text_color_dark="#17191c", body_text_color_subdued="#505760", body_text_color_subdued_dark="#505760", background_fill_primary="#ffffff", background_fill_primary_dark="#ffffff", background_fill_secondary="#f8f8f6", background_fill_secondary_dark="#f8f8f6", border_color_primary="#d6d8d3", border_color_primary_dark="#d6d8d3", block_background_fill="#ffffff", block_background_fill_dark="#ffffff", block_border_color="#d6d8d3", block_border_color_dark="#d6d8d3", block_label_background_fill="#ffffff", block_label_background_fill_dark="#ffffff", block_label_text_color="#2f3439", block_label_text_color_dark="#2f3439", block_info_text_color="#505760", block_info_text_color_dark="#505760", input_background_fill="#ffffff", input_background_fill_dark="#ffffff", input_background_fill_focus="#ffffff", input_background_fill_focus_dark="#ffffff", input_border_color="#c8cbc6", input_border_color_dark="#c8cbc6", input_border_color_focus="#c95f20", input_border_color_focus_dark="#c95f20", input_placeholder_color="#68707a", input_placeholder_color_dark="#68707a", checkbox_label_background_fill="#ffffff", checkbox_label_background_fill_dark="#ffffff", checkbox_label_background_fill_hover="#faf7f4", checkbox_label_background_fill_hover_dark="#faf7f4", checkbox_label_background_fill_selected="#fff4ec", checkbox_label_background_fill_selected_dark="#fff4ec", checkbox_label_border_color="#cfd2cd", checkbox_label_border_color_dark="#cfd2cd", checkbox_label_border_color_hover="#b8bcb6", checkbox_label_border_color_hover_dark="#b8bcb6", checkbox_label_border_color_selected="#c95f20", checkbox_label_border_color_selected_dark="#c95f20", checkbox_label_text_color="#24282d", checkbox_label_text_color_dark="#24282d", checkbox_label_text_color_selected="#17191c", checkbox_label_text_color_selected_dark="#17191c", button_primary_background_fill="#d76522", button_primary_background_fill_hover="#bd561a", button_primary_background_fill_dark="#d76522", button_primary_background_fill_hover_dark="#bd561a", button_primary_border_color="#d76522", button_primary_border_color_dark="#d76522", button_primary_text_color="#ffffff", button_primary_text_color_dark="#ffffff", accordion_text_color="#2f3439", accordion_text_color_dark="#2f3439", ) CSS = r""" :root { --ink: #17191c; --muted: #505760; --line: #d6d8d3; --line-strong: #c4c7c1; --paper: #ffffff; --soft: #f4f4f1; --accent: #d76522; --accent-soft: #f1cdb6; --option-soft: #f6eee8; --option-hover: #f3dfd1; } html, body { min-width: 100%; background: var(--soft) !important; color-scheme: light !important; } * { box-sizing: border-box; } .gradio-container { width: min(1460px, calc(100vw - 32px)) !important; max-width: none !important; margin: 0 auto !important; padding: 18px 0 30px !important; font-family: Aptos, "Segoe UI Variable", "Segoe UI", sans-serif !important; color: var(--ink) !important; } .gradio-container > * { width: 100% !important; max-width: none !important; } #app-header { margin: 0 0 14px; text-align: left; } #app-header h1 { margin: 0; color: var(--ink); font-size: clamp(28px, 3vw, 38px); font-weight: 760; line-height: 1.02; letter-spacing: -.04em; } #app-header p { margin: 5px 0 0; color: var(--muted); font-size: 13px; font-weight: 500; } .workspace-shell, .results-row { width: 100% !important; gap: 14px !important; align-items: stretch !important; } .workspace-shell > div, .results-row > div { min-width: 0 !important; } .flow-card, .result-card { padding: 16px !important; border: 1px solid var(--line) !important; border-radius: 14px !important; background: var(--paper) !important; box-shadow: 0 1px 2px rgba(24, 27, 30, .04) !important; } .section-heading { display: flex; align-items: center; gap: 8px; margin-bottom: 12px; color: var(--ink); font-size: 14px; font-weight: 760; } .step-number { display: inline-grid; width: 22px; height: 22px; flex: 0 0 22px; place-items: center; border-radius: 999px; background: #f3dfd1; color: #9d4614; font-size: 11px; font-weight: 800; } .field-title { margin: 12px 0 6px; color: #2d3237; font-size: 11px; font-weight: 760; letter-spacing: .01em; } .field-title:first-of-type { margin-top: 0; } /* Gradio 6 renders Radio choices inside div.wrap. */ .source-tabs .wrap, .option-grid .wrap { display: grid !important; gap: 6px !important; width: 100% !important; min-width: 0 !important; overflow: visible !important; } .source-tabs .wrap { grid-template-columns: repeat(2, minmax(0, 1fr)) !important; } .option-grid .wrap, .model-grid .wrap { grid-template-columns: repeat(2, minmax(0, 1fr)) !important; } .sample-rate-grid .wrap { grid-template-columns: repeat(3, minmax(0, 1fr)) !important; } .attn-grid .wrap { grid-template-columns: repeat(4, minmax(0, 1fr)) !important; gap: 6px !important; } .source-tabs .wrap label, .option-grid .wrap label { position: relative; display: flex !important; align-items: center !important; justify-content: center !important; width: 100% !important; min-width: 0 !important; margin: 0 !important; border: 1px solid #dfcfc3 !important; border-radius: 9px !important; background: var(--option-soft) !important; color: #23272c !important; box-shadow: none !important; cursor: pointer; transition: border-color .14s ease, background-color .14s ease, box-shadow .14s ease; } .source-tabs .wrap label { min-height: 36px; padding: 7px 10px !important; } .option-grid .wrap label { min-height: 42px; padding: 8px 9px !important; } .model-grid .wrap label { min-height: 40px; justify-content: flex-start !important; padding: 8px 10px !important; text-align: left !important; } .attn-grid .wrap label { min-height: 40px; padding: 6px 4px !important; } .source-tabs .wrap label:hover, .option-grid .wrap label:hover { border-color: #c99a79 !important; background: var(--option-hover) !important; } .source-tabs .wrap label.selected, .option-grid .wrap label.selected, .source-tabs .wrap label:has(input:checked), .option-grid .wrap label:has(input:checked) { border-color: var(--accent) !important; background: var(--accent-soft) !important; box-shadow: inset 0 0 0 1px var(--accent) !important; } .source-tabs .wrap label span, .option-grid .wrap label span, .source-tabs .wrap label, .option-grid .wrap label { color: #23272c !important; opacity: 1 !important; font-size: 12px !important; font-weight: 700 !important; line-height: 1.2 !important; } .attn-grid .wrap label, .attn-grid .wrap label span { white-space: nowrap !important; font-size: 10.5px !important; letter-spacing: -.015em !important; text-align: center !important; } /* Gradio may wrap the option text in an inner span. Give that span the full tile width so every attenuation label, including Unlimited, is optically centered rather than centered only within its text width. */ .attn-grid .wrap label > span, .attn-grid .wrap label span { display: flex !important; width: 100% !important; min-width: 0 !important; align-items: center !important; justify-content: center !important; text-align: center !important; } /* Two clear, restrained motion languages ----------------------------------- Model capability: a rotating conic-gradient lives behind an inset surface, so only a short illuminated segment is visible on the true perimeter. Noise reduction: one centered dotted sine-wave drifts horizontally behind the selected option. Stronger reduction makes it calmer and fainter. */ .model-grid .wrap label, .attn-grid .wrap label { position: relative; overflow: hidden !important; isolation: isolate; } /* -------------------------------------------------------------------------- MODEL CAPABILITY — text shimmer -------------------------------------------------------------------------- */ .model-grid .wrap label { --model-fill: rgb(215 101 34 / .030); --model-selected-fill: #f5d8c5; --model-glow: .06; --shimmer-red: rgb(214 86 39); --shimmer-red-soft: rgb(235 141 113); background: var(--model-fill) !important; border: 1px solid #dfcfc3 !important; } /* Capability changes only the shimmer warmth/glow, not the speed. */ .model-grid .wrap label:has(input[value*="baseline"]) { --model-glow: .050; --shimmer-red: rgb(207 112 82); --shimmer-red-soft: rgb(236 198 184); } .model-grid .wrap label:has(input[value*="dpdfnet2"]) { --model-glow: .090; --shimmer-red: rgb(212 97 61); --shimmer-red-soft: rgb(241 183 159); } .model-grid .wrap label:has(input[value*="dpdfnet4"]) { --model-glow: .145; --shimmer-red: rgb(217 77 44); --shimmer-red-soft: rgb(242 152 123); } .model-grid .wrap label:has(input[value*="dpdfnet8"]) { --model-glow: .210; --shimmer-red: rgb(219 55 28); --shimmer-red-soft: rgb(246 123 93); } .model-grid .wrap label > span { position: relative; z-index: 3; color: #23272c !important; } .model-grid .wrap label:hover { border-color: #c99a79 !important; background: rgb(215 101 34 / .060) !important; } .model-grid .wrap label.selected, .model-grid .wrap label:has(input:checked) { border-color: rgb(215 101 34 / .54) !important; background: var(--model-selected-fill) !important; box-shadow: inset 0 0 0 1px rgb(255 255 255 / .30), 0 3px 12px rgb(215 101 34 / calc(var(--model-glow) * .34)) !important; } /* Selected model names get a moving light pass. The shimmer always runs at the same speed; capability is expressed only by warmer red tones and stronger glow. */ .model-grid .wrap label.selected > span, .model-grid .wrap label:has(input:checked) > span { color: transparent !important; background-image: linear-gradient( 100deg, #5c2d1b 0%, #7b3b23 18%, var(--shimmer-red) 34%, #fffaf7 48%, #ffffff 50%, #fff3ee 54%, var(--shimmer-red-soft) 67%, var(--shimmer-red) 82%, #6b321d 100% ); background-size: 220% 100%; background-position: 130% 50%; -webkit-background-clip: text; background-clip: text; filter: drop-shadow(0 0 3px rgb(215 101 34 / calc(var(--model-glow) * .58))); animation: model-text-shimmer 2.4s linear infinite; } @keyframes model-text-shimmer { from { background-position: 130% 50%; } to { background-position: -40% 50%; } } /* -------------------------------------------------------------------------- NOISE REDUCTION — one calm dotted wave -------------------------------------------------------------------------- */ .attn-grid .wrap label { --reduction-fill: rgb(215 101 34 / .032); --reduction-selected-fill: rgb(215 101 34 / .098); --wave-opacity: .34; --wave-speed: 5.2s; background: var(--reduction-fill) !important; } /* More reduction => less visible and less active background motion. */ .attn-grid .wrap label:nth-child(1) { --reduction-fill: rgb(215 101 34 / .032); --reduction-selected-fill: rgb(215 101 34 / .100); --wave-opacity: .48; --wave-speed: 4.7s; } .attn-grid .wrap label:nth-child(2) { --reduction-fill: rgb(215 101 34 / .038); --reduction-selected-fill: rgb(215 101 34 / .112); --wave-opacity: .36; --wave-speed: 6.4s; } .attn-grid .wrap label:nth-child(3) { --reduction-fill: rgb(215 101 34 / .044); --reduction-selected-fill: rgb(215 101 34 / .124); --wave-opacity: .25; --wave-speed: 8.5s; } .attn-grid .wrap label:nth-child(4) { --reduction-fill: rgb(215 101 34 / .050); --reduction-selected-fill: rgb(215 101 34 / .136); --wave-opacity: .15; --wave-speed: 11.5s; } .attn-grid .wrap label::before { content: ""; position: absolute; inset: 0; z-index: 0; pointer-events: none; opacity: 0; background-image: url("data:image/svg+xml,%3Csvg xmlns='http://www.w3.org/2000/svg' width='160' height='30' viewBox='0 0 160 30'%3E%3Cg fill='%23231f1c'%3E%3Ccircle cx='0' cy='15.00' r='1.2'/%3E%3Ccircle cx='10' cy='12.13' r='1.2'/%3E%3Ccircle cx='20' cy='9.70' r='1.2'/%3E%3Ccircle cx='30' cy='8.08' r='1.2'/%3E%3Ccircle cx='40' cy='7.50' r='1.2'/%3E%3Ccircle cx='50' cy='8.08' r='1.2'/%3E%3Ccircle cx='60' cy='9.70' r='1.2'/%3E%3Ccircle cx='70' cy='12.13' r='1.2'/%3E%3Ccircle cx='80' cy='15.00' r='1.2'/%3E%3Ccircle cx='90' cy='17.87' r='1.2'/%3E%3Ccircle cx='100' cy='20.30' r='1.2'/%3E%3Ccircle cx='110' cy='21.92' r='1.2'/%3E%3Ccircle cx='120' cy='22.50' r='1.2'/%3E%3Ccircle cx='130' cy='21.92' r='1.2'/%3E%3Ccircle cx='140' cy='20.30' r='1.2'/%3E%3Ccircle cx='150' cy='17.87' r='1.2'/%3E%3Ccircle cx='160' cy='15.00' r='1.2'/%3E%3C/g%3E%3C/svg%3E"); background-repeat: repeat-x; background-size: 160px 30px; background-position: 0 50%; } .attn-grid .wrap label::after { content: none; } .attn-grid .wrap label > span { position: relative; z-index: 2; } .attn-grid .wrap label:hover { background: rgb(215 101 34 / .075) !important; } .attn-grid .wrap label.selected, .attn-grid .wrap label:has(input:checked) { border-color: var(--accent) !important; background: var(--reduction-selected-fill) !important; box-shadow: inset 0 0 0 1px rgb(215 101 34 / .24), 0 3px 10px rgb(215 101 34 / .045) !important; } .attn-grid .wrap label.selected::before, .attn-grid .wrap label:has(input:checked)::before { opacity: var(--wave-opacity); animation: reduction-wave var(--wave-speed) linear infinite; } @keyframes reduction-wave { from { background-position: 0 50%; } to { background-position: 160px 50%; } } @media (prefers-reduced-motion: reduce) { .model-grid .wrap label.selected > span, .model-grid .wrap label:has(input:checked) > span, .attn-grid .wrap label::before { animation: none !important; } .model-grid .wrap label.selected > span, .model-grid .wrap label:has(input:checked) > span { background-position: 50% 50%; } .attn-grid .wrap label.selected::before, .attn-grid .wrap label:has(input:checked)::before { opacity: calc(var(--wave-opacity) * .82); background-position: 40px 50%; } } .source-tabs input[type="radio"], .option-grid input[type="radio"] { position: absolute !important; width: 1px !important; height: 1px !important; opacity: 0 !important; pointer-events: none !important; } .audio-input, .audio-output, .spec-card { overflow: hidden; border: 1px solid var(--line) !important; border-radius: 10px !important; background: #fafaf8 !important; box-shadow: none !important; } .audio-input { min-height: 205px; } .audio-output { min-height: 108px; } .audio-input .wrap, .audio-output .wrap { border: 0 !important; background: transparent !important; } .audio-input .upload-container, .audio-input .upload-container *, .audio-input .empty, .audio-input .empty * { color: #444b53 !important; opacity: 1 !important; } .audio-input svg, .audio-output svg { color: #606871 !important; stroke: #606871 !important; } .gradio-container label, .gradio-container .label-wrap, .gradio-container .label-wrap span, .gradio-container .block-label, .gradio-container .block-label span { color: #30353a !important; opacity: 1 !important; } #run_btn { min-height: 42px; margin-top: 12px; border: 1px solid var(--accent) !important; border-radius: 9px !important; background: var(--accent) !important; color: #ffffff !important; box-shadow: none !important; font-size: 13px !important; font-weight: 780 !important; } #run_btn:hover { background: #bd561a !important; border-color: #bd561a !important; } #status_md { min-height: 18px; margin-top: 6px; color: #4f565e !important; font-size: 11px; font-weight: 600; text-align: center; } #status_md p { margin: 0; color: inherit !important; } .results-title { margin: 18px 2px 8px; color: var(--ink); font-size: 15px; font-weight: 760; letter-spacing: -.015em; } .result-tag { margin-bottom: 6px; color: #525960; font-size: 11px; font-weight: 760; } .result-tag.after { color: #a64a16; } .spec-card img { border-radius: 9px !important; } .gradio-container .accordion { margin-top: 8px !important; border: 1px solid var(--line) !important; border-radius: 10px !important; background: var(--paper) !important; box-shadow: none !important; } .gradio-container .accordion > .label-wrap { padding: 9px 11px !important; color: #3f464d !important; font-size: 11px !important; font-weight: 720 !important; } /* Compact Gradio's internal spacing without changing the workflow. */ .flow-card > .gap, .result-card > .gap { gap: 8px !important; } .audio-input .upload-container, .audio-input .empty { min-height: 165px !important; padding: 12px !important; } .audio-input .upload-container p, .audio-input .empty p { margin: 2px 0 !important; font-size: 12px !important; line-height: 1.3 !important; } .audio-input .upload-container svg, .audio-input .empty svg { width: 23px !important; height: 23px !important; } .audio-output audio { min-height: 38px !important; } @media (max-width: 980px) { .gradio-container { width: min(100% - 20px, 960px) !important; padding: 16px 0 28px !important; } .workspace-shell { flex-direction: column !important; } } @media (max-width: 620px) { .gradio-container { width: calc(100% - 20px) !important; } .flow-card, .result-card { padding: 13px !important; } .option-grid .wrap { grid-template-columns: repeat(2, minmax(0, 1fr)) !important; } .sample-rate-grid .wrap { grid-template-columns: repeat(3, minmax(0, 1fr)) !important; } .attn-grid .wrap { grid-template-columns: repeat(2, minmax(0, 1fr)) !important; } #app-header h1 { font-size: 29px; } } """ HEADER_HTML = """

DPDFNet

""" try: _GRADIO_MAJOR = int(gr.__version__.split(".", 1)[0]) except (AttributeError, ValueError): _GRADIO_MAJOR = 5 _BLOCKS_KWARGS = { "title": "DPDFNet Speech Enhancement", "fill_width": True, } _LAUNCH_KWARGS = {} if _GRADIO_MAJOR >= 6: _LAUNCH_KWARGS.update(theme=THEME, css=CSS) else: _BLOCKS_KWARGS.update(theme=THEME, css=CSS) with gr.Blocks(**_BLOCKS_KWARGS) as demo: gr.HTML(HEADER_HTML) with gr.Row(elem_classes=["workspace-shell"], equal_height=False): with gr.Column(scale=7, min_width=560, elem_classes=["flow-card"]): gr.HTML( '
1Add audio
' ) source = gr.Radio( choices=[("↑ Upload", "Upload"), ("● Record", "Microphone")], value="Upload", show_label=False, container=False, elem_classes=["source-tabs"], ) file_audio = gr.Audio( sources=["upload"], type="filepath", format="wav", label="Audio file", visible=True, buttons=["download"], elem_classes=["audio-input"], ) mic_audio = gr.Audio( sources=["microphone"], type="filepath", format="wav", label="Record", visible=False, buttons=["download"], elem_classes=["audio-input"], ) with gr.Column(scale=5, min_width=470, elem_classes=["flow-card"]): gr.HTML( '
2Enhance
' ) gr.HTML('
Sample rate
') sample_rate = gr.Radio( choices=SAMPLE_RATE_CHOICES, value=str(DEFAULT_SAMPLE_RATE) if DEFAULT_SAMPLE_RATE is not None else None, show_label=False, container=False, interactive=True, elem_classes=["option-grid", "sample-rate-grid"], ) gr.HTML('
Model
') model_key = gr.Radio( choices=DEFAULT_MODEL_CHOICES, value=DEFAULT_MODEL_KEY, show_label=False, container=False, interactive=True, elem_classes=["option-grid", "model-grid"], ) gr.HTML('
Noise reduction
') attn_choice = gr.Radio( choices=ATTN_RADIO_CHOICES, value=DEFAULT_ATTN_CHOICE, show_label=False, container=False, interactive=True, elem_classes=["option-grid", "attn-grid"], ) run_btn = gr.Button( "Remove noise", variant="primary", elem_id="run_btn", ) status = gr.Markdown("Ready", elem_id="status_md") gr.HTML('
3. Compare
') with gr.Row(equal_height=True, elem_classes=["results-row"]): with gr.Column(elem_classes=["result-card"]): gr.HTML('
Original
') out_noisy = gr.Audio( label="Before", interactive=False, format="wav", buttons=["download"], elem_classes=["audio-output"], ) with gr.Column(elem_classes=["result-card"]): gr.HTML('
Enhanced
') out_enh = gr.Audio( label="After", interactive=False, format="wav", buttons=["download"], elem_classes=["audio-output"], ) with gr.Accordion("View spectrograms", open=False): with gr.Row(): img_noisy = gr.Image( label="Original", elem_classes=["spec-card"], ) img_enh = gr.Image( label="Enhanced", elem_classes=["spec-card"], ) source.change( fn=set_source_visibility, inputs=source, outputs=[mic_audio, file_audio], ) sample_rate.change( fn=set_model_choices, inputs=sample_rate, outputs=model_key, ) run_btn.click( fn=run_enhancement, inputs=[source, mic_audio, file_audio, model_key, attn_choice], outputs=[out_noisy, out_enh, img_noisy, img_enh, status], api_name="enhance", ) if __name__ == "__main__": demo.queue(max_size=32).launch(**_LAUNCH_KWARGS)