Nekochu's picture
fix VE tab: restore info descriptions, remove Target Name + Sample Rate, results on right side
86624de
Raw History Blame Contribute Delete
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())