Spaces:
Sleeping
Sleeping
Download app.py from Luminia/audiosplitter_whisper: direct link, hf CLI and curl.
- Browser
- Download file 39.1 kB
-
https://huggingface.co/spaces/Luminia/audiosplitter_whisper/resolve/main/app.py
- Command line
-
hf download hf://spaces/Luminia/audiosplitter_whisper/app.py
-
curl -L -o app.py https://huggingface.co/spaces/Luminia/audiosplitter_whisper/resolve/main/app.py
39.1 kB
| """ | |
| Audio Splitter with Speaker Diarization | |
| Two speaker identification modes: | |
| 1. Diarization (default) - Labels who speaks when, best for any number of speakers | |
| 2. Speech Separation - Physically separates audio into 2 clean tracks (2 speakers only) | |
| No pyannote/HF token required - pure ONNX inference. | |
| """ | |
| import os | |
| import re | |
| import tempfile | |
| import time | |
| import unicodedata | |
| import zipfile | |
| import shutil | |
| from pathlib import Path | |
| from typing import List, Dict, Tuple, Optional | |
| from collections import defaultdict | |
| import numpy as np | |
| import gradio as gr | |
| from pydub import AudioSegment | |
| from onnx_diarize_v2 import ONNXDiarizerV2 | |
| # Optional imports | |
| try: | |
| from faster_whisper import WhisperModel | |
| FASTER_WHISPER_AVAILABLE = True | |
| except ImportError: | |
| FASTER_WHISPER_AVAILABLE = False | |
| print("[Warning] faster-whisper not installed") | |
| try: | |
| from speech_separation import DualPathRNNSeparator | |
| DPRNN_AVAILABLE = True | |
| except ImportError: | |
| DPRNN_AVAILABLE = False | |
| print("[Warning] DPRNN speech separation not available") | |
| try: | |
| from mossformer2_separation import MossFormer2Separator | |
| MOSSFORMER2_AVAILABLE = True | |
| except ImportError: | |
| MOSSFORMER2_AVAILABLE = False | |
| print("[Warning] MossFormer2 speech separation not available") | |
| # Unified separation availability | |
| SEPARATION_AVAILABLE = DPRNN_AVAILABLE or MOSSFORMER2_AVAILABLE | |
| def sanitize_filename(text: str, max_length: int = 50) -> str: | |
| """Sanitize text for use as filename.""" | |
| if not text: | |
| return "segment" | |
| text = unicodedata.normalize('NFKC', text) | |
| text = re.sub(r'[<>:"/\\|?*]', '_', text) | |
| text = ''.join(c for c in text if unicodedata.category(c) != 'Cc') | |
| text = text[:max_length] | |
| text = re.sub(r'[\s.]+$', '', text) | |
| text = re.sub(r'^[\s.]+', '', text) | |
| return text if text else "segment" | |
| # Global model cache | |
| _whisper_cache = {} | |
| _diarizer_cache = {} | |
| _separator_cache = {} | |
| def get_device_info(): | |
| """Detect available device.""" | |
| try: | |
| import torch | |
| if torch.cuda.is_available(): | |
| return "cuda", "float16" | |
| except ImportError: | |
| pass | |
| return "cpu", "int8" | |
| def load_whisper_model(model_size: str, device: str, compute_type: str): | |
| """Load faster-whisper model with caching.""" | |
| cache_key = f"{model_size}_{device}_{compute_type}" | |
| if cache_key not in _whisper_cache: | |
| print(f"[Whisper] Loading {model_size} on {device} ({compute_type})...") | |
| _whisper_cache[cache_key] = WhisperModel( | |
| model_size, | |
| device=device, | |
| compute_type=compute_type | |
| ) | |
| return _whisper_cache[cache_key] | |
| def load_diarizer(): | |
| """Load ONNX diarizer v2 with caching.""" | |
| if "diarizer" not in _diarizer_cache: | |
| base_dir = os.path.dirname(__file__) | |
| # Primary: models/audiosplitter/ layout | |
| seg_model = os.path.join(base_dir, "models", "audiosplitter", "pyannote", | |
| "sherpa-onnx-segmentation-3-0", "model.onnx") | |
| emb_model = os.path.join(base_dir, "models", "audiosplitter", "3dspeaker_eres2net.onnx") | |
| # Fallback: flat layout (legacy) | |
| if not os.path.exists(seg_model): | |
| seg_model = os.path.join(base_dir, "sherpa-onnx-pyannote-segmentation-3-0", "model.onnx") | |
| if not os.path.exists(emb_model): | |
| emb_model = os.path.join(base_dir, "3dspeaker_speech_eres2net_base_sv_zh-cn_3dspeaker_16k.onnx") | |
| # Fallback: models/ directory (legacy) | |
| if not os.path.exists(seg_model): | |
| seg_model = os.path.join(base_dir, "models", "segmentation.onnx") | |
| if not os.path.exists(emb_model): | |
| emb_model = os.path.join(base_dir, "models", "embedding.onnx") | |
| if os.path.exists(seg_model) and os.path.exists(emb_model): | |
| print(f"[Diarizer] Loading ONNX v2 models...") | |
| _diarizer_cache["diarizer"] = ONNXDiarizerV2(seg_model, emb_model) | |
| else: | |
| print("[Warning] ONNX diarization models not found!") | |
| print(f" Expected: {seg_model}") | |
| print(f" Expected: {emb_model}") | |
| _diarizer_cache["diarizer"] = None | |
| return _diarizer_cache["diarizer"] | |
| def load_separator(): | |
| """Load speech separator with caching.""" | |
| if "separator" not in _separator_cache: | |
| if not SEPARATION_AVAILABLE: | |
| print("[Warning] Speech separation not available") | |
| _separator_cache["separator"] = None | |
| else: | |
| base_dir = os.path.dirname(__file__) | |
| onnx_path = os.path.join(base_dir, "models", "audiosplitter", "dprnn_separator.onnx") | |
| pt_path = os.path.join(base_dir, "dual_path_rnn", "Dual-Path-RNN-portable", | |
| "Dual-Path-RNN", "Dual-Path-RNN-model-best.pt") | |
| # Fallback: flat layout (legacy) | |
| if not os.path.exists(onnx_path): | |
| onnx_path = os.path.join(base_dir, "dprnn_separator.onnx") | |
| if os.path.exists(onnx_path): | |
| print(f"[Separator] Loading ONNX model...") | |
| _separator_cache["separator"] = DualPathRNNSeparator( | |
| model_path=onnx_path, use_onnx=True | |
| ) | |
| elif os.path.exists(pt_path): | |
| print(f"[Separator] Loading PyTorch model...") | |
| _separator_cache["separator"] = DualPathRNNSeparator( | |
| model_path=pt_path, use_onnx=False | |
| ) | |
| else: | |
| print("[Warning] Speech separation model not found!") | |
| print(f" Expected: {onnx_path} or {pt_path}") | |
| _separator_cache["separator"] = None | |
| return _separator_cache["separator"] | |
| def load_mossformer2_separator(model_name: str = "mossformer2-whamr-2spk"): | |
| """Load MossFormer2 speech separator with caching.""" | |
| cache_key = f"mossformer2_{model_name}" | |
| if cache_key not in _separator_cache: | |
| if not MOSSFORMER2_AVAILABLE: | |
| print("[Warning] MossFormer2 not available") | |
| _separator_cache[cache_key] = None | |
| else: | |
| print(f"[MossFormer2] Loading {model_name}...") | |
| _separator_cache[cache_key] = MossFormer2Separator(model_name=model_name) | |
| return _separator_cache[cache_key] | |
| def load_audio_for_diarization(audio_path: str) -> np.ndarray: | |
| """Load audio as numpy array at 16kHz mono for diarization.""" | |
| audio = AudioSegment.from_file(audio_path) | |
| audio = audio.set_frame_rate(16000).set_channels(1) | |
| samples = np.array(audio.get_array_of_samples(), dtype=np.float32) | |
| samples = samples / 32768.0 # Normalize int16 to float | |
| return samples | |
| def transcribe_audio( | |
| audio_path: str, | |
| model: WhisperModel, | |
| language: str = "auto" | |
| ) -> List[Dict]: | |
| """Transcribe audio with faster-whisper.""" | |
| lang = None if language == "auto" else language | |
| segments, info = model.transcribe( | |
| audio_path, | |
| language=lang, | |
| beam_size=1, # Faster decoding (was 5) | |
| word_timestamps=True, | |
| vad_filter=True, | |
| vad_parameters=dict(min_silence_duration_ms=500) | |
| ) | |
| result = [] | |
| for seg in segments: | |
| result.append({ | |
| "start": seg.start, | |
| "end": seg.end, | |
| "text": seg.text.strip(), | |
| }) | |
| print(f"[Whisper] Detected language: {info.language}, {len(result)} segments") | |
| return result | |
| def assign_speakers_to_segments(transcript_segments: List[Dict], diar_segments) -> List[Dict]: | |
| """Assign speakers to transcript segments based on time overlap.""" | |
| for seg in transcript_segments: | |
| seg_start = seg["start"] | |
| seg_end = seg["end"] | |
| # Find overlapping diarization segments | |
| speaker_times = defaultdict(float) | |
| for diar_seg in diar_segments: | |
| overlap_start = max(seg_start, diar_seg.start) | |
| overlap_end = min(seg_end, diar_seg.end) | |
| overlap = max(0, overlap_end - overlap_start) | |
| if overlap > 0: | |
| speaker_times[diar_seg.speaker] += overlap | |
| # Assign speaker with most overlap | |
| if speaker_times: | |
| seg["speaker"] = max(speaker_times.keys(), key=lambda k: speaker_times[k]) | |
| else: | |
| seg["speaker"] = "SPEAKER_00" | |
| return transcript_segments | |
| def split_audio_with_diarization( | |
| audio_path: str, | |
| segments: List[Dict], | |
| output_dir: str, | |
| diarizer: Optional[ONNXDiarizerV2], | |
| min_duration: float = 0.5, | |
| max_duration: float = 10.0, | |
| padding_ms: int = 200 | |
| ) -> Tuple[List[Dict], str]: | |
| """ | |
| Split audio by segments and organize by speaker. | |
| Args: | |
| padding_ms: Milliseconds of audio to include before/after each segment | |
| to prevent word clipping (default: 200ms) | |
| """ | |
| audio = AudioSegment.from_file(audio_path) | |
| audio_duration_ms = len(audio) | |
| # Run diarization if available | |
| if diarizer: | |
| print("[Diarization] Running ONNX v2 speaker diarization...") | |
| audio_np = load_audio_for_diarization(audio_path) | |
| diar_segments = diarizer.diarize(audio_np) | |
| print(f"[Diarization] Found {len(diar_segments)} speaker segments") | |
| segments = assign_speakers_to_segments(segments, diar_segments) | |
| else: | |
| for seg in segments: | |
| seg["speaker"] = "SPEAKER_00" | |
| # Group segments by speaker | |
| by_speaker = defaultdict(list) | |
| for seg in segments: | |
| duration = seg["end"] - seg["start"] | |
| if min_duration <= duration <= max_duration: | |
| by_speaker[seg.get("speaker", "SPEAKER_00")].append(seg) | |
| results = [] | |
| transcript_lines = [] | |
| for speaker in sorted(by_speaker.keys()): | |
| speaker_dir = os.path.join(output_dir, speaker) | |
| os.makedirs(speaker_dir, exist_ok=True) | |
| for i, seg in enumerate(by_speaker[speaker]): | |
| # Add padding to prevent word clipping, with bounds checking | |
| start_ms = max(0, int(seg["start"] * 1000) - padding_ms) | |
| end_ms = min(audio_duration_ms, int(seg["end"] * 1000) + padding_ms) | |
| text = seg.get("text", "").strip() | |
| segment_audio = audio[start_ms:end_ms] | |
| audio_filename = f"{i+1:04d}.wav" | |
| audio_path_out = os.path.join(speaker_dir, audio_filename) | |
| segment_audio.export(audio_path_out, format="wav") | |
| txt_filename = f"{i+1:04d}.txt" | |
| txt_path = os.path.join(speaker_dir, txt_filename) | |
| with open(txt_path, "w", encoding="utf-8") as f: | |
| f.write(text) | |
| results.append({ | |
| "speaker": speaker, | |
| "audio": audio_path_out, | |
| "text": text, | |
| "start": seg["start"], | |
| "end": seg["end"] | |
| }) | |
| transcript_lines.append(f"[{speaker}] {audio_filename}: {text}") | |
| transcript_text = "\n".join(transcript_lines) | |
| return results, transcript_text | |
| def process_with_separation( | |
| audio_path: str, | |
| model: WhisperModel, | |
| language: str, | |
| separator: DualPathRNNSeparator, | |
| output_dir: str, | |
| min_duration: float, | |
| max_duration: float, | |
| padding_ms: int = 200, | |
| progress_fn=None | |
| ) -> Tuple[List[Dict], str]: | |
| """ | |
| Process audio using speech separation (2 speakers only). | |
| Pipeline: | |
| 1. Separate audio into 2 speaker tracks | |
| 2. Run ASR on each track | |
| 3. Save segments organized by speaker | |
| """ | |
| if progress_fn: | |
| progress_fn(0.2, desc="Separating speakers...") | |
| # Load audio | |
| audio = AudioSegment.from_file(audio_path) | |
| audio = audio.set_channels(1) # Mono | |
| sample_rate = audio.frame_rate | |
| audio_duration_ms = len(audio) | |
| # Convert to numpy | |
| audio_np = np.array(audio.get_array_of_samples(), dtype=np.float32) | |
| audio_np = audio_np / 32768.0 # Normalize int16 | |
| # Separate | |
| sep_start = time.time() | |
| spk1_audio, spk2_audio = separator.separate(audio_np, sample_rate) | |
| sep_time = time.time() - sep_start | |
| print(f"[Separation]: {sep_time:.1f}s") | |
| if progress_fn: | |
| progress_fn(0.4, desc="Transcribing speaker 1...") | |
| # Create temp files for each speaker | |
| temp_dir = tempfile.mkdtemp() | |
| spk1_path = os.path.join(temp_dir, "spk1.wav") | |
| spk2_path = os.path.join(temp_dir, "spk2.wav") | |
| # Save speaker audio files | |
| for path, audio_data in [(spk1_path, spk1_audio), (spk2_path, spk2_audio)]: | |
| audio_int16 = (audio_data * 32768).astype(np.int16) | |
| seg = AudioSegment( | |
| audio_int16.tobytes(), | |
| frame_rate=sample_rate, | |
| sample_width=2, | |
| channels=1 | |
| ) | |
| seg.export(path, format="wav") | |
| # Transcribe each speaker | |
| lang = None if language == "auto" else language | |
| asr_start = time.time() | |
| segs1, info1 = model.transcribe(spk1_path, language=lang, beam_size=1, | |
| word_timestamps=True, vad_filter=True, | |
| vad_parameters=dict(min_silence_duration_ms=500)) | |
| segments_spk1 = [{"start": s.start, "end": s.end, "text": s.text.strip(), "speaker": "SPEAKER_00"} | |
| for s in segs1] | |
| print(f"[ASR Speaker 1]: {len(segments_spk1)} segments") | |
| if progress_fn: | |
| progress_fn(0.6, desc="Transcribing speaker 2...") | |
| segs2, info2 = model.transcribe(spk2_path, language=lang, beam_size=1, | |
| word_timestamps=True, vad_filter=True, | |
| vad_parameters=dict(min_silence_duration_ms=500)) | |
| segments_spk2 = [{"start": s.start, "end": s.end, "text": s.text.strip(), "speaker": "SPEAKER_01"} | |
| for s in segs2] | |
| asr_time = time.time() - asr_start | |
| print(f"[ASR Speaker 2]: {len(segments_spk2)} segments (total ASR: {asr_time:.1f}s)") | |
| # Clean up temp files | |
| os.remove(spk1_path) | |
| os.remove(spk2_path) | |
| os.rmdir(temp_dir) | |
| if progress_fn: | |
| progress_fn(0.8, desc="Splitting audio by speaker...") | |
| # Combine and sort by time | |
| all_segments = segments_spk1 + segments_spk2 | |
| all_segments.sort(key=lambda x: x["start"]) | |
| # Load original audio for splitting | |
| orig_audio = AudioSegment.from_file(audio_path) | |
| # Group segments by speaker and filter by duration | |
| by_speaker = defaultdict(list) | |
| for seg in all_segments: | |
| duration = seg["end"] - seg["start"] | |
| if min_duration <= duration <= max_duration: | |
| by_speaker[seg["speaker"]].append(seg) | |
| results = [] | |
| transcript_lines = [] | |
| for speaker in sorted(by_speaker.keys()): | |
| speaker_dir = os.path.join(output_dir, speaker) | |
| os.makedirs(speaker_dir, exist_ok=True) | |
| for i, seg in enumerate(by_speaker[speaker]): | |
| start_ms = max(0, int(seg["start"] * 1000) - padding_ms) | |
| end_ms = min(audio_duration_ms, int(seg["end"] * 1000) + padding_ms) | |
| text = seg.get("text", "").strip() | |
| segment_audio = orig_audio[start_ms:end_ms] | |
| audio_filename = f"{i+1:04d}.wav" | |
| audio_path_out = os.path.join(speaker_dir, audio_filename) | |
| segment_audio.export(audio_path_out, format="wav") | |
| txt_filename = f"{i+1:04d}.txt" | |
| txt_path = os.path.join(speaker_dir, txt_filename) | |
| with open(txt_path, "w", encoding="utf-8") as f: | |
| f.write(text) | |
| results.append({ | |
| "speaker": speaker, | |
| "audio": audio_path_out, | |
| "text": text, | |
| "start": seg["start"], | |
| "end": seg["end"] | |
| }) | |
| transcript_lines.append(f"[{speaker}] {audio_filename}: {text}") | |
| transcript_text = "\n".join(transcript_lines) | |
| return results, transcript_text | |
| def process_with_mossformer2( | |
| audio_path: str, | |
| model: WhisperModel, | |
| language: str, | |
| separator: MossFormer2Separator, | |
| output_dir: str, | |
| min_duration: float, | |
| max_duration: float, | |
| padding_ms: int = 200, | |
| progress_fn=None | |
| ) -> Tuple[List[Dict], str]: | |
| """ | |
| Process audio using MossFormer2 speech separation (2 speakers, SOTA quality). | |
| Pipeline: | |
| 1. Separate audio into 2 speaker tracks using MossFormer2 | |
| 2. Run ASR on each track | |
| 3. Save segments organized by speaker | |
| """ | |
| if progress_fn: | |
| progress_fn(0.2, desc="Separating speakers (MossFormer2)...") | |
| # Load audio | |
| audio = AudioSegment.from_file(audio_path) | |
| audio = audio.set_channels(1) # Mono | |
| sample_rate = audio.frame_rate | |
| audio_duration_ms = len(audio) | |
| # Convert to numpy | |
| audio_np = np.array(audio.get_array_of_samples(), dtype=np.float32) | |
| audio_np = audio_np / 32768.0 # Normalize int16 | |
| # Separate using MossFormer2 | |
| sep_start = time.time() | |
| speaker_audios = separator.separate(audio_np, sample_rate) | |
| sep_time = time.time() - sep_start | |
| print(f"[MossFormer2 Separation]: {sep_time:.1f}s, {len(speaker_audios)} speakers") | |
| if progress_fn: | |
| progress_fn(0.4, desc="Transcribing speaker 1...") | |
| # Create temp files for each speaker | |
| temp_dir = tempfile.mkdtemp() | |
| spk_paths = [] | |
| for i, spk_audio in enumerate(speaker_audios): | |
| path = os.path.join(temp_dir, f"spk{i+1}.wav") | |
| audio_int16 = (spk_audio * 32768).clip(-32768, 32767).astype(np.int16) | |
| seg = AudioSegment( | |
| audio_int16.tobytes(), | |
| frame_rate=sample_rate, | |
| sample_width=2, | |
| channels=1 | |
| ) | |
| seg.export(path, format="wav") | |
| spk_paths.append(path) | |
| # Transcribe each speaker | |
| lang = None if language == "auto" else language | |
| all_segments = [] | |
| asr_start = time.time() | |
| for i, spk_path in enumerate(spk_paths): | |
| if progress_fn: | |
| progress_fn(0.4 + 0.2 * i, desc=f"Transcribing speaker {i+1}...") | |
| segs, info = model.transcribe(spk_path, language=lang, beam_size=1, | |
| word_timestamps=True, vad_filter=True, | |
| vad_parameters=dict(min_silence_duration_ms=500)) | |
| segments_spk = [{"start": s.start, "end": s.end, "text": s.text.strip(), | |
| "speaker": f"SPEAKER_{i:02d}"} for s in segs] | |
| print(f"[ASR Speaker {i+1}]: {len(segments_spk)} segments") | |
| all_segments.extend(segments_spk) | |
| asr_time = time.time() - asr_start | |
| print(f"[ASR Total]: {asr_time:.1f}s") | |
| # Clean up temp files | |
| for path in spk_paths: | |
| os.remove(path) | |
| os.rmdir(temp_dir) | |
| if progress_fn: | |
| progress_fn(0.8, desc="Splitting audio by speaker...") | |
| # Sort by time | |
| all_segments.sort(key=lambda x: x["start"]) | |
| # Load original audio for splitting | |
| orig_audio = AudioSegment.from_file(audio_path) | |
| # Group segments by speaker and filter by duration | |
| by_speaker = defaultdict(list) | |
| for seg in all_segments: | |
| duration = seg["end"] - seg["start"] | |
| if min_duration <= duration <= max_duration: | |
| by_speaker[seg["speaker"]].append(seg) | |
| results = [] | |
| transcript_lines = [] | |
| for speaker in sorted(by_speaker.keys()): | |
| speaker_dir = os.path.join(output_dir, speaker) | |
| os.makedirs(speaker_dir, exist_ok=True) | |
| for i, seg in enumerate(by_speaker[speaker]): | |
| start_ms = max(0, int(seg["start"] * 1000) - padding_ms) | |
| end_ms = min(audio_duration_ms, int(seg["end"] * 1000) + padding_ms) | |
| text = seg.get("text", "").strip() | |
| segment_audio = orig_audio[start_ms:end_ms] | |
| audio_filename = f"{i+1:04d}.wav" | |
| audio_path_out = os.path.join(speaker_dir, audio_filename) | |
| segment_audio.export(audio_path_out, format="wav") | |
| txt_filename = f"{i+1:04d}.txt" | |
| txt_path = os.path.join(speaker_dir, txt_filename) | |
| with open(txt_path, "w", encoding="utf-8") as f: | |
| f.write(text) | |
| results.append({ | |
| "speaker": speaker, | |
| "audio": audio_path_out, | |
| "text": text, | |
| "start": seg["start"], | |
| "end": seg["end"] | |
| }) | |
| transcript_lines.append(f"[{speaker}] {audio_filename}: {text}") | |
| transcript_text = "\n".join(transcript_lines) | |
| return results, transcript_text | |
| def process_audio( | |
| audio_file, | |
| model_size: str = "large-v3", | |
| language: str = "auto", | |
| speaker_mode: str = "diarization", | |
| min_duration: float = 0.5, | |
| max_duration: float = 10.0, | |
| progress=gr.Progress() | |
| ): | |
| """Main processing function for Gradio interface. | |
| Args: | |
| speaker_mode: "diarization", "separation", "mossformer2", or "none" | |
| """ | |
| if audio_file is None: | |
| return None, "Please upload an audio file.", 0 | |
| total_start = time.time() | |
| mode_str = speaker_mode if speaker_mode != "none" else "no speaker ID" | |
| print(f"\n{'='*50}") | |
| print(f"[Pipeline] Starting - ASR: faster-whisper, Model: {model_size}, Mode: {mode_str}") | |
| print(f"{'='*50}") | |
| if not FASTER_WHISPER_AVAILABLE: | |
| return None, "faster-whisper not available. Install: pip install faster-whisper", 0 | |
| device, compute_type = get_device_info() | |
| progress(0.1, desc="Loading Whisper model...") | |
| load_start = time.time() | |
| model = load_whisper_model(model_size, device, compute_type) | |
| load_time = time.time() - load_start | |
| print(f"[Model Load {model_size}]: {load_time:.1f}s") | |
| temp_dir = tempfile.mkdtemp() | |
| output_dir = os.path.join(temp_dir, "output") | |
| os.makedirs(output_dir, exist_ok=True) | |
| # Speech Separation mode (2 speakers only) - DPRNN | |
| if speaker_mode == "separation": | |
| separator = load_separator() | |
| if separator is None: | |
| shutil.rmtree(temp_dir) | |
| return None, "DPRNN speech separation model not available.", 0 | |
| results, transcript_text = process_with_separation( | |
| audio_file, model, language, separator, output_dir, | |
| min_duration, max_duration, progress_fn=progress | |
| ) | |
| # MossFormer2 mode (2 speakers, SOTA quality) | |
| elif speaker_mode == "mossformer2": | |
| separator = load_mossformer2_separator() | |
| if separator is None: | |
| shutil.rmtree(temp_dir) | |
| return None, "MossFormer2 speech separation not available.", 0 | |
| results, transcript_text = process_with_mossformer2( | |
| audio_file, model, language, separator, output_dir, | |
| min_duration, max_duration, progress_fn=progress | |
| ) | |
| else: | |
| # Standard diarization or no speaker ID | |
| progress(0.3, desc="Transcribing audio...") | |
| asr_start = time.time() | |
| segments = transcribe_audio(audio_file, model, language) | |
| asr_time = time.time() - asr_start | |
| print(f"[ASR faster-whisper {model_size}]: {asr_time:.1f}s ({len(segments)} segments)") | |
| if not segments: | |
| shutil.rmtree(temp_dir) | |
| return None, "No speech detected in the audio.", 0 | |
| diarizer = None | |
| if speaker_mode == "diarization": | |
| progress(0.5, desc="Loading diarization model...") | |
| diarize_load_start = time.time() | |
| diarizer = load_diarizer() | |
| diarize_load_time = time.time() - diarize_load_start | |
| if diarizer: | |
| print(f"[Diarization Model Load]: {diarize_load_time:.1f}s") | |
| progress(0.6, desc="Splitting audio by speaker...") | |
| split_start = time.time() | |
| results, transcript_text = split_audio_with_diarization( | |
| audio_file, | |
| segments, | |
| output_dir, | |
| diarizer, | |
| min_duration=min_duration, | |
| max_duration=max_duration | |
| ) | |
| split_time = time.time() - split_start | |
| print(f"[Diarization + Split]: {split_time:.1f}s ({len(results)} output files)") | |
| if not results: | |
| shutil.rmtree(temp_dir) | |
| return None, "No segments within duration range found.", 0 | |
| progress(0.9, desc="Creating zip file...") | |
| transcript_path = os.path.join(output_dir, "transcript.txt") | |
| with open(transcript_path, "w", encoding="utf-8") as f: | |
| f.write(transcript_text) | |
| zip_path = os.path.join(temp_dir, "audio_segments.zip") | |
| with zipfile.ZipFile(zip_path, "w", zipfile.ZIP_DEFLATED) as zf: | |
| for root, dirs, files in os.walk(output_dir): | |
| for file in files: | |
| file_path = os.path.join(root, file) | |
| arcname = os.path.relpath(file_path, output_dir) | |
| zf.write(file_path, arcname) | |
| progress(1.0, desc="Done!") | |
| total_time = time.time() - total_start | |
| speakers = set(r["speaker"] for r in results) | |
| print(f"{'='*50}") | |
| print(f"[TOTAL]: {total_time:.1f}s | {len(results)} segments | {len(speakers)} speakers") | |
| print(f"{'='*50}\n") | |
| summary = f"Created {len(results)} segments from {len(speakers)} speaker(s) in {total_time:.1f}s" | |
| return zip_path, f"{summary}\n\n{transcript_text}", len(results) | |
| # Language options | |
| LANGUAGES = { | |
| "auto": "Auto-detect", | |
| "en": "English", | |
| "zh": "Chinese", | |
| "ja": "Japanese", | |
| "ko": "Korean", | |
| "es": "Spanish", | |
| "fr": "French", | |
| "de": "German", | |
| "it": "Italian", | |
| "pt": "Portuguese", | |
| "ru": "Russian", | |
| "ar": "Arabic", | |
| "hi": "Hindi", | |
| "th": "Thai", | |
| "vi": "Vietnamese", | |
| } | |
| MODELS = { | |
| "large-v3": "Large-v3 (best multilingual, recommended)", | |
| "medium": "Medium (faster, good quality)", | |
| "small": "Small (fast)", | |
| "distil-large-v3": "Distil-Large-v3 (fast, English only)", | |
| "large-v3-turbo": "Turbo (fast, manual lang only)", | |
| } | |
| def create_interface(): | |
| """Create compact Gradio interface with two tabs: Audio Splitter + Voice Extractor.""" | |
| from voice_extractor import run_voice_extractor_pipeline | |
| device, compute_type = get_device_info() | |
| base_dir = os.path.dirname(os.path.abspath(__file__)) | |
| # Speaker mode choices - build dynamically based on available models | |
| speaker_modes = ["Diarization (any speakers)"] | |
| if MOSSFORMER2_AVAILABLE: | |
| speaker_modes.append("MossFormer2 (2 speakers, SOTA)") | |
| if DPRNN_AVAILABLE: | |
| speaker_modes.append("DPRNN (2 speakers)") | |
| speaker_modes.append("None") | |
| with gr.Blocks(title="Audio Splitter") as demo: | |
| with gr.Tabs(): | |
| # ============================================================== | |
| # Tab 1: Audio Splitter (default) | |
| # ============================================================== | |
| with gr.TabItem("Audio Splitter"): | |
| # Compact header with processing time | |
| gr.Markdown(f"## [Audio Splitter](https://github.com/JarodMica/audiosplitter_whisper) with Speaker Diarization | {device.upper()} ({compute_type}) | ~1.5s per 1s audio") | |
| with gr.Row(): | |
| # Left column - inputs | |
| with gr.Column(scale=1): | |
| audio_input = gr.Audio(label="Audio", type="filepath", sources=["upload", "microphone"]) | |
| with gr.Row(): | |
| model_dropdown = gr.Dropdown( | |
| choices=list(MODELS.keys()), value="large-v3", label="Model", scale=2 | |
| ) | |
| language_dropdown = gr.Dropdown( | |
| choices=list(LANGUAGES.keys()), value="auto", label="Lang", scale=1 | |
| ) | |
| with gr.Row(): | |
| min_duration = gr.Slider(0.1, 60.0, 0.5, step=0.1, label="Min segment (s)") | |
| max_duration = gr.Slider(1.0, 600.0, 30.0, step=1.0, label="Max segment (s)") | |
| speaker_mode = gr.Radio( | |
| choices=speaker_modes, | |
| value="Diarization (any speakers)", | |
| label="Speaker ID Mode" | |
| ) | |
| process_btn = gr.Button("Process", variant="primary", size="lg") | |
| # Right column - outputs | |
| with gr.Column(scale=1): | |
| output_file = gr.File(label="Download ZIP") | |
| segment_count = gr.Number(label="Segments", interactive=False) | |
| transcript_output = gr.Textbox(label="Transcript", lines=10, interactive=False) | |
| # Wrapper to convert radio choice to mode string | |
| def process_wrapper(audio, model, lang, mode_choice, min_dur, max_dur, progress=gr.Progress()): | |
| # Convert UI choice to internal mode | |
| if "Diarization" in mode_choice: | |
| mode = "diarization" | |
| elif "MossFormer2" in mode_choice: | |
| mode = "mossformer2" | |
| elif "DPRNN" in mode_choice or "Separation" in mode_choice: | |
| mode = "separation" | |
| else: | |
| mode = "none" | |
| return process_audio(audio, model, lang, mode, min_dur, max_dur, progress) | |
| # Examples | |
| gr.Examples( | |
| examples=[ | |
| ["assets/57_years_apart.mp3", "large-v3", "auto", "Diarization (any speakers)", 0.5, 30.0], | |
| ], | |
| inputs=[audio_input, model_dropdown, language_dropdown, speaker_mode, min_duration, max_duration], | |
| outputs=[output_file, transcript_output, segment_count], | |
| fn=process_wrapper, | |
| cache_examples=True, | |
| label="Examples" | |
| ) | |
| process_btn.click( | |
| fn=process_wrapper, | |
| inputs=[audio_input, model_dropdown, language_dropdown, speaker_mode, min_duration, max_duration], | |
| outputs=[output_file, transcript_output, segment_count] | |
| ) | |
| # ============================================================== | |
| # Tab 2: Voice Extractor | |
| # ============================================================== | |
| with gr.TabItem("Voice Extractor"): | |
| gr.Markdown( | |
| "## Voice Extractor (CPU)\n" | |
| "Extract a target speaker from multi-speaker audio using a reference clip. " | |
| "Based on [Voice_Extractor](https://github.com/ReisCook/Voice_Extractor). " | |
| "PyAnnote + WeSpeaker + SpeechBrain + Whisper + Bandit-v2 (ONNX)." | |
| ) | |
| with gr.Row(): | |
| # --- Left: Inputs --- | |
| with gr.Column(scale=1): | |
| ve_input_audio = gr.Audio(label="Input Audio (multi-speaker)", type="filepath") | |
| ve_ref_audio = gr.Audio(label="Reference (target speaker clip)", type="filepath") | |
| with gr.Accordion("Settings", open=False): | |
| ve_whisper_model = gr.Dropdown( | |
| choices=["tiny", "base", "small"], value="base", | |
| label="Whisper Model (CPU)", | |
| info="Larger = more accurate but slower on CPU.", | |
| ) | |
| ve_language = gr.Dropdown( | |
| choices=["en", "es", "fr", "de", "ja", "zh", "ko", "auto"], value="en", | |
| label="Language", | |
| ) | |
| ve_threshold = gr.Slider( | |
| 0.0, 1.0, value=0.7, step=0.05, | |
| label="Verification Threshold", | |
| info="Higher = stricter speaker matching.", | |
| ) | |
| ve_min_dur = gr.Slider( | |
| 0.5, 5.0, value=1.0, step=0.25, | |
| label="Min Segment Duration (s)", | |
| ) | |
| ve_merge_gap = gr.Slider( | |
| 0.0, 2.0, value=0.25, step=0.05, | |
| label="Merge Gap (s)", | |
| info="Merge segments closer than this gap.", | |
| ) | |
| ve_use_sb = gr.Checkbox( | |
| value=True, | |
| label="SpeechBrain Verification", | |
| info="Adds ECAPA-TDNN speaker verification (slower but more accurate).", | |
| ) | |
| ve_dry_run = gr.Checkbox( | |
| value=False, | |
| label="Dry Run (first 60s only)", | |
| info="Quick test mode -- processes only the first minute.", | |
| ) | |
| ve_use_bandit = gr.Checkbox( | |
| value=False, | |
| label="Bandit-v2 Vocal Separation", | |
| info="Separate vocals from music/SFX before diarization. Very slow on CPU (~25x realtime).", | |
| ) | |
| ve_run_btn = gr.Button("Extract Voice", variant="primary", size="lg") | |
| # --- Right: Results --- | |
| with gr.Column(scale=1): | |
| ve_status = gr.Textbox(label="Status", interactive=False, max_lines=2) | |
| ve_audio_preview = gr.Audio(label="Extracted Audio Preview", type="filepath") | |
| ve_output_files = gr.File(label="Output Files", file_count="multiple") | |
| ve_transcript = gr.Textbox(label="Transcripts", lines=6, interactive=False) | |
| def ve_wrapper(inp_audio, ref_audio, wmodel, lang, thresh, | |
| min_d, mgap, usb, dry, bandit, progress=gr.Progress(track_tqdm=True)): | |
| return run_voice_extractor_pipeline( | |
| input_audio_path=inp_audio, | |
| reference_audio_path=ref_audio, | |
| target_name="TargetSpeaker", | |
| whisper_model=wmodel, | |
| language=lang, | |
| verification_threshold=thresh, | |
| min_duration=min_d, | |
| merge_gap=mgap, | |
| output_sr=44100, | |
| use_speechbrain=usb, | |
| dry_run=dry, | |
| use_bandit=bandit, | |
| progress=progress, | |
| ) | |
| ve_run_btn.click( | |
| fn=ve_wrapper, | |
| inputs=[ | |
| ve_input_audio, ve_ref_audio, | |
| ve_whisper_model, ve_language, ve_threshold, | |
| ve_min_dur, ve_merge_gap, | |
| ve_use_sb, ve_dry_run, ve_use_bandit, | |
| ], | |
| outputs=[ve_status, ve_audio_preview, ve_output_files, ve_transcript], | |
| ) | |
| # --- Examples --- | |
| ve_sample_audio = os.path.join(base_dir, "assets", "57_years_apart.mp3") | |
| ve_sample_ref = os.path.join(base_dir, "assets", "ref_speaker.wav") | |
| if os.path.exists(ve_sample_audio) and os.path.exists(ve_sample_ref): | |
| gr.Examples( | |
| examples=[ | |
| [ve_sample_audio, ve_sample_ref, "tiny", "en", 0.5, 0.5, 0.5, False, True, False], | |
| ], | |
| inputs=[ | |
| ve_input_audio, ve_ref_audio, | |
| ve_whisper_model, ve_language, ve_threshold, | |
| ve_min_dur, ve_merge_gap, | |
| ve_use_sb, ve_dry_run, ve_use_bandit, | |
| ], | |
| cache_examples=False, | |
| ) | |
| return demo | |
| def cli_mode(audio_file: str, model_size: str = "large-v3", language: str = "auto", | |
| speaker_mode: str = "diarization", min_dur: float = 0.5, max_dur: float = 10.0, | |
| output_dir: str = None): | |
| """CLI mode for testing.""" | |
| import sys | |
| if not os.path.exists(audio_file): | |
| print(f"Error: File not found: {audio_file}") | |
| sys.exit(1) | |
| print(f"Processing: {audio_file}") | |
| print(f"Model: {model_size}, Language: {language}, Speaker mode: {speaker_mode}") | |
| print(f"Duration: {min_dur}-{max_dur}s") | |
| class DummyProgress: | |
| def __call__(self, val, desc=""): | |
| print(f"[{int(val*100):3d}%] {desc}") | |
| zip_path, transcript, count = process_audio( | |
| audio_file, model_size, language, speaker_mode, min_dur, max_dur, DummyProgress() | |
| ) | |
| if zip_path: | |
| final_path = output_dir or os.path.dirname(audio_file) or "." | |
| final_zip = os.path.join(final_path, "audio_segments.zip") | |
| shutil.move(zip_path, final_zip) | |
| print(f"\n[OK] Created {count} segments") | |
| print(f"[OK] Saved to: {final_zip}") | |
| print(f"\nTranscript:\n{transcript}") | |
| else: | |
| print(f"\n[FAIL] {transcript}") | |
| sys.exit(1) | |
| if __name__ == "__main__": | |
| import sys | |
| if len(sys.argv) > 1 and sys.argv[1] == "cli": | |
| if len(sys.argv) < 3: | |
| print("Usage: python app.py cli <audio_file> [model] [language] [mode] [min_dur] [max_dur]") | |
| print(" model: large-v3 (default), medium, small, distil-large-v3, large-v3-turbo") | |
| print(" mode: diarization (default), mossformer2, separation (DPRNN), none") | |
| sys.exit(1) | |
| audio_file = sys.argv[2] | |
| model = sys.argv[3] if len(sys.argv) > 3 else "large-v3" | |
| lang = sys.argv[4] if len(sys.argv) > 4 else "auto" | |
| mode = sys.argv[5] if len(sys.argv) > 5 else "diarization" | |
| min_d = float(sys.argv[6]) if len(sys.argv) > 6 else 0.5 | |
| max_d = float(sys.argv[7]) if len(sys.argv) > 7 else 10.0 | |
| cli_mode(audio_file, model, lang, mode, min_d, max_d) | |
| else: | |
| demo = create_interface() | |
| demo.launch(mcp_server=True, show_error=True, ssr_mode=False, theme=gr.themes.Soft()) | |