SyntheticMDProductions's picture
Some of Adams structure
e0265b9 verified
Raw
History Blame Contribute Delete
14.5 kB
from __future__ import annotations
import json
import re
from dataclasses import dataclass
from pathlib import Path
from typing import Any
from adam.models import ExecutionPlan, PlanStep
from adam.registry import ToolRegistry, ToolSpec
IMAGE_EXTENSIONS = {".png", ".jpg", ".jpeg", ".webp", ".bmp"}
@dataclass(frozen=True, slots=True)
class ChatGenerationRequest:
"""Generation settings recognized from a Command Center message."""
prompt: str
subject: str = ""
provider_hint: str = ""
model_query: str = ""
base_model_query: str = ""
negative_prompt: str = ""
image_count: int | None = None
steps: int | None = None
sampler: str = ""
aspect_ratio: str = ""
seed: int | None = None
cfg_scale: float | None = None
lora_strength: float | None = None
denoise_strength: float | None = None
reference_strength: int | None = None
reference_image: str = ""
has_positive_prompt: bool = False
_QUOTED = r'["\u201c\u201d]([^"\u201c\u201d]+)["\u201c\u201d]'
def _clean_chat_value(value: str) -> str:
return value.strip().strip('"\u201c\u201d').strip(" ,.;")
def generation_model_match_score(query: str, model_name: str) -> int:
"""Score whether conversational subject text clearly names a saved model."""
def words(value: str) -> list[str]:
value = re.sub(r"(?<=[a-z0-9])(?=[A-Z])", " ", value)
ignored = {"a", "an", "the", "of", "image", "picture", "model", "ddpm", "flow", "matching", "lora"}
return [word for word in re.findall(r"[a-z0-9]+", value.casefold()) if word not in ignored]
query_words = words(query)
model_words = words(model_name)
if not query_words or not model_words:
return 0
query_compact = "".join(query_words)
model_compact = "".join(model_words)
if query_compact == model_compact:
return 120
if query_compact in model_compact:
return 100 + min(10, len(query_words))
shared = len(set(query_words) & set(model_words))
if shared == len(set(query_words)):
return 90 + shared
coverage = shared / len(set(query_words))
return 60 + shared if shared >= 2 and coverage >= 0.7 else 0
def parse_chat_generation_request(text: str) -> ChatGenerationRequest | None:
"""Recognize a concise natural-language image generation command.
This intentionally requires both a creation verb and the word image/picture so
ordinary planning requests continue through the regular Command Center planner.
"""
request = " ".join(text.strip().split())
if not request or not re.search(r"\b(generate|create|make)\b", request, re.I):
return None
if not re.search(r"\b(image|images|picture|pictures)\b", request, re.I):
return None
provider_hint = ""
provider_match = re.search(
r"\b(ddpm|ddim|flow(?:\s+matching)?|lora)\b[\"\u201c\u201d]?(?=\s+(?:image|picture))",
request,
re.I,
)
if provider_match:
hint = provider_match.group(1).casefold()
provider_hint = "ddpm" if hint in {"ddpm", "ddim"} else "flow" if hint.startswith("flow") else "lora"
# Support natural phrasing such as "Generate an image of LoRA OrangeCat".
lora_subject_match = re.search(
rf"\b(?:image|picture)s?\s+of\s+(?:a\s+)?LoRA\s+{_QUOTED}",
request,
re.I,
)
if not lora_subject_match:
lora_subject_match = re.search(
r"\b(?:image|picture)s?\s+of\s+(?:a\s+)?LoRA\s+(.+?)(?=\s*(?:,|with\s+base\s+model|base\s+model|positive\s+prompt|negative\s+prompt|\d+\s+steps?|$))",
request,
re.I,
)
if lora_subject_match:
provider_hint = "lora"
subject = ""
prompt_match = re.search(
rf"\b(?:image|picture)s?\s+(?:of|showing|depicting)\s+{_QUOTED}",
request,
re.I,
)
if prompt_match:
subject = _clean_chat_value(prompt_match.group(1))
else:
prompt_match = re.search(
r"\b(?:image|picture)s?\s+(?:of|showing|depicting)\s+(.+?)(?=\s+(?:for|using|with|on|at)\s+|,|$)",
request,
re.I,
)
if prompt_match:
subject = _clean_chat_value(prompt_match.group(1))
# In promptless commands, a provider suffix is usually part of the saved
# model name (for example, "Minecraft Flow"), not prompt prose.
if not provider_hint and subject:
if re.search(r"\bflow(?:\s+match(?:ing)?)?\s*$", subject, re.I):
provider_hint = "flow"
elif re.search(r"\bddpm\s*$", subject, re.I):
provider_hint = "ddpm"
positive_match = re.search(
rf"\bpositive\s+prompt(?:\s+of|\s*=|\s*:)?\s*{_QUOTED}",
request,
re.I,
)
prompt = _clean_chat_value(positive_match.group(1)) if positive_match else subject
negative_match = re.search(
rf"\bnegative\s+prompt(?:\s+of|\s*=|\s*:)?\s*{_QUOTED}",
request,
re.I,
)
negative_prompt = _clean_chat_value(negative_match.group(1)) if negative_match else ""
# A quoted value after "using/with model" is treated as a completed model name.
model_query = ""
model_match = re.search(
rf"\b(?:using|with)\s+(?:the\s+)?(?:model\s+)?{_QUOTED}(?:\s+model)?",
request,
re.I,
)
if model_match:
model_query = _clean_chat_value(model_match.group(1))
if lora_subject_match:
model_query = _clean_chat_value(lora_subject_match.group(1))
base_model_match = re.search(
rf"\bbase\s+model(?:\s+of|\s*=|\s*:)?\s*{_QUOTED}",
request,
re.I,
)
base_model_query = _clean_chat_value(base_model_match.group(1)) if base_model_match else ""
count_match = re.search(r"\b(?:generate|create|make)\s+[\"\u201c]?([1-9]\d*)[\"\u201d]?\s+(?:images|pictures)\b", request, re.I)
steps_match = re.search(r"\b[\"\u201c]?(\d{1,4})[\"\u201d]?\s+(?:inference\s+)?steps?\b", request, re.I)
seed_match = re.search(r"\bseed(?:\s+of|\s*=|\s*:)?\s*[\"\u201c]?(\d{1,10})[\"\u201d]?", request, re.I)
sampler_match = re.search(
r"\b(?:on|with|using)\s+[\"\u201c]?(DDIM|DDPM|Heun|Euler(?:\s+a)?|DPM\+\+\s*2M)[\"\u201d]?\s+sampler\b"
r"|\bsampler(?:\s+of|\s*=|\s*:)?\s+[\"\u201c]?(DDIM|DDPM|Heun|Euler(?:\s+a)?|DPM\+\+\s*2M)",
request,
re.I,
)
aspect_match = re.search(r"\b(?:aspect\s+ratio(?:\s+of)?|ratio)\s*[\"\u201c]?(\d+\s*:\s*\d+)", request, re.I)
cfg_match = re.search(r"\bCFG(?:\s+scale)?(?:\s+of|\s*=|\s*:)?\s*[\"\u201c]?(\d+(?:\.\d+)?)", request, re.I)
lora_strength_match = re.search(r"\bLoRA\s+strength(?:\s+of|\s*=|\s*:)?\s*[\"\u201c]?(\d+(?:\.\d+)?)", request, re.I)
denoise_match = re.search(r"\bdenoise(?:\s+strength)?(?:\s+of|\s*=|\s*:)?\s*[\"\u201c]?(\d+(?:\.\d+)?)", request, re.I)
reference_strength_match = re.search(r"\breference\s+strength(?:\s+of|\s*=|\s*:)?\s*[\"\u201c]?(\d{1,3})\s*%?", request, re.I)
sampler_value = (sampler_match.group(1) or sampler_match.group(2)) if sampler_match else ""
sampler = sampler_value.upper() if sampler_value else ""
sampler = {"EULER A": "Euler a", "EULER": "Euler", "HEUN": "Heun", "DPM++ 2M": "DPM++ 2M"}.get(sampler, sampler)
aspect_ratio = re.sub(r"\s+", "", aspect_match.group(1)) if aspect_match else ""
# If the only DDPM/DDIM mention labels the sampler, do not let it choose a provider.
if not provider_match and re.search(r"\b(?:DDPM|DDIM)\b\s+sampler", request, re.I):
provider_hint = ""
return ChatGenerationRequest(
prompt=prompt,
subject=subject,
provider_hint=provider_hint,
model_query=model_query,
base_model_query=base_model_query,
negative_prompt=negative_prompt,
image_count=int(count_match.group(1)) if count_match else None,
steps=int(steps_match.group(1)) if steps_match else None,
sampler=sampler,
aspect_ratio=aspect_ratio,
seed=int(seed_match.group(1)) if seed_match else None,
cfg_scale=float(cfg_match.group(1)) if cfg_match else None,
lora_strength=float(lora_strength_match.group(1)) if lora_strength_match else None,
denoise_strength=float(denoise_match.group(1)) if denoise_match else None,
reference_strength=int(reference_strength_match.group(1)) if reference_strength_match else None,
has_positive_prompt=positive_match is not None,
)
def generation_output_folder(root: Path, provider_id: str, model_name: str) -> Path:
"""Return the browseable folder shared by all output from one model."""
def safe(value: str, fallback: str) -> str:
value = re.sub(r"[<>:\"/\\\\|?*\x00-\x1f]+", " ", value.strip())
return re.sub(r"\s+", " ", value).strip(" .")[:96] or fallback
folder = root.resolve() / "data" / "generations" / safe(provider_id, "generator") / safe(model_name, "model")
folder.mkdir(parents=True, exist_ok=True)
return folder
def generation_metadata_path(folder: Path, timestamp: str, job_id: str) -> Path:
return folder / f"generation_{timestamp}_{job_id}.json"
@dataclass(frozen=True, slots=True)
class GenerationRecord:
folder: Path
images: tuple[Path, ...]
provider_id: str
provider_name: str
model_name: str
model_path: str
prompt: str
seed: int
steps: int
sampler: str
aspect_ratio: str
created_at: str
@classmethod
def from_metadata(cls, metadata_path: Path) -> "GenerationRecord | None":
try:
payload = json.loads(metadata_path.read_text(encoding="utf-8"))
except (OSError, ValueError, TypeError, json.JSONDecodeError):
return None
folder = metadata_path.parent
listed_images = payload.get("images", [])
images = tuple(Path(str(path)) for path in listed_images if Path(str(path)).is_file())
if not images: # Support generation folders created by older ADAM versions.
images = tuple(path for path in sorted(folder.iterdir()) if path.is_file() and path.suffix.casefold() in IMAGE_EXTENSIONS)
if not images:
return None
return cls(
folder=folder,
images=images,
provider_id=str(payload.get("provider_id", "")),
provider_name=str(payload.get("provider_name", "Unknown generator")),
model_name=str(payload.get("model_name", folder.name)),
model_path=str(payload.get("model_path", "")),
prompt=str(payload.get("prompt", "")),
seed=int(payload.get("seed", 0) or 0),
steps=int(payload.get("steps", 0) or 0),
sampler=str(payload.get("sampler", "")),
aspect_ratio=str(payload.get("aspect_ratio", "")),
created_at=str(payload.get("created_at", "")),
)
def generation_tools(registry: ToolRegistry) -> list[ToolSpec]:
return [
tool
for tool in registry.enabled()
if "image_generation" in tool.capabilities
]
def load_generation_history(root: Path, *, limit: int = 200) -> list[GenerationRecord]:
history_root = root.resolve() / "data" / "generations"
if not history_root.is_dir():
return []
records = [
record
for metadata_path in history_root.rglob("generation*.json")
for record in [GenerationRecord.from_metadata(metadata_path)]
if record is not None
]
records.sort(key=lambda item: item.created_at or item.folder.name, reverse=True)
return records[: max(1, int(limit))]
def build_generation_plan(
tool: ToolSpec,
*,
model_name: str,
model_path: str,
prompt: str,
image_count: int,
steps: int,
seed: int,
sampler: str,
aspect_ratio: str,
extra_arguments: dict[str, Any] | None = None,
) -> ExecutionPlan:
if "image_generation" not in tool.capabilities:
raise ValueError(f"{tool.name} is not registered for image generation.")
safe_name = model_name.strip() or Path(model_path).name
arguments: dict[str, Any] = {
"model_name": safe_name,
"model_path": model_path,
"prompt": prompt.strip(),
"image_count": int(image_count),
"steps": int(steps),
"seed": int(seed),
"sampler": sampler,
"aspect_ratio": aspect_ratio,
}
if extra_arguments:
arguments.update(extra_arguments)
return ExecutionPlan(
request=f"Generate {image_count} image(s) with {safe_name}",
summary=f"Generate {image_count} image(s) using {tool.name} and {safe_name}.",
steps=[
PlanStep(
tool_id=tool.id,
title="Generate images",
description=f"Create a reproducible image batch with {safe_name}.",
arguments=arguments,
)
],
requires_confirmation=tool.requires_confirmation,
confirmation_reason=(
"This generator is configured to require approval before it runs."
if tool.requires_confirmation
else ""
),
project_name=f"{safe_name} generation",
)
def combine_generation_plans(
plans: list[ExecutionPlan],
*,
display_seconds: int = 5,
show_labels: bool = True,
loop: bool = False,
) -> ExecutionPlan:
"""Combine per-model image plans into one sequential presentation cycle."""
usable = [plan for plan in plans if plan.steps]
if not usable:
raise ValueError("A generation cycle needs at least one model.")
model_names = [
str(plan.steps[0].arguments.get("model_name", plan.project_name))
for plan in usable
]
cycle_settings = {
"display_seconds": max(1, int(display_seconds)),
"show_labels": bool(show_labels),
"loop": bool(loop),
"models": model_names,
}
reasons = [plan.confirmation_reason for plan in usable if plan.confirmation_reason]
return ExecutionPlan(
request=f"Generate a presentation cycle with {len(usable)} models.",
summary=(
f"Generate images sequentially with {len(usable)} models, then play them "
f"for {cycle_settings['display_seconds']} seconds each"
+ (" with model labels." if show_labels else ".")
),
steps=[step for plan in usable for step in plan.steps],
requires_confirmation=any(plan.requires_confirmation for plan in usable),
confirmation_reason="; ".join(dict.fromkeys(reasons)),
project_name="Generation Cycle",
)