# 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"]