File size: 9,738 Bytes
7c89374
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
# feature_extraction_mini_whisper.py
import numpy as np
import torch


class MiniWhisperFeatureExtractor:
    """Feature extractor matching Whisper exactly."""
    
    model_input_names = ["input_features"]
    
    def __init__(
        self,
        feature_size=80,
        sampling_rate=16000,
        hop_length=160,
        chunk_length=30,
        n_fft=400,
        padding_value=0.0,
        dither=0.0,
        return_attention_mask=False,
        **kwargs,
    ):
        self.feature_size = feature_size
        self.sampling_rate = sampling_rate
        self.hop_length = hop_length
        self.chunk_length = chunk_length
        self.n_fft = n_fft
        self.n_samples = chunk_length * sampling_rate
        self.nb_max_frames = self.n_samples // hop_length
        self.padding_value = padding_value
        self.dither = dither
        self.return_attention_mask = return_attention_mask
        
        self.mel_filters = self._mel_filter_bank()
    
    def _mel_filter_bank(self):
        """Create mel filter bank matching Whisper's parameters."""
        try:
            from torchaudio.functional import melscale_fbanks
            n_freq_bins = 1 + self.n_fft // 2
            mel_filters = melscale_fbanks(
                n_freqs=n_freq_bins,
                f_min=0.0,
                f_max=8000.0,
                n_mels=self.feature_size,
                sample_rate=self.sampling_rate,
                norm="slaney",
                mel_scale="slaney",
            )
            return mel_filters.numpy()
        except ImportError:
            return self._numpy_mel_filter_bank()
    
    def _numpy_mel_filter_bank(self):
        """NumPy fallback for mel filter bank."""
        def hz_to_mel(hz):
            return 2595 * np.log10(1 + hz / 700)
        
        def mel_to_hz(mel):
            return 700 * (10 ** (mel / 2595) - 1)
        
        n_freq_bins = 1 + self.n_fft // 2
        fmin, fmax = 0.0, 8000.0
        
        mel_min = hz_to_mel(fmin)
        mel_max = hz_to_mel(fmax)
        mel_points = np.linspace(mel_min, mel_max, self.feature_size + 2)
        hz_points = mel_to_hz(mel_points)
        
        bin_points = np.floor((self.n_fft + 1) * hz_points / self.sampling_rate).astype(int)
        
        filters = np.zeros((self.feature_size, n_freq_bins))
        for i in range(self.feature_size):
            left, center, right = bin_points[i], bin_points[i + 1], bin_points[i + 2]
            for j in range(left, center):
                filters[i, j] = (j - left) / (center - left)
            for j in range(center, right):
                filters[i, j] = (right - j) / (right - center)
        
        filters = filters / (filters.sum(axis=1, keepdims=True) + 1e-10)
        return filters
    
    def _np_extract_fbank_features(self, waveform_batch, device="cpu"):
        """NumPy implementation matching Whisper exactly."""
        log_spec_batch = []
        for waveform in waveform_batch:
            # Apply dither
            if self.dither != 0.0:
                waveform = waveform + self.dither * np.random.randn(*waveform.shape).astype(np.float32)
            
            # Compute STFT
            n_frames = 1 + (len(waveform) - self.n_fft) // self.hop_length
            spec = np.zeros((self.n_fft // 2 + 1, n_frames), dtype=np.float32)
            
            window = np.hanning(self.n_fft).astype(np.float32)
            for i in range(n_frames):
                start = i * self.hop_length
                frame = waveform[start:start + self.n_fft]
                if len(frame) < self.n_fft:
                    frame = np.pad(frame, (0, self.n_fft - len(frame)))
                spec[:, i] = np.abs(np.fft.rfft(frame * window)) ** 2
            
            # Apply mel filters: mel_filters is [n_mels, n_freqs]
            mel_spec = self.mel_filters @ spec
            
            # Log10 with floor
            log_spec = np.log10(np.maximum(mel_spec, 1e-10))
            
            # Whisper uses: np.maximum(log_spec, log_spec.max() - 8.0)
            log_spec = np.maximum(log_spec, log_spec.max() - 8.0)
            
            # Normalize: (log_spec + 4.0) / 4.0
            log_spec = (log_spec + 4.0) / 4.0
            
            # Remove last frame to match Whisper
            log_spec = log_spec[:, :-1]
            
            log_spec_batch.append(log_spec)
        
        return np.array(log_spec_batch)
    
    def _torch_extract_fbank_features(self, waveform, device="cpu"):
        """PyTorch implementation matching Whisper exactly."""
        waveform = torch.from_numpy(waveform).to(device, torch.float32)
        window = torch.hann_window(self.n_fft, device=device)
        
        if self.dither != 0.0:
            waveform += self.dither * torch.randn(waveform.shape, dtype=waveform.dtype, device=waveform.device)
        
        stft = torch.stft(waveform, self.n_fft, self.hop_length, window=window, return_complex=True)
        magnitudes = stft[..., :-1].abs() ** 2  # Remove last frame
        
        # mel_filters shape: [n_mels, n_freqs] from mel_filter_bank
        mel_filters = torch.from_numpy(self.mel_filters).to(device, torch.float32)
        # Transpose and matmul: mel_filters.T @ magnitudes
        mel_spec = mel_filters.T @ magnitudes
        
        log_spec = torch.clamp(mel_spec, min=1e-10).log10()
        
        if waveform.dim() == 2:
            max_val = log_spec.max(dim=2, keepdim=True)[0].max(dim=1, keepdim=True)[0]
            log_spec = torch.maximum(log_spec, max_val - 8.0)
        else:
            log_spec = torch.maximum(log_spec, log_spec.max() - 8.0)
        
        log_spec = (log_spec + 4.0) / 4.0
        
        if device != "cpu":
            log_spec = log_spec.detach().cpu()
        
        return log_spec.numpy()
    
    @staticmethod
    def zero_mean_unit_var_norm(input_values, attention_mask, padding_value=0.0):
        if attention_mask is not None:
            attention_mask = np.array(attention_mask, np.int32)
            normed = []
            for vector, length in zip(input_values, attention_mask.sum(-1)):
                normed_slice = (vector - vector[:length].mean()) / (np.sqrt(vector[:length].var() + 1e-7))
                if length < normed_slice.shape[0]:
                    normed_slice[length:] = padding_value
                normed.append(normed_slice)
        else:
            normed = [(x - x.mean()) / (np.sqrt(x.var() + 1e-7)) for x in input_values]
        return normed
    
    def __call__(
        self,
        raw_speech,
        sampling_rate=None,
        truncation=True,
        return_tensors="pt",
        return_attention_mask=None,
        do_normalize=False,
        **kwargs,
    ):
        if sampling_rate is not None and sampling_rate != self.sampling_rate:
            raise ValueError(f"Expected {self.sampling_rate}, got {sampling_rate}")
        
        is_batched = isinstance(raw_speech, list) or (isinstance(raw_speech, np.ndarray) and raw_speech.ndim > 1)
        
        if not is_batched:
            raw_speech = [raw_speech]
        
        # Convert to numpy float32
        processed = []
        for audio in raw_speech:
            if isinstance(audio, np.ndarray) and audio.dtype == np.float64:
                audio = audio.astype(np.float32)
            elif not isinstance(audio, np.ndarray):
                audio = np.array(audio, dtype=np.float32)
            
            # Pad or truncate
            if len(audio) > self.n_samples:
                if truncation:
                    audio = audio[:self.n_samples]
            elif len(audio) < self.n_samples:
                audio = np.pad(audio, (0, self.n_samples - len(audio)), constant_values=self.padding_value)
            
            processed.append(audio)
        
        waveform_batch = np.stack(processed, axis=0)
        
        # Use torch if available
        if torch.cuda.is_available():
            input_features = self._torch_extract_fbank_features(waveform_batch, device="cuda")
            # Result is [batch, n_mels, n_frames]
        else:
            input_features = self._np_extract_fbank_features(waveform_batch)
        
        result = {"input_features": input_features}
        
        if return_attention_mask is None:
            return_attention_mask = self.return_attention_mask
        
        if return_attention_mask:
            result["attention_mask"] = np.ones(input_features.shape[0], dtype=np.int64)
        
        if return_tensors == "pt":
            result["input_features"] = torch.from_numpy(input_features).float()
            if "attention_mask" in result:
                result["attention_mask"] = torch.from_numpy(result["attention_mask"]).long()
        
        return result
    
    def to_dict(self):
        return {
            "feature_size": self.feature_size,
            "sampling_rate": self.sampling_rate,
            "hop_length": self.hop_length,
            "n_fft": self.n_fft,
            "chunk_length": self.chunk_length,
            "padding_value": self.padding_value,
            "dither": self.dither,
        }
    
    def save_pretrained(self, save_directory):
        import json, os
        os.makedirs(save_directory, exist_ok=True)
        with open(os.path.join(save_directory, "preprocessor_config.json"), "w") as f:
            json.dump(self.to_dict(), f, indent=2)
    
    @classmethod
    def from_pretrained(cls, save_directory):
        import json, os
        path = os.path.join(save_directory, "preprocessor_config.json")
        if os.path.exists(path):
            with open(path) as f:
                return cls(**json.load(f))
        return cls()


__all__ = ["MiniWhisperFeatureExtractor"]