Reza2kn's picture
Publish complete experimental Audio8 Q4 model and runtimes
21fd722 verified
Raw History Blame Contribute Delete
2.73 kB
"""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:]