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