"""Pinned centered 400-point log-mel followed by stateless causal convolutions.""" import struct from pathlib import Path import numpy as np import mlx.core as mx from .math import gelu from .weights import require, sha256 class Frontend: def __init__(self, window, filters, provenance=None): require(window.shape == (400,) and filters.shape == (128, 201), 'frontend coefficient shapes') require(np.isfinite(window).all() and np.isfinite(filters).all(), 'nonfinite coefficients') self.window = mx.array(window.astype(np.float32)) self.filters = mx.array(filters.T.astype(np.float32)) self.provenance = provenance or {} mx.eval(self.window, self.filters) @classmethod def from_fixture(cls, path): path = Path(path) with path.open('rb') as stream: require(stream.read(8) == b'A8MEL001', 'frontend artifact format') require(1 <= struct.unpack('= len(waveform), 2 * len(waveform) - 2 - positions, positions) x = mx.array(waveform.astype(np.float32, copy=False))[mx.array(positions.astype(np.int32))] spectrum = mx.fft.rfft(x * self.window, n=400, axis=-1) power = mx.square(mx.abs(spectrum)) mel = power @ self.filters # Pinned global maximum 1.5. Never use a whole-utterance maximum. return (mx.maximum(mx.log10(mx.maximum(mel, 1e-10)), -6.5) + 4.0) / 4.0 def convolve(self, mel, weights, expected_frames, dtype): x = mel[None].astype(dtype) for number, stride, left in [(1, 1, 2), (2, 2, 1)]: prefix = f'audio_tower.embedder.conv{number}' w = weights.tensor(prefix + '.weight').transpose(0, 2, 1).astype(dtype) b = weights.tensor(prefix + '.bias').astype(dtype) x = gelu(mx.conv1d(mx.pad(x, [(0, 0), (left, 0), (0, 0)]), w, stride=stride) + b) require(0 < expected_frames <= x.shape[1] <= expected_frames + 2, 'convolution frame alignment') return x[:, -expected_frames:]