Download processing_sori_speech.py from Seungyoun/Sori-4B-FC: direct link, hf CLI and curl.
- Browser
- Download file 8.14 kB
-
https://huggingface.co/Seungyoun/Sori-4B-FC/resolve/main/processing_sori_speech.py
- Command line
-
hf download hf://Seungyoun/Sori-4B-FC/processing_sori_speech.py
-
curl -L -o processing_sori_speech.py https://huggingface.co/Seungyoun/Sori-4B-FC/resolve/main/processing_sori_speech.py
8.14 kB
| """ | |
| 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 | |
| ) | |
| 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) | |