AI_Development_Automation_Manager / tests /test_model_inspector.py
SyntheticMDProductions's picture
Update ADAM safety, UI, and model workflows (#1)
c61c435
Raw
History Blame Contribute Delete
2.34 kB
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()