File size: 4,665 Bytes
e0265b9
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
f8c73f9
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
e0265b9
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
from __future__ import annotations

from pathlib import Path
import json

import pytest

from adam.assets import AssetRegistry
from adam.commands import CommandValidationError, TrainingCommand


class FakeConfig:
    def __init__(self, values: dict) -> None:
        self.values = values

    def get(self, key: str, default=None):
        return self.values.get(key, default)


def test_asset_registry_prefers_exact_friendly_name(tmp_path: Path) -> None:
    exact = tmp_path / "Mario"
    similar = tmp_path / "Mario 2"
    exact.mkdir()
    similar.mkdir()
    registry = AssetRegistry(tmp_path)
    registry.register(kind="dataset", name="Mario", path=str(exact))
    registry.register(kind="dataset", name="Mario 2", path=str(similar))

    assert [item.name for item in registry.find("dataset", "Mario")] == ["Mario"]


def test_asset_discovery_removes_models_whose_paths_were_deleted(tmp_path: Path) -> None:
    model = tmp_path / "LoRA output" / "Crystal_Biter"
    model.mkdir(parents=True)
    registry = AssetRegistry(tmp_path)
    registry.register(kind="model", name="Crystal_Biter", path=str(model), trainer="lora")

    model.rmdir()
    registry.discover(FakeConfig({"tool_folders": {}}))

    assert registry.find("model", "Crystal_Biter", trainer="lora") == []


def test_lora_discovery_registers_weights_in_nested_trainer_output(tmp_path: Path) -> None:
    trainer = tmp_path / "LoRATrainer"
    weight = trainer / "output" / "Named run" / "adapter" / "My_LoRA.safetensors"
    weight.parent.mkdir(parents=True)
    weight.write_bytes(b"weights")
    (weight.parent.parent / "model_info.json").write_text(
        '{"trigger_word": "my_lora"}', encoding="utf-8"
    )

    registry = AssetRegistry(tmp_path)
    registry.discover(FakeConfig({"tool_folders": {"lora_trainer": str(trainer)}}))

    model = registry.find("model", "My LoRA", trainer="lora")[0]
    assert Path(model.path) == weight.resolve()
    assert model.metadata == {"trigger_word": "my_lora"}


def test_lora_discovery_excludes_intermediate_epoch_checkpoints(tmp_path: Path) -> None:
    trainer = tmp_path / "LoRATrainer"
    output = trainer / "output" / "Named run"
    output.mkdir(parents=True)
    (output / "My_LoRA.safetensors").write_bytes(b"final")
    (output / "My_LoRA_epoch_0050.safetensors").write_bytes(b"checkpoint")
    (output / "checkpoint-e50_s100.safetensors").write_bytes(b"checkpoint")

    registry = AssetRegistry(tmp_path)
    # Simulate an index written by an older ADAM release.
    registry.register(
        kind="model", name="My_LoRA_epoch_0050",
        path=str(output / "My_LoRA_epoch_0050.safetensors"), trainer="lora",
    )
    registry.discover(FakeConfig({"tool_folders": {"lora_trainer": str(trainer)}}))

    assert [item.name for item in registry.assets if item.trainer == "lora"] == ["My_LoRA"]


def test_training_command_rejects_uncontrolled_fields() -> None:
    with pytest.raises(CommandValidationError, match="Unsupported command fields"):
        TrainingCommand.from_dict(
            {
                "action": "train",
                "trainer": "ddpm",
                "dataset": "dataset",
                "model_name": "model",
                "epochs": 10,
                "shell_command": "unsafe",
            }
        )


def test_training_command_accepts_only_safe_ddpm_options() -> None:
    command = TrainingCommand.from_dict({
        "action": "train", "trainer": "ddpm", "dataset": "dataset", "model_name": "model", "epochs": 10,
        "training_options": {"resolution": 256, "batch_size": 2, "learning_rate": 0.0001},
    })
    assert command.training_options["resolution"] == 256


def test_flow_discovery_recovers_dataset_from_adam_job_history(tmp_path: Path) -> None:
    flow_root = tmp_path / "Flow"
    dataset = tmp_path / "Dataset"
    model = flow_root / "output_flow_models" / "Model"
    dataset.mkdir()
    (model / "unet").mkdir(parents=True)
    (model / "unet" / "config.json").write_text("{}", encoding="utf-8")
    (model / "flow_model_info.json").write_text(
        '{"model_type":"rectified_flow","name":"Friendly Flow","resolution":128}',
        encoding="utf-8",
    )
    (tmp_path / "data").mkdir()
    (tmp_path / "data" / "jobs.json").write_text(json.dumps({"jobs": [{"plan": {"steps": [{
        "tool_id": "flow_trainer", "arguments": {
            "output_dir": str(model), "dataset_dir": str(dataset),
        },
    }]}}]}), encoding="utf-8")
    registry = AssetRegistry(tmp_path)

    registry.discover(FakeConfig({"tool_folders": {"flow_trainer": str(flow_root)}}))

    flow = registry.find("model", "Friendly Flow", trainer="flow")[0]
    assert flow.dataset_id