botp
/

Solomon / mlx /tests /test_tiny_decoder.py
orz99's picture ArcherHume's picture
Duplicate from DoccyHealth/Solomon
1d2de8a
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