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