SyntheticMDProductions's picture
Some of Adams structure
e0265b9 verified
Raw
History Blame
23.6 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_tools,
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.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_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_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)