"""
Unit tests for the ProcessWeights class.
Comprehensive test coverage for all weight processing functions extracted from HookedTransformer.
"""
from unittest.mock import Mock, patch
import einops
import pytest
import torch
from transformer_lens.config.transformer_lens_config import TransformerLensConfig
from transformer_lens.weight_processing import ProcessWeights
def deep_copy_state_dict(state_dict):
"""Create a deep copy of a state dict with cloned tensors.
Args:
state_dict: State dict to copy
Returns:
Deep copy of state dict with cloned tensors
"""
return {k: v.clone() if isinstance(v, torch.Tensor) else v for k, v in state_dict.items()}
def assert_state_dicts_equal(dict1, dict2):
"""Compare two state dicts containing tensors.
Args:
dict1: First state dict
dict2: Second state dict
Raises:
AssertionError: If dicts are not equal
"""
assert set(dict1.keys()) == set(
dict2.keys()
), f"Keys differ: {set(dict1.keys()) ^ set(dict2.keys())}"
for key in dict1.keys():
val1, val2 = dict1[key], dict2[key]
if isinstance(val1, torch.Tensor) and isinstance(val2, torch.Tensor):
assert torch.equal(val1, val2), f"Tensors at key '{key}' are not equal"
else:
assert val1 == val2, f"Values at key '{key}' are not equal: {val1} != {val2}"
def create_test_config(**kwargs):
"""Create a test configuration with default values."""
defaults = {
"d_model": 8,
"d_head": 2,
"n_layers": 2,
"n_ctx": 50,
"n_heads": 4,
"d_mlp": 16,
"n_key_value_heads": None,
"attn_only": False,
"gated_mlp": False,
"act_fn": "relu",
"final_rms": False,
"positional_embedding_type": "standard",
"normalization_type": "LN",
"num_experts": None,
}
defaults.update(kwargs)
return TransformerLensConfig(**defaults)
@pytest.fixture
def basic_config():
"""Basic test configuration."""
return create_test_config()
@pytest.fixture
def gqa_config():
"""Configuration with Grouped Query Attention."""
return create_test_config(n_key_value_heads=2)
@pytest.fixture
def attn_only_config():
"""Attention-only configuration."""
return create_test_config(attn_only=True)
@pytest.fixture
def gated_mlp_config():
"""Configuration with gated MLP."""
return create_test_config(gated_mlp=True)
@pytest.fixture
def solu_config():
"""Configuration with SoLU activation."""
return create_test_config(act_fn="solu_ln")
@pytest.fixture
def basic_state_dict(basic_config):
"""Create a basic state dict for testing."""
cfg = basic_config
state_dict = {}
state_dict["embed.W_E"] = torch.randn(100, cfg.d_model)
state_dict["pos_embed.W_pos"] = torch.randn(50, cfg.d_model)
state_dict["unembed.W_U"] = torch.randn(cfg.d_model, 100)
state_dict["unembed.b_U"] = torch.randn(100)
state_dict["ln_final.w"] = torch.randn(cfg.d_model)
state_dict["ln_final.b"] = torch.randn(cfg.d_model)
for l in range(cfg.n_layers):
state_dict[f"blocks.{l}.ln1.w"] = torch.randn(cfg.d_model)
state_dict[f"blocks.{l}.ln1.b"] = torch.randn(cfg.d_model)
state_dict[f"blocks.{l}.ln2.w"] = torch.randn(cfg.d_model)
state_dict[f"blocks.{l}.ln2.b"] = torch.randn(cfg.d_model)
state_dict[f"blocks.{l}.attn.W_Q"] = torch.randn(cfg.n_heads, cfg.d_model, cfg.d_head)
state_dict[f"blocks.{l}.attn.W_K"] = torch.randn(cfg.n_heads, cfg.d_model, cfg.d_head)
state_dict[f"blocks.{l}.attn.W_V"] = torch.randn(cfg.n_heads, cfg.d_model, cfg.d_head)
state_dict[f"blocks.{l}.attn.W_O"] = torch.randn(cfg.n_heads, cfg.d_head, cfg.d_model)
state_dict[f"blocks.{l}.attn.b_Q"] = torch.randn(cfg.n_heads, cfg.d_head)
state_dict[f"blocks.{l}.attn.b_K"] = torch.randn(cfg.n_heads, cfg.d_head)
state_dict[f"blocks.{l}.attn.b_V"] = torch.randn(cfg.n_heads, cfg.d_head)
state_dict[f"blocks.{l}.attn.b_O"] = torch.randn(cfg.d_model)
state_dict[f"blocks.{l}.mlp.W_in"] = torch.randn(cfg.d_model, cfg.d_mlp)
state_dict[f"blocks.{l}.mlp.W_out"] = torch.randn(cfg.d_mlp, cfg.d_model)
state_dict[f"blocks.{l}.mlp.b_in"] = torch.randn(cfg.d_mlp)
state_dict[f"blocks.{l}.mlp.b_out"] = torch.randn(cfg.d_model)
return state_dict
@pytest.fixture
def gqa_state_dict(gqa_config):
"""Create a state dict for GQA testing."""
cfg = gqa_config
state_dict = {}
state_dict["embed.W_E"] = torch.randn(100, cfg.d_model)
state_dict["pos_embed.W_pos"] = torch.randn(50, cfg.d_model)
state_dict["unembed.W_U"] = torch.randn(cfg.d_model, 100)
state_dict["unembed.b_U"] = torch.randn(100)
state_dict["ln_final.w"] = torch.randn(cfg.d_model)
state_dict["ln_final.b"] = torch.randn(cfg.d_model)
for l in range(cfg.n_layers):
state_dict[f"blocks.{l}.ln1.w"] = torch.randn(cfg.d_model)
state_dict[f"blocks.{l}.ln1.b"] = torch.randn(cfg.d_model)
state_dict[f"blocks.{l}.ln2.w"] = torch.randn(cfg.d_model)
state_dict[f"blocks.{l}.ln2.b"] = torch.randn(cfg.d_model)
state_dict[f"blocks.{l}.attn.W_Q"] = torch.randn(cfg.n_heads, cfg.d_model, cfg.d_head)
state_dict[f"blocks.{l}.attn.b_Q"] = torch.randn(cfg.n_heads, cfg.d_head)
state_dict[f"blocks.{l}.attn._W_K"] = torch.randn(
cfg.n_key_value_heads, cfg.d_model, cfg.d_head
)
state_dict[f"blocks.{l}.attn._W_V"] = torch.randn(
cfg.n_key_value_heads, cfg.d_model, cfg.d_head
)
state_dict[f"blocks.{l}.attn._b_K"] = torch.randn(cfg.n_key_value_heads, cfg.d_head)
state_dict[f"blocks.{l}.attn._b_V"] = torch.randn(cfg.n_key_value_heads, cfg.d_head)
state_dict[f"blocks.{l}.attn.W_O"] = torch.randn(cfg.n_heads, cfg.d_head, cfg.d_model)
state_dict[f"blocks.{l}.attn.b_O"] = torch.randn(cfg.d_model)
state_dict[f"blocks.{l}.mlp.W_in"] = torch.randn(cfg.d_model, cfg.d_mlp)
state_dict[f"blocks.{l}.mlp.W_out"] = torch.randn(cfg.d_mlp, cfg.d_model)
state_dict[f"blocks.{l}.mlp.b_in"] = torch.randn(cfg.d_mlp)
state_dict[f"blocks.{l}.mlp.b_out"] = torch.randn(cfg.d_model)
return state_dict
class TestProcessWeights:
"""Test cases for the ProcessWeights class."""
def test_fold_layer_norm_basic(self, basic_config, basic_state_dict):
"""Test basic LayerNorm folding functionality."""
original_dict = deep_copy_state_dict(basic_state_dict)
processed_dict = ProcessWeights.fold_layer_norm(basic_state_dict, basic_config)
assert_state_dicts_equal(basic_state_dict, original_dict)
for l in range(basic_config.n_layers):
assert f"blocks.{l}.ln1.w" in processed_dict
assert torch.allclose(
processed_dict[f"blocks.{l}.ln1.w"],
torch.ones_like(processed_dict[f"blocks.{l}.ln1.w"]),
)
assert f"blocks.{l}.ln1.b" in processed_dict
assert torch.allclose(
processed_dict[f"blocks.{l}.ln1.b"],
torch.zeros_like(processed_dict[f"blocks.{l}.ln1.b"]),
)
assert f"blocks.{l}.ln2.w" in processed_dict
assert torch.allclose(
processed_dict[f"blocks.{l}.ln2.w"],
torch.ones_like(processed_dict[f"blocks.{l}.ln2.w"]),
)
assert f"blocks.{l}.ln2.b" in processed_dict
assert torch.allclose(
processed_dict[f"blocks.{l}.ln2.b"],
torch.zeros_like(processed_dict[f"blocks.{l}.ln2.b"]),
)
assert "ln_final.w" in processed_dict
assert torch.allclose(
processed_dict["ln_final.w"], torch.ones_like(processed_dict["ln_final.w"])
)
assert "ln_final.b" in processed_dict
assert torch.allclose(
processed_dict["ln_final.b"], torch.zeros_like(processed_dict["ln_final.b"])
)
for l in range(basic_config.n_layers):
assert f"blocks.{l}.attn.W_Q" in processed_dict
assert f"blocks.{l}.attn.W_K" in processed_dict
assert f"blocks.{l}.attn.W_V" in processed_dict
assert f"blocks.{l}.mlp.W_in" in processed_dict
assert "unembed.W_U" in processed_dict
def test_fold_layer_norm_no_biases(self, basic_config, basic_state_dict):
"""Test LayerNorm folding without bias folding."""
processed_dict = ProcessWeights.fold_layer_norm(
basic_state_dict, basic_config, fold_biases=False
)
for l in range(basic_config.n_layers):
assert f"blocks.{l}.ln1.w" in processed_dict
assert torch.allclose(
processed_dict[f"blocks.{l}.ln1.w"],
torch.ones_like(processed_dict[f"blocks.{l}.ln1.w"]),
)
assert f"blocks.{l}.ln2.w" in processed_dict
assert torch.allclose(
processed_dict[f"blocks.{l}.ln2.w"],
torch.ones_like(processed_dict[f"blocks.{l}.ln2.w"]),
)
def test_fold_layer_norm_no_centering(self, basic_config, basic_state_dict):
"""Test LayerNorm folding without weight centering."""
processed_dict = ProcessWeights.fold_layer_norm(
basic_state_dict, basic_config, center_weights=False
)
for l in range(basic_config.n_layers):
assert f"blocks.{l}.ln1.w" in processed_dict
assert torch.allclose(
processed_dict[f"blocks.{l}.ln1.w"],
torch.ones_like(processed_dict[f"blocks.{l}.ln1.w"]),
)
assert f"blocks.{l}.attn.W_Q" in processed_dict
def test_fold_layer_norm_attn_only(self, attn_only_config, basic_state_dict):
"""Test LayerNorm folding with attention-only model."""
attn_only_dict = {k: v for k, v in basic_state_dict.items() if "mlp" not in k}
processed_dict = ProcessWeights.fold_layer_norm(attn_only_dict, attn_only_config)
for l in range(attn_only_config.n_layers):
assert f"blocks.{l}.attn.W_Q" in processed_dict
assert f"blocks.{l}.mlp.W_in" not in processed_dict
def test_fold_layer_norm_gated_mlp(self, gated_mlp_config):
"""Test LayerNorm folding with gated MLP."""
state_dict = {}
cfg = gated_mlp_config
for l in range(cfg.n_layers):
state_dict[f"blocks.{l}.ln1.w"] = torch.randn(cfg.d_model)
state_dict[f"blocks.{l}.ln1.b"] = torch.randn(cfg.d_model)
state_dict[f"blocks.{l}.ln2.w"] = torch.randn(cfg.d_model)
state_dict[f"blocks.{l}.ln2.b"] = torch.randn(cfg.d_model)
state_dict[f"blocks.{l}.attn.W_Q"] = torch.randn(cfg.n_heads, cfg.d_model, cfg.d_head)
state_dict[f"blocks.{l}.attn.W_K"] = torch.randn(cfg.n_heads, cfg.d_model, cfg.d_head)
state_dict[f"blocks.{l}.attn.W_V"] = torch.randn(cfg.n_heads, cfg.d_model, cfg.d_head)
state_dict[f"blocks.{l}.attn.b_Q"] = torch.randn(cfg.n_heads, cfg.d_head)
state_dict[f"blocks.{l}.attn.b_K"] = torch.randn(cfg.n_heads, cfg.d_head)
state_dict[f"blocks.{l}.attn.b_V"] = torch.randn(cfg.n_heads, cfg.d_head)
state_dict[f"blocks.{l}.mlp.W_in"] = torch.randn(cfg.d_model, cfg.d_mlp)
state_dict[f"blocks.{l}.mlp.W_gate"] = torch.randn(cfg.d_model, cfg.d_mlp)
state_dict[f"blocks.{l}.mlp.b_in"] = torch.randn(cfg.d_mlp)
state_dict["ln_final.w"] = torch.randn(cfg.d_model)
state_dict["ln_final.b"] = torch.randn(cfg.d_model)
state_dict["unembed.W_U"] = torch.randn(cfg.d_model, 100)
state_dict["unembed.b_U"] = torch.randn(100)
processed_dict = ProcessWeights.fold_layer_norm(state_dict, cfg)
for l in range(cfg.n_layers):
assert f"blocks.{l}.mlp.W_gate" in processed_dict
def test_fold_layer_norm_solu(self, solu_config):
"""Test LayerNorm folding with SoLU activation."""
state_dict = {}
cfg = solu_config
for l in range(cfg.n_layers):
state_dict[f"blocks.{l}.ln1.w"] = torch.randn(cfg.d_model)
state_dict[f"blocks.{l}.ln1.b"] = torch.randn(cfg.d_model)
state_dict[f"blocks.{l}.ln2.w"] = torch.randn(cfg.d_model)
state_dict[f"blocks.{l}.ln2.b"] = torch.randn(cfg.d_model)
state_dict[f"blocks.{l}.attn.W_Q"] = torch.randn(cfg.n_heads, cfg.d_model, cfg.d_head)
state_dict[f"blocks.{l}.attn.W_K"] = torch.randn(cfg.n_heads, cfg.d_model, cfg.d_head)
state_dict[f"blocks.{l}.attn.W_V"] = torch.randn(cfg.n_heads, cfg.d_model, cfg.d_head)
state_dict[f"blocks.{l}.attn.b_Q"] = torch.randn(cfg.n_heads, cfg.d_head)
state_dict[f"blocks.{l}.attn.b_K"] = torch.randn(cfg.n_heads, cfg.d_head)
state_dict[f"blocks.{l}.attn.b_V"] = torch.randn(cfg.n_heads, cfg.d_head)
state_dict[f"blocks.{l}.mlp.W_in"] = torch.randn(cfg.d_model, cfg.d_mlp)
state_dict[f"blocks.{l}.mlp.W_out"] = torch.randn(cfg.d_mlp, cfg.d_model)
state_dict[f"blocks.{l}.mlp.b_in"] = torch.randn(cfg.d_mlp)
state_dict[f"blocks.{l}.mlp.b_out"] = torch.randn(cfg.d_model)
state_dict[f"blocks.{l}.mlp.ln.w"] = torch.randn(cfg.d_mlp)
state_dict[f"blocks.{l}.mlp.ln.b"] = torch.randn(cfg.d_mlp)
state_dict["ln_final.w"] = torch.randn(cfg.d_model)
state_dict["ln_final.b"] = torch.randn(cfg.d_model)
state_dict["unembed.W_U"] = torch.randn(cfg.d_model, 100)
state_dict["unembed.b_U"] = torch.randn(100)
processed_dict = ProcessWeights.fold_layer_norm(state_dict, cfg)
for l in range(cfg.n_layers):
assert f"blocks.{l}.mlp.ln.w" in processed_dict
assert torch.allclose(
processed_dict[f"blocks.{l}.mlp.ln.w"],
torch.ones_like(processed_dict[f"blocks.{l}.mlp.ln.w"]),
)
assert f"blocks.{l}.mlp.ln.b" in processed_dict
assert torch.allclose(
processed_dict[f"blocks.{l}.mlp.ln.b"],
torch.zeros_like(processed_dict[f"blocks.{l}.mlp.ln.b"]),
)
def test_center_writing_weights(self, basic_config, basic_state_dict):
"""Test weight centering functionality."""
processed_dict = ProcessWeights.center_writing_weights(basic_state_dict, basic_config)
embed_mean = processed_dict["embed.W_E"].mean(-1, keepdim=True)
assert torch.allclose(embed_mean, torch.zeros_like(embed_mean), atol=1e-6)
pos_mean = processed_dict["pos_embed.W_pos"].mean(-1, keepdim=True)
assert torch.allclose(pos_mean, torch.zeros_like(pos_mean), atol=1e-6)
for l in range(basic_config.n_layers):
w_o_mean = processed_dict[f"blocks.{l}.attn.W_O"].mean(-1, keepdim=True)
assert torch.allclose(w_o_mean, torch.zeros_like(w_o_mean), atol=1e-6)
b_o_mean = processed_dict[f"blocks.{l}.attn.b_O"].mean()
assert torch.allclose(b_o_mean, torch.tensor(0.0), atol=1e-6)
mlp_out_mean = processed_dict[f"blocks.{l}.mlp.W_out"].mean(-1, keepdim=True)
assert torch.allclose(mlp_out_mean, torch.zeros_like(mlp_out_mean), atol=1e-6)
mlp_b_out_mean = processed_dict[f"blocks.{l}.mlp.b_out"].mean()
assert torch.allclose(mlp_b_out_mean, torch.tensor(0.0), atol=1e-6)
def test_center_writing_weights_rotary(self, basic_config, basic_state_dict):
"""Test weight centering with rotary embeddings."""
basic_config.positional_embedding_type = "rotary"
processed_dict = ProcessWeights.center_writing_weights(basic_state_dict, basic_config)
assert torch.equal(processed_dict["pos_embed.W_pos"], basic_state_dict["pos_embed.W_pos"])
def test_center_writing_weights_attn_only(self, attn_only_config, basic_state_dict):
"""Test weight centering with attention-only model."""
attn_only_dict = {k: v for k, v in basic_state_dict.items() if "mlp" not in k}
processed_dict = ProcessWeights.center_writing_weights(attn_only_dict, attn_only_config)
for l in range(attn_only_config.n_layers):
assert f"blocks.{l}.attn.W_O" in processed_dict
assert f"blocks.{l}.mlp.W_out" not in processed_dict
def test_center_unembed(self, basic_state_dict):
"""Test unembedding weight centering."""
processed_dict = ProcessWeights.center_unembed(basic_state_dict)
w_u_mean = processed_dict["unembed.W_U"].mean(-1, keepdim=True)
assert torch.allclose(w_u_mean, torch.zeros_like(w_u_mean), atol=1e-6)
b_u_mean = processed_dict["unembed.b_U"].mean()
assert torch.allclose(b_u_mean, torch.tensor(0.0), atol=1e-6)
def test_fold_value_biases_basic(self, basic_config, basic_state_dict):
"""Test value bias folding functionality."""
original_dict = deep_copy_state_dict(basic_state_dict)
processed_dict = ProcessWeights.fold_value_biases(basic_state_dict, basic_config)
assert_state_dicts_equal(basic_state_dict, original_dict)
for l in range(basic_config.n_layers):
b_v = processed_dict[f"blocks.{l}.attn.b_V"]
assert torch.allclose(b_v, torch.zeros_like(b_v), atol=1e-6)
assert f"blocks.{l}.attn.b_O" in processed_dict
def test_fold_value_biases_gqa(self, gqa_config, gqa_state_dict):
"""Test value bias folding with GQA."""
processed_dict = ProcessWeights.fold_value_biases(gqa_state_dict, gqa_config)
for l in range(gqa_config.n_layers):
b_v = processed_dict[f"blocks.{l}.attn._b_V"]
assert torch.allclose(b_v, torch.zeros_like(b_v), atol=1e-6)
@pytest.mark.skipif(
not (torch.cuda.is_available() or torch.backends.mps.is_available()),
reason="Cross-device test requires a non-CPU accelerator (CUDA or MPS)",
)
def test_fold_value_biases_cross_device_state_dict(self, basic_config, basic_state_dict):
"""fold_value_biases must align tensors when state_dict has mixed devices.
Regression for #904: an HF model loaded on GPU could end up with state_dict
tensors on different devices (b_V on GPU, b_O on CPU because a downstream
converter created it without an explicit device=). The b_V * W_O multiply
then failed with a cross-device RuntimeError.
"""
accelerator = "cuda" if torch.cuda.is_available() else "mps"
cross_device_dict = deep_copy_state_dict(basic_state_dict)
for l in range(basic_config.n_layers):
cross_device_dict[f"blocks.{l}.attn.b_V"] = cross_device_dict[
f"blocks.{l}.attn.b_V"
].to(accelerator)
cross_device_dict[f"blocks.{l}.attn.W_V"] = cross_device_dict[
f"blocks.{l}.attn.W_V"
].to(accelerator)
processed_dict = ProcessWeights.fold_value_biases(cross_device_dict, basic_config)
for l in range(basic_config.n_layers):
b_v = processed_dict[f"blocks.{l}.attn.b_V"]
assert torch.allclose(b_v, torch.zeros_like(b_v), atol=1e-6)
def test_refactor_factored_attn_matrices(self, basic_config, basic_state_dict):
"""Test attention matrix refactoring."""
original_dict = deep_copy_state_dict(basic_state_dict)
with patch("transformer_lens.weight_processing.FactoredMatrix") as mock_factored_matrix:
mock_instance = Mock()
mock_instance.make_even.return_value.pair = (
torch.randn(basic_config.n_heads, basic_config.d_model + 1, basic_config.d_head),
torch.randn(basic_config.n_heads, basic_config.d_head, basic_config.d_model + 1),
)
mock_factored_matrix.return_value = mock_instance
mock_ov_instance = Mock()
U = torch.randn(basic_config.n_heads, basic_config.d_model, basic_config.d_head)
S = torch.randn(basic_config.n_heads, basic_config.d_head)
Vh = torch.randn(basic_config.n_heads, basic_config.d_head, basic_config.d_model)
mock_ov_instance.svd.return_value = (U, S, Vh)
def factored_matrix_side_effect(*args):
if len(args) == 2 and args[1].shape[-1] == basic_config.d_model + 1:
return mock_instance
else:
return mock_ov_instance
mock_factored_matrix.side_effect = factored_matrix_side_effect
processed_dict = ProcessWeights.refactor_factored_attn_matrices(
basic_state_dict, basic_config
)
assert_state_dicts_equal(basic_state_dict, original_dict)
for l in range(basic_config.n_layers):
assert f"blocks.{l}.attn.W_Q" in processed_dict
assert f"blocks.{l}.attn.W_K" in processed_dict
assert f"blocks.{l}.attn.W_V" in processed_dict
assert f"blocks.{l}.attn.W_O" in processed_dict
b_v = processed_dict[f"blocks.{l}.attn.b_V"]
assert torch.allclose(b_v, torch.zeros_like(b_v), atol=1e-6)
def test_refactor_factored_attn_matrices_rotary_error(self, basic_config, basic_state_dict):
"""Test that refactoring fails with rotary embeddings."""
basic_config.positional_embedding_type = "rotary"
with pytest.raises(
AssertionError, match="You can't refactor the QK circuit when using rotary embeddings"
):
ProcessWeights.refactor_factored_attn_matrices(basic_state_dict, basic_config)
def test_process_weights_full_pipeline(self, basic_config, basic_state_dict):
"""Test the full weight processing pipeline."""
original_dict = deep_copy_state_dict(basic_state_dict)
processed_dict = ProcessWeights.process_weights(basic_state_dict, basic_config)
assert_state_dicts_equal(basic_state_dict, original_dict)
for l in range(basic_config.n_layers):
assert f"blocks.{l}.ln1.w" in processed_dict
assert torch.allclose(
processed_dict[f"blocks.{l}.ln1.w"],
torch.ones_like(processed_dict[f"blocks.{l}.ln1.w"]),
)
assert f"blocks.{l}.ln2.w" in processed_dict
assert torch.allclose(
processed_dict[f"blocks.{l}.ln2.w"],
torch.ones_like(processed_dict[f"blocks.{l}.ln2.w"]),
)
assert "ln_final.w" in processed_dict
assert torch.allclose(
processed_dict["ln_final.w"], torch.ones_like(processed_dict["ln_final.w"])
)
embed_mean = processed_dict["embed.W_E"].mean(-1, keepdim=True)
assert torch.allclose(embed_mean, torch.zeros_like(embed_mean), atol=1e-6)
w_u_mean = processed_dict["unembed.W_U"].mean(-1, keepdim=True)
assert torch.allclose(w_u_mean, torch.zeros_like(w_u_mean), atol=1e-6)
for l in range(basic_config.n_layers):
b_v = processed_dict[f"blocks.{l}.attn.b_V"]
assert torch.allclose(b_v, torch.zeros_like(b_v), atol=1e-6)
def test_process_weights_selective_processing(self, basic_config, basic_state_dict):
"""Test selective processing options."""
processed_dict = ProcessWeights.process_weights(
basic_state_dict,
basic_config,
fold_ln=False,
center_writing_weights=False,
center_unembed=False,
fold_value_biases=False,
)
assert "blocks.0.ln1.w" in processed_dict
embed_mean = processed_dict["embed.W_E"].mean(-1, keepdim=True)
assert not torch.allclose(embed_mean, torch.zeros_like(embed_mean), atol=1e-6)
b_v = processed_dict["blocks.0.attn.b_V"]
assert not torch.allclose(b_v, torch.zeros_like(b_v), atol=1e-6)
def test_process_weights_moe_model(self, basic_config, basic_state_dict):
"""Test processing with MoE model (should skip LayerNorm folding)."""
basic_config.num_experts = 8
processed_dict = ProcessWeights.process_weights(basic_state_dict, basic_config)
assert "blocks.0.ln1.w" in processed_dict
def test_process_weights_rms_norm(self, basic_config, basic_state_dict):
"""Test processing with RMS normalization."""
basic_config.normalization_type = "RMS"
processed_dict = ProcessWeights.process_weights(basic_state_dict, basic_config)
assert "blocks.0.ln1.w" in processed_dict
assert torch.allclose(
processed_dict["blocks.0.ln1.w"], torch.ones_like(processed_dict["blocks.0.ln1.w"])
)
def test_process_weights_final_rms(self, basic_config, basic_state_dict):
"""Test processing with final RMS (should skip writing weight centering)."""
basic_config.final_rms = True
processed_dict = ProcessWeights.process_weights(basic_state_dict, basic_config)
embed_mean = processed_dict["embed.W_E"].mean(-1, keepdim=True)
assert not torch.allclose(embed_mean, torch.zeros_like(embed_mean), atol=1e-6)
def test_tensor_shapes_preserved(self, basic_config, basic_state_dict):
"""Test that tensor shapes are preserved correctly."""
processed_dict = ProcessWeights.process_weights(basic_state_dict, basic_config)
assert processed_dict["embed.W_E"].shape == basic_state_dict["embed.W_E"].shape
assert processed_dict["unembed.W_U"].shape == basic_state_dict["unembed.W_U"].shape
for l in range(basic_config.n_layers):
assert (
processed_dict[f"blocks.{l}.attn.W_Q"].shape
== basic_state_dict[f"blocks.{l}.attn.W_Q"].shape
)
assert (
processed_dict[f"blocks.{l}.attn.b_O"].shape
== basic_state_dict[f"blocks.{l}.attn.b_O"].shape
)
def test_mathematical_correctness_layer_norm_folding(self, basic_config):
"""Test mathematical correctness of LayerNorm folding."""
cfg = basic_config
state_dict = {}
ln_w = torch.tensor([2.0, 3.0, 1.0, 0.5])
ln_b = torch.tensor([0.1, 0.2, 0.3, 0.4])
w_q = torch.ones(2, 4, 2)
b_q = torch.zeros(2, 2)
cfg.d_model = 4
cfg.n_heads = 2
cfg.d_head = 2
cfg.n_layers = 1
state_dict["blocks.0.ln1.w"] = ln_w
state_dict["blocks.0.ln1.b"] = ln_b
state_dict["blocks.0.attn.W_Q"] = w_q
state_dict["blocks.0.attn.b_Q"] = b_q
state_dict["blocks.0.ln2.w"] = torch.ones(4)
state_dict["blocks.0.ln2.b"] = torch.zeros(4)
state_dict["blocks.0.attn.W_K"] = torch.ones(2, 4, 2)
state_dict["blocks.0.attn.W_V"] = torch.ones(2, 4, 2)
state_dict["blocks.0.attn.b_K"] = torch.zeros(2, 2)
state_dict["blocks.0.attn.b_V"] = torch.zeros(2, 2)
state_dict["blocks.0.mlp.W_in"] = torch.ones(4, 8)
state_dict["blocks.0.mlp.b_in"] = torch.zeros(8)
state_dict["ln_final.w"] = torch.ones(4)
state_dict["ln_final.b"] = torch.zeros(4)
state_dict["unembed.W_U"] = torch.ones(4, 10)
state_dict["unembed.b_U"] = torch.zeros(10)
processed_dict = ProcessWeights.fold_layer_norm(state_dict, cfg, center_weights=False)
expected_w_q = w_q * ln_w[None, :, None]
expected_b_q = b_q + (w_q * ln_b[None, :, None]).sum(-2)
assert torch.allclose(processed_dict["blocks.0.attn.W_Q"], expected_w_q)
assert torch.allclose(processed_dict["blocks.0.attn.b_Q"], expected_b_q)
processed_dict_centered = ProcessWeights.fold_layer_norm(
state_dict, cfg, center_weights=True
)
w_q_centered = processed_dict_centered["blocks.0.attn.W_Q"]
w_q_mean = einops.reduce(
w_q_centered, "head_index d_model d_head -> head_index 1 d_head", "mean"
)
assert torch.allclose(w_q_mean, torch.zeros_like(w_q_mean), atol=1e-6)
def test_edge_cases_empty_state_dict(self, basic_config):
"""Test handling of edge cases like empty state dicts."""
empty_dict = {}
try:
ProcessWeights.center_unembed(empty_dict)
ProcessWeights.center_writing_weights(empty_dict, basic_config)
except KeyError:
pass
def test_fold_layer_no_adapter_transformer_lens_format(self, basic_config):
"""Test _fold_layer function with no adapter (TransformerLens format).
This test locks in the current behavior of _fold_layer when no adapter is provided,
ensuring that HookedTransformer models continue to work correctly.
"""
cfg = basic_config
cfg.n_layers = 1
cfg.d_model = 4
cfg.n_heads = 2
cfg.d_head = 2
cfg.d_mlp = 8
state_dict = {}
ln1_w = torch.tensor([2.0, 3.0, 1.0, 0.5])
ln1_b = torch.tensor([0.1, 0.2, 0.3, 0.4])
ln2_w = torch.tensor([1.5, 2.5, 0.8, 1.2])
ln2_b = torch.tensor([0.05, 0.15, 0.25, 0.35])
w_q = torch.tensor(
[
[[1.0, 2.0], [3.0, 4.0], [5.0, 6.0], [7.0, 8.0]],
[[2.0, 3.0], [4.0, 5.0], [6.0, 7.0], [8.0, 9.0]],
]
)
w_k = torch.tensor(
[
[[0.5, 1.0], [1.5, 2.0], [2.5, 3.0], [3.5, 4.0]],
[[1.0, 1.5], [2.0, 2.5], [3.0, 3.5], [4.0, 4.5]],
]
)
w_v = torch.tensor(
[
[[0.8, 1.2], [1.6, 2.0], [2.4, 2.8], [3.2, 3.6]],
[[1.2, 1.6], [2.0, 2.4], [2.8, 3.2], [3.6, 4.0]],
]
)
b_q = torch.tensor([[0.1, 0.2], [0.3, 0.4]])
b_k = torch.tensor([[0.05, 0.15], [0.25, 0.35]])
b_v = torch.tensor([[0.08, 0.12], [0.16, 0.20]])
w_in = torch.tensor(
[
[1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0],
[2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0],
[3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0, 10.0],
[4.0, 5.0, 6.0, 7.0, 8.0, 9.0, 10.0, 11.0],
]
)
b_in = torch.tensor([0.1, 0.2, 0.3, 0.4, 0.5, 0.6, 0.7, 0.8])
state_dict["blocks.0.ln1.w"] = ln1_w
state_dict["blocks.0.ln1.b"] = ln1_b
state_dict["blocks.0.ln2.w"] = ln2_w
state_dict["blocks.0.ln2.b"] = ln2_b
state_dict["blocks.0.attn.W_Q"] = w_q
state_dict["blocks.0.attn.W_K"] = w_k
state_dict["blocks.0.attn.W_V"] = w_v
state_dict["blocks.0.attn.b_Q"] = b_q
state_dict["blocks.0.attn.b_K"] = b_k
state_dict["blocks.0.attn.b_V"] = b_v
state_dict["blocks.0.mlp.W_in"] = w_in
state_dict["blocks.0.mlp.b_in"] = b_in
original_state_dict = {k: v.clone() for k, v in state_dict.items()}
ProcessWeights._fold_layer(
state_dict,
cfg,
layer_idx=0,
fold_biases=True,
center_weights=True,
adapter=None,
gqa="",
)
assert "blocks.0.ln1.w" in state_dict
assert torch.allclose(
state_dict["blocks.0.ln1.w"], torch.ones_like(state_dict["blocks.0.ln1.w"])
)
assert "blocks.0.ln1.b" in state_dict
assert torch.allclose(
state_dict["blocks.0.ln1.b"], torch.zeros_like(state_dict["blocks.0.ln1.b"])
)
assert "blocks.0.ln2.w" in state_dict
assert torch.allclose(
state_dict["blocks.0.ln2.w"], torch.ones_like(state_dict["blocks.0.ln2.w"])
)
assert "blocks.0.ln2.b" in state_dict
assert torch.allclose(
state_dict["blocks.0.ln2.b"], torch.zeros_like(state_dict["blocks.0.ln2.b"])
)
w_q_processed = state_dict["blocks.0.attn.W_Q"]
w_k_processed = state_dict["blocks.0.attn.W_K"]
w_v_processed = state_dict["blocks.0.attn.W_V"]
expected_w_q_folded = w_q * ln1_w[None, :, None]
expected_w_k_folded = w_k * ln1_w[None, :, None]
expected_w_v_folded = w_v * ln1_w[None, :, None]
w_q_mean = einops.reduce(
w_q_processed, "head_index d_model d_head -> head_index 1 d_head", "mean"
)
w_k_mean = einops.reduce(
w_k_processed, "head_index d_model d_head -> head_index 1 d_head", "mean"
)
w_v_mean = einops.reduce(
w_v_processed, "head_index d_model d_head -> head_index 1 d_head", "mean"
)
assert torch.allclose(w_q_mean, torch.zeros_like(w_q_mean), atol=1e-6)
assert torch.allclose(w_k_mean, torch.zeros_like(w_k_mean), atol=1e-6)
assert torch.allclose(w_v_mean, torch.zeros_like(w_v_mean), atol=1e-6)
b_q_processed = state_dict["blocks.0.attn.b_Q"]
b_k_processed = state_dict["blocks.0.attn.b_K"]
b_v_processed = state_dict["blocks.0.attn.b_V"]
expected_b_q_folded = b_q + (w_q * ln1_b[None, :, None]).sum(-2)
expected_b_k_folded = b_k + (w_k * ln1_b[None, :, None]).sum(-2)
expected_b_v_folded = b_v + (w_v * ln1_b[None, :, None]).sum(-2)
assert torch.allclose(b_q_processed, expected_b_q_folded, atol=1e-6)
assert torch.allclose(b_k_processed, expected_b_k_folded, atol=1e-6)
assert torch.allclose(b_v_processed, expected_b_v_folded, atol=1e-6)
w_in_processed = state_dict["blocks.0.mlp.W_in"]
b_in_processed = state_dict["blocks.0.mlp.b_in"]
expected_w_in_folded = w_in * ln2_w[:, None]
expected_w_in_centered = expected_w_in_folded - einops.reduce(
expected_w_in_folded, "d_model d_mlp -> 1 d_mlp", "mean"
)
assert torch.allclose(w_in_processed, expected_w_in_centered, atol=1e-6)
expected_b_in_folded = b_in + (w_in * ln2_b[:, None]).sum(-2)
assert torch.allclose(b_in_processed, expected_b_in_folded, atol=1e-6)
w_in_mean = einops.reduce(w_in_processed, "d_model d_mlp -> 1 d_mlp", "mean")
assert torch.allclose(w_in_mean, torch.zeros_like(w_in_mean), atol=1e-6)
for k, v in original_state_dict.items():
assert torch.equal(v, original_state_dict[k])
def test_fold_layer_no_adapter_without_centering(self, basic_config):
"""Test _fold_layer function without weight centering to verify pure folding behavior."""
cfg = basic_config
cfg.n_layers = 1
state_dict = {}
ln1_w = torch.tensor([2.0, 3.0, 1.0, 0.5])
ln1_b = torch.tensor([0.1, 0.2, 0.3, 0.4])
w_q = torch.ones(2, 4, 2)
b_q = torch.zeros(2, 2)
state_dict["blocks.0.ln1.w"] = ln1_w
state_dict["blocks.0.ln1.b"] = ln1_b
state_dict["blocks.0.attn.W_Q"] = w_q
state_dict["blocks.0.attn.b_Q"] = b_q
state_dict["blocks.0.ln2.w"] = torch.ones(4)
state_dict["blocks.0.ln2.b"] = torch.zeros(4)
state_dict["blocks.0.attn.W_K"] = torch.ones(2, 4, 2)
state_dict["blocks.0.attn.W_V"] = torch.ones(2, 4, 2)
state_dict["blocks.0.attn.b_K"] = torch.zeros(2, 2)
state_dict["blocks.0.attn.b_V"] = torch.zeros(2, 2)
state_dict["blocks.0.mlp.W_in"] = torch.ones(4, 8)
state_dict["blocks.0.mlp.b_in"] = torch.zeros(8)
ProcessWeights._fold_layer(
state_dict,
cfg,
layer_idx=0,
fold_biases=True,
center_weights=False,
adapter=None,
gqa="",
)
expected_w_q = w_q * ln1_w[None, :, None]
expected_b_q = b_q + (w_q * ln1_b[None, :, None]).sum(-2)
assert torch.allclose(state_dict["blocks.0.attn.W_Q"], expected_w_q, atol=1e-6)
assert torch.allclose(state_dict["blocks.0.attn.b_Q"], expected_b_q, atol=1e-6)
assert "blocks.0.ln1.w" in state_dict
assert torch.allclose(
state_dict["blocks.0.ln1.w"], torch.ones_like(state_dict["blocks.0.ln1.w"])
)
assert "blocks.0.ln1.b" in state_dict
assert torch.allclose(
state_dict["blocks.0.ln1.b"], torch.zeros_like(state_dict["blocks.0.ln1.b"])
)
assert "blocks.0.ln2.w" in state_dict
assert torch.allclose(
state_dict["blocks.0.ln2.w"], torch.ones_like(state_dict["blocks.0.ln2.w"])
)
assert "blocks.0.ln2.b" in state_dict
assert torch.allclose(
state_dict["blocks.0.ln2.b"], torch.zeros_like(state_dict["blocks.0.ln2.b"])
)
def test_fold_layer_no_adapter_without_bias_folding(self, basic_config):
"""Test _fold_layer function without bias folding."""
cfg = basic_config
cfg.n_layers = 1
state_dict = {}
ln1_w = torch.tensor([2.0, 3.0, 1.0, 0.5])
ln1_b = torch.tensor([0.1, 0.2, 0.3, 0.4])
w_q = torch.ones(2, 4, 2)
b_q = torch.zeros(2, 2)
state_dict["blocks.0.ln1.w"] = ln1_w
state_dict["blocks.0.ln1.b"] = ln1_b
state_dict["blocks.0.attn.W_Q"] = w_q
state_dict["blocks.0.attn.b_Q"] = b_q
state_dict["blocks.0.ln2.w"] = torch.ones(4)
state_dict["blocks.0.ln2.b"] = torch.zeros(4)
state_dict["blocks.0.attn.W_K"] = torch.ones(2, 4, 2)
state_dict["blocks.0.attn.W_V"] = torch.ones(2, 4, 2)
state_dict["blocks.0.attn.b_K"] = torch.zeros(2, 2)
state_dict["blocks.0.attn.b_V"] = torch.zeros(2, 2)
state_dict["blocks.0.mlp.W_in"] = torch.ones(4, 8)
state_dict["blocks.0.mlp.b_in"] = torch.zeros(8)
ProcessWeights._fold_layer(
state_dict,
cfg,
layer_idx=0,
fold_biases=False,
center_weights=True,
adapter=None,
gqa="",
)
expected_w_q_folded = w_q * ln1_w[None, :, None]
expected_w_q_centered = expected_w_q_folded - einops.reduce(
expected_w_q_folded, "head_index d_model d_head -> head_index 1 d_head", "mean"
)
assert torch.allclose(state_dict["blocks.0.attn.W_Q"], expected_w_q_centered, atol=1e-6)
assert torch.allclose(state_dict["blocks.0.attn.b_Q"], b_q, atol=1e-6)
assert "blocks.0.ln1.w" in state_dict
assert torch.allclose(
state_dict["blocks.0.ln1.w"], torch.ones_like(state_dict["blocks.0.ln1.w"])
)
assert "blocks.0.ln1.b" in state_dict
assert "blocks.0.ln2.w" in state_dict
assert torch.allclose(
state_dict["blocks.0.ln2.w"], torch.ones_like(state_dict["blocks.0.ln2.w"])
)
assert "blocks.0.ln2.b" in state_dict
def test_center_unembed_padded_vocab_orientation():
"""Padded-vocab unembeds (rows > cfg.d_vocab, e.g. HyenaDNA's 12->16) must
still center along the vocab axis — the d_vocab shape match fails, so
orientation falls back to the unambiguous d_model axis."""
from types import SimpleNamespace
import torch
from transformer_lens.weight_processing import ProcessWeights
torch.manual_seed(0)
w = torch.randn(16, 128)
cfg = SimpleNamespace(d_vocab=12, d_model=128)
out = ProcessWeights.center_unembed({"unembed.W_U": w.clone()}, cfg=cfg, adapter=None)
centered = out["unembed.W_U"]
assert centered.mean(dim=0).abs().max().item() < 1e-6
assert centered.shape == w.shape
def test_resolve_state_dict_key_dense_mlp_fallback():
"""Mixed MoE/dense adapters (Llama4) name a non-MoE layer's gated MLP
projections dense_gate/dense_in/dense_out. When the standard mlp.{in,gate,
out} key is absent, the resolver must fall back to the dense_ variant — but
only then, and only if the dense key actually exists (no false rewrite)."""
from transformer_lens.weight_processing import ProcessWeights
sd = {"blocks.0.mlp.in.weight": 1, "blocks.0.mlp.dense_in.weight": 2}
assert ProcessWeights._resolve_state_dict_key(sd, "blocks.0.mlp.in.weight", 0) == (
"blocks.0.mlp.in.weight"
)
for name in ("in", "gate", "out"):
sd2 = {f"blocks.0.mlp.dense_{name}.weight": 2}
assert (
ProcessWeights._resolve_state_dict_key(sd2, f"blocks.0.mlp.{name}.weight", 0)
== f"blocks.0.mlp.dense_{name}.weight"
)
sd3 = {"blocks.0.attn.q.weight": 1}
assert ProcessWeights._resolve_state_dict_key(sd3, "blocks.0.mlp.in.weight", 0) == (
"blocks.0.mlp.in.weight"
)
sd4 = {"blocks.0.mlp.gate.weight": 1, "blocks.0.mlp.dense_gate.weight": 2}
assert ProcessWeights._resolve_state_dict_key(sd4, "blocks.0.mlp.gate.weight", 0) == (
"blocks.0.mlp.gate.weight"
)
sd5 = {"blocks.0.mlp.shared_expert.dense_in.weight": 1}
assert (
ProcessWeights._resolve_state_dict_key(sd5, "blocks.0.mlp.shared_expert.in.weight", 0)
== "blocks.0.mlp.shared_expert.in.weight"
)