Reza2kn's picture
Publish complete experimental Audio8 Q4 model and runtimes
21fd722 verified
Raw
History Blame Contribute Delete
4.9 kB
"""Causal attention and bounded KV state. No framework/model import beyond MLX."""
import math
import mlx.core as mx
from .weights import require
def gelu(x):
return (x * 0.5) * (1 + mx.erf(x * math.sqrt(0.5)))
def rms_norm(x, weight, eps):
f = x.astype(mx.float32)
normalized = f * mx.rsqrt(mx.mean(f * f, axis=-1, keepdims=True) + eps)
return normalized.astype(x.dtype) * weight.astype(x.dtype)
def fast_rms_norm(x, weight, eps):
# Keep upstream's low-dtype rounding BEFORE the weight multiplication.
# Passing weight into the fused norm can instead multiply in FP32.
return mx.fast.rms_norm(x, None, eps) * weight.astype(x.dtype)
def rotary_factors(start, count, dim, dtype, theta=1_000_000.0):
require(dim % 2 == 0 and count > 0, 'rotary dimensions')
inv = mx.power(mx.array(theta, mx.float32), -mx.arange(0, dim, 2, dtype=mx.float32) / dim)
pos = mx.arange(start, start + count, dtype=mx.float32)
angle = pos[:, None] * inv[None, :]
return (mx.concatenate([mx.cos(angle)] * 2, axis=-1).astype(dtype),
mx.concatenate([mx.sin(angle)] * 2, axis=-1).astype(dtype))
def apply_rotary(x, factors):
cos, sin = factors
require(cos.shape == sin.shape == (x.shape[-2], x.shape[-1])
and cos.dtype == sin.dtype == x.dtype, 'rotary factor shape/dtype')
half = x.shape[-1] // 2
rotate = mx.concatenate([-x[..., half:], x[..., :half]], axis=-1)
return x * cos + rotate * sin
def rope(x, start, theta=1_000_000.0):
"""Noninterleaved half-rotation, FP32 angles, positions including negative rebases."""
dim = x.shape[-1]
require(dim % 2 == 0, 'odd rotary dimension')
return apply_rotary(x, rotary_factors(start, x.shape[-2], dim, x.dtype, theta))
def rotate_delta(x, delta, theta=1_000_000.0):
"""Uniform RoPE rebase, unlike rope() no increasing position ramp."""
dim = x.shape[-1]
inv = mx.power(mx.array(theta, mx.float32), -mx.arange(0, dim, 2, dtype=mx.float32) / dim)
angle = mx.array(delta, mx.float32) * inv
cos = mx.concatenate([mx.cos(angle)] * 2).astype(x.dtype)
sin = mx.concatenate([mx.sin(angle)] * 2).astype(x.dtype)
half = dim // 2
return x * cos + mx.concatenate([-x[..., half:], x[..., :half]], axis=-1) * sin
class KVCache:
def __init__(self, capacity, dtype, *, sliding=False):
require(type(capacity) is int and capacity > 0, 'cache capacity')
self.capacity, self.dtype, self.sliding = capacity, dtype, sliding
self.keys = self.values = None
@property
def length(self):
return 0 if self.keys is None else self.keys.shape[2]
def append(self, keys, values):
require(keys.shape == values.shape and keys.ndim == 4 and keys.shape[0] == 1, 'KV shape')
require(self.sliding or self.length + keys.shape[2] <= self.capacity, 'decoder cache limit; trim before query')
if self.keys is None:
all_k, all_v = keys, values
else:
require(self.keys.shape[1::2] == keys.shape[1::2], 'changed KV heads/dim')
all_k = mx.concatenate([self.keys.astype(keys.dtype), keys], axis=2)
all_v = mx.concatenate([self.values.astype(values.dtype), values], axis=2)
keep = min(all_k.shape[2], self.capacity)
# Force owned compact state: no persistent view retaining the full history.
self.keys = mx.contiguous(all_k[:, :, -keep:, :].astype(self.dtype))
self.values = mx.contiguous(all_v[:, :, -keep:, :].astype(self.dtype))
return all_k, all_v
def trim_decoder(self, drop, stable, theta):
require(0 <= stable < self.length and 0 < drop < self.length - stable, 'decoder trim dimensions')
prefix_k = self.keys[:, :, :stable]
suffix_k = rotate_delta(self.keys[:, :, stable + drop:].astype(mx.float32), -drop, theta).astype(self.dtype)
self.keys = mx.contiguous(mx.concatenate([prefix_k, suffix_k], axis=2))
self.values = mx.contiguous(mx.concatenate([self.values[:, :, :stable], self.values[:, :, stable + drop:]], axis=2))
def rebase_encoder(self, drop_frames, theta):
if self.keys is not None:
self.keys = mx.contiguous(rotate_delta(self.keys.astype(mx.float32), -drop_frames, theta).astype(self.dtype))
@property
def nbytes(self):
return 0 if self.keys is None else self.keys.nbytes + self.values.nbytes
def arrays(self):
return [] if self.keys is None else [self.keys, self.values]
def attention(q, k, v, cache, *, window=None):
prior = cache.length
keys, values = cache.append(k, v)
qpos = mx.arange(prior, prior + q.shape[2])[:, None]
kpos = mx.arange(keys.shape[2])[None, :]
mask = kpos <= qpos
if window is not None:
mask = mask & (kpos > qpos - window)
return mx.fast.scaled_dot_product_attention(q, keys, values,
scale=q.shape[-1] ** -0.5, mask=mask)