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