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