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:
Daniel Nakov
2025-09-26 15:22:12 -07:00
committed by GitHub
co-authored by dnakov Awni Hannun
parent 358b4d2ab5
commit dcb4b9ba6d
3 changed files with 81 additions and 7 deletions
+1 -1
View File
@@ -1,3 +1,3 @@
# Copyright © 2023-2025 Apple Inc.
__version__ = "0.28.0"
__version__ = "0.28.1"
+37 -6
View File
@@ -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
]
+43
View File
@@ -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))