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