Solomon / mlx /tests /test_tiny_decoder.py
kelseyway's picture
Make MLX adapter-only and include the reproducible BF16 converter
5c0a4a8
Raw
History Blame Contribute Delete
2.37 kB
"""Exercise actual attention, convolution and recurrent layers on Metal."""
import mlx.core as mx
import numpy as np
from mlx_vlm.models.qwen3_5.config import TextConfig
from mlx_vlm.models.qwen3_5.language import LanguageModel
from solomon_mlx.engine import Engine, fork_cache
def test_chunked_hybrid_decoder_and_cache_isolation():
mx.random.seed(8)
config = TextConfig(
model_type="qwen3_5_text",
hidden_size=128,
intermediate_size=192,
linear_num_value_heads=2,
linear_num_key_heads=2,
linear_key_head_dim=32,
linear_value_head_dim=32,
linear_conv_kernel_dim=4,
num_hidden_layers=4,
num_attention_heads=4,
num_key_value_heads=2,
head_dim=32,
rms_norm_eps=1e-6,
vocab_size=256,
max_position_embeddings=512,
rope_parameters={
"type": "default",
"mrope_section": [2, 1, 1],
"rope_theta": 100000,
"partial_rotary_factor": 0.25,
},
)
engine = Engine.__new__(Engine)
engine.lm = LanguageModel(config)
engine.lm.eval()
class ForbiddenVocabularyHead:
def __call__(self, *args, **kwargs):
raise AssertionError("Vocabulary projection must never run")
engine.lm.lm_head = ForbiddenVocabularyHead()
engine.chunk_size = 16
engine.context = {"start": None}
tokens = list(range(1, 74))
p = 51
prefix = engine.lm.make_cache()
engine.forward(tokens[:p], engine.positions(0, p), prefix)
before = [tuple(None if x is None else np.asarray(x).copy() for x in c.state) for c in prefix]
a, _ = engine.forward(tokens[p:], engine.positions(p, len(tokens) - p), fork_cache(prefix))
b, _ = engine.forward(tokens, engine.positions(0, len(tokens)), engine.lm.make_cache())
np.testing.assert_allclose(np.asarray(a), np.asarray(b), atol=5e-4, rtol=5e-4)
again, _ = engine.forward(tokens[p:], engine.positions(p, len(tokens) - p), fork_cache(prefix))
np.testing.assert_array_equal(np.asarray(a), np.asarray(again))
for original, saved in zip(prefix, before):
for x, y in zip(original.state, saved):
if y is not None:
np.testing.assert_array_equal(np.asarray(x), y)
for c in prefix:
if hasattr(c, "cache"):
assert c.cache[1].dtype == mx.float32