echo / MVP /validate_checkpoint.py
void0x14
feat: add qwen35 text checkpoint pruner
6466ca1 unverified
Raw
History Blame
3.08 kB
from __future__ import annotations
import json
import math
from dataclasses import asdict, dataclass
from pathlib import Path
from typing import Mapping, Sequence
from safetensors import safe_open
@dataclass(frozen=True)
class ValidationReport:
parameter_count: int
tensor_count: int
layer_count: int
visual_tensor_count: int
mtp_tensor_count: int
tied_lm_head_present: bool
valid: bool
def _numel(shape: Sequence[int]) -> int:
return math.prod(int(dimension) for dimension in shape)
def _layer_indices(keys: Sequence[str]) -> set[int]:
indices = set()
prefix = "model.layers."
for key in keys:
if key.startswith(prefix):
indices.add(int(key[len(prefix):].split(".", 1)[0]))
return indices
def validate_state_dict_keys(
shapes: Mapping[str, Sequence[int]],
config: Mapping[str, object],
minimum: int,
maximum: int,
) -> ValidationReport:
keys = tuple(shapes.keys())
visual_count = sum(key.startswith("model.visual.") for key in keys)
mtp_count = sum(key.startswith("mtp.") for key in keys)
tied_head = "lm_head.weight" in shapes and bool(config.get("tie_word_embeddings", False))
if visual_count or mtp_count or tied_head:
raise ValueError("vision/MTP tensors or a duplicate tied head are present")
if config.get("model_type") != "qwen3_5_text":
raise ValueError("standalone text config must use model_type qwen3_5_text")
layer_count = int(config["num_hidden_layers"])
expected_indices = set(range(layer_count))
actual_indices = _layer_indices(keys)
if actual_indices != expected_indices:
raise ValueError(f"layer indices are not contiguous: expected {expected_indices}, got {actual_indices}")
layer_types = tuple(config["layer_types"])
if len(layer_types) != layer_count:
raise ValueError("config layer_types length does not match num_hidden_layers")
parameter_count = sum(_numel(shape) for shape in shapes.values())
if not minimum <= parameter_count <= maximum:
raise ValueError(f"parameter count {parameter_count} is outside [{minimum}, {maximum}]")
return ValidationReport(
parameter_count=parameter_count,
tensor_count=len(keys),
layer_count=layer_count,
visual_tensor_count=visual_count,
mtp_tensor_count=mtp_count,
tied_lm_head_present=tied_head,
valid=True,
)
def validate_checkpoint(
config_path: str | Path,
weights_path: str | Path,
minimum: int,
maximum: int,
) -> ValidationReport:
config = json.loads(Path(config_path).read_text(encoding="utf-8"))
shapes = {}
with safe_open(str(weights_path), framework="pt", device="cpu") as handle:
for key in handle.keys():
shapes[key] = tuple(handle.get_slice(key).get_shape())
report = validate_state_dict_keys(shapes, config, minimum, maximum)
Path(config_path).with_name("validation.json").write_text(
json.dumps(asdict(report), indent=2, sort_keys=True) + "\n", encoding="utf-8"
)
return report