""" Processor for Sori Speech model. """ import torch from typing import List, Optional, Union, Dict, Any from transformers import AutoTokenizer, ProcessorMixin from transformers.processing_utils import ProcessorMixin as BaseProcessorMixin from sori_speech_utils import load_audio, audio_to_mel_spectrogram class SoriSpeechProcessor(BaseProcessorMixin): """ Processor for SoriSpeech model that handles both text and audio inputs. This processor: 1. Tokenizes text with special audio tokens 2. Converts audio files to mel spectrograms 3. Manages the integration of audio and text modalities """ attributes = ["tokenizer"] tokenizer_class = "AutoTokenizer" def __init__( self, tokenizer=None, audio_sample_rate: int = 16000, n_fft: int = 400, hop_length: int = 160, n_mels: int = 128, **kwargs ): """ Initialize the processor. Args: tokenizer: The tokenizer to use for text processing audio_sample_rate: Sample rate for audio processing n_fft: FFT size for mel spectrogram hop_length: Hop length for mel spectrogram n_mels: Number of mel bins """ self.tokenizer = tokenizer self.audio_sample_rate = audio_sample_rate self.n_fft = n_fft self.hop_length = hop_length self.n_mels = n_mels super().__init__(tokenizer) def __call__( self, text: Optional[Union[str, List[str]]] = None, audio: Optional[Union[str, List[str]]] = None, return_tensors: Optional[str] = None, padding: Union[bool, str] = False, **kwargs ) -> Dict[str, Any]: """ Process text and audio inputs. Args: text: Text string or list of strings (already formatted with chat template) audio: Audio file path(s) return_tensors: Type of tensors to return ('pt' for PyTorch) padding: Whether to pad sequences Returns: Dictionary with input_ids, attention_mask, input_features, feature_lens """ # Tokenize text if text is None: raise ValueError("text input is required") text_inputs = self.tokenizer( text, return_tensors=return_tensors, padding=padding, **kwargs ) # Process audio if provided if audio is not None: if isinstance(audio, str): audio = [audio] # Convert audio files to mel spectrograms mel_features_list = [] feature_lens_list = [] for audio_path in audio: # Load and convert audio audio_tensor = load_audio(audio_path, self.audio_sample_rate) mel_features = audio_to_mel_spectrogram( audio_tensor, sample_rate=self.audio_sample_rate, n_fft=self.n_fft, hop_length=self.hop_length, n_mels=self.n_mels, ) mel_features_list.append(mel_features) feature_lens_list.append(mel_features.shape[1]) # Stack mel features (for batch processing, we'll just use the first one for now) if return_tensors == "pt": # For simplicity, handle single audio for now # Note: dtype conversion will be handled when moving to device text_inputs["input_features"] = mel_features_list[0] text_inputs["feature_lens"] = torch.tensor(feature_lens_list) return text_inputs def batch_decode(self, *args, **kwargs): """Decode token ids to text.""" return self.tokenizer.batch_decode(*args, **kwargs) def decode(self, *args, **kwargs): """Decode token ids to text.""" return self.tokenizer.decode(*args, **kwargs) def apply_chat_template( self, conversation: List[Dict[str, Any]], add_generation_prompt: bool = False, tokenize: bool = True, **kwargs ) -> Union[str, List[int]]: """ Apply chat template to conversation. This method processes multimodal conversations and replaces audio placeholders with the appropriate number of <|audio|> tokens. Args: conversation: List of message dicts with role and content add_generation_prompt: Whether to add generation prompt tokenize: Whether to tokenize the output Returns: Formatted text string or token ids """ from sori_speech_utils import process_mm_info # Extract audio paths from conversation audios, _, _ = process_mm_info(conversation) # Calculate number of audio tokens needed audio_token_counts = [] if audios: for audio_path in audios: # Load audio and get mel spectrogram length audio_tensor = load_audio(audio_path, self.audio_sample_rate) mel_features = audio_to_mel_spectrogram( audio_tensor, sample_rate=self.audio_sample_rate, n_fft=self.n_fft, hop_length=self.hop_length, n_mels=self.n_mels, ) # Calculate output length from audio encoder # This is a simplified calculation - you may need to match the actual encoder logic feature_len = mel_features.shape[1] # Use the same logic as in _get_feat_extract_output_lengths input_lengths_leave = feature_len % 100 feat_lengths = (input_lengths_leave - 1) // 2 + 1 output_length = ((feat_lengths - 1) // 2 + 1 - 1) // 2 + 1 + (feature_len // 100) * 13 audio_token_counts.append(int(output_length)) # Process conversation to replace audio items with text placeholders processed_conversation = [] audio_idx = 0 for message in conversation: processed_message = {"role": message["role"]} content = message.get("content", "") if isinstance(content, str): processed_message["content"] = content elif isinstance(content, list): # Process multimodal content text_parts = [] for item in content: if isinstance(item, dict): if item.get("type") == "audio": # Replace audio with token placeholders if audio_idx < len(audio_token_counts): num_tokens = audio_token_counts[audio_idx] audio_placeholder = "<|audio|>" * num_tokens text_parts.append(f"<|audio_start|>{audio_placeholder}<|audio_end|>") audio_idx += 1 elif item.get("type") == "text": text_parts.append(item.get("text", "")) processed_message["content"] = "".join(text_parts) processed_conversation.append(processed_message) # Apply tokenizer's chat template return self.tokenizer.apply_chat_template( processed_conversation, add_generation_prompt=add_generation_prompt, tokenize=tokenize, **kwargs ) @classmethod def from_pretrained(cls, pretrained_model_name_or_path, **kwargs): """Load processor from pretrained model.""" tokenizer = AutoTokenizer.from_pretrained(pretrained_model_name_or_path, **kwargs) return cls(tokenizer=tokenizer) def save_pretrained(self, save_directory, **kwargs): """Save processor to directory.""" self.tokenizer.save_pretrained(save_directory, **kwargs) # Register for AutoProcessor from transformers import AutoProcessor AutoProcessor.register("SoriSpeechProcessor", SoriSpeechProcessor)