File size: 2,342 Bytes
c61c435
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
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()