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] = "" 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}