SyntheticMDProductions's picture
Update ADAM safety, UI, and model workflows (#1)
c61c435
Raw History Blame Contribute Delete
3.88 kB
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