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