Route the gpt_oss to fused sdpa (#356)
This commit is contained in:
+108
-74
@@ -8,7 +8,7 @@ from typing import Any, Optional
|
||||
import mlx.core as mx
|
||||
import mlx.nn as nn
|
||||
|
||||
from .base import BaseModelArgs, create_attention_mask
|
||||
from .base import BaseModelArgs, create_causal_mask, scaled_dot_product_attention
|
||||
from .cache import KVCache, RotatingKVCache
|
||||
from .rope_utils import initialize_rope
|
||||
from .switch_layers import SwitchGLU
|
||||
@@ -64,41 +64,6 @@ class SwiGLU(nn.Module):
|
||||
return swiglu(x, gate)
|
||||
|
||||
|
||||
# ref. eager_attention_forward in tfm impl
|
||||
def sdpa(
|
||||
Q: mx.array,
|
||||
K: mx.array,
|
||||
V: mx.array,
|
||||
S: mx.array,
|
||||
sm_scale: float,
|
||||
mask: mx.array,
|
||||
):
|
||||
# Q, K, V shapes: (batch, num_heads, seqlen, head_dim)
|
||||
batch, num_kv_heads, seqlen, head_dim = K.shape
|
||||
_, num_q_heads, q_len, _ = Q.shape
|
||||
|
||||
n_rep = num_q_heads // num_kv_heads
|
||||
Q = Q.reshape(batch, num_kv_heads, n_rep, q_len, head_dim)
|
||||
attn_weights = sm_scale * mx.matmul(Q, mx.expand_dims(K, axis=2).swapaxes(-1, -2))
|
||||
attn_weights = attn_weights.reshape(batch, head_dim, q_len, seqlen)
|
||||
|
||||
if mask.shape[-1] != K.shape[-2]:
|
||||
mask = mask[..., -K.shape[-2] :]
|
||||
attn_weights = mx.where(mask, attn_weights, -mx.inf)
|
||||
|
||||
sinks = mx.tile(S.reshape(1, -1, 1, 1), [batch, 1, q_len, 1])
|
||||
|
||||
combined_logits = mx.concatenate([attn_weights, sinks], axis=-1)
|
||||
probs = mx.softmax(combined_logits, axis=-1, precise=True)
|
||||
scores = probs[..., :-1].reshape(batch, num_kv_heads, n_rep, q_len, seqlen)
|
||||
attn_output = mx.matmul(scores, mx.expand_dims(V, axis=2))
|
||||
attn_output = attn_output.reshape(batch, num_q_heads, q_len, head_dim).swapaxes(
|
||||
1, 2
|
||||
)
|
||||
|
||||
return attn_output
|
||||
|
||||
|
||||
class AttentionBlock(nn.Module):
|
||||
def __init__(self, config: ModelArgs):
|
||||
super().__init__()
|
||||
@@ -135,38 +100,111 @@ class AttentionBlock(nn.Module):
|
||||
scaling_config=config.rope_scaling,
|
||||
)
|
||||
|
||||
def __call__(self, x: mx.array, mask: mx.array, cache=None) -> mx.array:
|
||||
input_shape = x.shape[:-1] # (batch, seqlen)
|
||||
# Cache the mask so we don't have to create it every time
|
||||
self._previous_mask = None
|
||||
|
||||
q = self.q_proj(x)
|
||||
k = self.k_proj(x)
|
||||
v = self.v_proj(x)
|
||||
def get_causal_mask(self, x, cache):
|
||||
_, L, _ = x.shape
|
||||
offset = cache.offset if cache is not None else 0
|
||||
offset = max(1, offset)
|
||||
|
||||
# (batch, seqlen, num_heads * head_dim) -> (batch, num_heads, seqlen, head_dim)
|
||||
q = q.reshape(*input_shape, self.num_attention_heads, self.head_dim).swapaxes(
|
||||
1, 2
|
||||
)
|
||||
k = k.reshape(*input_shape, self.num_key_value_heads, self.head_dim).swapaxes(
|
||||
1, 2
|
||||
)
|
||||
v = v.reshape(*input_shape, self.num_key_value_heads, self.head_dim).swapaxes(
|
||||
1, 2
|
||||
)
|
||||
def _make_mask(L, offset):
|
||||
zero = mx.array(0, dtype=x.dtype)
|
||||
neginf = mx.array(-mx.inf, dtype=x.dtype)
|
||||
mask = mx.where(create_causal_mask(L, offset - 1), zero, neginf)
|
||||
mask = mask.reshape(1, 1, L, -1)
|
||||
mask = mx.tile(mask, (1, self.num_attention_heads, 1, 1))
|
||||
sinks = mx.tile(self.sinks.reshape(1, -1, 1, 1), (1, 1, L, 1))
|
||||
mask = mx.concatenate([sinks, mask], axis=-1)
|
||||
return mask
|
||||
|
||||
if cache is not None:
|
||||
q = self.rope(q, offset=cache.offset)
|
||||
k = self.rope(k, offset=cache.offset)
|
||||
k, v = cache.update_and_fetch(k, v)
|
||||
# When training re-create the mask so that gradients flow to the sinks.
|
||||
# When L is large then recreate the mask because otherwise it will take
|
||||
# a pretty significant chunk of memory.
|
||||
if self.training or L > 8:
|
||||
self._previous_mask = None
|
||||
return _make_mask(L, offset)
|
||||
|
||||
# Create the mask once and try to reuse it. For this reason we round up
|
||||
# to the closest multiple of 512 so we can reuse the mask several times.
|
||||
length = ((L + offset + 511) // 512) * 512
|
||||
if (
|
||||
self._previous_mask is None
|
||||
or self._previous_mask.shape[-1] < length
|
||||
or self._previous_mask.shape[-2] != L
|
||||
):
|
||||
self._previous_mask = _make_mask(L, length - L)
|
||||
|
||||
return self._previous_mask[..., : L + offset]
|
||||
|
||||
def get_sliding_window_mask(self, x, cache, window_size):
|
||||
_, L, _ = x.shape
|
||||
offset = cache.offset if cache is not None else 0
|
||||
offset = max(1, offset)
|
||||
|
||||
def _make_mask(L, offset):
|
||||
zero = mx.array(0, dtype=x.dtype)
|
||||
neginf = mx.array(-mx.inf, dtype=x.dtype)
|
||||
mask = create_causal_mask(L, offset - 1, window_size)
|
||||
mask = mx.where(mask, zero, neginf)
|
||||
mask = mask.reshape(1, 1, L, -1)
|
||||
mask = mx.tile(mask, (1, self.num_attention_heads, 1, 1))
|
||||
sinks = mx.tile(self.sinks.reshape(1, -1, 1, 1), (1, 1, L, 1))
|
||||
mask = mx.concatenate([sinks, mask], axis=-1)
|
||||
return mask
|
||||
|
||||
# If we are training then simply re-create the mask every time to make
|
||||
# sure gradients flow to the sinks.
|
||||
#
|
||||
# For simplicity also re-create the mask if we have more than 1 query
|
||||
# for now.
|
||||
if self.training or L > 1:
|
||||
self._previous_mask = None
|
||||
return _make_mask(L, min(window_size + 1, offset))
|
||||
|
||||
# We are in inference so cache the mask and try to reuse it
|
||||
if self._previous_mask is None:
|
||||
self._previous_mask = _make_mask(L, window_size + 1)
|
||||
|
||||
return self._previous_mask[..., : min(L + offset, window_size + 1)]
|
||||
|
||||
def get_mask(self, x, cache, window_size, idx):
|
||||
if idx % 2 == 1:
|
||||
return self.get_causal_mask(x, cache)
|
||||
else:
|
||||
return self.get_sliding_window_mask(x, cache, window_size)
|
||||
|
||||
def __call__(self, x: mx.array, mask: mx.array, cache=None) -> mx.array:
|
||||
B, L, _ = x.shape
|
||||
D = self.head_dim
|
||||
Hk = self.num_key_value_heads
|
||||
|
||||
q = self.q_proj(x).reshape(B, L, -1, D).swapaxes(1, 2)
|
||||
k = self.k_proj(x).reshape(B, L, -1, D).swapaxes(1, 2)
|
||||
v = self.v_proj(x).reshape(B, L, -1, D).swapaxes(1, 2)
|
||||
|
||||
# If cache is None or the cache offset is 0 then we need to add a 0 key
|
||||
# and value to make some space for the sink
|
||||
if cache is None or cache.offset == 0:
|
||||
q = self.rope(q)
|
||||
k = self.rope(k)
|
||||
|
||||
attn_output = sdpa(q, k, v, self.sinks, self.sm_scale, mask=mask)
|
||||
zeros = mx.zeros((B, Hk, 1, D), dtype=k.dtype)
|
||||
k = mx.concatenate([zeros, k], axis=2)
|
||||
v = mx.concatenate([zeros, v], axis=2)
|
||||
if cache is not None:
|
||||
k, v = cache.update_and_fetch(k, v)
|
||||
|
||||
# Reshape back to original format: (batch, seqlen, hidden_size)
|
||||
attn_output = attn_output.reshape(*input_shape, -1)
|
||||
out = self.o_proj(attn_output)
|
||||
return out
|
||||
# We have already put the 0 in the cache no need to do anything special
|
||||
else:
|
||||
q = self.rope(q, offset=cache.offset - 1)
|
||||
k = self.rope(k, offset=cache.offset - 1)
|
||||
k, v = cache.update_and_fetch(k, v)
|
||||
|
||||
# NOTE: mask should contain the sink weights already
|
||||
v_hat = scaled_dot_product_attention(q, k, v, cache, self.sm_scale, mask=mask)
|
||||
|
||||
return self.o_proj(v_hat.swapaxes(1, 2).reshape(B, L, -1))
|
||||
|
||||
|
||||
class MLPBlock(nn.Module):
|
||||
@@ -229,8 +267,8 @@ class GptOssMoeModel(nn.Module):
|
||||
super().__init__()
|
||||
self.embed_tokens = nn.Embedding(args.vocab_size, args.hidden_size)
|
||||
self.norm = nn.RMSNorm(args.hidden_size, args.rms_norm_eps)
|
||||
|
||||
self.layers = [TransformerBlock(args) for _ in range(args.num_hidden_layers)]
|
||||
self.window_size = args.sliding_window
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
@@ -248,19 +286,15 @@ class GptOssMoeModel(nn.Module):
|
||||
cache = [None] * len(self.layers)
|
||||
|
||||
if mask is None:
|
||||
full_mask = create_attention_mask(x, cache[1:2], return_array=True)
|
||||
sliding_window_mask = create_attention_mask(x, cache, return_array=True)
|
||||
masks = [
|
||||
l.self_attn.get_mask(x, c, self.window_size, i)
|
||||
for i, (l, c) in enumerate(zip(self.layers, cache))
|
||||
]
|
||||
else:
|
||||
masks = [mask] * len(self.layers)
|
||||
|
||||
for i, (layer, c) in enumerate(zip(self.layers, cache)):
|
||||
local_mask = mask
|
||||
if mask is None and (i % 2 == 1):
|
||||
local_mask = full_mask
|
||||
elif mask is None:
|
||||
local_mask = sliding_window_mask
|
||||
if local_mask is None:
|
||||
local_mask = mx.array([True], dtype=mx.bool_)
|
||||
|
||||
x = layer(x, local_mask, c)
|
||||
for i, (layer, c, m) in enumerate(zip(self.layers, cache, masks)):
|
||||
x = layer(x, m, c)
|
||||
x = self.norm(x)
|
||||
return x
|
||||
|
||||
@@ -377,6 +411,6 @@ class Model(nn.Module):
|
||||
caches.append(KVCache())
|
||||
else:
|
||||
caches.append(
|
||||
RotatingKVCache(max_size=self.args.sliding_window, keep=0)
|
||||
RotatingKVCache(max_size=self.args.sliding_window + 1, keep=1)
|
||||
)
|
||||
return caches
|
||||
|
||||
Reference in New Issue
Block a user