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}"
        )