""" Utility functions for Sori Speech model. """ import numpy as np import torch import torchaudio from typing import List, Dict, Any, Tuple, Optional from transformers import WhisperFeatureExtractor # Global WhisperFeatureExtractor instance (Qwen3-Omni compatible: 128 mels) _whisper_fe = None def _get_whisper_fe(): global _whisper_fe if _whisper_fe is None: _whisper_fe = WhisperFeatureExtractor( feature_size=128, sampling_rate=16000, hop_length=160, n_fft=400, chunk_length=300, # 300s max like Qwen3-Omni padding_value=0.0, ) return _whisper_fe def process_mm_info( conversation: List[Dict[str, Any]], use_audio_in_video: bool = False ) -> Tuple[Optional[List[str]], Optional[List], Optional[List]]: """ Extract multimodal information (audio, images, videos) from conversation. Args: conversation: List of message dicts with role and content use_audio_in_video: Whether to extract audio from video files Returns: Tuple of (audio_paths, image_paths, video_paths) """ audios = [] images = [] videos = [] for message in conversation: if isinstance(message, dict): content = message.get("content", []) # Handle string content if isinstance(content, str): continue # Handle list content (multimodal) if isinstance(content, list): for item in content: if isinstance(item, dict): item_type = item.get("type", "") if item_type == "audio": audio_path = item.get("audio") if audio_path: audios.append(audio_path) elif item_type == "image": image_path = item.get("image") if image_path: images.append(image_path) elif item_type == "video": video_path = item.get("video") if video_path: videos.append(video_path) return ( audios if audios else None, images if images else None, videos if videos else None, ) def load_audio(audio_path: str, target_sr: int = 16000) -> torch.Tensor: """ Load audio file and convert to target sample rate. Args: audio_path: Path to audio file target_sr: Target sample rate (default: 16000) Returns: Audio tensor of shape (1, num_samples) """ audio, sr = torchaudio.load(audio_path) # Convert to mono if stereo if audio.shape[0] > 1: audio = audio.mean(dim=0, keepdim=True) # Resample if needed if sr != target_sr: audio = torchaudio.transforms.Resample(sr, target_sr)(audio) return audio def audio_to_mel_spectrogram( audio: torch.Tensor, sample_rate: int = 16000, **kwargs, ) -> torch.Tensor: """ Convert audio waveform to log-mel spectrogram using WhisperFeatureExtractor. Matches Qwen3-Omni's audio preprocessing exactly. Args: audio: Audio tensor of shape (1, num_samples) or (num_samples,) sample_rate: Audio sample rate Returns: Log mel spectrogram of shape (128, time) - NOT padded """ if audio.dim() == 2: audio = audio.squeeze(0) # Convert to numpy float32 waveform = audio.numpy().astype(np.float32) fe = _get_whisper_fe() # Use WhisperFeatureExtractor's torch extraction (matches Qwen3-Omni) # This does: STFT → power spec → slaney mel filterbank → log10 → clamp → normalize # Returns shape (128, T) without padding window = torch.hann_window(fe.n_fft) waveform_t = torch.from_numpy(waveform).float() stft = torch.stft(waveform_t, fe.n_fft, fe.hop_length, window=window, return_complex=True) magnitudes = stft[..., :-1].abs() ** 2 mel_filters = torch.from_numpy(fe.mel_filters).float() mel_spec = mel_filters.T @ magnitudes log_spec = torch.clamp(mel_spec, min=1e-10).log10() log_spec = torch.maximum(log_spec, log_spec.max() - 8.0) log_spec = (log_spec + 4.0) / 4.0 return log_spec # (128, time)