File size: 5,385 Bytes
c2e902a
 
 
 
6466ca1
c2e902a
 
 
 
6466ca1
c2e902a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
5d84e23
 
c2e902a
 
 
 
 
 
 
 
 
 
 
5d84e23
c2e902a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
6466ca1
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
import json
from pathlib import Path

import pytest
from safetensors.torch import save_file

from MVP.qwen35_prune import (
    ParameterReport,
    build_text_config,
    count_parameter_groups,
    choose_prefix,
    translate_text_key,
)
from MVP.validate_checkpoint import validate_state_dict_keys


FIXTURE = Path(__file__).parent / "fixtures" / "qwen35_metadata.json"


def load_fixture():
    return json.loads(FIXTURE.read_text())


def test_measured_n4_prefix_is_inside_required_interval():
    data = load_fixture()
    report = ParameterReport(
        embedding_params=data["embedding_params"],
        layer_params=tuple(
            [data["linear_attention_params"]] * 3
            + [data["full_attention_params"]]
            + [data["linear_attention_params"]] * 3
            + [data["full_attention_params"]] * 5
        ),
        layer_types=tuple(data["layer_types"]),
        final_norm_params=data["final_norm_params"],
        all_named_params=data["all_named_params"],
    )

    choice = choose_prefix(report, 330_000_000, 350_000_000)

    assert choice.layer_count == 4
    assert choice.parameter_count == 337_299_424
    assert 330_000_000 <= choice.parameter_count <= 350_000_000


def test_prefix_selection_rejects_incomplete_hybrid_block():
    data = load_fixture()
    report = ParameterReport(
        embedding_params=data["embedding_params"],
        layer_params=(data["linear_attention_params"],) * 24,
        layer_types=tuple(data["layer_types"]),
        final_norm_params=data["final_norm_params"],
        all_named_params=data["all_named_params"],
    )

    with pytest.raises(ValueError, match="complete hybrid"):
        choose_prefix(report, 330_000_000, 350_000_000, requested_layers=3)


def test_text_prefix_translation_drops_vision_and_mtp():
    assert (
        translate_text_key("model.language_model.layers.3.linear_attn.A_log")
        == "model.layers.3.linear_attn.A_log"
    )
    assert translate_text_key("model.visual.patch_embed.proj.weight") is None
    assert translate_text_key("mtp.layers.0.mlp.down_proj.weight") is None


def test_text_config_is_standalone_and_keeps_live_attention_fields():
    full = {
        "model_type": "qwen3_5",
        "tie_word_embeddings": True,
        "vision_config": {"hidden_size": 512},
        "text_config": {
            "model_type": "qwen3_5_text",
            "hidden_size": 1024,
            "intermediate_size": 3584,
            "num_hidden_layers": 24,
            "layer_types": list(load_fixture()["layer_types"]),
            "linear_num_key_heads": 16,
            "linear_num_value_heads": 16,
            "linear_key_head_dim": 128,
            "linear_value_head_dim": 128,
            "linear_conv_kernel_dim": 4,
            "vocab_size": 248320,
            "mtp_num_hidden_layers": 1,
            "mtp_use_dedicated_embeddings": False,
        },
    }

    text = build_text_config(full, 4)

    assert text["model_type"] == "qwen3_5_text"
    assert text["num_hidden_layers"] == 4
    assert text["layer_types"] == load_fixture()["layer_types"][:4]
    assert text["linear_num_value_heads"] == 16
    assert text["tie_word_embeddings"] is True
    assert "vision_config" not in text
    assert not any(key.startswith("mtp_") for key in text)


def test_validator_rejects_vision_mtp_and_duplicate_tied_head():
    config = {
        "model_type": "qwen3_5_text",
        "tie_word_embeddings": True,
        "num_hidden_layers": 4,
        "layer_types": load_fixture()["layer_types"][:4],
    }
    keys = {
        "model.embed_tokens.weight": (248320, 1024),
        "model.layers.0.input_layernorm.weight": (1024,),
        "model.layers.1.input_layernorm.weight": (1024,),
        "model.layers.2.input_layernorm.weight": (1024,),
        "model.layers.3.input_layernorm.weight": (1024,),
        "model.norm.weight": (1024,),
        "lm_head.weight": (248320, 1024),
        "model.visual.patch_embed.proj.weight": (1, 1),
        "mtp.fc.weight": (1, 1),
    }

    with pytest.raises(ValueError, match="vision|MTP|tied"):
        validate_state_dict_keys(keys, config, 330_000_000, 350_000_000)


def test_metadata_counter_uses_tensor_shapes_without_loading_model():
    data = load_fixture()
    config = {
        "text_config": {
            "layer_types": list(data["layer_types"][:4]),
        }
    }
    tensors = {
        "model.language_model.embed_tokens.weight": (2, 3),
        "model.language_model.layers.0.input_layernorm.weight": (2,),
        "model.language_model.layers.1.input_layernorm.weight": (3,),
        "model.language_model.layers.2.input_layernorm.weight": (4,),
        "model.language_model.layers.3.input_layernorm.weight": (5,),
        "model.language_model.norm.weight": (2,),
        "model.visual.patch_embed.proj.weight": (7,),
        "mtp.fc.weight": (11,),
    }

    def make_tensor(shape):
        import torch

        return torch.zeros(shape, dtype=torch.float32)

    path = FIXTURE.parent / "counter-fixture.safetensors"
    try:
        save_file({key: make_tensor(shape) for key, shape in tensors.items()}, str(path))
        report = count_parameter_groups(path, config)
    finally:
        path.unlink(missing_ok=True)

    assert report.embedding_params == 6
    assert report.layer_params == (2, 3, 4, 5)
    assert report.final_norm_params == 2
    assert report.all_named_params == 40