""" Speech Separation using Dual-Path-RNN Separates mixed audio into 2 speaker tracks using the Dual-Path-RNN model. Based on: https://github.com/JusperLee/Dual-Path-RNN-Pytorch Limitations: - Only supports exactly 2 speakers - Output requires 2x slowdown correction (known model issue) """ import os import sys from pathlib import Path from typing import Tuple, Optional import numpy as np # Add model directory to path MODEL_DIR = Path(__file__).parent / "dual_path_rnn" / "Dual-Path-RNN-portable" / "Dual-Path-RNN" if MODEL_DIR.exists(): sys.path.insert(0, str(MODEL_DIR)) # Try to import PyTorch (required for model loading) try: import torch TORCH_AVAILABLE = True except ImportError: TORCH_AVAILABLE = False print("[Warning] PyTorch not installed - speech separation unavailable") # Try ONNX runtime try: import onnxruntime as ort ONNX_AVAILABLE = True except ImportError: ONNX_AVAILABLE = False class DualPathRNNSeparator: """ Speech separator using Dual-Path-RNN. Usage: separator = DualPathRNNSeparator() spk1, spk2 = separator.separate(audio_np, sample_rate=16000) """ # Model config from train_rnn.yml MODEL_CONFIG = { "in_channels": 256, "out_channels": 64, "hidden_channels": 128, "kernel_size": 2, "rnn_type": "LSTM", "norm": "ln", "dropout": 0, "bidirectional": True, "num_layers": 6, "K": 250, "num_spks": 2, } INTERNAL_SAMPLE_RATE = 8000 # Model trained on 8kHz CHUNK_SECONDS = 30 # Process in 30-second chunks for ONNX OVERLAP_SECONDS = 2 # Overlap between chunks for smooth transitions def __init__(self, model_path: Optional[str] = None, use_onnx: bool = True): """ Initialize the separator. Args: model_path: Path to model file (.pt or .onnx) use_onnx: If True, use ONNX runtime; if False, use PyTorch """ self.use_onnx = use_onnx and ONNX_AVAILABLE self._model = None self._onnx_session = None # Find default model path if model_path is None: pt_path = MODEL_DIR / "Dual-Path-RNN-model-best.pt" onnx_path = Path(__file__).parent / "dprnn_separator.onnx" if self.use_onnx and onnx_path.exists(): model_path = str(onnx_path) elif pt_path.exists(): model_path = str(pt_path) else: raise FileNotFoundError(f"Model not found. Expected: {pt_path} or {onnx_path}") self.model_path = model_path def _ensure_loaded(self): """Load model on first use.""" if self._model is not None or self._onnx_session is not None: return if self.model_path.endswith(".onnx"): self._load_onnx() else: self._load_pytorch() def _load_pytorch(self): """Load PyTorch model.""" if not TORCH_AVAILABLE: raise RuntimeError("PyTorch not installed") print(f"[SpeechSep] Loading PyTorch model: {self.model_path}") # Import model class from model.model_rnn import Dual_RNN_model # Create model self._model = Dual_RNN_model(**self.MODEL_CONFIG) # Load weights checkpoint = torch.load(self.model_path, map_location='cpu') self._model.load_state_dict(checkpoint["model_state_dict"]) self._model.eval() print(f"[SpeechSep] Model loaded (epoch {checkpoint.get('epoch', 'unknown')})") def _load_onnx(self): """Load ONNX model.""" if not ONNX_AVAILABLE: raise RuntimeError("onnxruntime not installed") print(f"[SpeechSep] Loading ONNX model: {self.model_path}") # Create session with CPU provider sess_options = ort.SessionOptions() sess_options.graph_optimization_level = ort.GraphOptimizationLevel.ORT_ENABLE_ALL self._onnx_session = ort.InferenceSession( self.model_path, sess_options, providers=['CPUExecutionProvider'] ) print("[SpeechSep] ONNX model loaded") def separate( self, audio: np.ndarray, sample_rate: int = 16000 ) -> Tuple[np.ndarray, np.ndarray]: """ Separate mixed audio into 2 speaker tracks. Args: audio: Input audio as numpy array (mono, float32) sample_rate: Sample rate of input audio Returns: Tuple of (speaker1_audio, speaker2_audio) at original sample rate """ self._ensure_loaded() # Ensure float32 audio = audio.astype(np.float32) # Resample to 8kHz if needed if sample_rate != self.INTERNAL_SAMPLE_RATE: audio_8k = self._resample(audio, sample_rate, self.INTERNAL_SAMPLE_RATE) else: audio_8k = audio # Normalize norm_factor = np.max(np.abs(audio_8k)) + 1e-8 audio_8k = audio_8k / norm_factor # Run separation if self._onnx_session is not None: spk1, spk2 = self._separate_onnx(audio_8k) else: spk1, spk2 = self._separate_pytorch(audio_8k) # Fix 2x speed issue by resampling (model outputs at 2x speed) # Resample from 8kHz to 4kHz equivalent, then to target spk1 = self._fix_speed(spk1, self.INTERNAL_SAMPLE_RATE, sample_rate) spk2 = self._fix_speed(spk2, self.INTERNAL_SAMPLE_RATE, sample_rate) # Denormalize spk1 = spk1 * norm_factor spk2 = spk2 * norm_factor return spk1, spk2 def _separate_pytorch(self, audio: np.ndarray) -> Tuple[np.ndarray, np.ndarray]: """Run separation with PyTorch.""" with torch.no_grad(): # [L] -> [1, L] x = torch.from_numpy(audio).unsqueeze(0) # Forward pass returns list of 2 tensors outputs = self._model(x) # Extract and process spk1 = outputs[0].squeeze().numpy() spk2 = outputs[1].squeeze().numpy() # Normalize each output spk1 = self._normalize_output(spk1) spk2 = self._normalize_output(spk2) return spk1, spk2 def _separate_onnx(self, audio: np.ndarray) -> Tuple[np.ndarray, np.ndarray]: """Run separation with ONNX runtime, using chunked processing for long audio.""" chunk_samples = self.CHUNK_SECONDS * self.INTERNAL_SAMPLE_RATE overlap_samples = self.OVERLAP_SECONDS * self.INTERNAL_SAMPLE_RATE # If audio is short enough, process directly if len(audio) <= chunk_samples: return self._separate_onnx_single(audio) # Process in overlapping chunks print(f" [SpeechSep] Processing {len(audio)/self.INTERNAL_SAMPLE_RATE:.1f}s audio in {self.CHUNK_SECONDS}s chunks...") step = chunk_samples - overlap_samples num_chunks = int(np.ceil((len(audio) - overlap_samples) / step)) # Initialize output arrays spk1_out = np.zeros(len(audio), dtype=np.float32) spk2_out = np.zeros(len(audio), dtype=np.float32) weight_sum = np.zeros(len(audio), dtype=np.float32) # Create crossfade window for overlap regions fade_len = overlap_samples fade_in = np.linspace(0, 1, fade_len, dtype=np.float32) fade_out = np.linspace(1, 0, fade_len, dtype=np.float32) for i in range(num_chunks): start = i * step end = min(start + chunk_samples, len(audio)) chunk = audio[start:end] # Pad last chunk if needed if len(chunk) < chunk_samples // 2: # Too short, pad it chunk = np.pad(chunk, (0, chunk_samples // 2 - len(chunk)), mode='constant') # Process chunk spk1_chunk, spk2_chunk = self._separate_onnx_single(chunk) # Ensure output matches input length (model may have slight length variation) if len(spk1_chunk) != len(chunk): spk1_chunk = self._resample(spk1_chunk, len(spk1_chunk), len(chunk)) spk2_chunk = self._resample(spk2_chunk, len(spk2_chunk), len(chunk)) # Create weight window (1 in middle, fade at edges) weight = np.ones(len(chunk), dtype=np.float32) if i > 0: # Fade in at start (except first chunk) weight[:fade_len] = fade_in[:len(weight[:fade_len])] if i < num_chunks - 1: # Fade out at end (except last chunk) weight[-fade_len:] = fade_out[:len(weight[-fade_len:])] # Accumulate with weights actual_end = min(start + len(chunk), len(audio)) actual_len = actual_end - start spk1_out[start:actual_end] += spk1_chunk[:actual_len] * weight[:actual_len] spk2_out[start:actual_end] += spk2_chunk[:actual_len] * weight[:actual_len] weight_sum[start:actual_end] += weight[:actual_len] print(f" [SpeechSep] Chunk {i+1}/{num_chunks} done") # Normalize by weight sum weight_sum = np.maximum(weight_sum, 1e-8) spk1_out /= weight_sum spk2_out /= weight_sum spk1_out = self._normalize_output(spk1_out) spk2_out = self._normalize_output(spk2_out) return spk1_out, spk2_out def _separate_onnx_single(self, audio: np.ndarray) -> Tuple[np.ndarray, np.ndarray]: """Run separation with ONNX runtime on a single chunk.""" # [L] -> [1, L] x = audio.reshape(1, -1).astype(np.float32) # Get input name input_name = self._onnx_session.get_inputs()[0].name # Run inference outputs = self._onnx_session.run(None, {input_name: x}) # Process outputs spk1 = outputs[0].squeeze() spk2 = outputs[1].squeeze() return spk1, spk2 def _normalize_output(self, audio: np.ndarray) -> np.ndarray: """Normalize output audio.""" audio = audio - np.mean(audio) max_val = np.max(np.abs(audio)) + 1e-8 audio = audio / max_val return audio def _fix_speed( self, audio: np.ndarray, input_sr: int, target_sr: int ) -> np.ndarray: """ Fix 2x speed issue and resample to target rate. The model outputs audio at 2x speed, so we need to: 1. Stretch to 2x length (effectively halving the speed) 2. Resample to target sample rate """ # Stretch to 2x length (fix speed) stretched = self._resample(audio, input_sr, input_sr // 2) # Resample to target (from effective 4kHz to target) effective_sr = input_sr // 2 # 4kHz output = self._resample(stretched, effective_sr, target_sr) return output def _resample( self, audio: np.ndarray, orig_sr: int, target_sr: int ) -> np.ndarray: """Resample audio to target sample rate.""" if orig_sr == target_sr: return audio try: from scipy import signal num_samples = int(len(audio) * target_sr / orig_sr) resampled = signal.resample(audio, num_samples) return resampled.astype(np.float32) except ImportError: # Fallback: linear interpolation ratio = target_sr / orig_sr indices = np.arange(0, len(audio), 1/ratio).astype(int) indices = indices[indices < len(audio)] return audio[indices] def export_to_onnx(self, output_path: str = "dprnn_separator.onnx"): """ Export the PyTorch model to ONNX format. Args: output_path: Path to save ONNX model """ if not TORCH_AVAILABLE: raise RuntimeError("PyTorch required for ONNX export") # Ensure PyTorch model is loaded if self._model is None: # Force PyTorch loading old_path = self.model_path if self.model_path.endswith(".onnx"): pt_path = MODEL_DIR / "Dual-Path-RNN-model-best.pt" self.model_path = str(pt_path) self._load_pytorch() self.model_path = old_path print(f"[SpeechSep] Exporting to ONNX: {output_path}") # Create dummy input (1 second at 8kHz) dummy_input = torch.randn(1, 8000) # Export torch.onnx.export( self._model, dummy_input, output_path, input_names=['audio'], output_names=['speaker1', 'speaker2'], dynamic_axes={ 'audio': {1: 'length'}, 'speaker1': {0: 'length'}, 'speaker2': {0: 'length'}, }, opset_version=14, do_constant_folding=True, ) print(f"[SpeechSep] ONNX model saved: {output_path}") return output_path # Convenience function def separate_speakers( audio: np.ndarray, sample_rate: int = 16000, use_onnx: bool = True ) -> Tuple[np.ndarray, np.ndarray]: """ Separate mixed audio into 2 speaker tracks. Args: audio: Input audio as numpy array (mono, float32) sample_rate: Sample rate of input audio use_onnx: Use ONNX runtime if available Returns: Tuple of (speaker1_audio, speaker2_audio) """ separator = DualPathRNNSeparator(use_onnx=use_onnx) return separator.separate(audio, sample_rate) # Test/CLI if __name__ == "__main__": import sys if len(sys.argv) < 2: print("Usage: python speech_separation.py [output_dir]") print(" python speech_separation.py --export # Export to ONNX") sys.exit(1) if sys.argv[1] == "--export": # Export to ONNX separator = DualPathRNNSeparator(use_onnx=False) output_path = sys.argv[2] if len(sys.argv) > 2 else "dprnn_separator.onnx" separator.export_to_onnx(output_path) else: # Separate audio from pydub import AudioSegment audio_file = sys.argv[1] output_dir = sys.argv[2] if len(sys.argv) > 2 else "separated" print(f"Loading: {audio_file}") audio_seg = AudioSegment.from_file(audio_file) audio_seg = audio_seg.set_channels(1) # Mono sample_rate = audio_seg.frame_rate audio_np = np.array(audio_seg.get_array_of_samples(), dtype=np.float32) audio_np = audio_np / 32768.0 # Normalize int16 print(f"Input: {len(audio_np)/sample_rate:.1f}s @ {sample_rate}Hz") # Separate print("Separating speakers...") separator = DualPathRNNSeparator(use_onnx=False) # Use PyTorch for testing spk1, spk2 = separator.separate(audio_np, sample_rate) # Save os.makedirs(output_dir, exist_ok=True) for i, spk in enumerate([spk1, spk2], 1): # Convert to int16 spk_int16 = (spk * 32768).astype(np.int16) spk_seg = AudioSegment( spk_int16.tobytes(), frame_rate=sample_rate, sample_width=2, channels=1 ) out_path = os.path.join(output_dir, f"speaker{i}.wav") spk_seg.export(out_path, format="wav") print(f"Saved: {out_path} ({len(spk)/sample_rate:.1f}s)") print("Done!")