File size: 7,605 Bytes
e0265b9 c61c435 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 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 | 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_preflight_does_not_treat_empty_tool_path_as_connected(monkeypatch) -> None:
def unexpected_scan(*args, **kwargs):
raise AssertionError("An empty dataset path must not scan the working directory")
monkeypatch.setattr(Path, "rglob", unexpected_scan)
plan = ExecutionPlan(
request="train", summary="Train.",
steps=[PlanStep("ddpm_trainer", "Train DDPM", "Train", {"epochs": 10})],
)
append_preflight_summary(plan, FakeConfig())
assert "program folder is not connected" in plan.summary
assert "Train DDPM: connected" not in plan.summary
def test_orion_text_without_a_report_does_not_skip_review(tmp_path: Path) -> None:
plan = ExecutionPlan(
request="train", summary="ORION — mentioned in a user summary.",
steps=[
PlanStep("dataset_collector", "Collect", "Collect", {
"output_dir": str(tmp_path / "future"), "image_count": 2000,
}),
PlanStep("ddpm_trainer", "Train", "Train", {
"dataset_dir": str(tmp_path / "future"), "epochs": 600,
}),
],
)
append_preflight_summary(plan, FakeConfig())
assert plan.orion_review["level"] == "warning"
assert plan.requires_confirmation is True
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"
|