from __future__ import annotations import argparse import copy import json import math from dataclasses import asdict, dataclass from pathlib import Path from typing import Iterable from safetensors import safe_open from safetensors.torch import save_file HYBRID_BLOCK = ("linear_attention", "linear_attention", "linear_attention", "full_attention") DEFAULT_MINIMUM = 330_000_000 DEFAULT_MAXIMUM = 350_000_000 @dataclass(frozen=True) class ParameterReport: embedding_params: int layer_params: tuple[int, ...] layer_types: tuple[str, ...] final_norm_params: int all_named_params: int @property def text_backbone_params(self) -> int: return self.embedding_params + sum(self.layer_params) + self.final_norm_params @dataclass(frozen=True) class PrefixChoice: layer_count: int parameter_count: int def load_live_config(path: str | Path) -> dict: with Path(path).open(encoding="utf-8") as handle: return json.load(handle) def _numel(shape: Iterable[int]) -> int: return math.prod(int(dimension) for dimension in shape) def _layer_index(key: str) -> int | None: prefix = "model.language_model.layers." if not key.startswith(prefix): return None remainder = key[len(prefix):] index_text = remainder.split(".", 1)[0] return int(index_text) def count_parameter_groups(safetensors_path: str | Path, config: dict) -> ParameterReport: text_config = config.get("text_config", config) layer_types = tuple(text_config["layer_types"]) layer_params = [0 for _ in layer_types] embedding_params = 0 final_norm_params = 0 all_named_params = 0 with safe_open(str(safetensors_path), framework="pt", device="cpu") as handle: for key in handle.keys(): shape = handle.get_slice(key).get_shape() params = _numel(shape) all_named_params += params if key == "model.language_model.embed_tokens.weight": embedding_params += params elif key == "model.language_model.norm.weight": final_norm_params += params else: index = _layer_index(key) if index is not None: if index >= len(layer_params): raise ValueError(f"tensor layer index {index} exceeds config layer count") layer_params[index] += params return ParameterReport( embedding_params=embedding_params, layer_params=tuple(layer_params), layer_types=layer_types, final_norm_params=final_norm_params, all_named_params=all_named_params, ) def _is_complete_prefix(layer_types: tuple[str, ...], layer_count: int) -> bool: if layer_count == 0 or layer_count % len(HYBRID_BLOCK) != 0: return False return layer_types[:layer_count] == HYBRID_BLOCK * (layer_count // len(HYBRID_BLOCK)) def choose_prefix( report: ParameterReport, minimum: int, maximum: int, requested_layers: int | None = None, ) -> PrefixChoice: if minimum > maximum: raise ValueError("minimum parameter bound exceeds maximum") limit = min(len(report.layer_params), len(report.layer_types)) candidates = [requested_layers] if requested_layers is not None else range( len(HYBRID_BLOCK), limit + 1, len(HYBRID_BLOCK) ) choices: list[PrefixChoice] = [] for layer_count in candidates: if layer_count is None or layer_count > limit: continue if not _is_complete_prefix(report.layer_types, layer_count): if requested_layers is not None: raise ValueError("requested layer count is not a complete hybrid block") continue parameter_count = ( report.embedding_params + report.final_norm_params + sum(report.layer_params[:layer_count]) ) if minimum <= parameter_count <= maximum: choices.append(PrefixChoice(layer_count, parameter_count)) if not choices: raise ValueError("no complete hybrid prefix fits the parameter interval") return max(choices, key=lambda choice: choice.layer_count) def translate_text_key(key: str) -> str | None: if key.startswith("model.visual.") or key.startswith("mtp."): return None prefix = "model.language_model." if key.startswith(prefix): return "model." + key[len(prefix):] if key == "lm_head.weight": return None return None def build_text_config(full_config: dict, layer_count: int) -> dict: source = copy.deepcopy(full_config.get("text_config", full_config)) layer_types = list(source.get("layer_types", [])) if layer_count <= 0 or layer_count > len(layer_types): raise ValueError("layer count is outside the live text config") if not _is_complete_prefix(tuple(layer_types), layer_count): raise ValueError("layer count is not a complete hybrid block") source["model_type"] = "qwen3_5_text" source["num_hidden_layers"] = layer_count source["layer_types"] = layer_types[:layer_count] source["tie_word_embeddings"] = bool( full_config.get("tie_word_embeddings", source.get("tie_word_embeddings", False)) ) source["architectures"] = ["Qwen3_5ForCausalLM"] source.pop("vision_config", None) source.pop("mtp_config", None) source.pop("mtp_num_hidden_layers", None) source.pop("mtp_use_dedicated_embeddings", None) return source def _keep_source_key(key: str, layer_count: int) -> bool: if key == "model.language_model.embed_tokens.weight": return True if key == "model.language_model.norm.weight": return True index = _layer_index(key) return index is not None and index < layer_count def prune_checkpoint( input_weights: str | Path, input_config: str | Path, output_dir: str | Path, layer_count: int, ) -> dict: full_config = load_live_config(input_config) text_config = build_text_config(full_config, layer_count) output_path = Path(output_dir) output_path.mkdir(parents=True, exist_ok=True) tensors = {} with safe_open(str(input_weights), framework="pt", device="cpu") as handle: for key in handle.keys(): if not _keep_source_key(key, layer_count): continue output_key = translate_text_key(key) if output_key is None: raise ValueError(f"source tensor cannot be translated: {key}") tensors[output_key] = handle.get_tensor(key) output_weights = output_path / "model.safetensors" save_file(tensors, str(output_weights), metadata={"format": "pt"}) output_config = output_path / "config.json" output_config.write_text(json.dumps(text_config, indent=2, sort_keys=True) + "\n", encoding="utf-8") from MVP.validate_checkpoint import validate_checkpoint report = validate_checkpoint( output_config, output_weights, DEFAULT_MINIMUM, DEFAULT_MAXIMUM, ) return asdict(report) def _main() -> None: parser = argparse.ArgumentParser(description="Measure and prune Qwen3.5 text backbone tensors") subparsers = parser.add_subparsers(dest="command", required=True) measure = subparsers.add_parser("measure") measure.add_argument("--weights", required=True) measure.add_argument("--config", required=True) prune = subparsers.add_parser("prune") prune.add_argument("--weights", required=True) prune.add_argument("--config", required=True) prune.add_argument("--output", required=True) prune.add_argument("--layers", type=int, required=True) validate = subparsers.add_parser("validate") validate.add_argument("--weights", required=True) validate.add_argument("--config", required=True) validate.add_argument("--minimum", type=int, default=DEFAULT_MINIMUM) validate.add_argument("--maximum", type=int, default=DEFAULT_MAXIMUM) args = parser.parse_args() if args.command == "measure": report = count_parameter_groups(args.weights, load_live_config(args.config)) print(json.dumps(asdict(report) | {"text_backbone_params": report.text_backbone_params}, indent=2)) elif args.command == "prune": print(json.dumps(prune_checkpoint(args.weights, args.config, args.output, args.layers), indent=2)) else: from MVP.validate_checkpoint import validate_checkpoint print(json.dumps(asdict(validate_checkpoint(args.config, args.weights, args.minimum, args.maximum)), indent=2)) if __name__ == "__main__": _main()