File size: 3,081 Bytes
6466ca1
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
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