Download tests/test_generations.py from SyntheticMDProductions/AI_Development_Automation_Manager: direct link, hf CLI and curl.
- Browser
- Download file 33.5 kB
-
https://huggingface.co/SyntheticMDProductions/AI_Development_Automation_Manager/resolve/refs%2Fpr%2F1/tests/test_generations.py
- Command line
-
hf download hf://SyntheticMDProductions/AI_Development_Automation_Manager@refs/pr/1/tests/test_generations.py
-
curl -L -o test_generations.py https://huggingface.co/SyntheticMDProductions/AI_Development_Automation_Manager/resolve/refs%2Fpr%2F1/tests/test_generations.py
33.5 kB
| from __future__ import annotations | |
| import json | |
| import shutil | |
| import threading | |
| from types import ModuleType | |
| from pathlib import Path | |
| from adam.assets import AssetRegistry | |
| from adam.generations import ( | |
| build_generation_plan, | |
| generation_output_folder, | |
| generation_model_match_score, | |
| generation_model_key, | |
| generation_tools, | |
| group_generation_records, | |
| load_generation_history, | |
| parse_chat_generation_request, | |
| combine_generation_plans, | |
| ) | |
| from adam.registry import ToolRegistry | |
| from adam.config import ConfigManager | |
| from adam.executor import ToolContext | |
| from adam.generation_previews import accepts_preview_callback, publish_generation_preview | |
| from adam.image_preferences import PreferenceScore | |
| from adam.registry import ToolSpec | |
| from adam.tools import ddpm_generator | |
| from adam.tools import flow_generator | |
| from adam.tools import lora_generator | |
| from adam.showcase import build_showcase_plan | |
| ROOT = Path(__file__).resolve().parents[1] | |
| def test_generation_preview_protocol_keeps_only_latest_preview(tmp_path: Path) -> None: | |
| class FakeImage: | |
| def save(self, path: Path, *, format: str) -> None: | |
| assert format == "PNG" | |
| path.write_bytes(b"preview") | |
| published: list[dict] = [] | |
| running = threading.Event(); running.set() | |
| context = ToolContext( | |
| root=tmp_path, job_id="PREVIEW1", | |
| tool=ToolSpec("ddpm_generator", "DDPM", "test", "Output", "generate_ddpm_images"), | |
| cancel_event=threading.Event(), run_event=running, | |
| progress_callback=lambda *_args: None, log_callback=lambda *_args: None, | |
| preview_callback=published.append, | |
| ) | |
| output = tmp_path / "output"; output.mkdir() | |
| publish_generation_preview( | |
| context, output, FakeImage(), image_index=0, image_count=2, step=5, total_steps=20, | |
| ) | |
| assert (output / ".live_previews" / "PREVIEW1_latest.png").read_bytes() == b"preview" | |
| assert published == [{ | |
| "path": str(output / ".live_previews" / "PREVIEW1_latest.png"), | |
| "epoch": 0, "next_epoch": 0, "prompt": "", "seed": None, "steps": 0, | |
| "kind": "generation", "current": 5, "total": 20, "image_index": 1, "image_count": 2, | |
| }] | |
| def test_preview_callback_protocol_requires_explicit_parameter() -> None: | |
| def supported(*, preview_callback): | |
| return preview_callback | |
| def unsupported(**_settings): | |
| return None | |
| assert accepts_preview_callback(supported) | |
| assert not accepts_preview_callback(unsupported) | |
| def test_command_center_parses_quoted_ddpm_generation_request() -> None: | |
| parsed = parse_chat_generation_request( | |
| 'Generate a "DDPM" image of "Person" for "100" steps, on "DDIM" sampler, with aspect ratio of "16:9"' | |
| ) | |
| assert parsed is not None | |
| assert parsed.provider_hint == "ddpm" | |
| assert parsed.prompt == "Person" | |
| assert parsed.steps == 100 | |
| assert parsed.sampler == "DDIM" | |
| assert parsed.aspect_ratio == "16:9" | |
| def test_command_center_parses_batch_seed_and_model() -> None: | |
| parsed = parse_chat_generation_request( | |
| 'Create 3 images of rainy neon streets using model "City Nights" with seed 42 on DDPM sampler at ratio 1:1' | |
| ) | |
| assert parsed is not None | |
| assert parsed.model_query == "City Nights" | |
| assert parsed.prompt == "rainy neon streets" | |
| assert parsed.image_count == 3 | |
| assert parsed.seed == 42 | |
| assert parsed.sampler == "DDPM" | |
| def test_command_center_does_not_capture_non_image_plans() -> None: | |
| assert parse_chat_generation_request("Generate four training previews") is None | |
| def test_command_center_subject_matches_completed_model_name() -> None: | |
| assert generation_model_match_score("rouge the bat", "Rouge The Bat V2 MADA") > 0 | |
| assert generation_model_match_score("rouge the bat", "SpectrogramV3") == 0 | |
| def test_command_center_matches_compact_model_names() -> None: | |
| assert generation_model_match_score( | |
| "neon convenience stores at night", "NeonConvenienceStoresAtNight" | |
| ) > 0 | |
| def test_command_center_separates_lora_subject_and_positive_prompt() -> None: | |
| parsed = parse_chat_generation_request( | |
| 'Generate a LoRA image of OrangeCat, Base Model "novaFurryXL", Positive Prompt "OrangeCat, Anthro, female, pretty, sitting on bed, looking at viewer", Negative Prompt "Bad Quality, low effort, missing limbs, poor anatomy"' | |
| ) | |
| assert parsed is not None | |
| assert parsed.provider_hint == "lora" | |
| assert parsed.subject == "OrangeCat" | |
| assert parsed.base_model_query == "novaFurryXL" | |
| assert parsed.prompt == ( | |
| "OrangeCat, Anthro, female, pretty, sitting on bed, looking at viewer" | |
| ) | |
| assert parsed.negative_prompt == "Bad Quality, low effort, missing limbs, poor anatomy" | |
| def test_command_center_parses_advanced_generation_overrides() -> None: | |
| parsed = parse_chat_generation_request( | |
| 'Generate an image, Positive Prompt "city at night", CFG 6.5, ' | |
| 'LoRA strength 0.75, denoise strength 0.4, reference strength 70%' | |
| ) | |
| assert parsed is not None | |
| assert parsed.prompt == "city at night" | |
| assert parsed.cfg_scale == 6.5 | |
| assert parsed.lora_strength == 0.75 | |
| assert parsed.denoise_strength == 0.4 | |
| assert parsed.reference_strength == 70 | |
| def test_command_center_separates_lora_from_base_model_in_natural_phrasing() -> None: | |
| parsed = parse_chat_generation_request( | |
| 'Generate an image of LoRA OrangeCat, 30 steps, Base Model "waiIllustriousSDXL"' | |
| ) | |
| assert parsed is not None | |
| assert parsed.provider_hint == "lora" | |
| assert parsed.model_query == "OrangeCat" | |
| assert parsed.base_model_query == "waiIllustriousSDXL" | |
| def test_plain_subject_generation_is_not_marked_as_stable_diffusion_prompt() -> None: | |
| parsed = parse_chat_generation_request("Generate an image of Minecraft") | |
| assert parsed is not None | |
| assert parsed.subject == "Minecraft" | |
| assert parsed.has_positive_prompt is False | |
| assert parsed.provider_hint == "" | |
| def test_positive_prompt_is_an_explicit_stable_diffusion_signal() -> None: | |
| parsed = parse_chat_generation_request( | |
| 'Generate an image, Positive Prompt "Minecraft, blocky world, player"' | |
| ) | |
| assert parsed is not None | |
| assert parsed.prompt == "Minecraft, blocky world, player" | |
| assert parsed.has_positive_prompt is True | |
| def test_plain_flow_model_suffix_selects_flow_provider() -> None: | |
| parsed = parse_chat_generation_request("Generate an image of Minecraft Flow") | |
| assert parsed is not None | |
| assert parsed.subject == "Minecraft Flow" | |
| assert parsed.provider_hint == "flow" | |
| def test_external_lora_and_base_model_drop_folders_are_discovered(tmp_path: Path) -> None: | |
| lora_folder = tmp_path / "LoRAModelsHere" | |
| base_folder = tmp_path / "LoRA StableDiffusionModels Here" | |
| lora_folder.mkdir() | |
| base_folder.mkdir() | |
| (lora_folder / "OrangeCat.safetensors").write_bytes(b"lora") | |
| (base_folder / "sdxl-base.safetensors").write_bytes(b"base") | |
| assets = AssetRegistry(tmp_path) | |
| assets.discover({"tool_folders": {}}) | |
| assert any( | |
| asset.kind == "model" and asset.trainer == "lora" and asset.name == "OrangeCat" | |
| for asset in assets.assets | |
| ) | |
| assert any( | |
| asset.kind == "base_model" and asset.name == "sdxl-base" | |
| for asset in assets.assets | |
| ) | |
| def _registry(tmp_path: Path) -> ToolRegistry: | |
| (tmp_path / "config").mkdir() | |
| shutil.copy2(ROOT / "config" / "tools.json", tmp_path / "config" / "tools.json") | |
| return ToolRegistry(tmp_path) | |
| def test_registry_declares_ddpm_image_generator(tmp_path: Path) -> None: | |
| registry = _registry(tmp_path) | |
| tool = registry.get("ddpm_generator") | |
| assert "image_generation" in tool.capabilities | |
| assert tool.model_trainers == ("ddpm",) | |
| assert [item.id for item in generation_tools(registry)] == [ | |
| "ddpm_generator", | |
| "flow_generator", | |
| "lora_generator", | |
| ] | |
| def test_registry_declares_lora_image_generator(tmp_path: Path) -> None: | |
| tool = _registry(tmp_path).get("lora_generator") | |
| assert "image_generation" in tool.capabilities | |
| assert "text_prompt" in tool.capabilities | |
| assert tool.model_trainers == ("lora",) | |
| assert "DPM++ 2M" in tool.generation_options["samplers"] | |
| assert "negative_prompt" in tool.arguments | |
| assert "base_model_path" in tool.arguments | |
| def test_lora_adapter_accepts_cancelled_checkpoints(tmp_path: Path) -> None: | |
| cancelled = tmp_path / "Character_cancelled.safetensors" | |
| cancelled.write_bytes(b"test") | |
| assert lora_generator._lora_file(cancelled) == cancelled | |
| def test_registry_declares_flow_image_generator(tmp_path: Path) -> None: | |
| tool = _registry(tmp_path).get("flow_generator") | |
| assert "image_generation" in tool.capabilities | |
| assert tool.model_trainers == ("flow",) | |
| assert tool.generation_options["samplers"] == ["Heun", "Euler"] | |
| def test_generation_plan_preserves_reproducible_settings(tmp_path: Path) -> None: | |
| tool = _registry(tmp_path).get("ddpm_generator") | |
| plan = build_generation_plan( | |
| tool, | |
| model_name="Mario V2", | |
| model_path="D:/DDPM/output/Mario", | |
| prompt="Explore colorful shapes", | |
| image_count=4, | |
| steps=75, | |
| seed=1234, | |
| sampler="DDIM", | |
| aspect_ratio="16:9 (Widescreen)", | |
| ) | |
| assert plan.requires_confirmation is False | |
| assert plan.steps[0].tool_id == "ddpm_generator" | |
| assert plan.steps[0].arguments == { | |
| "model_name": "Mario V2", | |
| "model_path": "D:/DDPM/output/Mario", | |
| "prompt": "Explore colorful shapes", | |
| "image_count": 4, | |
| "steps": 75, | |
| "seed": 1234, | |
| "sampler": "DDIM", | |
| "aspect_ratio": "16:9 (Widescreen)", | |
| } | |
| def test_generation_cycle_keeps_models_in_selected_order(tmp_path: Path) -> None: | |
| tool = _registry(tmp_path).get("ddpm_generator") | |
| plans = [ | |
| build_generation_plan(tool, model_name=name, model_path=f"D:/{name}", | |
| prompt="", image_count=3, steps=20, seed=10, | |
| sampler="DDIM", aspect_ratio="1:1 (Square)") | |
| for name in ("Minecraft", "Roblox") | |
| ] | |
| cycle = combine_generation_plans(plans, display_seconds=7, show_labels=True) | |
| assert cycle.project_name == "Generation Cycle" | |
| assert [step.arguments["model_name"] for step in cycle.steps] == ["Minecraft", "Roblox"] | |
| assert "7 seconds" in cycle.summary | |
| def test_showcase_plan_generates_then_renders_mp4(tmp_path: Path) -> None: | |
| registry = _registry(tmp_path) | |
| ddpm = registry.get("ddpm_generator") | |
| flow = registry.get("flow_generator") | |
| plans = [ | |
| build_generation_plan( | |
| ddpm, model_name="Windows XP", model_path="D:/DDPM/WindowsXP", | |
| prompt="", image_count=18, steps=30, seed=100, | |
| sampler="DDIM", aspect_ratio="16:9 (Widescreen)", | |
| ), | |
| build_generation_plan( | |
| flow, model_name="Adventure Time", model_path="D:/Flow/AdventureTime", | |
| prompt="", image_count=18, steps=30, seed=118, | |
| sampler="Heun", aspect_ratio="16:9 (Widescreen)", | |
| ), | |
| ] | |
| settings = [ | |
| {"name": "Windows XP", "trainer": "ddpm", "trainer_label": "DDPM", "steps": 30, "sampler": "DDIM", "aspect_ratio": "16:9 (Widescreen)"}, | |
| {"name": "Adventure Time", "trainer": "flow", "trainer_label": "Flow Matching", "steps": 30, "sampler": "Heun", "aspect_ratio": "16:9 (Widescreen)"}, | |
| ] | |
| showcase = build_showcase_plan( | |
| plans, title="21 Requests", display_seconds=4, | |
| resolution="1080p", model_settings=settings, | |
| ) | |
| assert showcase.project_name == "Showcase Video" | |
| assert [step.tool_id for step in showcase.steps] == [ | |
| "ddpm_generator", "flow_generator", "showcase_video_renderer" | |
| ] | |
| assert showcase.steps[-1].arguments["models"] == settings | |
| assert "36 images" in showcase.summary | |
| assert "4 seconds" in showcase.summary | |
| def test_showcase_rejects_unsupported_image_duration(tmp_path: Path) -> None: | |
| tool = _registry(tmp_path).get("ddpm_generator") | |
| plan = build_generation_plan( | |
| tool, model_name="Test", model_path="D:/Test", prompt="", | |
| image_count=12, steps=30, seed=1, sampler="DDIM", | |
| aspect_ratio="16:9 (Widescreen)", | |
| ) | |
| try: | |
| build_showcase_plan( | |
| [plan], title="Test", display_seconds=2, | |
| resolution="720p", model_settings=[{"name": "Test"}], | |
| ) | |
| except ValueError as exc: | |
| assert "3, 4, or 5" in str(exc) | |
| else: | |
| raise AssertionError("Unsupported showcase duration was accepted") | |
| def test_generation_history_reads_images_and_ignores_broken_batches(tmp_path: Path) -> None: | |
| good = tmp_path / "data" / "generations" / "good" | |
| broken = tmp_path / "data" / "generations" / "broken" | |
| good.mkdir(parents=True) | |
| broken.mkdir() | |
| (good / "image_001.png").write_bytes(b"not decoded by the history loader") | |
| (good / "generation.json").write_text( | |
| json.dumps( | |
| { | |
| "provider_id": "ddpm_generator", | |
| "provider_name": "DDPM Generator", | |
| "model_name": "Mario V2", | |
| "model_path": "D:/DDPM/output/Mario", | |
| "prompt": "Color study", | |
| "seed": 42, | |
| "steps": 50, | |
| "sampler": "DDIM", | |
| "aspect_ratio": "1:1 (Square)", | |
| "created_at": "2026-07-31T12:00:00+00:00", | |
| } | |
| ), | |
| encoding="utf-8", | |
| ) | |
| (broken / "generation.json").write_text("not json", encoding="utf-8") | |
| records = load_generation_history(tmp_path) | |
| assert len(records) == 1 | |
| assert records[0].model_name == "Mario V2" | |
| assert records[0].seed == 42 | |
| assert records[0].images == (good / "image_001.png",) | |
| def test_generation_history_reads_new_model_folder_layout(tmp_path: Path) -> None: | |
| folder = generation_output_folder(tmp_path, "ddpm_generator", "Mario V2") | |
| image = folder / "20260802_120000_TEST_DDIM_seed_42.png" | |
| image.write_bytes(b"image") | |
| (folder / "generation_20260802_120000_TEST.json").write_text( | |
| json.dumps({"model_name": "Mario V2", "images": [str(image)], "created_at": "2026-08-02T12:00:00+00:00"}), | |
| encoding="utf-8", | |
| ) | |
| records = load_generation_history(tmp_path) | |
| assert len(records) == 1 | |
| assert records[0].folder == folder | |
| assert records[0].images == (image,) | |
| def test_generation_history_groups_model_folders_with_latest_image_cover(tmp_path: Path) -> None: | |
| mario_model = tmp_path / "models" / "Mario V2" | |
| luigi_model = tmp_path / "models" / "Luigi" | |
| mario_model.mkdir(parents=True) | |
| luigi_model.mkdir(parents=True) | |
| mario_folder = generation_output_folder(tmp_path, "ddpm_generator", "Mario V2") | |
| luigi_folder = generation_output_folder(tmp_path, "flow_generator", "Luigi") | |
| def add_batch(folder: Path, stamp: str, model_name: str, model_path: Path, images: int) -> list[Path]: | |
| paths = [] | |
| for index in range(images): | |
| path = folder / f"{stamp}_{index}.png" | |
| path.write_bytes(b"image") | |
| paths.append(path) | |
| (folder / f"generation_{stamp}.json").write_text(json.dumps({ | |
| "provider_id": "ddpm_generator" if "Mario" in model_name else "flow_generator", | |
| "provider_name": "DDPM Generator" if "Mario" in model_name else "Flow Generator", | |
| "model_name": model_name, | |
| "model_path": str(model_path), | |
| "images": [str(path) for path in paths], | |
| "created_at": stamp, | |
| }), encoding="utf-8") | |
| return paths | |
| old_mario = add_batch(mario_folder, "2026-08-01", "Mario V2", mario_model, 2) | |
| new_mario = add_batch(mario_folder, "2026-08-03", "Mario V2", mario_model, 1) | |
| add_batch(luigi_folder, "2026-08-02", "Luigi", luigi_model, 1) | |
| records = load_generation_history(tmp_path) | |
| folders = group_generation_records(records) | |
| assert [folder.model_name for folder in folders] == ["Mario V2", "Luigi"] | |
| assert folders[0].image_count == 3 | |
| assert folders[0].cover_image == new_mario[0] | |
| assert tuple(record.created_at for record in folders[0].records) == ( | |
| "2026-08-03", "2026-08-01", | |
| ) | |
| assert generation_model_key(folders[0].records[0]).startswith("path:") | |
| assert old_mario[0] in folders[0].records[1].images | |
| def test_ddpm_adapter_writes_images_and_reproducibility_metadata( | |
| tmp_path: Path, monkeypatch | |
| ) -> None: | |
| trainer_root = tmp_path / "connected-ddpm" | |
| model = trainer_root / "output" / "Mario" | |
| model.mkdir(parents=True) | |
| (model / "model_index.json").write_text("{}", encoding="utf-8") | |
| script = trainer_root / "appStableDiffusion.py" | |
| script.write_text("# test backend", encoding="utf-8") | |
| ConfigManager(tmp_path).update( | |
| {"tool_folders": {"ddpm_trainer": str(trainer_root)}} | |
| ) | |
| calls: list[dict] = [] | |
| class FakeImage: | |
| def save(self, path: Path, *, format: str) -> None: | |
| assert format == "PNG" | |
| path.write_bytes(b"fake png") | |
| backend = ModuleType("fake_ddpm") | |
| def generate_images(_model: str, **settings): | |
| calls.append(settings) | |
| return [FakeImage()] | |
| backend.generate_images = generate_images # type: ignore[attr-defined] | |
| monkeypatch.setattr(ddpm_generator, "_backend_module", backend) | |
| monkeypatch.setattr(ddpm_generator, "_backend_script", script.resolve()) | |
| monkeypatch.setattr(ddpm_generator.importlib.util, "find_spec", lambda _name: object()) | |
| run_event = threading.Event() | |
| run_event.set() | |
| tool = ToolSpec( | |
| id="ddpm_generator", | |
| name="DDPM Generator", | |
| description="test", | |
| category="Output", | |
| entry_function="generate_ddpm_images", | |
| ) | |
| context = ToolContext( | |
| root=tmp_path, | |
| job_id="TEST1234", | |
| tool=tool, | |
| cancel_event=threading.Event(), | |
| run_event=run_event, | |
| progress_callback=lambda _percent, _message: None, | |
| log_callback=lambda _message: None, | |
| step_delay=0, | |
| ) | |
| result = ddpm_generator.generate_ddpm_images( | |
| context, | |
| model_name="Mario", | |
| model_path=str(model), | |
| prompt="Color study", | |
| image_count=2, | |
| steps=20, | |
| seed=100, | |
| sampler="DDIM", | |
| aspect_ratio="16:9 (Widescreen)", | |
| ) | |
| output = Path(str(result["output_folder"])) | |
| metadata = json.loads(next(output.glob("generation_*.json")).read_text(encoding="utf-8")) | |
| assert metadata["image_seeds"] == [100, 101] | |
| assert metadata["prompt_behavior"] == "label_only" | |
| assert len(list(output.glob("*.png"))) == 2 | |
| assert [call["seed"] for call in calls] == [100, 101] | |
| def test_ddpm_adapter_uses_native_reference_image_generation(tmp_path: Path, monkeypatch) -> None: | |
| trainer_root = tmp_path / "connected-ddpm" | |
| model = trainer_root / "output" / "Mario" | |
| model.mkdir(parents=True) | |
| (model / "model_index.json").write_text("{}", encoding="utf-8") | |
| script = trainer_root / "appStableDiffusion.py" | |
| script.write_text("# test backend", encoding="utf-8") | |
| reference = tmp_path / "reference.png" | |
| reference.write_bytes(b"fake reference") | |
| ConfigManager(tmp_path).update({"tool_folders": {"ddpm_trainer": str(trainer_root)}}) | |
| calls: list[dict] = [] | |
| class FakeImage: | |
| def save(self, path: Path, *, format: str) -> None: | |
| path.write_bytes(b"fake png") | |
| backend = ModuleType("fake_ddpm_reference") | |
| backend.generate_images = lambda *_args, **_kwargs: [FakeImage()] # type: ignore[attr-defined] | |
| def generate_reference_images(_model: str, image: str, **settings): | |
| calls.append({"image": image, **settings}) | |
| return [FakeImage()] | |
| backend.generate_reference_images = generate_reference_images # type: ignore[attr-defined] | |
| monkeypatch.setattr(ddpm_generator, "_backend_module", backend) | |
| monkeypatch.setattr(ddpm_generator, "_backend_script", script.resolve()) | |
| monkeypatch.setattr(ddpm_generator.importlib.util, "find_spec", lambda _name: object()) | |
| event = threading.Event(); event.set() | |
| context = ToolContext(tmp_path, "REF1234", ToolSpec("ddpm_generator", "DDPM Generator", "test", "Output", "generate_ddpm_images"), threading.Event(), event, lambda *_args: None, lambda *_args: None, 0) | |
| result = ddpm_generator.generate_ddpm_images( | |
| context, "Mario", str(model), "Reference study", 1, 20, 100, "DDIM", | |
| "1:1 (Square)", str(reference), 72, | |
| ) | |
| metadata = json.loads(next(Path(str(result["output_folder"])).glob("generation_*.json")).read_text(encoding="utf-8")) | |
| assert calls == [{"image": str(reference.resolve()), "seed": 100, "num_inference_steps": 20, "batch_size": 1, "sampler": "DDIM", "aspect_ratio": "1:1 (Square)", "reference_strength": 72}] | |
| assert metadata["reference_image"] == str(reference.resolve()) | |
| assert metadata["reference_strength"] == 72 | |
| def test_ddpm_adapter_passes_enabled_custom_dimensions(tmp_path: Path, monkeypatch) -> None: | |
| trainer_root = tmp_path / "connected-ddpm" | |
| model = trainer_root / "output" / "Mario" | |
| model.mkdir(parents=True) | |
| (model / "model_index.json").write_text("{}", encoding="utf-8") | |
| script = trainer_root / "appStableDiffusion.py" | |
| script.write_text("# test backend", encoding="utf-8") | |
| ConfigManager(tmp_path).update({"tool_folders": {"ddpm_trainer": str(trainer_root)}}) | |
| calls: list[dict] = [] | |
| class FakeImage: | |
| def save(self, path: Path, *, format: str) -> None: | |
| path.write_bytes(b"fake png") | |
| backend = ModuleType("fake_ddpm_size") | |
| def generate_images(_model: str, **settings): | |
| calls.append(settings) | |
| return [FakeImage()] | |
| backend.generate_images = generate_images # type: ignore[attr-defined] | |
| monkeypatch.setattr(ddpm_generator, "_backend_module", backend) | |
| monkeypatch.setattr(ddpm_generator, "_backend_script", script.resolve()) | |
| monkeypatch.setattr(ddpm_generator.importlib.util, "find_spec", lambda _name: object()) | |
| event = threading.Event(); event.set() | |
| context = ToolContext(tmp_path, "SIZE1234", ToolSpec("ddpm_generator", "DDPM Generator", "test", "Output", "generate_ddpm_images"), threading.Event(), event, lambda *_args: None, lambda *_args: None, 0) | |
| result = ddpm_generator.generate_ddpm_images( | |
| context, "Mario", str(model), "Size study", 1, 20, 100, "DDIM", | |
| "1:1 (Square)", width=320, height=192, | |
| ) | |
| metadata = json.loads(next(Path(str(result["output_folder"])).glob("generation_*.json")).read_text(encoding="utf-8")) | |
| assert calls[0]["width"] == 320 | |
| assert calls[0]["height"] == 192 | |
| assert metadata["width"] == 320 | |
| assert metadata["height"] == 192 | |
| def test_ddpm_smart_generation_stops_after_enough_passing_candidates(tmp_path: Path, monkeypatch) -> None: | |
| trainer_root = tmp_path / "connected-ddpm" | |
| model = trainer_root / "output" / "Mario" | |
| model.mkdir(parents=True) | |
| (model / "model_index.json").write_text("{}", encoding="utf-8") | |
| script = trainer_root / "appStableDiffusion.py" | |
| script.write_text("# test backend", encoding="utf-8") | |
| ConfigManager(tmp_path).update({"tool_folders": {"ddpm_trainer": str(trainer_root)}}) | |
| class FakeImage: | |
| def save(self, path: Path, *, format: str) -> None: | |
| path.write_bytes(b"fake png") | |
| backend = ModuleType("fake_ddpm_smart") | |
| backend.generate_images = lambda *_args, **_settings: [FakeImage()] # type: ignore[attr-defined] | |
| monkeypatch.setattr(ddpm_generator, "_backend_module", backend) | |
| monkeypatch.setattr(ddpm_generator, "_backend_script", script.resolve()) | |
| monkeypatch.setattr(ddpm_generator.importlib.util, "find_spec", lambda _name: object()) | |
| class FakeVision: | |
| def unload(self) -> None: | |
| pass | |
| class FakeEvaluator: | |
| def __init__(self, _root): | |
| self.vision = FakeVision() | |
| def score(self, _profile, paths, *, keep_threshold=None, reject_threshold=None): | |
| seed = int(str(paths[0]).rsplit("_seed_", 1)[1].split(".", 1)[0]) | |
| score = {10: 0.2, 11: 0.8, 12: 0.9}.get(seed, 0.1) | |
| category = "Strong Keep" if score >= float(keep_threshold or 0.7) else "Likely Reject" | |
| return [PreferenceScore(str(Path(paths[0]).resolve()), score, score, category)] | |
| monkeypatch.setattr(ddpm_generator, "GenerationPreferenceEvaluator", FakeEvaluator) | |
| event = threading.Event(); event.set() | |
| context = ToolContext(tmp_path, "SMARTDDPM", ToolSpec("ddpm_generator", "DDPM Generator", "test", "Output", "generate_ddpm_images"), threading.Event(), event, lambda *_args: None, lambda *_args: None, 0) | |
| result = ddpm_generator.generate_ddpm_images( | |
| context, "Mario", str(model), "Smart study", 2, 20, 10, "DDIM", | |
| "1:1 (Square)", smart_generation=True, smart_wanted_results=2, | |
| smart_max_candidates=5, smart_min_score=0.7, | |
| ) | |
| output = Path(str(result["output_folder"])) | |
| metadata = json.loads(next(output.glob("generation_*.json")).read_text(encoding="utf-8")) | |
| assert metadata["smart_generation"]["selected_count"] == 2 | |
| assert metadata["smart_generation"]["candidate_count"] == 3 | |
| assert len(metadata["images"]) == 3 | |
| assert "_seed_11" in metadata["images"][0] | |
| assert "_seed_12" in metadata["images"][1] | |
| def test_ddpm_smart_generation_reports_insufficient_passing_candidates(tmp_path: Path, monkeypatch) -> None: | |
| trainer_root = tmp_path / "connected-ddpm" | |
| model = trainer_root / "output" / "Mario" | |
| model.mkdir(parents=True) | |
| (model / "model_index.json").write_text("{}", encoding="utf-8") | |
| script = trainer_root / "appStableDiffusion.py" | |
| script.write_text("# test backend", encoding="utf-8") | |
| ConfigManager(tmp_path).update({"tool_folders": {"ddpm_trainer": str(trainer_root)}}) | |
| class FakeImage: | |
| def save(self, path: Path, *, format: str) -> None: | |
| path.write_bytes(b"fake png") | |
| backend = ModuleType("fake_ddpm_smart_insufficient") | |
| backend.generate_images = lambda *_args, **_settings: [FakeImage()] # type: ignore[attr-defined] | |
| monkeypatch.setattr(ddpm_generator, "_backend_module", backend) | |
| monkeypatch.setattr(ddpm_generator, "_backend_script", script.resolve()) | |
| monkeypatch.setattr(ddpm_generator.importlib.util, "find_spec", lambda _name: object()) | |
| class FakeVision: | |
| def unload(self) -> None: | |
| pass | |
| class FakeEvaluator: | |
| def __init__(self, _root): | |
| self.vision = FakeVision() | |
| def score(self, _profile, paths, *, keep_threshold=None, reject_threshold=None): | |
| return [PreferenceScore(str(Path(paths[0]).resolve()), 0.4, 0.6, "Needs Review")] | |
| monkeypatch.setattr(ddpm_generator, "GenerationPreferenceEvaluator", FakeEvaluator) | |
| event = threading.Event(); event.set() | |
| context = ToolContext(tmp_path, "SMARTLOW", ToolSpec("ddpm_generator", "DDPM Generator", "test", "Output", "generate_ddpm_images"), threading.Event(), event, lambda *_args: None, lambda *_args: None, 0) | |
| result = ddpm_generator.generate_ddpm_images( | |
| context, "Mario", str(model), "Smart study", 3, 20, 20, "DDIM", | |
| "1:1 (Square)", smart_generation=True, smart_wanted_results=3, | |
| smart_max_candidates=4, smart_min_score=0.7, | |
| ) | |
| metadata = json.loads(next(Path(str(result["output_folder"])).glob("generation_*.json")).read_text(encoding="utf-8")) | |
| assert metadata["smart_generation"]["selected_count"] == 0 | |
| assert metadata["smart_generation"]["candidate_count"] == 4 | |
| def test_flow_adapter_uses_registered_model_and_tracks_each_seed( | |
| tmp_path: Path, monkeypatch | |
| ) -> None: | |
| flow_root = tmp_path / "connected-flow" | |
| model = flow_root / "output_flow_models" / "Rooms" | |
| (model / "unet").mkdir(parents=True) | |
| (model / "unet" / "config.json").write_text("{}", encoding="utf-8") | |
| (model / "flow_model_info.json").write_text( | |
| json.dumps({"model_type": "rectified_flow", "model_name": "Rooms"}), | |
| encoding="utf-8", | |
| ) | |
| script = flow_root / "flow_matching_app.py" | |
| script.write_text("# test backend", encoding="utf-8") | |
| ConfigManager(tmp_path).update( | |
| {"tool_folders": {"flow_trainer": str(flow_root)}} | |
| ) | |
| calls: list[dict] = [] | |
| class FakeImage: | |
| def save(self, path: Path, *, format: str) -> None: | |
| assert format == "PNG" | |
| path.write_bytes(b"fake flow png") | |
| backend = ModuleType("fake_flow") | |
| backend.load_unet = lambda *_args, **_kwargs: object() # type: ignore[attr-defined] | |
| def sample_flow(_model, _count, steps, _device, _dtype, seed, method, progress, **settings): | |
| progress(steps, steps) | |
| calls.append({"seed": seed, "method": method, **settings}) | |
| return [FakeImage()] | |
| backend.sample_flow = sample_flow # type: ignore[attr-defined] | |
| monkeypatch.setattr(flow_generator, "_backend_module", backend) | |
| monkeypatch.setattr(flow_generator, "_backend_script", script.resolve()) | |
| monkeypatch.setattr(flow_generator, "_loaded_model", None) | |
| monkeypatch.setattr(flow_generator, "_loaded_model_path", None) | |
| monkeypatch.setattr(flow_generator.importlib.util, "find_spec", lambda _name: object()) | |
| run_event = threading.Event() | |
| run_event.set() | |
| context = ToolContext( | |
| root=tmp_path, | |
| job_id="FLOW1234", | |
| tool=ToolSpec( | |
| id="flow_generator", | |
| name="Flow Matching Generator", | |
| description="test", | |
| category="Output", | |
| entry_function="generate_flow_images", | |
| ), | |
| cancel_event=threading.Event(), | |
| run_event=run_event, | |
| progress_callback=lambda _percent, _message: None, | |
| log_callback=lambda _message: None, | |
| step_delay=0, | |
| ) | |
| result = flow_generator.generate_flow_images( | |
| context, | |
| model_name="Rooms", | |
| model_path=str(model), | |
| prompt="Room study", | |
| image_count=2, | |
| steps=8, | |
| seed=700, | |
| sampler="Heun", | |
| aspect_ratio="4:3 (Landscape)", | |
| ) | |
| output = Path(str(result["output_folder"])) | |
| metadata = json.loads(next(output.glob("generation_*.json")).read_text(encoding="utf-8")) | |
| assert metadata["provider_id"] == "flow_generator" | |
| assert metadata["image_seeds"] == [700, 701] | |
| assert [call["seed"] for call in calls] == [700, 701] | |
| assert all(call["method"] == "Heun" for call in calls) | |
| assert all(call["aspect_ratio"] == "4:3 (Landscape)" for call in calls) | |
| def test_flow_smart_generation_can_return_top_ranked_pool(tmp_path: Path, monkeypatch) -> None: | |
| flow_root = tmp_path / "connected-flow" | |
| model = flow_root / "output_flow_models" / "Rooms" | |
| (model / "unet").mkdir(parents=True) | |
| (model / "unet" / "config.json").write_text("{}", encoding="utf-8") | |
| (model / "flow_model_info.json").write_text( | |
| json.dumps({"model_type": "rectified_flow", "model_name": "Rooms"}), | |
| encoding="utf-8", | |
| ) | |
| script = flow_root / "flow_matching_app.py" | |
| script.write_text("# test backend", encoding="utf-8") | |
| ConfigManager(tmp_path).update({"tool_folders": {"flow_trainer": str(flow_root)}}) | |
| class FakeImage: | |
| def save(self, path: Path, *, format: str) -> None: | |
| path.write_bytes(b"fake flow png") | |
| backend = ModuleType("fake_flow_smart") | |
| backend.load_unet = lambda *_args, **_kwargs: object() # type: ignore[attr-defined] | |
| backend.sample_flow = lambda _model, _count, steps, _device, _dtype, seed, method, progress, **_settings: (progress(steps, steps), [FakeImage()])[1] # type: ignore[attr-defined] | |
| monkeypatch.setattr(flow_generator, "_backend_module", backend) | |
| monkeypatch.setattr(flow_generator, "_backend_script", script.resolve()) | |
| monkeypatch.setattr(flow_generator, "_loaded_model", None) | |
| monkeypatch.setattr(flow_generator, "_loaded_model_path", None) | |
| monkeypatch.setattr(flow_generator.importlib.util, "find_spec", lambda _name: object()) | |
| class FakeVision: | |
| def unload(self) -> None: | |
| pass | |
| class FakeEvaluator: | |
| def __init__(self, _root): | |
| self.vision = FakeVision() | |
| def score(self, _profile, paths, *, keep_threshold=None, reject_threshold=None): | |
| seed = int(str(paths[0]).rsplit("_seed_", 1)[1].split(".", 1)[0]) | |
| score = {50: 0.1, 51: 0.9, 52: 0.3, 53: 0.8}[seed] | |
| return [PreferenceScore(str(Path(paths[0]).resolve()), score, score, "Strong Keep" if score >= 0.7 else "Needs Review")] | |
| monkeypatch.setattr(flow_generator, "GenerationPreferenceEvaluator", FakeEvaluator) | |
| event = threading.Event(); event.set() | |
| context = ToolContext(tmp_path, "SMARTFLOW", ToolSpec("flow_generator", "Flow Matching Generator", "test", "Output", "generate_flow_images"), threading.Event(), event, lambda *_args: None, lambda *_args: None, 0) | |
| result = flow_generator.generate_flow_images( | |
| context, "Rooms", str(model), "Room study", 2, 8, 50, "Heun", | |
| "4:3 (Landscape)", smart_generation=True, smart_wanted_results=2, | |
| smart_max_candidates=4, smart_min_score=0.7, smart_mode="top_n", | |
| ) | |
| metadata = json.loads(next(Path(str(result["output_folder"])).glob("generation_*.json")).read_text(encoding="utf-8")) | |
| assert metadata["smart_generation"]["candidate_count"] == 4 | |
| assert metadata["smart_generation"]["selected_count"] == 2 | |
| assert "_seed_51" in metadata["images"][0] | |
| assert "_seed_53" in metadata["images"][1] | |