File size: 2,729 Bytes
21fd722
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""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('<I', stream.read(4))[0] <= 16, 'frontend fixture count')
            window = np.frombuffer(stream.read(400 * 4), '<f4').copy()
            filters = np.frombuffer(stream.read(128 * 201 * 4), '<f4').reshape(128, 201).copy()
        return cls(window, filters, {'artifact_sha256': sha256(path)})

    def mel(self, waveform):
        require(waveform.ndim == 1 and 200 < len(waveform) <= 65536
                and np.isfinite(waveform).all(), 'waveform must be bounded finite mono 16kHz FP32')
        frames = len(waveform) // 160
        positions = np.arange(frames)[:, None] * 160 + np.arange(400)[None, :] - 200
        positions = np.where(positions < 0, -positions, positions)
        positions = np.where(positions >= 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:]