#!/usr/bin/env python3 from __future__ import annotations import json import os from dataclasses import dataclass from pathlib import Path from typing import Any os.environ.setdefault("MPLCONFIGDIR", "/tmp/matplotlib") import gradio as gr import matplotlib matplotlib.use("Agg") import matplotlib.pyplot as plt import numpy as np import onnxruntime as ort import soundfile as sf import torch from scipy.signal import resample_poly ROOT = Path(__file__).resolve().parent MODELS = ROOT / "models" IS_HF_SPACE = bool(os.environ.get("SPACE_ID")) os.environ.setdefault("HF_HOME", "/tmp/huggingface") os.environ.setdefault("TRANSFORMERS_CACHE", str(Path(os.environ["HF_HOME"]) / "transformers")) RUNS = Path(os.environ.get("HALF_DUPLEX_RUN_DIR", "/tmp/half_duplex_runs" if IS_HF_SPACE else str(ROOT / "runs"))) DEFAULT_REPO = Path("/Utilisateurs/tnguye28/vad-lstm") DEFAULT_ENROLL = DEFAULT_REPO / "debug" / "audio (6).wav" DEFAULT_MIC = DEFAULT_REPO / "debug" / "audio (7).wav" DEFAULT_ASSISTANT = DEFAULT_REPO / "debug" / "test-target-spk4.wav" DEFAULT_PVAD_ONNX = MODELS / "pvad_core.onnx" DEFAULT_PVAD_H256_ONNX = MODELS / "pvad_core_h256.onnx" DEFAULT_SILERO_JIT = MODELS / "silero_vad.jit" DEFAULT_SMARTTURN_ONNX = MODELS / "smartturn-v3.1.onnx" LOCAL_SOTA_PREVBEST_CK50_INT8 = MODELS / "sota_prevbest_incw110_ck50_staticcalib8.onnx" LOCAL_SOTA_PREVBEST_CK100_INT8 = MODELS / "sota_prevbest_incw110_ck100_staticcalib8.onnx" LOCAL_SOTA_PREVBEST_CK150_INT8 = MODELS / "sota_prevbest_incw110_ck150_staticcalib8.onnx" LOCAL_SOTA_HARDNEG4K_CK50_INT8 = MODELS / "sota_hardneg4k_v2_incw110_ck50_staticcalib8.onnx" SOTA_ONNX_ROOT = Path("/Utilisateurs/tnguye28/smartturn-vn/outputs/lumi_turn/onnx_exports") REMOTE_SOTA_PREVBEST_CK50_INT8 = SOTA_ONNX_ROOT / "sota_prevbest_incw110_ck50_staticcalib8" / "model_int8_static_calib8.onnx" REMOTE_SOTA_PREVBEST_CK100_INT8 = SOTA_ONNX_ROOT / "sota_prevbest_incw110_ck100_staticcalib8" / "model_int8_static_calib8.onnx" REMOTE_SOTA_PREVBEST_CK150_INT8 = SOTA_ONNX_ROOT / "sota_prevbest_incw110_ck150_staticcalib8" / "model_int8_static_calib8.onnx" REMOTE_SOTA_HARDNEG4K_CK50_INT8 = SOTA_ONNX_ROOT / "sota_hardneg4k_v2_incw110_ck50_staticcalib8" / "model_int8_static_calib8.onnx" DUALTURN_MODEL_ID = "anyreach-ai/dualturn-qwen2.5-mimi-0.5B" DEFAULT_DEVICE = "cuda" if torch.cuda.is_available() else "cpu" STATE_COLORS = { "ACTIVE": "#d62828", "HOLD": "#f4a261", "SOFT_END": "#2563eb", "END": "#2a9d8f", "INTERRUPT": "#7b2cbf", "UNKNOWN": "#8d99ae", } def load_wav_16k(path: str | Path | None) -> np.ndarray: if path is None: return np.zeros(0, dtype=np.float32) audio, sr = sf.read(str(path), dtype="float32", always_2d=False) if audio.ndim > 1: audio = np.mean(audio, axis=1) audio = np.asarray(audio, dtype=np.float32).reshape(-1) if sr != 16000: gcd = np.gcd(sr, 16000) audio = resample_poly(audio, 16000 // gcd, sr // gcd).astype(np.float32) peak = float(np.max(np.abs(audio))) if len(audio) else 0.0 if peak > 1.0: audio = audio / peak return audio def default_audio_value(path: Path) -> str | None: return str(path) if path.exists() else None def prefer_existing(local_path: Path, remote_path: Path) -> Path: return local_path if local_path.exists() else remote_path def write_wav(path: Path, audio: np.ndarray, sample_rate: int = 16000) -> str: path.parent.mkdir(parents=True, exist_ok=True) sf.write(str(path), np.asarray(audio, dtype=np.float32), sample_rate) return str(path) def make_ort_session(path: str | Path) -> ort.InferenceSession: path = str(path).strip() opts = ort.SessionOptions() opts.execution_mode = ort.ExecutionMode.ORT_SEQUENTIAL opts.intra_op_num_threads = 1 opts.inter_op_num_threads = 1 opts.graph_optimization_level = ort.GraphOptimizationLevel.ORT_ENABLE_ALL return ort.InferenceSession(path, sess_options=opts, providers=["CPUExecutionProvider"]) def frame_audio(audio: np.ndarray, frame_size: int = 512) -> np.ndarray: if len(audio) == 0: return np.zeros((1, frame_size), dtype=np.float32) n = int(np.ceil(len(audio) / frame_size)) padded = np.zeros(n * frame_size, dtype=np.float32) padded[: len(audio)] = audio return padded.reshape(n, frame_size) def resample_audio(audio: np.ndarray, src_sr: int, dst_sr: int) -> np.ndarray: audio = np.asarray(audio, dtype=np.float32).reshape(-1) if src_sr == dst_sr: return audio gcd = np.gcd(src_sr, dst_sr) return resample_poly(audio, dst_sr // gcd, src_sr // gcd).astype(np.float32) class PvadOnnx: def __init__(self, pvad_path: str | Path, silero_path: str | Path): self.session = make_ort_session(pvad_path) self.num_layers, self.hidden_dim = self._infer_recurrent_shape() jit = torch.jit.load(str(silero_path), map_location="cpu") self.silero = jit._model if hasattr(jit, "_model") else jit self.silero.eval() def _infer_recurrent_shape(self) -> tuple[int, int]: inputs = {inp.name: inp for inp in self.session.get_inputs()} h0 = inputs.get("h0") if h0 is None or len(h0.shape) != 3: raise ValueError("PVAD ONNX must expose h0 input with shape [layers, batch, hidden]") layers, _batch, hidden = h0.shape if not isinstance(layers, int) or not isinstance(hidden, int): raise ValueError(f"PVAD ONNX h0 shape must have static layers/hidden dims, got {h0.shape}") return int(layers), int(hidden) def silero_scores(self, frames: np.ndarray) -> np.ndarray: context = torch.zeros(1, 64, dtype=torch.float32) state = torch.zeros(2, 1, 128, dtype=torch.float32) scores = [] with torch.no_grad(): for frame in frames: frame_t = torch.from_numpy(frame.reshape(1, 512).astype(np.float32)) score, state = self.silero(torch.cat([context, frame_t], dim=1), state) context = frame_t[:, -64:] scores.append(float(score.reshape(-1)[0].item())) return np.asarray(scores, dtype=np.float32).reshape(-1, 1) def run_core( self, frames: np.ndarray, target_vector: np.ndarray, vad_scores: np.ndarray, ) -> tuple[np.ndarray, np.ndarray]: h = np.zeros((self.num_layers, 1, self.hidden_dim), dtype=np.float32) c = np.zeros((self.num_layers, 1, self.hidden_dim), dtype=np.float32) target = target_vector.reshape(1, 16).astype(np.float32) probs = [] embeds = [] for frame, vad_score in zip(frames, vad_scores): final_prob, _raw_prob, embed, h, c = self.session.run( None, { "frame_pcm": frame.reshape(1, 512).astype(np.float32), "target_vector": target, "vad_score": vad_score.reshape(1, 1).astype(np.float32), "h0": h, "c0": c, }, ) probs.append(final_prob[0]) embeds.append(embed[0]) return np.asarray(probs, dtype=np.float32), np.asarray(embeds, dtype=np.float32) def target_vector(self, enroll_audio: np.ndarray) -> np.ndarray: frames = frame_audio(enroll_audio) scores = self.silero_scores(frames) _probs, embeds = self.run_core(frames, np.zeros(16, dtype=np.float32), scores) weights = scores.reshape(-1, 1) pooled = np.sum(embeds * weights, axis=0) / (float(np.sum(weights)) + 1e-8) norm = float(np.linalg.norm(pooled)) return (pooled / max(norm, 1e-8)).astype(np.float32) def predict_target_probs(self, mic_audio: np.ndarray, enroll_audio: np.ndarray) -> np.ndarray: max_val = float(np.max(np.abs(mic_audio)) + 1e-8) if len(mic_audio) else 1.0 frames = frame_audio(mic_audio / max_val) target = self.target_vector(enroll_audio) scores = self.silero_scores(frames) probs, _embeds = self.run_core(frames, target, scores) return probs[:, 0].astype(np.float32) class SmartTurnOnnx: def __init__(self, model_path: str | Path): from transformers import WhisperFeatureExtractor self.session = make_ort_session(model_path) self.feature_extractor = WhisperFeatureExtractor(chunk_length=8) def predict_prob(self, audio_16k: np.ndarray) -> float: samples = np.asarray(audio_16k, dtype=np.float32).reshape(-1) max_samples = 8 * 16000 if len(samples) > max_samples: samples = samples[-max_samples:] elif len(samples) < max_samples: samples = np.pad(samples, (max_samples - len(samples), 0), mode="constant") inputs = self.feature_extractor( samples, sampling_rate=16000, return_tensors="np", padding="max_length", max_length=max_samples, truncation=True, do_normalize=True, ) features = inputs.input_features.squeeze(0).astype(np.float32)[None, ...] outputs = self.session.run(None, {"input_features": features}) return float(outputs[0][0].item()) class DualTurnHF: def __init__(self, model_id: str, device: str): from transformers import AutoModel use_device = "cuda" if device == "cuda" and torch.cuda.is_available() else "cpu" self.device = torch.device(use_device) self.model = AutoModel.from_pretrained(model_id, trust_remote_code=True) self.model.to(self.device) self.model.eval() @torch.no_grad() def predict_channels(self, ch0_16k: np.ndarray, ch1_16k: np.ndarray) -> dict[str, np.ndarray]: ch0 = resample_audio(ch0_16k, 16000, 24000) ch1 = resample_audio(ch1_16k, 16000, 24000) n = min(len(ch0), len(ch1)) if n <= 0: ch0 = np.zeros(1, dtype=np.float32) ch1 = np.zeros(1, dtype=np.float32) else: ch0 = ch0[:n] ch1 = ch1[:n] stereo = torch.from_numpy(np.stack([ch0, ch1], axis=0)).to(self.device) out = self.model(stereo, sr=24000) fvad = out.fvad_probs.detach().float().cpu().numpy() if fvad.ndim == 3: fvad = fvad[0] return { "vad": self._user_np(out.vad_probs), "hold": self._user_np(out.hold_probs), "eot": self._user_np(out.eot_probs), "fvad_240": np.asarray(fvad[:, 0], dtype=np.float32).reshape(-1), "fvad_480": np.asarray(fvad[:, 1], dtype=np.float32).reshape(-1), "fvad_960": np.asarray(fvad[:, 2], dtype=np.float32).reshape(-1), "fvad_2000": np.asarray(fvad[:, 3], dtype=np.float32).reshape(-1), } @staticmethod def _user_np(tensor: torch.Tensor) -> np.ndarray: arr = tensor.detach().float().cpu().numpy() if arr.ndim == 3: arr = arr[0] if arr.ndim == 2: arr = arr[:, 0] return np.asarray(arr, dtype=np.float32).reshape(-1) def dualturn_state(row: dict[str, float]) -> str: if row["vad"] >= 0.5: return "ACTIVE" if row["hold"] >= 0.5: return "HOLD" if max(row["fvad_240"], row["fvad_480"], row["fvad_960"], row["fvad_2000"]) >= 0.5: return "CONTINUE" if row["eot"] >= 0.6 and row["hold"] < 0.4 and row["fvad_480"] < 0.35 and row["fvad_960"] < 0.35: return "END" return "UNKNOWN" def latest_dualturn_row(outputs: dict[str, np.ndarray], time_ms: int) -> dict[str, Any]: n = min(len(v) for v in outputs.values()) idx = max(0, n - 1) row = {k: float(v[idx]) if len(v) else 0.0 for k, v in outputs.items()} row["time_ms"] = int(time_ms) row["dualturn_state"] = dualturn_state(row) return row def make_pvad_target_audio(audio: np.ndarray, probs_32ms: np.ndarray, threshold: float) -> np.ndarray: out = np.zeros_like(audio, dtype=np.float32) frame = 512 for idx, prob in enumerate(probs_32ms): start = idx * frame end = min(len(audio), start + frame) if end <= start: break if prob >= threshold: out[start:end] = audio[start:end] return out def pvad_active_between(probs_32ms: np.ndarray, start_sample: int, end_sample: int, threshold: float) -> bool: start_idx = max(0, start_sample // 512) end_idx = min(len(probs_32ms), int(np.ceil(end_sample / 512))) if end_idx <= start_idx: return False return bool(np.max(probs_32ms[start_idx:end_idx]) >= threshold) def frame_is_active(audio: np.ndarray, threshold: float = 0.01) -> bool: if len(audio) == 0: return False return float(np.sqrt(np.mean(np.asarray(audio, dtype=np.float32) ** 2))) >= threshold def active_audio_sec(target_audio: np.ndarray, start: int, end: int) -> float: part = target_audio[start:end] return float(np.count_nonzero(np.abs(part) > 1e-8) / 16000.0) def append_row(rows: list[dict[str, Any]], row: dict[str, Any]) -> None: if rows and rows[-1]["state"] == row["state"] and rows[-1]["source"] == row["source"]: return rows.append(row) @dataclass class RunResult: rows: list[dict[str, Any]] pvad_probs: np.ndarray pvad_target_audio: np.ndarray assistant_track: np.ndarray smart_probs: list[tuple[float, float]] cuts: list[dict[str, Any]] def run_pipeline( enroll_path: str, mic_path: str, assistant_path: str | None, *, mode: str, smartturn_threshold: float, pvad_threshold: float, pvad_model_path: str, smartturn_model_path: str, min_active_target_ms: float, silence_fallback_ms: float, asr_cut_silence_ms: float, model_check_interval_ms: float, append_assistant_after_end: bool, assistant_max_playback_sec: float, device: str, ) -> RunResult: enroll = load_wav_16k(enroll_path) mic = load_wav_16k(mic_path) assistant = load_wav_16k(assistant_path) if assistant_path else np.zeros(0, dtype=np.float32) if assistant_max_playback_sec > 0: assistant = assistant[: int(round(assistant_max_playback_sec * 16000))] pvad = PvadOnnx(pvad_model_path, DEFAULT_SILERO_JIT) smartturn = SmartTurnOnnx(smartturn_model_path) dualturn = DualTurnHF(DUALTURN_MODEL_ID, device) pvad_probs = pvad.predict_target_probs(mic, enroll) target_audio = mic if mode == "raw" else make_pvad_target_audio(mic, pvad_probs, pvad_threshold) assistant_track = np.zeros_like(mic, dtype=np.float32) activity_mask = np.zeros_like(mic, dtype=np.float32) frame_samples = int(round(0.080 * 16000)) check_frames = max(1, int(round(model_check_interval_ms / 80.0))) total_frames = max(1, int(np.ceil(len(mic) / frame_samples))) min_active_sec = min_active_target_ms / 1000.0 silence_fallback_frames = 0 if silence_fallback_ms <= 0 else max(1, int(round(silence_fallback_ms / 80.0))) asr_cut_silence_frames = 0 if asr_cut_silence_ms <= 0 else max(1, int(round(asr_cut_silence_ms / 80.0))) rows: list[dict[str, Any]] = [] smart_probs: list[tuple[float, float]] = [] cuts: list[dict[str, Any]] = [] turn_start: int | None = None asr_cut_start: int | None = None last_soft_cut_end: int | None = None last_check_idx: int | None = None silence_frames = 0 assistant_playing = False assistant_pos = 0 append_row(rows, {"time_sec": 0.0, "state": "UNKNOWN", "source": "idle", "smartturn": 0.0, "dualturn": "UNKNOWN", "assistant": False}) for idx in range(total_frames): start = idx * frame_samples end = min(len(mic), start + frame_samples) if end <= start: break time_sec = start / 16000.0 assistant_active = assistant_playing and assistant_pos < len(assistant) if assistant_active: take = min(end - start, len(assistant) - assistant_pos) assistant_track[start : start + take] += assistant[assistant_pos : assistant_pos + take] assistant_pos += take if assistant_pos >= len(assistant): assistant_playing = False ch0_frame = target_audio[start:end] active = frame_is_active(ch0_frame, threshold=0.01) if active: activity_mask[start:end] = 1.0 if active and turn_start is None: turn_start = start asr_cut_start = start last_soft_cut_end = None last_check_idx = None silence_frames = 0 elif active: if asr_cut_start is None: asr_cut_start = start last_soft_cut_end = None silence_frames = 0 elif turn_start is not None: silence_frames += 1 if turn_start is None: continue periodic_due = last_check_idx is None or (idx - last_check_idx) >= check_frames silence_due = silence_fallback_frames > 0 and silence_frames >= silence_fallback_frames asr_flush_due = asr_cut_silence_frames > 0 and silence_frames >= asr_cut_silence_frames if not periodic_due and not silence_due and not asr_flush_due: continue buffer_ch0 = target_audio[turn_start:end].copy() buffer_ch0 *= activity_mask[turn_start:end] buffer_ch1 = assistant_track[turn_start:end] dual_outputs = dualturn.predict_channels(buffer_ch0, buffer_ch1) dual = latest_dualturn_row(dual_outputs, int(round(time_sec * 1000))) smartturn_due = dual["dualturn_state"] in {"HOLD", "END", "UNKNOWN"} or ( (silence_due or asr_flush_due) and dual["dualturn_state"] != "ACTIVE" ) smart_prob = smartturn.predict_prob(buffer_ch0) if smartturn_due else 0.0 smart_probs.append((time_sec, smart_prob)) last_check_idx = idx state = "HOLD" if dual["dualturn_state"] in {"CONTINUE", "END"} else dual["dualturn_state"] source = "dualturn" active_sec = float(np.sum(activity_mask[turn_start:end] > 0.0) / 16000.0) future_voice_count = int(float(dual.get("fvad_480", 0.0)) >= 0.5) + int(float(dual.get("fvad_960", 0.0)) >= 0.5) smart_end = smartturn_due and smart_prob >= smartturn_threshold if assistant_active and state == "ACTIVE": state = "INTERRUPT" source = "assistant_overlap" assistant_playing = False assistant_pos = 0 elif smart_end: if future_voice_count >= 2: state = "HOLD" source = "future_voice_guard" elif future_voice_count == 1: state = "SOFT_END" source = "smartturn_soft_partial_future_voice" else: state = "END" source = "smartturn" elif silence_due and active_sec >= min_active_sec and state != "ACTIVE": if future_voice_count >= 2: state = "HOLD" source = "future_voice_guard" elif future_voice_count == 1: state = "SOFT_END" source = "silence_soft_partial_future_voice" else: state = "END" source = "silence_fallback" if state == "END" and asr_cut_start is not None and end - asr_cut_start < int(round(min_active_sec * 16000.0)): state = "HOLD" source = "short_asr_cut_guard" append_row( rows, { "time_sec": time_sec, "state": state, "source": source, "smartturn": smart_prob, "dualturn": dual["dualturn_state"], "assistant": bool(assistant_active), "vad": dual["vad"], "hold": dual["hold"], "eot": dual["eot"], "fvad_480": dual["fvad_480"], "fvad_960": dual["fvad_960"], }, ) if state == "SOFT_END" and asr_cut_start is not None: can_emit_soft_cut = ( end - asr_cut_start >= int(round(min_active_sec * 16000.0)) and last_soft_cut_end is None ) if can_emit_soft_cut: cut_audio = target_audio[asr_cut_start:end].copy() * activity_mask[asr_cut_start:end] cuts.append( { "start": asr_cut_start / 16000.0, "end": end / 16000.0, "duration": (end - asr_cut_start) / 16000.0, "smartturn": smart_prob, "dualturn": dual["dualturn_state"], "vad": dual["vad"], "hold": dual["hold"], "eot": dual["eot"], "fvad_480": dual["fvad_480"], "fvad_960": dual["fvad_960"], "note": "soft_end_future_voice", "audio": cut_audio, } ) last_soft_cut_end = end if state == "END": if asr_cut_start is not None: cut_start = asr_cut_start cut_audio = target_audio[cut_start:end].copy() * activity_mask[cut_start:end] cuts.append( { "start": cut_start / 16000.0, "end": end / 16000.0, "duration": (end - cut_start) / 16000.0, "smartturn": smart_prob, "dualturn": dual["dualturn_state"], "vad": dual["vad"], "hold": dual["hold"], "eot": dual["eot"], "fvad_480": dual["fvad_480"], "fvad_960": dual["fvad_960"], "note": "hard_end", "audio": cut_audio, } ) if append_assistant_after_end and len(assistant): assistant_playing = True assistant_pos = 0 turn_start = None asr_cut_start = None last_soft_cut_end = None last_check_idx = None silence_frames = 0 elif asr_flush_due and asr_cut_start is not None: cut_audio = target_audio[asr_cut_start:end].copy() * activity_mask[asr_cut_start:end] cuts.append( { "start": asr_cut_start / 16000.0, "end": end / 16000.0, "duration": (end - asr_cut_start) / 16000.0, "smartturn": smart_prob, "dualturn": dual["dualturn_state"], "vad": dual["vad"], "hold": dual["hold"], "eot": dual["eot"], "fvad_480": dual["fvad_480"], "fvad_960": dual["fvad_960"], "note": "silence_asr_flush", "audio": cut_audio, } ) asr_cut_start = None turn_start = None last_soft_cut_end = None last_check_idx = None silence_frames = 0 if asr_cut_start is not None and asr_cut_start < len(mic): cut_audio = target_audio[asr_cut_start:].copy() * activity_mask[asr_cut_start:] if turn_start is not None and np.count_nonzero(np.abs(cut_audio) > 1e-8) > 0: buffer_ch0 = target_audio[turn_start:].copy() * activity_mask[turn_start:] buffer_ch1 = assistant_track[turn_start:] dual_outputs = dualturn.predict_channels(buffer_ch0, buffer_ch1) dual = latest_dualturn_row(dual_outputs, int(round(len(mic) / 16000.0 * 1000))) smart_prob = smartturn.predict_prob(buffer_ch0) smart_probs.append((len(mic) / 16000.0, smart_prob)) else: dual = {"dualturn_state": "FINAL", "vad": 0.0, "hold": 0.0, "eot": 0.0, "fvad_480": 0.0, "fvad_960": 0.0} smart_prob = 0.0 append_row( rows, { "time_sec": len(mic) / 16000.0, "state": "FINAL", "source": "final_flush", "smartturn": smart_prob, "dualturn": dual["dualturn_state"], "assistant": False, "vad": dual["vad"], "hold": dual["hold"], "eot": dual["eot"], "fvad_480": dual["fvad_480"], "fvad_960": dual["fvad_960"], }, ) cuts.append( { "start": asr_cut_start / 16000.0, "end": len(mic) / 16000.0, "duration": (len(mic) - asr_cut_start) / 16000.0, "smartturn": smart_prob, "dualturn": dual["dualturn_state"], "vad": dual["vad"], "hold": dual["hold"], "eot": dual["eot"], "fvad_480": dual["fvad_480"], "fvad_960": dual["fvad_960"], "note": "final_flush", "audio": cut_audio, } ) return RunResult(rows, pvad_probs, target_audio, assistant_track, smart_probs, cuts) def plot_result(result: RunResult, duration_sec: float, out_path: Path) -> str: out_path.parent.mkdir(parents=True, exist_ok=True) times_pvad = np.arange(len(result.pvad_probs), dtype=np.float32) * 0.032 smart_t = [x[0] for x in result.smart_probs] smart_y = [x[1] for x in result.smart_probs] fig, axes = plt.subplots(3, 1, figsize=(13, 7.2), sharex=True, gridspec_kw={"height_ratios": [1.3, 1.6, 0.65]}) axes[0].plot(times_pvad, result.pvad_probs, color="#1d4ed8", linewidth=1.2, label="PVAD target") axes[0].axhline(0.5, color="#64748b", linestyle="--", linewidth=0.8) axes[0].set_ylim(-0.02, 1.02) axes[0].legend(loc="upper right") axes[0].grid(True, alpha=0.25) dual_times = [r["time_sec"] for r in result.rows if "vad" in r] for key, color, style in [ ("vad", "#d62828", "-"), ("hold", "#f4a261", "-"), ("eot", "#2a9d8f", "-"), ("fvad_480", "#7c3aed", "--"), ("fvad_960", "#0891b2", "--"), ]: axes[1].plot( dual_times, [r.get(key, np.nan) for r in result.rows if "vad" in r], label=f"DualTurn {key}", color=color, linestyle=style, marker="o", markersize=2, linewidth=1.1, ) axes[1].plot(smart_t, smart_y, color="#2d6a4f", marker="s", markersize=2.5, linewidth=1.2, label="SmartTurn END") axes[1].axhline(0.9, color="#2d6a4f", linestyle="--", linewidth=0.8, alpha=0.55) axes[1].set_ylim(-0.02, 1.02) axes[1].legend(loc="upper right", ncol=2) axes[1].grid(True, alpha=0.25) ax = axes[2] ax.set_ylim(0, 1) ax.set_yticks([]) ax.set_xlim(0, max(duration_sec, 0.1)) for i, row in enumerate(result.rows): start = float(row["time_sec"]) end = float(result.rows[i + 1]["time_sec"]) if i + 1 < len(result.rows) else duration_sec if end <= start: end = start + 0.08 state = row["state"] ax.axvspan(start, end, color=STATE_COLORS.get(state, "#8d99ae"), alpha=0.85) if end - start >= 0.35: ax.text((start + end) / 2, 0.5, state, ha="center", va="center", fontsize=8, color="white") ax.axvline(start, color="#111827", linewidth=0.6, alpha=0.35) ax.set_xlabel("Time (sec)") fig.tight_layout() fig.savefig(out_path, dpi=140) plt.close(fig) return str(out_path) def run_gradio( enroll_audio: str, mic_audio: str, assistant_audio: str | None, mode: str, smartturn_threshold: float, pvad_threshold: float, pvad_model_path: str, smartturn_model_selection: str, smartturn_custom_model_path: str, min_active_target_ms: float, silence_fallback_ms: float, asr_cut_silence_ms: float, model_check_interval_ms: float, append_assistant_after_end: bool, assistant_max_playback_sec: float, device: str, ) -> tuple[str, str, str, str, str, Any, str | None]: if not enroll_audio: raise gr.Error("Upload an enrollment audio file.") if not mic_audio: raise gr.Error("Upload a mic audio file.") smartturn_model_path = resolve_smartturn_model_path(smartturn_model_selection, smartturn_custom_model_path) result = run_pipeline( enroll_audio, mic_audio, assistant_audio, mode=mode, smartturn_threshold=float(smartturn_threshold), pvad_threshold=float(pvad_threshold), pvad_model_path=pvad_model_path, smartturn_model_path=smartturn_model_path, min_active_target_ms=float(min_active_target_ms), silence_fallback_ms=float(silence_fallback_ms), asr_cut_silence_ms=float(asr_cut_silence_ms), model_check_interval_ms=float(model_check_interval_ms), append_assistant_after_end=bool(append_assistant_after_end), assistant_max_playback_sec=float(assistant_max_playback_sec), device=device, ) RUNS.mkdir(parents=True, exist_ok=True) mic = load_wav_16k(mic_audio) timeline = plot_result(result, len(mic) / 16000.0, RUNS / "timeline.png") target_wav = write_wav(RUNS / "pvad_target_timeline.wav", result.pvad_target_audio) assistant_wav = write_wav(RUNS / "assistant_channel.wav", result.assistant_track) mic_with_assistant_wav = write_wav(RUNS / "mic_with_assistant_echo.wav", mic + result.assistant_track) cut_paths = [] cut_choices = [] cut_dir = RUNS / "turn_cuts" for idx, cut in enumerate(result.cuts, start=1): path = cut_dir / f"{idx:03d}_{cut['start']:.2f}_{cut['end']:.2f}.wav" write_wav(path, cut["audio"]) path_str = str(path) cut_paths.append(path_str) label = ( f"{idx:03d} | {cut['start']:.2f}s-{cut['end']:.2f}s | " f"{cut['duration']:.2f}s | smart={cut['smartturn']:.3f} | " f"dual={cut.get('dualturn', '')} " f"vad={cut.get('vad', 0.0):.2f} hold={cut.get('hold', 0.0):.2f} " f"eot={cut.get('eot', 0.0):.2f} | {cut.get('note', 'cut')}" ) cut_choices.append((label, path_str)) rows = [ { "time_sec": round(r["time_sec"], 3), "state": r["state"], "source": r["source"], "smartturn": round(float(r["smartturn"]), 3), "dualturn": r["dualturn"], "dual_vad": round(float(r.get("vad", 0.0)), 3), "dual_hold": round(float(r.get("hold", 0.0)), 3), "dual_eot": round(float(r.get("eot", 0.0)), 3), "dual_fvad_480": round(float(r.get("fvad_480", 0.0)), 3), "dual_fvad_960": round(float(r.get("fvad_960", 0.0)), 3), "assistant": r["assistant"], } for r in result.rows ] summary = { "timeline": timeline, "target_audio": target_wav, "assistant_channel": assistant_wav, "mic_with_assistant": mic_with_assistant_wav, "turn_cuts": cut_paths, "models": { "pvad": str(pvad_model_path), "silero": str(DEFAULT_SILERO_JIT), "smartturn": str(smartturn_model_path), "dualturn": DUALTURN_MODEL_ID, }, "state_counts": {state: sum(1 for r in result.rows if r["state"] == state) for state in STATE_COLORS}, } first_cut = cut_paths[0] if cut_paths else None return ( timeline, json.dumps(summary, indent=2), json.dumps(rows, indent=2), target_wav, mic_with_assistant_wav, gr.update(choices=cut_choices, value=first_cut), first_cut, ) def select_cut_audio(cut_path: str | None) -> str | None: return cut_path or None def resolve_smartturn_model_path(selection: str, custom_path: str | None) -> str: selection = str(selection).strip() if selection == "custom": path = str(custom_path or "").strip() if not path: raise gr.Error("Paste a SmartTurn ONNX model path or choose one from the dropdown.") return path return selection def plot_smartturn_only(rows: list[dict[str, Any]], duration_sec: float, out_path: Path) -> str: out_path.parent.mkdir(parents=True, exist_ok=True) times = [float(r["time_sec"]) for r in rows] probs = [float(r["smartturn"]) for r in rows] states = [1.0 if r["state"] == "END" else 0.0 for r in rows] fig, axes = plt.subplots(2, 1, figsize=(12, 5), sharex=True) axes[0].plot(times, probs, color="#2563eb", marker="o", markersize=3, linewidth=1.2) axes[0].axhline(0.5, color="#d62828", linestyle="--", linewidth=1.0) axes[0].set_ylim(0, 1) axes[0].set_ylabel("SmartTurn") axes[0].grid(True, alpha=0.25) axes[1].step(times, states, where="post", color="#2a9d8f", linewidth=1.4) axes[1].set_ylim(-0.1, 1.1) axes[1].set_yticks([0, 1], ["RUN", "END"]) axes[1].set_xlabel("Time (sec)") axes[1].grid(True, alpha=0.25) axes[1].set_xlim(0, max(duration_sec, 0.24)) fig.tight_layout() fig.savefig(out_path, dpi=140) plt.close(fig) return str(out_path) def run_smartturn_only_gradio( enroll_audio: str, mic_audio: str, pvad_model_path: str, pvad_threshold: float, smartturn_model_selection: str, smartturn_custom_model_path: str, threshold: float, pvad_silence_ms: float, min_active_target_ms: float, ) -> tuple[str, str, str, str, Any, str | None]: if not enroll_audio: raise gr.Error("Upload an enrollment audio file.") if not mic_audio: raise gr.Error("Upload a mic audio file.") smartturn_model_path = resolve_smartturn_model_path(smartturn_model_selection, smartturn_custom_model_path) enroll = load_wav_16k(enroll_audio) mic = load_wav_16k(mic_audio) pvad = PvadOnnx(pvad_model_path, DEFAULT_SILERO_JIT) smartturn = SmartTurnOnnx(smartturn_model_path) pvad_probs = pvad.predict_target_probs(mic, enroll) target_audio = make_pvad_target_audio(mic, pvad_probs, float(pvad_threshold)) frame_samples = 512 frame_ms = frame_samples / 16000.0 * 1000.0 silence_frames_required = 0 if pvad_silence_ms <= 0 else max(1, int(np.ceil(float(pvad_silence_ms) / frame_ms))) total_frames = max(1, int(np.ceil(len(mic) / frame_samples))) min_active_sec = float(min_active_target_ms) / 1000.0 threshold = float(threshold) activity_mask = np.zeros_like(target_audio, dtype=np.float32) rows: list[dict[str, Any]] = [] cuts: list[dict[str, Any]] = [] turn_start: int | None = None silence_frames = 0 checked_current_silence = False for idx in range(total_frames): start = idx * frame_samples end = min(len(mic), start + frame_samples) if end <= start: break time_sec = start / 16000.0 active = frame_is_active(target_audio[start:end], threshold=0.01) if active: activity_mask[start:end] = 1.0 if active and turn_start is None: turn_start = start silence_frames = 0 checked_current_silence = False elif active: silence_frames = 0 checked_current_silence = False elif turn_start is not None: silence_frames += 1 if turn_start is None: rows.append( { "time_sec": round(time_sec, 3), "state": "IDLE", "source": "idle", "smartturn": 0.0, "target_active_sec": 0.0, "target_samples": 0, } ) continue silence_due = silence_frames >= silence_frames_required and not checked_current_silence if not silence_due: continue buffer = target_audio[turn_start:end].copy() * activity_mask[turn_start:end] prob = smartturn.predict_prob(buffer) active_samples = int(np.count_nonzero(np.abs(buffer) > 1e-8)) active_sec = active_samples / 16000.0 state = "END" if active_samples > 0 and prob > threshold and active_sec >= min_active_sec else "RUN" source = "pvad_silence_smartturn" if state == "END" else "pvad_silence_incomplete" rows.append( { "time_sec": round(time_sec, 3), "state": state, "source": source, "smartturn": round(prob, 4), "pvad_silence_sec": round(silence_frames * frame_samples / 16000.0, 3), "target_active_sec": round(active_sec, 3), "target_samples": active_samples, "turn_start_sec": round(turn_start / 16000.0, 3), } ) checked_current_silence = True if state == "END": cuts.append( { "start": turn_start / 16000.0, "end": end / 16000.0, "smartturn_end": end / 16000.0, "duration": (end - turn_start) / 16000.0, "smartturn": prob, "note": "hard_end", "audio": target_audio[turn_start:end].copy(), } ) turn_start = None silence_frames = 0 checked_current_silence = False if turn_start is not None and turn_start < len(target_audio): end = len(target_audio) buffer = target_audio[turn_start:end].copy() * activity_mask[turn_start:end] active_samples = int(np.count_nonzero(np.abs(buffer) > 1e-8)) active_sec = active_samples / 16000.0 if active_samples > 0: prob = smartturn.predict_prob(buffer) state = "END" if prob > threshold and active_sec >= min_active_sec else "FLUSH" rows.append( { "time_sec": round(end / 16000.0, 3), "state": state, "source": "final_flush", "smartturn": round(prob, 4), "pvad_silence_sec": round(silence_frames * frame_samples / 16000.0, 3), "target_active_sec": round(active_sec, 3), "target_samples": active_samples, "turn_start_sec": round(turn_start / 16000.0, 3), } ) cuts.append( { "start": turn_start / 16000.0, "end": end / 16000.0, "smartturn_end": end / 16000.0, "duration": (end - turn_start) / 16000.0, "smartturn": prob, "note": "final_flush", "audio": target_audio[turn_start:end].copy(), } ) run_dir = RUNS / "smartturn_only" run_dir.mkdir(parents=True, exist_ok=True) timeline = plot_smartturn_only(rows, len(mic) / 16000.0, run_dir / "timeline.png") target_wav = write_wav(run_dir / "pvad_target_audio.wav", target_audio) cut_paths = [] cut_choices = [] cut_dir = run_dir / "turn_cuts" for idx, cut in enumerate(cuts, start=1): path = cut_dir / f"{idx:03d}_{cut['start']:.2f}_{cut['end']:.2f}.wav" write_wav(path, cut["audio"]) path_str = str(path) cut_paths.append(path_str) cut_choices.append( ( f"{idx:03d} | {cut['start']:.2f}s-{cut['end']:.2f}s | " f"{cut['duration']:.2f}s | end={cut['smartturn_end']:.2f}s | smart={cut['smartturn']:.3f} | {cut.get('note', 'cut')}", path_str, ) ) first_cut = cut_paths[0] if cut_paths else None first_end_time = next((float(r["time_sec"]) for r in rows if r["state"] == "END"), None) has_final_flush = any(cut.get("note") == "final_flush" for cut in cuts) summary = { "state": "END_FOUND" if first_end_time is not None else ("FINAL_FLUSH" if has_final_flush else "NO_END"), "num_cuts": len(cuts), "turn_cuts": cut_paths, "audio_duration_sec": round(len(mic) / 16000.0, 3), "first_end_time_sec": None if first_end_time is None else round(first_end_time, 3), "threshold": threshold, "pvad_silence_before_smartturn_ms": float(pvad_silence_ms), "pvad_silence_frames_required": silence_frames_required, "pvad_frame_ms": 32.0, "min_active_target_ms": float(min_active_target_ms), "pvad_threshold": float(pvad_threshold), "pvad_model": str(pvad_model_path), "silero": str(DEFAULT_SILERO_JIT), "smartturn": str(smartturn_model_path), "rule": "PVAD-gate mic audio first. Run SmartTurn only after PVAD target silence. No DualTurn.", } return timeline, json.dumps(summary, indent=2), json.dumps(rows, indent=2), target_wav, gr.update(choices=cut_choices, value=first_cut), first_cut def build_app() -> gr.Blocks: sota_ck150 = prefer_existing(LOCAL_SOTA_PREVBEST_CK150_INT8, REMOTE_SOTA_PREVBEST_CK150_INT8) sota_ck50 = prefer_existing(LOCAL_SOTA_PREVBEST_CK50_INT8, REMOTE_SOTA_PREVBEST_CK50_INT8) sota_ck100 = prefer_existing(LOCAL_SOTA_PREVBEST_CK100_INT8, REMOTE_SOTA_PREVBEST_CK100_INT8) sota_hardneg4k_ck50 = prefer_existing(LOCAL_SOTA_HARDNEG4K_CK50_INT8, REMOTE_SOTA_HARDNEG4K_CK50_INT8) smartturn_choices = [ ("base pretrained smartturn-v3.1 int8", str(DEFAULT_SMARTTURN_ONNX)), ("top1 F1/FNR prev_best_1.1 ck150 | F1 0.8558 FPR 0.1822 FNR 0.1142", str(sota_ck150)), ("top2 F1 prev_best_1.1 ck50 | F1 0.8555 FPR 0.1802 FNR 0.1162", str(sota_ck50)), ("top1 FPR hardneg4k ck50 | F1 0.8150 FPR 0.1663 FNR 0.1964", str(sota_hardneg4k_ck50)), ("top2 FPR prev_best_1.1 ck100 | F1 0.8471 FPR 0.1762 FNR 0.1343", str(sota_ck100)), ("top1 FNR prev_best_1.1 ck150 | F1 0.8558 FPR 0.1822 FNR 0.1142", str(sota_ck150)), ("top2 FNR prev_best_1.1 ck50 | F1 0.8555 FPR 0.1802 FNR 0.1162", str(sota_ck50)), ("custom", "custom"), ] with gr.Blocks(title="Half Duuplex Demo") as demo: gr.Markdown("## Half Duuplex Demo") with gr.Tabs(): with gr.Tab("SmartTurn + DualTurn"): with gr.Row(): enroll = gr.Audio(value=default_audio_value(DEFAULT_ENROLL), label="Enrollment", type="filepath") mic = gr.Audio(value=default_audio_value(DEFAULT_MIC), label="Mic", type="filepath") assistant = gr.Audio(value=default_audio_value(DEFAULT_ASSISTANT), label="Assistant echo", type="filepath") with gr.Row(): mode = gr.Radio(["pvad_gated", "raw"], value="pvad_gated", label="Audio mode") device = gr.Radio(["cuda", "cpu"], value=DEFAULT_DEVICE, label="DualTurn device") append_assistant = gr.Checkbox(value=True, label="Append assistant after END") with gr.Row(): smart_threshold = gr.Slider(0, 1, value=0.5, step=0.01, label="SmartTurn END threshold") pvad_threshold = gr.Slider(0, 1, value=0.5, step=0.01, label="PVAD target threshold") pvad_model = gr.Dropdown( choices=[ ("h64 pvad_core.onnx", str(DEFAULT_PVAD_ONNX)), ("h256 pvad_core_h256.onnx", str(DEFAULT_PVAD_H256_ONNX)), ], value=str(DEFAULT_PVAD_ONNX), label="PVAD ONNX model", ) smartturn_model = gr.Dropdown( choices=smartturn_choices, value=str(DEFAULT_SMARTTURN_ONNX), label="SmartTurn ONNX model", ) smartturn_custom_model = gr.Textbox( value="", label="Custom SmartTurn ONNX path", placeholder="/path/to/model.onnx", ) with gr.Row(): min_active = gr.Slider(0, 2000, value=300, step=50, label="Min active target ms") check_ms = gr.Slider(80, 1000, value=240, step=80, label="Model check interval ms") silence_fallback = gr.Slider(0, 3000, value=800, step=100, label="Silence fallback END ms") with gr.Row(): asr_cut_silence = gr.Slider(0, 3000, value=2000, step=100, label="ASR cut silence ms") assistant_max = gr.Slider(0.5, 12, value=3, step=0.5, label="Assistant max playback sec") run = gr.Button("Run", variant="primary") timeline = gr.Image(label="Timeline", type="filepath") with gr.Row(): target_audio = gr.Audio(label="PVAD target audio", type="filepath") mic_assistant_audio = gr.Audio(label="Mic with assistant echo", type="filepath") with gr.Row(): cut_selector = gr.Dropdown(label="Model / ASR input audio cuts", choices=[], value=None) cut_audio = gr.Audio(label="Selected cut audio", type="filepath") summary = gr.Code(label="Summary JSON", language="json") rows = gr.Code(label="State transitions", language="json") run.click( run_gradio, inputs=[ enroll, mic, assistant, mode, smart_threshold, pvad_threshold, pvad_model, smartturn_model, smartturn_custom_model, min_active, silence_fallback, asr_cut_silence, check_ms, append_assistant, assistant_max, device, ], outputs=[timeline, summary, rows, target_audio, mic_assistant_audio, cut_selector, cut_audio], ) cut_selector.change(select_cut_audio, inputs=[cut_selector], outputs=[cut_audio]) with gr.Tab("SmartTurn Only"): with gr.Row(): st_enroll = gr.Audio(value=default_audio_value(DEFAULT_ENROLL), label="Enrollment", type="filepath") st_mic = gr.Audio(value=default_audio_value(DEFAULT_MIC), label="Mic", type="filepath") with gr.Row(): st_pvad_model = gr.Dropdown( choices=[ ("h64 pvad_core.onnx", str(DEFAULT_PVAD_ONNX)), ("h256 pvad_core_h256.onnx", str(DEFAULT_PVAD_H256_ONNX)), ], value=str(DEFAULT_PVAD_ONNX), label="PVAD ONNX model", ) st_pvad_threshold = gr.Slider(0, 1, value=0.5, step=0.01, label="PVAD target threshold") st_smartturn_model = gr.Dropdown( choices=smartturn_choices, value=str(DEFAULT_SMARTTURN_ONNX), label="SmartTurn ONNX model", ) st_smartturn_custom_model = gr.Textbox( value="", label="Custom SmartTurn ONNX path", placeholder="/path/to/model.onnx", ) with gr.Row(): st_threshold = gr.Slider(0, 1, value=0.5, step=0.01, label="SmartTurn END threshold") st_pvad_silence = gr.Slider(0, 1000, value=200, step=20, label="PVAD silence before SmartTurn ms") st_min_active = gr.Slider(0, 2000, value=300, step=50, label="Min active target ms") st_run = gr.Button("Run SmartTurn Only", variant="primary") st_timeline = gr.Image(label="Timeline", type="filepath") st_target_audio = gr.Audio(label="PVAD target audio passed to SmartTurn", type="filepath") with gr.Row(): st_cut_selector = gr.Dropdown(label="SmartTurn END cuts", choices=[], value=None) st_cut_audio = gr.Audio(label="Selected cut audio", type="filepath") st_summary = gr.Code(label="Summary JSON", language="json") st_rows = gr.Code(label="Checks", language="json") st_run.click( run_smartturn_only_gradio, inputs=[ st_enroll, st_mic, st_pvad_model, st_pvad_threshold, st_smartturn_model, st_smartturn_custom_model, st_threshold, st_pvad_silence, st_min_active, ], outputs=[st_timeline, st_summary, st_rows, st_target_audio, st_cut_selector, st_cut_audio], ) st_cut_selector.change(select_cut_audio, inputs=[st_cut_selector], outputs=[st_cut_audio]) return demo if __name__ == "__main__": server_port = int(os.environ.get("PORT", os.environ.get("GRADIO_SERVER_PORT", "7860"))) build_app().queue().launch(server_name="0.0.0.0", server_port=server_port)