SyntheticMDProductions's picture
Update ADAM safety, UI, and model workflows (#1)
c61c435
Raw History Blame Contribute Delete
6.9 kB
from __future__ import annotations
import json
from pathlib import Path
from typing import Any
from .base import ModelComparison, TensorComparison, is_cancelled, report
from .detector import inspect_model
from .generic import _extract_state_dict, _weight_files
from .statistics import component_for_name
def _iter_named_tensors(path: Path):
files = _weight_files(path)
for file in files:
if file.suffix.casefold() == ".safetensors":
from safetensors import safe_open
with safe_open(str(file), framework="pt", device="cpu") as handle:
for key in handle.keys():
yield key, handle.get_tensor(key)
else:
import torch
payload = torch.load(str(file), map_location="cpu", weights_only=False)
state = _extract_state_dict(payload)
for key, value in state.items():
if hasattr(value, "shape"):
yield str(key), value
def _tensor_map(path: Path) -> dict[str, Any]:
return {name: tensor for name, tensor in _iter_named_tensors(path)}
def _config_differences(configs_a: dict[str, Any], configs_b: dict[str, Any], *, limit: int = 40) -> list[str]:
diffs: list[str] = []
keys = sorted(set(configs_a) | set(configs_b))
for key in keys:
if key not in configs_a:
diffs.append(f"Only B has config {key}")
elif key not in configs_b:
diffs.append(f"Only A has config {key}")
elif json.dumps(configs_a[key], sort_keys=True, default=str) != json.dumps(configs_b[key], sort_keys=True, default=str):
diffs.append(f"Config differs: {key}")
if len(diffs) >= limit:
diffs.append("Additional config differences omitted.")
break
return diffs
def compare_models(
path_a: str | Path,
path_b: str | Path,
*,
arch_a: str = "",
arch_b: str = "",
settings_a: dict[str, Any] | None = None,
settings_b: dict[str, Any] | None = None,
progress=None,
cancelled=None,
) -> ModelComparison:
report(progress, 2, "Inspecting first model")
summary_a = inspect_model(path_a, recorded_architecture=arch_a, run_settings=settings_a, progress=progress, cancelled=cancelled)
report(progress, 30, "Inspecting second model")
summary_b = inspect_model(path_b, recorded_architecture=arch_b, run_settings=settings_b, progress=progress, cancelled=cancelled)
report(progress, 55, "Loading comparable tensors")
tensors_a = _tensor_map(Path(summary_a.resolved_path))
if is_cancelled(cancelled):
raise RuntimeError("Comparison cancelled.")
tensors_b = _tensor_map(Path(summary_b.resolved_path))
names_a = set(tensors_a)
names_b = set(tensors_b)
common = sorted(names_a & names_b)
only_a = sorted(names_a - names_b)[:200]
only_b = sorted(names_b - names_a)[:200]
shape_mismatches = []
comparable = []
for name in common:
if tuple(tensors_a[name].shape) != tuple(tensors_b[name].shape):
shape_mismatches.append(name)
else:
comparable.append(name)
comparisons: list[TensorComparison] = []
import torch
for index, name in enumerate(comparable):
if is_cancelled(cancelled):
raise RuntimeError("Comparison cancelled.")
if index % 10 == 0:
report(progress, 58 + int(36 * index / max(1, len(comparable))), f"Comparing {index + 1} of {len(comparable)} tensors")
with torch.no_grad():
a = tensors_a[name].detach().to(device="cpu").float().reshape(-1)
b = tensors_b[name].detach().to(device="cpu").float().reshape(-1)
if a.numel() == 0:
continue
limit = 1_000_000
if a.numel() > limit:
stride = max(1, a.numel() // limit)
a = a[::stride][:limit]
b = b[::stride][:limit]
delta = b - a
mean_abs = delta.abs().mean().item()
base_abs = a.abs().mean().item()
relative = mean_abs / (base_abs + 1e-12)
l2 = torch.linalg.vector_norm(delta).item()
norm_a = torch.linalg.vector_norm(a).item()
norm_b = torch.linalg.vector_norm(b).item()
cosine = torch.nn.functional.cosine_similarity(a, b, dim=0).item() if norm_a and norm_b else None
drift = (norm_b - norm_a) / (norm_a + 1e-12) if norm_a else None
score = relative * 0.7 + (1 - cosine if cosine is not None else 0) * 0.3
comparisons.append(
TensorComparison(
name=name,
shape=tuple(int(dim) for dim in tensors_a[name].shape),
component=component_for_name(name),
mean_abs_difference=float(mean_abs),
relative_difference=float(relative),
cosine_similarity=float(cosine) if cosine is not None else None,
l2_distance=float(l2),
drift=float(drift) if drift is not None else None,
change_score=float(score),
)
)
comparisons.sort(key=lambda item: item.change_score or 0, reverse=True)
groups: dict[str, dict[str, float]] = {}
for item in comparisons:
group = groups.setdefault(item.component, {"tensors": 0, "mean_change_score": 0.0, "mean_abs_difference": 0.0})
group["tensors"] += 1
group["mean_change_score"] += item.change_score or 0
group["mean_abs_difference"] += item.mean_abs_difference or 0
for group in groups.values():
count = max(1, int(group["tensors"]))
group["mean_change_score"] /= count
group["mean_abs_difference"] /= count
messages = [
"Change Score is a statistical weight-change metric; it does not directly equal behavioral importance."
]
if shape_mismatches:
messages.append("Some tensor comparisons are unavailable because tensor shapes differ.")
report(progress, 100, "Comparison complete")
return ModelComparison(
path_a=summary_a.resolved_path,
path_b=summary_b.resolved_path,
architecture_a=summary_a.architecture,
architecture_b=summary_b.architecture,
architecture_match=summary_a.architecture == summary_b.architecture,
config_differences=_config_differences(summary_a.configs, summary_b.configs),
resolution_difference=(
(summary_a.resolution, summary_b.resolution)
if summary_a.resolution != summary_b.resolution else None
),
parameter_count_difference=summary_b.total_parameters - summary_a.total_parameters,
only_a=only_a,
only_b=only_b,
shape_mismatches=shape_mismatches[:200],
tensor_comparisons=comparisons[:200],
group_comparisons=groups,
messages=messages,
)