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, )