Add Code World Model support (#505)
* Add sliding-window support to LLaMA * nits * version --------- Co-authored-by: dnakov <[email protected]> Co-authored-by: Awni Hannun <[email protected]>
This commit is contained in:
co-authored by
dnakov
Awni Hannun
parent
358b4d2ab5
commit
dcb4b9ba6d
+1
-1
@@ -1,3 +1,3 @@
|
||||
# Copyright © 2023-2025 Apple Inc.
|
||||
|
||||
__version__ = "0.28.0"
|
||||
__version__ = "0.28.1"
|
||||
|
||||
+37
-6
@@ -1,12 +1,13 @@
|
||||
# Copyright © 2023-2024 Apple Inc.
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Dict, Optional, Union
|
||||
from typing import Any, Dict, List, Optional, Union
|
||||
|
||||
import mlx.core as mx
|
||||
import mlx.nn as nn
|
||||
|
||||
from .base import BaseModelArgs, create_attention_mask, scaled_dot_product_attention
|
||||
from .cache import KVCache, RotatingKVCache
|
||||
from .rope_utils import initialize_rope
|
||||
|
||||
|
||||
@@ -28,11 +29,16 @@ class ModelArgs(BaseModelArgs):
|
||||
rope_traditional: bool = False
|
||||
rope_scaling: Optional[Dict[str, Union[float, str]]] = None
|
||||
tie_word_embeddings: bool = True
|
||||
layer_types: Optional[List[str]] = None
|
||||
sliding_window: Optional[int] = None
|
||||
|
||||
def __post_init__(self):
|
||||
if self.num_key_value_heads is None:
|
||||
self.num_key_value_heads = self.num_attention_heads
|
||||
|
||||
if self.layer_types is None:
|
||||
self.layer_types = ["full_attention"] * self.num_hidden_layers
|
||||
|
||||
|
||||
class Attention(nn.Module):
|
||||
def __init__(self, args: ModelArgs):
|
||||
@@ -114,10 +120,11 @@ class MLP(nn.Module):
|
||||
|
||||
|
||||
class TransformerBlock(nn.Module):
|
||||
def __init__(self, args: ModelArgs):
|
||||
def __init__(self, args: ModelArgs, use_sliding: bool = False):
|
||||
super().__init__()
|
||||
self.num_attention_heads = args.num_attention_heads
|
||||
self.hidden_size = args.hidden_size
|
||||
self.use_sliding = use_sliding
|
||||
self.self_attn = Attention(args)
|
||||
self.mlp = MLP(args)
|
||||
self.input_layernorm = nn.RMSNorm(args.hidden_size, eps=args.rms_norm_eps)
|
||||
@@ -145,12 +152,21 @@ class LlamaModel(nn.Module):
|
||||
self.args = args
|
||||
self.vocab_size = args.vocab_size
|
||||
self.num_hidden_layers = args.num_hidden_layers
|
||||
self.layer_types = args.layer_types
|
||||
self.sliding_window = args.sliding_window
|
||||
assert self.vocab_size > 0
|
||||
self.embed_tokens = nn.Embedding(args.vocab_size, args.hidden_size)
|
||||
self.layers = [
|
||||
TransformerBlock(args=args) for _ in range(args.num_hidden_layers)
|
||||
TransformerBlock(args=args, use_sliding=layer_type == "sliding_attention")
|
||||
for layer_type in self.layer_types
|
||||
]
|
||||
self.norm = nn.RMSNorm(args.hidden_size, eps=args.rms_norm_eps)
|
||||
self.fa_idx = self.layer_types.index("full_attention")
|
||||
self.swa_idx = None
|
||||
for e, l in enumerate(self.layers):
|
||||
if l.use_sliding:
|
||||
self.swa_idx = e
|
||||
break
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
@@ -166,10 +182,15 @@ class LlamaModel(nn.Module):
|
||||
if cache is None:
|
||||
cache = [None] * len(self.layers)
|
||||
|
||||
mask = create_attention_mask(h, cache[0])
|
||||
fa_mask = create_attention_mask(h, cache[self.fa_idx])
|
||||
if self.swa_idx is not None:
|
||||
swa_mask = create_attention_mask(
|
||||
h, cache[self.swa_idx], window_size=self.sliding_window
|
||||
)
|
||||
|
||||
for layer, c in zip(self.layers, cache):
|
||||
h = layer(h, mask, cache=c)
|
||||
for layer, cache in zip(self.layers, cache):
|
||||
mask = swa_mask if layer.use_sliding else fa_mask
|
||||
h = layer(h, mask, cache=cache)
|
||||
|
||||
return self.norm(h)
|
||||
|
||||
@@ -208,3 +229,13 @@ class Model(nn.Module):
|
||||
@property
|
||||
def layers(self):
|
||||
return self.model.layers
|
||||
|
||||
def make_cache(self):
|
||||
return [
|
||||
(
|
||||
RotatingKVCache(max_size=self.model.sliding_window)
|
||||
if layer.use_sliding
|
||||
else KVCache()
|
||||
)
|
||||
for layer in self.layers
|
||||
]
|
||||
|
||||
@@ -175,6 +175,49 @@ class TestModels(unittest.TestCase):
|
||||
sums = mask.sum(axis=1)
|
||||
self.assertTrue(mx.array_equal(sums, expected_sums))
|
||||
|
||||
def test_llama_model_sliding_attention(self):
|
||||
from mlx_lm.models import llama
|
||||
|
||||
args = llama.ModelArgs(
|
||||
model_type="llama",
|
||||
hidden_size=64,
|
||||
num_hidden_layers=4,
|
||||
intermediate_size=256,
|
||||
num_attention_heads=8,
|
||||
num_key_value_heads=4,
|
||||
rms_norm_eps=1e-5,
|
||||
vocab_size=128,
|
||||
sliding_window=4,
|
||||
layer_types=[
|
||||
"full_attention",
|
||||
"sliding_attention",
|
||||
"sliding_attention",
|
||||
"full_attention",
|
||||
],
|
||||
tie_word_embeddings=False,
|
||||
rope_theta=10000.0,
|
||||
)
|
||||
model = llama.Model(args)
|
||||
|
||||
tokens = mx.array([[1, 2, 3, 4, 5]], dtype=mx.int32)
|
||||
out = model(tokens)
|
||||
mx.eval(out)
|
||||
self.assertEqual(out.shape, (1, 5, args.vocab_size))
|
||||
|
||||
caches = model.make_cache()
|
||||
self.assertIsInstance(caches[0], KVCache)
|
||||
self.assertIsInstance(caches[1], RotatingKVCache)
|
||||
self.assertIsInstance(caches[2], RotatingKVCache)
|
||||
self.assertIsInstance(caches[3], KVCache)
|
||||
|
||||
caches = model.make_cache()
|
||||
step = model(tokens[:, :2], cache=caches)
|
||||
mx.eval(step)
|
||||
step = model(tokens[:, 2:3], cache=caches)
|
||||
mx.eval(step)
|
||||
self.assertEqual(caches[0].offset, 3)
|
||||
self.assertEqual(caches[1].offset, 3)
|
||||
|
||||
def test_rope(self):
|
||||
rope = rope_utils.initialize_rope(32, base=100, traditional=False)
|
||||
self.assertTrue(isinstance(rope, nn.RoPE))
|
||||
|
||||
Reference in New Issue
Block a user