mirror of
https://github.com/elder-plinius/OBLITERATUS.git
synced 2026-08-18 00:47:23 +02:00
171 lines
6.9 KiB
Python
171 lines
6.9 KiB
Python
"""Simulated MLX contracts that run without Apple hardware or MLX packages."""
|
|
|
|
from __future__ import annotations
|
|
|
|
import sys
|
|
from types import ModuleType, SimpleNamespace
|
|
from unittest.mock import Mock
|
|
|
|
import numpy as np
|
|
import pytest
|
|
import torch
|
|
|
|
from obliteratus import mlx_backend
|
|
|
|
|
|
@pytest.fixture
|
|
def fake_mlx(monkeypatch):
|
|
core = ModuleType("mlx.core")
|
|
core.array = np.array
|
|
core.eval = Mock()
|
|
core.matmul = np.matmul
|
|
core.save_safetensors = Mock()
|
|
nn = ModuleType("mlx.nn")
|
|
package = ModuleType("mlx")
|
|
package.core = core
|
|
package.nn = nn
|
|
lm = ModuleType("mlx_lm")
|
|
lm.load = Mock(return_value=(SimpleNamespace(config={"kind": "fixture"}), SimpleNamespace()))
|
|
lm.generate = Mock(return_value="generated")
|
|
lm.upload_to_hub = Mock()
|
|
monkeypatch.setitem(sys.modules, "mlx", package)
|
|
monkeypatch.setitem(sys.modules, "mlx.core", core)
|
|
monkeypatch.setitem(sys.modules, "mlx.nn", nn)
|
|
monkeypatch.setitem(sys.modules, "mlx_lm", lm)
|
|
monkeypatch.setattr(mlx_backend, "MLX_AVAILABLE", True)
|
|
monkeypatch.setattr(mlx_backend, "_mx", core)
|
|
monkeypatch.setattr(mlx_backend, "_mlx_nn", nn)
|
|
monkeypatch.setattr(mlx_backend, "_mlx_lm", lm)
|
|
return SimpleNamespace(core=core, lm=lm)
|
|
|
|
|
|
@pytest.mark.parametrize(
|
|
"call",
|
|
[
|
|
lambda: mlx_backend.load_model("model"),
|
|
lambda: mlx_backend.generate(SimpleNamespace(), "prompt"),
|
|
lambda: mlx_backend.get_activations(SimpleNamespace(), [], []),
|
|
lambda: mlx_backend.get_weight(SimpleNamespace(), 0, "weight"),
|
|
lambda: mlx_backend.modify_weights(SimpleNamespace(), 0, "weight", lambda value: value),
|
|
lambda: mlx_backend.project_out_direction(None, None),
|
|
lambda: mlx_backend.save_model(SimpleNamespace(), "out"),
|
|
lambda: mlx_backend.torch_tensor_to_mlx(torch.ones(1)),
|
|
],
|
|
)
|
|
def test_missing_mlx_fails_with_install_guidance(monkeypatch, call):
|
|
monkeypatch.setattr(mlx_backend, "MLX_AVAILABLE", False)
|
|
with pytest.raises(RuntimeError, match="pip install mlx>=0.22 mlx-lm>=0.20"):
|
|
call()
|
|
|
|
|
|
def test_load_generate_and_handle_config(fake_mlx):
|
|
handle = mlx_backend.load_model("local/model", dtype="bfloat16")
|
|
assert handle.model_name == "local/model"
|
|
assert handle.config == {"kind": "fixture"}
|
|
fake_mlx.lm.load.assert_called_once_with("local/model")
|
|
assert mlx_backend.generate(
|
|
handle,
|
|
"prompt",
|
|
max_tokens=4,
|
|
temperature=0.2,
|
|
top_p=0.8,
|
|
repetition_penalty=1.1,
|
|
) == "generated"
|
|
assert fake_mlx.lm.generate.call_args.kwargs == {
|
|
"prompt": "prompt",
|
|
"max_tokens": 4,
|
|
"temp": 0.2,
|
|
"top_p": 0.8,
|
|
"repetition_penalty": 1.1,
|
|
}
|
|
mlx_backend.generate(handle, "prompt", repetition_penalty=None)
|
|
assert "repetition_penalty" not in fake_mlx.lm.generate.call_args.kwargs
|
|
|
|
|
|
class _Layer:
|
|
def __call__(self, value):
|
|
return (value + 1, "attention")
|
|
|
|
|
|
def test_activation_capture_supports_llama_and_truncation(fake_mlx):
|
|
inner = SimpleNamespace(
|
|
layers=[_Layer(), _Layer()],
|
|
embed_tokens=lambda ids: np.repeat(ids[:, :, None], 3, axis=2),
|
|
)
|
|
handle = mlx_backend.MLXModelHandle(
|
|
model=SimpleNamespace(model=inner),
|
|
tokenizer=SimpleNamespace(encode=lambda _prompt: [1, 2, 3, 4]),
|
|
model_name="fixture",
|
|
)
|
|
activations = mlx_backend.get_activations(handle, ["a", "b"], [0, 1], max_length=2)
|
|
assert np.array_equal(activations[0][0], np.array([3, 3, 3]))
|
|
assert np.array_equal(activations[1][1], np.array([4, 4, 4]))
|
|
assert fake_mlx.core.eval.call_count == 4
|
|
|
|
|
|
def test_activation_capture_rejects_unknown_layers_and_embeddings(fake_mlx):
|
|
handle = mlx_backend.MLXModelHandle(SimpleNamespace(), SimpleNamespace(), "unknown")
|
|
with pytest.raises(RuntimeError, match="Cannot locate transformer layers"):
|
|
mlx_backend.get_activations(handle, ["x"], [0])
|
|
|
|
model = SimpleNamespace(model=SimpleNamespace(layers=[_Layer()]))
|
|
handle = mlx_backend.MLXModelHandle(model, SimpleNamespace(encode=lambda _prompt: [1]), "unknown")
|
|
with pytest.raises(RuntimeError, match="Cannot find embedding layer"):
|
|
mlx_backend.get_activations(handle, ["x"], [0])
|
|
|
|
|
|
def test_get_and_modify_weights_support_update_and_assignment(fake_mlx):
|
|
update_parent = SimpleNamespace(weight=np.array([1.0]), update=Mock())
|
|
assign_parent = SimpleNamespace(weight=np.array([2.0]))
|
|
layers = [SimpleNamespace(attention=update_parent), SimpleNamespace(attention=assign_parent)]
|
|
handle = mlx_backend.MLXModelHandle(
|
|
SimpleNamespace(transformer=SimpleNamespace(h=layers)),
|
|
SimpleNamespace(),
|
|
"fixture",
|
|
)
|
|
assert np.array_equal(mlx_backend.get_weight(handle, 0, "attention.weight"), np.array([1.0]))
|
|
mlx_backend.modify_weights(handle, 0, "attention.weight", lambda value: value + 1)
|
|
update_parent.update.assert_called_once()
|
|
mlx_backend.modify_weights(handle, 1, "attention.weight", lambda value: value + 1)
|
|
assert np.array_equal(assign_parent.weight, np.array([3.0]))
|
|
assert fake_mlx.core.eval.call_count == 2
|
|
|
|
|
|
def test_projection_removes_requested_direction(fake_mlx):
|
|
weight = np.array([[2.0, 3.0], [4.0, 5.0]])
|
|
direction = np.array([1.0, 0.0])
|
|
projected = mlx_backend.project_out_direction(weight, direction)
|
|
assert np.array_equal(projected, np.array([[0.0, 3.0], [0.0, 5.0]]))
|
|
|
|
|
|
def test_save_model_native_fallback_and_upload(fake_mlx, monkeypatch, tmp_path):
|
|
tokenizer = SimpleNamespace(save_pretrained=Mock())
|
|
model = SimpleNamespace(parameters=lambda: {"block": {"weight": np.array([1.0])}}.items())
|
|
handle = mlx_backend.MLXModelHandle(model, tokenizer, "fixture")
|
|
fake_mlx.lm.save_model = Mock()
|
|
out = mlx_backend.save_model(handle, tmp_path / "native", upload_repo="org/model")
|
|
assert out.is_dir()
|
|
fake_mlx.lm.save_model.assert_called_once()
|
|
fake_mlx.lm.upload_to_hub.assert_called_once_with(str(out), "org/model")
|
|
|
|
del fake_mlx.lm.save_model
|
|
out = mlx_backend.save_model(handle, tmp_path / "fallback")
|
|
fake_mlx.core.save_safetensors.assert_called_once()
|
|
tokenizer.save_pretrained.assert_called_once_with(str(out))
|
|
|
|
|
|
def test_tensor_conversions_and_internal_helpers(fake_mlx):
|
|
source = torch.tensor([1.0, 2.0], requires_grad=True)
|
|
converted = mlx_backend.torch_tensor_to_mlx(source)
|
|
assert np.array_equal(converted, np.array([1.0, 2.0], dtype=np.float32))
|
|
restored = mlx_backend.mlx_to_torch_tensor(np.array([3.0]), device="cpu")
|
|
assert torch.equal(restored, torch.tensor([3.0], dtype=torch.float64))
|
|
|
|
layers = [object()]
|
|
assert mlx_backend._get_layers(SimpleNamespace(gpt_neox=SimpleNamespace(layers=layers))) is layers
|
|
with pytest.raises(RuntimeError, match="Cannot locate transformer layers"):
|
|
mlx_backend._get_layers(SimpleNamespace())
|
|
flattened = {}
|
|
mlx_backend._flatten_dict({"a": {"b": 1}, "c": 2}, "", flattened)
|
|
assert flattened == {"a.b": 1, "c": 2}
|