"""Apple Silicon MPS smoke tests for TransformerLens.
Design principles:
- All tests skip automatically on non-MPS runners (Linux, Windows, CPU-only Macs)
- Only float32 is used (bfloat16 is unsupported on MPS)
- Only small models are loaded (roneneldan/TinyStories-1M, ~50MB)
- torch.mps.empty_cache() + gc.collect() between tests to stay within memory budget
- TRANSFORMERLENS_ALLOW_MPS=1 must be set for get_device() to return "mps"
CI: These tests are run via the `mps-checks` job in .github/workflows/checks.yml
which sets TRANSFORMERLENS_ALLOW_MPS=1 and runs on macos-latest.
"""
import gc
import os
import warnings
import pytest
import torch
from transformer_lens.tools.analysis.jacobian_lens_coordinate_patch import (
solve_coordinate_patch,
)
from transformer_lens.tools.analysis.jacobian_lens_decomposition import (
get_sparse_decomposition,
)
from transformer_lens.tools.analysis.projection_kernel import (
SubspaceBasis,
projection_kernel,
)
pytestmark = pytest.mark.skipif(
not torch.backends.mps.is_available(),
reason="MPS not available on this runner — skipping Apple Silicon tests",
)
SMALL_MODEL = "roneneldan/TinyStories-1M"
def _load_tiny_model(device: str = "mps"):
"""Load TinyStories-1M on the given device with float32 (bfloat16 unsupported on MPS)."""
from transformer_lens import HookedTransformer
return HookedTransformer.from_pretrained(SMALL_MODEL, device=device, dtype=torch.float32)
def _cleanup(model=None):
"""Free GPU memory between tests."""
if model is not None:
del model
torch.mps.empty_cache()
gc.collect()
def test_mps_device_available():
"""Sanity check: MPS backend is present and built on this runner."""
assert torch.backends.mps.is_available(), "MPS not available"
assert torch.backends.mps.is_built(), "MPS not built into this PyTorch"
def test_mps_get_device_returns_mps_with_env_var():
"""get_device() auto-selects MPS when TRANSFORMERLENS_ALLOW_MPS=1 is set."""
from transformer_lens.utilities.devices import get_device
original = os.environ.get("TRANSFORMERLENS_ALLOW_MPS", "")
try:
os.environ["TRANSFORMERLENS_ALLOW_MPS"] = "1"
device = get_device()
assert isinstance(device, str)
assert device == "mps", f"Expected 'mps', got '{device}'"
finally:
if original:
os.environ["TRANSFORMERLENS_ALLOW_MPS"] = original
else:
os.environ.pop("TRANSFORMERLENS_ALLOW_MPS", None)
def test_mps_get_device_falls_back_to_cpu_without_env_var():
"""get_device() falls back to CPU when TRANSFORMERLENS_ALLOW_MPS is unset (safety default)."""
from transformer_lens.utilities.devices import get_device
original = os.environ.get("TRANSFORMERLENS_ALLOW_MPS", "")
try:
os.environ.pop("TRANSFORMERLENS_ALLOW_MPS", None)
device = get_device()
assert isinstance(device, str)
assert (
device == "cpu"
), f"Without TRANSFORMERLENS_ALLOW_MPS=1, get_device() should return 'cpu' not '{device}'"
finally:
if original:
os.environ["TRANSFORMERLENS_ALLOW_MPS"] = original
def test_mps_jspace_decomposition_moves_before_float64_conversion():
"""The NNLS work moves to CPU before its unsupported float64 conversion."""
dictionary = torch.eye(5, device="mps")[:4]
activation = 2.0 * dictionary[0] + 3.0 * dictionary[1]
result = get_sparse_decomposition(activation, dictionary, k=2)
assert result.reconstruction.device.type == "mps"
torch.testing.assert_close(result.reconstruction, activation)
_cleanup()
def test_mps_jspace_coordinate_patch_stays_on_device():
"""Coordinate patching keeps vector outputs on MPS while diagnostics run safely."""
dictionary = torch.eye(5, device="mps")[:4]
activation = 2.0 * dictionary[0] + 3.0 * dictionary[1]
decomposition = get_sparse_decomposition(activation, dictionary, k=2)
result = solve_coordinate_patch(activation, dictionary, 0, 2, decomposition=decomposition)
assert result.patched.device.type == "mps"
torch.testing.assert_close(
result.patched, torch.tensor([0.0, 3.0, 2.0, 0.0, 0.0], device="mps")
)
_cleanup()
def test_mps_warn_if_mps_emits_warning_without_env_var():
"""warn_if_mps() emits a UserWarning when MPS is used without the env var."""
import transformer_lens.utilities.devices as devices_module
from transformer_lens.utilities import warn_if_mps
original = os.environ.get("TRANSFORMERLENS_ALLOW_MPS", "")
original_warned = devices_module._mps_warned
try:
os.environ.pop("TRANSFORMERLENS_ALLOW_MPS", None)
devices_module._mps_warned = False
with warnings.catch_warnings(record=True) as w:
warnings.simplefilter("always")
warn_if_mps("mps")
assert any(
"MPS backend" in str(warning.message) for warning in w
), "Expected MPS warning but got: " + str([str(x.message) for x in w])
finally:
if original:
os.environ["TRANSFORMERLENS_ALLOW_MPS"] = original
devices_module._mps_warned = original_warned
def test_mps_tensor_basic_operations():
"""Basic tensor arithmetic runs on the Metal GPU without errors."""
x = torch.randn(16, 32, device="mps", dtype=torch.float32)
y = torch.randn(16, 32, device="mps", dtype=torch.float32)
z = x + y
assert z.device.type == "mps"
w = torch.matmul(x, y.T)
assert w.device.type == "mps"
assert w.shape == (16, 16)
z_cpu = z.cpu()
assert z_cpu.device.type == "cpu"
_cleanup()
def test_mps_projection_kernel_principal_angles():
"""Projection Kernel computes principal angles and preserves the MPS device."""
try:
basis = torch.eye(4, device="mps", dtype=torch.float32)[:, :2]
subspace = SubspaceBasis(
basis=basis,
singular_values=torch.ones(2, device="mps"),
rank=2,
measured_rank=2,
rtol=4 * torch.finfo(torch.float32).eps,
threshold=4 * torch.finfo(torch.float32).eps,
input_shape=(4, 2),
)
result = projection_kernel(subspace, subspace)
assert result.score.device.type == "mps"
assert result.normalized.device.type == "mps"
assert result.cosines.device.type == "mps"
assert result.angles.device.type == "mps"
assert result.cosines.cpu().tolist() == pytest.approx([1.0, 1.0])
finally:
_cleanup()
def test_mps_softmax_and_layernorm():
"""Softmax and LayerNorm — core transformer ops — work on MPS."""
x = torch.randn(4, 16, 64, device="mps", dtype=torch.float32)
softmax_out = torch.nn.functional.softmax(x, dim=-1)
assert softmax_out.device.type == "mps"
assert torch.allclose(softmax_out.sum(dim=-1), torch.ones(4, 16, device="mps"), atol=1e-5)
ln = torch.nn.LayerNorm(64).to("mps")
ln_out = ln(x)
assert ln_out.device.type == "mps"
_cleanup()
def test_mps_model_forward_pass():
"""TinyStories-1M loads and runs a forward pass on the Metal GPU."""
model = _load_tiny_model(device="mps")
tokens = model.to_tokens("Once upon a time")
assert tokens.device.type == "mps", f"Tokens should be on MPS, got {tokens.device}"
logits = model(tokens)
assert logits.device.type == "mps", f"Logits should be on MPS, got {logits.device}"
assert logits.shape[-1] == model.cfg.d_vocab
assert not torch.isnan(logits).any(), "NaN values in logits — possible MPS compute error"
_cleanup(model)
def test_mps_run_with_cache():
"""run_with_cache() returns cache tensors on the Metal GPU."""
model = _load_tiny_model(device="mps")
tokens = model.to_tokens("The quick brown fox")
logits, cache = model.run_with_cache(tokens)
assert logits.device.type == "mps"
hook_q = cache["blocks.0.attn.hook_q"]
assert hook_q.device.type == "mps", f"Cache tensor not on MPS: {hook_q.device}"
assert not torch.isnan(hook_q).any(), "NaN in attention query cache"
_cleanup(model)
def test_mps_activation_hook_fires_on_metal():
"""run_with_hooks() fires hooks and hook tensors are on the Metal GPU."""
model = _load_tiny_model(device="mps")
tokens = model.to_tokens("Apple Silicon rocks")
hook_devices = []
hook_shapes = []
def capture_hook(value, hook):
hook_devices.append(value.device.type)
hook_shapes.append(value.shape)
return value
model.run_with_hooks(
tokens,
fwd_hooks=[
("blocks.0.attn.hook_q", capture_hook),
("blocks.0.mlp.hook_post", capture_hook),
],
)
assert len(hook_devices) == 2, f"Expected 2 hooks to fire, got {len(hook_devices)}"
for device in hook_devices:
assert device == "mps", f"Hook tensor not on MPS: {device}"
_cleanup(model)
def test_mps_float32_inference():
"""Explicit float32 model loads and infers correctly on MPS."""
model = _load_tiny_model(device="mps")
for name, param in model.named_parameters():
assert param.dtype == torch.float32, f"Parameter {name} has wrong dtype: {param.dtype}"
tokens = model.to_tokens("Testing float32 on Metal")
logits = model(tokens)
assert logits.dtype == torch.float32
_cleanup(model)
def test_mps_loss_computation():
"""Loss computation (return_type='loss') works on MPS."""
model = _load_tiny_model(device="mps")
loss = model("Once upon a time in a land", return_type="loss")
assert isinstance(loss, torch.Tensor)
assert loss.device.type == "mps"
assert not torch.isnan(loss), f"NaN loss — possible MPS compute error: {loss}"
assert loss.item() > 0, "Loss should be positive"
_cleanup(model)