Spaces:
Sleeping
Sleeping
Download voice_extractor.py from Luminia/audiosplitter_whisper: direct link, hf CLI and curl.
- Browser
- Download file 32 kB
-
https://huggingface.co/spaces/Luminia/audiosplitter_whisper/resolve/main/voice_extractor.py
- Command line
-
hf download hf://spaces/Luminia/audiosplitter_whisper/voice_extractor.py
-
curl -L -o voice_extractor.py https://huggingface.co/spaces/Luminia/audiosplitter_whisper/resolve/main/voice_extractor.py
32 kB
| """ | |
| Voice Extractor CPU - Pipeline module | |
| Identifies, isolates, and transcribes a target speaker from multi-speaker audio. | |
| Based on https://github.com/ReisCook/Voice_Extractor | |
| Heavy imports (torch, torchaudio, pyannote, speechbrain, whisper) are LAZY — | |
| imported inside functions only. This keeps startup fast when only the | |
| audiosplitter tab is used. | |
| """ | |
| import os | |
| import sys | |
| import shutil | |
| import tempfile | |
| import logging | |
| import re | |
| import csv | |
| import time | |
| from pathlib import Path | |
| import numpy as np | |
| # --------------------------------------------------------------------------- | |
| # Force CPU | |
| # --------------------------------------------------------------------------- | |
| os.environ["CUDA_VISIBLE_DEVICES"] = "" | |
| # --------------------------------------------------------------------------- | |
| # Logging | |
| # --------------------------------------------------------------------------- | |
| logging.basicConfig(level=logging.INFO, format="%(asctime)s [%(levelname)s] %(message)s") | |
| log = logging.getLogger("voice_extractor_cpu") | |
| # --------------------------------------------------------------------------- | |
| # Model paths (bundled locally — no HF_TOKEN needed) | |
| # --------------------------------------------------------------------------- | |
| _APP_DIR = os.path.dirname(os.path.abspath(__file__)) | |
| PYANNOTE_DIR = os.path.join(_APP_DIR, "models", "audiosplitter", "pyannote") | |
| BANDIT_ONNX_PATH = os.path.join(_APP_DIR, "models", "voice_extractor", "bandit_v2_speech_fp32.onnx") | |
| HF_TOKEN = os.environ.get("HF_TOKEN") # fallback, not required | |
| # --------------------------------------------------------------------------- | |
| # Intercept hf_hub_download to serve bundled PyAnnote models locally. | |
| # Also fixes pyannote's use_auth_token → token compat issue. | |
| # --------------------------------------------------------------------------- | |
| import huggingface_hub as _hfhub | |
| _orig_hf_hub_download = _hfhub.hf_hub_download | |
| # Map HF repo IDs to local model directories | |
| _LOCAL_MODEL_MAP = { | |
| "pyannote/speaker-diarization-3.1": os.path.join(PYANNOTE_DIR, "speaker-diarization-3.1"), | |
| "pyannote/segmentation-3.0": os.path.join(PYANNOTE_DIR, "segmentation-3.0"), | |
| "pyannote/wespeaker-voxceleb-resnet34-LM": os.path.join(PYANNOTE_DIR, "wespeaker-voxceleb-resnet34-LM"), | |
| "pyannote/segmentation": os.path.join(PYANNOTE_DIR, "segmentation"), | |
| "pyannote/overlapped-speech-detection": os.path.join(PYANNOTE_DIR, "overlapped-speech-detection"), | |
| } | |
| def _patched_hf_hub_download(*args, **kwargs): | |
| # Fix use_auth_token → token | |
| if "use_auth_token" in kwargs: | |
| kwargs["token"] = kwargs.pop("use_auth_token") | |
| # Intercept PyAnnote model downloads → serve from local | |
| repo_id = args[0] if args else kwargs.get("repo_id", "") | |
| filename = args[1] if len(args) > 1 else kwargs.get("filename", "") | |
| if repo_id in _LOCAL_MODEL_MAP: | |
| local_path = os.path.join(_LOCAL_MODEL_MAP[repo_id], filename) | |
| if os.path.exists(local_path): | |
| log.info(f"[LOCAL] {repo_id}/{filename} → {local_path}") | |
| return local_path | |
| else: | |
| log.warning(f"[LOCAL] {repo_id}/{filename} not found locally, falling back to HF") | |
| return _orig_hf_hub_download(*args, **kwargs) | |
| _hfhub.hf_hub_download = _patched_hf_hub_download | |
| # Also patch cached_download if it exists (used by some older pyannote internals) | |
| if hasattr(_hfhub, "cached_download"): | |
| _orig_cached_download = _hfhub.cached_download | |
| def _patched_cached_download(*args, **kwargs): | |
| if "use_auth_token" in kwargs: | |
| kwargs["token"] = kwargs.pop("use_auth_token") | |
| return _orig_cached_download(*args, **kwargs) | |
| _hfhub.cached_download = _patched_cached_download | |
| # --------------------------------------------------------------------------- | |
| # Lazy-loaded models (singleton pattern) | |
| # --------------------------------------------------------------------------- | |
| _models = { | |
| "wespeaker_rvector": None, | |
| "wespeaker_gemini": None, | |
| "speechbrain": None, | |
| "whisper": None, | |
| } | |
| def _get_wespeaker_models(): | |
| """Load WeSpeaker models (cached).""" | |
| if _models["wespeaker_rvector"] is None: | |
| import wespeaker | |
| log.info("Loading WeSpeaker english model (r-vector)...") | |
| m = wespeaker.load_model("english") | |
| m.set_device("cpu") | |
| _models["wespeaker_rvector"] = m | |
| # Use same model for gemini slot on CPU to save memory | |
| _models["wespeaker_gemini"] = m | |
| log.info("WeSpeaker model loaded.") | |
| return {"rvector": _models["wespeaker_rvector"], "gemini": _models["wespeaker_gemini"]} | |
| def _get_speechbrain_model(): | |
| """Load SpeechBrain ECAPA-TDNN (cached).""" | |
| if _models["speechbrain"] is None: | |
| try: | |
| from speechbrain.inference.speaker import SpeakerRecognition | |
| log.info("Loading SpeechBrain ECAPA-TDNN...") | |
| cache_dir = Path(tempfile.gettempdir()) / "voice_extractor_sb_cache" / "spkrec-ecapa-voxceleb" | |
| cache_dir.mkdir(parents=True, exist_ok=True) | |
| model = SpeakerRecognition.from_hparams( | |
| source="speechbrain/spkrec-ecapa-voxceleb", | |
| savedir=str(cache_dir), | |
| run_opts={"device": "cpu"}, | |
| ) | |
| model.eval() | |
| _models["speechbrain"] = model | |
| log.info("SpeechBrain ECAPA-TDNN loaded.") | |
| except Exception as e: | |
| log.warning(f"SpeechBrain load failed: {e}. Verification will use WeSpeaker only.") | |
| _models["speechbrain"] = False # sentinel: tried and failed | |
| if _models["speechbrain"] is False: | |
| return None | |
| return _models["speechbrain"] | |
| def _get_whisper_model(model_name: str = "base"): | |
| """Load Whisper model (cached).""" | |
| if _models["whisper"] is None: | |
| import whisper | |
| log.info(f"Loading Whisper '{model_name}' on CPU...") | |
| _models["whisper"] = whisper.load_model(model_name, device="cpu") | |
| log.info("Whisper model loaded.") | |
| return _models["whisper"] | |
| # --------------------------------------------------------------------------- | |
| # Bandit-v2 vocal separation (ONNX, speech stem only) | |
| # --------------------------------------------------------------------------- | |
| _bandit_session = None | |
| def _get_bandit_session(): | |
| """Load Bandit-v2 ONNX session (cached).""" | |
| global _bandit_session | |
| if _bandit_session is None: | |
| if not os.path.exists(BANDIT_ONNX_PATH): | |
| log.warning(f"Bandit-v2 ONNX not found at {BANDIT_ONNX_PATH}") | |
| return None | |
| import onnxruntime as ort | |
| log.info("Loading Bandit-v2 ONNX model...") | |
| _bandit_session = ort.InferenceSession(BANDIT_ONNX_PATH, providers=["CPUExecutionProvider"]) | |
| log.info("Bandit-v2 ONNX loaded.") | |
| return _bandit_session | |
| def bandit_vocal_separation(input_path: Path, output_path: Path, | |
| chunk_seconds: float = 8.0, fs: int = 48000, | |
| progress_cb=None) -> Path | None: | |
| """Separate vocals from music/SFX using Bandit-v2 ONNX. | |
| STFT/iSTFT in PyTorch, core neural net in ONNX Runtime. | |
| Returns path to separated speech WAV, or None on failure. | |
| """ | |
| import torch | |
| import librosa | |
| import soundfile as sf | |
| sess = _get_bandit_session() | |
| if sess is None: | |
| return None | |
| import torchaudio.transforms as T | |
| if progress_cb: | |
| progress_cb(0.02, "Loading audio for vocal separation...") | |
| # Load and resample to 48kHz mono | |
| y, sr = librosa.load(str(input_path), sr=fs, mono=True) | |
| audio = torch.from_numpy(y).unsqueeze(0) # (1, samples) | |
| n_samples = audio.shape[-1] | |
| log.info(f"Bandit-v2: {n_samples/fs:.1f}s audio @ {fs}Hz") | |
| # STFT params (must match training config) | |
| n_fft, hop_length, win_length = 2048, 512, 2048 | |
| stft = T.Spectrogram(n_fft=n_fft, win_length=win_length, hop_length=hop_length, | |
| power=None, normalized=True, center=True, pad_mode="reflect") | |
| istft = T.InverseSpectrogram(n_fft=n_fft, win_length=win_length, hop_length=hop_length, | |
| normalized=True, center=True) | |
| # Chunked overlap-add inference | |
| chunk_samples = int(chunk_seconds * fs) | |
| hop_samples = fs # 1s hop | |
| overlap = chunk_samples - hop_samples | |
| front_pad = 2 * overlap | |
| window = torch.hann_window(chunk_samples) / (chunk_samples / (2 * hop_samples)) | |
| padded = torch.nn.functional.pad(audio, (front_pad, front_pad), mode="constant") | |
| total_len = padded.shape[-1] | |
| starts = list(range(0, total_len - chunk_samples + 1, hop_samples)) | |
| output_buf = torch.zeros_like(padded) | |
| log.info(f"Bandit-v2: {len(starts)} chunks, {chunk_seconds}s each") | |
| for idx, start in enumerate(starts): | |
| if progress_cb and idx % max(1, len(starts)//10) == 0: | |
| pct = 0.02 + 0.15 * (idx / len(starts)) | |
| progress_cb(pct, f"Vocal separation: chunk {idx+1}/{len(starts)}...") | |
| chunk = padded[:, start:start+chunk_samples] # (1, chunk_samples) | |
| # STFT → complex spectrogram | |
| spec = stft(chunk) # (1, freq, time) complex | |
| spec_real = spec.real.unsqueeze(0).numpy() # (1, 1, freq, time) | |
| spec_imag = spec.imag.unsqueeze(0).numpy() | |
| # ONNX core inference | |
| speech_real, speech_imag = sess.run(None, { | |
| "spec_real": spec_real.astype(np.float32), | |
| "spec_imag": spec_imag.astype(np.float32), | |
| }) | |
| # iSTFT → waveform | |
| masked_spec = torch.complex( | |
| torch.from_numpy(speech_real[0]), # (1, freq, time) | |
| torch.from_numpy(speech_imag[0]), | |
| ) | |
| chunk_out = istft(masked_spec, chunk_samples) # (1, chunk_samples) | |
| output_buf[:, start:start+chunk_samples] += chunk_out * window | |
| # Trim padding | |
| speech_audio = output_buf[:, front_pad:front_pad+n_samples] | |
| # Save | |
| sf.write(str(output_path), speech_audio.squeeze().numpy(), fs) | |
| log.info(f"Bandit-v2: saved {output_path.name}") | |
| if progress_cb: | |
| progress_cb(0.18, "Vocal separation complete.") | |
| return output_path | |
| # --------------------------------------------------------------------------- | |
| # Audio utilities | |
| # --------------------------------------------------------------------------- | |
| def ff_convert(src: Path, dst: Path, sr: int = 16000, ac: int = 1, | |
| start: float | None = None, end: float | None = None): | |
| """Convert/trim audio with ffmpeg.""" | |
| import ffmpeg | |
| inp_kwargs = {} | |
| if start is not None: | |
| inp_kwargs["ss"] = start | |
| if end is not None: | |
| inp_kwargs["to"] = end | |
| ( | |
| ffmpeg.input(str(src), **inp_kwargs) | |
| .output(str(dst), acodec="pcm_s16le", ac=ac, ar=sr) | |
| .overwrite_output() | |
| .run(quiet=True, capture_stdout=True, capture_stderr=True) | |
| ) | |
| def cosine_sim(a, b) -> float: | |
| na, nb = np.linalg.norm(a), np.linalg.norm(b) | |
| if na == 0 or nb == 0: | |
| return 0.0 | |
| return float(np.dot(a, b) / (na * nb)) | |
| # --------------------------------------------------------------------------- | |
| # Pipeline stages | |
| # --------------------------------------------------------------------------- | |
| def prepare_reference(ref_path: Path, tmp_dir: Path) -> Path: | |
| """Convert reference to 16 kHz mono WAV.""" | |
| out = tmp_dir / "reference_16k.wav" | |
| ff_convert(ref_path, out, sr=16000, ac=1) | |
| return out | |
| def diarize(audio_path: Path, tmp_dir: Path, dry_run_sec: int | None = None, | |
| progress_cb=None) -> "Annotation | None": | |
| """Run PyAnnote speaker diarization 3.1 (local models, no HF_TOKEN needed).""" | |
| import torch | |
| from pyannote.audio import Pipeline as PyannotePipeline | |
| DEVICE = torch.device("cpu") | |
| if progress_cb: | |
| progress_cb(0.10, "Loading diarization model...") | |
| # hf_hub_download interceptor serves local models automatically | |
| pipeline = PyannotePipeline.from_pretrained( | |
| "pyannote/speaker-diarization-3.1", use_auth_token=HF_TOKEN or True | |
| ) | |
| # Force CPU | |
| pipeline = pipeline.to(DEVICE) | |
| target = audio_path | |
| if dry_run_sec: | |
| cut = tmp_dir / "diar_cut.wav" | |
| ff_convert(audio_path, cut, sr=16000, ac=1, end=dry_run_sec) | |
| target = cut | |
| if progress_cb: | |
| progress_cb(0.15, "Diarizing speakers...") | |
| result = pipeline({"uri": target.stem, "audio": str(target)}) | |
| n = len(result.labels()) | |
| log.info(f"Diarization found {n} speakers.") | |
| return result | |
| def detect_overlaps(audio_path: Path, tmp_dir: Path, dry_run_sec: int | None = None, | |
| progress_cb=None) -> "Timeline": | |
| """Run PyAnnote overlapped speech detection.""" | |
| import torch | |
| from pyannote.audio import Pipeline as PyannotePipeline | |
| from pyannote.core import Timeline | |
| DEVICE = torch.device("cpu") | |
| if progress_cb: | |
| progress_cb(0.30, "Loading overlap detection model...") | |
| try: | |
| # hf_hub_download interceptor serves local models automatically | |
| osd = PyannotePipeline.from_pretrained( | |
| "pyannote/overlapped-speech-detection", use_auth_token=HF_TOKEN or True | |
| ) | |
| osd = osd.to(DEVICE) | |
| except Exception as e: | |
| log.warning(f"OSD load failed: {e}. Proceeding without overlap detection.") | |
| return Timeline() | |
| target = audio_path | |
| if dry_run_sec: | |
| cut = tmp_dir / "osd_cut.wav" | |
| ff_convert(audio_path, cut, sr=16000, ac=1, end=dry_run_sec) | |
| target = cut | |
| if progress_cb: | |
| progress_cb(0.35, "Detecting overlaps...") | |
| try: | |
| result = osd({"uri": target.stem, "audio": str(target)}) | |
| if isinstance(result, Timeline): | |
| return result.support() | |
| # If Annotation, extract overlap label | |
| from pyannote.core import Annotation | |
| if isinstance(result, Annotation): | |
| tl = Timeline() | |
| if "overlap" in result.labels(): | |
| tl.update(result.label_timeline("overlap")) | |
| return tl.support() | |
| return Timeline() | |
| except Exception as e: | |
| log.warning(f"OSD failed: {e}") | |
| return Timeline() | |
| def identify_target(diar_annotation, audio_path: Path, ref_16k: Path, | |
| target_name: str, progress_cb=None) -> str | None: | |
| """Identify which diarized speaker matches the reference.""" | |
| import ffmpeg | |
| ws_models = _get_wespeaker_models() | |
| ws = ws_models["rvector"] | |
| if ws is None: | |
| return None | |
| if progress_cb: | |
| progress_cb(0.45, "Identifying target speaker...") | |
| ref_emb = ws.extract_embedding(str(ref_16k)) | |
| labels = diar_annotation.labels() | |
| if not labels: | |
| return None | |
| best_label, best_score = None, -1.0 | |
| with tempfile.TemporaryDirectory(prefix="spk_id_") as td: | |
| td_path = Path(td) | |
| for label in labels: | |
| tl = diar_annotation.label_timeline(label) | |
| if not tl: | |
| continue | |
| # Gather up to 20s of audio for this speaker | |
| segs = [] | |
| total = 0.0 | |
| for seg in tl: | |
| if total >= 20.0: | |
| break | |
| seg_path = td_path / f"{label}_{len(segs)}.wav" | |
| try: | |
| ff_convert(audio_path, seg_path, sr=16000, ac=1, | |
| start=seg.start, end=seg.end) | |
| if seg_path.exists() and seg_path.stat().st_size > 0: | |
| segs.append(seg_path) | |
| total += seg.duration | |
| except Exception: | |
| continue | |
| if not segs: | |
| continue | |
| # Concatenate if multiple segments | |
| if len(segs) == 1: | |
| concat_path = segs[0] | |
| else: | |
| concat_path = td_path / f"{label}_concat.wav" | |
| list_file = td_path / f"{label}_list.txt" | |
| list_file.write_text( | |
| "\n".join(f"file '{p.resolve().as_posix()}'" for p in segs) | |
| ) | |
| try: | |
| ( | |
| ffmpeg.input(str(list_file), format="concat", safe=0) | |
| .output(str(concat_path), acodec="pcm_s16le", ar=16000, ac=1) | |
| .overwrite_output() | |
| .run(quiet=True, capture_stdout=True, capture_stderr=True) | |
| ) | |
| except Exception: | |
| concat_path = segs[0] | |
| try: | |
| spk_emb = ws.extract_embedding(str(concat_path)) | |
| sim = cosine_sim(ref_emb, spk_emb) | |
| log.info(f" Speaker {label}: similarity = {sim:.4f}") | |
| if sim > best_score: | |
| best_score = sim | |
| best_label = label | |
| except Exception as e: | |
| log.warning(f"Embedding failed for {label}: {e}") | |
| if best_label: | |
| log.info(f"Identified '{target_name}' as {best_label} (score {best_score:.4f})") | |
| return best_label | |
| def verify_segment(seg_path: Path, ref_path: Path, | |
| use_speechbrain: bool = True) -> tuple[float, dict]: | |
| """Multi-model speaker verification on a single segment.""" | |
| import torch | |
| import librosa | |
| ws_models = _get_wespeaker_models() | |
| scores = {"wespeaker": 0.0, "speechbrain": 0.0} | |
| # WeSpeaker | |
| try: | |
| ws = ws_models["rvector"] | |
| ref_emb = ws.extract_embedding(str(ref_path)) | |
| seg_emb = ws.extract_embedding(str(seg_path)) | |
| scores["wespeaker"] = cosine_sim(ref_emb, seg_emb) | |
| except Exception as e: | |
| log.debug(f"WeSpeaker verify fail: {e}") | |
| # SpeechBrain | |
| if use_speechbrain: | |
| sb = _get_speechbrain_model() | |
| if sb is not None: | |
| try: | |
| score_t, _ = sb.verify_files( | |
| str(ref_path.resolve()).replace("\\", "/"), | |
| str(seg_path.resolve()).replace("\\", "/"), | |
| ) | |
| scores["speechbrain"] = score_t.item() | |
| except Exception as e: | |
| log.debug(f"SpeechBrain verify fail: {e}") | |
| # VAD check | |
| vad_factor = 1.0 | |
| try: | |
| y, sr = librosa.load(seg_path, sr=16000, mono=True) | |
| if len(y) > 0: | |
| vad_model, utils = torch.hub.load( | |
| "snakers4/silero-vad", "silero_vad", | |
| force_reload=False, trust_repo=True, verbose=False, onnx=False, | |
| ) | |
| get_ts = utils[0] | |
| audio_t = torch.FloatTensor(y) | |
| ts = get_ts(audio_t, vad_model, sampling_rate=16000, threshold=0.5) | |
| speech_dur = sum(d["end"] - d["start"] for d in ts) / 16000 | |
| total_dur = len(y) / 16000 | |
| ratio = speech_dur / total_dur if total_dur > 0 else 0 | |
| vad_factor = 1.0 if ratio >= 0.6 else 0.1 | |
| except Exception: | |
| pass | |
| # Combine | |
| if scores["speechbrain"] > 0: | |
| combined = (scores["wespeaker"] * 0.5 + scores["speechbrain"] * 0.5) * vad_factor | |
| else: | |
| combined = scores["wespeaker"] * vad_factor | |
| return combined, scores | |
| def extract_and_verify( | |
| diar_annotation, target_label: str, overlap_tl, | |
| audio_path: Path, ref_16k: Path, target_name: str, | |
| output_dir: Path, tmp_dir: Path, | |
| threshold: float = 0.7, | |
| min_duration: float = 1.0, | |
| merge_gap: float = 0.25, | |
| output_sr: int = 44100, | |
| use_speechbrain: bool = True, | |
| progress_cb=None, | |
| ) -> tuple[list[Path], list[Path]]: | |
| """Slice target solo segments, verify identity, return verified/rejected paths.""" | |
| from pyannote.core import Segment, Timeline | |
| # Get target solo timeline (minus overlaps) | |
| target_tl = diar_annotation.label_timeline(target_label).support() | |
| solo_tl = target_tl.extrude(overlap_tl.support()) if overlap_tl else target_tl | |
| # Merge nearby segments | |
| segs = sorted(list(solo_tl), key=lambda s: s.start) | |
| merged = [] | |
| if segs: | |
| cur = segs[0] | |
| for nxt in segs[1:]: | |
| if nxt.start <= cur.end + merge_gap and nxt.end > cur.end: | |
| cur = Segment(cur.start, nxt.end) | |
| elif nxt.start > cur.end + merge_gap: | |
| merged.append(cur) | |
| cur = nxt | |
| merged.append(cur) | |
| # Duration filter | |
| merged = [s for s in merged if s.duration >= min_duration] | |
| if not merged: | |
| return [], [] | |
| verified_dir = output_dir / "verified" | |
| rejected_dir = output_dir / "rejected" | |
| verified_dir.mkdir(parents=True, exist_ok=True) | |
| rejected_dir.mkdir(parents=True, exist_ok=True) | |
| verif_tmp = tmp_dir / "verif_16k" | |
| hq_tmp = tmp_dir / "hq" | |
| verif_tmp.mkdir(parents=True, exist_ok=True) | |
| hq_tmp.mkdir(parents=True, exist_ok=True) | |
| verified_paths = [] | |
| rejected_paths = [] | |
| total = len(merged) | |
| for i, seg in enumerate(merged): | |
| if progress_cb: | |
| pct = 0.55 + 0.25 * (i / max(total, 1)) | |
| progress_cb(pct, f"Verifying segment {i+1}/{total}...") | |
| s_str = f"{seg.start:.3f}".replace(".", "p") | |
| e_str = f"{seg.end:.3f}".replace(".", "p") | |
| base = f"{i:04d}_{s_str}s_to_{e_str}s" | |
| seg_16k = verif_tmp / f"{base}.wav" | |
| seg_hq = hq_tmp / f"{base}_hq.wav" | |
| try: | |
| ff_convert(audio_path, seg_16k, sr=16000, ac=1, start=seg.start, end=seg.end) | |
| ff_convert(audio_path, seg_hq, sr=output_sr, ac=1, start=seg.start, end=seg.end) | |
| except Exception as e: | |
| log.warning(f"Slice failed for segment {i}: {e}") | |
| continue | |
| if not seg_16k.exists() or seg_16k.stat().st_size == 0: | |
| continue | |
| score, _ = verify_segment(seg_16k, ref_16k, use_speechbrain=use_speechbrain) | |
| safe_name = re.sub(r'[<>:"/\\|?*]', "", target_name).replace(" ", "_") | |
| if score >= threshold: | |
| dst = verified_dir / f"{safe_name}_verified_{base}.wav" | |
| shutil.copy(seg_hq, dst) | |
| verified_paths.append(dst) | |
| else: | |
| dst = rejected_dir / f"{safe_name}_rejected_{base}_score{score:.3f}.wav" | |
| shutil.copy(seg_hq, dst) | |
| rejected_paths.append(dst) | |
| # Cleanup temp | |
| seg_16k.unlink(missing_ok=True) | |
| seg_hq.unlink(missing_ok=True) | |
| log.info(f"Verified: {len(verified_paths)}, Rejected: {len(rejected_paths)}") | |
| return verified_paths, rejected_paths | |
| def transcribe_files( | |
| segment_paths: list[Path], | |
| whisper_model_name: str = "base", | |
| language: str = "en", | |
| progress_cb=None, | |
| ) -> list[dict]: | |
| """Transcribe audio segments with Whisper. Returns list of dicts.""" | |
| import librosa | |
| if not segment_paths: | |
| return [] | |
| model = _get_whisper_model(whisper_model_name) | |
| results = [] | |
| total = len(segment_paths) | |
| for i, p in enumerate(segment_paths): | |
| if progress_cb: | |
| pct = 0.82 + 0.15 * (i / max(total, 1)) | |
| progress_cb(pct, f"Transcribing {i+1}/{total}...") | |
| if not p.exists() or p.stat().st_size == 0: | |
| continue | |
| try: | |
| r = model.transcribe(str(p), fp16=False, language=language if language != "auto" else None) | |
| text = r["text"].strip() | |
| except Exception as e: | |
| text = f"[error: {e}]" | |
| log.warning(f"Transcribe fail for {p.name}: {e}") | |
| dur = librosa.get_duration(path=p) | |
| results.append({"file": p.name, "duration_s": round(dur, 2), "text": text}) | |
| return results | |
| def concatenate_verified(paths: list[Path], output_path: Path, | |
| silence_s: float = 0.25, sr: int = 44100): | |
| """Concatenate verified segments with silence gaps.""" | |
| import ffmpeg | |
| import soundfile as sf | |
| if not paths: | |
| return None | |
| # Sort by segment start time from filename | |
| def sort_key(p): | |
| m = re.search(r"(\d+p\d+)s_to_", p.name) | |
| if m: | |
| return float(m.group(1).replace("p", ".")) | |
| return 0.0 | |
| paths = sorted(paths, key=sort_key) | |
| with tempfile.TemporaryDirectory(prefix="concat_") as td: | |
| td_path = Path(td) | |
| # Create silence file | |
| silence_file = td_path / "silence.wav" | |
| if silence_s > 0: | |
| ( | |
| ffmpeg.input(f"anullsrc=channel_layout=mono:sample_rate={sr}", | |
| format="lavfi", t=str(silence_s)) | |
| .output(str(silence_file), acodec="pcm_s16le", ar=sr, ac=1) | |
| .overwrite_output() | |
| .run(quiet=True, capture_stdout=True, capture_stderr=True) | |
| ) | |
| # Build concat list | |
| lines = [] | |
| for i, p in enumerate(paths): | |
| if i > 0 and silence_s > 0 and silence_file.exists(): | |
| lines.append(f"file '{silence_file.resolve().as_posix()}'") | |
| lines.append(f"file '{p.resolve().as_posix()}'") | |
| list_file = td_path / "list.txt" | |
| list_file.write_text("\n".join(lines)) | |
| try: | |
| ( | |
| ffmpeg.input(str(list_file), format="concat", safe=0) | |
| .output(str(output_path), acodec="pcm_s16le", ar=sr, ac=1) | |
| .overwrite_output() | |
| .run(quiet=True, capture_stdout=True, capture_stderr=True) | |
| ) | |
| return output_path | |
| except Exception as e: | |
| log.error(f"Concat failed: {e}") | |
| return None | |
| # --------------------------------------------------------------------------- | |
| # Main pipeline (entry point for app.py) | |
| # --------------------------------------------------------------------------- | |
| def run_voice_extractor_pipeline( | |
| input_audio_path: str, | |
| reference_audio_path: str, | |
| target_name: str, | |
| whisper_model: str = "base", | |
| language: str = "en", | |
| verification_threshold: float = 0.7, | |
| min_duration: float = 1.0, | |
| merge_gap: float = 0.25, | |
| output_sr: int = 44100, | |
| use_speechbrain: bool = True, | |
| dry_run: bool = False, | |
| use_bandit: bool = False, | |
| progress=None, | |
| ): | |
| """Full voice extraction pipeline. Returns (status, audio_preview, output_files, transcript_text).""" | |
| import gradio as gr | |
| if not input_audio_path: | |
| raise gr.Error("Please upload an input audio file.") | |
| if not reference_audio_path: | |
| raise gr.Error("Please upload a reference audio clip of the target speaker.") | |
| if not target_name or not target_name.strip(): | |
| raise gr.Error("Please enter a target speaker name.") | |
| start_time = time.time() | |
| def progress_cb(pct, msg): | |
| try: | |
| if progress is not None: | |
| progress(pct, desc=msg) | |
| except Exception: | |
| pass | |
| progress_cb(0.02, "Setting up...") | |
| input_p = Path(input_audio_path) | |
| ref_p = Path(reference_audio_path) | |
| work_dir = Path(tempfile.mkdtemp(prefix="voice_ext_")) | |
| tmp_dir = work_dir / "tmp" | |
| output_dir = work_dir / "output" | |
| tmp_dir.mkdir(parents=True, exist_ok=True) | |
| output_dir.mkdir(parents=True, exist_ok=True) | |
| dry_sec = 60 if dry_run else None | |
| try: | |
| # Stage 0a: Trim audio for dry run BEFORE any processing | |
| actual_input = input_p | |
| if dry_sec: | |
| trimmed = tmp_dir / f"{input_p.stem}_trimmed.wav" | |
| ff_convert(input_p, trimmed, sr=16000, ac=1, end=dry_sec) | |
| actual_input = trimmed | |
| log.info(f"Dry run: trimmed to {dry_sec}s") | |
| # Stage 0b (optional): Bandit-v2 vocal separation | |
| source_audio = actual_input | |
| if use_bandit: | |
| progress_cb(0.02, "Starting vocal separation (Bandit-v2, slow on CPU)...") | |
| vocals_path = tmp_dir / f"{input_p.stem}_vocals.wav" | |
| result = bandit_vocal_separation(actual_input, vocals_path, progress_cb=progress_cb) | |
| if result and result.exists(): | |
| source_audio = result | |
| log.info(f"Using Bandit-v2 vocals for downstream: {result.name}") | |
| else: | |
| log.warning("Bandit-v2 failed, using original audio.") | |
| # Stage 1: Prepare reference | |
| progress_cb(0.05, "Preparing reference audio...") | |
| ref_16k = prepare_reference(ref_p, tmp_dir) | |
| # Stage 2: Diarization (audio already trimmed if dry_run) | |
| progress_cb(0.08, "Starting diarization...") | |
| diar = diarize(source_audio, tmp_dir, progress_cb=progress_cb) | |
| if diar is None or not diar.labels(): | |
| raise gr.Error("Diarization found no speakers. Check your audio file.") | |
| # Stage 3: Overlap detection | |
| progress_cb(0.28, "Detecting overlapping speech...") | |
| overlap_tl = detect_overlaps(source_audio, tmp_dir, progress_cb=progress_cb) | |
| # Stage 4: Identify target | |
| progress_cb(0.42, "Identifying target speaker...") | |
| target_label = identify_target(diar, source_audio, ref_16k, target_name, | |
| progress_cb=progress_cb) | |
| if not target_label: | |
| raise gr.Error( | |
| f"Could not identify '{target_name}' among diarized speakers. " | |
| "Try a cleaner/longer reference clip." | |
| ) | |
| # Stage 5: Extract & verify | |
| progress_cb(0.50, "Extracting and verifying segments...") | |
| verified, rejected = extract_and_verify( | |
| diar, target_label, overlap_tl, | |
| source_audio, ref_16k, target_name, | |
| output_dir, tmp_dir, | |
| threshold=verification_threshold, | |
| min_duration=min_duration, | |
| merge_gap=merge_gap, | |
| output_sr=output_sr, | |
| use_speechbrain=use_speechbrain, | |
| progress_cb=progress_cb, | |
| ) | |
| if not verified and not rejected: | |
| raise gr.Error("No speech segments found for the target speaker.") | |
| # Stage 6: Transcribe verified segments | |
| progress_cb(0.80, "Transcribing verified segments...") | |
| transcripts = transcribe_files( | |
| verified, whisper_model_name=whisper_model, | |
| language=language, progress_cb=progress_cb, | |
| ) | |
| # Stage 7: Concatenate verified | |
| concat_path = None | |
| if verified: | |
| progress_cb(0.96, "Concatenating verified segments...") | |
| concat_path = output_dir / "concatenated_verified.wav" | |
| concatenate_verified(verified, concat_path, silence_s=0.25, sr=output_sr) | |
| # Build outputs | |
| progress_cb(0.98, "Packaging results...") | |
| # Collect all output files for download | |
| output_files = [] | |
| if concat_path and concat_path.exists(): | |
| output_files.append(str(concat_path)) | |
| for p in sorted(verified): | |
| output_files.append(str(p)) | |
| # Transcript text | |
| transcript_lines = [] | |
| for t in transcripts: | |
| transcript_lines.append(f"[{t['file']}] ({t['duration_s']}s): {t['text']}") | |
| transcript_text = "\n\n".join(transcript_lines) if transcript_lines else "No transcripts generated." | |
| # Save transcript CSV | |
| if transcripts: | |
| csv_path = output_dir / "transcripts.csv" | |
| with csv_path.open("w", newline="", encoding="utf-8") as f: | |
| w = csv.writer(f) | |
| w.writerow(["filename", "duration_s", "transcript"]) | |
| for t in transcripts: | |
| w.writerow([t["file"], t["duration_s"], t["text"]]) | |
| output_files.append(str(csv_path)) | |
| elapsed = time.time() - start_time | |
| status = ( | |
| f"Done in {elapsed:.0f}s. " | |
| f"Found {len(diar.labels())} speakers. " | |
| f"Identified '{target_name}' as {target_label}. " | |
| f"Verified: {len(verified)} segments, Rejected: {len(rejected)} segments." | |
| ) | |
| # Audio preview = concatenated file if available, else first verified segment | |
| audio_preview = None | |
| if concat_path and concat_path.exists(): | |
| audio_preview = str(concat_path) | |
| elif verified: | |
| audio_preview = str(verified[0]) | |
| progress_cb(1.0, "Complete!") | |
| return status, audio_preview, output_files, transcript_text | |
| except gr.Error: | |
| raise | |
| except Exception as e: | |
| log.exception("Pipeline error") | |
| raise gr.Error(f"Pipeline error: {e}") | |