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"