File size: 8,138 Bytes
23b3c64
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
"""
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)