File size: 6,335 Bytes
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
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
from __future__ import annotations

from pathlib import Path
from types import SimpleNamespace

from adam.models import ExecutionPlan, PlanStep
from adam.training_assistant import (
    append_preflight_summary,
    build_fine_tune_request,
    build_training_request,
    combine_training_plans,
    completion_recommendation,
    presets_from_config,
    parse_model_batch_names,
    build_dataset_collection_request,
    suggest_existing_dataset,
)


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

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


def test_wizard_builds_specific_new_ddpm_request() -> None:
    request = build_training_request(
        trainer="ddpm",
        subject="Luigi",
        dataset_name="",
        create_dataset=True,
        epochs=120,
        image_count=75,
        model_name="Luigi V2",
    )
    assert "dataset of Luigi" in request
    assert "75 images" in request
    assert "Luigi V2" in request
    assert "120 epochs" in request


def test_wizard_embeds_validated_training_options() -> None:
    request = build_training_request(
        trainer="ddpm", subject="Luigi", dataset_name="", create_dataset=True,
        epochs=10, image_count=20, model_name="Luigi", training_options={"resolution": 256, "batch_size": 2},
    )
    assert 'ADAM_TRAINING_OPTIONS:{"batch_size": 2, "resolution": 256}' in request


def test_fine_tune_request_targets_an_existing_model() -> None:
    request = build_fine_tune_request(model_name="Luigi V2", trainer="ddpm", epochs=20)

    assert request.startswith("Fine-tune Luigi V2 for 20 epochs with DDPM.")
    assert '"dataset_mode": "original"' in request


def test_fine_tune_request_can_select_a_new_dataset_and_settings() -> None:
    request = build_fine_tune_request(
        model_name="Luigi V2", trainer="ddpm", epochs=20,
        dataset_mode="new", new_subject="Luigi artwork", image_count=80,
        training_options={"resolution": 256},
    )

    assert '"new_subject": "Luigi artwork"' in request
    assert '"image_count": 80' in request
    assert '"resolution": 256' in request


def test_wizard_can_request_every_available_image() -> None:
    request = build_training_request(
        trainer="ddpm",
        subject="Luigi",
        dataset_name="",
        create_dataset=True,
        epochs=120,
        image_count=75,
        model_name="Luigi V2",
        collection_mode="all_available",
    )

    assert "as many available images" in request


def test_custom_presets_extend_built_in_presets() -> None:
    presets = presets_from_config(
        FakeConfig({"training_presets": {"My Quick Run": {"trainer": "lora", "epochs": 5}}})
    )
    assert "Character LoRA" in presets
    assert presets["My Quick Run"]["epochs"] == 5


def test_preflight_is_saved_in_plan_summary(tmp_path: Path) -> None:
    trainer = tmp_path / "trainer"
    dataset = tmp_path / "dataset"
    output = tmp_path / "output" / "model"
    trainer.mkdir()
    dataset.mkdir()
    (dataset / "one.png").write_bytes(b"image")
    plan = ExecutionPlan(
        request="train",
        summary="Train a model.",
        steps=[
            PlanStep(
                "ddpm_trainer",
                "Train DDPM",
                "Train",
                {"dataset_dir": str(dataset), "output_dir": str(output)},
            )
        ],
    )
    config = FakeConfig({"tool_folders": {"ddpm_trainer": str(trainer)}})

    append_preflight_summary(plan, config)

    assert "Pre-flight:" in plan.summary
    assert "1 images found" in plan.summary
    assert "connected" in plan.summary


def test_training_completion_suggests_preview_review() -> None:
    plan = SimpleNamespace(steps=[SimpleNamespace(tool_id="lora_trainer")])
    assert "preview images" in completion_recommendation(plan)


def test_multiple_model_plans_are_combined_in_order() -> None:
    first = ExecutionPlan(
        request="first", summary="Train first.", project_name="First",
        steps=[PlanStep("ddpm_trainer", "First model", "Train first")],
        requires_confirmation=True,
    )
    second = ExecutionPlan(
        request="second", summary="Train second.", project_name="Second",
        steps=[PlanStep("lora_trainer", "Second model", "Train second")],
        requires_confirmation=True,
    )

    batch = combine_training_plans([first, second])

    assert batch.project_name == "Training batch (2 models)"
    assert [step.title for step in batch.steps] == ["First model", "Second model"]
    assert batch.requires_confirmation is True
    assert "failed step stops the batch" in batch.summary


def test_single_model_plan_is_not_wrapped_as_a_batch() -> None:
    plan = ExecutionPlan(
        request="one", summary="One.",
        steps=[PlanStep("ddpm_trainer", "One model", "Train one")],
    )

    assert combine_training_plans([plan]) is plan


def test_bulk_model_names_are_cleaned_and_deduplicated() -> None:
    assert parse_model_batch_names("1. Luigi\n- EarthBound\nluigi\n\n• South Park") == [
        "Luigi", "EarthBound", "South Park"
    ]


def test_batch_dataset_first_request_does_not_train() -> None:
    request = build_dataset_collection_request("Luigi", image_count=80)

    assert request == "Collect a dataset of 80 images of Luigi."
    assert "train" not in request.casefold()


def test_existing_dataset_match_links_a_clear_name_match(tmp_path: Path) -> None:
    folder = tmp_path / "Luigi S Mansion Dataset"
    folder.mkdir()
    dataset = SimpleNamespace(name=folder.name, path=str(folder))

    match = suggest_existing_dataset(
        {"model_name": "Luigi's Mansion", "subject": "Luigi's Mansion"}, [dataset]
    )

    assert match.status == "matched"
    assert match.dataset_name == "Luigi S Mansion Dataset"


def test_existing_dataset_match_keeps_ambiguous_names_for_manual_selection(tmp_path: Path) -> None:
    first = tmp_path / "Liminal Spaces Dataset"; first.mkdir()
    second = tmp_path / "Liminal Space Images Dataset"; second.mkdir()
    datasets = [
        SimpleNamespace(name=first.name, path=str(first)),
        SimpleNamespace(name=second.name, path=str(second)),
    ]

    match = suggest_existing_dataset({"model_name": "Liminal Space"}, datasets)

    assert match.status == "ambiguous"