File size: 7,742 Bytes
2dd5f57 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 | """
Unit tests for the loading module in the Speculators library.
"""
import pytest
import torch
from transformers import AutoModelForCausalLM
from speculators.utils.loading import (
_resolve_file,
_resolve_key,
is_config_only_dir,
load_model_layers,
)
# Test model from HuggingFace
TEST_MODEL_REPO = "nm-testing/tiny-testing-random-weights"
SMALL_MODEL_REPO = "nm-testing/tinysmokellama-3.2"
# is_config_only_dir Tests
@pytest.mark.smoke
def test_is_config_only_dir(tmp_path):
# Missing directory and a directory without config.json are not config-only.
assert is_config_only_dir(tmp_path / "does-not-exist") is False
assert is_config_only_dir(tmp_path) is False
# config.json present, no weights -> config-only.
(tmp_path / "config.json").write_text("{}")
assert is_config_only_dir(tmp_path) is True
# A weight file makes it a full checkpoint.
(tmp_path / "model.safetensors").write_text("")
assert is_config_only_dir(tmp_path) is False
@pytest.mark.smoke
def test_is_config_only_dir_detects_bin_weights(tmp_path):
(tmp_path / "config.json").write_text("{}")
(tmp_path / "pytorch_model.bin").write_text("")
assert is_config_only_dir(tmp_path) is False
@pytest.mark.smoke
@pytest.mark.parametrize(
"index_file",
["model.safetensors.index.json", "pytorch_model.bin.index.json"],
)
def test_is_config_only_dir_detects_sharded_index(tmp_path, index_file):
# A sharded-checkpoint manifest ends in .json (so it dodges the *.safetensors /
# *.bin globs); it must still count as weights, not config-only.
(tmp_path / "config.json").write_text("{}")
(tmp_path / index_file).write_text("{}")
assert is_config_only_dir(tmp_path) is False
# _resolve_key Tests
FAKE_WEIGHT_MAP = {
"model.embed_tokens.weight": "shard-0.safetensors",
"model.layers.0.self_attn.q_proj.weight": "shard-0.safetensors",
"tok_embeddings.weight": "shard-1.safetensors",
"output.weight": "shard-1.safetensors",
"norm.weight": "shard-1.safetensors",
}
@pytest.mark.smoke
def test_resolve_key_exact_match():
assert _resolve_key("model.embed_tokens.weight", FAKE_WEIGHT_MAP) == (
"model.embed_tokens.weight"
)
@pytest.mark.smoke
def test_resolve_key_suffix_match():
assert _resolve_key("self_attn.q_proj.weight", FAKE_WEIGHT_MAP) == (
"model.layers.0.self_attn.q_proj.weight"
)
@pytest.mark.smoke
def test_resolve_key_alias_exact():
wm = {"tok_embeddings.weight": "shard.safetensors"}
assert _resolve_key("embed_tokens.weight", wm) == "tok_embeddings.weight"
@pytest.mark.smoke
def test_resolve_key_alias_suffix():
wm = {"model.tok_embeddings.weight": "shard.safetensors"}
assert _resolve_key("embed_tokens.weight", wm) == "model.tok_embeddings.weight"
@pytest.mark.smoke
def test_resolve_key_all_aliases():
wm_lm = {"output.weight": "s.safetensors"}
assert _resolve_key("lm_head.weight", wm_lm) == "output.weight"
wm_norm = {"norm.weight": "s.safetensors"}
assert _resolve_key("model.norm.weight", wm_norm) == "norm.weight"
@pytest.mark.smoke
def test_resolve_key_miss():
assert _resolve_key("nonexistent.weight", FAKE_WEIGHT_MAP) is None
@pytest.mark.smoke
def test_resolve_key_prefers_exact_over_alias():
wm = {
"embed_tokens.weight": "shard-0.safetensors",
"tok_embeddings.weight": "shard-1.safetensors",
}
assert _resolve_key("embed_tokens.weight", wm) == "embed_tokens.weight"
@pytest.mark.smoke
def test_resolve_key_prefers_shortest_suffix():
"""When several keys share the searched suffix, the shortest (most specific)
one wins via ``min(matches, key=len)``.
None of the keys match an alias here, so resolution falls through to the
generic ``norm.weight`` suffix scan and the tie-break is exercised directly
rather than short-circuited by an alias hit (see ``test_resolve_key_llm_aliases``
for the alias path).
"""
wm = {
"model.audio.final_norm.weight": "shard-a.safetensors",
"model.text.norm.weight": "shard-b.safetensors",
}
# Both keys end in "norm.weight" (reached via the model.norm.weight ->
# norm.weight alias); the shorter, more-specific key must win over the audio
# tower's norm.
assert _resolve_key("model.norm.weight", wm) == "model.text.norm.weight"
@pytest.mark.smoke
def test_resolve_key_llm_aliases():
"""Inkling-style keys with llm. prefix resolve correctly."""
wm = {
"model.llm.embed.weight": "shard-0.safetensors",
"model.llm.unembed.weight": "shard-1.safetensors",
"model.llm.norm.weight": "shard-2.safetensors",
}
assert _resolve_key("embed_tokens.weight", wm) == "model.llm.embed.weight"
assert _resolve_key("lm_head.weight", wm) == "model.llm.unembed.weight"
assert _resolve_key("model.norm.weight", wm) == "model.llm.norm.weight"
# _resolve_file Tests
@pytest.mark.sanity
def test_resolve_file_hub_download():
"""Test resolving a file from HuggingFace Hub using real model."""
result = _resolve_file(TEST_MODEL_REPO, "config.json")
assert result.exists()
assert result.name == "config.json"
# load_model_layers Tests
@pytest.mark.sanity
@pytest.mark.parametrize(
"test_model_repo",
[
TEST_MODEL_REPO, # Multi-shard model
SMALL_MODEL_REPO, # Single-shard model
],
)
def test_load_model(test_model_repo: str):
"""Test loading layers from a model repository."""
result = load_model_layers(
["model.embed_tokens.weight", "lm_head.weight"],
test_model_repo,
)
assert len(result) == 2
assert "model.embed_tokens.weight" in result
assert "lm_head.weight" in result
assert isinstance(result["model.embed_tokens.weight"], torch.Tensor)
assert isinstance(result["lm_head.weight"], torch.Tensor)
# Both should have same vocab dimension
assert (
result["model.embed_tokens.weight"].shape[0]
== result["lm_head.weight"].shape[0]
)
# Verify CPU device
assert result["model.embed_tokens.weight"].device.type == "cpu"
@pytest.mark.sanity
def test_load_model_layers_matches_full_model():
"""Test that tensors loaded via utility match those from full model loading."""
# Load full model
full_model = AutoModelForCausalLM.from_pretrained(
TEST_MODEL_REPO,
torch_dtype="auto",
)
# Get state dict from full model
state_dict = full_model.state_dict()
# Load specific layers using our utility
layer_names = [
"model.embed_tokens.weight",
"lm_head.weight",
"model.norm.weight",
"model.layers.0.input_layernorm.weight",
"model.layers.0.mlp.gate_proj.weight",
"model.layers.1.mlp.down_proj.weight",
]
loaded_tensors = load_model_layers(layer_names, TEST_MODEL_REPO)
# Compare each tensor
for layer_name in layer_names:
assert layer_name in loaded_tensors, f"Layer {layer_name} not loaded"
assert layer_name in state_dict, f"Layer {layer_name} not in state_dict"
util_tensor = loaded_tensors[layer_name]
model_tensor = state_dict[layer_name]
# Check dtype matches
assert util_tensor.dtype == model_tensor.dtype, (
f"Dtype mismatch for {layer_name}: "
f"{util_tensor.dtype} vs {model_tensor.dtype}"
)
# Check shape matches
assert util_tensor.shape == model_tensor.shape, (
f"Shape mismatch for {layer_name}: "
f"{util_tensor.shape} vs {model_tensor.shape}"
)
# Check values are identical
assert torch.equal(util_tensor, model_tensor), (
f"Tensor values don't match for {layer_name}"
)
|