File size: 4,903 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
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
"""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)