| from __future__ import annotations | |
| from pathlib import Path | |
| import torch | |
| from adam.model_inspector import compare_models, inspect_model | |
| def test_inspector_reads_generic_torch_checkpoint(tmp_path: Path) -> None: | |
| checkpoint = tmp_path / "checkpoint.pt" | |
| torch.save( | |
| { | |
| "state_dict": { | |
| "encoder.weight": torch.tensor([[1.0, 2.0], [3.0, 4.0]]), | |
| "attention.query.bias": torch.zeros(2), | |
| }, | |
| "epoch": 3, | |
| }, | |
| checkpoint, | |
| ) | |
| summary = inspect_model(checkpoint) | |
| assert summary.architecture == "Generic / Unknown" | |
| assert summary.tensor_count == 2 | |
| assert summary.total_parameters == 6 | |
| assert summary.epoch == 3 | |
| assert summary.components["attention"] == 2 | |
| def test_inspector_detects_lora_and_adapter_details(tmp_path: Path) -> None: | |
| checkpoint = tmp_path / "adapter_model.safetensors" | |
| try: | |
| from safetensors.torch import save_file | |
| except Exception: | |
| return | |
| save_file( | |
| { | |
| "unet.block.lora_down.weight": torch.ones(4, 8), | |
| "unet.block.lora_up.weight": torch.ones(8, 4) * 0.5, | |
| }, | |
| str(checkpoint), | |
| ) | |
| summary = inspect_model(checkpoint, recorded_architecture="lora") | |
| assert summary.architecture == "LoRA" | |
| assert summary.lora["rank"] == "4" | |
| assert summary.trainable_parameters == 64 | |
| def test_compare_models_ranks_changed_tensors(tmp_path: Path) -> None: | |
| first = tmp_path / "a.pt" | |
| second = tmp_path / "b.pt" | |
| torch.save({"state_dict": {"layer.weight": torch.ones(4), "same.weight": torch.zeros(2)}}, first) | |
| torch.save({"state_dict": {"layer.weight": torch.ones(4) * 3, "same.weight": torch.zeros(2)}}, second) | |
| comparison = compare_models(first, second) | |
| assert comparison.architecture_match | |
| assert comparison.parameter_count_difference == 0 | |
| assert comparison.tensor_comparisons[0].name == "layer.weight" | |
| assert comparison.tensor_comparisons[0].change_score is not None | |
| def test_malformed_checkpoint_returns_warning(tmp_path: Path) -> None: | |
| checkpoint = tmp_path / "broken.pt" | |
| checkpoint.write_bytes(b"not a checkpoint") | |
| summary = inspect_model(checkpoint) | |
| assert summary.status == "warning" | |
| assert "no readable tensor checkpoint" in "\n".join(summary.health).casefold() | |