Files
OBLITERATUS/obliteratus/mlx_backend.py
T

467 lines
14 KiB
Python

"""Optional MLX backend for Apple Silicon native inference and weight editing.
MLX is Apple's array framework that runs natively on the Apple Neural Engine
and Metal GPU. When available, it provides significantly faster inference and
weight manipulation than PyTorch's MPS backend on Apple hardware.
This module is entirely optional — if ``mlx`` / ``mlx-lm`` are not installed,
``MLX_AVAILABLE`` is ``False`` and all public functions raise ``RuntimeError``.
Install with::
pip install mlx>=0.22 mlx-lm>=0.20
"""
from __future__ import annotations
import logging
from pathlib import Path
from typing import Any, Callable
logger = logging.getLogger(__name__)
# ---------------------------------------------------------------------------
# Availability check
# ---------------------------------------------------------------------------
MLX_AVAILABLE = False
_mx = None # mlx module
_mlx_lm = None # mlx-lm module
_mlx_nn = None # mlx.nn module
try:
import mlx.core as _mx_core # type: ignore[import-untyped]
import mlx.nn as _mlx_nn_mod # type: ignore[import-untyped]
import mlx_lm # type: ignore[import-untyped]
_mx = _mx_core
_mlx_nn = _mlx_nn_mod
_mlx_lm = mlx_lm
MLX_AVAILABLE = True
logger.info("MLX backend available (mlx %s)", _mx.__version__ if hasattr(_mx, "__version__") else "?")
except ImportError:
pass
def _require_mlx() -> None:
if not MLX_AVAILABLE:
raise RuntimeError(
"MLX backend is not available. Install with: pip install mlx>=0.22 mlx-lm>=0.20"
)
# ---------------------------------------------------------------------------
# Model loading
# ---------------------------------------------------------------------------
class MLXModelHandle:
"""Lightweight wrapper around an MLX-loaded model + tokenizer."""
def __init__(self, model: Any, tokenizer: Any, model_name: str):
self.model = model
self.tokenizer = tokenizer
self.model_name = model_name
@property
def config(self) -> Any:
return getattr(self.model, "config", None)
def load_model(
model_name: str,
dtype: str = "float16",
) -> MLXModelHandle:
"""Load a HuggingFace model via ``mlx-lm`` for Apple-native execution.
Parameters
----------
model_name : str
HuggingFace model identifier (e.g. ``"meta-llama/Llama-3.2-3B-Instruct"``).
dtype : str
One of ``"float16"``, ``"bfloat16"``, ``"float32"``.
Returns
-------
MLXModelHandle
Wrapper with ``.model`` and ``.tokenizer`` attributes.
"""
_require_mlx()
from mlx_lm import load # type: ignore[import-untyped]
logger.info("Loading %s via MLX (dtype=%s)", model_name, dtype)
model, tokenizer = load(model_name)
return MLXModelHandle(model=model, tokenizer=tokenizer, model_name=model_name)
# ---------------------------------------------------------------------------
# Inference
# ---------------------------------------------------------------------------
def generate(
handle: MLXModelHandle,
prompt: str,
max_tokens: int = 256,
temperature: float = 0.7,
top_p: float = 0.9,
repetition_penalty: float | None = None,
) -> str:
"""Generate text using the MLX model.
Parameters
----------
handle : MLXModelHandle
A loaded MLX model handle.
prompt : str
The input prompt string.
max_tokens : int
Maximum number of tokens to generate.
temperature : float
Sampling temperature.
top_p : float
Nucleus sampling threshold.
repetition_penalty : float or None
Repetition penalty (1.0 = no penalty).
Returns
-------
str
Generated text completion.
"""
_require_mlx()
from mlx_lm import generate as mlx_generate # type: ignore[import-untyped]
kwargs: dict[str, Any] = {
"max_tokens": max_tokens,
"temp": temperature,
"top_p": top_p,
}
if repetition_penalty is not None:
kwargs["repetition_penalty"] = repetition_penalty
return mlx_generate(
handle.model,
handle.tokenizer,
prompt=prompt,
**kwargs,
)
# ---------------------------------------------------------------------------
# Activation extraction
# ---------------------------------------------------------------------------
def get_activations(
handle: MLXModelHandle,
prompts: list[str],
layer_indices: list[int],
max_length: int = 256,
) -> dict[int, list[Any]]:
"""Extract hidden-state activations from specified layers.
Uses MLX's computation graph to capture intermediate outputs.
Parameters
----------
handle : MLXModelHandle
Loaded model.
prompts : list[str]
Input prompts to probe.
layer_indices : list[int]
Which transformer layers to capture.
max_length : int
Maximum sequence length for tokenization.
Returns
-------
dict[int, list[mlx.core.array]]
Mapping from layer index to list of activation arrays (one per prompt).
Each array has shape ``(hidden_dim,)`` — the last-token hidden state.
"""
_require_mlx()
import mlx.core as mx # type: ignore[import-untyped]
model = handle.model
tokenizer = handle.tokenizer
# Identify the transformer block list
layers = None
for attr in ("model.layers", "transformer.h", "gpt_neox.layers"):
obj = model
try:
for part in attr.split("."):
obj = getattr(obj, part)
layers = obj
break
except AttributeError:
continue
if layers is None:
raise RuntimeError(
"Cannot locate transformer layers in the MLX model. "
"Supported architectures: LLaMA, GPT-2, GPT-NeoX."
)
activations: dict[int, list[Any]] = {idx: [] for idx in layer_indices}
target_set = set(layer_indices)
for prompt in prompts:
tokens = tokenizer.encode(prompt)
if len(tokens) > max_length:
tokens = tokens[:max_length]
input_ids = mx.array([tokens])
# Forward through embedding
if hasattr(model, "model"):
# LLaMA-style: model.model.embed_tokens
embed_module = model.model
elif hasattr(model, "transformer"):
embed_module = model.transformer
else:
embed_module = model
if hasattr(embed_module, "embed_tokens"):
h = embed_module.embed_tokens(input_ids)
elif hasattr(embed_module, "wte"):
h = embed_module.wte(input_ids)
else:
raise RuntimeError("Cannot find embedding layer in MLX model")
# Walk through layers, capturing activations at target indices
for i, layer in enumerate(layers):
h = layer(h)
# Some layers return tuples (hidden, attention) — take first
if isinstance(h, tuple):
h = h[0]
if i in target_set:
# Last token hidden state
last_hidden = h[0, -1, :]
mx.eval(last_hidden) # force evaluation
activations[i].append(last_hidden)
return activations
# ---------------------------------------------------------------------------
# Weight manipulation
# ---------------------------------------------------------------------------
def get_weight(handle: MLXModelHandle, layer_idx: int, param_path: str) -> Any:
"""Retrieve a weight matrix from the model.
Parameters
----------
handle : MLXModelHandle
Loaded model.
layer_idx : int
Transformer layer index.
param_path : str
Dot-separated path within the layer, e.g. ``"self_attn.o_proj.weight"``.
Returns
-------
mlx.core.array
The weight tensor.
"""
_require_mlx()
model = handle.model
# Navigate to the layer
layers = _get_layers(model)
layer = layers[layer_idx]
# Navigate the param path
obj = layer
for part in param_path.split("."):
obj = getattr(obj, part)
return obj
def modify_weights(
handle: MLXModelHandle,
layer_idx: int,
param_path: str,
modifier_fn: Callable[[Any], Any],
) -> None:
"""Modify a weight matrix in-place using a function.
Parameters
----------
handle : MLXModelHandle
Loaded model.
layer_idx : int
Transformer layer index.
param_path : str
Dot-separated path within the layer to the weight, e.g.
``"self_attn.o_proj.weight"``.
modifier_fn : callable
Function that takes the current weight (mlx array) and returns the
modified weight (mlx array). For abliteration, this would project
out the refusal direction.
"""
_require_mlx()
import mlx.core as mx # type: ignore[import-untyped]
model = handle.model
layers = _get_layers(model)
layer = layers[layer_idx]
# Navigate to the parent module and leaf attribute
parts = param_path.split(".")
parent = layer
for part in parts[:-1]:
parent = getattr(parent, part)
leaf_name = parts[-1]
old_weight = getattr(parent, leaf_name)
new_weight = modifier_fn(old_weight)
# MLX uses a functional update pattern
if hasattr(parent, "update"):
parent.update({leaf_name: new_weight})
else:
setattr(parent, leaf_name, new_weight)
mx.eval(new_weight) # materialize
def project_out_direction(weight: Any, direction: Any) -> Any:
"""Project a direction out of a weight matrix (abliteration).
Given weight matrix W and unit direction d, computes::
W' = W - (W @ d) outer d
Parameters
----------
weight : mlx.core.array
Weight matrix, shape ``(out_features, in_features)``.
direction : mlx.core.array
Unit direction vector, shape ``(in_features,)``.
Returns
-------
mlx.core.array
Modified weight with direction projected out.
"""
_require_mlx()
import mlx.core as mx # type: ignore[import-untyped]
d = direction.astype(weight.dtype)
# W @ d gives the component along d for each row
proj = mx.matmul(weight, d[:, None]) # (out, 1)
return weight - mx.matmul(proj, d[None, :]) # (out, in)
# ---------------------------------------------------------------------------
# Save model
# ---------------------------------------------------------------------------
def save_model(
handle: MLXModelHandle,
output_dir: str | Path,
upload_repo: str | None = None,
) -> Path:
"""Save the (modified) MLX model to disk.
Saves in safetensors format compatible with both MLX and HuggingFace.
Parameters
----------
handle : MLXModelHandle
Model handle (possibly with modified weights).
output_dir : str or Path
Directory to save into.
upload_repo : str or None
If set, also uploads to HuggingFace Hub.
Returns
-------
Path
The output directory.
"""
_require_mlx()
out = Path(output_dir)
out.mkdir(parents=True, exist_ok=True)
# mlx-lm's save uses safetensors
if hasattr(_mlx_lm, "save_model"):
_mlx_lm.save_model(str(out), handle.model, handle.tokenizer)
else:
# Fallback: manual save via mlx.core.save_safetensors
import mlx.core as mx # type: ignore[import-untyped]
weights = dict(handle.model.parameters())
flat = {}
_flatten_dict(weights, "", flat)
mx.save_safetensors(str(out / "model.safetensors"), flat)
# Save tokenizer via transformers
handle.tokenizer.save_pretrained(str(out))
logger.info("MLX model saved to %s", out)
if upload_repo:
try:
from mlx_lm import upload_to_hub # type: ignore[import-untyped]
upload_to_hub(str(out), upload_repo)
logger.info("Uploaded to %s", upload_repo)
except (ImportError, AttributeError):
logger.warning("mlx-lm upload not available — push manually with huggingface-cli")
return out
# ---------------------------------------------------------------------------
# Conversion: PyTorch ↔ MLX
# ---------------------------------------------------------------------------
def torch_tensor_to_mlx(tensor: "torch.Tensor") -> Any: # noqa: F821
"""Convert a PyTorch tensor to an MLX array."""
_require_mlx()
import mlx.core as mx # type: ignore[import-untyped]
# Move to CPU and convert via numpy
np_array = tensor.detach().cpu().float().numpy()
return mx.array(np_array)
def mlx_to_torch_tensor(array: Any, device: str = "cpu") -> "torch.Tensor": # noqa: F821
"""Convert an MLX array to a PyTorch tensor."""
import numpy as np
import torch
np_array = np.array(array)
return torch.from_numpy(np_array).to(device)
# ---------------------------------------------------------------------------
# Internal helpers
# ---------------------------------------------------------------------------
def _get_layers(model: Any) -> Any:
"""Locate the transformer block list in an MLX model."""
for attr_path in ("model.layers", "transformer.h", "gpt_neox.layers"):
obj = model
try:
for part in attr_path.split("."):
obj = getattr(obj, part)
return obj
except AttributeError:
continue
raise RuntimeError("Cannot locate transformer layers in MLX model")
def _flatten_dict(d: dict, prefix: str, out: dict) -> None:
"""Flatten a nested dict with dot-separated keys."""
for k, v in d.items():
key = f"{prefix}{k}" if prefix else k
if isinstance(v, dict):
_flatten_dict(v, f"{key}.", out)
else:
out[key] = v