File size: 3,876 Bytes
c61c435
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
from __future__ import annotations

from dataclasses import dataclass, field
from pathlib import Path
from typing import Any, Callable, Iterable


MODEL_EXTENSIONS = {".safetensors", ".pt", ".pth", ".bin", ".ckpt"}
CONFIG_FILENAMES = {
    "model_index.json",
    "config.json",
    "scheduler_config.json",
    "flow_model_info.json",
    "adapter_config.json",
}


@dataclass(slots=True)
class TensorStats:
    name: str
    shape: tuple[int, ...]
    dtype: str
    parameter_count: int
    memory_bytes: int
    minimum: float | None = None
    maximum: float | None = None
    mean: float | None = None
    std: float | None = None
    abs_mean: float | None = None
    l2_norm: float | None = None
    zero_percent: float | None = None
    component: str = "other"
    health: list[str] = field(default_factory=list)


@dataclass(slots=True)
class ModelInspection:
    path: str
    resolved_path: str
    architecture: str
    confidence: float
    status: str
    size_bytes: int
    config_files: list[str]
    resolution: int | None
    epoch: int | None
    step: int | None
    tensor_count: int
    total_parameters: int
    trainable_parameters: int | None
    parameter_memory_bytes: int
    dtypes: dict[str, int]
    components: dict[str, int]
    largest_tensors: list[TensorStats]
    tensors: list[TensorStats]
    health: list[str]
    messages: list[str]
    lora: dict[str, Any] = field(default_factory=dict)
    configs: dict[str, Any] = field(default_factory=dict)
    histogram: dict[str, list[float]] = field(default_factory=dict)
    tensor_size_distribution: list[tuple[str, int]] = field(default_factory=list)
    checkpoints: list[str] = field(default_factory=list)
    loss_history: list[tuple[int, float]] = field(default_factory=list)


@dataclass(slots=True)
class TensorComparison:
    name: str
    shape: tuple[int, ...]
    component: str
    mean_abs_difference: float | None
    relative_difference: float | None
    cosine_similarity: float | None
    l2_distance: float | None
    drift: float | None
    change_score: float | None


@dataclass(slots=True)
class ModelComparison:
    path_a: str
    path_b: str
    architecture_a: str
    architecture_b: str
    architecture_match: bool
    config_differences: list[str]
    resolution_difference: tuple[int | None, int | None] | None
    parameter_count_difference: int
    only_a: list[str]
    only_b: list[str]
    shape_mismatches: list[str]
    tensor_comparisons: list[TensorComparison]
    group_comparisons: dict[str, dict[str, float]]
    messages: list[str]


ProgressCallback = Callable[[int, str], None]
CancelCallback = Callable[[], bool]


class InspectorError(RuntimeError):
    pass


class BaseModelInspector:
    architecture = "Generic / Unknown"

    def inspect(
        self,
        path: str | Path,
        *,
        recorded_architecture: str = "",
        run_settings: dict[str, Any] | None = None,
        progress: ProgressCallback | None = None,
        cancelled: CancelCallback | None = None,
    ) -> ModelInspection:
        raise NotImplementedError


def report(progress: ProgressCallback | None, value: int, message: str) -> None:
    if progress:
        progress(max(0, min(100, int(value))), message)


def is_cancelled(cancelled: CancelCallback | None) -> bool:
    return bool(cancelled and cancelled())


def parameter_count(shape: Iterable[int]) -> int:
    total = 1
    for dim in shape:
        total *= int(dim)
    return int(total)


def dtype_size(dtype: str) -> int:
    lowered = dtype.casefold()
    if "float64" in lowered or "int64" in lowered:
        return 8
    if "float32" in lowered or "int32" in lowered:
        return 4
    if "float16" in lowered or "bfloat16" in lowered or "int16" in lowered:
        return 2
    if "int8" in lowered or "uint8" in lowered or "bool" in lowered:
        return 1
    return 4