""" Speech Separation using MossFormer2 Separates mixed audio into speaker tracks using the MossFormer2 model. Based on: https://github.com/alibabasglab/MossFormer2 MossFormer2 achieves 24.1 dB SI-SNRi on WSJ0-2mix benchmark. """ import os import sys from pathlib import Path from typing import Tuple, Optional, List import numpy as np # Add mossformer2 model directory to path MOSSFORMER2_DIR = Path(__file__).parent / "mossformer2" if MOSSFORMER2_DIR.exists(): sys.path.insert(0, str(MOSSFORMER2_DIR.parent)) # Try to import PyTorch try: import torch import torchaudio TORCH_AVAILABLE = True except ImportError: TORCH_AVAILABLE = False print("[Warning] PyTorch not installed - MossFormer2 separation unavailable") class MossFormer2Separator: """ Speech separator using MossFormer2. Usage: separator = MossFormer2Separator() speakers = separator.separate(audio_np, sample_rate=16000) """ # Available model configs AVAILABLE_MODELS = { "mossformer2-whamr-2spk": { "num_spks": 2, "sample_rate": 8000, "description": "2-speaker separation trained on WHAMR (noisy+reverb)", }, "mossformer2-librimix-2spk": { "num_spks": 2, "sample_rate": 8000, "description": "2-speaker separation trained on LibriMix", }, "mossformer2-wsj0mix-3spk": { "num_spks": 3, "sample_rate": 8000, "description": "3-speaker separation trained on WSJ0-3mix", }, } CHUNK_SECONDS = 30 # Process in 30-second chunks OVERLAP_SECONDS = 2 # Overlap between chunks def __init__(self, model_name: str = "mossformer2-whamr-2spk"): """ Initialize the separator. Args: model_name: Model to use (mossformer2-whamr-2spk, mossformer2-librimix-2spk, etc.) """ if not TORCH_AVAILABLE: raise RuntimeError("PyTorch not installed") if model_name not in self.AVAILABLE_MODELS: raise ValueError(f"Unknown model: {model_name}. Available: {list(self.AVAILABLE_MODELS.keys())}") self.model_name = model_name self.model_info = self.AVAILABLE_MODELS[model_name] self.internal_sample_rate = self.model_info["sample_rate"] self.num_spks = self.model_info["num_spks"] self._model = None def _ensure_loaded(self): """Load model on first use.""" if self._model is not None: return print(f"[MossFormer2] Loading {self.model_name}...") # Import and load model from mossformer2.mossformer2 import Mossformer2Wrapper self._model = Mossformer2Wrapper.from_pretrained(f"alibabasglab/{self.model_name}") self._model.eval() print(f"[MossFormer2] Model loaded on {self._model.device}") def separate( self, audio: np.ndarray, sample_rate: int = 16000 ) -> List[np.ndarray]: """ Separate mixed audio into speaker tracks. Args: audio: Input audio as numpy array (mono, float32) sample_rate: Sample rate of input audio Returns: List of separated speaker audio arrays at original sample rate """ self._ensure_loaded() # Ensure float32 audio = audio.astype(np.float32) # Resample to internal rate 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 # Process (with chunking for long audio) chunk_samples = self.CHUNK_SECONDS * self.internal_sample_rate if len(audio_8k) > chunk_samples: separated = self._separate_chunked(audio_8k) else: separated = self._separate_single(audio_8k) # Resample back to original rate result = [] for spk_audio in separated: if sample_rate != self.internal_sample_rate: spk_audio = self._resample(spk_audio, self.internal_sample_rate, sample_rate) # Denormalize spk_audio = spk_audio * norm_factor result.append(spk_audio) return result def _separate_single(self, audio: np.ndarray) -> List[np.ndarray]: """Separate a single chunk of audio.""" with torch.no_grad(): # [L] -> [1, L] x = torch.from_numpy(audio).unsqueeze(0).to(self._model.device) # Forward pass: returns [1, L, num_spks] est_source = self._model.forward(x) # Extract speakers result = [] for i in range(self.num_spks): spk = est_source[0, :, i].cpu().numpy() spk = self._normalize_output(spk) result.append(spk) return result def _separate_chunked(self, audio: np.ndarray) -> List[np.ndarray]: """Separate long audio in chunks with overlap-add.""" chunk_samples = self.CHUNK_SECONDS * self.internal_sample_rate overlap_samples = self.OVERLAP_SECONDS * self.internal_sample_rate print(f" [MossFormer2] 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 for each speaker outputs = [np.zeros(len(audio), dtype=np.float32) for _ in range(self.num_spks)] weight_sum = np.zeros(len(audio), dtype=np.float32) # Create crossfade window 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 if too short if len(chunk) < chunk_samples // 4: chunk = np.pad(chunk, (0, chunk_samples // 4 - len(chunk)), mode='constant') # Process chunk separated_chunk = self._separate_single(chunk) # Create weight window weight = np.ones(len(chunk), dtype=np.float32) if i > 0: weight[:min(fade_len, len(weight))] = fade_in[:min(fade_len, len(weight))] if i < num_chunks - 1: weight[-min(fade_len, len(weight)):] = fade_out[-min(fade_len, len(weight)):] # Accumulate actual_end = min(start + len(chunk), len(audio)) actual_len = actual_end - start for spk_idx in range(self.num_spks): spk_chunk = separated_chunk[spk_idx] if len(spk_chunk) != len(chunk): spk_chunk = self._resample(spk_chunk, len(spk_chunk), len(chunk)) outputs[spk_idx][start:actual_end] += spk_chunk[:actual_len] * weight[:actual_len] weight_sum[start:actual_end] += weight[:actual_len] print(f" [MossFormer2] Chunk {i+1}/{num_chunks} done") # Normalize by weights weight_sum = np.maximum(weight_sum, 1e-8) result = [] for spk_idx in range(self.num_spks): outputs[spk_idx] /= weight_sum outputs[spk_idx] = self._normalize_output(outputs[spk_idx]) result.append(outputs[spk_idx]) return result 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.astype(np.float32) 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] # Convenience function def separate_speakers_mossformer2( audio: np.ndarray, sample_rate: int = 16000, model_name: str = "mossformer2-whamr-2spk" ) -> List[np.ndarray]: """ Separate mixed audio into speaker tracks using MossFormer2. Args: audio: Input audio as numpy array (mono, float32) sample_rate: Sample rate of input audio model_name: Model to use Returns: List of separated speaker audio arrays """ separator = MossFormer2Separator(model_name=model_name) return separator.separate(audio, sample_rate) # Test/CLI if __name__ == "__main__": import sys if len(sys.argv) < 2: print("Usage: python mossformer2_separation.py [output_dir] [model_name]") print("Models: mossformer2-whamr-2spk (default), mossformer2-librimix-2spk, mossformer2-wsj0mix-3spk") sys.exit(1) from pydub import AudioSegment audio_file = sys.argv[1] output_dir = sys.argv[2] if len(sys.argv) > 2 else "separated_mossformer2" model_name = sys.argv[3] if len(sys.argv) > 3 else "mossformer2-whamr-2spk" 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") print(f"Model: {model_name}") # Separate print("Separating speakers...") separator = MossFormer2Separator(model_name=model_name) speakers = separator.separate(audio_np, sample_rate) # Save os.makedirs(output_dir, exist_ok=True) for i, spk in enumerate(speakers, 1): # Convert to int16 spk_int16 = (spk * 32768).clip(-32768, 32767).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!")