echo / MVP /tests /test_pruning_contract.py
void0x14
fix: remove mtp fields from text config
5d84e23 unverified
Raw History Blame Contribute Delete
5.39 kB
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