from __future__ import annotations import math from pathlib import Path from typing import Any from .base import CONFIG_FILENAMES, MODEL_EXTENSIONS, TensorStats, dtype_size, parameter_count def bytes_label(size: int | float | None) -> str: if size is None: return "-" value = float(size) for unit in ("B", "KB", "MB", "GB", "TB"): if abs(value) < 1024 or unit == "TB": return f"{value:.1f} {unit}" if unit != "B" else f"{int(value)} B" value /= 1024 return f"{value:.1f} TB" def component_for_name(name: str) -> str: lowered = name.casefold() mapping = ( ("down_blocks", ("down_blocks", "down.", "downsample")), ("mid_block", ("mid_block", "middle_block", "mid.")), ("up_blocks", ("up_blocks", "up.", "upsample")), ("attention", ("attn", "attention", "to_q", "to_k", "to_v", "query", "key", "value")), ("embeddings", ("embed", "embedding", "position", "token")), ("transformer blocks", ("transformer", "blocks.", "layers.", "encoder", "decoder")), ("output layers", ("out.", "output", "proj_out", "lm_head", "conv_out")), ("LoRA adapters", ("lora", "hada", "lokr", "adapter")), ("normalization", ("norm", "bn", "ln", "group_norm", "layer_norm")), ) for component, tokens in mapping: if any(token in lowered for token in tokens): return component return name.split(".", 1)[0] if "." in name else "other" def safe_number(value: Any) -> float | None: try: number = float(value) except (TypeError, ValueError, OverflowError): return None return number if math.isfinite(number) else None def tensor_stats_from_torch(name: str, tensor: Any, *, sample_limit: int = 1_000_000) -> TensorStats: shape = tuple(int(dim) for dim in getattr(tensor, "shape", ())) dtype = str(getattr(tensor, "dtype", "unknown")).replace("torch.", "") count = parameter_count(shape) stat = TensorStats( name=name, shape=shape, dtype=dtype, parameter_count=count, memory_bytes=count * dtype_size(dtype), component=component_for_name(name), ) if count == 0: stat.health.append("Empty tensor") return stat try: import torch with torch.no_grad(): values = tensor.detach().to(device="cpu") if not values.is_floating_point() and not values.is_complex(): values = values.float() else: values = values.float() flat = values.reshape(-1) if flat.numel() > sample_limit: stride = max(1, flat.numel() // sample_limit) flat = flat[::stride][:sample_limit] finite = torch.isfinite(flat) if not bool(finite.all()): if bool(torch.isnan(flat).any()): stat.health.append("Invalid: NaN values found") if bool(torch.isinf(flat).any()): stat.health.append("Invalid: Inf values found") flat = flat[finite] if flat.numel() == 0: return stat stat.minimum = safe_number(flat.min().item()) stat.maximum = safe_number(flat.max().item()) stat.mean = safe_number(flat.mean().item()) stat.std = safe_number(flat.std(unbiased=False).item()) if flat.numel() > 1 else 0.0 stat.abs_mean = safe_number(flat.abs().mean().item()) stat.l2_norm = safe_number(torch.linalg.vector_norm(flat).item()) stat.zero_percent = safe_number((flat == 0).float().mean().item() * 100) except Exception as exc: stat.health.append(f"Statistics unavailable: {exc}") if stat.abs_mean is not None and stat.abs_mean > 100: stat.health.append("Unusual: very large average weight magnitude") if stat.maximum is not None and stat.minimum is not None and max(abs(stat.maximum), abs(stat.minimum)) > 1_000: stat.health.append("Unusual: very large absolute weight value") return stat def discover_config_files(path: Path) -> list[Path]: root = path if path.is_dir() else path.parent files: list[Path] = [] try: for item in root.rglob("*"): if item.is_file() and item.name in CONFIG_FILENAMES: files.append(item) except OSError: return [] return sorted(files) def discover_checkpoint_paths(path: Path) -> list[Path]: root = path if path.is_dir() else path.parent candidates: list[Path] = [] try: for item in root.rglob("*"): if item.is_file() and item.suffix.casefold() in MODEL_EXTENSIONS: candidates.append(item) elif item.is_dir() and item.name.startswith("checkpoint-"): candidates.append(item) except OSError: return [] return sorted(candidates, key=lambda item: (step_from_name(item.name) or -1, str(item))) def step_from_name(name: str) -> int | None: import re matches = re.findall(r"(?:step|checkpoint|epoch|e|s)[-_]?(\d+)", name, flags=re.I) if not matches: matches = re.findall(r"(\d+)", name) if not matches: return None try: return int(matches[-1]) except ValueError: return None def shape_label(shape: tuple[int, ...]) -> str: return " x ".join(str(dim) for dim in shape) if shape else "scalar"