File size: 14,755 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
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
from __future__ import annotations

import json
from collections import Counter
from pathlib import Path
from typing import Any, Iterator

from .base import BaseModelInspector, InspectorError, ModelInspection, TensorStats, is_cancelled, report
from .statistics import (
    discover_checkpoint_paths,
    discover_config_files,
    step_from_name,
    tensor_stats_from_torch,
)


def _folder_size(path: Path) -> int:
    if path.is_file():
        return path.stat().st_size
    total = 0
    try:
        for item in path.rglob("*"):
            if item.is_file():
                total += item.stat().st_size
    except OSError:
        return total
    return total


def _read_configs(files: list[Path], root: Path) -> dict[str, Any]:
    configs: dict[str, Any] = {}
    for file in files[:40]:
        try:
            key = str(file.relative_to(root if root.is_dir() else root.parent))
        except ValueError:
            key = file.name
        try:
            configs[key] = json.loads(file.read_text(encoding="utf-8"))
        except (OSError, UnicodeDecodeError, json.JSONDecodeError):
            configs[key] = "<unreadable>"
    return configs


def _resolution_from_configs(configs: dict[str, Any], settings: dict[str, Any] | None) -> int | None:
    for source in (settings or {}, *[value for value in configs.values() if isinstance(value, dict)]):
        for key in ("resolution", "sample_size", "image_size", "size"):
            value = source.get(key) if isinstance(source, dict) else None
            if isinstance(value, int):
                return value
            if isinstance(value, (list, tuple)) and value and isinstance(value[0], int):
                return int(value[0])
            try:
                if value:
                    return int(value)
            except (TypeError, ValueError):
                pass
    return None


def _iter_safetensors(file: Path) -> Iterator[tuple[str, Any, dict[str, Any]]]:
    from safetensors import safe_open

    with safe_open(str(file), framework="pt", device="cpu") as handle:
        metadata = handle.metadata() or {}
        for key in handle.keys():
            yield key, handle.get_tensor(key), metadata


def _extract_state_dict(payload: Any) -> dict[str, Any]:
    try:
        import torch
    except Exception:
        torch = None
    if torch is not None and hasattr(payload, "shape"):
        return {"tensor": payload}
    if isinstance(payload, dict):
        for key in ("state_dict", "model_state_dict", "model", "module", "unet", "network"):
            value = payload.get(key)
            if isinstance(value, dict) and any(hasattr(item, "shape") for item in value.values()):
                return value
        if any(hasattr(item, "shape") for item in payload.values()):
            return payload
    return {}


def _iter_torch_checkpoint(file: Path) -> Iterator[tuple[str, Any, dict[str, Any]]]:
    import torch

    try:
        payload = torch.load(str(file), map_location="cpu", weights_only=True)
    except TypeError:
        payload = torch.load(str(file), map_location="cpu")
    except Exception:
        payload = torch.load(str(file), map_location="cpu", weights_only=False)
    state = _extract_state_dict(payload)
    metadata = {key: value for key, value in payload.items() if key not in state} if isinstance(payload, dict) else {}
    for key, tensor in state.items():
        if hasattr(tensor, "shape"):
            yield str(key), tensor, metadata


def _weight_files(path: Path) -> list[Path]:
    if path.is_file():
        return [path]
    ignored_names = {"optimizer.bin", "scheduler.bin", "scaler.pt"}
    names = {
        "diffusion_pytorch_model.safetensors",
        "model.safetensors",
        "pytorch_model.bin",
        "adapter_model.safetensors",
        "adapter_model.bin",
        "checkpoint.pt",
        "best_checkpoint.pt",
    }
    files: list[Path] = []
    try:
        for item in path.rglob("*"):
            if item.is_file() and (item.name in names or item.suffix.casefold() in {".safetensors", ".pt", ".pth", ".bin", ".ckpt"}):
                if item.name.casefold() not in ignored_names:
                    files.append(item)
    except OSError:
        return []
    if (path / "model_index.json").is_file():
        final_files = [
            item for item in files
            if not any(part.startswith("checkpoint-") for part in item.relative_to(path).parts)
        ]
        if final_files:
            files = final_files
    return sorted(files, key=lambda item: (0 if item.name in names else 1, str(item)))


class GenericModelInspector(BaseModelInspector):
    architecture = "Generic / Unknown"

    def inspect(
        self,
        path: str | Path,
        *,
        recorded_architecture: str = "",
        run_settings: dict[str, Any] | None = None,
        progress=None,
        cancelled=None,
    ) -> ModelInspection:
        target = Path(path).expanduser()
        if not target.exists():
            raise InspectorError(f"Model path does not exist: {target}")
        target = target.resolve()
        report(progress, 3, "Finding model files")
        config_files = discover_config_files(target)
        configs = _read_configs(config_files, target)
        files = _weight_files(target)
        if not files:
            message = "Model contains no readable tensor checkpoint"
            return self._empty(target, recorded_architecture, run_settings, config_files, configs, message)

        tensors: list[TensorStats] = []
        dtypes: Counter[str] = Counter()
        components: Counter[str] = Counter()
        health: list[str] = []
        messages: list[str] = []
        metadata: dict[str, Any] = {}
        for file_index, file in enumerate(files):
            if is_cancelled(cancelled):
                raise InspectorError("Inspection cancelled.")
            report(progress, 8 + int(80 * file_index / max(1, len(files))), f"Reading {file.name}")
            try:
                if file.suffix.casefold() == ".safetensors":
                    iterator = _iter_safetensors(file)
                else:
                    iterator = _iter_torch_checkpoint(file)
                for name, tensor, file_metadata in iterator:
                    if is_cancelled(cancelled):
                        raise InspectorError("Inspection cancelled.")
                    prefix = file.parent.name if len(files) > 1 else ""
                    stat = tensor_stats_from_torch(f"{prefix}.{name}" if prefix and not name.startswith(prefix) else name, tensor)
                    tensors.append(stat)
                    dtypes[stat.dtype] += stat.parameter_count
                    components[stat.component] += stat.parameter_count
                    health.extend(f"{stat.name}: {item}" for item in stat.health)
                    if file_metadata:
                        metadata.update(file_metadata)
            except Exception as exc:
                health.append(f"{file.name}: unreadable checkpoint ({exc})")

        if not tensors:
            message = "Model contains no readable tensor checkpoint"
            return self._empty(target, recorded_architecture, run_settings, config_files, configs, message, [*(health or []), message])

        report(progress, 92, "Summarizing model")
        total_parameters = sum(tensor.parameter_count for tensor in tensors)
        parameter_memory = sum(tensor.memory_bytes for tensor in tensors)
        largest = sorted(tensors, key=lambda item: item.parameter_count, reverse=True)[:20]
        architecture, confidence, message = self._architecture_from_signals(
            target, recorded_architecture, configs, [tensor.name for tensor in tensors]
        )
        messages.append(message)
        duplicate_count = len(tensors) - len({tensor.name for tensor in tensors})
        if duplicate_count:
            health.append(f"Unusual: {duplicate_count} duplicate tensor names after folder merging")
        checkpoints = [str(item) for item in discover_checkpoint_paths(target)]
        return ModelInspection(
            path=str(path),
            resolved_path=str(target),
            architecture=architecture,
            confidence=confidence,
            status="ok",
            size_bytes=_folder_size(target),
            config_files=[str(item) for item in config_files],
            resolution=_resolution_from_configs(configs, run_settings),
            epoch=self._number_from_metadata(metadata, "epoch"),
            step=self._number_from_metadata(metadata, "step") or step_from_name(target.name),
            tensor_count=len(tensors),
            total_parameters=total_parameters,
            trainable_parameters=self._trainable_parameters(tensors, architecture),
            parameter_memory_bytes=parameter_memory,
            dtypes=dict(dtypes),
            components=dict(components),
            largest_tensors=largest,
            tensors=tensors,
            health=health or ["No invalid tensor values found in sampled statistics."],
            messages=messages,
            lora=self._lora_info(tensors, configs),
            configs=configs,
            histogram=self._histogram(tensors),
            tensor_size_distribution=[(tensor.name, tensor.parameter_count) for tensor in largest],
            checkpoints=checkpoints,
            loss_history=[],
        )

    def _empty(
        self,
        target: Path,
        recorded_architecture: str,
        run_settings: dict[str, Any] | None,
        config_files: list[Path],
        configs: dict[str, Any],
        message: str,
        health: list[str] | None = None,
    ) -> ModelInspection:
        architecture, confidence, detection_message = self._architecture_from_signals(target, recorded_architecture, configs, [])
        return ModelInspection(
            path=str(target),
            resolved_path=str(target),
            architecture=architecture,
            confidence=confidence,
            status="warning",
            size_bytes=_folder_size(target),
            config_files=[str(item) for item in config_files],
            resolution=_resolution_from_configs(configs, run_settings),
            epoch=None,
            step=step_from_name(target.name),
            tensor_count=0,
            total_parameters=0,
            trainable_parameters=None,
            parameter_memory_bytes=0,
            dtypes={},
            components={},
            largest_tensors=[],
            tensors=[],
            health=health or [message],
            messages=[detection_message, message],
            configs=configs,
            checkpoints=[str(item) for item in discover_checkpoint_paths(target)],
        )

    @staticmethod
    def _number_from_metadata(metadata: dict[str, Any], key: str) -> int | None:
        for candidate in (key, f"global_{key}", f"current_{key}"):
            try:
                value = metadata.get(candidate)
                if value is not None:
                    return int(value)
            except (TypeError, ValueError):
                pass
        return None

    @staticmethod
    def _trainable_parameters(tensors: list[TensorStats], architecture: str) -> int | None:
        if architecture == "LoRA":
            return sum(tensor.parameter_count for tensor in tensors)
        return None

    @staticmethod
    def _architecture_from_signals(
        target: Path,
        recorded_architecture: str,
        configs: dict[str, Any],
        tensor_names: list[str],
    ) -> tuple[str, float, str]:
        recorded = recorded_architecture.casefold()
        joined_names = "\n".join(tensor_names).casefold()
        config_text = json.dumps(configs, default=str).casefold()
        folder_text = str(target).casefold()
        signals = " ".join((joined_names, config_text, folder_text))
        if "lora" in recorded or "lora" in signals or "adapter_config" in signals:
            return "LoRA", 0.92, "Model recognized as LoRA"
        if "maskgit" in recorded or "maskgit" in signals:
            return "MaskGIT", 0.86, "Model recognized as MaskGIT"
        if "flow" in recorded or "rectified_flow" in signals or "flow_model_info" in signals:
            return "Flow Matching", 0.9, "Model recognized as Flow Matching"
        if "ddpm" in recorded or "diffusers" in config_text or "unet" in signals or "scheduler_config" in signals:
            return "DDPM / Diffusers", 0.88, "Model recognized as DDPM"
        return "Generic / Unknown", 0.35, "Model type uncertain - using generic tensor inspection"

    @staticmethod
    def _lora_info(tensors: list[TensorStats], configs: dict[str, Any]) -> dict[str, Any]:
        lora_tensors = [tensor for tensor in tensors if "lora" in tensor.name.casefold()]
        if not lora_tensors:
            return {}
        down = [tensor for tensor in lora_tensors if any(token in tensor.name.casefold() for token in ("down", "lora_a"))]
        up = [tensor for tensor in lora_tensors if any(token in tensor.name.casefold() for token in ("up", "lora_b"))]
        ranks = sorted({tensor.shape[0] for tensor in down if tensor.shape})
        alpha = None
        targets: set[str] = set()
        for config in configs.values():
            if isinstance(config, dict):
                alpha = config.get("lora_alpha", config.get("alpha", alpha))
                modules = config.get("target_modules")
                if isinstance(modules, list):
                    targets.update(str(item) for item in modules)
        if not targets:
            for tensor in lora_tensors:
                parts = tensor.name.split(".")
                if len(parts) > 2:
                    targets.add(parts[-3])
        return {
            "rank": ", ".join(str(item) for item in ranks[:8]) if ranks else "unknown",
            "alpha": alpha if alpha is not None else "unknown",
            "target_modules": sorted(targets)[:20],
            "down_matrices": len(down),
            "up_matrices": len(up),
            "adapter_parameter_count": sum(tensor.parameter_count for tensor in lora_tensors),
            "average_abs_mean": (
                sum(tensor.abs_mean or 0 for tensor in lora_tensors) / max(1, len(lora_tensors))
            ),
        }

    @staticmethod
    def _histogram(tensors: list[TensorStats]) -> dict[str, list[float]]:
        values = [tensor.abs_mean for tensor in tensors if tensor.abs_mean is not None]
        if not values:
            return {}
        buckets = [0.0] * 10
        high = max(values) or 1.0
        for value in values:
            index = min(9, int((value / high) * 10))
            buckets[index] += 1
        return {"abs_mean_bins": [round(high * index / 10, 6) for index in range(11)], "counts": buckets}