File size: 14,501 Bytes
e0265b9
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
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",
    )