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