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"
|