SyntheticMDProductions commited on
Commit
f8c73f9
·
verified ·
1 Parent(s): c61c435

ADAM October 2026 source release: PixelRow, INRFlow, Wan Video, Oasis player and field guide

Browse files

Update the September 5 source snapshot with current desktop workflows and plugins. Add clean default configuration, fresh screenshots, public architecture map, interactive guide, source ZIP, release notes and validation (322 passing tests across 33 isolated modules). Personal data, credentials, checkpoints, external trainer folders and build outputs are excluded.

This view is limited to 50 files because it contains too many changes.   See raw diff
Files changed (50) hide show
  1. .gitattributes +1 -0
  2. ADAM-source-2026-10-01.zip +3 -0
  3. ADAM.spec +1 -0
  4. README.md +591 -394
  5. SHA256SUMS.txt +1 -0
  6. adam/assets.py +72 -20
  7. adam/auto_training.py +109 -0
  8. adam/cnn_reviewer.py +172 -0
  9. adam/commands.py +1 -1
  10. adam/config.py +3 -0
  11. adam/dataset_registry.py +4 -3
  12. adam/generations.py +334 -14
  13. adam/intelligence.py +200 -0
  14. adam/job_manager.py +18 -3
  15. adam/model_plugins.py +3 -1
  16. adam/model_plugins_builtin/ddpm/manifest.py +3 -1
  17. adam/model_plugins_builtin/flow_matching/manifest.py +2 -0
  18. adam/model_plugins_builtin/inrflow/__init__.py +2 -0
  19. adam/model_plugins_builtin/inrflow/common.py +106 -0
  20. adam/model_plugins_builtin/inrflow/generator.py +305 -0
  21. adam/model_plugins_builtin/inrflow/manifest.py +394 -0
  22. adam/model_plugins_builtin/inrflow/model.py +400 -0
  23. adam/model_plugins_builtin/inrflow/trainer.py +523 -0
  24. adam/model_plugins_builtin/oasis/manifest.py +24 -9
  25. adam/model_plugins_builtin/pixelrow/__init__.py +2 -0
  26. adam/model_plugins_builtin/pixelrow/common.py +86 -0
  27. adam/model_plugins_builtin/pixelrow/generator.py +212 -0
  28. adam/model_plugins_builtin/pixelrow/manifest.py +301 -0
  29. adam/model_plugins_builtin/pixelrow/model.py +244 -0
  30. adam/model_plugins_builtin/pixelrow/trainer.py +411 -0
  31. adam/model_plugins_builtin/sdxl_lora/manifest.py +1 -1
  32. adam/model_plugins_builtin/wan_video/__init__.py +1 -0
  33. adam/model_plugins_builtin/wan_video/manifest.py +50 -0
  34. adam/nova.py +2 -0
  35. adam/oasis_dataset.py +94 -15
  36. adam/oasis_player.py +112 -0
  37. adam/ollama.py +56 -7
  38. adam/orion.py +85 -0
  39. adam/planner.py +254 -21
  40. adam/progressive_training.py +103 -0
  41. adam/recommendations.py +113 -15
  42. adam/remote_access.py +25 -4
  43. adam/remote_dashboard.py +22 -13
  44. adam/remote_v1.py +53 -35
  45. adam/tool_folders.py +5 -0
  46. adam/tools/ddpm_adapter.py +244 -28
  47. adam/tools/flow_adapter.py +119 -4
  48. adam/tools/flow_generator.py +19 -0
  49. adam/tools/lora_adapter.py +8 -1
  50. adam/tools/lora_generator.py +115 -16
.gitattributes CHANGED
@@ -36,3 +36,4 @@ saved_model/**/* filter=lfs diff=lfs merge=lfs -text
36
  assets/adam_atom.ico filter=lfs diff=lfs merge=lfs -text
37
  assets/adam_atom.png filter=lfs diff=lfs merge=lfs -text
38
  docs/screenshots/command-center.png filter=lfs diff=lfs merge=lfs -text
 
 
36
  assets/adam_atom.ico filter=lfs diff=lfs merge=lfs -text
37
  assets/adam_atom.png filter=lfs diff=lfs merge=lfs -text
38
  docs/screenshots/command-center.png filter=lfs diff=lfs merge=lfs -text
39
+ docs/field-guide/adam-map.png filter=lfs diff=lfs merge=lfs -text
ADAM-source-2026-10-01.zip ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:f0e60f14332f6b6c5445871e81133162df75d98f3113f8c9c43d6b2dfd8a5d29
3
+ size 1726930
ADAM.spec CHANGED
@@ -4,6 +4,7 @@ from PyInstaller.utils.hooks import collect_submodules
4
 
5
  hiddenimports = (
6
  collect_submodules("adam.tools")
 
7
  + collect_submodules("transformers.models.dinov2")
8
  + ["transformers", "torch", "PIL"]
9
  )
 
4
 
5
  hiddenimports = (
6
  collect_submodules("adam.tools")
7
+ + collect_submodules("adam.model_plugins_builtin")
8
  + collect_submodules("transformers.models.dinov2")
9
  + ["transformers", "torch", "PIL"]
10
  )
README.md CHANGED
@@ -1,394 +1,591 @@
1
- ---
2
- license: mit
3
- tags:
4
- - desktop-application
5
- - ai-tools
6
- - dataset-management
7
- - lora-training
8
- - windows
9
- ---
10
-
11
- # ADAM — AI Development and Automation Manager
12
-
13
- ADAM is a local, safety-first desktop hub for orchestrating AI project tools.
14
- It includes registered dataset, DDPM, and SDXL LoRA workflows with background
15
- planning, approval gates, progress reporting, and persistent asset history.
16
-
17
- ![ADAM command center](docs/screenshots/command-center.png)
18
-
19
- Dataset preparation, captioning, and preview placeholders remain clearly marked
20
- as demo tools. The connected Dataset Collector, DDPM trainer, and Local SDXL
21
- LoRA Trainer use real adapters and never fall back to simulated training.
22
-
23
- Existing program folders can be connected from **Settings → Tool folders**.
24
- ADAM stores only the path and scans for likely entry points; it does not copy or
25
- modify the external project. Folder assignments can also be pasted into chat:
26
-
27
- ```text
28
- DDPM Trainer: D:\AI\DDPM
29
- Flow Matching Trainer: D:\AI\FlowMatchImageGenerator
30
- ```
31
-
32
- On a new computer, install `requirements.txt` in your chosen Python environment
33
- before running `Launch ADAM.bat`. The launcher checks desktop dependencies and
34
- does not install packages automatically or depend on the developer's personal
35
- trainer folders. Install each optional trainer's dependencies according to that
36
- tool's setup instructions before using its ADAM workflow.
37
-
38
- Remote access is disabled by default. Devices with an access token can browse
39
- datasets, edit captions/review marks and submit work. Only enable it for trusted
40
- devices. The desktop **Allow remote job controls and approval changes** setting
41
- also permits remote confirmation, stopping and retrying jobs. A remote browser
42
- can enable training auto-approval only after that desktop permission is granted;
43
- it can always turn auto-approval off. Saved token changes and disabled access
44
- take effect for new requests without restarting the server.
45
-
46
- Use private Tailscale access for connections beyond a trusted local network;
47
- the built-in HTTP listener does not provide transport encryption by itself.
48
- Phone URLs and QR codes contain the access token and should be treated as
49
- credentials. Remote commands require JSON, have bounded request sizes and
50
- connection counts, and reject cross-site browser submissions. These controls
51
- do not sandbox installed Python plugins or connected trainers: install only
52
- code you trust.
53
-
54
- Detection does not automatically authorize training. A real training adapter
55
- remains gated until its dataset, model name, run settings, and output location
56
- are explicit.
57
-
58
- ## Training agents
59
-
60
- ![ADAM model creation settings](docs/screenshots/create-model.png)
61
-
62
- ADAM's training lifecycle is divided into four explainable responsibilities:
63
-
64
- - **EVE** reviews dataset membership and leaves uncertain images for the user.
65
- - **ORION** reviews planned epochs, batch size, resolution, image exposures, and
66
- estimated optimizer steps. He can require approval but never silently changes
67
- the requested settings. In the Model Creation Assistant, **ORION: apply a
68
- starting recipe** fills a conservative, editable draft from the image count
69
- and selected resolution before a plan is built.
70
- - **ATLAS** watches active training for non-finite loss, sustained critical GPU
71
- temperature, critically low disk space, stalls, and large runtime overruns.
72
- Critical conditions pause the trainer process tree so the user can inspect it.
73
- - **NOVA** examines available post-training previews and samples for unreadable
74
- files and exact-looking duplicate collapse. Her report explicitly separates
75
- technical sample health from subjective or subject-quality review.
76
-
77
- ORION, ATLAS, and NOVA reports are stored with each durable job record and are
78
- shown in Current Plan, Active Job, and Jobs / History respectively. ATLAS's
79
- default thresholds can be overridden in `config/settings.json` with the
80
- `atlas_*` settings defined in `adam/config.py`.
81
-
82
- Every new job passes through the shared preflight and ORION review before its
83
- queue state is chosen. Desktop plans, Remote prompts, and Remote training forms
84
- use the same review. Remote training auto-approval still applies to ordinary
85
- plans, but an ORION warning leaves the job awaiting explicit approval. Reviewing
86
- a plan does not change the requested training settings.
87
-
88
- ## Real image collection
89
-
90
- When a valid Dataset Collector folder is connected, the `dataset_collector`
91
- registry entry uses ADAM's real visible-browser adapter. After plan approval it:
92
-
93
- - opens Bing Images in a normal visible Chrome window;
94
- - waits when consent/CAPTCHA/human-verification text is detected;
95
- - resumes automatically after the user resolves the page;
96
- - downloads valid images at least 256×256;
97
- - removes exact duplicate downloads;
98
- - writes a matching `.txt` caption beside every image; and
99
- - records URLs, captions, sources, and dimensions in `metadata.csv`.
100
-
101
- No CAPTCHA or website restriction is bypassed. Closing Chrome or stopping the
102
- job ends collection safely. A new timestamped dataset folder is used rather
103
- than overwriting an existing collection.
104
-
105
- ADAM keeps an incomplete DDPM request in conversation memory. A follow-up such
106
- as `dataset folder Mario, model name Mario V2, epoch count 100, output D:\Runs`
107
- fills the pending fields and validates named datasets against the connected
108
- collector. It will not start if the dataset cannot be found.
109
-
110
- ## Showcase videos
111
-
112
- The **Showcase Video** workspace creates a finished MP4 directly from completed
113
- DDPM and Flow Matching models. Select and reorder the models, choose 12–24
114
- images per model, a 3-, 4-, or 5-second image duration, shared steps and aspect
115
- ratio, provider-compatible samplers, seed, and 720p or 1080p output. ADAM runs
116
- the image batches sequentially and then renders a request-list interface that
117
- tracks the active model, image number, trainer, steps, sampler, and aspect ratio.
118
- LoRA models are intentionally excluded from this streamlined workflow.
119
-
120
- When Ollama is reachable, messages that are not workflow commands receive a
121
- short conversational answer. Ollama may explain or plan, but it still cannot
122
- bypass the registry or confirmation gates.
123
-
124
- ## Web search in Chat Mode
125
-
126
- Chat Mode can give local Ollama current web context without an API key. Enable
127
- it in **Settings → Planning model**, then ask naturally, for example:
128
-
129
- ```text
130
- Search the web for Dandy's World character ideas.
131
- What are the latest Ollama release notes?
132
- Look up a reference for a cyberpunk city character.
133
- ```
134
-
135
- ADAM sends only that search query to Bing's public results feed, reads the
136
- result titles and snippets,
137
- and passes up to five titles, snippets, and links to Ollama. It does not open
138
- the result pages, download anything, or let web content run tools. Results are
139
- untrusted reference material, so ADAM is instructed to cite the links and flag
140
- uncertainty. Disable the setting to keep Chat Mode fully local.
141
-
142
- When you explicitly ask ADAM to **read**, **open**, or **research** result links,
143
- it can read up to three public HTML/text pages and give Ollama short extracts.
144
- For example: `Search the web for Undertale character ideas and read the most
145
- relevant links.` Direct links can be read with `Read https://example.com/ and
146
- summarize it.` Private/local addresses, non-web protocols, oversized pages,
147
- downloads, and more than three pages are blocked. This control can be disabled
148
- in Settings.
149
-
150
- Planning runs away from the interface thread, and conversational Ollama output
151
- is streamed into the chat. ADAM validates training commands against a strict
152
- schema and each registered trainer's declared capabilities before offering a
153
- job.
154
-
155
- In **Settings → Planning model**, **Chat response length** sets the maximum
156
- number of generated tokens for a Chat Mode reply. Higher values allow longer
157
- research summaries but use more time and GPU memory. The default is 1,024,
158
- which gives Qwen3 enough room to reason and still produce a visible response.
159
-
160
- ADAM stores friendly dataset/model names, paths, trainer types, epochs, and
161
- resume checkpoints in `data/assets.json`. Requests such as:
162
-
163
- ```text
164
- From the Mario dataset, train it on a DDPM for 300 epochs.
165
- With the Mario dataset, train it on a LoRA for 100 epochs.
166
- Continue the Mario model from the DDPM for 50 epochs.
167
- ```
168
-
169
- are resolved to real paths before approval. Continuation is offered only when a
170
- compatible checkpoint exists. New DDPM runs retain the latest resume checkpoint.
171
-
172
- ## Run
173
-
174
- ```powershell
175
- python main.py
176
- ```
177
-
178
- On Windows, you can also double-click `Launch ADAM.bat`.
179
-
180
- The app requires Python 3.10+ and PySide6. Optional integrations use `psutil`
181
- for system information and `pynvml` for NVIDIA GPU information.
182
-
183
- ```powershell
184
- python -m pip install -r requirements.txt
185
- ```
186
-
187
- Try:
188
-
189
- - Click **Create a model…** in Trainer Mode for the guided Model Creation Assistant.
190
- - `Adam, train a LoRA of Hatsune Miku`
191
- - `Adam, collect a dataset of liminal spaces`
192
- - `Adam, generate previews`
193
- - `Adam, check GPU status`
194
- - `From the Mario dataset, train it on a DDPM for 300 epochs`
195
- - `With the Mario dataset, train it on a LoRA for 100 epochs`
196
-
197
- Training and large collection plans are never started until you approve the
198
- plan. All actions are recorded in `logs/adam.log`, while project artifacts live
199
- under `data/projects/`.
200
-
201
- The Model Creation Assistant can start from a built-in Character LoRA, Style
202
- LoRA, DDPM, or Flow Matching preset. It can create a dataset or select a
203
- registered one, recommends starting values, and saves personal presets. The
204
- result still goes through ADAM's normal validated planner and approval gate.
205
- Use **+ Add model** to build a multi-model training batch. Each wide model tab
206
- keeps its own dataset, trainer, name, and settings; the minus button removes an
207
- unwanted model, and tabs can be dragged to change the run order. ADAM validates
208
- all models, presents one combined approval plan, and runs them sequentially so
209
- only one training workflow uses the GPU at a time. A failed step stops the batch
210
- before a later model starts.
211
- Before approval, ADAM adds checks for connected tools, dataset contents, the
212
- LoRA base model, and output-drive free space. Completed dataset and training
213
- jobs also include a suggested next step.
214
-
215
- ### Model Batch Builder
216
-
217
- Use **Create model batch…** to paste one requested subject per line. ADAM turns
218
- the list into editable model tabs, removes duplicate names, and lets the current
219
- trainer recipe be applied to any multi-selection of models. The batch is saved
220
- as a draft so it can be closed and resumed later.
221
-
222
- For a review-first workflow, choose **Collect missing datasets first**. This
223
- queues only sequential dataset collection and leaves training in the saved
224
- draft. After collection, reopen the draft, use **Find collected datasets**, and
225
- review each dataset in Training Studio. **Exclude rejected** moves rejected
226
- images out of the training folder into a recoverable quarantine, and **Restore
227
- excluded** reverses it. **Keep all images** marks the whole selected dataset as
228
- accepted in one action, after which individual bad images can still be rejected.
229
- Training remains locked until each model is explicitly
230
- marked as reviewed and ready. If every linked dataset is acceptable as-is,
231
- **Approve all datasets** marks the entire batch ready after one confirmation;
232
- it does not inspect individual images or apply pending rejection decisions.
233
-
234
- Completed Flow Matching models can be selected in **Fine-tune**. ADAM uses the
235
- saved Flow model folder as the continuation source, locks the continuation to
236
- the model's original resolution, and writes the fine-tuned result to a new
237
- output folder. This continues the saved weights while starting a fresh optimizer
238
- and learning-rate schedule; it does not overwrite the original model.
239
-
240
- ## Training Studio
241
-
242
- The **Training Studio** turns completed work into a reviewable experiment loop:
243
-
244
- - **Datasets** provides an image gallery, keep/reject decisions, caption editing,
245
- exact duplicate detection, and visually similar duplicate candidates.
246
- - **Experiments** compares job settings and outcomes, opens outputs, marks a
247
- preferred model, and converts successful settings into reusable recipes.
248
- - **Checkpoint Lab** browses model checkpoints and output images, records
249
- consistent prompt/seed evaluations, and sends preview requests through the
250
- normal approval-aware planner.
251
- - **Recipes** preserves training starting points and can import or export
252
- portable JSON recipe files.
253
-
254
- ### EVE AI Dataset Review
255
-
256
- In Training Studio → Datasets, **EVE AI Review…** performs a local reference-
257
- guided visual review. Add one or more good reference images and optional bad
258
- references, then choose Keep and Reject confidence thresholds. EVE uses a small
259
- DINOv2 vision model to divide the selected dataset into **Keep**, **Reject**, and
260
- **Uncertain** galleries with confidence scores. The model is downloaded once on
261
- first use and subsequent analysis stays local.
262
-
263
- Nothing is applied automatically. Inspect both sides, double-click images for a
264
- full view, and move selected results between the three groups before choosing
265
- **Apply EVE review**. EVE's decisions remain ordinary Training Studio review
266
- marks: they can be manually changed, and rejected files are not moved until
267
- **Exclude rejected** is selected. The latest proposal is also saved under
268
- `data/eve_reviews/` for auditing. Use **Select all in current group** (or
269
- Ctrl/Shift selection) to move many images at once; EVE transfers only the
270
- chosen thumbnails so manual sorting stays responsive on large datasets.
271
-
272
- Training panels show elapsed time, a progress-based ETA, recent logs, and a
273
- loss sparkline when the connected trainer reports `loss`. Preflight summaries
274
- include clearly labelled workload, duration, VRAM, and disk estimates. These
275
- estimates are planning hints rather than hardware guarantees.
276
-
277
- Create a Model also supports live training previews with a configurable
278
- epoch interval, prompt, and reproducible seed for each model tab. While a
279
- training job is active, its newest 256×256 preview appears in the right sidebar
280
- with the source epoch and next scheduled preview. The full-size trainer output
281
- can be opened from the card. Built-in adapters may publish previews directly;
282
- registered DDPM, Flow, LoRA, APVD, MaskGit, and other trainers can also
283
- participate by writing conventionally named `preview`, `sample`, or `epoch`
284
- images beneath their declared output folder.
285
-
286
- ## Generations
287
-
288
- The **Generations** workspace runs compatible registered image generators
289
- without opening their separate desktop interfaces. The connected DDPM and Flow
290
- Matching projects can generate from completed models with a reproducible seed,
291
- sampler or ODE method, step count, image count, and aspect ratio. Generation
292
- work uses the normal ADAM job queue, progress reporting, cancellation, and
293
- logging.
294
-
295
- Every completed batch is stored under `data/generations/` with its images and a
296
- `generation.json` sidecar. The history gallery can open an image or batch folder
297
- and restore the exact settings for another run. DDPM creative notes are stored
298
- with a batch for organization; they are not presented as text conditioning for
299
- an unconditional DDPM model.
300
-
301
- Generation history opens as automatic model folders. Each folder uses the
302
- registered model name and a recent generated image as its cover. Double-clicking
303
- a folder filters history to that model and selects its provider and model in the
304
- generation controls, so the next batch is generated into the same existing model
305
- directory. This view does not move or rewrite older generation files.
306
-
307
- **Generation Cycle…** selects multiple compatible completed models and queues
308
- one generation step per model. Choose images per model, a shared prompt or
309
- creative note, starting seed, slideshow duration, looping, fullscreen playback,
310
- and an optional model/trainer label. When the cycle finishes, ADAM opens the
311
- results as a local slideshow while preserving every ordinary generation record
312
- in history.
313
-
314
- If ADAM discovers a job interrupted by an unexpected shutdown, it offers to
315
- open Jobs & History. The previous record remains intact and can be retried as a
316
- new approval-gated job. Job logs can also be exported for troubleshooting.
317
-
318
- ## Connect an existing tool
319
-
320
- ADAM supports importable Python functions and command-line Python scripts.
321
- For a no-code setup, open **Settings → External Tools → Add external tool**.
322
- Choose the program folder, select its training entry script and important
323
- configuration files, then review ADAM's static compatibility and safety report.
324
- The report covers:
325
-
326
- - detected command-line options and required inputs;
327
- - likely dataset formats;
328
- - output and checkpoint behavior;
329
- - progress reporting;
330
- - resume-training support; and
331
- - potentially risky operations visible in the selected entry script.
332
-
333
- The 1–10 rating measures how clearly the script fits ADAM's safe command-line
334
- contract. It is not a guarantee that third-party code is harmless. ADAM does
335
- not execute a script while scanning it, external tools cannot replace built-in
336
- registry entries, and every external-tool run requires explicit approval.
337
-
338
- After registration, a tool can be planned with a request such as:
339
-
340
- ```text
341
- Run APVD Model Trainer with dataset=D:\DreamData, epochs=20, output=D:\APVD\output
342
- ```
343
-
344
- ADAM will ask for any required inputs that were omitted before it offers the
345
- approval plan.
346
-
347
- For manual registry configuration, edit the relevant item in
348
- `config/tools.json`:
349
-
350
- ```json
351
- {
352
- "backend": {
353
- "type": "python",
354
- "module": "my_tools.lora",
355
- "function": "train"
356
- },
357
- "demo": false
358
- }
359
- ```
360
-
361
- The function receives a `ToolContext` as its first argument and keyword
362
- arguments from the approved plan. This keeps training code in one place: your
363
- existing GUI and ADAM can both call the same backend.
364
-
365
- For scripts:
366
-
367
- ```json
368
- {
369
- "backend": {
370
- "type": "script",
371
- "path": "D:/AI/LoRATrainer/train.py"
372
- },
373
- "demo": false
374
- }
375
- ```
376
-
377
- ADAM invokes scripts directly with the current Python interpreter, captures
378
- stdout/stderr, and never drives another GUI with mouse clicks.
379
-
380
- ## Safety model
381
-
382
- - Plans are shown before execution.
383
- - Long, destructive, or high-volume work requires confirmation.
384
- - Unregistered tools cannot be invoked.
385
- - External paths and arguments are validated before execution.
386
- - The LLM may propose a plan, but only registered tools can execute it.
387
- - Pause, resume, and cancel controls are available for active jobs.
388
- - Every tool action and state transition is logged.
389
-
390
- ## Tests
391
-
392
- ```powershell
393
- python -m pytest -q
394
- ```
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ ---
2
+ license: mit
3
+ tags:
4
+ - desktop-application
5
+ - ai-tools
6
+ - dataset-management
7
+ - lora-training
8
+ - windows
9
+ - ai-workflow
10
+ - experiment-tracking
11
+ - model-plugins
12
+ - video-lora
13
+ ---
14
+
15
+ # ADAM — AI Development and Automation Manager
16
+
17
+ ADAM is a local desktop command center for AI experiments: describe an idea,
18
+ prepare and review a dataset, approve a training plan, generate samples, and
19
+ use saved experiment history to decide what to try next.
20
+
21
+ It began as a **Jarvis-inspired assistant for AI workflows** and has grown into
22
+ a Windows/PySide6 application with registered tools, model plugins, background
23
+ jobs, live previews, and optional local Ollama chat and Remote access.
24
+
25
+ **Source release: 1 October 2026 · 322 tests passed across 33 modules.** This repository contains the desktop
26
+ application and its source code. Model weights, personal datasets, saved jobs,
27
+ and connected external trainer projects are supplied separately by the user.
28
+
29
+ ![ADAM workflow and architecture map](docs/field-guide/adam-map.png)
30
+
31
+ **[Explore the interactive ADAM field guide](https://huggingface.co/spaces/SyntheticMDProductions/ADAM-Field-Guide)**
32
+ — click the workflow stages, reviewers and model architectures, then explore
33
+ four animated explanations of image construction. You can also
34
+ [download the standalone HTML guide](https://huggingface.co/SyntheticMDProductions/AI_Development_Automation_Manager/resolve/main/docs/field-guide/adam-field-guide.html?download=true)
35
+ and [the visual map](docs/field-guide/adam-map.png).
36
+ The animations explain mechanisms; they are not model inference or quality benchmarks.
37
+
38
+ ## What ADAM includes
39
+
40
+ | Workspace | What it does |
41
+ | --- | --- |
42
+ | Command center | Local Ollama chat, validated workflow planning, guided model creation, and sequential model batches |
43
+ | Training Studio / Dataset Lab | Dataset browsing, captions, keep/reject decisions, EVE proposals, recoverable exclusions, checkpoint review, and recipes |
44
+ | Jobs / History | Approval, scheduling, progress, previews, pause/cancel, retry, and persistent records |
45
+ | Generations / Showcase Video | Reproducible image batches, saved settings, generation history, and DDPM/Flow showcase MP4s |
46
+ | Video LoRA | Wan 2.1 T2V 1.3B clip preparation, character-reference suggestions, reviewed caption drafts, training, and video generation |
47
+ | Model Intelligence / Experiments | Training and generation evidence, run comparison, and follow-up experiment suggestions |
48
+ | Oasis player | Playable inference for compatible action-conditioned world models |
49
+ | Tools / Remote / System | Declared plugin settings, connected tool folders, optional authenticated device access, and hardware telemetry |
50
+
51
+ | Model family | Integration |
52
+ | --- | --- |
53
+ | DDPM, regular Flow Matching, SDXL LoRA | Adapters for separately connected local trainers and generators |
54
+ | PixelRow | Experimental built-in model that generates top to bottom, one row at a time |
55
+ | INRFlow | Experimental built-in coordinate-to-RGB flow model without a pretrained image compressor |
56
+ | Neural Cellular Automata | Included experimental custom plugin that learns image growth from a living seed |
57
+ | Oasis | Experimental connected action world model with temporal latent and temporal pixel-flow workflows |
58
+ | Wan video LoRA | Dedicated workspace and adapter for a connected LoRAVideoTrainer/Musubi environment |
59
+
60
+ See [release notes](docs/releases/2026-10-01.md) for changes since the September 5 release,
61
+ and [Oasis](docs/oasis_integration.md) / [INRFlow](docs/inrflow_integration.md)
62
+ for integration details. Experimental architectures are intended for local
63
+ exploration; this release makes no benchmark-quality claims.
64
+
65
+ ## Quick start
66
+
67
+ [Download the complete source ZIP](https://huggingface.co/SyntheticMDProductions/AI_Development_Automation_Manager/resolve/main/ADAM-source-2026-10-01.zip?download=true),
68
+ extract it, or clone this repository, then run:
69
+
70
+ ```powershell
71
+ python -m pip install -r requirements.txt
72
+ python main.py
73
+ ```
74
+
75
+ Use Python 3.10+ on Windows. You can also double-click `Launch ADAM.bat` after
76
+ installing the requirements. Choose a CUDA-compatible PyTorch build for GPU
77
+ training according to your hardware. Optional trainer projects keep their own
78
+ dependencies and model weights; connect them in **Settings → Tool folders**.
79
+ Ollama is optional for local chat and model-assisted planning. Chrome is needed
80
+ for the visible image collector; FFmpeg is needed for video collection and splitting.
81
+
82
+ The release starts with empty tool connections and Remote access disabled.
83
+ Application state is saved locally under `data/`, `logs/`, and `config/settings.json`.
84
+
85
+ ![ADAM command center](docs/screenshots/command-center.png)
86
+
87
+ Dataset preparation, captioning, and preview placeholders remain clearly marked
88
+ as demo tools. The connected Dataset Collector, DDPM trainer, and Local SDXL
89
+ LoRA Trainer use real adapters and never fall back to simulated training.
90
+
91
+ Existing program folders can be connected from **Settings → Tool folders**.
92
+ ADAM stores only the path and scans for likely entry points; it does not copy or
93
+ modify the external project. Folder assignments can also be pasted into chat:
94
+
95
+ ```text
96
+ DDPM Trainer: D:\AI\DDPM
97
+ Flow Matching Trainer: D:\AI\FlowMatchImageGenerator
98
+ ```
99
+
100
+ On a new computer, install `requirements.txt` in your chosen Python environment
101
+ before running `Launch ADAM.bat`. The launcher checks desktop dependencies and
102
+ does not install packages automatically or depend on the developer's personal
103
+ trainer folders. Install each optional trainer's dependencies according to that
104
+ tool's setup instructions before using its ADAM workflow.
105
+
106
+ Remote access is disabled by default. Devices with an access token can browse
107
+ datasets, edit captions/review marks and submit work. Only enable it for trusted
108
+ devices. The desktop **Allow remote job controls and approval changes** setting
109
+ also permits remote confirmation, stopping and retrying jobs. A remote browser
110
+ can enable training auto-approval only after that desktop permission is granted;
111
+ it can always turn auto-approval off. Saved token changes and disabled access
112
+ take effect for new requests without restarting the server.
113
+
114
+ Use private Tailscale access for connections beyond a trusted local network;
115
+ the built-in HTTP listener does not provide transport encryption by itself.
116
+ Phone URLs and QR codes contain the access token and should be treated as
117
+ credentials. Remote commands require JSON, have bounded request sizes and
118
+ connection counts, and reject cross-site browser submissions. These controls
119
+ do not sandbox installed Python plugins or connected trainers: install only
120
+ code you trust.
121
+
122
+ Detection does not automatically authorize training. A real training adapter
123
+ remains gated until its dataset, model name, run settings, and output location
124
+ are explicit.
125
+
126
+ ## Training agents
127
+
128
+ ![ADAM model creation settings](docs/screenshots/create-model.png)
129
+
130
+ ADAM's training lifecycle is divided into four explainable responsibilities:
131
+
132
+ - **EVE** reviews dataset membership and leaves uncertain images for the user.
133
+ - **ORION** reviews planned epochs, batch size, resolution, image exposures, and
134
+ estimated optimizer steps. He can require approval but never silently changes
135
+ the requested settings. In the Model Creation Assistant, **ORION: apply a
136
+ starting recipe** fills a conservative, editable draft from the image count
137
+ and selected resolution before a plan is built.
138
+ - **ATLAS** watches active training for non-finite loss, sustained critical GPU
139
+ temperature, critically low disk space, stalls, and large runtime overruns.
140
+ Critical conditions pause the trainer process tree so the user can inspect it.
141
+ - **NOVA** examines available post-training previews and samples for unreadable
142
+ files and exact-looking duplicate collapse. Her report explicitly separates
143
+ technical sample health from subjective or subject-quality review.
144
+
145
+ ORION, ATLAS, and NOVA reports are stored with each durable job record and are
146
+ shown in Current Plan, Active Job, and Jobs / History respectively. ATLAS's
147
+ default thresholds can be overridden in `config/settings.json` with the
148
+ `atlas_*` settings defined in `adam/config.py`.
149
+
150
+ Every new job passes through the shared preflight and ORION review before its
151
+ queue state is chosen. Desktop plans, Remote prompts, and Remote training forms
152
+ use the same review. Remote training auto-approval still applies to ordinary
153
+ plans, but an ORION warning leaves the job awaiting explicit approval. Reviewing
154
+ a plan does not change the requested training settings.
155
+
156
+ ## Real image collection
157
+
158
+ When a valid Dataset Collector folder is connected, the `dataset_collector`
159
+ registry entry uses ADAM's real visible-browser adapter. After plan approval it:
160
+
161
+ - opens Bing Images in a normal visible Chrome window;
162
+ - waits when consent/CAPTCHA/human-verification text is detected;
163
+ - resumes automatically after the user resolves the page;
164
+ - downloads valid images at least 256×256;
165
+ - removes exact duplicate downloads;
166
+ - writes a matching `.txt` caption beside every image; and
167
+ - records URLs, captions, sources, and dimensions in `metadata.csv`.
168
+
169
+ No CAPTCHA or website restriction is bypassed. Closing Chrome or stopping the
170
+ job ends collection safely. A new timestamped dataset folder is used rather
171
+ than overwriting an existing collection.
172
+
173
+ ADAM keeps an incomplete DDPM request in conversation memory. A follow-up such
174
+ as `dataset folder Mario, model name Mario V2, epoch count 100, output D:\Runs`
175
+ fills the pending fields and validates named datasets against the connected
176
+ collector. It will not start if the dataset cannot be found.
177
+
178
+ ## Showcase videos
179
+
180
+ The **Showcase Video** workspace creates a finished MP4 directly from completed
181
+ DDPM and Flow Matching models. Select and reorder the models, choose 12–24
182
+ images per model, a 3-, 4-, or 5-second image duration, shared steps and aspect
183
+ ratio, provider-compatible samplers, seed, and 720p or 1080p output. ADAM runs
184
+ the image batches sequentially and then renders a request-list interface that
185
+ tracks the active model, image number, trainer, steps, sampler, and aspect ratio.
186
+ LoRA models are intentionally excluded from this streamlined workflow.
187
+
188
+ When Ollama is reachable, messages that are not workflow commands receive a
189
+ conversational answer. In Chat Mode, attach an image with the **+** button and
190
+ a vision-capable Ollama model can describe it, suggest a caption, or answer
191
+ questions about visible details. Images stay on the local Ollama connection.
192
+
193
+ In Trainer Mode, the planning model can propose a registered ADAM action when
194
+ the request is not recognized by the built-in planner. ADAM validates the
195
+ proposed tool and every setting against its registry, and confirmation gates
196
+ still apply. This option can be disabled in **Settings → Safety & Notifications**.
197
+
198
+ ## Web search in Chat Mode
199
+
200
+ Chat Mode can give local Ollama current web context without an API key. Enable
201
+ it in **Settings → Planning model**, then ask naturally, for example:
202
+
203
+ ```text
204
+ Search the web for Dandy's World character ideas.
205
+ What are the latest Ollama release notes?
206
+ Look up a reference for a cyberpunk city character.
207
+ ```
208
+
209
+ ADAM sends only that search query to Bing's public results feed, reads the
210
+ result titles and snippets,
211
+ and passes up to five titles, snippets, and links to Ollama. It does not open
212
+ the result pages, download anything, or let web content run tools. Results are
213
+ untrusted reference material, so ADAM is instructed to cite the links and flag
214
+ uncertainty. Disable the setting to keep Chat Mode fully local.
215
+
216
+ When you explicitly ask ADAM to **read**, **open**, or **research** result links,
217
+ it can read up to three public HTML/text pages and give Ollama short extracts.
218
+ For example: `Search the web for Undertale character ideas and read the most
219
+ relevant links.` Direct links can be read with `Read https://example.com/ and
220
+ summarize it.` Private/local addresses, non-web protocols, oversized pages,
221
+ downloads, and more than three pages are blocked. This control can be disabled
222
+ in Settings.
223
+
224
+ Planning runs away from the interface thread, and conversational Ollama output
225
+ is streamed into the chat. ADAM validates training commands against a strict
226
+ schema and each registered trainer's declared capabilities before offering a
227
+ job.
228
+
229
+ In **Settings → Planning model**, choose an automatic, short, balanced, or
230
+ detailed response style. Automatic uses a smaller response for simple questions
231
+ and makes more room for image reviews, explanations, and planning. **Maximum
232
+ response length** remains a hard limit for response time and GPU memory; the
233
+ default is 1,024 tokens.
234
+
235
+ ADAM stores friendly dataset/model names, paths, trainer types, epochs, and
236
+ resume checkpoints in `data/assets.json`. Requests such as:
237
+
238
+ ```text
239
+ From the Mario dataset, train it on a DDPM for 300 epochs.
240
+ With the Mario dataset, train it on a LoRA for 100 epochs.
241
+ Continue the Mario model from the DDPM for 50 epochs.
242
+ ```
243
+
244
+ are resolved to real paths before approval. Continuation is offered only when a
245
+ compatible checkpoint exists. New DDPM runs retain the latest resume checkpoint.
246
+
247
+ ## Run
248
+
249
+ ```powershell
250
+ python main.py
251
+ ```
252
+
253
+ On Windows, you can also double-click `Launch ADAM.bat`.
254
+
255
+ The app requires Python 3.10+ and PySide6. Optional integrations use `psutil`
256
+ for system information and `pynvml` for NVIDIA GPU information.
257
+
258
+ ```powershell
259
+ python -m pip install -r requirements.txt
260
+ ```
261
+
262
+ Try:
263
+
264
+ - Click **Create a model…** in Trainer Mode for the guided Model Creation Assistant.
265
+ - `Adam, train a LoRA of Hatsune Miku`
266
+ - `Adam, collect a dataset of liminal spaces`
267
+ - `Adam, generate previews`
268
+ - `Adam, check GPU status`
269
+ - `From the Mario dataset, train it on a DDPM for 300 epochs`
270
+ - `With the Mario dataset, train it on a LoRA for 100 epochs`
271
+
272
+ Training and large collection plans are never started until you approve the
273
+ plan. All actions are recorded in `logs/adam.log`, while project artifacts live
274
+ under `data/projects/`.
275
+
276
+ The Model Creation Assistant can start from a built-in Character LoRA, Style
277
+ LoRA, DDPM, Flow Matching, or experimental PixelRow preset. It can create a dataset or select a
278
+ registered one, recommends starting values, and saves personal presets. The
279
+ result still goes through ADAM's normal validated planner and approval gate.
280
+ Use **+ Add model** to build a multi-model training batch. Each wide model tab
281
+ keeps its own dataset, trainer, name, and settings; the minus button removes an
282
+ unwanted model, and tabs can be dragged to change the run order. ADAM validates
283
+ all models, presents one combined approval plan, and runs them sequentially so
284
+ only one training workflow uses the GPU at a time. A failed step stops the batch
285
+ before a later model starts.
286
+ Before approval, ADAM adds checks for connected tools, dataset contents, the
287
+ LoRA base model, and output-drive free space. Completed dataset and training
288
+ jobs also include a suggested next step.
289
+
290
+ ### Model Batch Builder
291
+
292
+ Use **Create model batch…** to paste one requested subject per line. ADAM turns
293
+ the list into editable model tabs, removes duplicate names, and lets the current
294
+ trainer recipe be applied to any multi-selection of models. The batch is saved
295
+ as a draft so it can be closed and resumed later.
296
+
297
+ For a review-first workflow, choose **Collect missing datasets first**. This
298
+ queues only sequential dataset collection and leaves training in the saved
299
+ draft. After collection, reopen the draft, use **Find collected datasets**, and
300
+ review each dataset in Training Studio. **Exclude rejected** moves rejected
301
+ images out of the training folder into a recoverable quarantine, and **Restore
302
+ excluded** reverses it. **Keep all images** marks the whole selected dataset as
303
+ accepted in one action, after which individual bad images can still be rejected.
304
+ Training remains locked until each model is explicitly
305
+ marked as reviewed and ready. If every linked dataset is acceptable as-is,
306
+ **Approve all datasets** marks the entire batch ready after one confirmation;
307
+ it does not inspect individual images or apply pending rejection decisions.
308
+
309
+ Completed Flow Matching models can be selected in **Fine-tune**. ADAM uses the
310
+ saved Flow model folder as the continuation source, locks the continuation to
311
+ the model's original resolution, and writes the fine-tuned result to a new
312
+ output folder. This continues the saved weights while starting a fresh optimizer
313
+ and learning-rate schedule; it does not overwrite the original model.
314
+
315
+ ## PixelRow
316
+
317
+ PixelRow is ADAM's experimental top-to-bottom image architecture. It trains on
318
+ ordinary image folders and predicts one complete quantized RGB row from all
319
+ previous rows, without a diffusion noise schedule. Start with 64×64 images for
320
+ the first experiment; 128×128 is available but trains more slowly.
321
+
322
+ PixelRow generation uses one step per image row. The Generations page labels
323
+ these as **Rows** and offers creativity, top-color-choice, and row-frame
324
+ settings. Enabling **Save row-build frames** writes a PNG sequence beneath the
325
+ generation folder, making the construction process ready for a video or visual
326
+ comparison. Seeds reproduce both the finished image and its intermediate rows.
327
+
328
+ ## Wan Video LoRA
329
+
330
+ Open **Video LoRA** in the sidebar to use the connected LoRAVideoTrainer from
331
+ inside ADAM. This workspace is separate from the image-model creation assistant
332
+ and image Generations page. It targets **Wan 2.1 T2V 1.3B** only.
333
+
334
+ Connect **Settings → Tool folders → Wan Video LoRA Trainer** to your existing
335
+ LoRAVideoTrainer folder. An existing `external_loravideotrainer` connection is
336
+ recognized automatically. ADAM runs its `.venv/Scripts/python.exe` and installed
337
+ Musubi Tuner; use that project's setup instructions for its CUDA dependencies
338
+ and base weights. ADAM does not install or replace the trainer environment.
339
+
340
+ 1. In **Dataset**, choose a folder of short video clips with matching `.txt`
341
+ captions, import clips, or use **Split long video**. Splitting creates a new
342
+ folder of evenly spaced 49-frame clips at 12 FPS, preserving the source.
343
+ Select clips for looping preview and caption editing. Captions save on
344
+ selection/tab changes; include the exact trigger word from Training.
345
+ **Characters & recognition** maintains a local character library with one or
346
+ more reference images per character. It samples several frames from each
347
+ clip and uses local DINOv2 visual similarity to suggest zero or multiple
348
+ characters. Review the suggestions and uncheck false matches before applying
349
+ trigger words to captions; existing action descriptions are preserved.
350
+ Recognition suggestions are not applied automatically, and similarity
351
+ scores are only a review aid. Character references and the library are stored
352
+ under `data/video_characters/`.
353
+ To draft action captions, select one or more clips and choose **AI draft
354
+ captions**. ADAM samples six ordered frames per clip and asks the configured
355
+ Ollama model to describe visible actions and changes without guessing who is
356
+ present, in English. It retries once if the model returns CJK text. Review and
357
+ edit each draft as soon as it finishes while the next selected clip is being
358
+ processed. ADAM trims notes, alternate summaries, and repeated commentary to
359
+ one short caption sentence. Check the captions to keep, then save;
360
+ ADAM adds the active training trigger automatically. Failed clips remain
361
+ unchanged, and unchecked drafts are not written.
362
+ 2. In **Training**, enter a unique run name and review epochs, trigger, frame
363
+ buckets, training resolution, rank/alpha, learning rate and memory swapping.
364
+ **Review full training pipeline** creates a job awaiting approval in
365
+ **Jobs / History**. After approval it validates clips, caches video latents,
366
+ caches captions, then trains. Failures stop subsequent stages.
367
+ 3. In **Generate**, choose a compatible checkpoint and set prompt, strength,
368
+ landscape/portrait format, seconds, FPS, steps, seed and block swapping.
369
+ ADAM converts duration to Wan's `4N+1` frame count and displays the actual
370
+ duration. Start with **Fast preview preset** before trying longer clips.
371
+ 4. **Videos / takes** lists new MP4s with prompt/seed/settings and existing
372
+ samples from the connected trainer. Double-click to play a video.
373
+
374
+ Training outputs default to `data/video_models/<run name>/`, including an
375
+ isolated cache, dataset configuration, logs and model metadata. Existing output
376
+ folders cannot be overwritten. **Continue weights** loads a Wan adapter into a
377
+ new run with a fresh optimizer and schedule; its rank and alpha come from the
378
+ checkpoint. It does not restore a full interrupted optimizer state. Cancellation
379
+ keeps previously written checkpoints, but does not force a new checkpoint.
380
+
381
+ Generations are stored in `data/video_generations/`, with an MP4 and
382
+ `generation.json` recording the actual seed, model, prompt, settings and output
383
+ dimensions. Existing Wan checkpoints in LoRAVideoTrainer's `output/` are indexed
384
+ separately from SDXL. The first visit imports compatible local trainer settings;
385
+ subsequent changes are saved in ADAM's own configuration. The original trainer's
386
+ storyboard editor remains available through that application.
387
+
388
+ All work uses ADAM's shared job queue, logs and pause/cancel controls. Training
389
+ keeps ORION review and ATLAS supervision, with video-specific workload notes.
390
+ NOVA requests video samples instead of applying image-preview quality checks.
391
+ Chat requests mentioning Wan or video LoRA direct you to the dedicated workspace.
392
+
393
+ ## Training Studio
394
+
395
+ The **Training Studio** turns completed work into a reviewable experiment loop:
396
+
397
+ - **Datasets** provides an image gallery, keep/reject decisions, caption editing,
398
+ exact duplicate detection, and visually similar duplicate candidates.
399
+ - **Experiments** compares job settings and outcomes, opens outputs, marks a
400
+ preferred model, and converts successful settings into reusable recipes.
401
+ - **Checkpoint Lab** browses model checkpoints and output images, records
402
+ consistent prompt/seed evaluations, and sends preview requests through the
403
+ normal approval-aware planner.
404
+ - **Recipes** preserves training starting points and can import or export
405
+ portable JSON recipe files.
406
+
407
+ ### EVE AI Dataset Review
408
+
409
+ In Training Studio → Datasets, **EVE AI Review…** performs a local reference-
410
+ guided visual review. Add one or more good reference images and optional bad
411
+ references, then choose Keep and Reject confidence thresholds. EVE uses a small
412
+ DINOv2 vision model to divide the selected dataset into **Keep**, **Reject**, and
413
+ **Uncertain** galleries with confidence scores. The model is downloaded once on
414
+ first use and subsequent analysis stays local.
415
+
416
+ Nothing is applied automatically. Inspect both sides, double-click images for a
417
+ full view, and move selected results between the three groups before choosing
418
+ **Apply EVE review**. EVE's decisions remain ordinary Training Studio review
419
+ marks: they can be manually changed, and rejected files are not moved until
420
+ **Exclude rejected** is selected. The latest proposal is also saved under
421
+ `data/eve_reviews/` for auditing. Use **Select all in current group** (or
422
+ Ctrl/Shift selection) to move many images at once; EVE transfers only the
423
+ chosen thumbnails so manual sorting stays responsive on large datasets.
424
+
425
+ Training panels show elapsed time, a progress-based ETA, recent logs, and a
426
+ loss sparkline when the connected trainer reports `loss`. Preflight summaries
427
+ include clearly labelled workload, duration, VRAM, and disk estimates. These
428
+ estimates are planning hints rather than hardware guarantees.
429
+
430
+ Create a Model also supports live training previews with a configurable
431
+ epoch interval, prompt, and reproducible seed for each model tab. While a
432
+ training job is active, its newest 256×256 preview appears in the right sidebar
433
+ with the source epoch and next scheduled preview. The full-size trainer output
434
+ can be opened from the card. Built-in adapters may publish previews directly;
435
+ registered DDPM, Flow, LoRA, APVD, MaskGit, and other trainers can also
436
+ participate by writing conventionally named `preview`, `sample`, or `epoch`
437
+ images beneath their declared output folder.
438
+
439
+ ## INRFlow
440
+
441
+ ADAM includes an experimental, lightweight **INRFlow** trainer and generator as
442
+ a separate built-in model architecture. It follows the paper's ambient-space
443
+ design: an image is represented as coordinate-to-RGB pairs, spatial context
444
+ latents summarize the current noisy field, and a point decoder predicts the
445
+ flow velocity for independently sampled pixel queries. Training therefore uses
446
+ continuous flow matching directly on RGB values without a VAE or another
447
+ pretrained image compressor.
448
+
449
+ The implementation is intentionally scaled for local experiments rather than
450
+ the much larger published configurations. The 64px default is the recommended
451
+ starting point on an 8–12 GB GPU. Pixel-query subsampling lowers training memory;
452
+ batch size, query count, width, and resolution can be reduced further. Models
453
+ save ordinary checkpoints, EMA weights, resumable optimizer state, metadata,
454
+ and optional training previews beneath
455
+ `data/model_plugin_outputs/inrflow/`.
456
+
457
+ Completed INRFlow models appear in Generations next to regular Flow Matching,
458
+ with reproducible seeds, Euler or Heun integration, live ODE previews, Smart
459
+ Generation, and square resolution-flexible queries. A different output
460
+ resolution is coordinate-field extrapolation, so native resolution is the fair
461
+ default for model comparisons. Creative notes are metadata because the current
462
+ backend is unconditional.
463
+
464
+ This is an independent ADAM-sized implementation informed by the
465
+ [INRFlow paper](https://arxiv.org/abs/2412.03791) and
466
+ [Apple's reference repository](https://github.com/apple/ml-inrflow), not a copy
467
+ of the published training setup or a claim of reproducing its reported model
468
+ scale.
469
+
470
+ ## Generations
471
+
472
+ The **Generations** workspace runs compatible registered image generators
473
+ without opening their separate desktop interfaces. DDPM, regular Flow Matching,
474
+ INRFlow, and PixelRow can generate from completed models with reproducible
475
+ settings; PixelRow uses rows, while both flow backends use ODE sampling steps.
476
+ Generation work uses the normal ADAM job queue, progress reporting,
477
+ cancellation, and logging.
478
+
479
+ Every completed batch is stored under `data/generations/` with its images and a
480
+ `generation.json` sidecar. The history gallery can open an image or batch folder
481
+ and restore the exact settings for another run. DDPM creative notes are stored
482
+ with a batch for organization; they are not presented as text conditioning for
483
+ an unconditional DDPM model.
484
+
485
+ The workspace groups controls into **Model & Prompt**, **Image Settings**, and
486
+ **Advanced Settings**. Dimensions, presets, image count, and seed are available
487
+ in Image Settings; sampling controls are expanded by default. Collapsing Advanced
488
+ Settings preserves its values. Custom LoRA dimensions and reference images are
489
+ remembered when reopening the page.
490
+
491
+ Generation history opens with image cards for each generator and a preview of
492
+ the latest batch. Click a generator, then a model, to browse its images and select
493
+ that model for generation. **Recent Output**, **All Generations**, and **Favorites**
494
+ provide alternate history views. The thumbnail strip selects the image shown
495
+ alongside its metadata and reuse/save actions. In **Compare**, pin one image as
496
+ a reference and select another thumbnail to view them together. Long metadata
497
+ values are available in tooltips. These views do not move or rewrite older
498
+ generation files.
499
+
500
+ **Generation Cycle…** selects multiple compatible completed models and queues
501
+ one generation step per model. Choose images per model, a shared prompt or
502
+ creative note, starting seed, slideshow duration, looping, fullscreen playback,
503
+ and an optional model/trainer label. When the cycle finishes, ADAM opens the
504
+ results as a local slideshow while preserving every ordinary generation record
505
+ in history.
506
+
507
+ If ADAM discovers a job interrupted by an unexpected shutdown, it offers to
508
+ open Jobs & History. The previous record remains intact and can be retried as a
509
+ new approval-gated job. Job logs can also be exported for troubleshooting.
510
+
511
+ ## Connect an existing tool
512
+
513
+ ADAM supports importable Python functions and command-line Python scripts.
514
+ For a no-code setup, open **Settings → External Tools → Add external tool**.
515
+ Choose the program folder, select its training entry script and important
516
+ configuration files, then review ADAM's static compatibility and safety report.
517
+ The report covers:
518
+
519
+ - detected command-line options and required inputs;
520
+ - likely dataset formats;
521
+ - output and checkpoint behavior;
522
+ - progress reporting;
523
+ - resume-training support; and
524
+ - potentially risky operations visible in the selected entry script.
525
+
526
+ The 1–10 rating measures how clearly the script fits ADAM's safe command-line
527
+ contract. It is not a guarantee that third-party code is harmless. ADAM does
528
+ not execute a script while scanning it, external tools cannot replace built-in
529
+ registry entries, and every external-tool run requires explicit approval.
530
+
531
+ After registration, a tool can be planned with a request such as:
532
+
533
+ ```text
534
+ Run APVD Model Trainer with dataset=D:\DreamData, epochs=20, output=D:\APVD\output
535
+ ```
536
+
537
+ ADAM will ask for any required inputs that were omitted before it offers the
538
+ approval plan.
539
+
540
+ For manual registry configuration, edit the relevant item in
541
+ `config/tools.json`:
542
+
543
+ ```json
544
+ {
545
+ "backend": {
546
+ "type": "python",
547
+ "module": "my_tools.lora",
548
+ "function": "train"
549
+ },
550
+ "demo": false
551
+ }
552
+ ```
553
+
554
+ The function receives a `ToolContext` as its first argument and keyword
555
+ arguments from the approved plan. This keeps training code in one place: your
556
+ existing GUI and ADAM can both call the same backend.
557
+
558
+ For scripts:
559
+
560
+ ```json
561
+ {
562
+ "backend": {
563
+ "type": "script",
564
+ "path": "D:/AI/LoRATrainer/train.py"
565
+ },
566
+ "demo": false
567
+ }
568
+ ```
569
+
570
+ ADAM invokes scripts directly with the current Python interpreter, captures
571
+ stdout/stderr, and never drives another GUI with mouse clicks.
572
+
573
+ ## Safety model
574
+
575
+ - Plans are shown before execution.
576
+ - Long, destructive, or high-volume work requires confirmation.
577
+ - Unregistered tools cannot be invoked.
578
+ - External paths and arguments are validated before execution.
579
+ - The LLM may propose a plan, but only registered tools can execute it.
580
+ - Pause, resume, and cancel controls are available for active jobs.
581
+ - Every tool action and state transition is logged.
582
+
583
+ ## Tests
584
+
585
+ ```powershell
586
+ python scripts/run_tests.py
587
+ `
588
+
589
+ The runner executes every test module in its own process, keeping Qt application
590
+ lifetimes isolated. A single module can also be run with
591
+ python -m pytest tests/test_agents.py -q.``
SHA256SUMS.txt ADDED
@@ -0,0 +1 @@
 
 
1
+ f0e60f14332f6b6c5445871e81133162df75d98f3113f8c9c43d6b2dfd8a5d29 ADAM-source-2026-10-01.zip
adam/assets.py CHANGED
@@ -26,6 +26,20 @@ def _friendly_name(value: str, fallback: str) -> str:
26
  return text
27
 
28
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
29
  @dataclass(slots=True)
30
  class Asset:
31
  id: str
@@ -167,12 +181,30 @@ class AssetRegistry:
167
  matches.append(item)
168
  return exact or matches
169
 
170
- def discover(self, config: Any) -> None:
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
171
  folders = config.get("tool_folders", {})
172
  if not isinstance(folders, dict):
173
  return
174
  folders = dict(folders)
175
  app_root = self.path.parent.parent
 
 
176
  if not folders.get("oasis_trainer"):
177
  try:
178
  external = json.loads((app_root / "config" / "external_tools.json").read_text(encoding="utf-8"))
@@ -186,7 +218,11 @@ class AssetRegistry:
186
  external_lora_root = app_root / "LoRAModelsHere"
187
  if external_lora_root.is_dir():
188
  for path in external_lora_root.rglob("*.safetensors"):
189
- if path.is_file() and "_comfy" not in path.stem.casefold():
 
 
 
 
190
  self.register(
191
  kind="model",
192
  name=path.stem.removesuffix("_cancelled"),
@@ -230,6 +266,37 @@ class AssetRegistry:
230
  root = Path(str(folders.get(folder_name, ""))) / output_name
231
  if not root.is_dir():
232
  continue
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
233
  for folder in root.iterdir():
234
  if not folder.is_dir():
235
  continue
@@ -250,22 +317,6 @@ class AssetRegistry:
250
  if p.name.rsplit("-", 1)[-1].isdigit()
251
  else -1,
252
  )
253
- elif trainer == "lora":
254
- trigger_word = ""
255
- checkpoints = sorted(
256
- (
257
- path for path in folder.glob("*.safetensors")
258
- if "_comfy" not in path.stem.casefold()
259
- ),
260
- key=lambda p: p.stat().st_mtime,
261
- )
262
- if checkpoints:
263
- name = checkpoints[-1].stem.removesuffix("_cancelled")
264
- try:
265
- metadata = json.loads((folder / "model_info.json").read_text(encoding="utf-8"))
266
- trigger_word = str(metadata.get("trigger_word") or "")
267
- except (OSError, ValueError, TypeError, json.JSONDecodeError):
268
- trigger_word = ""
269
  elif trainer == "flow":
270
  checkpoints = []
271
  try:
@@ -333,10 +384,11 @@ class AssetRegistry:
333
  and item.metadata.get("dataset_location_id") not in valid_location_ids
334
  )
335
  ]
336
- dataset_registry.discover_into_assets(self, persist=False)
337
  except Exception:
338
  pass
339
- self.save()
 
340
 
341
  def _flow_dataset_paths(self) -> dict[str, str]:
342
  """Recover source datasets for Flow models created by ADAM in older runs."""
 
26
  return text
27
 
28
 
29
+ def _is_lora_training_checkpoint(path: Path) -> bool:
30
+ """Return whether a LoRA weight is an intermediate training snapshot.
31
+
32
+ The LoRA trainer writes both the finished adapter and periodic weights such
33
+ as ``name_epoch_0050.safetensors``. The latter are useful for recovery,
34
+ but are not independently selectable models in ADAM's model library.
35
+ """
36
+ name = path.stem.casefold()
37
+ return bool(re.search(
38
+ r"(?:^|[_\- ])(?:checkpoint(?:[_\- ]?(?:epoch|e|step))?|epoch|e|step)[_\- ]?\d+(?:[_\- ]|$)",
39
+ name,
40
+ ))
41
+
42
+
43
  @dataclass(slots=True)
44
  class Asset:
45
  id: str
 
181
  matches.append(item)
182
  return exact or matches
183
 
184
+ def discover(self, config: Any, *, persist: bool = True) -> None:
185
+ # Models are stored by their output folder (or the model file itself).
186
+ # Keep the registry in step with the filesystem so removing an old
187
+ # output cannot leave a ghost model that makes name matching ambiguous.
188
+ self.assets = [
189
+ item
190
+ for item in self.assets
191
+ if item.kind != "model" or (
192
+ item.path.strip() and Path(item.path).expanduser().exists()
193
+ )
194
+ # Old ADAM versions registered LoRA epoch snapshots. Prune those
195
+ # stale records as well as skipping them during new discovery.
196
+ and not (
197
+ item.trainer == "lora"
198
+ and _is_lora_training_checkpoint(Path(item.path))
199
+ )
200
+ ]
201
  folders = config.get("tool_folders", {})
202
  if not isinstance(folders, dict):
203
  return
204
  folders = dict(folders)
205
  app_root = self.path.parent.parent
206
+ from adam.video_lora import discover_assets as discover_video_assets
207
+ discover_video_assets(self, app_root, config)
208
  if not folders.get("oasis_trainer"):
209
  try:
210
  external = json.loads((app_root / "config" / "external_tools.json").read_text(encoding="utf-8"))
 
218
  external_lora_root = app_root / "LoRAModelsHere"
219
  if external_lora_root.is_dir():
220
  for path in external_lora_root.rglob("*.safetensors"):
221
+ if (
222
+ path.is_file()
223
+ and "_comfy" not in path.stem.casefold()
224
+ and not _is_lora_training_checkpoint(path)
225
+ ):
226
  self.register(
227
  kind="model",
228
  name=path.stem.removesuffix("_cancelled"),
 
266
  root = Path(str(folders.get(folder_name, ""))) / output_name
267
  if not root.is_dir():
268
  continue
269
+ # LoRA Trainer versions do not all agree on their output layout.
270
+ # Some write ``output/<run>/<name>.safetensors`` while others add
271
+ # a second folder below the run. Register the actual weight file
272
+ # in either layout so the generator can load it directly.
273
+ if trainer == "lora":
274
+ for checkpoint_path in root.rglob("*.safetensors"):
275
+ if (
276
+ not checkpoint_path.is_file()
277
+ or "_comfy" in checkpoint_path.stem.casefold()
278
+ or _is_lora_training_checkpoint(checkpoint_path)
279
+ ):
280
+ continue
281
+ trigger_word = ""
282
+ for metadata_path in (checkpoint_path.parent / "model_info.json", checkpoint_path.parent.parent / "model_info.json"):
283
+ try:
284
+ metadata = json.loads(metadata_path.read_text(encoding="utf-8"))
285
+ trigger_word = str(metadata.get("trigger_word") or "")
286
+ if trigger_word:
287
+ break
288
+ except (OSError, ValueError, TypeError, json.JSONDecodeError):
289
+ continue
290
+ self.register(
291
+ kind="model",
292
+ name=checkpoint_path.stem.removesuffix("_cancelled"),
293
+ path=str(checkpoint_path),
294
+ trainer="lora",
295
+ checkpoint=str(checkpoint_path),
296
+ metadata={"trigger_word": trigger_word or checkpoint_path.stem.removesuffix("_cancelled")},
297
+ persist=False,
298
+ )
299
+ continue
300
  for folder in root.iterdir():
301
  if not folder.is_dir():
302
  continue
 
317
  if p.name.rsplit("-", 1)[-1].isdigit()
318
  else -1,
319
  )
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
320
  elif trainer == "flow":
321
  checkpoints = []
322
  try:
 
384
  and item.metadata.get("dataset_location_id") not in valid_location_ids
385
  )
386
  ]
387
+ dataset_registry.discover_into_assets(self, persist=False, update_cache=persist)
388
  except Exception:
389
  pass
390
+ if persist:
391
+ self.save()
392
 
393
  def _flow_dataset_paths(self) -> dict[str, str]:
394
  """Recover source datasets for Flow models created by ADAM in older runs."""
adam/auto_training.py ADDED
@@ -0,0 +1,109 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Intent-level AUTO policies for ADAM training requests.
2
+
3
+ Natural-language parsing belongs in the planner. This module deliberately does
4
+ not inspect prompt wording: it turns an already-selected trainer, profile, and
5
+ dataset size into transparent, reproducible training settings.
6
+ """
7
+
8
+ from __future__ import annotations
9
+
10
+ from dataclasses import dataclass
11
+ from enum import Enum
12
+ from typing import Any
13
+
14
+ from adam.model_profiles import ModelProfile
15
+ from adam.models import SystemSnapshot
16
+ from adam.recommendations import SettingsRecommendation, recommend_for_profile
17
+
18
+
19
+ class TrainingProfile(str, Enum):
20
+ TEST = "test"
21
+ BALANCED = "balanced"
22
+ QUALITY = "quality"
23
+ OVERNIGHT = "overnight"
24
+
25
+
26
+ @dataclass(frozen=True, slots=True)
27
+ class AutoTrainingPlan:
28
+ """Resolved, explainable settings for one training run."""
29
+
30
+ profile: TrainingProfile
31
+ dataset_target: int
32
+ epochs: int
33
+ settings: dict[str, Any]
34
+ summary: str
35
+ reasons: tuple[str, ...] = ()
36
+ warnings: tuple[str, ...] = ()
37
+
38
+
39
+ def profile_from_request(request: str) -> TrainingProfile:
40
+ """Map only stable, user-facing intent modifiers to a named policy."""
41
+ lowered = request.casefold()
42
+ if any(word in lowered for word in ("overnight", "all night", "long run")):
43
+ return TrainingProfile.OVERNIGHT
44
+ if any(phrase in lowered for phrase in ("high quality", "best quality", "really good", "quality")):
45
+ return TrainingProfile.QUALITY
46
+ if any(word in lowered for word in ("quick", "quickly", "test", "small", "smoke test")):
47
+ return TrainingProfile.TEST
48
+ return TrainingProfile.BALANCED
49
+
50
+
51
+ def dataset_target_for(trainer: str, profile: TrainingProfile) -> int:
52
+ """Choose a collection target, never a random count.
53
+
54
+ These are conservative collection targets. The later recommendation is
55
+ calculated from the actual usable count when an existing dataset is known.
56
+ """
57
+ targets = {
58
+ "lora": {TrainingProfile.TEST: 40, TrainingProfile.BALANCED: 150, TrainingProfile.QUALITY: 300, TrainingProfile.OVERNIGHT: 500},
59
+ "ddpm": {TrainingProfile.TEST: 100, TrainingProfile.BALANCED: 400, TrainingProfile.QUALITY: 800, TrainingProfile.OVERNIGHT: 1_200},
60
+ "flow": {TrainingProfile.TEST: 100, TrainingProfile.BALANCED: 400, TrainingProfile.QUALITY: 800, TrainingProfile.OVERNIGHT: 1_200},
61
+ "inrflow": {TrainingProfile.TEST: 80, TrainingProfile.BALANCED: 300, TrainingProfile.QUALITY: 600, TrainingProfile.OVERNIGHT: 900},
62
+ }
63
+ return targets.get(trainer, targets["ddpm"])[profile]
64
+
65
+
66
+ def resolve_auto_training(
67
+ profile: ModelProfile,
68
+ *,
69
+ trainer: str,
70
+ policy: TrainingProfile,
71
+ dataset_items: int,
72
+ snapshot: SystemSnapshot | None = None,
73
+ ) -> AutoTrainingPlan:
74
+ """Resolve a named policy through ADAM's existing exposure-aware recommender."""
75
+ recommendation: SettingsRecommendation = recommend_for_profile(
76
+ profile,
77
+ dataset_items=max(10, dataset_items),
78
+ snapshot=snapshot,
79
+ )
80
+ epoch_multiplier = {
81
+ TrainingProfile.TEST: 0.20,
82
+ TrainingProfile.BALANCED: 1.00,
83
+ TrainingProfile.QUALITY: 1.35,
84
+ TrainingProfile.OVERNIGHT: 1.80,
85
+ }[policy]
86
+ # Preserve safe bounds from the recommendation; profiles express their own
87
+ # architecture-specific baseline rather than sharing a global epoch range.
88
+ minimum = 3 if policy is TrainingProfile.TEST else 10
89
+ maximum = 1_000 if policy is TrainingProfile.OVERNIGHT else 600
90
+ epochs = max(minimum, min(maximum, round(recommendation.epochs * epoch_multiplier)))
91
+ settings = dict(recommendation.settings)
92
+ if "save_every" in settings:
93
+ settings["save_every"] = max(1, min(int(settings["save_every"]), max(1, epochs // 4)))
94
+ if "preview_every" in settings:
95
+ settings["preview_every"] = max(1, min(int(settings["preview_every"]), max(1, epochs // 5)))
96
+ target = dataset_target_for(trainer, policy)
97
+ summary = (
98
+ f"{policy.value.title()} AUTO policy: target {target:,} source images; "
99
+ f"{epochs:,} epoch budget based on {max(10, dataset_items):,} expected usable items."
100
+ )
101
+ return AutoTrainingPlan(
102
+ profile=policy,
103
+ dataset_target=target,
104
+ epochs=epochs,
105
+ settings=settings,
106
+ summary=summary,
107
+ reasons=tuple(recommendation.reasons),
108
+ warnings=tuple(recommendation.warnings),
109
+ )
adam/cnn_reviewer.py ADDED
@@ -0,0 +1,172 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """A small, local CNN that learns which gameplay frames are worth reviewing.
2
+
3
+ This is intentionally independent from Oasis. Its job is to prioritize and
4
+ quality-check collected frames; it never supplies inputs to an Oasis checkpoint
5
+ or changes an Oasis model.
6
+ """
7
+ from __future__ import annotations
8
+
9
+ from dataclasses import dataclass
10
+ from pathlib import Path
11
+ from typing import Callable
12
+
13
+
14
+ IMAGE_SUFFIXES = {".png", ".jpg", ".jpeg", ".webp", ".bmp"}
15
+ Progress = Callable[[str], None]
16
+
17
+
18
+ @dataclass(slots=True)
19
+ class ReviewerTrainingResult:
20
+ checkpoint: str
21
+ kept_examples: int
22
+ rejected_examples: int
23
+ epochs: int
24
+
25
+
26
+ @dataclass(slots=True)
27
+ class FrameScore:
28
+ path: str
29
+ keep_probability: float
30
+ suggestion: str
31
+
32
+
33
+ def image_paths(folder: str | Path, *, limit: int = 5_000) -> list[Path]:
34
+ root = Path(folder).expanduser()
35
+ if not root.is_dir():
36
+ return []
37
+ return [
38
+ path for path in sorted(root.rglob("*"))
39
+ if path.is_file() and path.suffix.casefold() in IMAGE_SUFFIXES
40
+ ][:limit]
41
+
42
+
43
+ def reviewer_checkpoint(root: str | Path, dataset_folder: str | Path) -> Path:
44
+ """Return an ADAM-owned checkpoint path, separate from dataset and Oasis."""
45
+ import hashlib
46
+
47
+ dataset = str(Path(dataset_folder).expanduser().resolve()).encode("utf-8")
48
+ identifier = hashlib.sha1(dataset).hexdigest()[:12]
49
+ return Path(root).expanduser().resolve() / "data" / "cnn_reviewers" / f"{identifier}.pt"
50
+
51
+
52
+ def _torch():
53
+ try:
54
+ import torch
55
+ from torch import nn
56
+ except ImportError as exc: # pragma: no cover - controlled by application install
57
+ raise RuntimeError("CNN Reviewer needs PyTorch. Install the ADAM requirements first.") from exc
58
+ return torch, nn
59
+
60
+
61
+ def make_reviewer_model():
62
+ """Build a deliberately small binary CNN for local frame triage."""
63
+ _torch_module, nn = _torch()
64
+ return nn.Sequential(
65
+ nn.Conv2d(3, 16, kernel_size=5, stride=2, padding=2), nn.ReLU(),
66
+ nn.Conv2d(16, 32, kernel_size=3, stride=2, padding=1), nn.ReLU(),
67
+ nn.Conv2d(32, 48, kernel_size=3, stride=2, padding=1), nn.ReLU(),
68
+ nn.AdaptiveAvgPool2d((1, 1)), nn.Flatten(), nn.Linear(48, 1),
69
+ )
70
+
71
+
72
+ def _load_image(path: Path):
73
+ torch, _nn = _torch()
74
+ from PIL import Image
75
+
76
+ with Image.open(path) as image:
77
+ image = image.convert("RGB").resize((128, 72))
78
+ # No torchvision transform is required, which keeps this feature portable.
79
+ pixels = image.get_flattened_data() if hasattr(image, "get_flattened_data") else image.getdata()
80
+ values = torch.tensor(list(pixels), dtype=torch.float32)
81
+ return values.reshape(72, 128, 3).permute(2, 0, 1).div_(255.0)
82
+
83
+
84
+ def _labeled_paths(decisions: dict[str, str], *, per_class_limit: int = 250) -> tuple[list[Path], list[float]]:
85
+ grouped: dict[str, list[Path]] = {"keep": [], "reject": []}
86
+ for raw_path, decision in decisions.items():
87
+ if decision not in grouped:
88
+ continue
89
+ path = Path(raw_path).expanduser()
90
+ if path.is_file() and path.suffix.casefold() in IMAGE_SUFFIXES:
91
+ grouped[decision].append(path)
92
+ kept = sorted(grouped["keep"])[:per_class_limit]
93
+ rejected = sorted(grouped["reject"])[:per_class_limit]
94
+ return kept + rejected, [1.0] * len(kept) + [0.0] * len(rejected)
95
+
96
+
97
+ def train_reviewer(
98
+ root: str | Path,
99
+ dataset_folder: str | Path,
100
+ decisions: dict[str, str],
101
+ *,
102
+ epochs: int = 8,
103
+ progress: Progress | None = None,
104
+ ) -> ReviewerTrainingResult:
105
+ """Train a frame-quality CNN from explicit Keep and Reject review decisions."""
106
+ torch, nn = _torch()
107
+ paths, labels = _labeled_paths(decisions)
108
+ keep_count = int(sum(labels))
109
+ reject_count = len(labels) - keep_count
110
+ if min(keep_count, reject_count) < 8:
111
+ raise ValueError("Review at least 8 Keep and 8 Reject frames before training the CNN reviewer.")
112
+ if progress:
113
+ progress(f"Loading {len(paths)} reviewed frame(s)…")
114
+ images = []
115
+ valid_labels = []
116
+ for path, label in zip(paths, labels):
117
+ try:
118
+ images.append(_load_image(path))
119
+ valid_labels.append(label)
120
+ except Exception:
121
+ continue
122
+ keep_count = int(sum(valid_labels))
123
+ reject_count = len(valid_labels) - keep_count
124
+ if min(keep_count, reject_count) < 8:
125
+ raise ValueError("Some reviewed images could not be read; at least 8 valid Keep and Reject frames are needed.")
126
+ torch.manual_seed(7)
127
+ inputs = torch.stack(images)
128
+ targets = torch.tensor(valid_labels, dtype=torch.float32).unsqueeze(1)
129
+ model = make_reviewer_model()
130
+ optimizer = torch.optim.Adam(model.parameters(), lr=0.001)
131
+ loss_fn = nn.BCEWithLogitsLoss()
132
+ model.train()
133
+ for epoch in range(max(1, min(int(epochs), 50))):
134
+ order = torch.randperm(len(inputs))
135
+ loss_value = 0.0
136
+ batches = 0
137
+ for start in range(0, len(order), 16):
138
+ batch = order[start:start + 16]
139
+ optimizer.zero_grad()
140
+ loss = loss_fn(model(inputs[batch]), targets[batch])
141
+ loss.backward()
142
+ optimizer.step()
143
+ loss_value += float(loss.detach())
144
+ batches += 1
145
+ if progress:
146
+ progress(f"CNN reviewer epoch {epoch + 1}/{epochs} · loss {loss_value / max(1, batches):.3f}")
147
+ checkpoint = reviewer_checkpoint(root, dataset_folder)
148
+ checkpoint.parent.mkdir(parents=True, exist_ok=True)
149
+ torch.save({"state_dict": model.state_dict(), "image_size": [128, 72]}, checkpoint)
150
+ return ReviewerTrainingResult(str(checkpoint), keep_count, reject_count, epochs)
151
+
152
+
153
+ def score_frames(checkpoint: str | Path, folder: str | Path, *, limit: int = 5_000, progress: Progress | None = None) -> list[FrameScore]:
154
+ """Return review suggestions without modifying the dataset or its decisions."""
155
+ torch, _nn = _torch()
156
+ saved = torch.load(Path(checkpoint), map_location="cpu", weights_only=True)
157
+ model = make_reviewer_model()
158
+ model.load_state_dict(saved["state_dict"])
159
+ model.eval()
160
+ paths = image_paths(folder, limit=limit)
161
+ scores: list[FrameScore] = []
162
+ with torch.no_grad():
163
+ for index, path in enumerate(paths, 1):
164
+ try:
165
+ probability = float(torch.sigmoid(model(_load_image(path).unsqueeze(0))).item())
166
+ except Exception:
167
+ continue
168
+ suggestion = "keep" if probability >= 0.70 else "reject" if probability <= 0.30 else "review"
169
+ scores.append(FrameScore(str(path.resolve()), probability, suggestion))
170
+ if progress and (index % 100 == 0 or index == len(paths)):
171
+ progress(f"CNN reviewer scored {index}/{len(paths)} frame(s)…")
172
+ return scores
adam/commands.py CHANGED
@@ -83,7 +83,7 @@ class TrainingCommand:
83
  "resolution", "batch_size", "learning_rate", "gradient_accumulation_steps",
84
  "dataloader_num_workers", "mixed_precision", "save_every", "preview_steps",
85
  "training_intensity", "preview_enabled", "preview_every", "preview_prompt",
86
- "preview_seed",
87
  },
88
  "flow": {
89
  "resolution", "batch_size", "learning_rate", "gradient_accumulation",
 
83
  "resolution", "batch_size", "learning_rate", "gradient_accumulation_steps",
84
  "dataloader_num_workers", "mixed_precision", "save_every", "preview_steps",
85
  "training_intensity", "preview_enabled", "preview_every", "preview_prompt",
86
+ "preview_seed", "training_aspect_ratio", "resize_mode",
87
  },
88
  "flow": {
89
  "resolution", "batch_size", "learning_rate", "gradient_accumulation",
adam/config.py CHANGED
@@ -11,6 +11,8 @@ DEFAULT_SETTINGS: dict[str, Any] = {
11
  "ollama_url": "http://localhost:11434",
12
  "ollama_model": "qwen2.5:1.5b",
13
  "ollama_chat_max_tokens": 1024,
 
 
14
  "web_search_enabled": True,
15
  "web_link_reading_enabled": True,
16
  "command_center_mode": "trainer",
@@ -41,6 +43,7 @@ DEFAULT_SETTINGS: dict[str, Any] = {
41
  "ddpm_trainer": "",
42
  "flow_trainer": "",
43
  "oasis_trainer": "",
 
44
  "preview_generator": "",
45
  },
46
  }
 
11
  "ollama_url": "http://localhost:11434",
12
  "ollama_model": "qwen2.5:1.5b",
13
  "ollama_chat_max_tokens": 1024,
14
+ "ollama_chat_response_length": "automatic",
15
+ "ollama_proposed_actions": True,
16
  "web_search_enabled": True,
17
  "web_link_reading_enabled": True,
18
  "command_center_mode": "trainer",
 
43
  "ddpm_trainer": "",
44
  "flow_trainer": "",
45
  "oasis_trainer": "",
46
+ "wan_video_trainer": "",
47
  "preview_generator": "",
48
  },
49
  }
adam/dataset_registry.py CHANGED
@@ -210,9 +210,9 @@ class DatasetRegistry:
210
  )
211
  return locations
212
 
213
- def discover_into_assets(self, assets: "AssetRegistry", *, persist: bool = False) -> list["Asset"]:
214
  discovered: list[Asset] = []
215
- records = self.discover(asset_registry=assets, refresh_missing=False)
216
  for record in records:
217
  if not record.exists:
218
  continue
@@ -236,6 +236,7 @@ class DatasetRegistry:
236
  *,
237
  asset_registry: "AssetRegistry | None" = None,
238
  refresh_missing: bool = True,
 
239
  ) -> list[DatasetRecord]:
240
  self.load()
241
  changed = False
@@ -271,7 +272,7 @@ class DatasetRegistry:
271
  record.exists = Path(record.path).is_dir()
272
  if refresh_missing and record.exists and self._needs_refresh(record):
273
  self.refresh_async(record.path, source=record.source, location_id=record.location_id)
274
- if changed:
275
  self.save()
276
  return self.sorted_records()
277
 
 
210
  )
211
  return locations
212
 
213
+ def discover_into_assets(self, assets: "AssetRegistry", *, persist: bool = False, update_cache: bool = True) -> list["Asset"]:
214
  discovered: list[Asset] = []
215
+ records = self.discover(asset_registry=assets, refresh_missing=False, persist=update_cache)
216
  for record in records:
217
  if not record.exists:
218
  continue
 
236
  *,
237
  asset_registry: "AssetRegistry | None" = None,
238
  refresh_missing: bool = True,
239
+ persist: bool = True,
240
  ) -> list[DatasetRecord]:
241
  self.load()
242
  changed = False
 
272
  record.exists = Path(record.path).is_dir()
273
  if refresh_missing and record.exists and self._needs_refresh(record):
274
  self.refresh_async(record.path, source=record.source, location_id=record.location_id)
275
+ if changed and persist:
276
  self.save()
277
  return self.sorted_records()
278
 
adam/generations.py CHANGED
@@ -2,6 +2,7 @@ from __future__ import annotations
2
 
3
  import json
4
  import re
 
5
  from dataclasses import dataclass
6
  from pathlib import Path
7
  from typing import Any
@@ -13,6 +14,220 @@ from adam.registry import ToolRegistry, ToolSpec
13
  IMAGE_EXTENSIONS = {".png", ".jpg", ".jpeg", ".webp", ".bmp"}
14
 
15
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
16
  @dataclass(frozen=True, slots=True)
17
  class ChatGenerationRequest:
18
  """Generation settings recognized from a Command Center message."""
@@ -34,6 +249,11 @@ class ChatGenerationRequest:
34
  reference_strength: int | None = None
35
  reference_image: str = ""
36
  has_positive_prompt: bool = False
 
 
 
 
 
37
 
38
 
39
  _QUOTED = r'["\u201c\u201d]([^"\u201c\u201d]+)["\u201c\u201d]'
@@ -47,7 +267,9 @@ def generation_model_match_score(query: str, model_name: str) -> int:
47
  """Score whether conversational subject text clearly names a saved model."""
48
  def words(value: str) -> list[str]:
49
  value = re.sub(r"(?<=[a-z0-9])(?=[A-Z])", " ", value)
50
- ignored = {"a", "an", "the", "of", "image", "picture", "model", "ddpm", "flow", "matching", "lora"}
 
 
51
  return [word for word in re.findall(r"[a-z0-9]+", value.casefold()) if word not in ignored]
52
 
53
  query_words = words(query)
@@ -73,21 +295,43 @@ def parse_chat_generation_request(text: str) -> ChatGenerationRequest | None:
73
  This intentionally requires both a creation verb and the word image/picture so
74
  ordinary planning requests continue through the regular Command Center planner.
75
  """
 
 
 
 
 
 
76
  request = " ".join(text.strip().split())
77
- if not request or not re.search(r"\b(generate|create|make)\b", request, re.I):
78
- return None
79
- if not re.search(r"\b(image|images|picture|pictures)\b", request, re.I):
 
 
 
 
 
 
 
 
 
 
80
  return None
81
 
82
  provider_hint = ""
83
  provider_match = re.search(
84
- r"\b(ddpm|ddim|flow(?:\s+matching)?|lora)\b[\"\u201c\u201d]?(?=\s+(?:image|picture))",
85
  request,
86
  re.I,
87
  )
88
  if provider_match:
89
  hint = provider_match.group(1).casefold()
90
- provider_hint = "ddpm" if hint in {"ddpm", "ddim"} else "flow" if hint.startswith("flow") else "lora"
 
 
 
 
 
 
91
  # Support natural phrasing such as "Generate an image of LoRA OrangeCat".
92
  lora_subject_match = re.search(
93
  rf"\b(?:image|picture)s?\s+of\s+(?:a\s+)?LoRA\s+{_QUOTED}",
@@ -123,10 +367,14 @@ def parse_chat_generation_request(text: str) -> ChatGenerationRequest | None:
123
  # In promptless commands, a provider suffix is usually part of the saved
124
  # model name (for example, "Minecraft Flow"), not prompt prose.
125
  if not provider_hint and subject:
126
- if re.search(r"\bflow(?:\s+match(?:ing)?)?\s*$", subject, re.I):
 
 
127
  provider_hint = "flow"
128
  elif re.search(r"\bddpm\s*$", subject, re.I):
129
  provider_hint = "ddpm"
 
 
130
 
131
  positive_match = re.search(
132
  rf"\bpositive\s+prompt(?:\s+of|\s*=|\s*:)?\s*{_QUOTED}",
@@ -286,6 +534,66 @@ class GenerationModelFolder:
286
  latest_at: str
287
 
288
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
289
  def generation_model_key(record: GenerationRecord) -> str:
290
  """Keep renamed or duplicated display names separated by model identity."""
291
  raw_path = str(record.model_path or "").strip()
@@ -345,14 +653,26 @@ def load_generation_history(root: Path, *, limit: int = 200) -> list[GenerationR
345
  history_root = root.resolve() / "data" / "generations"
346
  if not history_root.is_dir():
347
  return []
348
- records = [
349
- record
350
- for metadata_path in history_root.rglob("generation*.json")
351
- for record in [GenerationRecord.from_metadata(metadata_path)]
352
- if record is not None
353
- ]
 
 
 
 
 
 
 
 
 
 
 
 
354
  records.sort(key=lambda item: item.created_at or item.folder.name, reverse=True)
355
- return records[: max(1, int(limit))]
356
 
357
 
358
  def build_generation_plan(
 
2
 
3
  import json
4
  import re
5
+ from html import unescape
6
  from dataclasses import dataclass
7
  from pathlib import Path
8
  from typing import Any
 
14
  IMAGE_EXTENSIONS = {".png", ".jpg", ".jpeg", ".webp", ".bmp"}
15
 
16
 
17
+ @dataclass(frozen=True, slots=True)
18
+ class ImportedLoRAMetadata:
19
+ """Portable LoRA-generation settings copied from an image or another app."""
20
+
21
+ prompt: str
22
+ negative_prompt: str
23
+ seed: int
24
+ steps: int
25
+ cfg_scale: float | None
26
+ base_model_path: str
27
+ lora_path: str
28
+ lora_strength: float | None
29
+ sampler: str
30
+ width: int | None = None
31
+ height: int | None = None
32
+ sampler_note: str = ""
33
+
34
+
35
+ def parse_lora_generation_metadata(text: str) -> ImportedLoRAMetadata:
36
+ """Parse pasted LoRA metadata without requiring an intermediate JSON file.
37
+
38
+ Browser copy/paste sometimes HTML-escapes JSON (for example ``&#x20;`` for a
39
+ space), so decode those entities before loading it. This intentionally
40
+ accepts only the small, reproducible LoRA schema ADAM understands.
41
+ """
42
+ source = str(text or "").strip()
43
+ # A few sites escape twice when metadata is copied out of a code block.
44
+ for _ in range(2):
45
+ decoded = unescape(source)
46
+ if decoded == source:
47
+ break
48
+ source = decoded
49
+ try:
50
+ payload = json.loads(source)
51
+ except (TypeError, ValueError, json.JSONDecodeError):
52
+ # Some viewers copy a display block rather than strict JSON, omitting
53
+ # braces, commas, or quotes around values. Recover its named fields.
54
+ fields: dict[str, Any] = {}
55
+ for key in ("prompt", "negative_prompt", "seed", "steps", "cfg_scale", "model", "sampler"):
56
+ match = re.search(rf'["\']?{key}["\']?\s*:\s*(?:["\']([^"\']*)["\']|([^\r\n]+))', source, re.I)
57
+ if match:
58
+ fields[key] = (match.group(1) if match.group(1) is not None else match.group(2)).strip().rstrip(",").strip()
59
+ lora_path_match = re.search(r'["\']?path["\']?\s*:\s*["\']?([^"\',\r\n}\]]+)', source, re.I)
60
+ strength_match = re.search(r'["\']?strength["\']?\s*:\s*([^,\r\n}\]]+)', source, re.I)
61
+ if lora_path_match:
62
+ fields["loras"] = [{
63
+ "path": lora_path_match.group(1).strip(),
64
+ "strength": strength_match.group(1).strip() if strength_match else None,
65
+ }]
66
+ payload = fields
67
+ if not isinstance(payload, dict):
68
+ raise ValueError("Metadata must be a JSON object.")
69
+
70
+ # Source apps vary between `width`/`height` and `Width`/`Height`.
71
+ payload = {str(key).casefold(): value for key, value in payload.items()}
72
+
73
+ def text_value(key: str, *, required: bool = False) -> str:
74
+ value = payload.get(key, "")
75
+ if value is None:
76
+ value = ""
77
+ if not isinstance(value, str):
78
+ raise ValueError(f"{key.replace('_', ' ').title()} must be text.")
79
+ value = value.strip()
80
+ if required and not value:
81
+ raise ValueError(f"Metadata is missing {key.replace('_', ' ')}.")
82
+ return value
83
+
84
+ loras = payload.get("loras")
85
+ if not isinstance(loras, list):
86
+ raise ValueError("Metadata needs a 'loras' list (use [] when no LoRA was used).")
87
+ if loras and not isinstance(loras[0], dict):
88
+ raise ValueError("The first LoRA entry must be an object.")
89
+ lora_path = str(loras[0].get("path", "")).strip() if loras else ""
90
+ if loras and not lora_path:
91
+ raise ValueError("The first LoRA entry needs a path.")
92
+ try:
93
+ seed = int(payload.get("seed", 0))
94
+ steps = int(payload.get("steps", 30))
95
+ except (TypeError, ValueError) as exc:
96
+ raise ValueError("Seed and steps must be whole numbers.") from exc
97
+ # Automatic1111-style metadata commonly uses -1 for a fresh random seed.
98
+ # ADAM uses 0 for the same behavior in generation plans.
99
+ if seed == -1:
100
+ seed = 0
101
+ if not 0 <= seed <= 2_147_483_647 or steps < 1:
102
+ raise ValueError("Seed or steps is outside ADAM's supported range.")
103
+ cfg_value = payload.get("cfg_scale")
104
+ try:
105
+ cfg_scale = float(cfg_value) if cfg_value is not None else None
106
+ except (TypeError, ValueError) as exc:
107
+ raise ValueError("CFG scale must be a number.") from exc
108
+ strength_value = loras[0].get("strength") if loras else None
109
+ try:
110
+ lora_strength = float(strength_value) if strength_value is not None else None
111
+ except (TypeError, ValueError) as exc:
112
+ raise ValueError("LoRA strength must be a number.") from exc
113
+ def dimension(key: str) -> int | None:
114
+ value = payload.get(key)
115
+ if value is None or value == "":
116
+ return None
117
+ try:
118
+ result = int(value)
119
+ except (TypeError, ValueError) as exc:
120
+ raise ValueError(f"{key.title()} must be a whole number.") from exc
121
+ if not 256 <= result <= 2048:
122
+ raise ValueError(f"{key.title()} must be between 256 and 2048 pixels.")
123
+ return result
124
+
125
+ sampler = text_value("sampler") or "DPM++ 2M"
126
+ supported = {"DPM++ 2M", "DPM++ 2M Karras", "DPM++ 2M SDE", "DPM++ 2M SDE Karras", "DPM++ SDE", "DPM++ SDE Karras", "Euler", "Euler a", "Heun", "LMS", "DDIM"}
127
+ sampler_note = ""
128
+ if sampler not in supported:
129
+ normalized = sampler.casefold()
130
+ replacement = "DPM++ SDE" if "sde" in normalized else "DPM++ 2M" if "2m" in normalized else ""
131
+ if not replacement:
132
+ raise ValueError(f"Sampler '{sampler}' is not supported by ADAM's LoRA generator.")
133
+ sampler_note = f"Sampler '{sampler}' was mapped to '{replacement}' because the connected generator does not support it."
134
+ sampler = replacement
135
+ return ImportedLoRAMetadata(
136
+ prompt=text_value("prompt", required=True),
137
+ negative_prompt=text_value("negative_prompt"),
138
+ seed=seed,
139
+ steps=steps,
140
+ cfg_scale=cfg_scale,
141
+ base_model_path=text_value("model", required=True),
142
+ lora_path=lora_path,
143
+ lora_strength=lora_strength,
144
+ sampler=sampler,
145
+ width=dimension("width"),
146
+ height=dimension("height"),
147
+ sampler_note=sampler_note,
148
+ )
149
+
150
+
151
+ def parse_pasted_lora_metadata_request(text: str) -> "ChatGenerationRequest | None":
152
+ """Recognize copied LoRA metadata as an unambiguous prompt-box command."""
153
+ source = str(text or "")
154
+ if not re.search(r'["\']?prompt["\']?\s*:', source, re.I) or not re.search(
155
+ r'["\']?(?:loras|negative_prompt|sampler)["\']?\s*:', source, re.I
156
+ ):
157
+ return None
158
+ try:
159
+ metadata = parse_lora_generation_metadata(source)
160
+ except ValueError:
161
+ return None
162
+ return ChatGenerationRequest(
163
+ prompt=metadata.prompt,
164
+ subject=Path(metadata.lora_path).stem if metadata.lora_path else "",
165
+ provider_hint="lora",
166
+ model_query=Path(metadata.lora_path).stem if metadata.lora_path else "",
167
+ base_model_query=Path(metadata.base_model_path).stem,
168
+ negative_prompt=metadata.negative_prompt,
169
+ steps=metadata.steps,
170
+ sampler=metadata.sampler,
171
+ seed=metadata.seed,
172
+ cfg_scale=metadata.cfg_scale,
173
+ lora_strength=metadata.lora_strength,
174
+ width=metadata.width,
175
+ height=metadata.height,
176
+ has_positive_prompt=True,
177
+ is_pasted_metadata=True,
178
+ metadata_model_path=metadata.lora_path,
179
+ metadata_base_model_path=metadata.base_model_path,
180
+ )
181
+
182
+
183
+ def parse_plain_generation_metadata(text: str) -> "ChatGenerationRequest | None":
184
+ """Read the common CivitAI/A1111 and PixAI copied-text metadata layouts."""
185
+ source = str(text or "").replace("\r\n", "\n").strip()
186
+ if not source:
187
+ return None
188
+ civitai = re.search(r"\bNegative\s+prompt\s*:", source, re.I)
189
+ pixai = re.search(r"\b(?:Sampling\s+Steps|Original\s+Prompt)\b", source, re.I)
190
+ if not civitai and not pixai:
191
+ return None
192
+ prompt = ""
193
+ negative = ""
194
+ if civitai:
195
+ prompt = source[:civitai.start()].strip(" ,\n")
196
+ tail = source[civitai.end():]
197
+ settings = re.search(r"\b(?:Steps|Size)\s*:", tail, re.I)
198
+ negative = tail[:settings.start()].strip(" ,\n") if settings else tail.strip(" ,\n")
199
+ else:
200
+ original = re.search(r"\bOriginal\s+Prompt\s*\n+(.+?)(?=\n+\s*Size\s*\n)", source, re.I | re.S)
201
+ prompt = (original.group(1) if original else source.split("\n\n", 1)[0]).strip(" ,\n")
202
+ negative_match = re.search(r"\n\s*Negative\s*\n+(.+?)(?=\n\s*(?:Prompt\s+Helper|#|$))", source, re.I | re.S)
203
+ negative = negative_match.group(1).strip(" ,\n") if negative_match else ""
204
+
205
+ def number(pattern: str, kind):
206
+ match = re.search(pattern, source, re.I)
207
+ return kind(match.group(1)) if match else None
208
+ steps = number(r"\b(?:Sampling\s+)?Steps\s*:?\s*(\d+)", int)
209
+ cfg = number(r"\bCFG\s*(?:Scale)?\s*:?\s*(\d+(?:\.\d+)?)", float)
210
+ seed = number(r"\bSeed\s*:?\s*(-?\d+)", int)
211
+ if seed == -1:
212
+ seed = 0
213
+ size = re.search(r"\bSize\s*:?\s*(\d+)\s*[x×]\s*(\d+)", source, re.I)
214
+ sampler_match = re.search(r"\b(?:Sampling\s+Method|Sampler)\s*:?\s*([^\n,]+)", source, re.I)
215
+ sampler = sampler_match.group(1).strip() if sampler_match else ""
216
+ if sampler:
217
+ folded = sampler.casefold()
218
+ sampler = next((name for name in ("DPM++ 2M SDE Karras", "DPM++ 2M SDE", "DPM++ 2M Karras", "DPM++ SDE Karras", "DPM++ SDE", "DPM++ 2M") if name.casefold() in folded), sampler)
219
+ loras = re.findall(r"<lora:([^:>]+)(?::([\d.]+))?>", prompt, re.I)
220
+ if loras:
221
+ prompt = re.sub(r"\s*<lora:[^>]+>", "", prompt, flags=re.I).strip(" ,")
222
+ return ChatGenerationRequest(
223
+ prompt=prompt, provider_hint="lora", model_query=loras[0][0].strip() if len(loras) == 1 else "",
224
+ negative_prompt=negative, steps=steps, seed=seed, sampler=sampler, cfg_scale=cfg,
225
+ lora_strength=float(loras[0][1]) if len(loras) == 1 and loras[0][1] else None,
226
+ width=int(size.group(1)) if size else None, height=int(size.group(2)) if size else None,
227
+ has_positive_prompt=True, is_pasted_metadata=True,
228
+ )
229
+
230
+
231
  @dataclass(frozen=True, slots=True)
232
  class ChatGenerationRequest:
233
  """Generation settings recognized from a Command Center message."""
 
249
  reference_strength: int | None = None
250
  reference_image: str = ""
251
  has_positive_prompt: bool = False
252
+ is_pasted_metadata: bool = False
253
+ metadata_model_path: str = ""
254
+ metadata_base_model_path: str = ""
255
+ width: int | None = None
256
+ height: int | None = None
257
 
258
 
259
  _QUOTED = r'["\u201c\u201d]([^"\u201c\u201d]+)["\u201c\u201d]'
 
267
  """Score whether conversational subject text clearly names a saved model."""
268
  def words(value: str) -> list[str]:
269
  value = re.sub(r"(?<=[a-z0-9])(?=[A-Z])", " ", value)
270
+ value = re.sub(r"\bpixel\s+row\b", " ", value, flags=re.I)
271
+ value = re.sub(r"\binr\s*flow\b", " ", value, flags=re.I)
272
+ ignored = {"a", "an", "the", "of", "image", "picture", "model", "ddpm", "flow", "matching", "lora", "inr"}
273
  return [word for word in re.findall(r"[a-z0-9]+", value.casefold()) if word not in ignored]
274
 
275
  query_words = words(query)
 
295
  This intentionally requires both a creation verb and the word image/picture so
296
  ordinary planning requests continue through the regular Command Center planner.
297
  """
298
+ pasted_metadata = parse_pasted_lora_metadata_request(text)
299
+ if pasted_metadata is not None:
300
+ return pasted_metadata
301
+ plain_metadata = parse_plain_generation_metadata(text)
302
+ if plain_metadata is not None:
303
+ return plain_metadata
304
  request = " ".join(text.strip().split())
305
+ # Do not treat any request that happens to contain both words as an image
306
+ # generation command. Dataset requests commonly say things such as
307
+ # "image mode" and "generate captions"; those must continue to the
308
+ # regular planner (and, in particular, the video dataset collector).
309
+ # Require the creation verb to directly introduce the image noun instead.
310
+ generation_command = re.compile(
311
+ r"\b(?:generate|create|make)\s+"
312
+ r"(?:(?:an?|the|\d+)\s+)?"
313
+ r"(?:[\"\u201c\u201d]?(?:ddpm|ddim|inr\s*flow|flow(?:\s+matching)?|pixel\s*row|lora)[\"\u201c\u201d]?\s+)?"
314
+ r"(?:images?|pictures?)\b",
315
+ re.I,
316
+ )
317
+ if not request or not generation_command.search(request):
318
  return None
319
 
320
  provider_hint = ""
321
  provider_match = re.search(
322
+ r"\b(ddpm|ddim|inr\s*flow|flow(?:\s+matching)?|pixel\s*row|lora)\b[\"\u201c\u201d]?(?=\s+(?:image|picture))",
323
  request,
324
  re.I,
325
  )
326
  if provider_match:
327
  hint = provider_match.group(1).casefold()
328
+ provider_hint = (
329
+ "ddpm" if hint in {"ddpm", "ddim"}
330
+ else "inrflow" if hint.replace(" ", "") == "inrflow"
331
+ else "flow" if hint.startswith("flow")
332
+ else "pixelrow" if hint.replace(" ", "") == "pixelrow"
333
+ else "lora"
334
+ )
335
  # Support natural phrasing such as "Generate an image of LoRA OrangeCat".
336
  lora_subject_match = re.search(
337
  rf"\b(?:image|picture)s?\s+of\s+(?:a\s+)?LoRA\s+{_QUOTED}",
 
367
  # In promptless commands, a provider suffix is usually part of the saved
368
  # model name (for example, "Minecraft Flow"), not prompt prose.
369
  if not provider_hint and subject:
370
+ if re.search(r"\binr\s*flow\s*$", subject, re.I):
371
+ provider_hint = "inrflow"
372
+ elif re.search(r"\bflow(?:\s+match(?:ing)?)?\s*$", subject, re.I):
373
  provider_hint = "flow"
374
  elif re.search(r"\bddpm\s*$", subject, re.I):
375
  provider_hint = "ddpm"
376
+ elif re.search(r"\bpixel\s*row\s*$", subject, re.I):
377
+ provider_hint = "pixelrow"
378
 
379
  positive_match = re.search(
380
  rf"\bpositive\s+prompt(?:\s+of|\s*=|\s*:)?\s*{_QUOTED}",
 
534
  latest_at: str
535
 
536
 
537
+ @dataclass(frozen=True, slots=True)
538
+ class GenerationProviderFolder:
539
+ """A generator-centered view over the on-disk generation folders.
540
+
541
+ Output is stored as ``generations/<generator>/<model>/...``. This view
542
+ intentionally exposes that first directory level in the UI, keeping all
543
+ images made by one generator together without moving any user files.
544
+ """
545
+
546
+ key: str
547
+ provider_id: str
548
+ provider_name: str
549
+ records: tuple[GenerationRecord, ...]
550
+ model_count: int
551
+ image_count: int
552
+ latest_at: str
553
+
554
+
555
+ def generation_provider_key(record: GenerationRecord) -> str:
556
+ """Return the stable key for the generator directory containing a batch."""
557
+ provider_id = str(record.provider_id or "").strip()
558
+ if not provider_id:
559
+ # Metadata written by older versions may not have a provider id. Its
560
+ # parent is still the generator directory in the current file layout.
561
+ provider_id = record.folder.parent.name
562
+ return provider_id.casefold()
563
+
564
+
565
+ def group_generation_providers(
566
+ records: list[GenerationRecord],
567
+ ) -> list[GenerationProviderFolder]:
568
+ """Build newest-first generator folders from existing generation records."""
569
+ grouped: dict[str, list[GenerationRecord]] = {}
570
+ for record in records:
571
+ grouped.setdefault(generation_provider_key(record), []).append(record)
572
+ folders: list[GenerationProviderFolder] = []
573
+ for key, provider_records in grouped.items():
574
+ newest_first = sorted(
575
+ provider_records,
576
+ key=lambda item: item.created_at or item.folder.name,
577
+ reverse=True,
578
+ )
579
+ latest = newest_first[0]
580
+ provider_id = str(latest.provider_id or latest.folder.parent.name)
581
+ provider_name = str(latest.provider_name or provider_id)
582
+ folders.append(
583
+ GenerationProviderFolder(
584
+ key=key,
585
+ provider_id=provider_id,
586
+ provider_name=provider_name,
587
+ records=tuple(newest_first),
588
+ model_count=len({generation_model_key(record) for record in newest_first}),
589
+ image_count=sum(len(record.images) for record in newest_first),
590
+ latest_at=latest.created_at,
591
+ )
592
+ )
593
+ folders.sort(key=lambda item: (item.latest_at, item.provider_name.casefold()), reverse=True)
594
+ return folders
595
+
596
+
597
  def generation_model_key(record: GenerationRecord) -> str:
598
  """Keep renamed or duplicated display names separated by model identity."""
599
  raw_path = str(record.model_path or "").strip()
 
653
  history_root = root.resolve() / "data" / "generations"
654
  if not history_root.is_dir():
655
  return []
656
+ # History can grow into thousands of image batches. Sort inexpensive file
657
+ # metadata first, then decode only the newest records requested by the UI.
658
+ # This keeps a page refresh responsive without moving or rewriting history.
659
+ try:
660
+ metadata_paths = sorted(
661
+ history_root.rglob("generation*.json"),
662
+ key=lambda path: path.stat().st_mtime,
663
+ reverse=True,
664
+ )
665
+ except OSError:
666
+ metadata_paths = list(history_root.rglob("generation*.json"))
667
+ records = []
668
+ for metadata_path in metadata_paths:
669
+ record = GenerationRecord.from_metadata(metadata_path)
670
+ if record is not None:
671
+ records.append(record)
672
+ if len(records) >= max(1, int(limit)):
673
+ break
674
  records.sort(key=lambda item: item.created_at or item.folder.name, reverse=True)
675
+ return records
676
 
677
 
678
  def build_generation_plan(
adam/intelligence.py ADDED
@@ -0,0 +1,200 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Evidence-led summaries for ADAM's saved training and generation history.
2
+
3
+ This module intentionally stays independent of the Qt interface so its advice can
4
+ be tested and reused by a future remote surface. It does not assess a model's
5
+ absolute quality: it identifies useful follow-up experiments from the evidence
6
+ ADAM has recorded locally.
7
+ """
8
+ from __future__ import annotations
9
+
10
+ import json
11
+ from dataclasses import dataclass
12
+ from pathlib import Path
13
+
14
+ from adam.experiment_tracker import ExperimentRun
15
+ from adam.generations import GenerationRecord
16
+ from adam.image_preferences import PreferenceProfile
17
+
18
+
19
+ @dataclass(frozen=True, slots=True)
20
+ class ModelIntelligence:
21
+ key: str
22
+ model_name: str
23
+ architecture: str
24
+ runs: tuple[ExperimentRun, ...]
25
+ generations: tuple[GenerationRecord, ...]
26
+ rated_images: int
27
+ positive_ratings: int
28
+ rejected_images: int
29
+ state: str
30
+ diagnosis: str
31
+ recommendation: str
32
+ recommended_epochs: int
33
+
34
+
35
+ def _normalized(value: str) -> str:
36
+ return " ".join(str(value or "").casefold().split())
37
+
38
+
39
+ def _paths_match(left: str, right: str, cache: dict[str, Path] | None = None) -> bool:
40
+ if not left or not right:
41
+ return False
42
+ try:
43
+ cache = cache if cache is not None else {}
44
+ for value in (left, right):
45
+ if value not in cache:
46
+ cache[value] = Path(value).expanduser().resolve()
47
+ left_path, right_path = cache[left], cache[right]
48
+ return left_path == right_path or left_path in right_path.parents or right_path in left_path.parents
49
+ except OSError:
50
+ return _normalized(left) == _normalized(right)
51
+
52
+
53
+ def _generation_matches(run: ExperimentRun, record: GenerationRecord, cache: dict[str, Path] | None = None) -> bool:
54
+ if _paths_match(run.output_folder, record.model_path, cache):
55
+ return True
56
+ if any(_paths_match(path, record.model_path, cache) for path in run.checkpoint_paths):
57
+ return True
58
+ return _normalized(run.model_name) == _normalized(record.model_name)
59
+
60
+
61
+ def _ratings(records: tuple[GenerationRecord, ...], root: Path | None = None) -> tuple[int, int, int]:
62
+ rated = positive = rejected = 0
63
+ seen: set[str] = set()
64
+ preferences_by_model: dict[tuple[str, str], PreferenceProfile] = {}
65
+ image_keys: dict[Path, str] = {}
66
+ for record in records:
67
+ preferences = None
68
+ if root:
69
+ key = (record.provider_id, record.model_path)
70
+ if key not in preferences_by_model:
71
+ preferences_by_model[key] = PreferenceProfile(root, record.provider_id, record.model_name, record.model_path)
72
+ preferences = preferences_by_model[key]
73
+ for image in record.images:
74
+ if image not in image_keys:
75
+ image_keys[image] = str(image.expanduser().resolve())
76
+ image_key = image_keys[image]
77
+ if image_key in seen:
78
+ continue
79
+ seen.add(image_key)
80
+ evaluation = record.image_evaluations.get(image_key, {})
81
+ saved_rating = preferences.rating_for(image) if preferences else None
82
+ rating = saved_rating.rating if saved_rating else str(evaluation.get("rating", "")).casefold()
83
+ # Preference ratings are persisted separately today. The evaluator
84
+ # score still counts as review evidence when it is available.
85
+ score = evaluation.get("score")
86
+ if rating or isinstance(score, (int, float)):
87
+ rated += 1
88
+ if rating in {"favorite", "keep"} or (isinstance(score, (int, float)) and score >= 0.70):
89
+ positive += 1
90
+ if rating == "reject" or (isinstance(score, (int, float)) and score <= 0.35):
91
+ rejected += 1
92
+ return rated, positive, rejected
93
+
94
+
95
+ def _diagnose(
96
+ runs: tuple[ExperimentRun, ...], generations: tuple[GenerationRecord, ...], rated: int, positive: int, rejected: int,
97
+ ) -> tuple[str, str, str, int]:
98
+ latest = runs[0]
99
+ finished = [run for run in runs if run.status.casefold() == "finished"]
100
+ epoch_budget = max(1, latest.epochs)
101
+ previous = runs[1] if len(runs) > 1 else None
102
+
103
+ if not finished:
104
+ return (
105
+ "Needs a completed run",
106
+ "ADAM has not recorded a finished training run for this model yet, so it cannot judge training behavior.",
107
+ "Finish one run and generate a small, fixed-prompt test batch before changing several settings at once.",
108
+ epoch_budget,
109
+ )
110
+ if rated == 0:
111
+ return (
112
+ "Needs visual review",
113
+ "Training history exists, but there are no scored test generations linked to this model. Loss alone cannot tell ADAM which checkpoint you prefer.",
114
+ "Generate 4–8 images with one repeatable prompt and seed, then rate the results in Generations before starting a follow-up.",
115
+ epoch_budget,
116
+ )
117
+ if rejected > positive and rated >= 3:
118
+ return (
119
+ "Review data or settings",
120
+ f"{rejected} of {rated} reviewed generated images were rejected or scored low. More epochs by themselves are unlikely to be the best first change.",
121
+ "Review dataset variety, captions, and the fixed-prompt gallery. Keep the epoch count similar for the next controlled test, changing only one training setting.",
122
+ epoch_budget,
123
+ )
124
+ if previous and latest.final_loss is not None and previous.final_loss is not None:
125
+ loss_change = latest.final_loss - previous.final_loss
126
+ if abs(loss_change) <= max(0.0001, abs(previous.final_loss) * 0.03):
127
+ return (
128
+ "Likely plateau",
129
+ "The last two recorded losses changed very little. That is a plateau signal, not proof that the model has stopped improving visually.",
130
+ "Run a shorter follow-up (about 25% fewer epochs) with a lower learning rate or improved data; compare it using the same evaluation prompt and seed.",
131
+ max(1, round(epoch_budget * 0.75)),
132
+ )
133
+ if positive >= max(2, rejected * 2):
134
+ return (
135
+ "Promising",
136
+ f"{positive} reviewed generated images look positive versus {rejected} rejected. The model has enough signal for a focused continuation test.",
137
+ "Preserve this run as a baseline. Try a modest continuation of about 25% more epochs, then compare the same prompt-and-seed gallery before committing further.",
138
+ max(epoch_budget + 1, round(epoch_budget * 1.25)),
139
+ )
140
+ return (
141
+ "Gather one more comparison",
142
+ "ADAM has mixed review evidence. A single outcome can be affected by prompt choice, seed, or dataset coverage.",
143
+ "Make another small fixed-prompt generation batch, rate it, and change only one setting in the next run so the result is interpretable.",
144
+ epoch_budget,
145
+ )
146
+
147
+
148
+ def build_model_intelligence(
149
+ runs: list[ExperimentRun], generations: list[GenerationRecord], *, root: Path | None = None,
150
+ ) -> list[ModelIntelligence]:
151
+ """Group local records into newest-first, model-centered intelligence cards."""
152
+ grouped: dict[tuple[str, str], list[ExperimentRun]] = {}
153
+ for run in runs:
154
+ key = (_normalized(run.model_name), _normalized(run.model_architecture))
155
+ if key[0]:
156
+ grouped.setdefault(key, []).append(run)
157
+
158
+ profiles: list[ModelIntelligence] = []
159
+ path_cache: dict[str, Path] = {}
160
+ for (name_key, architecture_key), raw_runs in grouped.items():
161
+ model_runs = tuple(sorted(raw_runs, key=lambda run: run.timestamp, reverse=True))
162
+ model_generations = tuple(
163
+ record for record in generations
164
+ if any(_generation_matches(run, record, path_cache) for run in model_runs)
165
+ )
166
+ rated, positive, rejected = _ratings(model_generations, root)
167
+ state, diagnosis, recommendation, epochs = _diagnose(
168
+ model_runs, model_generations, rated, positive, rejected
169
+ )
170
+ profiles.append(ModelIntelligence(
171
+ key=f"{architecture_key}:{name_key}",
172
+ model_name=model_runs[0].model_name,
173
+ architecture=model_runs[0].model_architecture,
174
+ runs=model_runs,
175
+ generations=model_generations,
176
+ rated_images=rated,
177
+ positive_ratings=positive,
178
+ rejected_images=rejected,
179
+ state=state,
180
+ diagnosis=diagnosis,
181
+ recommendation=recommendation,
182
+ recommended_epochs=epochs,
183
+ ))
184
+ return sorted(profiles, key=lambda profile: profile.runs[0].timestamp, reverse=True)
185
+
186
+
187
+ def recommended_training_request(profile: ModelIntelligence) -> str:
188
+ """Create an approval-aware follow-up request using the latest run as a baseline."""
189
+ run = profile.runs[0]
190
+ options = {
191
+ key: value for key, value in run.settings.items()
192
+ if key not in {"dataset_dir", "model_name", "epochs", "output_dir", "resume_from"}
193
+ }
194
+ return (
195
+ f"From the {run.dataset_name or run.dataset_path} dataset, train a "
196
+ f"{run.model_architecture.upper()} model for {profile.recommended_epochs} epochs. "
197
+ f"Name the model {run.model_name} Follow-up. "
198
+ "[ADAM_TRAINING_OPTIONS:" + json.dumps(options, sort_keys=True) + "] "
199
+ "[ADAM_TRAINER:" + run.model_architecture + "]"
200
+ )
adam/job_manager.py CHANGED
@@ -73,7 +73,7 @@ class JobWorker(QThread):
73
 
74
  preview_state = {"epoch": 0, "path": ""}
75
  last_progress_emit = {"time": 0.0, "overall": -1, "message": ""}
76
- progress_samples: list[dict[str, float]] = []
77
 
78
  def on_progress(percent: int, message: str, step_index: int = index, **details: Any) -> None:
79
  overall = int(((step_index + percent / 100) / total_steps) * 100)
@@ -175,7 +175,7 @@ class JobWorker(QThread):
175
  @staticmethod
176
  def _estimate_step_eta(
177
  details: dict[str, Any],
178
- samples: list[dict[str, float]],
179
  now: float,
180
  ) -> dict[str, Any]:
181
  """Estimate remaining runtime from real step cadence instead of percent alone."""
@@ -184,6 +184,8 @@ class JobWorker(QThread):
184
  unit = str(details.get("unit", "step") or "step")
185
  epoch = _safe_int(details.get("epoch"))
186
  total_epochs = _safe_int(details.get("total_epochs"))
 
 
187
  if (not current or not total) and epoch and total_epochs:
188
  current, total, unit = epoch, total_epochs, "epoch"
189
  payload: dict[str, Any] = {
@@ -194,9 +196,18 @@ class JobWorker(QThread):
194
  if not current or not total or current >= total:
195
  return payload
196
  last = samples[-1] if samples else None
 
 
 
 
 
 
 
 
 
197
  if last and current <= last["current"]:
198
  return payload
199
- samples.append({"time": now, "current": float(current)})
200
  del samples[:-25]
201
  if len(samples) < 2:
202
  return payload
@@ -380,6 +391,10 @@ class JobManager(QObject):
380
  step = job.plan.steps[job.current_step]
381
  if step.tool_id != "ddpm_trainer":
382
  raise ValueError("Safe epoch-boundary adjustment currently supports DDPM training.")
 
 
 
 
383
  allowed = {"batch_size", "training_intensity", "gradient_accumulation_steps"}
384
  cleaned = {key: int(value) for key, value in updates.items() if key in allowed}
385
  if not cleaned or not 1 <= cleaned.get("batch_size", 1) <= 64 \
 
73
 
74
  preview_state = {"epoch": 0, "path": ""}
75
  last_progress_emit = {"time": 0.0, "overall": -1, "message": ""}
76
+ progress_samples: list[dict[str, Any]] = []
77
 
78
  def on_progress(percent: int, message: str, step_index: int = index, **details: Any) -> None:
79
  overall = int(((step_index + percent / 100) / total_steps) * 100)
 
175
  @staticmethod
176
  def _estimate_step_eta(
177
  details: dict[str, Any],
178
+ samples: list[dict[str, Any]],
179
  now: float,
180
  ) -> dict[str, Any]:
181
  """Estimate remaining runtime from real step cadence instead of percent alone."""
 
184
  unit = str(details.get("unit", "step") or "step")
185
  epoch = _safe_int(details.get("epoch"))
186
  total_epochs = _safe_int(details.get("total_epochs"))
187
+ if details.get("reset_eta"):
188
+ samples.clear()
189
  if (not current or not total) and epoch and total_epochs:
190
  current, total, unit = epoch, total_epochs, "epoch"
191
  payload: dict[str, Any] = {
 
196
  if not current or not total or current >= total:
197
  return payload
198
  last = samples[-1] if samples else None
199
+ if last and (
200
+ int(last.get("total", total)) != total
201
+ or str(last.get("unit", unit)) != unit
202
+ or current < last["current"]
203
+ ):
204
+ # A different tqdm operation or a restarted counter needs a fresh
205
+ # cadence; carrying the old rate creates wildly incorrect ETAs.
206
+ samples.clear()
207
+ last = None
208
  if last and current <= last["current"]:
209
  return payload
210
+ samples.append({"time": now, "current": float(current), "total": float(total), "unit": unit})
211
  del samples[:-25]
212
  if len(samples) < 2:
213
  return payload
 
391
  step = job.plan.steps[job.current_step]
392
  if step.tool_id != "ddpm_trainer":
393
  raise ValueError("Safe epoch-boundary adjustment currently supports DDPM training.")
394
+ if step.arguments.get("progressive_stages"):
395
+ raise ValueError(
396
+ "Change batch settings before starting a progressive run; each stage manages its own saved handoff."
397
+ )
398
  allowed = {"batch_size", "training_intensity", "gradient_accumulation_steps"}
399
  cleaned = {key: int(value) for key, value in updates.items() if key in allowed}
400
  if not cleaned or not 1 <= cleaned.get("batch_size", 1) <= 64 \
adam/model_plugins.py CHANGED
@@ -237,7 +237,9 @@ class ModelPluginRegistry:
237
  "demo": False,
238
  }
239
  defaults.update(tool)
240
- defaults["arguments"] = list(defaults.get("arguments") or [*core_arguments, *list(schema)])
 
 
241
  defaults["required_arguments"] = list(defaults.get("required_arguments") or [])
242
  return defaults
243
 
 
237
  "demo": False,
238
  }
239
  defaults.update(tool)
240
+ defaults["arguments"] = list(dict.fromkeys(
241
+ defaults.get("arguments") or [*core_arguments, *list(schema)]
242
+ ))
243
  defaults["required_arguments"] = list(defaults.get("required_arguments") or [])
244
  return defaults
245
 
adam/model_plugins_builtin/ddpm/manifest.py CHANGED
@@ -16,7 +16,9 @@ MODEL_INFO = {
16
  }
17
 
18
  TRAINING_SETTINGS = {
19
- "resolution": {"label": "Resolution", "type": "choice", "options": [64, 128, 256, 384, 512], "default": 128, "group": "Basic"},
 
 
20
  "batch_size": {"label": "Batch size", "type": "int", "default": 1, "min": 1, "max": 64, "group": "Basic"},
21
  "learning_rate": {"label": "Learning rate", "type": "float", "default": 0.0001, "min": 0.0000001, "max": 0.1, "decimals": 7, "step": 0.00005, "group": "Optimization"},
22
  "gradient_accumulation_steps": {"label": "Gradient accumulation", "type": "int", "default": 1, "min": 1, "max": 64, "group": "Optimization"},
 
16
  }
17
 
18
  TRAINING_SETTINGS = {
19
+ "resolution": {"label": "Longest edge", "type": "choice", "options": [64, 128, 256, 384, 512], "default": 128, "group": "Basic", "description": "The longest side of the native training canvas."},
20
+ "training_aspect_ratio": {"label": "Training aspect ratio", "type": "choice", "options": ["Dataset (Auto)", "1:1 (Square)", "16:9 (Widescreen)", "9:16 (Portrait)", "4:3 (Classic)", "3:4 (Portrait Classic)", "3:2 (Photo)", "2:3 (Portrait Photo)"], "default": "Dataset (Auto)", "group": "Basic", "description": "Dataset Auto uses the median source-image aspect ratio; 256 with 16:9 creates a 256x144 model."},
21
+ "resize_mode": {"label": "Image fitting", "type": "choice", "options": ["fit", "fill", "stretch"], "default": "fit", "group": "Dataset", "description": "fit preserves the entire image and edge-pads only when needed; fill crops; stretch changes proportions."},
22
  "batch_size": {"label": "Batch size", "type": "int", "default": 1, "min": 1, "max": 64, "group": "Basic"},
23
  "learning_rate": {"label": "Learning rate", "type": "float", "default": 0.0001, "min": 0.0000001, "max": 0.1, "decimals": 7, "step": 0.00005, "group": "Optimization"},
24
  "gradient_accumulation_steps": {"label": "Gradient accumulation", "type": "int", "default": 1, "min": 1, "max": 64, "group": "Optimization"},
adam/model_plugins_builtin/flow_matching/manifest.py CHANGED
@@ -37,6 +37,8 @@ GENERATION_SETTINGS = {
37
  "steps": {"label": "ODE steps", "type": "int", "default": 20, "min": 1, "max": 200, "group": "Generation"},
38
  "sampler": {"label": "Method", "type": "choice", "options": ["Heun", "Euler"], "default": "Heun", "group": "Generation"},
39
  "aspect_ratio": {"label": "Aspect ratio", "type": "choice", "options": ["1:1 (Square)", "4:3 (Landscape)", "3:4 (Portrait)", "3:2 (Landscape)", "2:3 (Portrait)", "16:9 (Widescreen)", "9:16 (Vertical)"], "default": "1:1 (Square)", "group": "Generation"},
 
 
40
  "seed": {"label": "Seed", "type": "int", "default": 0, "min": 0, "max": 2147483647, "group": "Generation"},
41
  "preview_interval": {"label": "Steps per preview", "type": "int", "default": 0, "min": 0, "max": 500, "group": "Preview"},
42
  "smart_generation": {"label": "Smart Generation", "type": "bool", "default": False, "group": "Smart Generation"},
 
37
  "steps": {"label": "ODE steps", "type": "int", "default": 20, "min": 1, "max": 200, "group": "Generation"},
38
  "sampler": {"label": "Method", "type": "choice", "options": ["Heun", "Euler"], "default": "Heun", "group": "Generation"},
39
  "aspect_ratio": {"label": "Aspect ratio", "type": "choice", "options": ["1:1 (Square)", "4:3 (Landscape)", "3:4 (Portrait)", "3:2 (Landscape)", "2:3 (Portrait)", "16:9 (Widescreen)", "9:16 (Vertical)"], "default": "1:1 (Square)", "group": "Generation"},
40
+ "width": {"label": "Width", "type": "int", "default": 0, "min": 0, "max": 2048, "group": "Generation"},
41
+ "height": {"label": "Height", "type": "int", "default": 0, "min": 0, "max": 2048, "group": "Generation"},
42
  "seed": {"label": "Seed", "type": "int", "default": 0, "min": 0, "max": 2147483647, "group": "Generation"},
43
  "preview_interval": {"label": "Steps per preview", "type": "int", "default": 0, "min": 0, "max": 500, "group": "Preview"},
44
  "smart_generation": {"label": "Smart Generation", "type": "bool", "default": False, "group": "Smart Generation"},
adam/model_plugins_builtin/inrflow/__init__.py ADDED
@@ -0,0 +1,2 @@
 
 
 
1
+ """ADAM-sized INRFlow image model plugin."""
2
+
adam/model_plugins_builtin/inrflow/common.py ADDED
@@ -0,0 +1,106 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from __future__ import annotations
2
+
3
+ import json
4
+ import re
5
+ from pathlib import Path
6
+ from typing import Any
7
+
8
+ import torch
9
+ from PIL import Image
10
+
11
+ from adam.executor import ToolExecutionError
12
+
13
+ from .model import MODEL_FORMAT_VERSION, INRFlowConfig, INRFlowModel
14
+
15
+
16
+ IMAGE_EXTENSIONS = {".jpg", ".jpeg", ".png", ".webp", ".bmp"}
17
+ FINAL_CHECKPOINT_NAME = "inrflow_model.pt"
18
+ CONFIG_NAME = "inrflow_config.json"
19
+
20
+
21
+ def safe_model_name(value: str) -> str:
22
+ name = re.sub(r"\s+", " ", value.strip())
23
+ if not name or len(name) > 96 or any(character in name for character in '<>:"/\\|?*\x00'):
24
+ raise ToolExecutionError(
25
+ "Choose a short INRFlow model name without reserved filename characters."
26
+ )
27
+ return name
28
+
29
+
30
+ def ensure_below(path: Path, root: Path, label: str) -> Path:
31
+ resolved = path.expanduser().resolve()
32
+ try:
33
+ resolved.relative_to(root.expanduser().resolve())
34
+ except ValueError as exc:
35
+ raise ToolExecutionError(f"{label} must stay inside {root.resolve()}.") from exc
36
+ return resolved
37
+
38
+
39
+ def image_files(folder: Path) -> list[Path]:
40
+ try:
41
+ return sorted(
42
+ path
43
+ for path in folder.rglob("*")
44
+ if path.is_file() and path.suffix.casefold() in IMAGE_EXTENSIONS
45
+ )
46
+ except OSError:
47
+ return []
48
+
49
+
50
+ def resolve_checkpoint(path: Path) -> Path:
51
+ candidate = path.expanduser().resolve()
52
+ if candidate.is_dir():
53
+ candidate = candidate / FINAL_CHECKPOINT_NAME
54
+ if not candidate.is_file():
55
+ raise ToolExecutionError(
56
+ "The selected INRFlow checkpoint does not exist or is incomplete."
57
+ )
58
+ return candidate
59
+
60
+
61
+ def load_checkpoint(
62
+ path: Path,
63
+ device: torch.device,
64
+ *,
65
+ prefer_ema: bool = True,
66
+ ) -> tuple[INRFlowModel, dict[str, Any]]:
67
+ checkpoint_path = resolve_checkpoint(path)
68
+ try:
69
+ payload = torch.load(checkpoint_path, map_location=device, weights_only=True)
70
+ except (OSError, RuntimeError, ValueError, TypeError) as exc:
71
+ raise ToolExecutionError(f"Could not load the INRFlow checkpoint: {exc}") from exc
72
+ if not isinstance(payload, dict) or "model_state" not in payload or "config" not in payload:
73
+ raise ToolExecutionError("The selected file is not a valid INRFlow checkpoint.")
74
+ if int(payload.get("format_version", 0)) != MODEL_FORMAT_VERSION:
75
+ raise ToolExecutionError("This INRFlow checkpoint uses an unsupported format version.")
76
+ try:
77
+ config = INRFlowConfig.from_dict(dict(payload["config"]))
78
+ model = INRFlowModel(config).to(device)
79
+ state = payload.get("ema_state") if prefer_ema else None
80
+ model.load_state_dict(state if isinstance(state, dict) else payload["model_state"], strict=True)
81
+ except (KeyError, TypeError, ValueError, RuntimeError) as exc:
82
+ raise ToolExecutionError(f"The INRFlow checkpoint is incompatible: {exc}") from exc
83
+ return model, payload
84
+
85
+
86
+ def save_image(image: torch.Tensor, path: Path) -> None:
87
+ pixels = (
88
+ image.detach()
89
+ .float()
90
+ .cpu()
91
+ .clamp(-1.0, 1.0)
92
+ .add(1.0)
93
+ .mul(127.5)
94
+ .round()
95
+ .to(torch.uint8)
96
+ .numpy()
97
+ )
98
+ path.parent.mkdir(parents=True, exist_ok=True)
99
+ Image.fromarray(pixels, mode="RGB").save(path, format="PNG")
100
+
101
+
102
+ def write_json(path: Path, payload: dict[str, Any]) -> None:
103
+ path.parent.mkdir(parents=True, exist_ok=True)
104
+ temporary = path.with_suffix(path.suffix + ".tmp")
105
+ temporary.write_text(json.dumps(payload, indent=2), encoding="utf-8")
106
+ temporary.replace(path)
adam/model_plugins_builtin/inrflow/generator.py ADDED
@@ -0,0 +1,305 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from __future__ import annotations
2
+
3
+ import random
4
+ from datetime import datetime, timezone
5
+ from pathlib import Path
6
+ from typing import Any
7
+
8
+ import torch
9
+
10
+ from adam.executor import ToolExecutionError
11
+ from adam.generations import generation_metadata_path, generation_output_folder
12
+ from adam.image_preferences import GenerationPreferenceEvaluator, PreferenceProfile
13
+
14
+ from .common import ensure_below, load_checkpoint, safe_model_name, save_image, write_json
15
+ from .model import sample_image
16
+
17
+
18
+ def generate(
19
+ context,
20
+ model_name: str,
21
+ model_path: str,
22
+ prompt: str,
23
+ image_count: int,
24
+ steps: int,
25
+ seed: int,
26
+ sampler: str,
27
+ aspect_ratio: str,
28
+ output_resolution: str = "Native",
29
+ noise_scale: float = 1.0,
30
+ query_chunk_size: int = 1024,
31
+ preview_interval: int = 5,
32
+ smart_generation: bool = False,
33
+ smart_wanted_results: int = 8,
34
+ smart_max_candidates: int = 32,
35
+ smart_min_score: float = 0.7,
36
+ smart_mode: str = "threshold",
37
+ smart_keep_rejected: bool = True,
38
+ ) -> dict[str, Any]:
39
+ """Generate images by integrating the learned ambient-space velocity field."""
40
+ name = safe_model_name(model_name)
41
+ model_root = (
42
+ context.root.resolve() / "data" / "model_plugin_outputs" / "inrflow"
43
+ ).resolve()
44
+ selected = ensure_below(Path(model_path), model_root, "INRFlow model")
45
+ if not selected.exists():
46
+ raise ToolExecutionError("The selected INRFlow model no longer exists.")
47
+ count = int(image_count)
48
+ step_count = int(steps)
49
+ if not 1 <= count <= 48:
50
+ raise ToolExecutionError("INRFlow image count must be between 1 and 48.")
51
+ if not 2 <= step_count <= 200:
52
+ raise ToolExecutionError("INRFlow ODE steps must be between 2 and 200.")
53
+ method = sampler.strip().title()
54
+ if method not in {"Euler", "Heun"}:
55
+ raise ToolExecutionError("INRFlow supports the Euler and Heun ODE methods.")
56
+ if aspect_ratio != "1:1 (Coordinate Field)":
57
+ raise ToolExecutionError("INRFlow currently generates square coordinate fields.")
58
+ if not 0.1 <= float(noise_scale) <= 2.0:
59
+ raise ToolExecutionError("INRFlow starting noise scale must be between 0.1 and 2.0.")
60
+ if int(query_chunk_size) not in {256, 512, 1024, 2048, 4096}:
61
+ raise ToolExecutionError("Choose a supported INRFlow query chunk size.")
62
+ if not 0 <= int(preview_interval) <= step_count:
63
+ raise ToolExecutionError("Preview interval must be between 0 and the ODE step count.")
64
+ if len(prompt) > 500:
65
+ raise ToolExecutionError("The INRFlow creative note must be 500 characters or shorter.")
66
+
67
+ smart_enabled = bool(smart_generation)
68
+ wanted_results = int(smart_wanted_results or count)
69
+ max_candidates = int(smart_max_candidates or count)
70
+ threshold = float(smart_min_score)
71
+ top_n_mode = str(smart_mode).casefold() == "top_n"
72
+ if smart_enabled:
73
+ if not 1 <= wanted_results <= 48:
74
+ raise ToolExecutionError("Wanted Smart Generation results must be between 1 and 48.")
75
+ if not wanted_results <= max_candidates <= 256:
76
+ raise ToolExecutionError(
77
+ "Maximum Smart Generation candidates must be between wanted results and 256."
78
+ )
79
+ if not 0.0 <= threshold <= 1.0:
80
+ raise ToolExecutionError("Minimum Smart Generation score must be between 0 and 1.")
81
+ if str(smart_mode).casefold() not in {"threshold", "top_n"}:
82
+ raise ToolExecutionError("Smart Generation mode must be threshold or top_n.")
83
+ count = wanted_results
84
+
85
+ device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
86
+ model, checkpoint = load_checkpoint(selected, device, prefer_ema=True)
87
+ if str(output_resolution) == "Native":
88
+ resolution = model.config.resolution
89
+ else:
90
+ try:
91
+ resolution = int(output_resolution)
92
+ except (TypeError, ValueError) as exc:
93
+ raise ToolExecutionError("Choose Native or a supported INRFlow output resolution.") from exc
94
+ if resolution not in {32, 64, 128, 256}:
95
+ raise ToolExecutionError("INRFlow output resolution must be 32, 64, 128, or 256.")
96
+ if resolution % model.config.patch_size:
97
+ raise ToolExecutionError(
98
+ "That output resolution is not divisible by this model's spatial latent patch."
99
+ )
100
+ if resolution != model.config.resolution:
101
+ context.log(
102
+ f"Querying the learned coordinate field at {resolution}px; it was trained at "
103
+ f"{model.config.resolution}px, so this is resolution extrapolation."
104
+ )
105
+
106
+ generated_total = max_candidates if smart_enabled else count
107
+ base_seed = int(seed)
108
+ if base_seed <= 0:
109
+ base_seed = random.SystemRandom().randint(
110
+ 1, 2_147_483_647 - generated_total
111
+ )
112
+ if base_seed + generated_total - 1 > 2_147_483_647:
113
+ raise ToolExecutionError("The INRFlow seed is too large for this image count.")
114
+
115
+ output = generation_output_folder(context.root, context.tool.id, name)
116
+ timestamp = datetime.now(timezone.utc).strftime("%Y%m%d_%H%M%S")
117
+ image_paths: list[str] = []
118
+ selected_paths: list[str] = []
119
+ image_evaluations: dict[str, dict[str, object]] = {}
120
+ profile = (
121
+ PreferenceProfile(context.root, context.tool.id, name, str(selected))
122
+ if smart_enabled
123
+ else None
124
+ )
125
+ evaluator = GenerationPreferenceEvaluator(context.root) if smart_enabled else None
126
+ context.log(
127
+ f"Loaded {name}: trained at {model.config.resolution}px, generating at "
128
+ f"{resolution}px with {method} on {device}."
129
+ )
130
+ context.log(
131
+ "INRFlow is unconditional; the creative note is saved with the result but is not a text prompt."
132
+ )
133
+
134
+ try:
135
+ for image_index in range(generated_total):
136
+ context.checkpoint()
137
+ current_seed = base_seed + image_index
138
+ generator = torch.Generator(device=device)
139
+ generator.manual_seed(current_seed)
140
+ live_path = output / ".live" / context.job_id / f"image_{image_index + 1:03d}.png"
141
+
142
+ def on_step(done: int, total: int, image: torch.Tensor) -> None:
143
+ context.checkpoint()
144
+ overall = (image_index + done / max(1, total)) / generated_total
145
+ context.progress(
146
+ max(1, min(99, round(overall * 100))),
147
+ f"Image {image_index + 1} of {generated_total} · ODE step {done} of {total}",
148
+ current=done,
149
+ total=total,
150
+ image_index=image_index,
151
+ image_count=generated_total,
152
+ unit="step",
153
+ )
154
+ if int(preview_interval) > 0 and (
155
+ done % int(preview_interval) == 0 or done == total
156
+ ):
157
+ save_image(image, live_path)
158
+ context.preview(
159
+ live_path,
160
+ kind="generation",
161
+ current=done,
162
+ total=total,
163
+ image_index=image_index,
164
+ image_count=generated_total,
165
+ seed=current_seed,
166
+ steps=step_count,
167
+ )
168
+
169
+ image = sample_image(
170
+ model,
171
+ resolution=resolution,
172
+ steps=step_count,
173
+ method=method,
174
+ noise_scale=float(noise_scale),
175
+ query_chunk_size=int(query_chunk_size),
176
+ generator=generator,
177
+ step_callback=on_step,
178
+ )
179
+ destination = output / (
180
+ f"{timestamp}_{context.job_id}_INRFlow_{method}_seed_{current_seed}_"
181
+ f"{resolution}px.png"
182
+ )
183
+ save_image(image, destination)
184
+ image_paths.append(str(destination))
185
+ if smart_enabled and profile is not None and evaluator is not None:
186
+ score = evaluator.score(
187
+ profile,
188
+ [destination],
189
+ keep_threshold=threshold,
190
+ reject_threshold=profile.reject_threshold,
191
+ )[0]
192
+ image_evaluations[str(destination.resolve())] = {
193
+ "score": score.score,
194
+ "confidence": score.confidence,
195
+ "category": score.category,
196
+ "reason": score.reason,
197
+ }
198
+ if not top_n_mode and score.score is not None and score.score >= threshold:
199
+ selected_paths.append(str(destination))
200
+ if len(selected_paths) >= wanted_results:
201
+ break
202
+ except torch.cuda.OutOfMemoryError as exc:
203
+ if device.type == "cuda":
204
+ torch.cuda.empty_cache()
205
+ raise ToolExecutionError(
206
+ "INRFlow ran out of VRAM while generating. Lower output resolution or query chunk size."
207
+ ) from exc
208
+ finally:
209
+ if evaluator is not None:
210
+ evaluator.vision.unload()
211
+
212
+ if smart_enabled and top_n_mode:
213
+ ranked = sorted(
214
+ image_paths,
215
+ key=lambda path: float(
216
+ image_evaluations.get(str(Path(path).resolve()), {}).get("score") or -1.0
217
+ ),
218
+ reverse=True,
219
+ )
220
+ selected_paths = ranked[:wanted_results]
221
+ if smart_enabled:
222
+ chosen = set(selected_paths)
223
+ ordered_images = [*selected_paths, *[path for path in image_paths if path not in chosen]]
224
+ saved_images = ordered_images if bool(smart_keep_rejected) else selected_paths
225
+ else:
226
+ ordered_images = image_paths
227
+ saved_images = image_paths
228
+
229
+ metadata = {
230
+ "version": 1,
231
+ "provider_id": context.tool.id,
232
+ "provider_name": context.tool.name,
233
+ "model_name": name,
234
+ "model_path": str(selected),
235
+ "model_type": "inrflow",
236
+ "architecture": "inrflow_ambient_space",
237
+ "prompt": prompt.strip(),
238
+ "prompt_behavior": "label_only",
239
+ "seed": base_seed,
240
+ "image_seeds": [base_seed + index for index in range(len(image_paths))],
241
+ "image_count": len(saved_images),
242
+ "steps": step_count,
243
+ "sampler": method,
244
+ "aspect_ratio": aspect_ratio,
245
+ "training_resolution": model.config.resolution,
246
+ "output_resolution": resolution,
247
+ "noise_scale": float(noise_scale),
248
+ "query_chunk_size": int(query_chunk_size),
249
+ "preview_interval": int(preview_interval),
250
+ "images": saved_images,
251
+ "image_evaluations": image_evaluations,
252
+ "checkpoint_epoch": int(checkpoint.get("completed_epochs", 0) or 0),
253
+ "used_ema_weights": isinstance(checkpoint.get("ema_state"), dict),
254
+ "uses_pretrained_compressor": False,
255
+ "smart_generation": {
256
+ "enabled": smart_enabled,
257
+ "mode": str(smart_mode),
258
+ "wanted_results": wanted_results if smart_enabled else count,
259
+ "maximum_candidates": max_candidates if smart_enabled else count,
260
+ "minimum_score": threshold,
261
+ "selected_count": len(selected_paths) if smart_enabled else count,
262
+ "candidate_count": len(image_paths),
263
+ "profile_id": profile.id if profile else "",
264
+ "keep_rejected_candidates": bool(smart_keep_rejected),
265
+ },
266
+ "created_at": datetime.now(timezone.utc).isoformat(),
267
+ }
268
+ write_json(generation_metadata_path(output, timestamp, context.job_id), metadata)
269
+ live_folder = output / ".live" / context.job_id
270
+ if live_folder.is_dir():
271
+ for path in live_folder.glob("*.png"):
272
+ try:
273
+ path.unlink()
274
+ except OSError:
275
+ pass
276
+ try:
277
+ live_folder.rmdir()
278
+ live_folder.parent.rmdir()
279
+ except OSError:
280
+ pass
281
+
282
+ if smart_enabled:
283
+ context.progress(
284
+ 100,
285
+ f"Smart Generation selected {len(selected_paths)} of {wanted_results} requested "
286
+ f"image(s) from {len(image_paths)} candidate(s)",
287
+ )
288
+ else:
289
+ context.progress(100, f"Generated {count} INRFlow image(s)")
290
+ return {
291
+ "output_folder": str(output),
292
+ "assets": [
293
+ {
294
+ "kind": "generation",
295
+ "name": f"{name} · {timestamp}",
296
+ "path": str(output),
297
+ "trainer": "inrflow",
298
+ "metadata": {
299
+ "resolution": resolution,
300
+ "steps": step_count,
301
+ "sampler": method,
302
+ },
303
+ }
304
+ ],
305
+ }
adam/model_plugins_builtin/inrflow/manifest.py ADDED
@@ -0,0 +1,394 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ PLUGIN_ID = "inrflow"
2
+
3
+ MODEL_INFO = {
4
+ "name": "INRFlow (Ambient Space)",
5
+ "version": "0.1",
6
+ "category": "Image Generation",
7
+ "description": (
8
+ "ADAM-sized INRFlow: coordinate-to-RGB flow matching directly in pixel space, "
9
+ "with spatial context latents and no pretrained image compressor."
10
+ ),
11
+ "architecture": "inrflow_ambient_space",
12
+ "status": "experimental",
13
+ "output_type": "image",
14
+ "capabilities": [
15
+ "fresh_training",
16
+ "resume_training",
17
+ "image_generation",
18
+ "smart_generation",
19
+ "live_preview",
20
+ "resolution_flexible_generation",
21
+ ],
22
+ "input_formats": ["image folder"],
23
+ "output_formats": ["INRFlow checkpoint", "INRFlow metadata", "PNG preview"],
24
+ "hardware": {"recommended_vram_gb": 8, "recommended_system_ram_gb": 16},
25
+ "vram_behavior": {
26
+ "scales_with": [
27
+ "resolution", "batch_size", "hidden_size", "depth",
28
+ "query_points", "sampling_steps",
29
+ ],
30
+ "estimate": (
31
+ "Designed for 64px experiments on 8–12 GB GPUs. At 128px, reduce batch "
32
+ "size and query points before shrinking the model."
33
+ ),
34
+ },
35
+ "method_reference": "https://arxiv.org/abs/2412.03791",
36
+ "reference_implementation": "https://github.com/apple/ml-inrflow",
37
+ }
38
+
39
+ TRAINING_SETTINGS = {
40
+ "resolution": {
41
+ "label": "Training resolution",
42
+ "type": "choice",
43
+ "options": [32, 64, 128, 256],
44
+ "default": 64,
45
+ "group": "Basic",
46
+ "description": "Start at 64px for an architecture comparison on an RTX 3060.",
47
+ },
48
+ "resize_mode": {
49
+ "label": "Image fitting",
50
+ "type": "choice",
51
+ "options": ["fill", "fit", "stretch"],
52
+ "default": "fill",
53
+ "group": "Dataset",
54
+ },
55
+ "horizontal_flip": {
56
+ "label": "Random horizontal flip",
57
+ "type": "bool",
58
+ "default": True,
59
+ "group": "Dataset",
60
+ },
61
+ "batch_size": {
62
+ "label": "Batch size",
63
+ "type": "int",
64
+ "default": 4,
65
+ "min": 1,
66
+ "max": 32,
67
+ "group": "Basic",
68
+ },
69
+ "learning_rate": {
70
+ "label": "Learning rate",
71
+ "type": "float",
72
+ "default": 0.0001,
73
+ "min": 0.0000001,
74
+ "max": 0.1,
75
+ "decimals": 7,
76
+ "step": 0.00005,
77
+ "group": "Optimization",
78
+ },
79
+ "weight_decay": {
80
+ "label": "Weight decay",
81
+ "type": "float",
82
+ "default": 0.0,
83
+ "min": 0.0,
84
+ "max": 1.0,
85
+ "decimals": 5,
86
+ "step": 0.001,
87
+ "group": "Optimization",
88
+ "advanced": True,
89
+ },
90
+ "gradient_accumulation_steps": {
91
+ "label": "Gradient accumulation",
92
+ "type": "int",
93
+ "default": 1,
94
+ "min": 1,
95
+ "max": 64,
96
+ "group": "Optimization",
97
+ },
98
+ "workers": {
99
+ "label": "Loader workers",
100
+ "type": "int",
101
+ "default": 0,
102
+ "min": 0,
103
+ "max": 16,
104
+ "group": "Dataset",
105
+ },
106
+ "mixed_precision": {
107
+ "label": "Precision",
108
+ "type": "choice",
109
+ "options": ["fp16", "bf16", "no"],
110
+ "default": "fp16",
111
+ "group": "Optimization",
112
+ },
113
+ "patch_size": {
114
+ "label": "Spatial latent patch",
115
+ "type": "choice",
116
+ "options": [4, 8, 16],
117
+ "default": 8,
118
+ "group": "INRFlow",
119
+ "description": "Each spatial context latent attends to the coordinate-value pairs in one patch.",
120
+ },
121
+ "hidden_size": {
122
+ "label": "Transformer width",
123
+ "type": "choice",
124
+ "options": [128, 192, 256, 384],
125
+ "default": 256,
126
+ "group": "INRFlow",
127
+ "advanced": True,
128
+ },
129
+ "depth": {
130
+ "label": "Transformer layers",
131
+ "type": "choice",
132
+ "options": [2, 4, 6, 8],
133
+ "default": 4,
134
+ "group": "INRFlow",
135
+ "advanced": True,
136
+ },
137
+ "num_heads": {
138
+ "label": "Attention heads",
139
+ "type": "choice",
140
+ "options": [4, 8],
141
+ "default": 8,
142
+ "group": "INRFlow",
143
+ "advanced": True,
144
+ },
145
+ "decoder_layers": {
146
+ "label": "Point decoder layers",
147
+ "type": "choice",
148
+ "options": [1, 2],
149
+ "default": 1,
150
+ "group": "INRFlow",
151
+ "advanced": True,
152
+ },
153
+ "fourier_frequencies": {
154
+ "label": "Coordinate frequencies",
155
+ "type": "choice",
156
+ "options": [4, 6, 8, 10],
157
+ "default": 8,
158
+ "group": "INRFlow",
159
+ "advanced": True,
160
+ },
161
+ "query_points": {
162
+ "label": "Pixel queries per image",
163
+ "type": "choice",
164
+ "options": [256, 512, 1024, 2048, 4096],
165
+ "default": 1024,
166
+ "group": "INRFlow",
167
+ "description": "Point-wise subsampling is a defining INRFlow training advantage.",
168
+ },
169
+ "time_sampling": {
170
+ "label": "Flow-time sampling",
171
+ "type": "choice",
172
+ "options": ["logit_normal", "uniform"],
173
+ "default": "logit_normal",
174
+ "group": "INRFlow",
175
+ "advanced": True,
176
+ },
177
+ "ema_decay": {
178
+ "label": "EMA decay",
179
+ "type": "float",
180
+ "default": 0.999,
181
+ "min": 0.9,
182
+ "max": 0.99999,
183
+ "decimals": 5,
184
+ "step": 0.0001,
185
+ "group": "Optimization",
186
+ "advanced": True,
187
+ },
188
+ "save_every": {
189
+ "label": "Save every",
190
+ "type": "int",
191
+ "default": 10,
192
+ "min": 1,
193
+ "max": 1000,
194
+ "group": "Checkpoints",
195
+ },
196
+ "preview_enabled": {
197
+ "label": "Generate previews while training",
198
+ "type": "bool",
199
+ "default": True,
200
+ "group": "Preview",
201
+ },
202
+ "preview_every": {
203
+ "label": "Preview interval",
204
+ "type": "int",
205
+ "default": 5,
206
+ "min": 1,
207
+ "max": 100000,
208
+ "group": "Preview",
209
+ },
210
+ "preview_steps": {
211
+ "label": "Preview flow steps",
212
+ "type": "int",
213
+ "default": 20,
214
+ "min": 2,
215
+ "max": 200,
216
+ "group": "Preview",
217
+ },
218
+ "preview_prompt": {
219
+ "label": "Preview note",
220
+ "type": "text",
221
+ "default": "",
222
+ "group": "Preview",
223
+ },
224
+ "preview_seed": {
225
+ "label": "Preview seed",
226
+ "type": "int",
227
+ "default": 123456789,
228
+ "min": 0,
229
+ "max": 2147483647,
230
+ "group": "Preview",
231
+ },
232
+ }
233
+
234
+ GENERATION_SETTINGS = {
235
+ "prompt": {
236
+ "label": "Creative note",
237
+ "type": "multiline_text",
238
+ "default": "",
239
+ "group": "Generation",
240
+ },
241
+ "image_count": {
242
+ "label": "Images",
243
+ "type": "int",
244
+ "default": 1,
245
+ "min": 1,
246
+ "max": 48,
247
+ "group": "Generation",
248
+ },
249
+ "steps": {
250
+ "label": "ODE steps",
251
+ "type": "int",
252
+ "default": 50,
253
+ "min": 2,
254
+ "max": 200,
255
+ "group": "Generation",
256
+ },
257
+ "sampler": {
258
+ "label": "ODE method",
259
+ "type": "choice",
260
+ "options": ["Euler", "Heun"],
261
+ "default": "Euler",
262
+ "group": "Generation",
263
+ },
264
+ "aspect_ratio": {
265
+ "label": "Aspect ratio",
266
+ "type": "choice",
267
+ "options": ["1:1 (Coordinate Field)"],
268
+ "default": "1:1 (Coordinate Field)",
269
+ "group": "Generation",
270
+ },
271
+ "seed": {
272
+ "label": "Seed",
273
+ "type": "int",
274
+ "default": 0,
275
+ "min": 0,
276
+ "max": 2147483647,
277
+ "group": "Generation",
278
+ },
279
+ "output_resolution": {
280
+ "label": "Output resolution",
281
+ "type": "choice",
282
+ "options": ["Native", "32", "64", "128", "256"],
283
+ "default": "Native",
284
+ "group": "Coordinate Field",
285
+ "description": "INRFlow can query the learned coordinate field at a different resolution.",
286
+ },
287
+ "noise_scale": {
288
+ "label": "Starting noise scale",
289
+ "type": "float",
290
+ "default": 1.0,
291
+ "min": 0.1,
292
+ "max": 2.0,
293
+ "decimals": 2,
294
+ "step": 0.05,
295
+ "group": "Generation",
296
+ },
297
+ "query_chunk_size": {
298
+ "label": "Query chunk size",
299
+ "type": "choice",
300
+ "options": [256, 512, 1024, 2048, 4096],
301
+ "default": 1024,
302
+ "group": "Advanced",
303
+ "advanced": True,
304
+ "description": "Reduce this if resolution-flexible generation runs out of VRAM.",
305
+ },
306
+ "preview_interval": {
307
+ "label": "Steps per live preview",
308
+ "type": "int",
309
+ "default": 5,
310
+ "min": 0,
311
+ "max": 200,
312
+ "group": "Preview",
313
+ },
314
+ "smart_generation": {
315
+ "label": "Smart Generation",
316
+ "type": "bool",
317
+ "default": False,
318
+ "group": "Smart Generation",
319
+ },
320
+ "smart_wanted_results": {
321
+ "label": "Wanted results",
322
+ "type": "int",
323
+ "default": 8,
324
+ "min": 1,
325
+ "max": 48,
326
+ "group": "Smart Generation",
327
+ },
328
+ "smart_max_candidates": {
329
+ "label": "Maximum candidates",
330
+ "type": "int",
331
+ "default": 32,
332
+ "min": 1,
333
+ "max": 256,
334
+ "group": "Smart Generation",
335
+ },
336
+ "smart_min_score": {
337
+ "label": "Minimum score",
338
+ "type": "float",
339
+ "default": 0.7,
340
+ "min": 0,
341
+ "max": 1,
342
+ "group": "Smart Generation",
343
+ },
344
+ "smart_mode": {
345
+ "label": "Selection mode",
346
+ "type": "choice",
347
+ "options": ["threshold", "top_n"],
348
+ "default": "threshold",
349
+ "group": "Smart Generation",
350
+ },
351
+ "smart_keep_rejected": {
352
+ "label": "Keep rejected candidates",
353
+ "type": "bool",
354
+ "default": True,
355
+ "group": "Smart Generation",
356
+ },
357
+ }
358
+
359
+ TRAINING_TOOL = {
360
+ "id": "inrflow_trainer",
361
+ "name": "INRFlow Trainer",
362
+ "description": "Trains coordinate-to-RGB flow matching directly in ambient image space.",
363
+ "capabilities": ["fresh_training", "resume_training", "progress", "pause", "cancel", "live_preview"],
364
+ "backend": {
365
+ "type": "python",
366
+ "module": "adam.model_plugins_builtin.inrflow.trainer",
367
+ "function": "train",
368
+ },
369
+ }
370
+
371
+ GENERATION_TOOL = {
372
+ "id": "inrflow_generator",
373
+ "name": "INRFlow Generator",
374
+ "description": "Integrates an INRFlow coordinate field from Gaussian noise to an image.",
375
+ "model_trainers": ["inrflow"],
376
+ "capabilities": [
377
+ "image_generation", "smart_generation", "seed", "ode_method", "batch",
378
+ "resolution_flexible_generation", "live_preview", "progress", "cancel",
379
+ ],
380
+ "generation_options": {
381
+ "samplers": ["Euler", "Heun"],
382
+ "aspect_ratios": ["1:1 (Coordinate Field)"],
383
+ "step_min": 2,
384
+ "step_max": 200,
385
+ "step_default": 50,
386
+ "preview_step_default": 5,
387
+ },
388
+ "backend": {
389
+ "type": "python",
390
+ "module": "adam.model_plugins_builtin.inrflow.generator",
391
+ "function": "generate",
392
+ },
393
+ }
394
+
adam/model_plugins_builtin/inrflow/model.py ADDED
@@ -0,0 +1,400 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from __future__ import annotations
2
+
3
+ import math
4
+ from dataclasses import asdict, dataclass
5
+ from typing import Any, Callable
6
+
7
+ import torch
8
+ from torch import Tensor, nn
9
+ from torch.nn import functional as F
10
+
11
+
12
+ MODEL_FORMAT_VERSION = 1
13
+
14
+
15
+ @dataclass(frozen=True, slots=True)
16
+ class INRFlowConfig:
17
+ resolution: int = 64
18
+ patch_size: int = 8
19
+ hidden_size: int = 256
20
+ depth: int = 4
21
+ num_heads: int = 8
22
+ decoder_layers: int = 1
23
+ fourier_frequencies: int = 8
24
+
25
+ def __post_init__(self) -> None:
26
+ if self.resolution not in {32, 64, 128, 256}:
27
+ raise ValueError("INRFlow resolution must be 32, 64, 128, or 256.")
28
+ if self.patch_size not in {4, 8, 16} or self.resolution % self.patch_size:
29
+ raise ValueError("INRFlow patch size must be 4, 8, or 16 and divide the resolution.")
30
+ if not 64 <= self.hidden_size <= 768:
31
+ raise ValueError("INRFlow transformer width must be between 64 and 768.")
32
+ if not 1 <= self.depth <= 12 or not 1 <= self.decoder_layers <= 4:
33
+ raise ValueError("INRFlow transformer depth is outside the supported range.")
34
+ if self.num_heads not in {2, 4, 8, 16} or self.hidden_size % self.num_heads:
35
+ raise ValueError("INRFlow width must be divisible by its attention-head count.")
36
+ if not 2 <= self.fourier_frequencies <= 16:
37
+ raise ValueError("INRFlow coordinate frequencies must be between 2 and 16.")
38
+
39
+ def to_dict(self) -> dict[str, int]:
40
+ return asdict(self)
41
+
42
+ @classmethod
43
+ def from_dict(cls, payload: dict[str, Any]) -> "INRFlowConfig":
44
+ return cls(
45
+ resolution=int(payload.get("resolution", 64)),
46
+ patch_size=int(payload.get("patch_size", 8)),
47
+ hidden_size=int(payload.get("hidden_size", 256)),
48
+ depth=int(payload.get("depth", 4)),
49
+ num_heads=int(payload.get("num_heads", 8)),
50
+ decoder_layers=int(payload.get("decoder_layers", 1)),
51
+ fourier_frequencies=int(payload.get("fourier_frequencies", 8)),
52
+ )
53
+
54
+
55
+ def coordinate_grid(
56
+ height: int,
57
+ width: int,
58
+ *,
59
+ device: torch.device | str | None = None,
60
+ ) -> Tensor:
61
+ """Return normalized x/y coordinates as a flattened Nx2 field."""
62
+ y = torch.linspace(0.0, 1.0, height, device=device)
63
+ x = torch.linspace(0.0, 1.0, width, device=device)
64
+ yy, xx = torch.meshgrid(y, x, indexing="ij")
65
+ return torch.stack((xx, yy), dim=-1).reshape(height * width, 2)
66
+
67
+
68
+ class FourierCoordinates(nn.Module):
69
+ def __init__(self, frequencies: int) -> None:
70
+ super().__init__()
71
+ bands = torch.pow(2.0, torch.arange(frequencies, dtype=torch.float32)) * math.pi
72
+ self.register_buffer("bands", bands, persistent=False)
73
+ self.output_size = 2 + 4 * frequencies
74
+
75
+ def forward(self, coordinates: Tensor) -> Tensor:
76
+ phases = coordinates.unsqueeze(-1) * self.bands
77
+ return torch.cat(
78
+ (coordinates, phases.sin().flatten(-2), phases.cos().flatten(-2)), dim=-1
79
+ )
80
+
81
+
82
+ class TimeEmbedding(nn.Module):
83
+ def __init__(self, hidden_size: int, frequency_size: int = 64) -> None:
84
+ super().__init__()
85
+ self.frequency_size = frequency_size
86
+ self.mlp = nn.Sequential(
87
+ nn.Linear(frequency_size, hidden_size),
88
+ nn.SiLU(),
89
+ nn.Linear(hidden_size, hidden_size),
90
+ )
91
+
92
+ def forward(self, time: Tensor) -> Tensor:
93
+ half = self.frequency_size // 2
94
+ frequencies = torch.exp(
95
+ -math.log(10_000.0)
96
+ * torch.arange(half, device=time.device, dtype=torch.float32)
97
+ / max(1, half)
98
+ )
99
+ phases = time.float().unsqueeze(1) * frequencies.unsqueeze(0)
100
+ embedding = torch.cat((phases.cos(), phases.sin()), dim=1)
101
+ return self.mlp(embedding)
102
+
103
+
104
+ class PatchContextEncoder(nn.Module):
105
+ """Cross-attend one spatial latent to nearby coordinate/value pairs."""
106
+
107
+ def __init__(self, config: INRFlowConfig, coordinates: FourierCoordinates) -> None:
108
+ super().__init__()
109
+ hidden = config.hidden_size
110
+ self.patch_size = config.patch_size
111
+ self.coordinates = coordinates
112
+ self.point_projection = nn.Sequential(
113
+ nn.Linear(coordinates.output_size + 3, hidden),
114
+ nn.LayerNorm(hidden),
115
+ nn.SiLU(),
116
+ )
117
+ self.center_projection = nn.Linear(coordinates.output_size, hidden)
118
+ self.latent_seed = nn.Parameter(torch.randn(1, 1, hidden) * 0.02)
119
+ self.attention = nn.MultiheadAttention(
120
+ hidden, config.num_heads, batch_first=True
121
+ )
122
+ self.norm1 = nn.LayerNorm(hidden)
123
+ self.norm2 = nn.LayerNorm(hidden)
124
+ self.mlp = nn.Sequential(
125
+ nn.Linear(hidden, hidden * 2), nn.GELU(), nn.Linear(hidden * 2, hidden)
126
+ )
127
+
128
+ @staticmethod
129
+ def _patchify(values: Tensor, height: int, width: int, patch: int) -> Tensor:
130
+ batch, points, channels = values.shape
131
+ if points != height * width or height % patch or width % patch:
132
+ raise ValueError("INRFlow context field does not match its patch grid.")
133
+ return values.reshape(
134
+ batch, height // patch, patch, width // patch, patch, channels
135
+ ).permute(0, 1, 3, 2, 4, 5).reshape(
136
+ batch, (height // patch) * (width // patch), patch * patch, channels
137
+ )
138
+
139
+ def forward(
140
+ self,
141
+ context_coordinates: Tensor,
142
+ context_values: Tensor,
143
+ *,
144
+ height: int,
145
+ width: int,
146
+ ) -> tuple[Tensor, Tensor]:
147
+ batch = context_values.shape[0]
148
+ encoded_coordinates = self.coordinates(context_coordinates)
149
+ point_features = self.point_projection(
150
+ torch.cat((encoded_coordinates, context_values), dim=-1)
151
+ )
152
+ point_patches = self._patchify(
153
+ point_features, height, width, self.patch_size
154
+ )
155
+ coordinate_patches = self._patchify(
156
+ context_coordinates, height, width, self.patch_size
157
+ )
158
+ centers = coordinate_patches.mean(dim=2)
159
+ latent_queries = self.latent_seed + self.center_projection(
160
+ self.coordinates(centers)
161
+ )
162
+ latent_count = point_patches.shape[1]
163
+ queries = latent_queries.reshape(batch * latent_count, 1, -1)
164
+ points = point_patches.reshape(
165
+ batch * latent_count, self.patch_size * self.patch_size, -1
166
+ )
167
+ attended, _weights = self.attention(
168
+ queries, points, points, need_weights=False
169
+ )
170
+ latents = self.norm1(queries + attended)
171
+ latents = latents + self.mlp(self.norm2(latents))
172
+ return latents.reshape(batch, latent_count, -1), centers
173
+
174
+
175
+ def _modulate(value: Tensor, shift: Tensor, scale: Tensor) -> Tensor:
176
+ return value * (1.0 + scale.unsqueeze(1)) + shift.unsqueeze(1)
177
+
178
+
179
+ class TimeConditionedBlock(nn.Module):
180
+ def __init__(self, hidden_size: int, heads: int) -> None:
181
+ super().__init__()
182
+ self.norm1 = nn.LayerNorm(hidden_size, elementwise_affine=False)
183
+ self.norm2 = nn.LayerNorm(hidden_size, elementwise_affine=False)
184
+ self.attention = nn.MultiheadAttention(hidden_size, heads, batch_first=True)
185
+ self.mlp = nn.Sequential(
186
+ nn.Linear(hidden_size, hidden_size * 4),
187
+ nn.GELU(approximate="tanh"),
188
+ nn.Linear(hidden_size * 4, hidden_size),
189
+ )
190
+ self.modulation = nn.Sequential(nn.SiLU(), nn.Linear(hidden_size, hidden_size * 4))
191
+ nn.init.zeros_(self.modulation[-1].weight)
192
+ nn.init.zeros_(self.modulation[-1].bias)
193
+
194
+ def forward(self, latents: Tensor, time_embedding: Tensor) -> Tensor:
195
+ shift1, scale1, shift2, scale2 = self.modulation(time_embedding).chunk(4, dim=-1)
196
+ attended, _weights = self.attention(
197
+ _modulate(self.norm1(latents), shift1, scale1),
198
+ _modulate(self.norm1(latents), shift1, scale1),
199
+ _modulate(self.norm1(latents), shift1, scale1),
200
+ need_weights=False,
201
+ )
202
+ latents = latents + attended
203
+ return latents + self.mlp(_modulate(self.norm2(latents), shift2, scale2))
204
+
205
+
206
+ class QueryDecoderBlock(nn.Module):
207
+ def __init__(self, hidden_size: int, heads: int) -> None:
208
+ super().__init__()
209
+ self.query_norm = nn.LayerNorm(hidden_size)
210
+ self.latent_norm = nn.LayerNorm(hidden_size)
211
+ self.attention = nn.MultiheadAttention(hidden_size, heads, batch_first=True)
212
+ self.output_norm = nn.LayerNorm(hidden_size)
213
+ self.mlp = nn.Sequential(
214
+ nn.Linear(hidden_size, hidden_size * 2),
215
+ nn.GELU(approximate="tanh"),
216
+ nn.Linear(hidden_size * 2, hidden_size),
217
+ )
218
+
219
+ def forward(self, queries: Tensor, latents: Tensor) -> Tensor:
220
+ attended, _weights = self.attention(
221
+ self.query_norm(queries),
222
+ self.latent_norm(latents),
223
+ self.latent_norm(latents),
224
+ need_weights=False,
225
+ )
226
+ queries = queries + attended
227
+ return queries + self.mlp(self.output_norm(queries))
228
+
229
+
230
+ class INRFlowModel(nn.Module):
231
+ """Coordinate-query flow model following INRFlow's ambient-space structure."""
232
+
233
+ def __init__(self, config: INRFlowConfig) -> None:
234
+ super().__init__()
235
+ self.config = config
236
+ self.coordinate_embedding = FourierCoordinates(config.fourier_frequencies)
237
+ self.time_embedding = TimeEmbedding(config.hidden_size)
238
+ self.context_encoder = PatchContextEncoder(config, self.coordinate_embedding)
239
+ self.latent_coordinate_projection = nn.Linear(
240
+ self.coordinate_embedding.output_size, config.hidden_size
241
+ )
242
+ self.trunk = nn.ModuleList([
243
+ TimeConditionedBlock(config.hidden_size, config.num_heads)
244
+ for _ in range(config.depth)
245
+ ])
246
+ self.query_projection = nn.Sequential(
247
+ nn.Linear(self.coordinate_embedding.output_size + 3, config.hidden_size),
248
+ nn.LayerNorm(config.hidden_size),
249
+ nn.SiLU(),
250
+ )
251
+ self.decoder = nn.ModuleList([
252
+ QueryDecoderBlock(config.hidden_size, config.num_heads)
253
+ for _ in range(config.decoder_layers)
254
+ ])
255
+ self.output = nn.Sequential(
256
+ nn.LayerNorm(config.hidden_size), nn.Linear(config.hidden_size, 3)
257
+ )
258
+ nn.init.zeros_(self.output[-1].weight)
259
+ nn.init.zeros_(self.output[-1].bias)
260
+
261
+ def encode_context(
262
+ self,
263
+ context_coordinates: Tensor,
264
+ context_values: Tensor,
265
+ time: Tensor,
266
+ *,
267
+ height: int,
268
+ width: int,
269
+ ) -> tuple[Tensor, Tensor]:
270
+ latents, centers = self.context_encoder(
271
+ context_coordinates, context_values, height=height, width=width
272
+ )
273
+ time_embedding = self.time_embedding(time)
274
+ latents = latents + self.latent_coordinate_projection(
275
+ self.coordinate_embedding(centers)
276
+ )
277
+ for block in self.trunk:
278
+ latents = block(latents, time_embedding)
279
+ return latents, time_embedding
280
+
281
+ def decode_queries(
282
+ self,
283
+ latents: Tensor,
284
+ time_embedding: Tensor,
285
+ query_coordinates: Tensor,
286
+ query_values: Tensor,
287
+ ) -> Tensor:
288
+ queries = self.query_projection(torch.cat((
289
+ self.coordinate_embedding(query_coordinates), query_values
290
+ ), dim=-1))
291
+ queries = queries + time_embedding.unsqueeze(1)
292
+ for block in self.decoder:
293
+ queries = block(queries, latents)
294
+ return self.output(queries)
295
+
296
+ def forward(
297
+ self,
298
+ context_coordinates: Tensor,
299
+ context_values: Tensor,
300
+ time: Tensor,
301
+ query_coordinates: Tensor,
302
+ query_values: Tensor,
303
+ *,
304
+ height: int,
305
+ width: int,
306
+ ) -> Tensor:
307
+ latents, time_embedding = self.encode_context(
308
+ context_coordinates, context_values, time, height=height, width=width
309
+ )
310
+ return self.decode_queries(
311
+ latents, time_embedding, query_coordinates, query_values
312
+ )
313
+
314
+ @torch.inference_mode()
315
+ def velocity_field(
316
+ self,
317
+ coordinates: Tensor,
318
+ values: Tensor,
319
+ time: Tensor,
320
+ *,
321
+ height: int,
322
+ width: int,
323
+ query_chunk_size: int = 1024,
324
+ ) -> Tensor:
325
+ latents, time_embedding = self.encode_context(
326
+ coordinates, values, time, height=height, width=width
327
+ )
328
+ outputs = []
329
+ for start in range(0, coordinates.shape[1], query_chunk_size):
330
+ stop = min(coordinates.shape[1], start + query_chunk_size)
331
+ outputs.append(self.decode_queries(
332
+ latents,
333
+ time_embedding,
334
+ coordinates[:, start:stop],
335
+ values[:, start:stop],
336
+ ))
337
+ return torch.cat(outputs, dim=1)
338
+
339
+
340
+ @torch.inference_mode()
341
+ def sample_image(
342
+ model: INRFlowModel,
343
+ *,
344
+ resolution: int,
345
+ steps: int,
346
+ method: str,
347
+ noise_scale: float,
348
+ query_chunk_size: int,
349
+ generator: torch.Generator,
350
+ step_callback: Callable[[int, int, Tensor], None] | None = None,
351
+ ) -> Tensor:
352
+ """Integrate the learned velocity from Gaussian noise (t=0) to data (t=1)."""
353
+ if resolution % model.config.patch_size:
354
+ raise ValueError("Output resolution must be divisible by the trained patch size.")
355
+ model.eval()
356
+ device = next(model.parameters()).device
357
+ coordinates = coordinate_grid(resolution, resolution, device=device).unsqueeze(0)
358
+ values = torch.randn(
359
+ 1, resolution * resolution, 3, device=device, generator=generator
360
+ ) * float(noise_scale)
361
+ times = torch.linspace(0.0, 1.0, int(steps) + 1, device=device)
362
+ for index in range(int(steps)):
363
+ time = times[index].expand(1)
364
+ next_time = times[index + 1].expand(1)
365
+ delta = times[index + 1] - times[index]
366
+ with torch.autocast(
367
+ device_type=device.type,
368
+ dtype=torch.float16,
369
+ enabled=device.type == "cuda",
370
+ ):
371
+ first = model.velocity_field(
372
+ coordinates,
373
+ values,
374
+ time,
375
+ height=resolution,
376
+ width=resolution,
377
+ query_chunk_size=query_chunk_size,
378
+ )
379
+ if method == "Heun":
380
+ predicted = values + delta * first
381
+ second = model.velocity_field(
382
+ coordinates,
383
+ predicted,
384
+ next_time,
385
+ height=resolution,
386
+ width=resolution,
387
+ query_chunk_size=query_chunk_size,
388
+ )
389
+ if method == "Heun":
390
+ values = values + delta * 0.5 * (first + second)
391
+ else:
392
+ values = values + delta * first
393
+ if step_callback is not None:
394
+ image = values[0].reshape(resolution, resolution, 3).clamp(-1, 1)
395
+ step_callback(index + 1, int(steps), image)
396
+ return values[0].reshape(resolution, resolution, 3).clamp(-1, 1)
397
+
398
+
399
+ def parameter_count(model: nn.Module) -> int:
400
+ return sum(parameter.numel() for parameter in model.parameters())
adam/model_plugins_builtin/inrflow/trainer.py ADDED
@@ -0,0 +1,523 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from __future__ import annotations
2
+
3
+ import copy
4
+ import random
5
+ from contextlib import nullcontext
6
+ from datetime import datetime, timezone
7
+ from pathlib import Path
8
+ from typing import Any
9
+
10
+ import torch
11
+ from PIL import Image, ImageOps, ImageStat
12
+ from torch.utils.data import DataLoader, Dataset
13
+
14
+ from adam.executor import ToolExecutionError
15
+
16
+ from .common import (
17
+ CONFIG_NAME,
18
+ FINAL_CHECKPOINT_NAME,
19
+ ensure_below,
20
+ image_files,
21
+ load_checkpoint,
22
+ resolve_checkpoint,
23
+ safe_model_name,
24
+ save_image,
25
+ write_json,
26
+ )
27
+ from .model import (
28
+ MODEL_FORMAT_VERSION,
29
+ INRFlowConfig,
30
+ INRFlowModel,
31
+ coordinate_grid,
32
+ parameter_count,
33
+ sample_image,
34
+ )
35
+
36
+
37
+ class INRFlowImageDataset(Dataset[torch.Tensor]):
38
+ def __init__(
39
+ self,
40
+ paths: list[Path],
41
+ *,
42
+ resolution: int,
43
+ resize_mode: str,
44
+ horizontal_flip: bool,
45
+ ) -> None:
46
+ self.paths = paths
47
+ self.resolution = resolution
48
+ self.resize_mode = resize_mode
49
+ self.horizontal_flip = horizontal_flip
50
+
51
+ def __len__(self) -> int:
52
+ return len(self.paths)
53
+
54
+ def __getitem__(self, index: int) -> torch.Tensor:
55
+ path = self.paths[index]
56
+ try:
57
+ with Image.open(path) as opened:
58
+ image = opened.convert("RGB")
59
+ size = (self.resolution, self.resolution)
60
+ if self.resize_mode == "fill":
61
+ image = ImageOps.fit(image, size, method=Image.Resampling.LANCZOS)
62
+ elif self.resize_mode == "fit":
63
+ mean = tuple(
64
+ int(value) for value in ImageStat.Stat(image.resize((1, 1))).mean
65
+ )
66
+ image = ImageOps.pad(
67
+ image, size, method=Image.Resampling.LANCZOS, color=mean
68
+ )
69
+ else:
70
+ image = image.resize(size, Image.Resampling.LANCZOS)
71
+ if self.horizontal_flip and random.random() < 0.5:
72
+ image = image.transpose(Image.Transpose.FLIP_LEFT_RIGHT)
73
+ buffer = bytearray(image.tobytes())
74
+ except (OSError, ValueError) as exc:
75
+ raise RuntimeError(f"Could not read training image {path.name}: {exc}") from exc
76
+ pixels = torch.frombuffer(buffer, dtype=torch.uint8).reshape(
77
+ self.resolution, self.resolution, 3
78
+ )
79
+ return pixels.float().div(127.5).sub(1.0).permute(2, 0, 1)
80
+
81
+
82
+ def _checkpoint_payload(
83
+ model: INRFlowModel,
84
+ ema_model: INRFlowModel,
85
+ optimizer: torch.optim.Optimizer,
86
+ *,
87
+ model_name: str,
88
+ dataset_dir: Path,
89
+ completed_epochs: int,
90
+ global_step: int,
91
+ training_settings: dict[str, Any],
92
+ ) -> dict[str, Any]:
93
+ return {
94
+ "format_version": MODEL_FORMAT_VERSION,
95
+ "architecture": "inrflow_ambient_space",
96
+ "method": "conditionally_independent_continuous_flow_matching",
97
+ "model_name": model_name,
98
+ "config": model.config.to_dict(),
99
+ "model_state": model.state_dict(),
100
+ "ema_state": ema_model.state_dict(),
101
+ "optimizer_state": optimizer.state_dict(),
102
+ "completed_epochs": int(completed_epochs),
103
+ "global_step": int(global_step),
104
+ "dataset_dir": str(dataset_dir),
105
+ "training_settings": training_settings,
106
+ "saved_at": datetime.now(timezone.utc).isoformat(),
107
+ }
108
+
109
+
110
+ def _save_checkpoint(path: Path, payload: dict[str, Any]) -> None:
111
+ path.parent.mkdir(parents=True, exist_ok=True)
112
+ temporary = path.with_suffix(path.suffix + ".tmp")
113
+ torch.save(payload, temporary)
114
+ temporary.replace(path)
115
+
116
+
117
+ @torch.no_grad()
118
+ def _update_ema(ema_model: INRFlowModel, model: INRFlowModel, decay: float) -> None:
119
+ source = model.state_dict()
120
+ for name, value in ema_model.state_dict().items():
121
+ incoming = source[name]
122
+ if value.is_floating_point():
123
+ value.mul_(decay).add_(incoming, alpha=1.0 - decay)
124
+ else:
125
+ value.copy_(incoming)
126
+
127
+
128
+ def _preview(
129
+ context,
130
+ model: INRFlowModel,
131
+ output: Path,
132
+ *,
133
+ epoch: int,
134
+ next_epoch: int,
135
+ steps: int,
136
+ seed: int,
137
+ prompt: str,
138
+ ) -> None:
139
+ device = next(model.parameters()).device
140
+ generator = torch.Generator(device=device)
141
+ generator.manual_seed(int(seed))
142
+ image = sample_image(
143
+ model,
144
+ resolution=model.config.resolution,
145
+ steps=int(steps),
146
+ method="Euler",
147
+ noise_scale=1.0,
148
+ query_chunk_size=1024,
149
+ generator=generator,
150
+ step_callback=lambda _done, _total, _image: context.checkpoint(),
151
+ )
152
+ destination = output / "previews" / f"preview_epoch_{epoch:06d}.png"
153
+ save_image(image, destination)
154
+ context.preview(
155
+ destination,
156
+ epoch=epoch,
157
+ next_epoch=next_epoch,
158
+ prompt=prompt,
159
+ seed=seed,
160
+ steps=int(steps),
161
+ )
162
+
163
+
164
+ def train(
165
+ context,
166
+ dataset_dir: str,
167
+ model_name: str,
168
+ epochs: int,
169
+ output_dir: str,
170
+ resume_from: str = "",
171
+ resolution: int = 64,
172
+ resize_mode: str = "fill",
173
+ horizontal_flip: bool = True,
174
+ batch_size: int = 4,
175
+ learning_rate: float = 0.0001,
176
+ weight_decay: float = 0.0,
177
+ gradient_accumulation_steps: int = 1,
178
+ workers: int = 0,
179
+ mixed_precision: str = "fp16",
180
+ patch_size: int = 8,
181
+ hidden_size: int = 256,
182
+ depth: int = 4,
183
+ num_heads: int = 8,
184
+ decoder_layers: int = 1,
185
+ fourier_frequencies: int = 8,
186
+ query_points: int = 1024,
187
+ time_sampling: str = "logit_normal",
188
+ ema_decay: float = 0.999,
189
+ save_every: int = 10,
190
+ preview_enabled: bool = True,
191
+ preview_every: int = 5,
192
+ preview_steps: int = 20,
193
+ preview_prompt: str = "",
194
+ preview_seed: int = 123456789,
195
+ ) -> dict[str, Any]:
196
+ """Train an ADAM-sized INRFlow model directly on RGB coordinate fields."""
197
+ name = safe_model_name(model_name)
198
+ dataset = Path(dataset_dir).expanduser().resolve()
199
+ if not dataset.is_dir():
200
+ raise ToolExecutionError("The selected INRFlow dataset folder no longer exists.")
201
+ training_dataset = dataset
202
+ accepted_frames = dataset / "frames"
203
+ frame_paths = image_files(accepted_frames) if accepted_frames.is_dir() else []
204
+ if frame_paths:
205
+ training_dataset, paths = accepted_frames, frame_paths
206
+ else:
207
+ paths = image_files(dataset)
208
+ if len(paths) < 2:
209
+ raise ToolExecutionError(
210
+ "INRFlow needs at least two readable image files before training can start."
211
+ )
212
+
213
+ output_root = (
214
+ context.root.resolve() / "data" / "model_plugin_outputs" / "inrflow"
215
+ ).resolve()
216
+ output = ensure_below(Path(output_dir), output_root, "INRFlow output")
217
+ if output.exists() and not output.is_dir():
218
+ raise ToolExecutionError("The INRFlow output path must be a folder.")
219
+ if output.exists() and any(output.iterdir()):
220
+ raise ToolExecutionError(
221
+ "The INRFlow output folder is not empty. Choose a new model output."
222
+ )
223
+ output.mkdir(parents=True, exist_ok=True)
224
+
225
+ if resize_mode not in {"fill", "fit", "stretch"}:
226
+ raise ToolExecutionError("INRFlow image fitting must be fill, fit, or stretch.")
227
+ if not 1 <= int(epochs) <= 100_000:
228
+ raise ToolExecutionError("INRFlow epochs must be between 1 and 100000.")
229
+ if not 1 <= int(batch_size) <= 32 or not 1 <= int(gradient_accumulation_steps) <= 64:
230
+ raise ToolExecutionError("INRFlow batch size or gradient accumulation is invalid.")
231
+ if not 1e-7 <= float(learning_rate) <= 0.1 or not 0.0 <= float(weight_decay) <= 1.0:
232
+ raise ToolExecutionError("INRFlow learning rate or weight decay is invalid.")
233
+ if not 0 <= int(workers) <= 16 or mixed_precision not in {"fp16", "bf16", "no"}:
234
+ raise ToolExecutionError("INRFlow loader workers or precision is invalid.")
235
+ if int(query_points) < 1 or int(query_points) > int(resolution) ** 2:
236
+ raise ToolExecutionError("Pixel queries cannot exceed the number of training pixels.")
237
+ if time_sampling not in {"logit_normal", "uniform"}:
238
+ raise ToolExecutionError("INRFlow time sampling must be logit_normal or uniform.")
239
+ if not 0.9 <= float(ema_decay) <= 0.99999:
240
+ raise ToolExecutionError("INRFlow EMA decay must be between 0.9 and 0.99999.")
241
+ if not 1 <= int(save_every) <= 1000 or not 1 <= int(preview_every) <= 100_000:
242
+ raise ToolExecutionError("INRFlow save and preview intervals must be positive.")
243
+ if not 2 <= int(preview_steps) <= 200:
244
+ raise ToolExecutionError("INRFlow preview steps must be between 2 and 200.")
245
+
246
+ device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
247
+ resume_payload: dict[str, Any] | None = None
248
+ if resume_from.strip():
249
+ resume_path = ensure_below(
250
+ resolve_checkpoint(Path(resume_from)), output_root, "INRFlow resume checkpoint"
251
+ )
252
+ model, resume_payload = load_checkpoint(resume_path, device, prefer_ema=False)
253
+ config = model.config
254
+ context.log(
255
+ "Continuing with the checkpoint architecture: "
256
+ f"{config.resolution}px, width {config.hidden_size}, depth {config.depth}."
257
+ )
258
+ else:
259
+ try:
260
+ config = INRFlowConfig(
261
+ resolution=int(resolution),
262
+ patch_size=int(patch_size),
263
+ hidden_size=int(hidden_size),
264
+ depth=int(depth),
265
+ num_heads=int(num_heads),
266
+ decoder_layers=int(decoder_layers),
267
+ fourier_frequencies=int(fourier_frequencies),
268
+ )
269
+ except ValueError as exc:
270
+ raise ToolExecutionError(str(exc)) from exc
271
+ model = INRFlowModel(config).to(device)
272
+
273
+ if int(query_points) > config.resolution ** 2:
274
+ raise ToolExecutionError(
275
+ "Pixel queries cannot exceed the resumed model's training resolution."
276
+ )
277
+ ema_model = copy.deepcopy(model).to(device).eval()
278
+ if resume_payload is not None and isinstance(resume_payload.get("ema_state"), dict):
279
+ try:
280
+ ema_model.load_state_dict(resume_payload["ema_state"], strict=True)
281
+ except RuntimeError:
282
+ context.log("The previous EMA weights were incompatible; EMA restarted from the model.")
283
+
284
+ dataset_object = INRFlowImageDataset(
285
+ paths,
286
+ resolution=config.resolution,
287
+ resize_mode=resize_mode,
288
+ horizontal_flip=bool(horizontal_flip),
289
+ )
290
+ loader = DataLoader(
291
+ dataset_object,
292
+ batch_size=int(batch_size),
293
+ shuffle=True,
294
+ num_workers=int(workers),
295
+ pin_memory=device.type == "cuda",
296
+ drop_last=False,
297
+ )
298
+ optimizer = torch.optim.AdamW(
299
+ model.parameters(),
300
+ lr=float(learning_rate),
301
+ betas=(0.9, 0.95),
302
+ weight_decay=float(weight_decay),
303
+ )
304
+ start_epoch = 0
305
+ global_step = 0
306
+ if resume_payload is not None:
307
+ start_epoch = int(resume_payload.get("completed_epochs", 0) or 0)
308
+ global_step = int(resume_payload.get("global_step", 0) or 0)
309
+ if isinstance(resume_payload.get("optimizer_state"), dict):
310
+ try:
311
+ optimizer.load_state_dict(resume_payload["optimizer_state"])
312
+ for group in optimizer.param_groups:
313
+ group["lr"] = float(learning_rate)
314
+ group["weight_decay"] = float(weight_decay)
315
+ except (ValueError, RuntimeError):
316
+ context.log("The old optimizer state was incompatible; using a fresh optimizer.")
317
+
318
+ use_fp16 = mixed_precision == "fp16" and device.type == "cuda"
319
+ use_bf16 = (
320
+ mixed_precision == "bf16"
321
+ and device.type == "cuda"
322
+ and bool(getattr(torch.cuda, "is_bf16_supported", lambda: False)())
323
+ )
324
+ if mixed_precision != "no" and not (use_fp16 or use_bf16):
325
+ context.log(f"{mixed_precision.upper()} is unavailable here; INRFlow will use full precision.")
326
+ autocast_dtype = torch.bfloat16 if use_bf16 else torch.float16
327
+ try:
328
+ scaler = torch.amp.GradScaler("cuda", enabled=use_fp16)
329
+ except (AttributeError, TypeError):
330
+ scaler = torch.cuda.amp.GradScaler(enabled=use_fp16)
331
+
332
+ accumulation = int(gradient_accumulation_steps)
333
+ requested_epochs = int(epochs)
334
+ final_epoch = start_epoch + requested_epochs
335
+ batches_per_epoch = max(1, len(loader))
336
+ total_batches = requested_epochs * batches_per_epoch
337
+ coordinates = coordinate_grid(config.resolution, config.resolution, device=device)
338
+ settings = {
339
+ "resolution": config.resolution,
340
+ "resize_mode": resize_mode,
341
+ "horizontal_flip": bool(horizontal_flip),
342
+ "batch_size": int(batch_size),
343
+ "learning_rate": float(learning_rate),
344
+ "weight_decay": float(weight_decay),
345
+ "gradient_accumulation_steps": accumulation,
346
+ "workers": int(workers),
347
+ "mixed_precision": mixed_precision,
348
+ "query_points": int(query_points),
349
+ "time_sampling": time_sampling,
350
+ "ema_decay": float(ema_decay),
351
+ **config.to_dict(),
352
+ }
353
+ write_json(
354
+ output / CONFIG_NAME,
355
+ {
356
+ "format_version": MODEL_FORMAT_VERSION,
357
+ "model_type": "inrflow",
358
+ "model_name": name,
359
+ **config.to_dict(),
360
+ },
361
+ )
362
+ context.log(
363
+ f"Training INRFlow on {len(paths)} images from {training_dataset} at "
364
+ f"{config.resolution}x{config.resolution}, {parameter_count(model):,} parameters, "
365
+ f"batch {batch_size}, device {device}."
366
+ )
367
+ context.log(
368
+ "Images stay in RGB coordinate space: no VAE or other pretrained image compressor is used."
369
+ )
370
+ optimizer.zero_grad(set_to_none=True)
371
+ processed_batches = 0
372
+ last_loss = 0.0
373
+ try:
374
+ for epoch in range(start_epoch + 1, final_epoch + 1):
375
+ model.train()
376
+ epoch_loss = 0.0
377
+ for batch_index, images in enumerate(loader, 1):
378
+ context.checkpoint()
379
+ images = images.to(device, non_blocking=device.type == "cuda")
380
+ batch = images.shape[0]
381
+ clean = images.permute(0, 2, 3, 1).reshape(batch, -1, 3)
382
+ noise = torch.randn_like(clean)
383
+ if time_sampling == "logit_normal":
384
+ time = torch.sigmoid(torch.randn(batch, device=device))
385
+ else:
386
+ time = torch.rand(batch, device=device)
387
+ mixed = (1.0 - time[:, None, None]) * noise + time[:, None, None] * clean
388
+ target = clean - noise
389
+ sample_count = min(int(query_points), clean.shape[1])
390
+ indices = torch.randperm(clean.shape[1], device=device)[:sample_count]
391
+ all_coordinates = coordinates.unsqueeze(0).expand(batch, -1, -1)
392
+ amp = (
393
+ torch.autocast(
394
+ device_type=device.type,
395
+ dtype=autocast_dtype,
396
+ enabled=use_fp16 or use_bf16,
397
+ )
398
+ if device.type in {"cuda", "cpu"}
399
+ else nullcontext()
400
+ )
401
+ with amp:
402
+ velocity = model(
403
+ all_coordinates,
404
+ mixed,
405
+ time,
406
+ all_coordinates[:, indices],
407
+ mixed[:, indices],
408
+ height=config.resolution,
409
+ width=config.resolution,
410
+ )
411
+ loss = torch.nn.functional.mse_loss(velocity, target[:, indices])
412
+ scaled_loss = loss / accumulation
413
+ scaler.scale(scaled_loss).backward()
414
+ if batch_index % accumulation == 0 or batch_index == batches_per_epoch:
415
+ scaler.unscale_(optimizer)
416
+ torch.nn.utils.clip_grad_norm_(model.parameters(), 2.0)
417
+ scaler.step(optimizer)
418
+ scaler.update()
419
+ optimizer.zero_grad(set_to_none=True)
420
+ global_step += 1
421
+ _update_ema(ema_model, model, float(ema_decay))
422
+ last_loss = float(loss.detach().item())
423
+ epoch_loss += last_loss
424
+ processed_batches += 1
425
+ percent = max(1, min(99, round(processed_batches * 100 / total_batches)))
426
+ context.progress(
427
+ percent,
428
+ f"Epoch {epoch} of {final_epoch} · flow loss {last_loss:.4f}",
429
+ epoch=epoch,
430
+ total_epochs=final_epoch,
431
+ current_step=processed_batches,
432
+ total_steps=total_batches,
433
+ unit="batch",
434
+ loss=last_loss,
435
+ )
436
+
437
+ payload = _checkpoint_payload(
438
+ model,
439
+ ema_model,
440
+ optimizer,
441
+ model_name=name,
442
+ dataset_dir=dataset,
443
+ completed_epochs=epoch,
444
+ global_step=global_step,
445
+ training_settings=settings,
446
+ )
447
+ if epoch % int(save_every) == 0:
448
+ _save_checkpoint(output / "checkpoints" / f"epoch_{epoch:06d}.pt", payload)
449
+ if bool(preview_enabled) and epoch % int(preview_every) == 0:
450
+ _preview(
451
+ context,
452
+ ema_model,
453
+ output,
454
+ epoch=epoch,
455
+ next_epoch=min(final_epoch, epoch + int(preview_every)),
456
+ steps=int(preview_steps),
457
+ seed=int(preview_seed),
458
+ prompt=preview_prompt,
459
+ )
460
+ context.log(f"Finished epoch {epoch}; average loss {epoch_loss / batches_per_epoch:.4f}.")
461
+ except torch.cuda.OutOfMemoryError as exc:
462
+ if device.type == "cuda":
463
+ torch.cuda.empty_cache()
464
+ raise ToolExecutionError(
465
+ "INRFlow ran out of VRAM. Reduce batch size, then pixel queries, resolution, or model width."
466
+ ) from exc
467
+
468
+ final_payload = _checkpoint_payload(
469
+ model,
470
+ ema_model,
471
+ optimizer,
472
+ model_name=name,
473
+ dataset_dir=dataset,
474
+ completed_epochs=final_epoch,
475
+ global_step=global_step,
476
+ training_settings=settings,
477
+ )
478
+ final_checkpoint = output / FINAL_CHECKPOINT_NAME
479
+ _save_checkpoint(final_checkpoint, final_payload)
480
+ write_json(
481
+ output / "training_metadata.json",
482
+ {
483
+ "format_version": MODEL_FORMAT_VERSION,
484
+ "model_type": "inrflow",
485
+ "architecture": "inrflow_ambient_space",
486
+ "method": "conditionally_independent_continuous_flow_matching",
487
+ "model_name": name,
488
+ "dataset_dir": str(dataset),
489
+ "image_count": len(paths),
490
+ "completed_epochs": final_epoch,
491
+ "epochs_this_run": requested_epochs,
492
+ "global_step": global_step,
493
+ "final_loss": last_loss,
494
+ "parameter_count": parameter_count(model),
495
+ "checkpoint": str(final_checkpoint),
496
+ "uses_pretrained_compressor": False,
497
+ "method_reference": "https://arxiv.org/abs/2412.03791",
498
+ "settings": settings,
499
+ "finished_at": datetime.now(timezone.utc).isoformat(),
500
+ },
501
+ )
502
+ context.progress(100, "INRFlow training completed")
503
+ return {
504
+ "output_folder": str(output),
505
+ "model_name": name,
506
+ "assets": [
507
+ {
508
+ "kind": "model",
509
+ "name": name,
510
+ "path": str(output),
511
+ "trainer": "inrflow",
512
+ "dataset_path": str(dataset),
513
+ "checkpoint": str(final_checkpoint),
514
+ "epochs": final_epoch,
515
+ "metadata": {
516
+ "architecture": "inrflow_ambient_space",
517
+ "resolution": config.resolution,
518
+ "parameter_count": parameter_count(model),
519
+ "uses_pretrained_compressor": False,
520
+ },
521
+ }
522
+ ],
523
+ }
adam/model_plugins_builtin/oasis/manifest.py CHANGED
@@ -2,35 +2,40 @@ PLUGIN_ID = "oasis"
2
 
3
  MODEL_INFO = {
4
  "name": "Oasis Action World Model",
5
- "version": "1.0",
6
  "category": "Playable World Models",
7
  "description": "Action-conditioned playable world model trainer for gameplay frame sequences.",
8
- "architecture": "action_conditioned_rectified_flow_video",
9
  "status": "experimental",
10
  "output_type": "playable_world",
11
  "capabilities": ["fresh_training", "resume_training", "playable_inference", "live_preview"],
12
  "input_formats": ["Oasis action dataset folder", "semicolon-separated Oasis dataset folders"],
13
- "output_formats": ["action_flow_model_info.json", "diffusers unet folder", "png preview"],
14
  "hardware": {"recommended_vram_gb": 12, "recommended_system_ram_gb": 32},
15
  "vram_behavior": {
16
- "scales_with": ["resolution", "batch_size", "sequence_context"],
17
  "estimate": "High; 256x144 with batch 2 is the conservative RTX 3060 starting point.",
18
  },
19
  }
20
 
21
  TRAINING_SETTINGS = {
22
- "resolution": {"label": "Resolution", "type": "choice", "options": ["128x72", "256x144", "384x216", "512x288"], "default": "256x144", "group": "Basic"},
 
 
23
  "batch_size": {"label": "Batch size", "type": "int", "default": 2, "min": 1, "max": 16, "group": "Basic"},
24
  "learning_rate": {"label": "Learning rate", "type": "float", "default": 0.00002, "min": 0.0000001, "max": 0.01, "decimals": 7, "step": 0.00001, "group": "Optimization"},
25
  "workers": {"label": "Loader workers", "type": "int", "default": 2, "min": 0, "max": 8, "group": "Dataset"},
26
- "mixed_precision": {"label": "Precision", "type": "choice", "options": ["fp32", "fp16", "no"], "default": "fp32", "group": "Optimization"},
27
  "gradient_accumulation": {"label": "Gradient accumulation", "type": "int", "default": 1, "min": 1, "max": 16, "group": "Optimization"},
28
- "frame_gap": {"label": "Prediction horizon", "type": "int", "default": 3, "min": 1, "max": 60, "group": "Sequence"},
29
- "sequence_context": {"label": "Context length", "type": "int", "default": 1, "min": 1, "max": 32, "group": "Sequence"},
 
 
30
  "action_aggregation": {"label": "Action aggregation", "type": "choice", "options": ["window", "mean", "last"], "default": "window", "group": "Sequence"},
31
  "validation_split": {"label": "Validation split", "type": "float", "default": 0.1, "min": 0.01, "max": 0.5, "decimals": 3, "step": 0.01, "group": "Dataset"},
32
  "validation_batches": {"label": "Validation batches", "type": "int", "default": 8, "min": 0, "max": 128, "group": "Dataset"},
33
  "save_every": {"label": "Save every", "type": "int", "default": 5, "min": 1, "max": 1000, "group": "Checkpoints"},
 
34
  "preview_enabled": {"label": "Generate previews while training", "type": "bool", "default": True, "group": "Preview"},
35
  "preview_every": {"label": "Preview interval", "type": "int", "default": 5, "min": 1, "max": 100000, "group": "Preview"},
36
  "preview_steps": {"label": "Preview steps", "type": "int", "default": 1, "min": 1, "max": 50, "group": "Preview"},
@@ -43,8 +48,18 @@ TRAINING_SETTINGS = {
43
  "neutral_action_dropout": {"label": "Neutral action dropout", "type": "float", "default": 0.15, "min": 0.0, "max": 0.9, "decimals": 3, "step": 0.05, "group": "Advanced", "advanced": True},
44
  "action_contrast_weight": {"label": "Action contrast weight", "type": "float", "default": 0.35, "min": 0.0, "max": 5.0, "decimals": 3, "step": 0.05, "group": "Advanced", "advanced": True},
45
  "action_contrast_margin": {"label": "Action contrast margin", "type": "float", "default": 0.02, "min": 0.0, "max": 1.0, "decimals": 4, "step": 0.005, "group": "Advanced", "advanced": True},
 
 
 
 
 
 
 
 
 
 
46
  "gradient_checkpointing": {"label": "Gradient checkpointing", "type": "bool", "default": False, "group": "Advanced", "advanced": True},
47
- "balance_actions": {"label": "Balance rare actions", "type": "bool", "default": False, "group": "Advanced", "advanced": True},
48
  }
49
 
50
  GENERATION_SETTINGS = {
 
2
 
3
  MODEL_INFO = {
4
  "name": "Oasis Action World Model",
5
+ "version": "1.1",
6
  "category": "Playable World Models",
7
  "description": "Action-conditioned playable world model trainer for gameplay frame sequences.",
8
+ "architecture": "action_conditioned_rectified_flow_video | action_conditioned_latent_vae_flow_video | action_conditioned_temporal_latent_flow | action_conditioned_temporal_pixel_flow",
9
  "status": "experimental",
10
  "output_type": "playable_world",
11
  "capabilities": ["fresh_training", "resume_training", "playable_inference", "live_preview"],
12
  "input_formats": ["Oasis action dataset folder", "semicolon-separated Oasis dataset folders"],
13
+ "output_formats": ["action_flow_model_info.json", "diffusers unet folder", "VAE folder for latent models", "png preview"],
14
  "hardware": {"recommended_vram_gb": 12, "recommended_system_ram_gb": 32},
15
  "vram_behavior": {
16
+ "scales_with": ["resolution", "batch_size", "history frames"],
17
  "estimate": "High; 256x144 with batch 2 is the conservative RTX 3060 starting point.",
18
  },
19
  }
20
 
21
  TRAINING_SETTINGS = {
22
+ "model_engine": {"label": "Model engine", "type": "choice", "options": ["pixel_flow", "vae_cpu_lite", "temporal_latent", "temporal_pixel_flow"], "option_labels": {"pixel_flow": "Pixel Flow (RGB)", "vae_cpu_lite": "VAE CPU Lite", "temporal_latent": "Temporal Latent (VAE + history)", "temporal_pixel_flow": "Temporal Pixel Flow"}, "default": "temporal_latent", "group": "Basic", "description": "Temporal Pixel Flow keeps recent-frame and timed-input conditioning while predicting full RGB frames without a VAE. It needs more VRAM than Temporal Latent."},
23
+ "resolution": {"label": "Resolution", "type": "choice", "options": ["256x144", "512x288"], "default": "256x144", "group": "Basic", "description": "Must be 16:9 with both dimensions divisible by 16."},
24
+ "frame_gap": {"label": "Frame gap", "type": "int", "default": 1, "min": 1, "max": 60, "group": "Basic", "description": "How many recorded frames one generated game frame spans. Use 1 for a 12–15 FPS recording when you want the most responsive native AI FPS."},
25
  "batch_size": {"label": "Batch size", "type": "int", "default": 2, "min": 1, "max": 16, "group": "Basic"},
26
  "learning_rate": {"label": "Learning rate", "type": "float", "default": 0.00002, "min": 0.0000001, "max": 0.01, "decimals": 7, "step": 0.00001, "group": "Optimization"},
27
  "workers": {"label": "Loader workers", "type": "int", "default": 2, "min": 0, "max": 8, "group": "Dataset"},
28
+ "mixed_precision": {"label": "Precision", "type": "choice", "options": ["fp16", "no"], "default": "fp16", "group": "Optimization", "description": "The connected Oasis trainer supports FP16 or full precision (no AMP)."},
29
  "gradient_accumulation": {"label": "Gradient accumulation", "type": "int", "default": 1, "min": 1, "max": 16, "group": "Optimization"},
30
+ "sequence_context": {"label": "Legacy context length", "type": "int", "default": 1, "min": 1, "max": 32, "group": "Sequence", "advanced": True, "description": "Only used by the older pixel-flow engines."},
31
+ "context_frames": {"label": "History frames", "type": "choice", "options": [1, 4, 8], "default": 4, "group": "Sequence", "description": "How much recent visual and input history the temporal model uses."},
32
+ "rollout_frames": {"label": "Future frames", "type": "choice", "options": [1, 3], "default": 3, "group": "Sequence", "description": "How many future frames each temporal training example predicts."},
33
+ "vae_epochs": {"label": "VAE warm-up epochs", "type": "int", "default": 5, "min": 1, "max": 100, "group": "Sequence", "description": "Initial epochs used to learn the compact visual representation for latent engines. Not used by pixel_flow or temporal_pixel_flow."},
34
  "action_aggregation": {"label": "Action aggregation", "type": "choice", "options": ["window", "mean", "last"], "default": "window", "group": "Sequence"},
35
  "validation_split": {"label": "Validation split", "type": "float", "default": 0.1, "min": 0.01, "max": 0.5, "decimals": 3, "step": 0.01, "group": "Dataset"},
36
  "validation_batches": {"label": "Validation batches", "type": "int", "default": 8, "min": 0, "max": 128, "group": "Dataset"},
37
  "save_every": {"label": "Save every", "type": "int", "default": 5, "min": 1, "max": 1000, "group": "Checkpoints"},
38
+ "best_checkpoint_min_improvement": {"label": "Best-checkpoint improvement (%)", "type": "float", "default": 0.1, "min": 0.0, "max": 20.0, "decimals": 2, "step": 0.1, "group": "Checkpoints", "advanced": True, "description": "Save a separate best-validation checkpoint only when held-out quality improves by this percentage. ADAM uses that checkpoint automatically when playing the completed model."},
39
  "preview_enabled": {"label": "Generate previews while training", "type": "bool", "default": True, "group": "Preview"},
40
  "preview_every": {"label": "Preview interval", "type": "int", "default": 5, "min": 1, "max": 100000, "group": "Preview"},
41
  "preview_steps": {"label": "Preview steps", "type": "int", "default": 1, "min": 1, "max": 50, "group": "Preview"},
 
48
  "neutral_action_dropout": {"label": "Neutral action dropout", "type": "float", "default": 0.15, "min": 0.0, "max": 0.9, "decimals": 3, "step": 0.05, "group": "Advanced", "advanced": True},
49
  "action_contrast_weight": {"label": "Action contrast weight", "type": "float", "default": 0.35, "min": 0.0, "max": 5.0, "decimals": 3, "step": 0.05, "group": "Advanced", "advanced": True},
50
  "action_contrast_margin": {"label": "Action contrast margin", "type": "float", "default": 0.02, "min": 0.0, "max": 1.0, "decimals": 4, "step": 0.005, "group": "Advanced", "advanced": True},
51
+ "contrast_every": {"label": "Contrast every N batches", "type": "int", "default": 4, "min": 1, "max": 128, "group": "Advanced", "advanced": True},
52
+ "contrast_samples": {"label": "Contrast samples", "type": "int", "default": 2, "min": 2, "max": 32, "group": "Advanced", "advanced": True},
53
+ "chunk_size": {"label": "Transitions per training chunk", "type": "int", "default": 0, "min": 0, "max": 1000000, "group": "Dataset", "description": "0 uses every transition each epoch. A balanced chunk bounds large-dataset training time."},
54
+ "chunk_mode": {"label": "Chunk selection", "type": "choice", "options": ["balanced", "random", "sequential"], "default": "balanced", "group": "Dataset"},
55
+ "chunk_offset": {"label": "Chunk offset", "type": "int", "default": 0, "min": 0, "max": 100000000, "group": "Dataset", "advanced": True},
56
+ "replay_older_percent": {"label": "Older-data replay (%)", "type": "float", "default": 50.0, "min": 0.0, "max": 500.0, "decimals": 1, "step": 25.0, "group": "Dataset", "advanced": True},
57
+ "include_older_data": {"label": "Mix selected older datasets", "type": "bool", "default": True, "group": "Dataset", "advanced": True},
58
+ "recovery_minutes": {"label": "Emergency recovery every minutes", "type": "int", "default": 30, "min": 0, "max": 240, "group": "Checkpoints", "advanced": True},
59
+ "benchmark_batches": {"label": "Speed-test batches", "type": "int", "default": 8, "min": 1, "max": 100, "group": "Advanced", "advanced": True},
60
+ "tf32": {"label": "Use TF32 acceleration", "type": "bool", "default": True, "group": "Optimization", "advanced": True},
61
  "gradient_checkpointing": {"label": "Gradient checkpointing", "type": "bool", "default": False, "group": "Advanced", "advanced": True},
62
+ "balance_actions": {"label": "Balance rare actions", "type": "bool", "default": True, "group": "Advanced", "advanced": True},
63
  }
64
 
65
  GENERATION_SETTINGS = {
adam/model_plugins_builtin/pixelrow/__init__.py ADDED
@@ -0,0 +1,2 @@
 
 
 
1
+ """PixelRow autoregressive image model plugin."""
2
+
adam/model_plugins_builtin/pixelrow/common.py ADDED
@@ -0,0 +1,86 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from __future__ import annotations
2
+
3
+ import json
4
+ import re
5
+ from pathlib import Path
6
+ from typing import Any
7
+
8
+ import torch
9
+ from PIL import Image
10
+
11
+ from adam.executor import ToolExecutionError
12
+
13
+ from .model import MODEL_FORMAT_VERSION, PixelRowConfig, PixelRowModel, canvas_to_uint8
14
+
15
+
16
+ IMAGE_EXTENSIONS = {".jpg", ".jpeg", ".png", ".webp", ".bmp"}
17
+ FINAL_CHECKPOINT_NAME = "pixelrow_model.pt"
18
+ CONFIG_NAME = "pixelrow_config.json"
19
+
20
+
21
+ def safe_model_name(value: str) -> str:
22
+ name = re.sub(r"\s+", " ", value.strip())
23
+ if not name or len(name) > 96 or any(character in name for character in '<>:"/\\|?*\x00'):
24
+ raise ToolExecutionError("Choose a short PixelRow model name without reserved filename characters.")
25
+ return name
26
+
27
+
28
+ def ensure_below(path: Path, root: Path, label: str) -> Path:
29
+ resolved = path.expanduser().resolve()
30
+ try:
31
+ resolved.relative_to(root.expanduser().resolve())
32
+ except ValueError as exc:
33
+ raise ToolExecutionError(f"{label} must stay inside {root.resolve()}.") from exc
34
+ return resolved
35
+
36
+
37
+ def image_files(folder: Path) -> list[Path]:
38
+ try:
39
+ return sorted(
40
+ path for path in folder.rglob("*")
41
+ if path.is_file() and path.suffix.casefold() in IMAGE_EXTENSIONS
42
+ )
43
+ except OSError:
44
+ return []
45
+
46
+
47
+ def resolve_checkpoint(path: Path) -> Path:
48
+ candidate = path.expanduser().resolve()
49
+ if candidate.is_dir():
50
+ candidate = candidate / FINAL_CHECKPOINT_NAME
51
+ if not candidate.is_file():
52
+ raise ToolExecutionError("The selected PixelRow checkpoint does not exist or is incomplete.")
53
+ return candidate
54
+
55
+
56
+ def load_checkpoint(path: Path, device: torch.device) -> tuple[PixelRowModel, dict[str, Any]]:
57
+ checkpoint_path = resolve_checkpoint(path)
58
+ try:
59
+ payload = torch.load(checkpoint_path, map_location=device, weights_only=False)
60
+ except (OSError, RuntimeError, ValueError, TypeError) as exc:
61
+ raise ToolExecutionError(f"Could not load the PixelRow checkpoint: {exc}") from exc
62
+ if not isinstance(payload, dict) or "model_state" not in payload or "config" not in payload:
63
+ raise ToolExecutionError("The selected file is not a valid PixelRow checkpoint.")
64
+ if int(payload.get("format_version", 0)) != MODEL_FORMAT_VERSION:
65
+ raise ToolExecutionError("This PixelRow checkpoint uses an unsupported format version.")
66
+ try:
67
+ config = PixelRowConfig.from_dict(dict(payload["config"]))
68
+ model = PixelRowModel(config).to(device)
69
+ model.load_state_dict(payload["model_state"], strict=True)
70
+ except (KeyError, TypeError, ValueError, RuntimeError) as exc:
71
+ raise ToolExecutionError(f"The PixelRow checkpoint is incompatible: {exc}") from exc
72
+ return model, payload
73
+
74
+
75
+ def save_canvas(canvas: torch.Tensor, path: Path, completed_rows: int | None = None) -> None:
76
+ pixels = canvas_to_uint8(canvas, completed_rows=completed_rows).numpy()
77
+ path.parent.mkdir(parents=True, exist_ok=True)
78
+ Image.fromarray(pixels, mode="RGB").save(path, format="PNG")
79
+
80
+
81
+ def write_json(path: Path, payload: dict[str, Any]) -> None:
82
+ path.parent.mkdir(parents=True, exist_ok=True)
83
+ temporary = path.with_suffix(path.suffix + ".tmp")
84
+ temporary.write_text(json.dumps(payload, indent=2), encoding="utf-8")
85
+ temporary.replace(path)
86
+
adam/model_plugins_builtin/pixelrow/generator.py ADDED
@@ -0,0 +1,212 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from __future__ import annotations
2
+
3
+ import random
4
+ from datetime import datetime, timezone
5
+ from pathlib import Path
6
+ from typing import Any
7
+
8
+ import torch
9
+
10
+ from adam.executor import ToolExecutionError
11
+ from adam.generations import generation_metadata_path, generation_output_folder
12
+
13
+ from .common import ensure_below, load_checkpoint, safe_model_name, save_canvas, write_json
14
+
15
+
16
+ def generate(
17
+ context,
18
+ model_name: str,
19
+ model_path: str,
20
+ prompt: str,
21
+ image_count: int,
22
+ steps: int,
23
+ seed: int,
24
+ sampler: str,
25
+ aspect_ratio: str,
26
+ temperature: float = 0.85,
27
+ top_k: int = 4,
28
+ save_progress_frames: bool = True,
29
+ frame_interval: int = 2,
30
+ preview_interval: int = 4,
31
+ ) -> dict[str, Any]:
32
+ """Generate images one complete RGB row at a time."""
33
+ name = safe_model_name(model_name)
34
+ model_root = (context.root.resolve() / "data" / "model_plugin_outputs" / "pixelrow").resolve()
35
+ selected = ensure_below(Path(model_path), model_root, "PixelRow model")
36
+ if not selected.exists():
37
+ raise ToolExecutionError("The selected PixelRow model no longer exists.")
38
+ if not 1 <= int(image_count) <= 48:
39
+ raise ToolExecutionError("PixelRow image count must be between 1 and 48.")
40
+ if not 1 <= int(steps) <= 128:
41
+ raise ToolExecutionError("PixelRow rows to generate must be between 1 and 128.")
42
+ if sampler != "Categorical":
43
+ raise ToolExecutionError("PixelRow currently supports categorical row sampling.")
44
+ if aspect_ratio != "1:1 (Native)":
45
+ raise ToolExecutionError("PixelRow currently generates at its native square resolution.")
46
+ if not 0.05 <= float(temperature) <= 3.0:
47
+ raise ToolExecutionError("PixelRow creativity must be between 0.05 and 3.0.")
48
+ if not 1 <= int(top_k) <= 64:
49
+ raise ToolExecutionError("PixelRow top color choices must be between 1 and 64.")
50
+ if int(frame_interval) not in {1, 2, 4, 8, 16}:
51
+ raise ToolExecutionError("PixelRow frame interval must be 1, 2, 4, 8, or 16 rows.")
52
+ if not 0 <= int(preview_interval) <= 128:
53
+ raise ToolExecutionError("PixelRow live preview interval must be between 0 and 128 rows.")
54
+ if len(prompt) > 500:
55
+ raise ToolExecutionError("The PixelRow creative note must be 500 characters or shorter.")
56
+
57
+ device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
58
+ model, checkpoint = load_checkpoint(selected, device)
59
+ height = model.config.resolution
60
+ rows_to_generate = min(int(steps), height)
61
+ effective_top_k = min(int(top_k), model.config.color_bins)
62
+ if int(steps) > height:
63
+ context.log(
64
+ f"This model is {height}px tall, so PixelRow will stop after its {height} native rows."
65
+ )
66
+ if effective_top_k != int(top_k):
67
+ context.log(
68
+ f"This model has {model.config.color_bins} color levels; top color choices was capped to that value."
69
+ )
70
+
71
+ count = int(image_count)
72
+ base_seed = int(seed)
73
+ if base_seed <= 0:
74
+ base_seed = random.SystemRandom().randint(1, 2_147_483_647 - count)
75
+ if base_seed + count - 1 > 2_147_483_647:
76
+ raise ToolExecutionError("The PixelRow seed is too large for this image count.")
77
+
78
+ output = generation_output_folder(context.root, context.tool.id, name)
79
+ timestamp = datetime.now(timezone.utc).strftime("%Y%m%d_%H%M%S")
80
+ image_paths: list[str] = []
81
+ frame_folders: list[str] = []
82
+ context.log(
83
+ f"Loaded {name}: {height}x{height}, {model.config.color_bins} color levels, device {device}."
84
+ )
85
+ context.log("PixelRow uses the creative note as metadata; generation is unconditional.")
86
+
87
+ try:
88
+ for image_index in range(count):
89
+ context.checkpoint()
90
+ current_seed = base_seed + image_index
91
+ torch_generator = torch.Generator(device=device)
92
+ torch_generator.manual_seed(current_seed)
93
+ frame_folder = output / "row_progress" / f"{timestamp}_{context.job_id}_seed_{current_seed}"
94
+ last_frame_path: Path | None = None
95
+
96
+ def row_ready(completed_rows: int, canvas: torch.Tensor) -> None:
97
+ nonlocal last_frame_path
98
+ context.checkpoint()
99
+ overall = ((image_index * rows_to_generate) + completed_rows) / max(1, count * rows_to_generate)
100
+ context.progress(
101
+ max(1, min(99, round(overall * 100))),
102
+ f"Image {image_index + 1} of {count} · row {completed_rows} of {rows_to_generate}",
103
+ current=completed_rows,
104
+ total=rows_to_generate,
105
+ image_index=image_index,
106
+ image_count=count,
107
+ unit="row",
108
+ )
109
+ save_frame = bool(save_progress_frames) and (
110
+ completed_rows % int(frame_interval) == 0 or completed_rows == rows_to_generate
111
+ )
112
+ publish_preview = int(preview_interval) > 0 and (
113
+ completed_rows % int(preview_interval) == 0 or completed_rows == rows_to_generate
114
+ )
115
+ if save_frame:
116
+ last_frame_path = frame_folder / f"row_{completed_rows:04d}.png"
117
+ save_canvas(canvas, last_frame_path, completed_rows=completed_rows)
118
+ if publish_preview:
119
+ preview_path = last_frame_path
120
+ if preview_path is None or not preview_path.is_file():
121
+ preview_path = output / ".live" / context.job_id / f"image_{image_index + 1:03d}.png"
122
+ save_canvas(canvas, preview_path, completed_rows=completed_rows)
123
+ context.preview(
124
+ preview_path,
125
+ kind="generation",
126
+ current=completed_rows,
127
+ total=rows_to_generate,
128
+ image_index=image_index,
129
+ image_count=count,
130
+ seed=current_seed,
131
+ steps=rows_to_generate,
132
+ )
133
+
134
+ canvas = model.generate(
135
+ rows=rows_to_generate,
136
+ temperature=float(temperature),
137
+ top_k=effective_top_k,
138
+ generator=torch_generator,
139
+ row_callback=row_ready,
140
+ )
141
+ destination = output / (
142
+ f"{timestamp}_{context.job_id}_PixelRow_seed_{current_seed}_rows_{rows_to_generate}.png"
143
+ )
144
+ save_canvas(canvas, destination, completed_rows=rows_to_generate)
145
+ image_paths.append(str(destination))
146
+ if bool(save_progress_frames):
147
+ frame_folders.append(str(frame_folder))
148
+ except torch.cuda.OutOfMemoryError as exc:
149
+ if device.type == "cuda":
150
+ torch.cuda.empty_cache()
151
+ raise ToolExecutionError(
152
+ "PixelRow ran out of VRAM while generating. Generate fewer images in one batch."
153
+ ) from exc
154
+
155
+ created_at = datetime.now(timezone.utc).isoformat()
156
+ metadata = {
157
+ "version": 1,
158
+ "provider_id": context.tool.id,
159
+ "provider_name": context.tool.name,
160
+ "model_name": name,
161
+ "model_path": str(selected),
162
+ "model_type": "pixelrow",
163
+ "prompt": prompt.strip(),
164
+ "prompt_behavior": "label_only",
165
+ "seed": base_seed,
166
+ "image_seeds": [base_seed + index for index in range(count)],
167
+ "image_count": count,
168
+ "steps": rows_to_generate,
169
+ "sampler": sampler,
170
+ "aspect_ratio": aspect_ratio,
171
+ "resolution": height,
172
+ "temperature": float(temperature),
173
+ "top_k": effective_top_k,
174
+ "save_progress_frames": bool(save_progress_frames),
175
+ "frame_interval": int(frame_interval),
176
+ "row_progress_folders": frame_folders,
177
+ "preview_interval": int(preview_interval),
178
+ "images": image_paths,
179
+ "checkpoint_epoch": int(checkpoint.get("completed_epochs", 0) or 0),
180
+ "created_at": created_at,
181
+ }
182
+ write_json(generation_metadata_path(output, timestamp, context.job_id), metadata)
183
+ # Leave only durable showcase frames; the live-preview file is an implementation detail.
184
+ live_folder = output / ".live" / context.job_id
185
+ if live_folder.is_dir():
186
+ for path in live_folder.glob("*.png"):
187
+ try:
188
+ path.unlink()
189
+ except OSError:
190
+ pass
191
+ try:
192
+ live_folder.rmdir()
193
+ except OSError:
194
+ pass
195
+ try:
196
+ live_folder.parent.rmdir()
197
+ except OSError:
198
+ pass
199
+ context.progress(100, f"Generated {count} PixelRow image(s)")
200
+ return {
201
+ "output_folder": str(output),
202
+ "assets": [{
203
+ "kind": "generation",
204
+ "name": f"{name} · {timestamp}",
205
+ "path": str(output),
206
+ "trainer": "pixelrow",
207
+ "metadata": {
208
+ "row_progress_folders": frame_folders,
209
+ "rows_generated": rows_to_generate,
210
+ },
211
+ }],
212
+ }
adam/model_plugins_builtin/pixelrow/manifest.py ADDED
@@ -0,0 +1,301 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ PLUGIN_ID = "pixelrow"
2
+
3
+ MODEL_INFO = {
4
+ "name": "PixelRow",
5
+ "version": "0.1",
6
+ "category": "Image Generation",
7
+ "description": (
8
+ "Experimental autoregressive image model that constructs pictures from "
9
+ "top to bottom, predicting one complete RGB row at a time."
10
+ ),
11
+ "architecture": "autoregressive_rows",
12
+ "status": "experimental",
13
+ "output_type": "image",
14
+ "capabilities": [
15
+ "fresh_training",
16
+ "resume_training",
17
+ "image_generation",
18
+ "live_preview",
19
+ "row_progress_frames",
20
+ ],
21
+ "input_formats": ["image folder"],
22
+ "output_formats": ["PixelRow checkpoint", "PNG image", "PNG row progress frames"],
23
+ "hardware": {"recommended_vram_gb": 6, "recommended_system_ram_gb": 16},
24
+ "vram_behavior": {
25
+ "scales_with": ["resolution", "batch_size", "color_bins", "hidden_size"],
26
+ "estimate": "Moderate at 64px; reduce batch size first when training at 128px.",
27
+ },
28
+ }
29
+
30
+ TRAINING_SETTINGS = {
31
+ "resolution": {
32
+ "label": "Resolution",
33
+ "type": "choice",
34
+ "options": [32, 64, 128],
35
+ "default": 64,
36
+ "group": "Basic",
37
+ "description": "PixelRow currently learns square images. Start at 64px for the first car experiment.",
38
+ },
39
+ "resize_mode": {
40
+ "label": "Image fitting",
41
+ "type": "choice",
42
+ "options": ["fill", "fit", "stretch"],
43
+ "default": "fill",
44
+ "group": "Dataset",
45
+ "description": "Fill preserves proportions and center-crops; fit pads; stretch changes proportions.",
46
+ },
47
+ "horizontal_flip": {
48
+ "label": "Random horizontal flip",
49
+ "type": "bool",
50
+ "default": True,
51
+ "group": "Dataset",
52
+ },
53
+ "batch_size": {
54
+ "label": "Batch size",
55
+ "type": "int",
56
+ "default": 8,
57
+ "min": 1,
58
+ "max": 64,
59
+ "group": "Basic",
60
+ },
61
+ "learning_rate": {
62
+ "label": "Learning rate",
63
+ "type": "float",
64
+ "default": 0.0002,
65
+ "min": 0.0000001,
66
+ "max": 0.1,
67
+ "decimals": 7,
68
+ "step": 0.00005,
69
+ "group": "Optimization",
70
+ },
71
+ "gradient_accumulation_steps": {
72
+ "label": "Gradient accumulation",
73
+ "type": "int",
74
+ "default": 1,
75
+ "min": 1,
76
+ "max": 64,
77
+ "group": "Optimization",
78
+ },
79
+ "workers": {
80
+ "label": "Loader workers",
81
+ "type": "int",
82
+ "default": 0,
83
+ "min": 0,
84
+ "max": 16,
85
+ "group": "Dataset",
86
+ "description": "Zero is the safest choice for the Windows desktop app.",
87
+ },
88
+ "mixed_precision": {
89
+ "label": "Precision",
90
+ "type": "choice",
91
+ "options": ["fp16", "no"],
92
+ "default": "fp16",
93
+ "group": "Optimization",
94
+ },
95
+ "hidden_size": {
96
+ "label": "Spatial memory channels",
97
+ "type": "choice",
98
+ "options": [64, 128, 192, 256],
99
+ "default": 128,
100
+ "group": "PixelRow",
101
+ "advanced": True,
102
+ "description": "Column-aware memory carried from completed rows into the next-row prediction.",
103
+ },
104
+ "recurrent_layers": {
105
+ "label": "Sequence layers",
106
+ "type": "int",
107
+ "default": 2,
108
+ "min": 1,
109
+ "max": 4,
110
+ "group": "PixelRow",
111
+ "advanced": True,
112
+ },
113
+ "row_channels": {
114
+ "label": "Row feature channels",
115
+ "type": "choice",
116
+ "options": [32, 64, 96, 128],
117
+ "default": 64,
118
+ "group": "PixelRow",
119
+ "advanced": True,
120
+ },
121
+ "color_bins": {
122
+ "label": "Color levels per channel",
123
+ "type": "choice",
124
+ "options": [16, 32, 64],
125
+ "default": 32,
126
+ "group": "PixelRow",
127
+ "advanced": True,
128
+ "description": "PixelRow predicts a color category instead of averaging raw RGB values.",
129
+ },
130
+ "edge_loss_weight": {
131
+ "label": "Line-detail strength",
132
+ "type": "float",
133
+ "default": 0.05,
134
+ "min": 0.0,
135
+ "max": 1.0,
136
+ "decimals": 3,
137
+ "step": 0.01,
138
+ "group": "PixelRow",
139
+ "advanced": True,
140
+ "description": "Encourages horizontal and vertical color boundaries to match the training images.",
141
+ },
142
+ "save_every": {
143
+ "label": "Save every",
144
+ "type": "int",
145
+ "default": 10,
146
+ "min": 1,
147
+ "max": 1000,
148
+ "group": "Checkpoints",
149
+ },
150
+ "preview_enabled": {
151
+ "label": "Generate previews while training",
152
+ "type": "bool",
153
+ "default": True,
154
+ "group": "Preview",
155
+ },
156
+ "preview_every": {
157
+ "label": "Preview interval",
158
+ "type": "int",
159
+ "default": 5,
160
+ "min": 1,
161
+ "max": 100000,
162
+ "group": "Preview",
163
+ },
164
+ "preview_prompt": {
165
+ "label": "Preview note",
166
+ "type": "text",
167
+ "default": "",
168
+ "group": "Preview",
169
+ },
170
+ "preview_seed": {
171
+ "label": "Preview seed",
172
+ "type": "int",
173
+ "default": 123456789,
174
+ "min": 0,
175
+ "max": 2147483647,
176
+ "group": "Preview",
177
+ },
178
+ }
179
+
180
+ GENERATION_SETTINGS = {
181
+ "prompt": {
182
+ "label": "Creative note",
183
+ "type": "multiline_text",
184
+ "default": "",
185
+ "group": "Generation",
186
+ },
187
+ "image_count": {
188
+ "label": "Images",
189
+ "type": "int",
190
+ "default": 1,
191
+ "min": 1,
192
+ "max": 48,
193
+ "group": "Generation",
194
+ },
195
+ "steps": {
196
+ "label": "Rows to generate",
197
+ "type": "int",
198
+ "default": 128,
199
+ "min": 1,
200
+ "max": 128,
201
+ "group": "Generation",
202
+ "description": "Values above the trained image height automatically produce the complete image.",
203
+ },
204
+ "sampler": {
205
+ "label": "Row sampling",
206
+ "type": "choice",
207
+ "options": ["Categorical"],
208
+ "default": "Categorical",
209
+ "group": "Generation",
210
+ },
211
+ "aspect_ratio": {
212
+ "label": "Aspect ratio",
213
+ "type": "choice",
214
+ "options": ["1:1 (Native)"],
215
+ "default": "1:1 (Native)",
216
+ "group": "Generation",
217
+ },
218
+ "seed": {
219
+ "label": "Seed",
220
+ "type": "int",
221
+ "default": 0,
222
+ "min": 0,
223
+ "max": 2147483647,
224
+ "group": "Generation",
225
+ },
226
+ "temperature": {
227
+ "label": "Creativity",
228
+ "type": "float",
229
+ "default": 0.85,
230
+ "min": 0.05,
231
+ "max": 3.0,
232
+ "decimals": 2,
233
+ "step": 0.05,
234
+ "group": "PixelRow",
235
+ },
236
+ "top_k": {
237
+ "label": "Top color choices",
238
+ "type": "int",
239
+ "default": 4,
240
+ "min": 1,
241
+ "max": 64,
242
+ "group": "PixelRow",
243
+ "description": "Smaller values are more conservative; larger values add variation.",
244
+ },
245
+ "save_progress_frames": {
246
+ "label": "Save row-build frames",
247
+ "type": "bool",
248
+ "default": True,
249
+ "group": "Row Showcase",
250
+ },
251
+ "frame_interval": {
252
+ "label": "Save every N rows",
253
+ "type": "choice",
254
+ "options": [1, 2, 4, 8, 16],
255
+ "default": 2,
256
+ "group": "Row Showcase",
257
+ },
258
+ "preview_interval": {
259
+ "label": "Rows per live preview",
260
+ "type": "int",
261
+ "default": 4,
262
+ "min": 0,
263
+ "max": 128,
264
+ "group": "Preview",
265
+ },
266
+ }
267
+
268
+ TRAINING_TOOL = {
269
+ "id": "pixelrow_trainer",
270
+ "name": "PixelRow Trainer",
271
+ "description": "Trains an experimental model to construct images one RGB row at a time.",
272
+ "capabilities": ["fresh_training", "resume_training", "progress", "pause", "cancel", "live_preview"],
273
+ "backend": {
274
+ "type": "python",
275
+ "module": "adam.model_plugins_builtin.pixelrow.trainer",
276
+ "function": "train",
277
+ },
278
+ }
279
+
280
+ GENERATION_TOOL = {
281
+ "id": "pixelrow_generator",
282
+ "name": "PixelRow Generator",
283
+ "description": "Builds images from top to bottom and can save the visible row-by-row process.",
284
+ "model_trainers": ["pixelrow"],
285
+ "capabilities": ["image_generation", "seed", "batch", "row_progress_frames", "live_preview", "progress", "cancel"],
286
+ "generation_options": {
287
+ "samplers": ["Categorical"],
288
+ "aspect_ratios": ["1:1 (Native)"],
289
+ "step_min": 1,
290
+ "step_max": 128,
291
+ "step_default": 128,
292
+ "step_label": "Rows",
293
+ "preview_step_label": "Rows / preview",
294
+ "preview_step_default": 4,
295
+ },
296
+ "backend": {
297
+ "type": "python",
298
+ "module": "adam.model_plugins_builtin.pixelrow.generator",
299
+ "function": "generate",
300
+ },
301
+ }
adam/model_plugins_builtin/pixelrow/model.py ADDED
@@ -0,0 +1,244 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from __future__ import annotations
2
+
3
+ from dataclasses import asdict, dataclass
4
+ from typing import Any, Callable
5
+
6
+ import torch
7
+ from torch import Tensor, nn
8
+ from torch.nn import functional as F
9
+
10
+
11
+ MODEL_FORMAT_VERSION = 1
12
+
13
+
14
+ @dataclass(frozen=True, slots=True)
15
+ class PixelRowConfig:
16
+ resolution: int = 64
17
+ hidden_size: int = 128
18
+ recurrent_layers: int = 2
19
+ row_channels: int = 64
20
+ color_bins: int = 32
21
+
22
+ def __post_init__(self) -> None:
23
+ if self.resolution not in {32, 64, 128}:
24
+ raise ValueError("PixelRow resolution must be 32, 64, or 128.")
25
+ if not 64 <= self.hidden_size <= 1024:
26
+ raise ValueError("PixelRow hidden size must be between 64 and 1024.")
27
+ if not 1 <= self.recurrent_layers <= 4:
28
+ raise ValueError("PixelRow recurrent layers must be between 1 and 4.")
29
+ if not 16 <= self.row_channels <= 256:
30
+ raise ValueError("PixelRow row channels must be between 16 and 256.")
31
+ if self.color_bins not in {16, 32, 64}:
32
+ raise ValueError("PixelRow color bins must be 16, 32, or 64.")
33
+
34
+ def to_dict(self) -> dict[str, int]:
35
+ return asdict(self)
36
+
37
+ @classmethod
38
+ def from_dict(cls, payload: dict[str, Any]) -> "PixelRowConfig":
39
+ return cls(
40
+ resolution=int(payload.get("resolution", 64)),
41
+ hidden_size=int(payload.get("hidden_size", 128)),
42
+ recurrent_layers=int(payload.get("recurrent_layers", 2)),
43
+ row_channels=int(payload.get("row_channels", 64)),
44
+ color_bins=int(payload.get("color_bins", 32)),
45
+ )
46
+
47
+
48
+ class ResidualRowBlock(nn.Module):
49
+ def __init__(self, channels: int) -> None:
50
+ super().__init__()
51
+ groups = max(1, min(8, channels // 8))
52
+ self.norm = nn.GroupNorm(groups, channels)
53
+ self.conv1 = nn.Conv1d(channels, channels, 5, padding=2)
54
+ self.conv2 = nn.Conv1d(channels, channels, 3, padding=1)
55
+
56
+ def forward(self, value: Tensor) -> Tensor:
57
+ residual = value
58
+ value = self.conv1(F.gelu(self.norm(value)))
59
+ value = self.conv2(F.gelu(value))
60
+ return value + residual
61
+
62
+
63
+ class PixelRowModel(nn.Module):
64
+ """A row-level autoregressive model with categorical RGB outputs.
65
+
66
+ During training the GRU sees only the rows above each target row. The row
67
+ decoder predicts all pixels in the next row together, so generation takes
68
+ exactly one autoregressive decision per image row.
69
+ """
70
+
71
+ def __init__(self, config: PixelRowConfig) -> None:
72
+ super().__init__()
73
+ self.config = config
74
+ self.row_encoder = nn.Sequential(
75
+ nn.Conv1d(3, config.hidden_size, 5, padding=2),
76
+ nn.GELU(),
77
+ nn.Conv1d(config.hidden_size, config.hidden_size, 5, padding=2),
78
+ nn.GELU(),
79
+ )
80
+ self.start_embedding = nn.Parameter(
81
+ torch.zeros(1, 1, config.hidden_size, config.resolution)
82
+ )
83
+ self.row_position = nn.Embedding(config.resolution, config.hidden_size)
84
+ self.sequence = nn.GRU(
85
+ input_size=config.hidden_size,
86
+ hidden_size=config.hidden_size,
87
+ num_layers=config.recurrent_layers,
88
+ batch_first=True,
89
+ dropout=0.1 if config.recurrent_layers > 1 else 0.0,
90
+ )
91
+ self.hidden_to_row = nn.Conv1d(config.hidden_size, config.row_channels, 1)
92
+ self.column_features = nn.Parameter(
93
+ torch.randn(1, config.row_channels, config.resolution) * 0.02
94
+ )
95
+ self.row_decoder = nn.Sequential(
96
+ ResidualRowBlock(config.row_channels),
97
+ ResidualRowBlock(config.row_channels),
98
+ nn.GroupNorm(max(1, min(8, config.row_channels // 8)), config.row_channels),
99
+ nn.GELU(),
100
+ nn.Conv1d(config.row_channels, 3 * config.color_bins, 1),
101
+ )
102
+ nn.init.normal_(self.start_embedding, std=0.02)
103
+
104
+ def encode_rows(self, rows: Tensor) -> Tensor:
105
+ """Encode BxHx3xW normalized RGB rows into BxHxCxW features."""
106
+ batch, height, channels, width = rows.shape
107
+ if channels != 3 or width != self.config.resolution:
108
+ raise ValueError("PixelRow input rows do not match the model configuration.")
109
+ encoded = self.row_encoder(rows.reshape(batch * height, channels, width))
110
+ return encoded.reshape(batch, height, self.config.hidden_size, width)
111
+
112
+ def decode_hidden(self, hidden: Tensor) -> Tensor:
113
+ """Decode BxHxCxW states to BxHx3xBinsxW logits."""
114
+ batch, height, channels, width = hidden.shape
115
+ features = self.hidden_to_row(hidden.reshape(batch * height, channels, width))
116
+ features = features + self.column_features
117
+ logits = self.row_decoder(features)
118
+ return logits.reshape(
119
+ batch,
120
+ height,
121
+ 3,
122
+ self.config.color_bins,
123
+ self.config.resolution,
124
+ )
125
+
126
+ def forward(self, target_rows: Tensor) -> Tensor:
127
+ """Teacher-force the image while preserving strict top-to-bottom causality."""
128
+ batch, height, channels, width = target_rows.shape
129
+ if height != self.config.resolution or channels != 3 or width != self.config.resolution:
130
+ raise ValueError("PixelRow expects square BxHx3xW tensors at its trained resolution.")
131
+ encoded = self.encode_rows(target_rows)
132
+ inputs = torch.cat(
133
+ (self.start_embedding.expand(batch, -1, -1, -1), encoded[:, :-1]), dim=1
134
+ )
135
+ positions = self.row_position(torch.arange(height, device=target_rows.device))
136
+ inputs = inputs + positions.view(1, height, self.config.hidden_size, 1)
137
+ # Each column gets a recurrent sequence, while the row encoder and
138
+ # decoder exchange local horizontal context through 1D convolutions.
139
+ column_sequences = inputs.permute(0, 3, 1, 2).reshape(
140
+ batch * width, height, self.config.hidden_size
141
+ )
142
+ sequence_output, _state = self.sequence(column_sequences)
143
+ spatial_output = sequence_output.reshape(
144
+ batch, width, height, self.config.hidden_size
145
+ ).permute(0, 2, 3, 1).contiguous()
146
+ return self.decode_hidden(spatial_output)
147
+
148
+ def loss(self, images: Tensor, *, edge_loss_weight: float = 0.0) -> tuple[Tensor, dict[str, float]]:
149
+ targets = quantize_images(images, self.config.color_bins)
150
+ normalized = dequantize_images(targets, self.config.color_bins)
151
+ logits = self(normalized)
152
+ categorical = F.cross_entropy(
153
+ logits.permute(0, 1, 2, 4, 3).reshape(-1, self.config.color_bins),
154
+ targets.reshape(-1),
155
+ )
156
+ edge_loss = categorical.new_zeros(())
157
+ if edge_loss_weight > 0:
158
+ levels = torch.linspace(-1.0, 1.0, self.config.color_bins, device=images.device)
159
+ expected = (logits.softmax(dim=3) * levels.view(1, 1, 1, -1, 1)).sum(dim=3)
160
+ horizontal = F.l1_loss(expected[..., 1:] - expected[..., :-1], normalized[..., 1:] - normalized[..., :-1])
161
+ vertical = F.l1_loss(expected[:, 1:] - expected[:, :-1], normalized[:, 1:] - normalized[:, :-1])
162
+ edge_loss = (horizontal + vertical) * 0.5
163
+ total = categorical + float(edge_loss_weight) * edge_loss
164
+ return total, {
165
+ "categorical": float(categorical.detach().item()),
166
+ "edge": float(edge_loss.detach().item()),
167
+ }
168
+
169
+ @torch.inference_mode()
170
+ def generate(
171
+ self,
172
+ *,
173
+ rows: int | None = None,
174
+ temperature: float = 1.0,
175
+ top_k: int = 8,
176
+ generator: torch.Generator | None = None,
177
+ row_callback: Callable[[int, Tensor], None] | None = None,
178
+ ) -> Tensor:
179
+ """Generate one image and optionally report its partially completed canvas."""
180
+ self.eval()
181
+ total_rows = min(max(1, int(rows or self.config.resolution)), self.config.resolution)
182
+ temperature = max(0.05, float(temperature))
183
+ top_k = min(max(1, int(top_k)), self.config.color_bins)
184
+ device = next(self.parameters()).device
185
+ canvas = torch.zeros(1, self.config.resolution, 3, self.config.resolution, device=device)
186
+ recurrent_state: Tensor | None = None
187
+ previous_embedding: Tensor | None = None
188
+ for row_index in range(total_rows):
189
+ if previous_embedding is None:
190
+ step_input = self.start_embedding[:, 0]
191
+ else:
192
+ step_input = previous_embedding
193
+ position = self.row_position(torch.tensor([row_index], device=device)).unsqueeze(-1)
194
+ column_input = (step_input + position).permute(0, 2, 1).reshape(
195
+ self.config.resolution, 1, self.config.hidden_size
196
+ )
197
+ column_output, recurrent_state = self.sequence(column_input, recurrent_state)
198
+ spatial_output = column_output.reshape(
199
+ 1, self.config.resolution, 1, self.config.hidden_size
200
+ ).permute(0, 2, 3, 1).contiguous()
201
+ logits = self.decode_hidden(spatial_output)[:, 0] / temperature
202
+ if top_k < self.config.color_bins:
203
+ best_values, best_indices = torch.topk(logits, top_k, dim=2)
204
+ probabilities = best_values.softmax(dim=2)
205
+ sampled_offset = torch.multinomial(
206
+ probabilities.permute(0, 1, 3, 2).reshape(-1, top_k),
207
+ 1,
208
+ generator=generator,
209
+ ).reshape(1, 3, self.config.resolution)
210
+ sampled = best_indices.permute(0, 1, 3, 2).gather(
211
+ 3, sampled_offset.unsqueeze(-1)
212
+ ).squeeze(-1)
213
+ else:
214
+ probabilities = logits.softmax(dim=2)
215
+ sampled = torch.multinomial(
216
+ probabilities.permute(0, 1, 3, 2).reshape(-1, self.config.color_bins),
217
+ 1,
218
+ generator=generator,
219
+ ).reshape(1, 3, self.config.resolution)
220
+ normalized_row = dequantize_images(sampled, self.config.color_bins)
221
+ canvas[:, row_index] = normalized_row
222
+ previous_embedding = self.row_encoder(normalized_row)
223
+ if row_callback is not None:
224
+ row_callback(row_index + 1, canvas[0].detach())
225
+ return canvas[0]
226
+
227
+
228
+ def quantize_images(images: Tensor, color_bins: int) -> Tensor:
229
+ """Convert normalized RGB values in [-1, 1] to categorical color levels."""
230
+ return ((images.clamp(-1, 1) + 1.0) * 0.5 * (color_bins - 1)).round().long()
231
+
232
+
233
+ def dequantize_images(indices: Tensor, color_bins: int) -> Tensor:
234
+ """Convert categorical color levels back to normalized RGB values."""
235
+ return indices.float() * (2.0 / (color_bins - 1)) - 1.0
236
+
237
+
238
+ def canvas_to_uint8(canvas: Tensor, completed_rows: int | None = None) -> Tensor:
239
+ """Convert Hx3xW normalized rows to a display-ready HxWx3 byte tensor."""
240
+ image = ((canvas.detach().float().cpu().clamp(-1, 1) + 1.0) * 127.5).round().byte()
241
+ image = image.permute(0, 2, 1).contiguous()
242
+ if completed_rows is not None and completed_rows < image.shape[0]:
243
+ image[completed_rows:] = 32
244
+ return image
adam/model_plugins_builtin/pixelrow/trainer.py ADDED
@@ -0,0 +1,411 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from __future__ import annotations
2
+
3
+ import random
4
+ from datetime import datetime, timezone
5
+ from pathlib import Path
6
+ from typing import Any
7
+
8
+ import torch
9
+ from PIL import Image, ImageOps, ImageStat
10
+ from torch.utils.data import DataLoader, Dataset
11
+
12
+ from adam.executor import ToolExecutionError
13
+
14
+ from .common import (
15
+ CONFIG_NAME,
16
+ FINAL_CHECKPOINT_NAME,
17
+ ensure_below,
18
+ image_files,
19
+ load_checkpoint,
20
+ resolve_checkpoint,
21
+ safe_model_name,
22
+ save_canvas,
23
+ write_json,
24
+ )
25
+ from .model import MODEL_FORMAT_VERSION, PixelRowConfig, PixelRowModel
26
+
27
+
28
+ class PixelRowImageDataset(Dataset[torch.Tensor]):
29
+ def __init__(
30
+ self,
31
+ paths: list[Path],
32
+ *,
33
+ resolution: int,
34
+ resize_mode: str,
35
+ horizontal_flip: bool,
36
+ ) -> None:
37
+ self.paths = paths
38
+ self.resolution = resolution
39
+ self.resize_mode = resize_mode
40
+ self.horizontal_flip = horizontal_flip
41
+
42
+ def __len__(self) -> int:
43
+ return len(self.paths)
44
+
45
+ def __getitem__(self, index: int) -> torch.Tensor:
46
+ path = self.paths[index]
47
+ try:
48
+ with Image.open(path) as opened:
49
+ image = opened.convert("RGB")
50
+ if self.resize_mode == "fill":
51
+ image = ImageOps.fit(
52
+ image,
53
+ (self.resolution, self.resolution),
54
+ method=Image.Resampling.LANCZOS,
55
+ )
56
+ elif self.resize_mode == "fit":
57
+ mean = tuple(int(value) for value in ImageStat.Stat(image.resize((1, 1))).mean)
58
+ image = ImageOps.pad(
59
+ image,
60
+ (self.resolution, self.resolution),
61
+ method=Image.Resampling.LANCZOS,
62
+ color=mean,
63
+ )
64
+ else:
65
+ image = image.resize(
66
+ (self.resolution, self.resolution), Image.Resampling.LANCZOS
67
+ )
68
+ if self.horizontal_flip and random.random() < 0.5:
69
+ image = image.transpose(Image.Transpose.FLIP_LEFT_RIGHT)
70
+ buffer = bytearray(image.tobytes())
71
+ except (OSError, ValueError) as exc:
72
+ raise RuntimeError(f"Could not read training image {path.name}: {exc}") from exc
73
+ pixels = torch.frombuffer(buffer, dtype=torch.uint8).reshape(
74
+ self.resolution, self.resolution, 3
75
+ )
76
+ # Model layout is H rows x RGB channels x W columns.
77
+ return pixels.permute(0, 2, 1).float().div(127.5).sub(1.0)
78
+
79
+
80
+ def _checkpoint_payload(
81
+ model: PixelRowModel,
82
+ optimizer: torch.optim.Optimizer,
83
+ *,
84
+ model_name: str,
85
+ dataset_dir: Path,
86
+ completed_epochs: int,
87
+ global_step: int,
88
+ training_settings: dict[str, Any],
89
+ ) -> dict[str, Any]:
90
+ return {
91
+ "format_version": MODEL_FORMAT_VERSION,
92
+ "architecture": "autoregressive_rows",
93
+ "model_name": model_name,
94
+ "config": model.config.to_dict(),
95
+ "model_state": model.state_dict(),
96
+ "optimizer_state": optimizer.state_dict(),
97
+ "completed_epochs": int(completed_epochs),
98
+ "global_step": int(global_step),
99
+ "dataset_dir": str(dataset_dir),
100
+ "training_settings": training_settings,
101
+ "saved_at": datetime.now(timezone.utc).isoformat(),
102
+ }
103
+
104
+
105
+ def _save_checkpoint(path: Path, payload: dict[str, Any]) -> None:
106
+ path.parent.mkdir(parents=True, exist_ok=True)
107
+ temporary = path.with_suffix(path.suffix + ".tmp")
108
+ torch.save(payload, temporary)
109
+ temporary.replace(path)
110
+
111
+
112
+ def _preview(
113
+ context,
114
+ model: PixelRowModel,
115
+ output: Path,
116
+ *,
117
+ epoch: int,
118
+ next_epoch: int,
119
+ seed: int,
120
+ prompt: str,
121
+ ) -> None:
122
+ device = next(model.parameters()).device
123
+ generator = torch.Generator(device=device)
124
+ generator.manual_seed(int(seed))
125
+ canvas = model.generate(
126
+ rows=model.config.resolution,
127
+ temperature=0.85,
128
+ top_k=min(4, model.config.color_bins),
129
+ generator=generator,
130
+ )
131
+ destination = output / "previews" / f"preview_epoch_{epoch:06d}.png"
132
+ save_canvas(canvas, destination, completed_rows=model.config.resolution)
133
+ context.preview(
134
+ destination,
135
+ epoch=epoch,
136
+ next_epoch=next_epoch,
137
+ prompt=prompt,
138
+ seed=seed,
139
+ steps=model.config.resolution,
140
+ )
141
+
142
+
143
+ def train(
144
+ context,
145
+ dataset_dir: str,
146
+ model_name: str,
147
+ epochs: int,
148
+ output_dir: str,
149
+ resume_from: str = "",
150
+ resolution: int = 64,
151
+ resize_mode: str = "fill",
152
+ horizontal_flip: bool = True,
153
+ batch_size: int = 8,
154
+ learning_rate: float = 0.0002,
155
+ gradient_accumulation_steps: int = 1,
156
+ workers: int = 0,
157
+ mixed_precision: str = "fp16",
158
+ hidden_size: int = 128,
159
+ recurrent_layers: int = 2,
160
+ row_channels: int = 64,
161
+ color_bins: int = 32,
162
+ edge_loss_weight: float = 0.05,
163
+ save_every: int = 10,
164
+ preview_enabled: bool = True,
165
+ preview_every: int = 5,
166
+ preview_prompt: str = "",
167
+ preview_seed: int = 123456789,
168
+ ) -> dict[str, Any]:
169
+ """Train PixelRow inside ADAM's managed plugin-output area."""
170
+ name = safe_model_name(model_name)
171
+ dataset = Path(dataset_dir).expanduser().resolve()
172
+ if not dataset.is_dir():
173
+ raise ToolExecutionError("The selected PixelRow dataset folder no longer exists.")
174
+ training_dataset = dataset
175
+ accepted_frames = dataset / "frames"
176
+ frame_paths = image_files(accepted_frames) if accepted_frames.is_dir() else []
177
+ if frame_paths:
178
+ training_dataset = accepted_frames
179
+ paths = frame_paths
180
+ else:
181
+ paths = image_files(dataset)
182
+ if len(paths) < 2:
183
+ raise ToolExecutionError("PixelRow needs at least two readable image files before training can start.")
184
+
185
+ output_root = (context.root.resolve() / "data" / "model_plugin_outputs" / "pixelrow").resolve()
186
+ output = ensure_below(Path(output_dir), output_root, "PixelRow output")
187
+ if output.exists() and not output.is_dir():
188
+ raise ToolExecutionError("The PixelRow output path must be a folder.")
189
+ if output.exists() and any(output.iterdir()):
190
+ raise ToolExecutionError("The PixelRow output folder is not empty. Choose a new model output.")
191
+ output.mkdir(parents=True, exist_ok=True)
192
+
193
+ if resize_mode not in {"fill", "fit", "stretch"}:
194
+ raise ToolExecutionError("PixelRow image fitting must be fill, fit, or stretch.")
195
+ if not 1 <= int(epochs) <= 100_000:
196
+ raise ToolExecutionError("PixelRow epochs must be between 1 and 100000.")
197
+ if not 1 <= int(batch_size) <= 64 or not 1 <= int(gradient_accumulation_steps) <= 64:
198
+ raise ToolExecutionError("PixelRow batch size and gradient accumulation must be between 1 and 64.")
199
+ if not 1e-7 <= float(learning_rate) <= 0.1:
200
+ raise ToolExecutionError("PixelRow learning rate must be between 0.0000001 and 0.1.")
201
+ if not 0 <= int(workers) <= 16 or mixed_precision not in {"fp16", "no"}:
202
+ raise ToolExecutionError("PixelRow loader workers or precision is outside the supported range.")
203
+ if not 0.0 <= float(edge_loss_weight) <= 1.0:
204
+ raise ToolExecutionError("PixelRow line-detail strength must be between 0 and 1.")
205
+ if not 1 <= int(save_every) <= 1000 or not 1 <= int(preview_every) <= 100_000:
206
+ raise ToolExecutionError("PixelRow save and preview intervals must be positive.")
207
+
208
+ device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
209
+ resume_payload: dict[str, Any] | None = None
210
+ if resume_from.strip():
211
+ resume_path = ensure_below(
212
+ resolve_checkpoint(Path(resume_from)), output_root, "PixelRow resume checkpoint"
213
+ )
214
+ model, resume_payload = load_checkpoint(resume_path, device)
215
+ config = model.config
216
+ context.log(
217
+ "Continuing with the checkpoint architecture: "
218
+ f"{config.resolution}px, width {config.hidden_size}, {config.color_bins} color levels."
219
+ )
220
+ else:
221
+ try:
222
+ config = PixelRowConfig(
223
+ resolution=int(resolution),
224
+ hidden_size=int(hidden_size),
225
+ recurrent_layers=int(recurrent_layers),
226
+ row_channels=int(row_channels),
227
+ color_bins=int(color_bins),
228
+ )
229
+ except ValueError as exc:
230
+ raise ToolExecutionError(str(exc)) from exc
231
+ model = PixelRowModel(config).to(device)
232
+
233
+ dataset_object = PixelRowImageDataset(
234
+ paths,
235
+ resolution=config.resolution,
236
+ resize_mode=resize_mode,
237
+ horizontal_flip=bool(horizontal_flip),
238
+ )
239
+ loader = DataLoader(
240
+ dataset_object,
241
+ batch_size=int(batch_size),
242
+ shuffle=True,
243
+ num_workers=int(workers),
244
+ pin_memory=device.type == "cuda",
245
+ drop_last=False,
246
+ )
247
+ optimizer = torch.optim.AdamW(model.parameters(), lr=float(learning_rate), betas=(0.9, 0.95))
248
+ start_epoch = 0
249
+ global_step = 0
250
+ if resume_payload is not None:
251
+ start_epoch = int(resume_payload.get("completed_epochs", 0) or 0)
252
+ global_step = int(resume_payload.get("global_step", 0) or 0)
253
+ optimizer_state = resume_payload.get("optimizer_state")
254
+ if isinstance(optimizer_state, dict):
255
+ try:
256
+ optimizer.load_state_dict(optimizer_state)
257
+ for group in optimizer.param_groups:
258
+ group["lr"] = float(learning_rate)
259
+ except (ValueError, RuntimeError):
260
+ context.log("The previous optimizer state was incompatible; continuing with a fresh optimizer.")
261
+
262
+ use_fp16 = mixed_precision == "fp16" and device.type == "cuda"
263
+ if mixed_precision == "fp16" and not use_fp16:
264
+ context.log("FP16 requires CUDA; PixelRow will train in full precision on this device.")
265
+ try:
266
+ scaler = torch.amp.GradScaler("cuda", enabled=use_fp16)
267
+ except (AttributeError, TypeError): # PyTorch 2.2 compatibility.
268
+ scaler = torch.cuda.amp.GradScaler(enabled=use_fp16)
269
+ accumulation = int(gradient_accumulation_steps)
270
+ requested_epochs = int(epochs)
271
+ final_epoch = start_epoch + requested_epochs
272
+ batches_per_epoch = max(1, len(loader))
273
+ total_batches = requested_epochs * batches_per_epoch
274
+ settings = {
275
+ "resolution": config.resolution,
276
+ "resize_mode": resize_mode,
277
+ "horizontal_flip": bool(horizontal_flip),
278
+ "batch_size": int(batch_size),
279
+ "learning_rate": float(learning_rate),
280
+ "gradient_accumulation_steps": accumulation,
281
+ "workers": int(workers),
282
+ "mixed_precision": mixed_precision,
283
+ "hidden_size": config.hidden_size,
284
+ "recurrent_layers": config.recurrent_layers,
285
+ "row_channels": config.row_channels,
286
+ "color_bins": config.color_bins,
287
+ "edge_loss_weight": float(edge_loss_weight),
288
+ }
289
+ write_json(output / CONFIG_NAME, {
290
+ "format_version": MODEL_FORMAT_VERSION,
291
+ "model_type": "pixelrow",
292
+ "model_name": name,
293
+ **config.to_dict(),
294
+ })
295
+ context.log(
296
+ f"Training PixelRow on {len(paths)} images from {training_dataset} at "
297
+ f"{config.resolution}x{config.resolution}, "
298
+ f"batch {batch_size}, learning rate {learning_rate}, device {device}."
299
+ )
300
+ optimizer.zero_grad(set_to_none=True)
301
+ processed_batches = 0
302
+ last_loss = 0.0
303
+ try:
304
+ for epoch in range(start_epoch + 1, final_epoch + 1):
305
+ model.train()
306
+ epoch_loss = 0.0
307
+ for batch_index, images in enumerate(loader, 1):
308
+ context.checkpoint()
309
+ images = images.to(device, non_blocking=device.type == "cuda")
310
+ with torch.autocast(device_type=device.type, dtype=torch.float16, enabled=use_fp16):
311
+ loss, parts = model.loss(images, edge_loss_weight=float(edge_loss_weight))
312
+ scaled_loss = loss / accumulation
313
+ scaler.scale(scaled_loss).backward()
314
+ if batch_index % accumulation == 0 or batch_index == batches_per_epoch:
315
+ scaler.unscale_(optimizer)
316
+ torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
317
+ scaler.step(optimizer)
318
+ scaler.update()
319
+ optimizer.zero_grad(set_to_none=True)
320
+ global_step += 1
321
+ last_loss = float(loss.detach().item())
322
+ epoch_loss += last_loss
323
+ processed_batches += 1
324
+ percent = max(1, min(99, round(processed_batches * 100 / total_batches)))
325
+ context.progress(
326
+ percent,
327
+ f"Epoch {epoch} of {final_epoch} · loss {last_loss:.4f}",
328
+ epoch=epoch,
329
+ total_epochs=final_epoch,
330
+ current_step=processed_batches,
331
+ total_steps=total_batches,
332
+ unit="batch",
333
+ loss=last_loss,
334
+ categorical_loss=parts["categorical"],
335
+ edge_loss=parts["edge"],
336
+ )
337
+
338
+ payload = _checkpoint_payload(
339
+ model,
340
+ optimizer,
341
+ model_name=name,
342
+ dataset_dir=dataset,
343
+ completed_epochs=epoch,
344
+ global_step=global_step,
345
+ training_settings=settings,
346
+ )
347
+ if epoch % int(save_every) == 0:
348
+ _save_checkpoint(output / "checkpoints" / f"epoch_{epoch:06d}.pt", payload)
349
+ if bool(preview_enabled) and epoch % int(preview_every) == 0:
350
+ _preview(
351
+ context,
352
+ model,
353
+ output,
354
+ epoch=epoch,
355
+ next_epoch=min(final_epoch, epoch + int(preview_every)),
356
+ seed=int(preview_seed),
357
+ prompt=preview_prompt,
358
+ )
359
+ context.log(f"Finished epoch {epoch}; average loss {epoch_loss / batches_per_epoch:.4f}.")
360
+ except torch.cuda.OutOfMemoryError as exc:
361
+ if device.type == "cuda":
362
+ torch.cuda.empty_cache()
363
+ raise ToolExecutionError(
364
+ "PixelRow ran out of VRAM. Reduce batch size first, then sequence width or resolution."
365
+ ) from exc
366
+
367
+ final_payload = _checkpoint_payload(
368
+ model,
369
+ optimizer,
370
+ model_name=name,
371
+ dataset_dir=dataset,
372
+ completed_epochs=final_epoch,
373
+ global_step=global_step,
374
+ training_settings=settings,
375
+ )
376
+ final_checkpoint = output / FINAL_CHECKPOINT_NAME
377
+ _save_checkpoint(final_checkpoint, final_payload)
378
+ write_json(output / "training_metadata.json", {
379
+ "format_version": MODEL_FORMAT_VERSION,
380
+ "model_type": "pixelrow",
381
+ "architecture": "autoregressive_rows",
382
+ "model_name": name,
383
+ "dataset_dir": str(dataset),
384
+ "image_count": len(paths),
385
+ "completed_epochs": final_epoch,
386
+ "epochs_this_run": requested_epochs,
387
+ "global_step": global_step,
388
+ "final_loss": last_loss,
389
+ "checkpoint": str(final_checkpoint),
390
+ "settings": settings,
391
+ "finished_at": datetime.now(timezone.utc).isoformat(),
392
+ })
393
+ context.progress(100, "PixelRow training completed")
394
+ return {
395
+ "output_folder": str(output),
396
+ "model_name": name,
397
+ "assets": [{
398
+ "kind": "model",
399
+ "name": name,
400
+ "path": str(output),
401
+ "trainer": "pixelrow",
402
+ "dataset_path": str(dataset),
403
+ "checkpoint": str(final_checkpoint),
404
+ "epochs": final_epoch,
405
+ "metadata": {
406
+ "architecture": "autoregressive_rows",
407
+ "resolution": config.resolution,
408
+ "color_bins": config.color_bins,
409
+ },
410
+ }],
411
+ }
adam/model_plugins_builtin/sdxl_lora/manifest.py CHANGED
@@ -45,7 +45,7 @@ GENERATION_SETTINGS = {
45
  "image_count": {"label": "Images", "type": "int", "default": 1, "min": 1, "max": 48, "group": "Generation"},
46
  "steps": {"label": "Steps", "type": "int", "default": 30, "min": 1, "max": 150, "group": "Generation"},
47
  "cfg_scale": {"label": "CFG scale", "type": "float", "default": 7.0, "min": 0.1, "max": 30.0, "decimals": 2, "step": 0.5, "group": "Generation"},
48
- "sampler": {"label": "Sampler", "type": "choice", "options": ["DPM++ 2M", "DPM++ SDE", "Euler", "Euler a", "DDIM"], "default": "DPM++ 2M", "group": "Generation"},
49
  "aspect_ratio": {"label": "Aspect ratio", "type": "choice", "options": ["1:1 (Square)", "4:3 (Landscape)", "3:4 (Portrait)", "3:2 (Landscape)", "2:3 (Portrait)", "16:9 (Widescreen)", "9:16 (Vertical)"], "default": "1:1 (Square)", "group": "Generation"},
50
  "width": {"label": "Width", "type": "int", "default": 1024, "min": 256, "max": 2048, "group": "Generation"},
51
  "height": {"label": "Height", "type": "int", "default": 1024, "min": 256, "max": 2048, "group": "Generation"},
 
45
  "image_count": {"label": "Images", "type": "int", "default": 1, "min": 1, "max": 48, "group": "Generation"},
46
  "steps": {"label": "Steps", "type": "int", "default": 30, "min": 1, "max": 150, "group": "Generation"},
47
  "cfg_scale": {"label": "CFG scale", "type": "float", "default": 7.0, "min": 0.1, "max": 30.0, "decimals": 2, "step": 0.5, "group": "Generation"},
48
+ "sampler": {"label": "Sampler", "type": "choice", "options": ["DPM++ 2M", "DPM++ 2M Karras", "DPM++ 2M SDE", "DPM++ 2M SDE Karras", "DPM++ SDE", "DPM++ SDE Karras", "Euler", "Euler a", "Heun", "LMS", "DDIM"], "default": "DPM++ 2M", "group": "Generation"},
49
  "aspect_ratio": {"label": "Aspect ratio", "type": "choice", "options": ["1:1 (Square)", "4:3 (Landscape)", "3:4 (Portrait)", "3:2 (Landscape)", "2:3 (Portrait)", "16:9 (Widescreen)", "9:16 (Vertical)"], "default": "1:1 (Square)", "group": "Generation"},
50
  "width": {"label": "Width", "type": "int", "default": 1024, "min": 256, "max": 2048, "group": "Generation"},
51
  "height": {"label": "Height", "type": "int", "default": 1024, "min": 256, "max": 2048, "group": "Generation"},
adam/model_plugins_builtin/wan_video/__init__.py ADDED
@@ -0,0 +1 @@
 
 
1
+ """Wan 2.1 video LoRA integration."""
adam/model_plugins_builtin/wan_video/manifest.py ADDED
@@ -0,0 +1,50 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ PLUGIN_ID = "wan_video"
2
+
3
+ MODEL_INFO = {
4
+ "name": "Wan Video LoRA", "version": "1.0", "category": "Video Generation",
5
+ "description": "Train Wan 2.1 T2V 1.3B LoRAs on captioned clips and generate MP4 videos with the connected LoRAVideoTrainer.",
6
+ "architecture": "wan21_t2v_1_3b_lora", "status": "experimental", "output_type": "video",
7
+ "workspace": "video_lora",
8
+ "input_formats": ["captioned video folder"], "output_formats": ["safetensors", "mp4"],
9
+ "hardware": {"recommended_vram_gb": 12, "recommended_system_ram_gb": 32},
10
+ "capabilities": ["fresh_training", "resume_training", "video_generation"],
11
+ }
12
+
13
+ TRAINING_SETTINGS = {
14
+ "trigger_word": {"label": "Trigger word", "type": "text", "default": "subject_token", "required": True, "group": "Dataset"},
15
+ "resolution": {"label": "Training resolution", "type": "choice", "options": ["448x256", "256x448"], "default": "448x256", "group": "Dataset"},
16
+ "target_frames": {"label": "Frame buckets", "type": "choice", "options": ["25", "49", "25,49"], "default": "25,49", "group": "Dataset", "description": "A clip must have at least the shortest selected frame count. Longer clips can supply multiple buckets."},
17
+ "batch_size": {"label": "Batch size", "type": "int", "default": 1, "min": 1, "max": 4, "group": "Training"},
18
+ "learning_rate": {"label": "Learning rate", "type": "float", "default": 0.0001, "min": 0.0000001, "max": 0.01, "decimals": 7, "group": "Training"},
19
+ "rank": {"label": "LoRA rank", "type": "int", "default": 16, "min": 1, "max": 128, "group": "Training"},
20
+ "alpha": {"label": "LoRA alpha", "type": "int", "default": 16, "min": 1, "max": 128, "group": "Training"},
21
+ "blocks_to_swap": {"label": "Training blocks to swap", "type": "int", "default": 20, "min": 0, "max": 29, "group": "Memory", "description": "More swapping reduces GPU memory use and increases CPU transfer time."},
22
+ "save_every": {"label": "Save every N epochs", "type": "int", "default": 2, "min": 1, "max": 1000, "group": "Checkpoints"},
23
+ "seed": {"label": "Training seed", "type": "int", "default": 42, "min": 0, "max": 2147483647, "group": "Advanced", "advanced": True},
24
+ }
25
+
26
+ GENERATION_SETTINGS = {
27
+ "prompt": {"label": "Prompt", "type": "multiline_text", "default": "", "required": True, "group": "Prompt"},
28
+ "format": {"label": "Video format", "type": "choice", "options": ["Landscape 832x480", "Portrait 480x832"], "default": "Landscape 832x480", "group": "Video"},
29
+ "duration": {"label": "Requested seconds", "type": "float", "default": 2.0, "min": 2.0, "max": 15.0, "decimals": 2, "step": 0.5, "group": "Video"},
30
+ "fps": {"label": "Playback FPS", "type": "int", "default": 12, "min": 4, "max": 60, "group": "Video"},
31
+ "steps": {"label": "Inference steps", "type": "int", "default": 20, "min": 1, "max": 100, "group": "Generation"},
32
+ "lora_strength": {"label": "LoRA strength", "type": "float", "default": 0.8, "min": 0.0, "max": 2.0, "decimals": 2, "step": 0.05, "group": "Generation"},
33
+ "seed": {"label": "Seed", "type": "int", "default": 1701, "min": 0, "max": 2147483647, "group": "Generation"},
34
+ "randomize_seed": {"label": "Randomize seed", "type": "bool", "default": False, "group": "Generation"},
35
+ "blocks_to_swap": {"label": "Generation blocks to swap", "type": "int", "default": 24, "min": 0, "max": 29, "group": "Memory"},
36
+ "experimental_speed": {"label": "Experimental TF32 speed mode", "type": "bool", "default": False, "group": "Advanced", "advanced": True},
37
+ }
38
+
39
+ TRAINING_TOOL = {
40
+ "id": "wan_video_trainer", "name": "Wan Video LoRA Trainer",
41
+ "capabilities": ["fresh_training", "resume_training", "progress", "pause", "cancel"],
42
+ "backend": {"type": "python", "module": "adam.tools.wan_video_adapter", "function": "train"},
43
+ }
44
+ GENERATION_TOOL = {
45
+ "id": "wan_video_generator", "name": "Wan Video Generator", "model_trainers": ["wan_video"],
46
+ "arguments": ["model_name", "model_path", *GENERATION_SETTINGS],
47
+ "required_arguments": ["model_path", "prompt"],
48
+ "capabilities": ["video_generation", "progress", "pause", "cancel"],
49
+ "backend": {"type": "python", "module": "adam.tools.wan_video_adapter", "function": "generate"},
50
+ }
adam/nova.py CHANGED
@@ -36,6 +36,8 @@ def _candidate_images(job: Job, limit: int = 64) -> list[Path]:
36
 
37
  def evaluate_job_output(job: Job) -> dict[str, Any]:
38
  """Evaluate technical sample health without claiming to judge artistic quality."""
 
 
39
  if not any(step.tool_id.endswith("_trainer") for step in job.plan.steps):
40
  return {}
41
  paths = _candidate_images(job)
 
36
 
37
  def evaluate_job_output(job: Job) -> dict[str, Any]:
38
  """Evaluate technical sample health without claiming to judge artistic quality."""
39
+ if any(step.tool_id == "wan_video_trainer" for step in job.plan.steps):
40
+ return {"agent": "NOVA", "status": "NEEDS VIDEO SAMPLES", "summary": "Wan training saved adapter weights. Generate a fixed-seed video in Video LoRA and review motion, subject consistency and flicker.", "sample_count": 0}
41
  if not any(step.tool_id.endswith("_trainer") for step in job.plan.steps):
42
  return {}
43
  paths = _candidate_images(job)
adam/oasis_dataset.py CHANGED
@@ -24,6 +24,7 @@ DERIVED_ACTIONS = {
24
  SUPPORTED_ACTIONS = BINARY_ACTIONS | CONTINUOUS_ACTIONS | DERIVED_ACTIONS
25
  REQUIRED_CANONICAL_ACTIONS = {"w", "a", "s", "d", "jump"}
26
  IMAGE_EXTENSIONS = {".jpg", ".jpeg", ".png", ".webp", ".bmp"}
 
27
  METADATA_FIELDS = {
28
  "session_id", "session_started_at", "frame_index", "filename",
29
  "timestamp_seconds", "camera_encoding",
@@ -39,7 +40,11 @@ class OasisDatasetReport:
39
  valid_transitions: int = 0
40
  sessions: int = 0
41
  resolution: str = ""
 
 
 
42
  action_counts: dict[str, int] = field(default_factory=dict)
 
43
  errors: list[str] = field(default_factory=list)
44
  warnings: list[str] = field(default_factory=list)
45
 
@@ -48,6 +53,61 @@ class OasisDatasetReport:
48
  return not self.errors
49
 
50
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
51
  def dataset_directories(value: str | list[str] | tuple[str, ...]) -> list[Path]:
52
  entries = value if isinstance(value, (list, tuple)) else str(value or "").split(";")
53
  directories: list[Path] = []
@@ -66,15 +126,22 @@ def _numeric_frame_index(path: Path) -> int | None:
66
  return int(match.group(1)) if match else None
67
 
68
 
69
- def validate_oasis_dataset(value: str | list[str] | tuple[str, ...], *, frame_gap: int = 1) -> OasisDatasetReport:
 
 
 
70
  report = OasisDatasetReport()
71
  frame_gap = max(1, int(frame_gap))
 
 
 
 
72
  directories = dataset_directories(value)
73
  if not directories:
74
  report.errors.append("Select at least one Oasis action dataset folder.")
75
  return report
76
  seen_resolution: tuple[int, int] | None = None
77
- action_counts = {name: 0 for name in sorted(REQUIRED_CANONICAL_ACTIONS | CONTINUOUS_ACTIONS)}
78
  transition_total = 0
79
  session_ids: set[str] = set()
80
 
@@ -145,19 +212,20 @@ def validate_oasis_dataset(value: str | list[str] | tuple[str, ...], *, frame_ga
145
  f"{actions_path.name} line {line_number} points to missing frame {filename}; skipping row."
146
  )
147
  continue
148
- try:
149
- with Image.open(frame_path) as image:
150
- image.verify()
151
- with Image.open(frame_path) as image:
152
- size = image.size
153
- except Exception as exc:
154
- report.errors.append(f"Broken image file {frame_path.name}: {exc}")
155
- continue
156
- if seen_resolution is None:
157
- seen_resolution = size
158
- report.resolution = f"{size[0]}x{size[1]}"
159
- elif size != seen_resolution:
160
- report.errors.append(f"Inconsistent frame resolution: {frame_path.name} is {size[0]}x{size[1]}, expected {seen_resolution[0]}x{seen_resolution[1]}.")
 
161
  unexpected = sorted(set(row) - SUPPORTED_ACTIONS - METADATA_FIELDS)
162
  if unexpected:
163
  report.errors.append(f"{actions_path.name} line {line_number} contains unsupported action field(s): {', '.join(unexpected[:6])}.")
@@ -176,6 +244,7 @@ def validate_oasis_dataset(value: str | list[str] | tuple[str, ...], *, frame_ga
176
  report.errors.append(f"{actions_path.name} repeats frame_index {frame_index} in session {session_id}.")
177
  continue
178
  seen_keys.add(key)
 
179
  for name in action_counts:
180
  try:
181
  value = float(row.get(name, 0))
@@ -184,6 +253,9 @@ def validate_oasis_dataset(value: str | list[str] | tuple[str, ...], *, frame_ga
184
  value = 0.0
185
  if abs(value) > (0.5 if name in BINARY_ACTIONS else 0.02):
186
  action_counts[name] += 1
 
 
 
187
  row["_session_id"] = session_id
188
  rows_by_session.setdefault(session_id, []).append(row)
189
  session_ids.add(f"{resolved}:{session_id}")
@@ -213,3 +285,10 @@ def validate_oasis_dataset(value: str | list[str] | tuple[str, ...], *, frame_ga
213
  if report.valid_rows and not any(action_counts.values()):
214
  report.errors.append("No non-idle action labels were found. Record idle plus at least one active control.")
215
  return report
 
 
 
 
 
 
 
 
24
  SUPPORTED_ACTIONS = BINARY_ACTIONS | CONTINUOUS_ACTIONS | DERIVED_ACTIONS
25
  REQUIRED_CANONICAL_ACTIONS = {"w", "a", "s", "d", "jump"}
26
  IMAGE_EXTENSIONS = {".jpg", ".jpeg", ".png", ".webp", ".bmp"}
27
+ RECORDED_ACTIONS = BINARY_ACTIONS | CONTINUOUS_ACTIONS
28
  METADATA_FIELDS = {
29
  "session_id", "session_started_at", "frame_index", "filename",
30
  "timestamp_seconds", "camera_encoding",
 
40
  valid_transitions: int = 0
41
  sessions: int = 0
42
  resolution: str = ""
43
+ capture_fps: float | None = None
44
+ recommended_frame_gap: int | None = None
45
+ native_ai_fps: float | None = None
46
  action_counts: dict[str, int] = field(default_factory=dict)
47
+ idle_rows: int = 0
48
  errors: list[str] = field(default_factory=list)
49
  warnings: list[str] = field(default_factory=list)
50
 
 
53
  return not self.errors
54
 
55
 
56
+ def _read_dataset_info(directory: Path) -> dict[str, Any]:
57
+ """Read optional recorder metadata without making it a dataset requirement."""
58
+ try:
59
+ value = json.loads((directory / "dataset_info.json").read_text(encoding="utf-8"))
60
+ return value if isinstance(value, dict) else {}
61
+ except (OSError, json.JSONDecodeError):
62
+ return {}
63
+
64
+
65
+ def _recommended_frame_gap(capture_fps: float, _camera_encoding: str) -> int:
66
+ """Choose a playable horizon that targets about 12 genuine AI frames per second.
67
+
68
+ This intentionally differs from the external trainer's older movement-focused
69
+ heuristic. A 12–15 FPS recording should train at gap 1, rather than being
70
+ slowed to a 3–5 FPS playable world before GPU speed is even considered.
71
+ """
72
+ target_ai_fps = 12.0
73
+ return max(1, min(12, round(capture_fps / target_ai_fps)))
74
+
75
+
76
+ def oasis_pace(
77
+ value: str | list[str] | tuple[str, ...], *, frame_gap: int,
78
+ ) -> dict[str, float | int | None]:
79
+ """Return portable pacing information derived from connected recorder metadata.
80
+
81
+ A model only makes one genuine frame for every prediction horizon. Display
82
+ interpolation can look smoother, but cannot make controls more responsive.
83
+ """
84
+ rates: list[float] = []
85
+ camera_encodings: list[str] = []
86
+ for directory in dataset_directories(value):
87
+ info = _read_dataset_info(directory)
88
+ try:
89
+ rate = float(info.get("capture_fps"))
90
+ except (TypeError, ValueError):
91
+ continue
92
+ if rate > 0:
93
+ rates.append(rate)
94
+ camera_encodings.append(str(info.get("camera_encoding", "legacy_pixels")))
95
+ if not rates or len({round(rate, 6) for rate in rates}) != 1:
96
+ return {"capture_fps": None, "recommended_frame_gap": None, "native_ai_fps": None}
97
+ capture_fps = rates[0]
98
+ camera_encoding = (
99
+ "relative_degrees_v1"
100
+ if "relative_degrees_v1" in camera_encodings
101
+ else camera_encodings[0]
102
+ )
103
+ gap = max(1, int(frame_gap))
104
+ return {
105
+ "capture_fps": capture_fps,
106
+ "recommended_frame_gap": _recommended_frame_gap(capture_fps, camera_encoding),
107
+ "native_ai_fps": capture_fps / gap,
108
+ }
109
+
110
+
111
  def dataset_directories(value: str | list[str] | tuple[str, ...]) -> list[Path]:
112
  entries = value if isinstance(value, (list, tuple)) else str(value or "").split(";")
113
  directories: list[Path] = []
 
126
  return int(match.group(1)) if match else None
127
 
128
 
129
+ def validate_oasis_dataset(
130
+ value: str | list[str] | tuple[str, ...], *, frame_gap: int = 1,
131
+ verify_images: bool = True,
132
+ ) -> OasisDatasetReport:
133
  report = OasisDatasetReport()
134
  frame_gap = max(1, int(frame_gap))
135
+ pace = oasis_pace(value, frame_gap=frame_gap)
136
+ report.capture_fps = pace["capture_fps"] # type: ignore[assignment]
137
+ report.recommended_frame_gap = pace["recommended_frame_gap"] # type: ignore[assignment]
138
+ report.native_ai_fps = pace["native_ai_fps"] # type: ignore[assignment]
139
  directories = dataset_directories(value)
140
  if not directories:
141
  report.errors.append("Select at least one Oasis action dataset folder.")
142
  return report
143
  seen_resolution: tuple[int, int] | None = None
144
+ action_counts = {name: 0 for name in sorted(RECORDED_ACTIONS)}
145
  transition_total = 0
146
  session_ids: set[str] = set()
147
 
 
212
  f"{actions_path.name} line {line_number} points to missing frame {filename}; skipping row."
213
  )
214
  continue
215
+ if verify_images:
216
+ try:
217
+ with Image.open(frame_path) as image:
218
+ image.verify()
219
+ with Image.open(frame_path) as image:
220
+ size = image.size
221
+ except Exception as exc:
222
+ report.errors.append(f"Broken image file {frame_path.name}: {exc}")
223
+ continue
224
+ if seen_resolution is None:
225
+ seen_resolution = size
226
+ report.resolution = f"{size[0]}x{size[1]}"
227
+ elif size != seen_resolution:
228
+ report.errors.append(f"Inconsistent frame resolution: {frame_path.name} is {size[0]}x{size[1]}, expected {seen_resolution[0]}x{seen_resolution[1]}.")
229
  unexpected = sorted(set(row) - SUPPORTED_ACTIONS - METADATA_FIELDS)
230
  if unexpected:
231
  report.errors.append(f"{actions_path.name} line {line_number} contains unsupported action field(s): {', '.join(unexpected[:6])}.")
 
244
  report.errors.append(f"{actions_path.name} repeats frame_index {frame_index} in session {session_id}.")
245
  continue
246
  seen_keys.add(key)
247
+ action_active = False
248
  for name in action_counts:
249
  try:
250
  value = float(row.get(name, 0))
 
253
  value = 0.0
254
  if abs(value) > (0.5 if name in BINARY_ACTIONS else 0.02):
255
  action_counts[name] += 1
256
+ action_active = True
257
+ if not action_active:
258
+ report.idle_rows += 1
259
  row["_session_id"] = session_id
260
  rows_by_session.setdefault(session_id, []).append(row)
261
  session_ids.add(f"{resolved}:{session_id}")
 
285
  if report.valid_rows and not any(action_counts.values()):
286
  report.errors.append("No non-idle action labels were found. Record idle plus at least one active control.")
287
  return report
288
+
289
+
290
+ def inspect_oasis_dataset(
291
+ value: str | list[str] | tuple[str, ...], *, frame_gap: int = 1,
292
+ ) -> OasisDatasetReport:
293
+ """Quickly inspect training-relevant labels and pacing without decoding images."""
294
+ return validate_oasis_dataset(value, frame_gap=frame_gap, verify_images=False)
adam/oasis_player.py ADDED
@@ -0,0 +1,112 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Small, UI-independent helpers for ADAM's native Oasis Player page."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import json
6
+ import random
7
+ from pathlib import Path
8
+
9
+
10
+ IMAGE_EXTENSIONS = {".png", ".jpg", ".jpeg", ".webp", ".bmp"}
11
+ ACTION_MODEL_TYPES = {
12
+ "action_conditioned_rectified_flow_video",
13
+ "action_conditioned_latent_vae_flow_video",
14
+ "action_conditioned_temporal_latent_flow",
15
+ "action_conditioned_temporal_pixel_flow",
16
+ }
17
+ LATENT_MODEL_TYPES = {
18
+ "action_conditioned_latent_vae_flow_video",
19
+ "action_conditioned_temporal_latent_flow",
20
+ }
21
+
22
+
23
+ def is_action_model(path: str | Path) -> bool:
24
+ """Return whether *path* is a complete Oasis action-world-model folder."""
25
+ folder = Path(path)
26
+ try:
27
+ info = json.loads((folder / "action_flow_model_info.json").read_text(encoding="utf-8"))
28
+ except (OSError, ValueError, TypeError):
29
+ return False
30
+ weights = (
31
+ folder / "unet" / "diffusion_pytorch_model.safetensors",
32
+ folder / "unet" / "diffusion_pytorch_model.bin",
33
+ )
34
+ model_type = info.get("model_type")
35
+ has_unet = (folder / "unet" / "config.json").is_file() and any(candidate.is_file() for candidate in weights)
36
+ has_vae = (folder / "vae" / "config.json").is_file() and (folder / "vae" / "pytorch_model.bin").is_file()
37
+ return model_type in ACTION_MODEL_TYPES and has_unet and (model_type not in LATENT_MODEL_TYPES or has_vae)
38
+
39
+
40
+ def model_info(path: str | Path) -> dict:
41
+ """Read compatible model metadata, raising a useful error for the UI."""
42
+ folder = Path(path)
43
+ if not is_action_model(folder):
44
+ raise ValueError("Choose a complete Oasis action model folder.")
45
+ return json.loads((folder / "action_flow_model_info.json").read_text(encoding="utf-8"))
46
+
47
+
48
+ def frame_size(info: dict) -> tuple[int, int]:
49
+ """Return the trained (width, height), including legacy checkpoint metadata."""
50
+ if info.get("width") and info.get("height"):
51
+ return int(info["width"]), int(info["height"])
52
+ value = str(info.get("resolution", "256x144")).lower().replace("×", "x")
53
+ if "x" in value:
54
+ width, height = (int(part.strip()) for part in value.split("x", 1))
55
+ return width, height
56
+ width = int(value)
57
+ return width, round(width * 9 / 16)
58
+
59
+
60
+ def discover_models(oasis_root: str | Path, assets=()) -> list[tuple[str, Path]]:
61
+ """Find local Oasis models without copying, moving, or modifying them."""
62
+ found: dict[Path, str] = {}
63
+ root = Path(oasis_root)
64
+ library = root / "output_action_flow_models"
65
+ if library.is_dir():
66
+ for folder in library.iterdir():
67
+ if folder.is_dir() and is_action_model(folder):
68
+ found[folder.resolve()] = folder.name
69
+ for asset in assets:
70
+ if getattr(asset, "kind", "") != "model" or getattr(asset, "trainer", "") != "oasis":
71
+ continue
72
+ folder = Path(getattr(asset, "path", ""))
73
+ if is_action_model(folder):
74
+ found[folder.resolve()] = getattr(asset, "name", "") or folder.name
75
+ return sorted(((name, path) for path, name in found.items()), key=lambda item: item[0].casefold())
76
+
77
+
78
+ def capture_path(root: str | Path, prefix: str = "oasis_frame") -> Path:
79
+ """Choose an ADAM-owned, collision-resistant PNG destination."""
80
+ from datetime import datetime
81
+ from uuid import uuid4
82
+
83
+ folder = Path(root) / "data" / "oasis_captures"
84
+ folder.mkdir(parents=True, exist_ok=True)
85
+ return folder / f"{prefix}_{datetime.now().strftime('%Y%m%d_%H%M%S_%f')}_{uuid4().hex[:6]}.png"
86
+
87
+
88
+ def random_roblox_dataset_frame(oasis_root: str | Path, chooser=None) -> Path:
89
+ """Pick one image from a deeply nested Roblox Dataset without assuming its layout.
90
+
91
+ Oasis recordings often live in several named recording folders, each with a
92
+ ``frames`` folder. Reservoir sampling avoids holding every path in memory.
93
+ ``chooser`` is injectable for deterministic tests.
94
+ """
95
+ root = Path(oasis_root) / "OldDatasets" / "Roblox Dataset"
96
+ if not root.is_dir():
97
+ raise FileNotFoundError("The connected Oasis Trainer has no OldDatasets/Roblox Dataset folder.")
98
+ pick = None
99
+ count = 0
100
+ chooser = chooser or random.randrange
101
+ try:
102
+ for path in root.rglob("*"):
103
+ if not path.is_file() or path.suffix.casefold() not in IMAGE_EXTENSIONS:
104
+ continue
105
+ count += 1
106
+ if chooser(count) == 0:
107
+ pick = path
108
+ except OSError as exc:
109
+ raise OSError(f"ADAM could not read the Roblox Dataset: {exc}") from exc
110
+ if pick is None:
111
+ raise FileNotFoundError("No PNG, JPG, WEBP, or BMP frames were found in the Roblox Dataset.")
112
+ return pick
adam/ollama.py CHANGED
@@ -1,8 +1,10 @@
1
  from __future__ import annotations
2
 
3
  import json
 
4
  import urllib.error
5
  import urllib.request
 
6
  from typing import Any, Callable
7
 
8
 
@@ -17,22 +19,56 @@ class OllamaClient:
17
  model: str,
18
  timeout: float = 2.5,
19
  chat_max_tokens: int | None = None,
 
20
  ) -> None:
21
  self.base_url = base_url.rstrip("/")
22
  self.model = model
23
  self.timeout = timeout
24
  self.chat_max_tokens = chat_max_tokens
 
25
 
26
  def _chat_system(self, system: str) -> str:
27
  """Return the application system prompt unchanged."""
28
  return system
29
 
30
- def _num_predict(self, default: int, *, chat: bool = True) -> int:
31
  """Qwen3's reasoning commonly needs more than a short-chat token budget."""
32
  if chat and self.chat_max_tokens is not None:
33
- return self.chat_max_tokens
 
 
 
 
 
 
 
 
 
 
 
 
 
 
34
  return 1024 if self.model.casefold().startswith("qwen3") else default
35
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
36
  def is_available(self, timeout: float = 0.35) -> bool:
37
  request = urllib.request.Request(f"{self.base_url}/api/tags", method="GET")
38
  try:
@@ -64,8 +100,10 @@ class OllamaClient:
64
  raise OllamaError("Ollama returned an unsupported plan shape.")
65
  return parsed
66
 
67
- def generate_text(self, system: str, prompt: str) -> str:
68
- response = self._generate(system, prompt, json_format=False).strip()
 
 
69
  if not response:
70
  raise OllamaError("Ollama returned an empty response.")
71
  return response
@@ -75,14 +113,21 @@ class OllamaClient:
75
  system: str,
76
  prompt: str,
77
  on_chunk: Callable[[str], None],
 
78
  ) -> str:
 
79
  payload: dict[str, Any] = {
80
  "model": self.model,
81
  "system": self._chat_system(system),
82
  "prompt": prompt,
 
 
 
83
  "stream": True,
84
- "options": {"temperature": 0.35, "num_predict": self._num_predict(180)},
85
  }
 
 
86
  request = urllib.request.Request(
87
  f"{self.base_url}/api/generate",
88
  data=json.dumps(payload).encode("utf-8"),
@@ -109,19 +154,23 @@ class OllamaClient:
109
  raise OllamaError("Ollama returned an empty response.")
110
  return result
111
 
112
- def _generate(self, system: str, prompt: str, *, json_format: bool) -> str:
 
113
  payload: dict[str, Any] = {
114
  "model": self.model,
115
  "system": self._chat_system(system),
116
  "prompt": prompt,
 
117
  "stream": False,
118
  "options": {
119
  "temperature": 0.1 if json_format else 0.35,
120
- "num_predict": self._num_predict(300 if json_format else 180, chat=not json_format),
121
  },
122
  }
123
  if json_format:
124
  payload["format"] = "json"
 
 
125
  body = json.dumps(
126
  payload
127
  ).encode("utf-8")
 
1
  from __future__ import annotations
2
 
3
  import json
4
+ import base64
5
  import urllib.error
6
  import urllib.request
7
+ from pathlib import Path
8
  from typing import Any, Callable
9
 
10
 
 
19
  model: str,
20
  timeout: float = 2.5,
21
  chat_max_tokens: int | None = None,
22
+ chat_response_length: str = "automatic",
23
  ) -> None:
24
  self.base_url = base_url.rstrip("/")
25
  self.model = model
26
  self.timeout = timeout
27
  self.chat_max_tokens = chat_max_tokens
28
+ self.chat_response_length = str(chat_response_length or "automatic").casefold()
29
 
30
  def _chat_system(self, system: str) -> str:
31
  """Return the application system prompt unchanged."""
32
  return system
33
 
34
+ def _num_predict(self, default: int, *, chat: bool = True, prompt: str = "", image_count: int = 0) -> int:
35
  """Qwen3's reasoning commonly needs more than a short-chat token budget."""
36
  if chat and self.chat_max_tokens is not None:
37
+ limit = max(64, int(self.chat_max_tokens))
38
+ choices = {"short": 256, "balanced": 512, "detailed": limit}
39
+ selected = choices.get(self.chat_response_length)
40
+ if selected is not None:
41
+ return min(limit, selected)
42
+ text = prompt.casefold()
43
+ if image_count:
44
+ budget = 384
45
+ elif any(word in text for word in ("compare", "explain", "research", "plan", "review", "why", "how")):
46
+ budget = 768
47
+ elif any(word in text for word in ("summarize", "details", "ideas", "examples")):
48
+ budget = 512
49
+ else:
50
+ budget = 256
51
+ return min(limit, max(128, budget))
52
  return 1024 if self.model.casefold().startswith("qwen3") else default
53
 
54
+ @staticmethod
55
+ def prepare_image(path: str | Path, *, maximum_side: int = 1536) -> str:
56
+ """Return a compact PNG attachment without changing the user's source image."""
57
+ source = Path(path).expanduser()
58
+ if not source.is_file():
59
+ raise OllamaError("The attached image no longer exists.")
60
+ try:
61
+ from PIL import Image, ImageOps
62
+ with Image.open(source) as opened:
63
+ image = ImageOps.exif_transpose(opened).convert("RGB")
64
+ image.thumbnail((maximum_side, maximum_side), Image.Resampling.LANCZOS)
65
+ from io import BytesIO
66
+ buffer = BytesIO()
67
+ image.save(buffer, format="PNG", optimize=True)
68
+ except (OSError, ValueError) as exc:
69
+ raise OllamaError(f"ADAM could not read {source.name} as an image.") from exc
70
+ return base64.b64encode(buffer.getvalue()).decode("ascii")
71
+
72
  def is_available(self, timeout: float = 0.35) -> bool:
73
  request = urllib.request.Request(f"{self.base_url}/api/tags", method="GET")
74
  try:
 
100
  raise OllamaError("Ollama returned an unsupported plan shape.")
101
  return parsed
102
 
103
+ def generate_text(
104
+ self, system: str, prompt: str, *, image_paths: list[str | Path] | None = None,
105
+ ) -> str:
106
+ response = self._generate(system, prompt, json_format=False, image_paths=image_paths).strip()
107
  if not response:
108
  raise OllamaError("Ollama returned an empty response.")
109
  return response
 
113
  system: str,
114
  prompt: str,
115
  on_chunk: Callable[[str], None],
116
+ image_paths: list[str | Path] | None = None,
117
  ) -> str:
118
+ images = [self.prepare_image(path) for path in (image_paths or [])]
119
  payload: dict[str, Any] = {
120
  "model": self.model,
121
  "system": self._chat_system(system),
122
  "prompt": prompt,
123
+ # Free GPU memory for image generation and other local workloads as
124
+ # soon as this one-shot text request finishes.
125
+ "keep_alive": 0,
126
  "stream": True,
127
+ "options": {"temperature": 0.35, "num_predict": self._num_predict(180, prompt=prompt, image_count=len(images))},
128
  }
129
+ if images:
130
+ payload["images"] = images
131
  request = urllib.request.Request(
132
  f"{self.base_url}/api/generate",
133
  data=json.dumps(payload).encode("utf-8"),
 
154
  raise OllamaError("Ollama returned an empty response.")
155
  return result
156
 
157
+ def _generate(self, system: str, prompt: str, *, json_format: bool, image_paths: list[str | Path] | None = None) -> str:
158
+ images = [self.prepare_image(path) for path in (image_paths or [])]
159
  payload: dict[str, Any] = {
160
  "model": self.model,
161
  "system": self._chat_system(system),
162
  "prompt": prompt,
163
+ "keep_alive": 0,
164
  "stream": False,
165
  "options": {
166
  "temperature": 0.1 if json_format else 0.35,
167
+ "num_predict": self._num_predict(300 if json_format else 180, chat=not json_format, prompt=prompt, image_count=len(images)),
168
  },
169
  }
170
  if json_format:
171
  payload["format"] = "json"
172
+ if images:
173
+ payload["images"] = images
174
  body = json.dumps(
175
  payload
176
  ).encode("utf-8")
adam/orion.py CHANGED
@@ -138,6 +138,91 @@ def review_training_plan(plan: Any) -> dict[str, Any]:
138
 
139
  for step in training_steps:
140
  args = step.arguments
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
141
  dataset_key = str(Path(str(args.get("dataset_dir", ""))).expanduser())
142
  images = dataset_image_count(args.get("dataset_dir")) or projected_counts.get(dataset_key, 0)
143
  epochs = max(1, int(args.get("epochs", 1) or 1))
 
138
 
139
  for step in training_steps:
140
  args = step.arguments
141
+ if step.tool_id == "wan_video_trainer":
142
+ from adam.video_lora import clips_in
143
+ count = len(clips_in(Path(str(args.get("dataset_dir", "")))))
144
+ epochs = int(args.get("epochs", 1))
145
+ findings.append({"level": "warning" if epochs > 100 or count < 20 else "ready", "message": (
146
+ f"Wan video: {count} clips, {epochs} epochs, frame buckets {args.get('target_frames', '25,49')}. "
147
+ "Each clip can contribute multiple frame buckets. Inspect motion and captions; video runtime requires a measured run."
148
+ )})
149
+ continue
150
+ if step.tool_id == "oasis_trainer":
151
+ from adam.oasis_dataset import inspect_oasis_dataset, oasis_pace
152
+
153
+ gap = max(1, int(args.get("frame_gap", 1) or 1))
154
+ pace = oasis_pace(args.get("dataset_dir", ""), frame_gap=gap)
155
+ recommended_gap = pace["recommended_frame_gap"]
156
+ capture_fps = pace["capture_fps"]
157
+ native_fps = pace["native_ai_fps"]
158
+ if isinstance(capture_fps, (int, float)) and isinstance(native_fps, (int, float)):
159
+ message = (
160
+ f"{step.title or 'Oasis'}: {float(capture_fps):g} FPS capture with "
161
+ f"prediction gap {gap} trains at a native pace of {float(native_fps):g} AI FPS."
162
+ )
163
+ if isinstance(recommended_gap, int) and recommended_gap != gap:
164
+ message += f" Dataset metadata recommends gap {recommended_gap} for responsive control."
165
+ findings.append({"level": "warning", "message": message})
166
+ else:
167
+ findings.append({"level": "ready", "message": message})
168
+ report = inspect_oasis_dataset(args.get("dataset_dir", ""), frame_gap=gap)
169
+ if report.ok and report.valid_transitions:
170
+ batch = max(1, int(args.get("batch_size", 1) or 1))
171
+ accumulation = max(1, int(args.get("gradient_accumulation", 1) or 1))
172
+ requested_chunk = max(0, int(args.get("chunk_size", 0) or 0))
173
+ transitions_per_epoch = min(report.valid_transitions, requested_chunk) if requested_chunk else report.valid_transitions
174
+ steps_per_epoch = math.ceil(transitions_per_epoch / batch / accumulation)
175
+ epochs = max(1, int(args.get("epochs", 1) or 1))
176
+ optimizer_steps = steps_per_epoch * epochs
177
+ total_steps += optimizer_steps
178
+ label = step.title or "Oasis"
179
+ findings.append({
180
+ "level": "ready",
181
+ "message": (
182
+ f"{label}: {report.valid_transitions:,} valid transitions; "
183
+ f"{transitions_per_epoch:,} used per epoch; about "
184
+ f"{steps_per_epoch:,} optimizer steps per epoch."
185
+ ),
186
+ })
187
+ if report.valid_transitions >= 7_500 and not requested_chunk:
188
+ findings.append({
189
+ "level": "warning",
190
+ "message": (
191
+ f"{label}: every epoch uses all {report.valid_transitions:,} transitions. "
192
+ "Use a balanced 5,000-transition chunk or explicitly confirm the longer run."
193
+ ),
194
+ })
195
+ if optimizer_steps >= 100_000:
196
+ findings.append({
197
+ "level": "warning",
198
+ "message": (
199
+ f"{label}: this plan schedules about {optimizer_steps:,} optimizer steps. "
200
+ "Run the short benchmark and inspect rollout previews before committing."
201
+ ),
202
+ })
203
+ idle_ratio = report.idle_rows / max(1, report.valid_rows)
204
+ if idle_ratio < 0.05:
205
+ findings.append({
206
+ "level": "warning",
207
+ "message": (
208
+ f"{label}: only {idle_ratio:.1%} of labelled frames are idle. "
209
+ "Record more no-input gameplay to improve stable pauses."
210
+ ),
211
+ })
212
+ rare_threshold = max(10, math.ceil(report.valid_rows * 0.01))
213
+ rare_controls = [
214
+ name for name, count in report.action_counts.items()
215
+ if 0 < count < rare_threshold
216
+ ]
217
+ if rare_controls and not bool(args.get("balance_actions", False)):
218
+ findings.append({
219
+ "level": "warning",
220
+ "message": (
221
+ "Rare controls are present (" + ", ".join(rare_controls[:5])
222
+ + "); turn on Balance rare actions or record more examples."
223
+ ),
224
+ })
225
+ continue
226
  dataset_key = str(Path(str(args.get("dataset_dir", ""))).expanduser())
227
  images = dataset_image_count(args.get("dataset_dir")) or projected_counts.get(dataset_key, 0)
228
  epochs = max(1, int(args.get("epochs", 1) or 1))
adam/planner.py CHANGED
@@ -7,9 +7,12 @@ from pathlib import Path
7
  from typing import Any
8
  from collections.abc import Callable
9
 
 
10
  from adam.assets import Asset, AssetRegistry
11
  from adam.commands import CommandValidationError, TrainingCommand
12
  from adam.config import ConfigManager
 
 
13
  from adam.models import ExecutionPlan, PlanStep
14
  from adam.ollama import OllamaClient, OllamaError
15
  from adam.registry import RegistryError, ToolRegistry
@@ -47,6 +50,7 @@ def _trainer_label(trainer: str) -> str:
47
  return {
48
  "ddpm": "DDPM",
49
  "flow": "Flow Matching",
 
50
  "lora": "LoRA",
51
  "oasis": "Oasis Action World Model",
52
  }.get(trainer, trainer.replace("_", " ").title())
@@ -99,10 +103,18 @@ class Planner:
99
  if not request:
100
  raise PlanningError("Tell ADAM what you want to accomplish.")
101
 
 
 
 
 
102
  if self.pending_request and self._looks_like_pending_details(request):
103
  return self._continue_pending_request(request)
104
 
105
  self.assets.discover(self.config)
 
 
 
 
106
  external = self._external_tool_plan(request)
107
  if external:
108
  self.last_mode = "Validated external tool"
@@ -130,7 +142,7 @@ class Planner:
130
  project_name="Conversation",
131
  )
132
 
133
- if self.config.get("provider") == "ollama":
134
  try:
135
  generated = self._ollama_plan(request)
136
  self.last_mode = "Ollama + registry validation"
@@ -164,6 +176,7 @@ class Planner:
164
  request: str,
165
  history: list[dict[str, str]] | None = None,
166
  stream_callback: Callable[[str], None] | None = None,
 
167
  ) -> str:
168
  """Answer conversationally without creating or executing a workflow."""
169
  request = request.strip()
@@ -174,6 +187,7 @@ class Planner:
174
  self.config.get("ollama_model"),
175
  timeout=45.0,
176
  chat_max_tokens=int(self.config.get("ollama_chat_max_tokens", 1024)),
 
177
  )
178
  if self.config.get("provider") != "ollama":
179
  raise PlanningError(
@@ -199,10 +213,17 @@ class Planner:
199
  "about AI datasets, captions, LoRA, DDPM, Flow Matching, model training, previews, "
200
  "and the workflows registered in ADAM. This is Chat Mode: you cannot run tools, "
201
  "change files, start jobs, or claim that work occurred. If the user asks you to "
202
- "perform an action, explain that they should switch to Trainer Mode. Never invent "
203
- "job results or capabilities. Registered read-only capability summary:\n"
 
204
  + json.dumps(capabilities, ensure_ascii=False)
205
  )
 
 
 
 
 
 
206
  recent = (history or [])[-10:]
207
  transcript = "\n".join(
208
  f"{'User' if item.get('role') == 'user' else 'ADAM'}: "
@@ -233,9 +254,9 @@ class Planner:
233
  )
234
  try:
235
  response = (
236
- client.generate_text_stream(system, prompt, stream_callback)
237
  if stream_callback
238
- else client.generate_text(system, prompt)
239
  )
240
  except OllamaError as exc:
241
  raise PlanningError(f"Ollama could not answer: {exc}") from exc
@@ -582,6 +603,7 @@ class Planner:
582
  ) or (
583
  "lora" if re.search(r"\blora\b", lowered)
584
  else "ddpm" if re.search(r"\bddpm\b", lowered)
 
585
  else "flow" if re.search(r"\bflow(?:\s+matching)?\b", lowered)
586
  else "oasis" if re.search(
587
  r"\b(oasis|action[- ]conditioned|playable\s+ai\s+games?|world\s+models?|gameplay[- ]frame|wasd|w/a/s/d)\b",
@@ -601,7 +623,7 @@ class Planner:
601
  model_query = ""
602
  resume_match = re.search(
603
  r"\b(?:fine[- ]?tune|retrain|continue|resume)\s+(?:the\s+)?(.+?)"
604
- r"(?:\s+model)?\s+(?:from|on|with)\s+(?:the\s+)?(?:ddpm|lora|oasis)\b",
605
  request,
606
  re.I,
607
  )
@@ -619,7 +641,7 @@ class Planner:
619
  model_query = _clean_subject(match.group(1)) if match else ""
620
  natural_resume = re.search(
621
  r"\b(?:fine[- ]?tune|retrain|continue|resume)\s+(?:the\s+)?(.+?)\s+"
622
- r"from\s+(?:my|our|the)\s+(?:ddpm|lora)\s+model\b",
623
  request,
624
  re.I,
625
  )
@@ -627,7 +649,7 @@ class Planner:
627
  model_query = _clean_subject(natural_resume.group(1))
628
  model_of_resume = re.search(
629
  r"\b(?:fine[- ]?tune|retrain|continue|resume)\s+(?:the\s+)?"
630
- r"(?:ddpm\s+|lora\s+)?model\s+of\s+(.+?)(?:\s+for\b|,|$)",
631
  request,
632
  re.I,
633
  )
@@ -636,7 +658,7 @@ class Planner:
636
  # Natural phrasing such as "fine-tune Hatsune Miku from our DDPM model"
637
  # should search for "Hatsune Miku", not the whole explanatory clause.
638
  model_query = re.sub(
639
- r"\s+from\s+(?:my|our|the)?\s*(?:ddpm|lora)\s+model\s*$",
640
  "",
641
  model_query,
642
  flags=re.I,
@@ -710,7 +732,18 @@ class Planner:
710
  steps=[],
711
  project_name="Resume training",
712
  )
713
- resumed_model_name = _friendly_model_name(model)
 
 
 
 
 
 
 
 
 
 
 
714
  command = TrainingCommand.from_dict(
715
  {
716
  "action": "resume_training",
@@ -718,10 +751,9 @@ class Planner:
718
  "dataset": dataset.path,
719
  "model_name": resumed_model_name,
720
  "epochs": epochs,
721
- "output": (
722
- str(self._training_output(trainer, f"{resumed_model_name} Fine Tune") or model.path)
723
- if trainer == "flow" else model.path
724
- ),
725
  # The DDPM adapter can safely branch from a complete pipeline when
726
  # its exact Accelerate checkpoint has been cleaned up.
727
  "resume_from": model.checkpoint or model.path,
@@ -779,6 +811,141 @@ class Planner:
779
  raise PlanningError(str(exc)) from exc
780
  return self._plan_training_command(request, command)
781
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
782
  @staticmethod
783
  def _fine_tune_payload(request: str) -> dict[str, Any]:
784
  match = re.search(r"\[ADAM_FINE_TUNE:(\{.*\})\]\s*$", request, re.S)
@@ -792,6 +959,20 @@ class Planner:
792
  raise PlanningError("Fine-tune settings must be an object.")
793
  return payload
794
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
795
  def _dataset_for_model(self, model: Asset) -> Asset | None:
796
  if model.dataset_id:
797
  linked = next(
@@ -827,14 +1008,15 @@ class Planner:
827
  if dataset_dir.exists():
828
  dataset_dir = dataset_dir.with_name(f"{dataset_dir.name} {datetime.now().strftime('%Y%m%d_%H%M%S')}")
829
  image_count = max(10, min(int(payload.get("image_count", 60)), 5000))
830
- model_name = _friendly_model_name(model)
 
 
 
 
831
  arguments: dict[str, Any] = {
832
  "dataset_dir": str(dataset_dir), "model_name": model_name,
833
  "epochs": epochs,
834
- "output_dir": (
835
- str(self._training_output(trainer, f"{model_name} Fine Tune") or model.path)
836
- if trainer == "flow" else model.path
837
- ),
838
  "resume_from": model.checkpoint or model.path, **training_options,
839
  }
840
  if trainer == "lora":
@@ -969,13 +1151,44 @@ class Planner:
969
  return existing
970
  if matches:
971
  return matches
972
- return self.assets.find("model", raw_query, trainer=trainer)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
973
 
974
  def _plan_training_command(
975
  self,
976
  request: str,
977
  command: TrainingCommand,
978
  ) -> ExecutionPlan:
 
 
 
 
 
 
 
 
979
  tool_id = f"{command.trainer}_trainer"
980
  spec = self.registry.get(tool_id)
981
  capability = (
@@ -1028,6 +1241,23 @@ class Planner:
1028
  ) from exc
1029
  if command.resume_from and not Path(command.resume_from).exists():
1030
  raise PlanningError("The validated resume checkpoint does not exist.")
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1031
  arguments: dict[str, Any] = {
1032
  "dataset_dir": dataset_argument,
1033
  "model_name": command.model_name,
@@ -1413,6 +1643,9 @@ class Planner:
1413
  self.config.get("ollama_model"),
1414
  timeout=30.0,
1415
  chat_max_tokens=int(self.config.get("ollama_chat_max_tokens", 1024)),
 
 
 
1416
  )
1417
  if re.search(r"\bollama\b.*\b(working|online|reachable|running)\b", request, re.I):
1418
  return (
@@ -1893,7 +2126,7 @@ class Planner:
1893
 
1894
  catalog = self.registry.safe_llm_catalog()
1895
  system = (
1896
- "You are ADAM's planning component. You only plan; you never execute. "
1897
  "Return strict JSON with summary, project_name, requires_confirmation, "
1898
  "confirmation_reason, and steps. Each step has tool_id, title, "
1899
  "description, and arguments. Use only listed tool IDs and only their "
 
7
  from typing import Any
8
  from collections.abc import Callable
9
 
10
+ from adam.auto_training import profile_from_request, resolve_auto_training
11
  from adam.assets import Asset, AssetRegistry
12
  from adam.commands import CommandValidationError, TrainingCommand
13
  from adam.config import ConfigManager
14
+ from adam.dataset_lab import scan_dataset
15
+ from adam.model_profiles import ModelProfileRegistry
16
  from adam.models import ExecutionPlan, PlanStep
17
  from adam.ollama import OllamaClient, OllamaError
18
  from adam.registry import RegistryError, ToolRegistry
 
50
  return {
51
  "ddpm": "DDPM",
52
  "flow": "Flow Matching",
53
+ "inrflow": "INRFlow",
54
  "lora": "LoRA",
55
  "oasis": "Oasis Action World Model",
56
  }.get(trainer, trainer.replace("_", " ").title())
 
103
  if not request:
104
  raise PlanningError("Tell ADAM what you want to accomplish.")
105
 
106
+ if re.search(r"\b(?:wan(?:\s*2[.]1)?|video\s+lora|lora\s+video)\b", request, re.I):
107
+ self.last_mode = "Video LoRA workspace"
108
+ return ExecutionPlan(request=request, summary="Open Video LoRA in the sidebar to select captioned clips, review a Wan training pipeline, continue weights, or generate an MP4. Wan settings and checkpoints are separate from SDXL LoRA.", steps=[], project_name="Video LoRA")
109
+
110
  if self.pending_request and self._looks_like_pending_details(request):
111
  return self._continue_pending_request(request)
112
 
113
  self.assets.discover(self.config)
114
+ auto_training = self._auto_training_plan(request)
115
+ if auto_training:
116
+ self.last_mode = "Intent-based AUTO training"
117
+ return auto_training
118
  external = self._external_tool_plan(request)
119
  if external:
120
  self.last_mode = "Validated external tool"
 
142
  project_name="Conversation",
143
  )
144
 
145
+ if self.config.get("provider") == "ollama" and self.config.get("ollama_proposed_actions", True):
146
  try:
147
  generated = self._ollama_plan(request)
148
  self.last_mode = "Ollama + registry validation"
 
176
  request: str,
177
  history: list[dict[str, str]] | None = None,
178
  stream_callback: Callable[[str], None] | None = None,
179
+ image_paths: list[str | Path] | None = None,
180
  ) -> str:
181
  """Answer conversationally without creating or executing a workflow."""
182
  request = request.strip()
 
187
  self.config.get("ollama_model"),
188
  timeout=45.0,
189
  chat_max_tokens=int(self.config.get("ollama_chat_max_tokens", 1024)),
190
+ chat_response_length=str(self.config.get("ollama_chat_response_length", "automatic")),
191
  )
192
  if self.config.get("provider") != "ollama":
193
  raise PlanningError(
 
213
  "about AI datasets, captions, LoRA, DDPM, Flow Matching, model training, previews, "
214
  "and the workflows registered in ADAM. This is Chat Mode: you cannot run tools, "
215
  "change files, start jobs, or claim that work occurred. If the user asks you to "
216
+ "perform an action, explain that ADAM must create a validated plan before anything "
217
+ "can happen. Never invent job results, capability status, or a running/completed job. "
218
+ "Registered read-only capability summary:\n"
219
  + json.dumps(capabilities, ensure_ascii=False)
220
  )
221
+ if image_paths:
222
+ system += (
223
+ " The user attached image(s). Describe only visible evidence, distinguish "
224
+ "uncertainty from facts, and offer an editable caption when useful. Do not "
225
+ "claim the image was added to a dataset or used for training."
226
+ )
227
  recent = (history or [])[-10:]
228
  transcript = "\n".join(
229
  f"{'User' if item.get('role') == 'user' else 'ADAM'}: "
 
254
  )
255
  try:
256
  response = (
257
+ client.generate_text_stream(system, prompt, stream_callback, image_paths=image_paths)
258
  if stream_callback
259
+ else client.generate_text(system, prompt, image_paths=image_paths)
260
  )
261
  except OllamaError as exc:
262
  raise PlanningError(f"Ollama could not answer: {exc}") from exc
 
603
  ) or (
604
  "lora" if re.search(r"\blora\b", lowered)
605
  else "ddpm" if re.search(r"\bddpm\b", lowered)
606
+ else "inrflow" if re.search(r"\binr\s*flow\b", lowered)
607
  else "flow" if re.search(r"\bflow(?:\s+matching)?\b", lowered)
608
  else "oasis" if re.search(
609
  r"\b(oasis|action[- ]conditioned|playable\s+ai\s+games?|world\s+models?|gameplay[- ]frame|wasd|w/a/s/d)\b",
 
623
  model_query = ""
624
  resume_match = re.search(
625
  r"\b(?:fine[- ]?tune|retrain|continue|resume)\s+(?:the\s+)?(.+?)"
626
+ r"(?:\s+model)?\s+(?:from|on|with)\s+(?:the\s+)?(?:ddpm|lora|oasis|inr\s*flow|flow)\b",
627
  request,
628
  re.I,
629
  )
 
641
  model_query = _clean_subject(match.group(1)) if match else ""
642
  natural_resume = re.search(
643
  r"\b(?:fine[- ]?tune|retrain|continue|resume)\s+(?:the\s+)?(.+?)\s+"
644
+ r"from\s+(?:my|our|the)\s+(?:ddpm|lora|inr\s*flow|flow)\s+model\b",
645
  request,
646
  re.I,
647
  )
 
649
  model_query = _clean_subject(natural_resume.group(1))
650
  model_of_resume = re.search(
651
  r"\b(?:fine[- ]?tune|retrain|continue|resume)\s+(?:the\s+)?"
652
+ r"(?:(?:ddpm|lora|inr\s*flow|flow)\s+)?model\s+of\s+(.+?)(?:\s+for\b|,|$)",
653
  request,
654
  re.I,
655
  )
 
658
  # Natural phrasing such as "fine-tune Hatsune Miku from our DDPM model"
659
  # should search for "Hatsune Miku", not the whole explanatory clause.
660
  model_query = re.sub(
661
+ r"\s+from\s+(?:my|our|the)?\s*(?:ddpm|lora|inr\s*flow|flow)\s+model\s*$",
662
  "",
663
  model_query,
664
  flags=re.I,
 
732
  steps=[],
733
  project_name="Resume training",
734
  )
735
+ requested_output_name = self._fine_tune_output_name(request, fine_tune_payload)
736
+ resumed_model_name = _clean_subject(requested_output_name) if requested_output_name else _friendly_model_name(model)
737
+ continuation_output = self._training_output(
738
+ trainer, f"{resumed_model_name} Fine Tune"
739
+ )
740
+ if not continuation_output:
741
+ return ExecutionPlan(
742
+ request=request,
743
+ summary=f"The {_trainer_label(trainer)} trainer folder is not connected.",
744
+ steps=[],
745
+ project_name="Resume training",
746
+ )
747
  command = TrainingCommand.from_dict(
748
  {
749
  "action": "resume_training",
 
751
  "dataset": dataset.path,
752
  "model_name": resumed_model_name,
753
  "epochs": epochs,
754
+ # A continuation is always a new model branch. Never send a
755
+ # resumed run back to the selected source model's folder.
756
+ "output": str(continuation_output),
 
757
  # The DDPM adapter can safely branch from a complete pipeline when
758
  # its exact Accelerate checkpoint has been cleaned up.
759
  "resume_from": model.checkpoint or model.path,
 
811
  raise PlanningError(str(exc)) from exc
812
  return self._plan_training_command(request, command)
813
 
814
+ def _auto_training_plan(self, request: str) -> ExecutionPlan | None:
815
+ """Plan a short, clear training request through a named AUTO policy.
816
+
817
+ This is intentionally narrow: a trainer and a subject must both be
818
+ explicit. Ambiguous requests continue through the existing planner,
819
+ which can ask a focused follow-up instead of guessing an architecture.
820
+ """
821
+ lowered = request.casefold()
822
+ if not re.search(r"\b(train|make|create|build|test)\b", lowered):
823
+ return None
824
+ # Fully specified legacy requests retain their established planner path.
825
+ # AUTO is for omitted decisions, not a replacement for explicit control.
826
+ if re.search(r"\b\d{1,5}\s*epochs?\b", request, re.I):
827
+ return None
828
+ match = re.search(
829
+ r"^\s*(?:(?:can|could|will|would)\s+you\s+(?:please\s+)?)?(?:quickly\s+)?"
830
+ r"(?:train|make|create|build|test)\s+(?:me\s+)?(?:a|an|the)?\s*"
831
+ r"(?:(?:quick|high\s+quality|best\s+quality|really\s+good|quality)\s+)?"
832
+ r"(ddpm|lora|inr\s*flow|flow(?:\s+matching)?)(?:\s+model)?\s+"
833
+ r"(?:on|of|for|using|with)\s+(?:an?\s+)?(?:dataset\s+(?:of|for)\s+)?(.+?)\s*$",
834
+ request,
835
+ re.I,
836
+ )
837
+ if not match:
838
+ return None
839
+ trainer_token, raw_subject = match.groups()
840
+ trainer = (
841
+ "inrflow" if trainer_token.casefold().replace(" ", "") == "inrflow"
842
+ else "flow" if trainer_token.casefold().startswith("flow")
843
+ else trainer_token.casefold()
844
+ )
845
+ # Remove only trailing presentation words/settings; the subject itself
846
+ # remains ordinary natural language and never becomes a hidden command.
847
+ subject = re.sub(
848
+ r"\s+(?:images?|pictures?|screenshots?)(?:\s+(?:for|with|overnight|quickly)\b.*)?$|"
849
+ r"\s+for\s+\d{1,5}\s+epochs?\b.*$",
850
+ "",
851
+ raw_subject,
852
+ flags=re.I,
853
+ )
854
+ subject = _clean_subject(subject)
855
+ if subject == "new subject":
856
+ return None
857
+ profile = ModelProfileRegistry(self.registry.model_plugins).get(trainer)
858
+ if profile is None:
859
+ return None
860
+
861
+ policy = profile_from_request(request)
862
+ options = self._training_options_from_request(request)
863
+ epoch_match = re.search(r"\b(\d{1,5})\s*epochs?\b", request, re.I)
864
+ count_match = re.search(r"\b(\d{1,6})\s+(?:images?|pictures?)\b", request, re.I)
865
+ existing_dataset = self._asset_dataset(subject)
866
+ expected_items = (
867
+ scan_dataset(existing_dataset.path, limit=1).image_count
868
+ if existing_dataset else 0
869
+ )
870
+ auto = resolve_auto_training(
871
+ profile,
872
+ trainer=trainer,
873
+ policy=policy,
874
+ dataset_items=expected_items or (int(count_match.group(1)) if count_match else 400),
875
+ )
876
+ epochs = int(epoch_match.group(1)) if epoch_match else auto.epochs
877
+ image_count = int(count_match.group(1)) if count_match else auto.dataset_target
878
+ # Explicit structured/manual settings override AUTO; AUTO supplies every
879
+ # remaining supported setting, making the final plan reproducible.
880
+ training_options = {**auto.settings, **options}
881
+ model_name = self._model_name_from_request(request) or subject
882
+
883
+ if existing_dataset:
884
+ output = self._training_output(trainer, model_name)
885
+ if not output:
886
+ return None
887
+ try:
888
+ command = TrainingCommand.from_dict({
889
+ "action": "train", "trainer": trainer, "dataset": existing_dataset.path,
890
+ "model_name": model_name, "epochs": epochs, "output": str(output),
891
+ "base_model": (
892
+ str(training_options.get("base_model") or self._lora_base_model())
893
+ if trainer == "lora" else ""
894
+ ),
895
+ "training_options": training_options,
896
+ })
897
+ except CommandValidationError as exc:
898
+ raise PlanningError(str(exc)) from exc
899
+ plan = self._plan_training_command(request, command)
900
+ plan.summary += f" {auto.summary}"
901
+ return plan
902
+
903
+ collector_root = self._configured_tool_folder("dataset_collector")
904
+ output = self._training_output(trainer, model_name)
905
+ if not collector_root or not output:
906
+ return None
907
+ project = _project_name(subject, "Dataset")
908
+ dataset_dir = (Path(collector_root) / "Datasets" / project).resolve()
909
+ if dataset_dir.exists():
910
+ dataset_dir = dataset_dir.with_name(
911
+ f"{dataset_dir.name} {datetime.now().strftime('%Y%m%d_%H%M%S')}"
912
+ )
913
+ arguments: dict[str, Any] = {
914
+ "dataset_dir": str(dataset_dir), "model_name": model_name,
915
+ "epochs": epochs, "output_dir": str(output), **training_options,
916
+ }
917
+ if trainer == "lora":
918
+ base_model = str(training_options.get("base_model") or self._lora_base_model())
919
+ if not base_model or not Path(base_model).is_file():
920
+ return ExecutionPlan(
921
+ request=request,
922
+ summary="Choose a valid SDXL base model in the LoRA settings before AUTO training.",
923
+ steps=[], project_name="LoRA training",
924
+ )
925
+ arguments["base_model"] = base_model
926
+ arguments["trigger_word"] = str(training_options.get("trigger_word") or model_name)
927
+ return ExecutionPlan(
928
+ request=request,
929
+ summary=(
930
+ f"AUTO plan: collect up to {image_count:,} images for {subject}, then train "
931
+ f"{model_name} with {_trainer_label(trainer)} for {epochs:,} epochs. {auto.summary}"
932
+ ),
933
+ steps=[
934
+ PlanStep("dataset_collector", "Collect dataset", "Collect a reviewable dataset for the requested subject.", {
935
+ "subject": subject, "image_count": max(10, min(image_count, 100_000)),
936
+ "collection_mode": _collection_mode(request), "project_name": project,
937
+ "output_dir": str(dataset_dir),
938
+ }),
939
+ PlanStep(f"{trainer}_trainer", f"Train {_trainer_label(trainer)} model", "Train using resolved AUTO settings.", arguments),
940
+ ],
941
+ requires_confirmation=True,
942
+ confirmation_reason=(
943
+ "This plan downloads a dataset and starts real GPU training. "
944
+ "The resolved AUTO settings are included in the training step."
945
+ ),
946
+ project_name=model_name[:64],
947
+ )
948
+
949
  @staticmethod
950
  def _fine_tune_payload(request: str) -> dict[str, Any]:
951
  match = re.search(r"\[ADAM_FINE_TUNE:(\{.*\})\]\s*$", request, re.S)
 
959
  raise PlanningError("Fine-tune settings must be an object.")
960
  return payload
961
 
962
+ @staticmethod
963
+ def _fine_tune_output_name(request: str, payload: dict[str, Any]) -> str:
964
+ """Return an explicit result name without confusing it with the source model."""
965
+ requested = str(payload.get("output_model_name", "")).strip()
966
+ if requested:
967
+ return requested
968
+ match = re.search(
969
+ r"\b(?:name|call)\s+(?:the\s+)?(?:fine[- ]?tuned\s+)?"
970
+ r"(?:model\s+)?(?:as|to)\s+[\"\u201c]?([^\"\u201d.,]+)",
971
+ request,
972
+ re.I,
973
+ )
974
+ return _clean_subject(match.group(1)) if match else ""
975
+
976
  def _dataset_for_model(self, model: Asset) -> Asset | None:
977
  if model.dataset_id:
978
  linked = next(
 
1008
  if dataset_dir.exists():
1009
  dataset_dir = dataset_dir.with_name(f"{dataset_dir.name} {datetime.now().strftime('%Y%m%d_%H%M%S')}")
1010
  image_count = max(10, min(int(payload.get("image_count", 60)), 5000))
1011
+ requested_output_name = self._fine_tune_output_name(request, payload)
1012
+ model_name = _clean_subject(requested_output_name) if requested_output_name else _friendly_model_name(model)
1013
+ continuation_output = self._training_output(trainer, f"{model_name} Fine Tune")
1014
+ if not continuation_output:
1015
+ return ExecutionPlan(request=request, summary=f"The {_trainer_label(trainer)} trainer folder is not connected.", steps=[], project_name="Resume training")
1016
  arguments: dict[str, Any] = {
1017
  "dataset_dir": str(dataset_dir), "model_name": model_name,
1018
  "epochs": epochs,
1019
+ "output_dir": str(continuation_output),
 
 
 
1020
  "resume_from": model.checkpoint or model.path, **training_options,
1021
  }
1022
  if trainer == "lora":
 
1151
  return existing
1152
  if matches:
1153
  return matches
1154
+ matches = self.assets.find("model", raw_query, trainer=trainer)
1155
+ # Asset discovery and an in-memory registration can legitimately refer
1156
+ # to the same saved model. One physical path is one continuation
1157
+ # candidate, not an ambiguity the user has to resolve.
1158
+ unique: list[Asset] = []
1159
+ seen_paths: set[str] = set()
1160
+ # Prefer a model folder over a checkpoint file inside that same folder.
1161
+ # Discovery can register both representations of one saved run.
1162
+ ordered = sorted(matches, key=lambda item: (not Path(item.path).is_dir(), len(str(item.path))))
1163
+ for asset in ordered:
1164
+ try:
1165
+ resolved = Path(asset.path).expanduser().resolve()
1166
+ key = str(resolved).casefold()
1167
+ except OSError:
1168
+ resolved = Path(asset.path)
1169
+ key = str(asset.path).casefold()
1170
+ nested_in_known_model = any(
1171
+ key.startswith(parent + "\\") or key.startswith(parent + "/")
1172
+ for parent in seen_paths
1173
+ )
1174
+ if key not in seen_paths and not nested_in_known_model:
1175
+ seen_paths.add(key)
1176
+ unique.append(asset)
1177
+ return unique
1178
 
1179
  def _plan_training_command(
1180
  self,
1181
  request: str,
1182
  command: TrainingCommand,
1183
  ) -> ExecutionPlan:
1184
+ if command.trainer == "wan_video":
1185
+ from adam.video_lora import training_plan
1186
+ try:
1187
+ return training_plan(self.root, command.dataset, command.model_name, command.epochs,
1188
+ command.training_options or {}, command.resume_from,
1189
+ "" if command.output == "default output" else command.output)
1190
+ except ValueError as exc:
1191
+ raise PlanningError(str(exc)) from exc
1192
  tool_id = f"{command.trainer}_trainer"
1193
  spec = self.registry.get(tool_id)
1194
  capability = (
 
1241
  ) from exc
1242
  if command.resume_from and not Path(command.resume_from).exists():
1243
  raise PlanningError("The validated resume checkpoint does not exist.")
1244
+ if command.resume_from:
1245
+ resume_path = Path(command.resume_from).expanduser().resolve()
1246
+ # A checkpoint may live below its model folder, so protect both
1247
+ # ancestors and descendants rather than checking simple equality.
1248
+ try:
1249
+ output_path.relative_to(resume_path)
1250
+ overlaps_resume = True
1251
+ except ValueError:
1252
+ try:
1253
+ resume_path.relative_to(output_path)
1254
+ overlaps_resume = True
1255
+ except ValueError:
1256
+ overlaps_resume = False
1257
+ if overlaps_resume:
1258
+ raise PlanningError(
1259
+ "Continuation output must be a new folder, separate from the source model."
1260
+ )
1261
  arguments: dict[str, Any] = {
1262
  "dataset_dir": dataset_argument,
1263
  "model_name": command.model_name,
 
1643
  self.config.get("ollama_model"),
1644
  timeout=30.0,
1645
  chat_max_tokens=int(self.config.get("ollama_chat_max_tokens", 1024)),
1646
+ chat_response_length=str(
1647
+ self.config.get("ollama_chat_response_length", "automatic")
1648
+ ),
1649
  )
1650
  if re.search(r"\bollama\b.*\b(working|online|reachable|running)\b", request, re.I):
1651
  return (
 
2126
 
2127
  catalog = self.registry.safe_llm_catalog()
2128
  system = (
2129
+ "You are ADAM's proposed-action component. You only propose plans; you never execute. "
2130
  "Return strict JSON with summary, project_name, requires_confirmation, "
2131
  "confirmation_reason, and steps. Each step has tool_id, title, "
2132
  "description, and arguments. Use only listed tool IDs and only their "
adam/progressive_training.py ADDED
@@ -0,0 +1,103 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Resolution-curriculum helpers shared by the DDPM and Flow adapters.
2
+
3
+ The adapters deliberately use conservative batch caps instead of trying an
4
+ out-of-memory probe in a real training job. A user can still turn the policy
5
+ off and enter every batch setting manually.
6
+ """
7
+
8
+ from __future__ import annotations
9
+
10
+ from dataclasses import dataclass
11
+ from typing import Any
12
+
13
+ from adam.executor import ToolExecutionError
14
+
15
+
16
+ @dataclass(frozen=True, slots=True)
17
+ class ResolutionStage:
18
+ resolution: int
19
+ epochs: int
20
+
21
+
22
+ _BATCH_CAPS = {
23
+ "ddpm": {64: 16, 128: 12, 256: 4, 384: 2, 512: 1},
24
+ "flow": {64: 12, 128: 8, 256: 4, 384: 2, 512: 1},
25
+ }
26
+
27
+
28
+ def parse_stages(value: Any, *, trainer: str, total_epochs: int) -> list[ResolutionStage]:
29
+ """Validate a JSON-friendly progressive-resolution schedule.
30
+
31
+ Stages are intentionally a list of small dictionaries so plans remain easy
32
+ to inspect and edit in saved job JSON.
33
+ """
34
+ if not isinstance(value, list) or len(value) < 2:
35
+ raise ToolExecutionError("Progressive training needs at least two resolution stages.")
36
+ multiple = 16 if trainer == "flow" else 8
37
+ stages: list[ResolutionStage] = []
38
+ previous = 0
39
+ for raw in value:
40
+ if not isinstance(raw, dict):
41
+ raise ToolExecutionError("Each progressive stage must include resolution and epochs.")
42
+ try:
43
+ resolution = int(raw.get("resolution", 0))
44
+ epochs = int(raw.get("epochs", 0))
45
+ except (TypeError, ValueError) as exc:
46
+ raise ToolExecutionError("Progressive stage resolution and epochs must be whole numbers.") from exc
47
+ if not 64 <= resolution <= 512 or resolution % multiple:
48
+ raise ToolExecutionError(
49
+ f"{trainer.upper()} progressive resolutions must be 64–512 and divisible by {multiple}."
50
+ )
51
+ if resolution <= previous:
52
+ raise ToolExecutionError("Progressive stages must increase from lower to higher resolution.")
53
+ if epochs < 1:
54
+ raise ToolExecutionError("Each progressive stage needs at least one epoch.")
55
+ stages.append(ResolutionStage(resolution, epochs))
56
+ previous = resolution
57
+ if sum(stage.epochs for stage in stages) != int(total_epochs):
58
+ raise ToolExecutionError(
59
+ f"Progressive stage epochs total {sum(stage.epochs for stage in stages):,}, "
60
+ f"but training length is {int(total_epochs):,}."
61
+ )
62
+ return stages
63
+
64
+
65
+ def suggested_stages(final_resolution: int, total_epochs: int) -> list[ResolutionStage]:
66
+ """Return an editable low-to-high schedule that always matches the budget."""
67
+ resolutions = [size for size in (64, 128, 256, 384, 512) if size <= int(final_resolution)]
68
+ if len(resolutions) < 2:
69
+ return [ResolutionStage(int(final_resolution), int(total_epochs))]
70
+ if int(total_epochs) < len(resolutions):
71
+ resolutions = resolutions[-int(total_epochs):]
72
+ return [ResolutionStage(resolution, 1) for resolution in resolutions]
73
+ # Front-load inexpensive structure learning while reserving final-resolution
74
+ # refinement. Normalizing lets the same policy work for any epoch budget.
75
+ weights = [0.60, 0.20, 0.10, 0.06, 0.04][-len(resolutions):]
76
+ allocation = [max(1, round(total_epochs * weight / sum(weights))) for weight in weights]
77
+ difference = int(total_epochs) - sum(allocation)
78
+ allocation[0] += difference
79
+ return [ResolutionStage(resolution, epochs) for resolution, epochs in zip(resolutions, allocation)]
80
+
81
+
82
+ def stage_batch_settings(
83
+ *, trainer: str, stage_resolution: int, final_resolution: int,
84
+ final_batch_size: int, base_accumulation: int, auto_batch: bool,
85
+ ) -> tuple[int, int]:
86
+ """Choose a conservative physical batch and matching accumulation count."""
87
+ requested_batch = max(1, min(64, int(final_batch_size)))
88
+ requested_accumulation = max(1, min(64, int(base_accumulation)))
89
+ if not auto_batch:
90
+ return requested_batch, requested_accumulation
91
+ caps = _BATCH_CAPS[trainer]
92
+ cap = caps[min(caps, key=lambda size: abs(size - int(stage_resolution)))]
93
+ # Scale from the user's final-stage batch according to image area, then
94
+ # enforce a trainer-specific safe cap. This never raises the 512px batch.
95
+ scaled = round(requested_batch * (int(final_resolution) / int(stage_resolution)) ** 2)
96
+ batch_size = max(1, min(cap, 64, scaled))
97
+ effective_batch = requested_batch * requested_accumulation
98
+ accumulation = max(1, min(64, -(-effective_batch // batch_size)))
99
+ return batch_size, accumulation
100
+
101
+
102
+ def stage_summary(stages: list[ResolutionStage]) -> str:
103
+ return ", ".join(f"{stage.resolution}px × {stage.epochs}" for stage in stages)
adam/recommendations.py CHANGED
@@ -74,6 +74,7 @@ def recommend_for_profile(
74
  profile: ModelProfile,
75
  *,
76
  dataset_items: int,
 
77
  resolution: int | str | None = None,
78
  snapshot: SystemSnapshot | None = None,
79
  base_model_gb: float = 0.0,
@@ -85,6 +86,7 @@ def recommend_for_profile(
85
  else:
86
  resolution = int(raw_resolution)
87
  reasons: list[str] = []
 
88
  vram_total = snapshot.vram_total_gb if snapshot and snapshot.vram_total_gb else None
89
  available_vram = (
90
  max(0.0, snapshot.vram_total_gb - snapshot.vram_used_gb)
@@ -92,12 +94,18 @@ def recommend_for_profile(
92
  else vram_total
93
  )
94
  architecture = profile.architecture.casefold()
95
- target_exposures = 80_000 if profile.id == "lora" else 180_000 if "diffusion" in architecture else 120_000
96
- max_epochs = 220 if profile.id == "lora" else 600 if "diffusion" in architecture else 300
97
- epochs = max(10 if profile.id == "lora" else 25, min(max_epochs, round(target_exposures / images)))
98
- reasons.append(
99
- f"Epochs target roughly {target_exposures:,} image exposures, then clamp to the profile's safe range."
100
- )
 
 
 
 
 
 
101
 
102
  batch_defaults = {
103
  64: 16,
@@ -110,6 +118,8 @@ def recommend_for_profile(
110
  }
111
  if profile.id == "flow":
112
  batch_defaults.update({64: 12, 128: 8, 256: 4})
 
 
113
  if profile.id == "oasis":
114
  batch_defaults.update({128: 4, 256: 2, 384: 1, 512: 1})
115
  if profile.id == "lora":
@@ -136,7 +146,7 @@ def recommend_for_profile(
136
  settings["learning_rate"] = _clamp_to_schema(
137
  profile,
138
  "learning_rate",
139
- 0.00002 if profile.id == "oasis" else 0.0001 if profile.id in {"ddpm", "lora"} else 0.0002,
140
  )
141
  workers = max(1, min(8, (os.cpu_count() or 4) // 2))
142
  for key in ("dataloader_num_workers", "workers"):
@@ -145,22 +155,103 @@ def recommend_for_profile(
145
  for key, value in {
146
  "gradient_accumulation_steps": 1,
147
  "gradient_accumulation": 1,
148
- "mixed_precision": "fp32" if profile.id == "oasis" else "fp16",
149
  "save_every": max(5, min(25, max(1, epochs // 10))),
150
  "preview_every": max(5, min(50, max(1, epochs // 10))),
151
  "training_intensity": 100,
152
  "gradient_checkpointing": resolution >= 384 or (available_vram is not None and available_vram < 8),
153
  "rank": 16,
154
  "alpha": 16,
155
- "frame_gap": 3,
156
  "sequence_context": 1,
157
  "preview_steps": 1 if profile.id == "oasis" else 50 if profile.id == "ddpm" else 10,
158
  }.items():
159
  if key in profile.training:
160
  settings[key] = _clamp_to_schema(profile, key, value)
161
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
162
  estimated = estimate_vram_gb(profile, resolution, int(settings.get("batch_size", batch_size)), base_model_gb)
163
- warnings: list[str] = []
164
  if available_vram is not None and estimated > available_vram * 0.9:
165
  warnings.append(
166
  f"Estimated VRAM need is about {estimated:.1f} GB, above the conservative {available_vram * 0.9:.1f} GB working limit."
@@ -178,11 +269,18 @@ def recommend_for_profile(
178
  memory_note = (
179
  f" using about {available_vram:.1f} GB available VRAM" if available_vram is not None else " without detected VRAM"
180
  )
181
- summary = (
182
- f"Recommended {epochs:,} epochs for {images:,} item(s), "
183
- f"batch {settings.get('batch_size', batch_size)} at {resolution}px{memory_note}. "
184
- "Treat this as a starting recipe, not a guarantee."
185
- )
 
 
 
 
 
 
 
186
  return SettingsRecommendation(
187
  profile_id=profile.id,
188
  epochs=epochs,
 
74
  profile: ModelProfile,
75
  *,
76
  dataset_items: int,
77
+ dataset_path: str = "",
78
  resolution: int | str | None = None,
79
  snapshot: SystemSnapshot | None = None,
80
  base_model_gb: float = 0.0,
 
86
  else:
87
  resolution = int(raw_resolution)
88
  reasons: list[str] = []
89
+ warnings: list[str] = []
90
  vram_total = snapshot.vram_total_gb if snapshot and snapshot.vram_total_gb else None
91
  available_vram = (
92
  max(0.0, snapshot.vram_total_gb - snapshot.vram_used_gb)
 
94
  else vram_total
95
  )
96
  architecture = profile.architecture.casefold()
97
+ if profile.id == "oasis":
98
+ # Oasis learns labelled transitions rather than independent images. A real
99
+ # dataset inspection below replaces this fallback whenever one is selected.
100
+ epochs = 20
101
+ reasons.append("Oasis starts from a transition-step budget, not an image-exposure target.")
102
+ else:
103
+ target_exposures = 80_000 if profile.id == "lora" else 180_000 if "diffusion" in architecture else 120_000
104
+ max_epochs = 220 if profile.id == "lora" else 600 if "diffusion" in architecture else 300
105
+ epochs = max(10 if profile.id == "lora" else 25, min(max_epochs, round(target_exposures / images)))
106
+ reasons.append(
107
+ f"Epochs target roughly {target_exposures:,} image exposures, then clamp to the profile's safe range."
108
+ )
109
 
110
  batch_defaults = {
111
  64: 16,
 
118
  }
119
  if profile.id == "flow":
120
  batch_defaults.update({64: 12, 128: 8, 256: 4})
121
+ if profile.id == "inrflow":
122
+ batch_defaults.update({64: 4, 128: 2, 256: 1})
123
  if profile.id == "oasis":
124
  batch_defaults.update({128: 4, 256: 2, 384: 1, 512: 1})
125
  if profile.id == "lora":
 
146
  settings["learning_rate"] = _clamp_to_schema(
147
  profile,
148
  "learning_rate",
149
+ 0.00002 if profile.id == "oasis" else 0.0001 if profile.id in {"ddpm", "lora", "inrflow"} else 0.0002,
150
  )
151
  workers = max(1, min(8, (os.cpu_count() or 4) // 2))
152
  for key in ("dataloader_num_workers", "workers"):
 
155
  for key, value in {
156
  "gradient_accumulation_steps": 1,
157
  "gradient_accumulation": 1,
158
+ "mixed_precision": "fp16",
159
  "save_every": max(5, min(25, max(1, epochs // 10))),
160
  "preview_every": max(5, min(50, max(1, epochs // 10))),
161
  "training_intensity": 100,
162
  "gradient_checkpointing": resolution >= 384 or (available_vram is not None and available_vram < 8),
163
  "rank": 16,
164
  "alpha": 16,
165
+ "frame_gap": 1,
166
  "sequence_context": 1,
167
  "preview_steps": 1 if profile.id == "oasis" else 50 if profile.id == "ddpm" else 10,
168
  }.items():
169
  if key in profile.training:
170
  settings[key] = _clamp_to_schema(profile, key, value)
171
 
172
+ if profile.id == "inrflow" and "query_points" in profile.training:
173
+ settings["query_points"] = _clamp_to_schema(
174
+ profile, "query_points", min(1024, resolution * resolution)
175
+ )
176
+ reasons.append(
177
+ "INRFlow starts with at most 1,024 decoded pixel queries per image to keep training memory practical."
178
+ )
179
+
180
+ if profile.id == "oasis" and dataset_path:
181
+ from adam.oasis_dataset import dataset_directories, inspect_oasis_dataset, oasis_pace
182
+
183
+ pace = oasis_pace(dataset_path, frame_gap=int(settings.get("frame_gap", 1)))
184
+ recommended_gap = pace["recommended_frame_gap"]
185
+ capture_fps = pace["capture_fps"]
186
+ if isinstance(recommended_gap, int) and isinstance(capture_fps, (int, float)):
187
+ settings["frame_gap"] = _clamp_to_schema(profile, "frame_gap", recommended_gap)
188
+ native_fps = float(capture_fps) / int(settings["frame_gap"])
189
+ reasons.append(
190
+ f"The dataset records at {float(capture_fps):g} FPS, so prediction gap "
191
+ f"{settings['frame_gap']} gives a native trained pace of {native_fps:g} AI FPS."
192
+ )
193
+ report = inspect_oasis_dataset(dataset_path, frame_gap=int(settings.get("frame_gap", 1)))
194
+ if report.ok and report.valid_transitions:
195
+ transition_count = report.valid_transitions
196
+ # Large datasets need bounded epochs and a rotating, balanced sample.
197
+ # This keeps the recommendation in tens of thousands of updates rather
198
+ # than silently turning 10K captured frames into a multi-day run.
199
+ chunk_size = 5_000 if transition_count >= 7_500 else 0
200
+ transitions_per_epoch = min(transition_count, chunk_size) if chunk_size else transition_count
201
+ optimizer_steps_per_epoch = math.ceil(
202
+ transitions_per_epoch
203
+ / max(1, int(settings.get("batch_size", batch_size)))
204
+ / max(1, int(settings.get("gradient_accumulation", 1)))
205
+ )
206
+ target_updates = 50_000 if transition_count >= 7_500 else 30_000
207
+ epochs = max(5, min(45, math.ceil(target_updates / max(1, optimizer_steps_per_epoch))))
208
+ for key, value in {
209
+ "chunk_size": chunk_size,
210
+ "chunk_mode": "balanced",
211
+ "chunk_offset": 0,
212
+ "balance_actions": True,
213
+ "tf32": True,
214
+ "contrast_every": 4,
215
+ "contrast_samples": 2,
216
+ "recovery_minutes": 30,
217
+ "save_every": max(5, min(10, max(1, epochs // 4))),
218
+ "preview_every": max(2, min(10, max(1, epochs // 5))),
219
+ }.items():
220
+ if key in profile.training:
221
+ settings[key] = _clamp_to_schema(profile, key, value)
222
+ if len(dataset_directories(dataset_path)) > 1:
223
+ for key, value in {"include_older_data": True, "replay_older_percent": 50.0}.items():
224
+ if key in profile.training:
225
+ settings[key] = _clamp_to_schema(profile, key, value)
226
+ reasons.append(
227
+ f"{transition_count:,} valid transitions use "
228
+ f"{transitions_per_epoch:,} transition(s) per epoch, about "
229
+ f"{optimizer_steps_per_epoch:,} optimizer steps per epoch, and a "
230
+ f"{target_updates:,}-step initial budget."
231
+ )
232
+ if chunk_size:
233
+ reasons.append(
234
+ "A balanced 5,000-transition chunk keeps rare controls represented; "
235
+ "increase the chunk offset on a later continuation to rotate the sample."
236
+ )
237
+ idle_ratio = report.idle_rows / max(1, report.valid_rows)
238
+ if idle_ratio < 0.05:
239
+ warnings.append(
240
+ f"Only {idle_ratio:.1%} of labelled frames are idle. Record more no-input gameplay "
241
+ "so the world can stay stable when the player releases controls."
242
+ )
243
+ rare_threshold = max(10, math.ceil(report.valid_rows * 0.01))
244
+ rare_controls = [
245
+ name for name, count in report.action_counts.items()
246
+ if 0 < count < rare_threshold
247
+ ]
248
+ if rare_controls:
249
+ warnings.append(
250
+ "Rare recorded controls: " + ", ".join(rare_controls[:5])
251
+ + ". Action balancing is enabled, but more examples are still safer."
252
+ )
253
+
254
  estimated = estimate_vram_gb(profile, resolution, int(settings.get("batch_size", batch_size)), base_model_gb)
 
255
  if available_vram is not None and estimated > available_vram * 0.9:
256
  warnings.append(
257
  f"Estimated VRAM need is about {estimated:.1f} GB, above the conservative {available_vram * 0.9:.1f} GB working limit."
 
269
  memory_note = (
270
  f" using about {available_vram:.1f} GB available VRAM" if available_vram is not None else " without detected VRAM"
271
  )
272
+ if profile.id == "oasis":
273
+ summary = (
274
+ f"Recommended {epochs:,} Oasis epochs, batch {settings.get('batch_size', batch_size)}, "
275
+ f"prediction gap {settings.get('frame_gap', 1)}{memory_note}. "
276
+ "The recipe uses labelled transitions and is a starting point, not a guarantee."
277
+ )
278
+ else:
279
+ summary = (
280
+ f"Recommended {epochs:,} epochs for {images:,} item(s), "
281
+ f"batch {settings.get('batch_size', batch_size)} at {resolution}px{memory_note}. "
282
+ "Treat this as a starting recipe, not a guarantee."
283
+ )
284
  return SettingsRecommendation(
285
  profile_id=profile.id,
286
  epochs=epochs,
adam/remote_access.py CHANGED
@@ -24,6 +24,7 @@ from adam.generations import (
24
  load_generation_history,
25
  parse_chat_generation_request,
26
  )
 
27
  from adam.remote_dispatcher import RemoteCommandDispatcher
28
  from adam.remote_media import OpaqueIdCodec, RemoteMediaStore
29
  from adam.remote_v1 import RemoteV1Service
@@ -2413,9 +2414,11 @@ class RemoteAccessService:
2413
  preferred_id = {
2414
  "ddpm": "ddpm_generator",
2415
  "flow": "flow_generator",
 
 
2416
  "lora": "lora_generator",
2417
  }.get(parsed.provider_hint, "")
2418
- if stable_diffusion_request and parsed.provider_hint not in {"ddpm", "flow"}:
2419
  preferred_id = "lora_generator"
2420
  preferred_tool = next((item for item in tools if item.id == preferred_id), None)
2421
  if parsed.provider_hint and preferred_tool is None:
@@ -2443,7 +2446,23 @@ class RemoteAccessService:
2443
  key=lambda item: item[0],
2444
  reverse=True,
2445
  )
2446
- model = scored[0][1] if scored and scored[0][0] > 0 else None
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
2447
  if model is None and not model_query and len(candidates) == 1:
2448
  model = candidates[0]
2449
  if model is None and plain_model_search:
@@ -2534,8 +2553,8 @@ class RemoteAccessService:
2534
  extra_arguments = {
2535
  "negative_prompt": parsed.negative_prompt or str(saved_generation.get("negative_prompt", "")),
2536
  "base_model_path": base_model_path,
2537
- "width": 0,
2538
- "height": 0,
2539
  "cfg_scale": parsed.cfg_scale if parsed.cfg_scale is not None else float(saved_generation.get("cfg_scale", 0) or 0),
2540
  "lora_strength": 0.0 if base_only else (
2541
  parsed.lora_strength if parsed.lora_strength is not None else float(saved_generation.get("lora_strength", 0) or 0)
@@ -2568,6 +2587,8 @@ class RemoteAccessService:
2568
  for asset in getattr(self.planner.assets, "assets", [])
2569
  if asset.kind == "base_model" and Path(asset.path).exists()
2570
  ]
 
 
2571
  if parsed.base_model_query:
2572
  scored_bases = sorted(
2573
  (
 
24
  load_generation_history,
25
  parse_chat_generation_request,
26
  )
27
+ from adam.assets import Asset
28
  from adam.remote_dispatcher import RemoteCommandDispatcher
29
  from adam.remote_media import OpaqueIdCodec, RemoteMediaStore
30
  from adam.remote_v1 import RemoteV1Service
 
2414
  preferred_id = {
2415
  "ddpm": "ddpm_generator",
2416
  "flow": "flow_generator",
2417
+ "inrflow": "inrflow_generator",
2418
+ "pixelrow": "pixelrow_generator",
2419
  "lora": "lora_generator",
2420
  }.get(parsed.provider_hint, "")
2421
+ if stable_diffusion_request and parsed.provider_hint not in {"ddpm", "flow", "inrflow", "pixelrow"}:
2422
  preferred_id = "lora_generator"
2423
  preferred_tool = next((item for item in tools if item.id == preferred_id), None)
2424
  if parsed.provider_hint and preferred_tool is None:
 
2446
  key=lambda item: item[0],
2447
  reverse=True,
2448
  )
2449
+ model = next(
2450
+ (
2451
+ asset for asset in candidates
2452
+ if parsed.metadata_model_path
2453
+ and Path(asset.path).resolve() == Path(parsed.metadata_model_path).expanduser().resolve()
2454
+ ),
2455
+ None,
2456
+ )
2457
+ if model is None and parsed.metadata_model_path:
2458
+ direct_path = Path(parsed.metadata_model_path).expanduser()
2459
+ if direct_path.is_file() and direct_path.suffix.casefold() == ".safetensors":
2460
+ model = Asset(
2461
+ id="pasted-metadata", kind="model", name=direct_path.stem,
2462
+ path=str(direct_path.resolve()), trainer="lora",
2463
+ )
2464
+ if model is None:
2465
+ model = scored[0][1] if scored and scored[0][0] > 0 else None
2466
  if model is None and not model_query and len(candidates) == 1:
2467
  model = candidates[0]
2468
  if model is None and plain_model_search:
 
2553
  extra_arguments = {
2554
  "negative_prompt": parsed.negative_prompt or str(saved_generation.get("negative_prompt", "")),
2555
  "base_model_path": base_model_path,
2556
+ "width": parsed.width or 0,
2557
+ "height": parsed.height or 0,
2558
  "cfg_scale": parsed.cfg_scale if parsed.cfg_scale is not None else float(saved_generation.get("cfg_scale", 0) or 0),
2559
  "lora_strength": 0.0 if base_only else (
2560
  parsed.lora_strength if parsed.lora_strength is not None else float(saved_generation.get("lora_strength", 0) or 0)
 
2587
  for asset in getattr(self.planner.assets, "assets", [])
2588
  if asset.kind == "base_model" and Path(asset.path).exists()
2589
  ]
2590
+ if parsed.metadata_base_model_path and Path(parsed.metadata_base_model_path).expanduser().is_file():
2591
+ return str(Path(parsed.metadata_base_model_path).expanduser().resolve())
2592
  if parsed.base_model_query:
2593
  scored_bases = sorted(
2594
  (
adam/remote_dashboard.py CHANGED
@@ -22,7 +22,8 @@ def remote_dashboard_app_html() -> str:
22
  <div class="grid wide">
23
  <article class="card"><div class="section">Active Job</div><div id="activeJob">Checking ADAM...</div><div class="progress"><div id="activeBar" class="bar"></div></div><div class="row" style="margin-top:10px"><span class="pill">Time Left <b id="timeLeft">-</b></span><span class="pill">Finish <b id="finishTime">-</b></span></div><div class="row" style="margin-top:10px"><button data-action="pause">Pause</button><button data-action="resume">Resume</button><button class="bad" data-action="cancel">Cancel</button></div></article>
24
  <article class="card"><div class="section">Live Preview</div><img id="preview" class="preview" alt="Live preview" style="display:none"><div id="previewNote" class="status">Waiting for a preview.</div></article>
25
- <article class="card"><div class="section">Prompt ADAM</div><textarea id="prompt" placeholder="Ask ADAM naturally."></textarea><div class="row"><button class="primary" id="sendPrompt">Send</button><button class="ghost" id="askCreate">Create Model</button><button class="ghost" id="askDataset">Collect Dataset</button><button class="ghost" id="askQuick">Quick</button></div><div id="promptStatus" class="status"></div></article>
 
26
  <article class="card"><div class="section">Latest Generation</div><div id="latestGeneration" class="thumbs"></div></article>
27
  </div>
28
  </section>
@@ -62,16 +63,18 @@ def remote_dashboard_app_html() -> str:
62
  <pre id="trainingReview" class="status"></pre>
63
  </div>
64
  </article>
65
- </section>
66
-
67
- <section id="jobsView" class="view"><article class="card"><div class="section">Jobs</div><div id="queues" class="grid"></div></article></section>
 
 
68
  <section id="settingsView" class="view"><div class="grid"><article class="card"><div class="section">Remote Control</div><label class="toggle"><span>Auto-approve remote training</span><input id="autoApproveTraining" type="checkbox"></label><label class="toggle"><span>Keep screen updated</span><input id="keepAwake" type="checkbox"></label><div id="settingsStatus" class="status"></div></article><article class="card"><div class="section">Remembered Locations</div><div id="locationsList" class="dataset-list"></div></article><article class="card"><div class="section">System</div><div id="system" class="grid two"></div></article></div></section>
69
  </main>
70
  <nav class="tabs"><button class="active" data-view="home"><b>^</b>Home</button><button data-view="datasetsView"><b>O</b>Datasets</button><button data-view="createView"><b>+</b>Create</button><button data-view="jobsView"><b>/</b>Jobs</button><button data-view="settingsView"><b>*</b>Settings</button></nav>
71
  <script>
72
  (function(){
73
  var queryString=window.location.search||"";
74
- var state={datasets:[],locations:[],trainingSchema:{trainers:[],base_models:[],presets:[]},selectedDataset:"",datasetPage:1,selectedItem:null,activeJobId:"",refreshMs:3000,locationFilter:""};
75
  var timer=null;
76
  function $(id){return document.getElementById(id)}
77
  function list(v){return Array.isArray(v)?v:[]}
@@ -93,18 +96,19 @@ window.addEventListener("error",function(e){text("connection","Phone app error:
93
  window.addEventListener("unhandledrejection",function(e){text("connection","Remote request error: "+errorMessage(e.reason))});
94
 
95
  function img(src,alt,note){var im=document.createElement("img");im.alt=alt||"";if(src)im.src=authUrl(src);im.onerror=function(){im.style.display="none";if(note)note.textContent="Preview unavailable"};return im}
96
- function renderStatus(p){p=p||{};text("connection",(p.app||"ADAM")+" online - "+(p.scope||"remote"));var perms=p.permissions||{};var job=p.active_job||null;state.activeJobId=job&&job.id?job.id:"";if(job){text("activeJob",(job.project||"Active job")+" - "+(job.status||"")+" - "+String(job.progress||0)+"% - Time left "+((job.timing||{}).remaining_label||""));text("timeLeft",(job.timing||{}).remaining_label||"-");text("finishTime",(job.timing||{}).finish_label||"-");$("activeBar").style.width=String(job.progress||0)+"%"}else{text("activeJob","No active job.");text("timeLeft","-");text("finishTime","-");$("activeBar").style.width="0%"}var buttons=document.querySelectorAll("[data-action]");for(var i=0;i<buttons.length;i++)buttons[i].disabled=!job||!perms.job_control;$("autoApproveTraining").checked=!!perms.auto_approve_training;$("autoApproveTraining").disabled=!perms.job_control&&!perms.auto_approve_training;renderPreview(p.preview||{});renderLatest(p.latest_generation||{});renderSystem(p.system||{});renderQueues(p)}
97
  function renderPreview(info){var p=$("preview");if(info.available){p.style.display="block";p.src=authUrl("/api/preview","t="+Date.now());text("previewNote",(info.kind||"preview")+" "+(info.current||info.epoch||"")+"/"+(info.total||""))}else{p.style.display="none";p.removeAttribute("src");text("previewNote","Waiting for a preview.")}}
98
  function renderLatest(latest){var root=$("latestGeneration");clear(root);var images=list(latest.images).slice(0,8);if(!latest.available||!images.length){var d=document.createElement("div");d.className="status";d.textContent="Finished generated images will appear here.";root.appendChild(d);return}images.forEach(function(item){var b=document.createElement("button");b.className="thumb";b.type="button";var n=document.createElement("div");n.textContent=latest.model_name||"Generated image";b.appendChild(img(item.url,"Generated image",n));b.appendChild(n);b.onclick=function(){window.open(authUrl(item.url,"t="+Date.now()),"_blank","noopener,noreferrer")};root.appendChild(b)})}
99
  function renderSystem(sys){var root=$("system");clear(root);[["CPU",sys.cpu_percent],["RAM",sys.memory_percent],["GPU",sys.gpu_percent],["VRAM",sys.vram_percent]].forEach(function(m){var d=document.createElement("div");d.className="card";appendText(d,"div",m[0],"meta");appendText(d,"b",m[1]==null?"-":m[1]+"%");root.appendChild(d)})}
100
- function renderQueues(p){var root=$("queues");clear(root);var jobs=list(p.queue).concat(list(p.completed_jobs),list(p.failed_jobs)).slice(0,40);if(!jobs.length){text("queues","No jobs yet.");return}jobs.forEach(function(job){var d=document.createElement("div");d.className="queue-item";appendText(d,"b",job.project||"ADAM Job");appendText(d,"div",(job.status||"")+" - "+String(job.progress||0)+"%","meta");if(job.current_step_title){var s=document.createElement("div");s.className="meta";s.textContent=job.current_step_title;d.appendChild(s)}if(job.error){var e=document.createElement("div");e.className="bad-text";e.textContent=job.error;d.appendChild(e)}root.appendChild(d)})}
101
 
102
  function load(){return getJson("/api/status").then(renderStatus).catch(function(e){text("connection","Offline: "+errorMessage(e))})}
103
  function loadData(){return Promise.all([
104
  getJson("/api/v1/datasets").then(function(p){state.datasets=list(p.datasets);renderDatasets();fillDatasets()}),
105
  getJson("/api/v1/datasets/locations").then(function(p){state.locations=list(p.locations);renderLocations()}),
106
- getJson("/api/v1/training/schema").then(function(p){state.trainingSchema=p||{trainers:[]};fillCreate()})
107
- ]).catch(function(e){text("trainingReview",errorMessage(e))})}
 
108
  function datasetMatches(d){var q=($("datasetSearch").value||"").toLowerCase();if(q&&(d.name||"").toLowerCase().indexOf(q)<0)return false;if(state.locationFilter&&d.location_id!==state.locationFilter)return false;return true}
109
  function datasetImage(d){return d.thumbnail_url?img(d.thumbnail_url,"Dataset thumbnail"):document.createElement("span")}
110
  function useDataset(id){if(!id)return;postJson("/api/v1/datasets/"+encodeURIComponent(id)+"/use",{}).then(function(p){var d=p.dataset||{};state.selectedDataset=d.id||id;fillDatasets();if(!$("modelName").value)$("modelName").value=d.name||"";updateSelectedDatasetNote();switchView("createView");loadData()}).catch(function(e){text("datasetStats",errorMessage(e))})}
@@ -117,21 +121,26 @@ function updateSelectedDatasetNote(){var id=state.selectedDataset||($("trainData
117
  function openDataset(id,page){state.selectedDataset=id;state.datasetPage=page;text("datasetStats","Loading preview...");$("datasetDetail").style.display="block";getJson("/api/v1/datasets/"+encodeURIComponent(id)+"/items?page="+encodeURIComponent(page)+"&page_size=24").then(function(p){var d=p.dataset||{};text("datasetTitle",d.name||"Dataset");text("datasetStats",String(d.image_count||d.item_count||0)+" items - "+String(d.caption_count||0)+" captions - "+String(d.missing_caption_count||0)+" missing captions - "+(d.dataset_format||"Dataset"));text("pageLabel","Page "+String((p.pagination||{}).page||page));$("prevPage").disabled=page<=1;$("nextPage").disabled=!(p.pagination||{}).has_next;var grid=$("datasetGrid");clear(grid);list(p.items).forEach(function(item){var b=document.createElement("button");b.type="button";b.className="thumb";var note=document.createElement("div");note.textContent=(item.display_name||"Image")+"\\n"+(item.decision||"unreviewed");b.appendChild(img(item.thumbnail_url,"Dataset image",note));b.appendChild(note);b.onclick=function(){showImage(item)};grid.appendChild(b)})}).catch(function(e){text("datasetStats",errorMessage(e))})}
118
  function showImage(item){state.selectedItem=item;$("imageDetail").style.display="block";$("detailImage").src=authUrl(item.preview_url);text("detailName",item.display_name||"Image");text("detailDims",String((item.dimensions||{}).width||0)+" x "+String((item.dimensions||{}).height||0));$("captionEditor").value=item.caption||"";text("imageStatus",item.decision||"unreviewed")}
119
 
120
- function fillCreate(){clear($("trainer"));list(state.trainingSchema.trainers).forEach(function(t){option($("trainer"),t.name,t.id)});clear($("preset"));option($("preset"),"Custom","");list(state.trainingSchema.presets).forEach(function(p){option($("preset"),p.name,p.id)});clear($("baseModel"));list(state.trainingSchema.base_models).forEach(function(m){option($("baseModel"),m.name,m.id)});fillDatasets();renderSettings();ensureEpochField()}
 
 
 
 
121
  function makeSetting(key,spec,advanced){if(key==="base_model")return null;var wrap=document.createElement("div");wrap.className=spec.type==="bool"?"toggle":"form-row";var input;if(spec.type==="bool"){var label=document.createElement("div");appendText(label,"b",spec.label||key);appendText(label,"div",spec.description||"Enable this setting","hint");input=document.createElement("input");input.type="checkbox";input.checked=spec.default!==false;wrap.appendChild(label);wrap.appendChild(input)}else{var label=document.createElement("label");label.textContent=spec.label||key;wrap.appendChild(label);if(spec.type==="choice"){input=document.createElement("select");list(spec.options).forEach(function(v){option(input,String(v),String(v))})}else{input=document.createElement("input");input.type=(spec.type==="int"||spec.type==="float"||spec.type==="slider")?"number":"text";if(spec.min!==undefined)input.min=spec.min;if(spec.max!==undefined)input.max=spec.max;if(spec.step!==undefined)input.step=spec.step}input.value=spec.default!==undefined?String(spec.default):"";wrap.appendChild(input)}input.id=fieldId(key);input.setAttribute("data-setting-key",key);input.setAttribute("data-setting-type",spec.type||"text");input.setAttribute("data-advanced",advanced?"1":"0");return wrap}
122
  function renderSettings(){var t=activeTrainer();var schema=t.settings||{};var basic=$("basicSettings"),advanced=$("advancedSettings");clear(basic);clear(advanced);var basicKeys=["resolution","batch_size","preview_enabled"];Object.keys(schema).forEach(function(key){var spec=schema[key]||{};var isAdvanced=!!spec.advanced||basicKeys.indexOf(key)<0;var node=makeSetting(key,spec,isAdvanced);if(!node)return;(isAdvanced?advanced:basic).appendChild(node)});$("baseModelWrap").style.display=schema.base_model?"grid":"none"}
123
  function settingValue(node){var type=node.getAttribute("data-setting-type");if(type==="bool")return !!node.checked;if(type==="int"||type==="slider")return Number(node.value||0);if(type==="float")return Number(node.value||0);var spec=schemaFor(node.getAttribute("data-setting-key"));if(spec.type==="choice"){var sample=list(spec.options)[0];if(typeof sample==="number")return Number(node.value)}return node.value}
124
  function trainingPayload(){var settings={};var nodes=document.querySelectorAll("[data-setting-key]");for(var i=0;i<nodes.length;i++){var key=nodes[i].getAttribute("data-setting-key");if(key==="epochs")continue;settings[key]=settingValue(nodes[i])}return{trainer:$("trainer").value,dataset_id:state.selectedDataset||$("trainDataset").value,base_model_id:$("baseModel").value,model_name:$("modelName").value,trigger_word:settings.trigger_word||"",epochs:Number($(fieldId("epochs"))&&$(fieldId("epochs")).value||10),settings:settings}}
125
  function applyPreset(){var id=$("preset").value;var preset=list(state.trainingSchema.presets).filter(function(p){return p.id===id})[0];if(!preset)return;if(preset.trainer){$("trainer").value=preset.trainer;renderSettings();ensureEpochField()}if(preset.epochs&&$(fieldId("epochs")))$(fieldId("epochs")).value=preset.epochs;Object.keys(preset.settings||{}).forEach(function(key){var n=$(fieldId(key));if(!n)return;if(n.type==="checkbox")n.checked=!!preset.settings[key];else n.value=String(preset.settings[key])})}
126
  function ensureEpochField(){if(!$(fieldId("epochs"))){var node=makeSetting("epochs",{label:"Epochs",type:"int",default:10,min:1,max:100000},false);$("basicSettings").appendChild(node)}}
127
- $("trainer").onchange=function(){renderSettings();ensureEpochField()};$("preset").onchange=applyPreset;$("datasetSearch").oninput=renderDatasets;$("filterDatasets").onclick=function(){var chips=$("locationChips");chips.style.display=chips.style.display==="none"?"flex":"none"};$("chooseDataset").onclick=function(){switchView("datasetsView")};$("createBack").onclick=function(){switchView("home")};$("createHelp").onclick=function(){text("trainingReview","Choose a preset, trainer, model name, and dataset. Advanced settings come from ADAM's trainer registry.")};$("askInstead").onclick=function(){switchView("home");$("prompt").focus()};$("useDataset").onclick=function(){useDataset(state.selectedDataset)};$("selectDatasetForCreate").onclick=$("useDataset").onclick;
128
  $("reviewPlan").onclick=function(){postJson("/api/v1/training/plan",trainingPayload()).then(function(p){$("trainingReview").textContent=JSON.stringify(p.plan||p,null,2)}).catch(function(e){text("trainingReview",errorMessage(e))})};
129
- $("startTraining").onclick=function(){if(!state.selectedDataset&&!$("trainDataset").value){text("trainingReview","Choose a dataset first.");return}if(!confirm("Start this training job on the PC?"))return;postJson("/api/v1/training/start",trainingPayload()).then(function(p){text("trainingReview",p.message||"Training queued.");load();loadData()}).catch(function(e){text("trainingReview",errorMessage(e))})};
 
130
  $("prevPage").onclick=function(){openDataset(state.selectedDataset,Math.max(1,state.datasetPage-1))};$("nextPage").onclick=function(){openDataset(state.selectedDataset,state.datasetPage+1)};
131
  $("saveCaption").onclick=function(){if(!state.selectedDataset||!state.selectedItem)return;postJson("/api/v1/datasets/"+encodeURIComponent(state.selectedDataset)+"/items/"+encodeURIComponent(state.selectedItem.id)+"/caption",{caption:$("captionEditor").value}).then(function(p){text("imageStatus",p.message||"Saved.")}).catch(function(e){text("imageStatus",errorMessage(e))})};
132
  function decide(decision){if(!state.selectedDataset||!state.selectedItem)return;postJson("/api/v1/datasets/"+encodeURIComponent(state.selectedDataset)+"/items/"+encodeURIComponent(state.selectedItem.id)+"/decision",{decision:decision}).then(function(p){text("imageStatus",p.message||decision);openDataset(state.selectedDataset,state.datasetPage)}).catch(function(e){text("imageStatus",errorMessage(e))})}
133
  $("keepImage").onclick=function(){decide("keep")};$("rejectImage").onclick=function(){decide("reject")};$("unreviewImage").onclick=function(){decide("unreviewed")};
134
- $("sendPrompt").onclick=function(){postJson("/api/prompt",{prompt:$("prompt").value}).then(function(p){text("promptStatus",p.message||"Sent.");$("prompt").value="";load()}).catch(function(e){text("promptStatus",errorMessage(e))})};$("askCreate").onclick=function(){switchView("createView")};$("askDataset").onclick=function(){$("prompt").value="Collect a new dataset for Example."};$("askQuick").onclick=function(){$("prompt").value="Generate an image of Example."};
135
  $("refresh").onclick=function(){load();loadData()};$("autoApproveTraining").onchange=function(e){postJson("/api/remote-settings",{auto_approve_training:e.target.checked}).then(function(p){text("settingsStatus",p.message||"Saved.");load()}).catch(function(err){text("settingsStatus",errorMessage(err))})};$("keepAwake").onchange=function(e){state.refreshMs=e.target.checked?1500:5000;startTimer();text("settingsStatus",e.target.checked?"Fast refresh is on.":"Quiet refresh is on.")};
136
  function startTimer(){if(timer)clearInterval(timer);timer=setInterval(load,state.refreshMs)}
137
  load();loadData();startTimer();
 
22
  <div class="grid wide">
23
  <article class="card"><div class="section">Active Job</div><div id="activeJob">Checking ADAM...</div><div class="progress"><div id="activeBar" class="bar"></div></div><div class="row" style="margin-top:10px"><span class="pill">Time Left <b id="timeLeft">-</b></span><span class="pill">Finish <b id="finishTime">-</b></span></div><div class="row" style="margin-top:10px"><button data-action="pause">Pause</button><button data-action="resume">Resume</button><button class="bad" data-action="cancel">Cancel</button></div></article>
24
  <article class="card"><div class="section">Live Preview</div><img id="preview" class="preview" alt="Live preview" style="display:none"><div id="previewNote" class="status">Waiting for a preview.</div></article>
25
+ <article class="card"><div class="section">Prompt ADAM</div><textarea id="prompt" placeholder="Ask ADAM naturally."></textarea><div class="row"><button class="primary" id="sendPrompt">Send</button><button class="ghost" id="askCreate">Create Model</button><button class="ghost" id="askDataset">Collect Dataset</button><button class="ghost" id="askQuick">Quick</button></div><div id="promptStatus" class="status"></div></article>
26
+ <article class="card"><div class="section">Image Generation</div><div class="status">Choose the model directly instead of describing a LoRA in a prompt.</div><button id="openGeneration" class="primary" style="margin-top:10px">Generate an Image</button></article>
27
  <article class="card"><div class="section">Latest Generation</div><div id="latestGeneration" class="thumbs"></div></article>
28
  </div>
29
  </section>
 
63
  <pre id="trainingReview" class="status"></pre>
64
  </div>
65
  </article>
66
+ </section>
67
+
68
+ <section id="generationView" class="view"><div class="titlebar"><button id="generationBack" class="icon" title="Back">&lt;</button><h1>Generate Image</h1><span></span></div><article class="panel"><div class="form"><div class="form-row"><label>Provider</label><select id="generationProvider"></select></div><div class="form-row stack"><label>Model</label><select id="generationModel"></select><div id="generationModelNote" class="hint"></div></div><div id="generationBaseWrap" class="form-row stack"><label>Base model</label><select id="generationBase"></select></div><div class="form-row stack"><label>Prompt</label><textarea id="generationPrompt" placeholder="Describe the image you want."></textarea></div><div class="form-row stack"><label>Negative prompt</label><textarea id="generationNegative" placeholder="Optional"></textarea></div><div class="settings-grid"><div class="form-row"><label>Images</label><input id="generationCount" type="number" min="1" max="8" value="1"></div><div class="form-row"><label>Steps</label><input id="generationSteps" type="number" min="1" value="30"></div><div class="form-row"><label>Seed</label><input id="generationSeed" type="number" min="0" value="0"></div><div class="form-row"><label>LoRA strength</label><input id="generationStrength" type="number" min="0" max="3" step="0.05" value="1"></div></div><button id="startGeneration" class="primary big-action">Generate</button><div id="generationStatus" class="status"></div></div></article></section>
69
+
70
+ <section id="jobsView" class="view"><article class="card"><div class="section">Jobs</div><div id="queues" class="grid"></div></article></section>
71
  <section id="settingsView" class="view"><div class="grid"><article class="card"><div class="section">Remote Control</div><label class="toggle"><span>Auto-approve remote training</span><input id="autoApproveTraining" type="checkbox"></label><label class="toggle"><span>Keep screen updated</span><input id="keepAwake" type="checkbox"></label><div id="settingsStatus" class="status"></div></article><article class="card"><div class="section">Remembered Locations</div><div id="locationsList" class="dataset-list"></div></article><article class="card"><div class="section">System</div><div id="system" class="grid two"></div></article></div></section>
72
  </main>
73
  <nav class="tabs"><button class="active" data-view="home"><b>^</b>Home</button><button data-view="datasetsView"><b>O</b>Datasets</button><button data-view="createView"><b>+</b>Create</button><button data-view="jobsView"><b>/</b>Jobs</button><button data-view="settingsView"><b>*</b>Settings</button></nav>
74
  <script>
75
  (function(){
76
  var queryString=window.location.search||"";
77
+ var state={datasets:[],locations:[],trainingSchema:{trainers:[],base_models:[],presets:[]},generationSchema:{providers:[],models:[],base_models:[]},selectedDataset:"",datasetPage:1,selectedItem:null,activeJobId:"",promptJobId:"",promptLocked:false,refreshMs:3000,locationFilter:""};
78
  var timer=null;
79
  function $(id){return document.getElementById(id)}
80
  function list(v){return Array.isArray(v)?v:[]}
 
96
  window.addEventListener("unhandledrejection",function(e){text("connection","Remote request error: "+errorMessage(e.reason))});
97
 
98
  function img(src,alt,note){var im=document.createElement("img");im.alt=alt||"";if(src)im.src=authUrl(src);im.onerror=function(){im.style.display="none";if(note)note.textContent="Preview unavailable"};return im}
99
+ function renderStatus(p){p=p||{};text("connection",(p.app||"ADAM")+" online - "+(p.scope||"remote"));var perms=p.permissions||{};var job=p.active_job||null;state.activeJobId=job&&job.id?job.id:"";if(state.promptLocked&&state.promptJobId&&job&&job.id===state.promptJobId){state.promptLocked=false;state.promptJobId="";$("sendPrompt").disabled=false;text("promptStatus","Job started. You can send another prompt.")}if(job){text("activeJob",(job.project||"Active job")+" - "+(job.status||"")+" - "+String(job.progress||0)+"% - Time left "+((job.timing||{}).remaining_label||""));text("timeLeft",(job.timing||{}).remaining_label||"-");text("finishTime",(job.timing||{}).finish_label||"-");$("activeBar").style.width=String(job.progress||0)+"%"}else{text("activeJob","No active job.");text("timeLeft","-");text("finishTime","-");$("activeBar").style.width="0%"}var buttons=document.querySelectorAll("[data-action]");for(var i=0;i<buttons.length;i++)buttons[i].disabled=!job||!perms.job_control;$("autoApproveTraining").checked=!!perms.auto_approve_training;$("autoApproveTraining").disabled=!perms.job_control&&!perms.auto_approve_training;renderPreview(p.preview||{});renderLatest(p.latest_generation||{});renderSystem(p.system||{});renderQueues(p,perms)}
100
  function renderPreview(info){var p=$("preview");if(info.available){p.style.display="block";p.src=authUrl("/api/preview","t="+Date.now());text("previewNote",(info.kind||"preview")+" "+(info.current||info.epoch||"")+"/"+(info.total||""))}else{p.style.display="none";p.removeAttribute("src");text("previewNote","Waiting for a preview.")}}
101
  function renderLatest(latest){var root=$("latestGeneration");clear(root);var images=list(latest.images).slice(0,8);if(!latest.available||!images.length){var d=document.createElement("div");d.className="status";d.textContent="Finished generated images will appear here.";root.appendChild(d);return}images.forEach(function(item){var b=document.createElement("button");b.className="thumb";b.type="button";var n=document.createElement("div");n.textContent=latest.model_name||"Generated image";b.appendChild(img(item.url,"Generated image",n));b.appendChild(n);b.onclick=function(){window.open(authUrl(item.url,"t="+Date.now()),"_blank","noopener,noreferrer")};root.appendChild(b)})}
102
  function renderSystem(sys){var root=$("system");clear(root);[["CPU",sys.cpu_percent],["RAM",sys.memory_percent],["GPU",sys.gpu_percent],["VRAM",sys.vram_percent]].forEach(function(m){var d=document.createElement("div");d.className="card";appendText(d,"div",m[0],"meta");appendText(d,"b",m[1]==null?"-":m[1]+"%");root.appendChild(d)})}
103
+ function renderQueues(p,perms){var root=$("queues");clear(root);var jobs=list(p.queue).concat(list(p.completed_jobs),list(p.failed_jobs)).slice(0,40);if(!jobs.length){text("queues","No jobs yet.");return}jobs.forEach(function(job){var d=document.createElement("div");d.className="queue-item";appendText(d,"b",job.project||"ADAM Job");appendText(d,"div",(job.status||"")+" - "+String(job.progress||0)+"%","meta");if(job.current_step_title){var s=document.createElement("div");s.className="meta";s.textContent=job.current_step_title;d.appendChild(s)}if(job.error){var e=document.createElement("div");e.className="bad-text";e.textContent=job.error;d.appendChild(e)}if(perms&&perms.job_control&&["Queued","Scheduled","Awaiting confirmation","Running","Paused","Interrupted"].indexOf(job.status)>=0){var cancel=document.createElement("button");cancel.className="bad";cancel.textContent="Cancel job";cancel.style.marginTop="8px";cancel.onclick=function(){cancel.disabled=true;postJson("/api/job",{job_id:job.id,action:"cancel"}).then(function(){load()}).catch(function(e){text("promptStatus",errorMessage(e));cancel.disabled=false})};d.appendChild(cancel)}root.appendChild(d)})}
104
 
105
  function load(){return getJson("/api/status").then(renderStatus).catch(function(e){text("connection","Offline: "+errorMessage(e))})}
106
  function loadData(){return Promise.all([
107
  getJson("/api/v1/datasets").then(function(p){state.datasets=list(p.datasets);renderDatasets();fillDatasets()}),
108
  getJson("/api/v1/datasets/locations").then(function(p){state.locations=list(p.locations);renderLocations()}),
109
+ getJson("/api/v1/training/schema").then(function(p){state.trainingSchema=p||{trainers:[]};fillCreate()}),
110
+ getJson("/api/v1/generation/schema").then(function(p){state.generationSchema=p||{providers:[],models:[],base_models:[]};fillGeneration()})
111
+ ]).catch(function(e){text("trainingReview",errorMessage(e))})}
112
  function datasetMatches(d){var q=($("datasetSearch").value||"").toLowerCase();if(q&&(d.name||"").toLowerCase().indexOf(q)<0)return false;if(state.locationFilter&&d.location_id!==state.locationFilter)return false;return true}
113
  function datasetImage(d){return d.thumbnail_url?img(d.thumbnail_url,"Dataset thumbnail"):document.createElement("span")}
114
  function useDataset(id){if(!id)return;postJson("/api/v1/datasets/"+encodeURIComponent(id)+"/use",{}).then(function(p){var d=p.dataset||{};state.selectedDataset=d.id||id;fillDatasets();if(!$("modelName").value)$("modelName").value=d.name||"";updateSelectedDatasetNote();switchView("createView");loadData()}).catch(function(e){text("datasetStats",errorMessage(e))})}
 
121
  function openDataset(id,page){state.selectedDataset=id;state.datasetPage=page;text("datasetStats","Loading preview...");$("datasetDetail").style.display="block";getJson("/api/v1/datasets/"+encodeURIComponent(id)+"/items?page="+encodeURIComponent(page)+"&page_size=24").then(function(p){var d=p.dataset||{};text("datasetTitle",d.name||"Dataset");text("datasetStats",String(d.image_count||d.item_count||0)+" items - "+String(d.caption_count||0)+" captions - "+String(d.missing_caption_count||0)+" missing captions - "+(d.dataset_format||"Dataset"));text("pageLabel","Page "+String((p.pagination||{}).page||page));$("prevPage").disabled=page<=1;$("nextPage").disabled=!(p.pagination||{}).has_next;var grid=$("datasetGrid");clear(grid);list(p.items).forEach(function(item){var b=document.createElement("button");b.type="button";b.className="thumb";var note=document.createElement("div");note.textContent=(item.display_name||"Image")+"\\n"+(item.decision||"unreviewed");b.appendChild(img(item.thumbnail_url,"Dataset image",note));b.appendChild(note);b.onclick=function(){showImage(item)};grid.appendChild(b)})}).catch(function(e){text("datasetStats",errorMessage(e))})}
122
  function showImage(item){state.selectedItem=item;$("imageDetail").style.display="block";$("detailImage").src=authUrl(item.preview_url);text("detailName",item.display_name||"Image");text("detailDims",String((item.dimensions||{}).width||0)+" x "+String((item.dimensions||{}).height||0));$("captionEditor").value=item.caption||"";text("imageStatus",item.decision||"unreviewed")}
123
 
124
+ function fillCreate(){clear($("trainer"));list(state.trainingSchema.trainers).forEach(function(t){option($("trainer"),t.name,t.id)});clear($("preset"));option($("preset"),"Custom","");list(state.trainingSchema.presets).forEach(function(p){option($("preset"),p.name,p.id)});clear($("baseModel"));list(state.trainingSchema.base_models).forEach(function(m){option($("baseModel"),m.name,m.id)});fillDatasets();renderSettings();ensureEpochField()}
125
+ function activeGenerationProvider(){return list(state.generationSchema.providers).filter(function(p){return p.id===$("generationProvider").value})[0]||{model_trainers:[],options:{}}}
126
+ function fillGeneration(){var current=$("generationProvider").value;clear($("generationProvider"));list(state.generationSchema.providers).forEach(function(p){option($("generationProvider"),p.name,p.id)});if(current)$("generationProvider").value=current;fillGenerationModels()}
127
+ function fillGenerationModels(){var provider=activeGenerationProvider(),current=$("generationModel").value,trainers=list(provider.model_trainers);clear($("generationModel"));list(state.generationSchema.models).filter(function(m){return !trainers.length||trainers.indexOf(m.trainer)>=0}).forEach(function(m){option($("generationModel"),m.name+(m.trigger_word?" · trigger: "+m.trigger_word:""),m.id)});if(current)$("generationModel").value=current;clear($("generationBase"));list(state.generationSchema.base_models).forEach(function(m){option($("generationBase"),m.name,m.id)});var isLora=provider.id==="lora_generator";$("generationBaseWrap").style.display=isLora?"grid":"none";$("generationStrength").parentNode.style.display=isLora?"grid":"none";var selected=list(state.generationSchema.models).filter(function(m){return m.id===$("generationModel").value})[0];text("generationModelNote",selected&&selected.trigger_word?"LoRA trigger word: "+selected.trigger_word:"Choose a completed model.");var opts=provider.options||{};if(opts.step_default)$("generationSteps").value=String(opts.step_default)}
128
+ function generationPayload(){var provider=activeGenerationProvider();return{provider_id:$("generationProvider").value,model_id:$("generationModel").value,base_model_id:$("generationBase").value,prompt:$("generationPrompt").value,negative_prompt:$("generationNegative").value,image_count:Number($("generationCount").value||1),steps:Number($("generationSteps").value||30),seed:Number($("generationSeed").value||0),sampler:(provider.options&&list(provider.options.samplers)[0])||"DDIM",aspect_ratio:(provider.options&&list(provider.options.aspect_ratios)[0])||"1:1 (Square)",lora_strength:Number($("generationStrength").value||1),prompt_weighting:true}}
129
  function makeSetting(key,spec,advanced){if(key==="base_model")return null;var wrap=document.createElement("div");wrap.className=spec.type==="bool"?"toggle":"form-row";var input;if(spec.type==="bool"){var label=document.createElement("div");appendText(label,"b",spec.label||key);appendText(label,"div",spec.description||"Enable this setting","hint");input=document.createElement("input");input.type="checkbox";input.checked=spec.default!==false;wrap.appendChild(label);wrap.appendChild(input)}else{var label=document.createElement("label");label.textContent=spec.label||key;wrap.appendChild(label);if(spec.type==="choice"){input=document.createElement("select");list(spec.options).forEach(function(v){option(input,String(v),String(v))})}else{input=document.createElement("input");input.type=(spec.type==="int"||spec.type==="float"||spec.type==="slider")?"number":"text";if(spec.min!==undefined)input.min=spec.min;if(spec.max!==undefined)input.max=spec.max;if(spec.step!==undefined)input.step=spec.step}input.value=spec.default!==undefined?String(spec.default):"";wrap.appendChild(input)}input.id=fieldId(key);input.setAttribute("data-setting-key",key);input.setAttribute("data-setting-type",spec.type||"text");input.setAttribute("data-advanced",advanced?"1":"0");return wrap}
130
  function renderSettings(){var t=activeTrainer();var schema=t.settings||{};var basic=$("basicSettings"),advanced=$("advancedSettings");clear(basic);clear(advanced);var basicKeys=["resolution","batch_size","preview_enabled"];Object.keys(schema).forEach(function(key){var spec=schema[key]||{};var isAdvanced=!!spec.advanced||basicKeys.indexOf(key)<0;var node=makeSetting(key,spec,isAdvanced);if(!node)return;(isAdvanced?advanced:basic).appendChild(node)});$("baseModelWrap").style.display=schema.base_model?"grid":"none"}
131
  function settingValue(node){var type=node.getAttribute("data-setting-type");if(type==="bool")return !!node.checked;if(type==="int"||type==="slider")return Number(node.value||0);if(type==="float")return Number(node.value||0);var spec=schemaFor(node.getAttribute("data-setting-key"));if(spec.type==="choice"){var sample=list(spec.options)[0];if(typeof sample==="number")return Number(node.value)}return node.value}
132
  function trainingPayload(){var settings={};var nodes=document.querySelectorAll("[data-setting-key]");for(var i=0;i<nodes.length;i++){var key=nodes[i].getAttribute("data-setting-key");if(key==="epochs")continue;settings[key]=settingValue(nodes[i])}return{trainer:$("trainer").value,dataset_id:state.selectedDataset||$("trainDataset").value,base_model_id:$("baseModel").value,model_name:$("modelName").value,trigger_word:settings.trigger_word||"",epochs:Number($(fieldId("epochs"))&&$(fieldId("epochs")).value||10),settings:settings}}
133
  function applyPreset(){var id=$("preset").value;var preset=list(state.trainingSchema.presets).filter(function(p){return p.id===id})[0];if(!preset)return;if(preset.trainer){$("trainer").value=preset.trainer;renderSettings();ensureEpochField()}if(preset.epochs&&$(fieldId("epochs")))$(fieldId("epochs")).value=preset.epochs;Object.keys(preset.settings||{}).forEach(function(key){var n=$(fieldId(key));if(!n)return;if(n.type==="checkbox")n.checked=!!preset.settings[key];else n.value=String(preset.settings[key])})}
134
  function ensureEpochField(){if(!$(fieldId("epochs"))){var node=makeSetting("epochs",{label:"Epochs",type:"int",default:10,min:1,max:100000},false);$("basicSettings").appendChild(node)}}
135
+ $("trainer").onchange=function(){renderSettings();ensureEpochField()};$("preset").onchange=applyPreset;$("datasetSearch").oninput=renderDatasets;$("filterDatasets").onclick=function(){var chips=$("locationChips");chips.style.display=chips.style.display==="none"?"flex":"none"};$("chooseDataset").onclick=function(){switchView("datasetsView")};$("createBack").onclick=function(){switchView("home")};$("createHelp").onclick=function(){text("trainingReview","Choose a preset, trainer, model name, and dataset. Advanced settings come from ADAM's trainer registry.")};$("askInstead").onclick=function(){switchView("home");$("prompt").focus()};$("useDataset").onclick=function(){useDataset(state.selectedDataset)};$("selectDatasetForCreate").onclick=$("useDataset").onclick;$("openGeneration").onclick=function(){switchView("generationView")};$("generationBack").onclick=function(){switchView("home")};$("generationProvider").onchange=fillGenerationModels;$("generationModel").onchange=fillGenerationModels;
136
  $("reviewPlan").onclick=function(){postJson("/api/v1/training/plan",trainingPayload()).then(function(p){$("trainingReview").textContent=JSON.stringify(p.plan||p,null,2)}).catch(function(e){text("trainingReview",errorMessage(e))})};
137
+ $("startTraining").onclick=function(){if(!state.selectedDataset&&!$("trainDataset").value){text("trainingReview","Choose a dataset first.");return}if(!confirm("Start this training job on the PC?"))return;postJson("/api/v1/training/start",trainingPayload()).then(function(p){text("trainingReview",p.message||"Training queued.");load();loadData()}).catch(function(e){text("trainingReview",errorMessage(e))})};
138
+ $("startGeneration").onclick=function(){var button=this;button.disabled=true;postJson("/api/v1/generation/start",generationPayload()).then(function(p){text("generationStatus",p.message||"Generation queued.");load()}).catch(function(e){text("generationStatus",errorMessage(e))}).then(function(){button.disabled=false})};
139
  $("prevPage").onclick=function(){openDataset(state.selectedDataset,Math.max(1,state.datasetPage-1))};$("nextPage").onclick=function(){openDataset(state.selectedDataset,state.datasetPage+1)};
140
  $("saveCaption").onclick=function(){if(!state.selectedDataset||!state.selectedItem)return;postJson("/api/v1/datasets/"+encodeURIComponent(state.selectedDataset)+"/items/"+encodeURIComponent(state.selectedItem.id)+"/caption",{caption:$("captionEditor").value}).then(function(p){text("imageStatus",p.message||"Saved.")}).catch(function(e){text("imageStatus",errorMessage(e))})};
141
  function decide(decision){if(!state.selectedDataset||!state.selectedItem)return;postJson("/api/v1/datasets/"+encodeURIComponent(state.selectedDataset)+"/items/"+encodeURIComponent(state.selectedItem.id)+"/decision",{decision:decision}).then(function(p){text("imageStatus",p.message||decision);openDataset(state.selectedDataset,state.datasetPage)}).catch(function(e){text("imageStatus",errorMessage(e))})}
142
  $("keepImage").onclick=function(){decide("keep")};$("rejectImage").onclick=function(){decide("reject")};$("unreviewImage").onclick=function(){decide("unreviewed")};
143
+ $("sendPrompt").onclick=function(){var button=this;if(state.promptLocked)return;button.disabled=true;postJson("/api/prompt",{prompt:$("prompt").value}).then(function(p){state.promptJobId=p.job_id||"";state.promptLocked=!!state.promptJobId;text("promptStatus",state.promptLocked?"Queued. Sending is available again when this job starts.":(p.message||"Sent."));$("prompt").value="";if(!state.promptLocked)button.disabled=false;load()}).catch(function(e){text("promptStatus",errorMessage(e));button.disabled=false})};$("askCreate").onclick=function(){switchView("createView")};$("askDataset").onclick=function(){$("prompt").value="Collect a new dataset for Example."};$("askQuick").onclick=function(){$("prompt").value="Generate an image of Example."};
144
  $("refresh").onclick=function(){load();loadData()};$("autoApproveTraining").onchange=function(e){postJson("/api/remote-settings",{auto_approve_training:e.target.checked}).then(function(p){text("settingsStatus",p.message||"Saved.");load()}).catch(function(err){text("settingsStatus",errorMessage(err))})};$("keepAwake").onchange=function(e){state.refreshMs=e.target.checked?1500:5000;startTimer();text("settingsStatus",e.target.checked?"Fast refresh is on.":"Quiet refresh is on.")};
145
  function startTimer(){if(timer)clearInterval(timer);timer=setInterval(load,state.refreshMs)}
146
  load();loadData();startTimer();
adam/remote_v1.py CHANGED
@@ -3,6 +3,8 @@ from __future__ import annotations
3
  from pathlib import Path
4
  import os
5
  import tempfile
 
 
6
  from typing import Any
7
  from urllib.parse import parse_qs
8
 
@@ -55,6 +57,12 @@ class RemoteV1Service:
55
  self.auto_approve_training = auto_approve_training
56
  self._asset_fallback: AssetRegistry | None = None
57
  self._studio: StudioStore | None = None
 
 
 
 
 
 
58
 
59
  def route(
60
  self,
@@ -424,40 +432,45 @@ class RemoteV1Service:
424
  return self._dataset_summary(asset, record)
425
 
426
  def models(self) -> list[dict[str, Any]]:
427
- assets = self._assets()
428
- experiment_by_model: dict[str, Any] = {}
429
- try:
430
- for run in getattr(getattr(self.jobs, "experiments", None), "list_runs", lambda limit=100: [])(limit=100):
431
- experiment_by_model.setdefault(run.model_name, run)
432
- except Exception:
433
- experiment_by_model = {}
434
- rows = []
435
- for asset in assets.assets:
436
- if asset.kind not in {"model", "base_model"} or not Path(asset.path).exists():
437
- continue
438
- dataset = next((item for item in assets.assets if item.id == asset.dataset_id), None)
439
- metadata = dict(asset.metadata or {})
440
- trigger_word = str(metadata.get("trigger_word") or (asset.name if asset.trainer == "lora" else ""))
441
- latest = experiment_by_model.get(asset.name)
442
- rows.append({
443
- "id": self._asset_public_id(asset),
444
- "name": asset.name,
445
- "kind": asset.kind,
446
- "architecture": asset.trainer or ("stable_diffusion" if asset.kind == "base_model" else ""),
447
- "trainer": asset.trainer,
448
- "checkpoint_name": Path(asset.checkpoint or asset.path).name,
449
- "dataset": None if dataset is None else {"id": self._asset_public_id(dataset), "name": dataset.name},
450
- "epochs": asset.epochs,
451
- "trigger_word": trigger_word,
452
- "latest_experiment": None if latest is None else {
453
- "id": latest.id,
454
- "status": latest.status,
455
- "resolution": latest.resolution,
456
- "dataset_name": latest.dataset_name,
457
- "trigger_word": getattr(latest, "trigger_word", ""),
458
- },
459
- })
460
- return rows[:300]
 
 
 
 
 
461
 
462
  def training_schema(self) -> dict[str, Any]:
463
  registry = self._registry()
@@ -641,7 +654,12 @@ class RemoteV1Service:
641
  "options": tool.generation_options,
642
  "settings": registry.model_plugins.generation_schema_for_tool(tool.id),
643
  })
644
- return {"providers": providers, "models": [item for item in self.models() if item["kind"] == "model"], "base_models": [item for item in self.models() if item["kind"] == "base_model"]}
 
 
 
 
 
645
 
646
  def start_generation(self, payload: dict[str, Any]) -> dict[str, Any]:
647
  if self.jobs is None:
 
3
  from pathlib import Path
4
  import os
5
  import tempfile
6
+ import threading
7
+ import time
8
  from typing import Any
9
  from urllib.parse import parse_qs
10
 
 
57
  self.auto_approve_training = auto_approve_training
58
  self._asset_fallback: AssetRegistry | None = None
59
  self._studio: StudioStore | None = None
60
+ # The dashboard requests both schemas at once. They share an asset
61
+ # registry, so build one short-lived catalog instead of asking two
62
+ # request threads to rediscover and rewrite it concurrently.
63
+ self._model_catalog_lock = threading.RLock()
64
+ self._model_catalog: list[dict[str, Any]] = []
65
+ self._model_catalog_at = 0.0
66
 
67
  def route(
68
  self,
 
432
  return self._dataset_summary(asset, record)
433
 
434
  def models(self) -> list[dict[str, Any]]:
435
+ with self._model_catalog_lock:
436
+ if time.monotonic() - self._model_catalog_at < 1.0:
437
+ return list(self._model_catalog)
438
+ assets = self._assets()
439
+ experiment_by_model: dict[str, Any] = {}
440
+ try:
441
+ for run in getattr(getattr(self.jobs, "experiments", None), "list_runs", lambda limit=100: [])(limit=100):
442
+ experiment_by_model.setdefault(run.model_name, run)
443
+ except Exception:
444
+ experiment_by_model = {}
445
+ rows = []
446
+ for asset in assets.assets:
447
+ if asset.kind not in {"model", "base_model"} or not Path(asset.path).exists():
448
+ continue
449
+ dataset = next((item for item in assets.assets if item.id == asset.dataset_id), None)
450
+ metadata = dict(asset.metadata or {})
451
+ trigger_word = str(metadata.get("trigger_word") or (asset.name if asset.trainer == "lora" else ""))
452
+ latest = experiment_by_model.get(asset.name)
453
+ rows.append({
454
+ "id": self._asset_public_id(asset),
455
+ "name": asset.name,
456
+ "kind": asset.kind,
457
+ "architecture": asset.trainer or ("stable_diffusion" if asset.kind == "base_model" else ""),
458
+ "trainer": asset.trainer,
459
+ "checkpoint_name": Path(asset.checkpoint or asset.path).name,
460
+ "dataset": None if dataset is None else {"id": self._asset_public_id(dataset), "name": dataset.name},
461
+ "epochs": asset.epochs,
462
+ "trigger_word": trigger_word,
463
+ "latest_experiment": None if latest is None else {
464
+ "id": latest.id,
465
+ "status": latest.status,
466
+ "resolution": latest.resolution,
467
+ "dataset_name": latest.dataset_name,
468
+ "trigger_word": getattr(latest, "trigger_word", ""),
469
+ },
470
+ })
471
+ self._model_catalog = rows[:300]
472
+ self._model_catalog_at = time.monotonic()
473
+ return list(self._model_catalog)
474
 
475
  def training_schema(self) -> dict[str, Any]:
476
  registry = self._registry()
 
654
  "options": tool.generation_options,
655
  "settings": registry.model_plugins.generation_schema_for_tool(tool.id),
656
  })
657
+ models = self.models()
658
+ return {
659
+ "providers": providers,
660
+ "models": [item for item in models if item["kind"] == "model"],
661
+ "base_models": [item for item in models if item["kind"] == "base_model"],
662
+ }
663
 
664
  def start_generation(self, payload: dict[str, Any]) -> dict[str, Any]:
665
  if self.jobs is None:
adam/tool_folders.py CHANGED
@@ -26,6 +26,10 @@ class ToolFolderStatus:
26
 
27
 
28
  TOOL_FOLDER_DEFINITIONS = (
 
 
 
 
29
  ToolFolderDefinition(
30
  "dataset_collector",
31
  "Dataset Collector",
@@ -155,6 +159,7 @@ class ToolFolderManager:
155
  def parse_assignments(self, text: str) -> dict[str, str]:
156
  """Recognize folder assignments pasted into chat without executing them."""
157
  patterns = {
 
158
  "ddpm_trainer": r"(?im)^\s*DDPM(?:\s+Trainer)?\s*:\s*(.+?)\s*$",
159
  "flow_trainer": (
160
  r"(?im)^\s*Flow(?:\s+Matching)?(?:\s+Trainer)?\s*:\s*(.+?)\s*$"
 
26
 
27
 
28
  TOOL_FOLDER_DEFINITIONS = (
29
+ ToolFolderDefinition(
30
+ "wan_video_trainer", "Wan Video LoRA Trainer",
31
+ ("trainer-engine/src/musubi_tuner/wan_train_network.py",),
32
+ ),
33
  ToolFolderDefinition(
34
  "dataset_collector",
35
  "Dataset Collector",
 
159
  def parse_assignments(self, text: str) -> dict[str, str]:
160
  """Recognize folder assignments pasted into chat without executing them."""
161
  patterns = {
162
+ "wan_video_trainer": r"(?im)^\s*(?:Wan(?:\s+Video)?(?:\s+LoRA)?(?:\s+Trainer)?|LoRA\s*Video\s*Trainer)\s*:\s*(.+?)\s*$",
163
  "ddpm_trainer": r"(?im)^\s*DDPM(?:\s+Trainer)?\s*:\s*(.+?)\s*$",
164
  "flow_trainer": (
165
  r"(?im)^\s*Flow(?:\s+Matching)?(?:\s+Trainer)?\s*:\s*(.+?)\s*$"
adam/tools/ddpm_adapter.py CHANGED
@@ -7,6 +7,8 @@ import importlib.util
7
  import math
8
  import queue
9
  import re
 
 
10
  import subprocess
11
  import sys
12
  import threading
@@ -15,6 +17,7 @@ from pathlib import Path
15
 
16
  from adam.config import ConfigManager
17
  from adam.executor import ToolAdjustmentRequested, ToolCancelled, ToolContext, ToolExecutionError
 
18
  from adam.process_control import set_process_tree_paused, terminate_process_tree
19
 
20
 
@@ -22,7 +25,34 @@ IMAGE_EXTENSIONS = {".jpg", ".jpeg", ".png", ".webp", ".bmp"}
22
  FORCE_STOP_TIMEOUT_SECONDS = 30
23
 
24
 
25
- def _saved_unet_resolution(path: Path) -> int | None:
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
26
  config_path = path / "unet" / "config.json"
27
  if not config_path.is_file():
28
  return None
@@ -31,14 +61,57 @@ def _saved_unet_resolution(path: Path) -> int | None:
31
  except (OSError, json.JSONDecodeError):
32
  return None
33
  sample_size = data.get("sample_size")
34
- if isinstance(sample_size, list):
35
- sample_size = sample_size[0] if sample_size else None
36
  try:
37
- return int(sample_size) if sample_size is not None else None
 
 
 
 
 
38
  except (TypeError, ValueError):
39
  return None
40
 
41
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
42
  def _latest_preview(folder: Path) -> Path | None:
43
  try:
44
  images = [
@@ -99,7 +172,7 @@ def _parse_progress(
99
  context.progress(100, "DDPM training completed")
100
 
101
 
102
- def train_ddpm(
103
  context: ToolContext,
104
  dataset_dir: str,
105
  model_name: str,
@@ -112,6 +185,7 @@ def train_ddpm(
112
  training_intensity: int = 100, preview_enabled: bool = True,
113
  preview_every: int = 5, preview_prompt: str = "", preview_seed: int = 123456789,
114
  completed_epochs: int = 0,
 
115
  ) -> dict[str, object]:
116
  """Run the registered DDPM project without shell interpolation or overwrites."""
117
  trainer_root = Path(str(ConfigManager(context.root).get("tool_folders", {}).get("ddpm_trainer", ""))).expanduser()
@@ -123,12 +197,15 @@ def train_ddpm(
123
  raise ToolExecutionError("DDPM train.py was not found. Re-scan the DDPM folder in Settings.")
124
  if not dataset.is_dir():
125
  raise ToolExecutionError("The selected DDPM dataset folder no longer exists.")
126
- image_count = sum(1 for item in dataset.iterdir() if item.is_file() and item.suffix.lower() in IMAGE_EXTENSIONS)
127
  if image_count < 2:
128
  raise ToolExecutionError("The DDPM dataset needs at least two image files before training can start.")
129
  requested_epochs = int(epochs)
130
  if not 64 <= int(resolution) <= 512 or int(resolution) % 8 or not 1 <= int(batch_size) <= 64:
131
  raise ToolExecutionError("DDPM resolution must be a multiple of 8 (64–512) and batch size 1–64.")
 
 
 
132
  if not 1e-7 <= float(learning_rate) <= 0.1 or not 1 <= int(gradient_accumulation_steps) <= 64:
133
  raise ToolExecutionError("DDPM learning rate or gradient accumulation is outside ADAM's safe range.")
134
  if not 0 <= int(dataloader_num_workers) <= 16 or not 1 <= int(save_every) <= 1000 or not 1 <= int(preview_steps) <= 500 or not 1 <= int(preview_every) <= 100_000 or not 10 <= int(training_intensity) <= 100 or mixed_precision not in {"fp16", "no"}:
@@ -170,34 +247,27 @@ def train_ddpm(
170
  and (resume / "scheduler.bin").is_file()
171
  )
172
  standalone_checkpoint = (resume / "pytorch_model.bin").is_file() or (resume / "model.safetensors").is_file()
 
173
  if not accelerate_checkpoint and not standalone_checkpoint:
174
- if not (output / "model_index.json").is_file():
175
  raise ToolExecutionError("The saved DDPM model is incomplete and cannot be fine-tuned safely.")
176
- pretrained_model = output
177
- timestamp = time.strftime("%Y%m%d_%H%M%S")
178
- output = output.with_name(f"{output.name}_finetuned_{timestamp}")
179
  resume = None
180
  context.log(
181
  "The exact resume checkpoint is incomplete. Creating a new fine-tuned model "
182
  "from the saved DDPM pipeline instead."
183
  )
184
  else:
185
- if not output.is_dir():
186
- raise ToolExecutionError("The DDPM model folder for resume no longer exists.")
187
- try:
188
- resume.relative_to(output)
189
- except ValueError as exc:
190
- raise ToolExecutionError("The DDPM checkpoint must be inside its model folder.") from exc
191
  if not resume.is_dir() or not resume.name.startswith("checkpoint-"):
192
  raise ToolExecutionError("A valid DDPM checkpoint-* folder is required to resume.")
193
- checkpoint_resolution = _saved_unet_resolution(resume)
194
- if checkpoint_resolution and checkpoint_resolution != int(resolution) and (output / "model_index.json").is_file():
195
- pretrained_model = output
196
- timestamp = time.strftime("%Y%m%d_%H%M%S")
197
- output = output.with_name(f"{output.name}_finetuned_{timestamp}")
198
  resume = None
199
  context.log(
200
- f"Changing DDPM resolution from {checkpoint_resolution}px to {int(resolution)}px. "
 
201
  "Starting a fresh fine-tune from the saved model weights instead of resuming the old optimizer schedule."
202
  )
203
  if resume is not None:
@@ -209,18 +279,25 @@ def train_ddpm(
209
  f"Continuing after approximately {prior_epochs} completed epochs "
210
  f"for {int(epochs) - prior_epochs} additional epochs."
211
  )
212
- if not resume:
213
- if output.exists():
214
- raise ToolExecutionError("The chosen DDPM output folder already exists; ADAM will not overwrite it.")
215
- output.mkdir(parents=True, exist_ok=False)
 
 
 
 
 
 
216
  # The connected trainer asks Accelerate/TensorBoard to write directly to
217
  # output/logs/train. Create it up front because its writer does not always
218
  # create the nested directory on Windows.
219
  (output / "logs" / "train").mkdir(parents=True, exist_ok=True)
220
  stop_file = output / ".adam_stop_training.flag"
221
  command = [
222
- sys.executable, str(script), "--train_data_dir", str(dataset), "--output_dir", str(output),
223
  "--model_name", model_name, "--resolution", str(int(resolution)), "--train_batch_size", str(int(batch_size)),
 
224
  "--num_epochs", str(int(epochs)), "--learning_rate", str(float(learning_rate)), "--mixed_precision", mixed_precision,
225
  "--ddpm_beta_schedule", "linear", "--tf32", "true", "--save_images_epochs", str(int(preview_every) if preview_enabled else int(epochs) + 1),
226
  "--save_model_epochs", str(int(save_every)), "--training_intensity", str(int(training_intensity)), "--dataloader_num_workers", str(int(dataloader_num_workers)),
@@ -235,7 +312,10 @@ def train_ddpm(
235
  command.extend(["--resume_completed_epochs", str(int(completed_epochs))])
236
  if pretrained_model:
237
  command.extend(["--pretrained_model_path", str(pretrained_model)])
238
- context.log(f"Starting real DDPM training with {image_count} images at {resolution}px, batch {batch_size}, lr {learning_rate}.")
 
 
 
239
  context.log(f"Output folder: {output}")
240
  process = subprocess.Popen(command, cwd=str(trainer_root), stdout=subprocess.PIPE, stderr=subprocess.STDOUT,
241
  text=True, encoding="utf-8", errors="replace", shell=False)
@@ -339,3 +419,139 @@ def train_ddpm(
339
  }
340
  ],
341
  }
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
7
  import math
8
  import queue
9
  import re
10
+ import shutil
11
+ import statistics
12
  import subprocess
13
  import sys
14
  import threading
 
17
 
18
  from adam.config import ConfigManager
19
  from adam.executor import ToolAdjustmentRequested, ToolCancelled, ToolContext, ToolExecutionError
20
+ from adam.progressive_training import parse_stages, stage_batch_settings, stage_summary
21
  from adam.process_control import set_process_tree_paused, terminate_process_tree
22
 
23
 
 
25
  FORCE_STOP_TIMEOUT_SECONDS = 30
26
 
27
 
28
+ def _image_files(folder: Path):
29
+ """Yield training images below a dataset folder without loading them into memory."""
30
+ try:
31
+ yield from (
32
+ path for path in folder.rglob("*")
33
+ if path.is_file() and path.suffix.casefold() in IMAGE_EXTENSIONS
34
+ )
35
+ except OSError:
36
+ return
37
+
38
+
39
+ def _training_image_folder(dataset: Path) -> tuple[Path, int]:
40
+ """Choose the accepted-frame tree for video datasets, or the dataset itself.
41
+
42
+ YouTube collections retain their provenance by storing accepted frames in
43
+ ``frames/<source>/`` and rejected candidates separately. Passing the root
44
+ to a recursive trainer would include rejected images, while only checking
45
+ the root makes the collection appear empty.
46
+ """
47
+ frames = dataset / "frames"
48
+ if frames.is_dir():
49
+ frame_count = sum(1 for _ in _image_files(frames))
50
+ if frame_count:
51
+ return frames, frame_count
52
+ return dataset, sum(1 for _ in _image_files(dataset))
53
+
54
+
55
+ def _saved_unet_size(path: Path) -> tuple[int, int] | None:
56
  config_path = path / "unet" / "config.json"
57
  if not config_path.is_file():
58
  return None
 
61
  except (OSError, json.JSONDecodeError):
62
  return None
63
  sample_size = data.get("sample_size")
 
 
64
  try:
65
+ if isinstance(sample_size, list) and len(sample_size) >= 2:
66
+ return int(sample_size[0]), int(sample_size[1])
67
+ if sample_size is not None:
68
+ size = int(sample_size)
69
+ return size, size
70
+ return None
71
  except (TypeError, ValueError):
72
  return None
73
 
74
 
75
+ def _snap_dimension(value: float, *, multiple: int = 16) -> int:
76
+ return max(64, min(512, int(round(value / multiple)) * multiple))
77
+
78
+
79
+ def _dataset_aspect_ratio(dataset: Path) -> float:
80
+ from PIL import Image
81
+
82
+ ratios: list[float] = []
83
+ for path in _image_files(dataset):
84
+ try:
85
+ with Image.open(path) as image:
86
+ if image.width > 0 and image.height > 0:
87
+ ratios.append(image.width / image.height)
88
+ except OSError:
89
+ continue
90
+ return statistics.median(ratios) if ratios else 1.0
91
+
92
+
93
+ def _training_canvas(dataset: Path, resolution: int, aspect_ratio: str) -> tuple[int, int]:
94
+ ratios = {
95
+ "1:1 (Square)": 1.0,
96
+ "16:9 (Widescreen)": 16 / 9,
97
+ "9:16 (Portrait)": 9 / 16,
98
+ "4:3 (Classic)": 4 / 3,
99
+ "3:4 (Portrait Classic)": 3 / 4,
100
+ "3:2 (Photo)": 3 / 2,
101
+ "2:3 (Portrait Photo)": 2 / 3,
102
+ }
103
+ ratio = _dataset_aspect_ratio(dataset) if aspect_ratio == "Dataset (Auto)" else ratios.get(aspect_ratio)
104
+ if ratio is None or ratio <= 0:
105
+ raise ToolExecutionError("Choose a supported DDPM training aspect ratio.")
106
+ if abs(ratio - 1.0) < 0.01:
107
+ return resolution, resolution
108
+ if ratio > 1:
109
+ width, height = resolution, _snap_dimension(resolution / ratio)
110
+ else:
111
+ width, height = _snap_dimension(resolution * ratio), resolution
112
+ return width, height
113
+
114
+
115
  def _latest_preview(folder: Path) -> Path | None:
116
  try:
117
  images = [
 
172
  context.progress(100, "DDPM training completed")
173
 
174
 
175
+ def _train_ddpm_stage(
176
  context: ToolContext,
177
  dataset_dir: str,
178
  model_name: str,
 
185
  training_intensity: int = 100, preview_enabled: bool = True,
186
  preview_every: int = 5, preview_prompt: str = "", preview_seed: int = 123456789,
187
  completed_epochs: int = 0,
188
+ training_aspect_ratio: str = "Dataset (Auto)", resize_mode: str = "fit",
189
  ) -> dict[str, object]:
190
  """Run the registered DDPM project without shell interpolation or overwrites."""
191
  trainer_root = Path(str(ConfigManager(context.root).get("tool_folders", {}).get("ddpm_trainer", ""))).expanduser()
 
197
  raise ToolExecutionError("DDPM train.py was not found. Re-scan the DDPM folder in Settings.")
198
  if not dataset.is_dir():
199
  raise ToolExecutionError("The selected DDPM dataset folder no longer exists.")
200
+ training_dataset, image_count = _training_image_folder(dataset)
201
  if image_count < 2:
202
  raise ToolExecutionError("The DDPM dataset needs at least two image files before training can start.")
203
  requested_epochs = int(epochs)
204
  if not 64 <= int(resolution) <= 512 or int(resolution) % 8 or not 1 <= int(batch_size) <= 64:
205
  raise ToolExecutionError("DDPM resolution must be a multiple of 8 (64–512) and batch size 1–64.")
206
+ canvas_width, canvas_height = _training_canvas(training_dataset, int(resolution), str(training_aspect_ratio))
207
+ if resize_mode not in {"fit", "fill", "stretch"}:
208
+ raise ToolExecutionError("DDPM resize mode must be fit, fill, or stretch.")
209
  if not 1e-7 <= float(learning_rate) <= 0.1 or not 1 <= int(gradient_accumulation_steps) <= 64:
210
  raise ToolExecutionError("DDPM learning rate or gradient accumulation is outside ADAM's safe range.")
211
  if not 0 <= int(dataloader_num_workers) <= 16 or not 1 <= int(save_every) <= 1000 or not 1 <= int(preview_steps) <= 500 or not 1 <= int(preview_every) <= 100_000 or not 10 <= int(training_intensity) <= 100 or mixed_precision not in {"fp16", "no"}:
 
247
  and (resume / "scheduler.bin").is_file()
248
  )
249
  standalone_checkpoint = (resume / "pytorch_model.bin").is_file() or (resume / "model.safetensors").is_file()
250
+ source_model = resume.parent if resume.name.startswith("checkpoint-") else resume
251
  if not accelerate_checkpoint and not standalone_checkpoint:
252
+ if not (source_model / "model_index.json").is_file():
253
  raise ToolExecutionError("The saved DDPM model is incomplete and cannot be fine-tuned safely.")
254
+ pretrained_model = source_model
 
 
255
  resume = None
256
  context.log(
257
  "The exact resume checkpoint is incomplete. Creating a new fine-tuned model "
258
  "from the saved DDPM pipeline instead."
259
  )
260
  else:
 
 
 
 
 
 
261
  if not resume.is_dir() or not resume.name.startswith("checkpoint-"):
262
  raise ToolExecutionError("A valid DDPM checkpoint-* folder is required to resume.")
263
+ checkpoint_size = _saved_unet_size(resume)
264
+ target_size = (canvas_height, canvas_width)
265
+ if checkpoint_size and checkpoint_size != target_size and (source_model / "model_index.json").is_file():
266
+ pretrained_model = source_model
 
267
  resume = None
268
  context.log(
269
+ f"Changing DDPM canvas from {checkpoint_size[1]}x{checkpoint_size[0]} to "
270
+ f"{canvas_width}x{canvas_height}. "
271
  "Starting a fresh fine-tune from the saved model weights instead of resuming the old optimizer schedule."
272
  )
273
  if resume is not None:
 
279
  f"Continuing after approximately {prior_epochs} completed epochs "
280
  f"for {int(epochs) - prior_epochs} additional epochs."
281
  )
282
+ if output.exists():
283
+ raise ToolExecutionError("The chosen DDPM output folder already exists; ADAM will not overwrite it.")
284
+ output.mkdir(parents=True, exist_ok=False)
285
+ if resume:
286
+ # The connected trainer resolves --resume_from_checkpoint inside its
287
+ # output directory. Copy only the checkpoint into this new branch so
288
+ # it can restore optimizer state without touching the source model.
289
+ resume_copy = output / resume.name
290
+ shutil.copytree(resume, resume_copy)
291
+ resume = resume_copy
292
  # The connected trainer asks Accelerate/TensorBoard to write directly to
293
  # output/logs/train. Create it up front because its writer does not always
294
  # create the nested directory on Windows.
295
  (output / "logs" / "train").mkdir(parents=True, exist_ok=True)
296
  stop_file = output / ".adam_stop_training.flag"
297
  command = [
298
+ sys.executable, str(script), "--train_data_dir", str(training_dataset), "--output_dir", str(output),
299
  "--model_name", model_name, "--resolution", str(int(resolution)), "--train_batch_size", str(int(batch_size)),
300
+ "--resolution_width", str(canvas_width), "--resolution_height", str(canvas_height), "--resize_mode", resize_mode,
301
  "--num_epochs", str(int(epochs)), "--learning_rate", str(float(learning_rate)), "--mixed_precision", mixed_precision,
302
  "--ddpm_beta_schedule", "linear", "--tf32", "true", "--save_images_epochs", str(int(preview_every) if preview_enabled else int(epochs) + 1),
303
  "--save_model_epochs", str(int(save_every)), "--training_intensity", str(int(training_intensity)), "--dataloader_num_workers", str(int(dataloader_num_workers)),
 
312
  command.extend(["--resume_completed_epochs", str(int(completed_epochs))])
313
  if pretrained_model:
314
  command.extend(["--pretrained_model_path", str(pretrained_model)])
315
+ context.log(
316
+ f"Starting real DDPM training with {image_count} images from {training_dataset} on a {canvas_width}x{canvas_height} "
317
+ f"{resize_mode} canvas, batch {batch_size}, lr {learning_rate}."
318
+ )
319
  context.log(f"Output folder: {output}")
320
  process = subprocess.Popen(command, cwd=str(trainer_root), stdout=subprocess.PIPE, stderr=subprocess.STDOUT,
321
  text=True, encoding="utf-8", errors="replace", shell=False)
 
419
  }
420
  ],
421
  }
422
+
423
+
424
+ def _stage_context(context: ToolContext, *, stage_index: int, stage_count: int) -> ToolContext:
425
+ """Map one stage's progress into the single job progress bar."""
426
+ def report(percent: int, message: str, **details: object) -> None:
427
+ overall = round(((stage_index + max(0, min(percent, 100)) / 100) / stage_count) * 100)
428
+ context.progress(
429
+ overall,
430
+ f"Stage {stage_index + 1}/{stage_count} · {message}",
431
+ **details,
432
+ )
433
+
434
+ return ToolContext(
435
+ root=context.root, job_id=context.job_id, tool=context.tool,
436
+ cancel_event=context.cancel_event, run_event=context.run_event,
437
+ progress_callback=report, log_callback=context.log_callback,
438
+ preview_callback=context.preview_callback,
439
+ # A settings-change request currently restarts a single DDPM process.
440
+ # Keep a curriculum stage atomic until that recovery path understands
441
+ # its stage manifest.
442
+ adjustment_event=None, adjustment_request=None, step_delay=context.step_delay,
443
+ )
444
+
445
+
446
+ def _stage_output(root: Path, stage_number: int, resolution: int) -> Path:
447
+ base = root / f"stage-{stage_number:02d}-{resolution}px"
448
+ if not base.exists():
449
+ return base
450
+ attempt = 2
451
+ while (candidate := root / f"{base.name}-retry-{attempt}").exists():
452
+ attempt += 1
453
+ return candidate
454
+
455
+
456
+ def train_ddpm(
457
+ context: ToolContext,
458
+ dataset_dir: str,
459
+ model_name: str,
460
+ epochs: int,
461
+ output_dir: str,
462
+ resume_from: str = "",
463
+ resolution: int = 128, batch_size: int = 1, learning_rate: float = 0.0001,
464
+ gradient_accumulation_steps: int = 1, dataloader_num_workers: int = 4,
465
+ mixed_precision: str = "fp16", save_every: int = 10, preview_steps: int = 50,
466
+ training_intensity: int = 100, preview_enabled: bool = True,
467
+ preview_every: int = 5, preview_prompt: str = "", preview_seed: int = 123456789,
468
+ completed_epochs: int = 0,
469
+ training_aspect_ratio: str = "Dataset (Auto)", resize_mode: str = "fit",
470
+ progressive_stages: list[dict[str, object]] | None = None,
471
+ progressive_auto_batch: bool = True,
472
+ ) -> dict[str, object]:
473
+ """Train once, or run a low-to-high resolution DDPM curriculum.
474
+
475
+ Every completed stage is retained in a hidden sibling folder. The public
476
+ output folder is created only after the final pipeline has completed, so a
477
+ partial curriculum can never replace a usable completed model.
478
+ """
479
+ if not progressive_stages:
480
+ return _train_ddpm_stage(
481
+ context, dataset_dir, model_name, epochs, output_dir, resume_from,
482
+ resolution, batch_size, learning_rate, gradient_accumulation_steps,
483
+ dataloader_num_workers, mixed_precision, save_every, preview_steps,
484
+ training_intensity, preview_enabled, preview_every, preview_prompt,
485
+ preview_seed, completed_epochs, training_aspect_ratio, resize_mode,
486
+ )
487
+ stages = parse_stages(progressive_stages, trainer="ddpm", total_epochs=epochs)
488
+ public_output = Path(output_dir).expanduser().resolve()
489
+ if public_output.exists():
490
+ raise ToolExecutionError("The chosen DDPM output folder already exists; ADAM will not overwrite it.")
491
+ stage_root = public_output.parent / f".{public_output.name}.progressive"
492
+ state_path = stage_root / "progressive_state.json"
493
+ stage_root.mkdir(parents=True, exist_ok=True)
494
+ try:
495
+ state = json.loads(state_path.read_text(encoding="utf-8"))
496
+ except (OSError, json.JSONDecodeError):
497
+ state = {"model_name": model_name, "stages": [], "completed": []}
498
+ completed = state.get("completed", []) if isinstance(state.get("completed"), list) else []
499
+ completed_by_index = {
500
+ int(item.get("index")): Path(str(item.get("output")))
501
+ for item in completed if isinstance(item, dict) and str(item.get("index", "")).isdigit()
502
+ }
503
+ final_resolution = stages[-1].resolution
504
+ prior_model = resume_from
505
+ final_result: dict[str, object] | None = None
506
+ context.log(
507
+ "Progressive DDPM schedule: " + stage_summary(stages) + ". "
508
+ + ("Auto batch caps are enabled." if progressive_auto_batch else "Using the same batch settings at every stage.")
509
+ )
510
+ for index, stage in enumerate(stages):
511
+ completed_output = completed_by_index.get(index)
512
+ if completed_output and completed_output.is_dir():
513
+ prior_model = str(completed_output)
514
+ context.log(f"Stage {index + 1}/{len(stages)} already completed; using its saved weights.")
515
+ continue
516
+ stage_output = _stage_output(stage_root, index + 1, stage.resolution)
517
+ stage_batch, stage_accumulation = stage_batch_settings(
518
+ trainer="ddpm", stage_resolution=stage.resolution, final_resolution=final_resolution,
519
+ final_batch_size=batch_size, base_accumulation=gradient_accumulation_steps,
520
+ auto_batch=bool(progressive_auto_batch),
521
+ )
522
+ context.log(
523
+ f"Stage {index + 1}/{len(stages)}: {stage.resolution}px for {stage.epochs} epochs; "
524
+ f"batch {stage_batch}, gradient accumulation {stage_accumulation}."
525
+ )
526
+ result = _train_ddpm_stage(
527
+ _stage_context(context, stage_index=index, stage_count=len(stages)),
528
+ dataset_dir, model_name, stage.epochs, str(stage_output), prior_model,
529
+ stage.resolution, stage_batch, learning_rate, stage_accumulation,
530
+ dataloader_num_workers, mixed_precision, min(save_every, stage.epochs), preview_steps,
531
+ training_intensity, preview_enabled, min(preview_every, stage.epochs), preview_prompt,
532
+ preview_seed, 0, training_aspect_ratio, resize_mode,
533
+ )
534
+ prior_model = str(stage_output)
535
+ final_result = result
536
+ completed.append({"index": index, "resolution": stage.resolution, "epochs": stage.epochs, "output": prior_model})
537
+ state.update({"stages": [{"resolution": item.resolution, "epochs": item.epochs} for item in stages], "completed": completed})
538
+ state_path.write_text(json.dumps(state, indent=2), encoding="utf-8")
539
+ if not prior_model or not Path(prior_model).is_dir():
540
+ raise ToolExecutionError("Progressive DDPM training did not produce a final stage model.")
541
+ shutil.move(prior_model, public_output)
542
+ final_checkpoint = sorted(
543
+ public_output.glob("checkpoint-*"),
544
+ key=lambda path: int(path.name.rsplit("-", 1)[-1])
545
+ if path.name.rsplit("-", 1)[-1].isdigit() else -1,
546
+ )[-1:]
547
+ checkpoint = str(final_checkpoint[0]) if final_checkpoint else ""
548
+ context.progress(100, "Progressive DDPM training completed")
549
+ return {
550
+ "output_folder": str(public_output), "model_name": model_name,
551
+ "progressive_stages": [{"resolution": stage.resolution, "epochs": stage.epochs} for stage in stages],
552
+ "assets": [{
553
+ "kind": "model", "name": model_name, "path": str(public_output), "trainer": "ddpm",
554
+ "dataset_path": str(Path(dataset_dir).expanduser().resolve()), "checkpoint": checkpoint,
555
+ "epochs": sum(stage.epochs for stage in stages),
556
+ }],
557
+ }
adam/tools/flow_adapter.py CHANGED
@@ -8,10 +8,12 @@ import subprocess
8
  import sys
9
  import threading
10
  import time
 
11
  from pathlib import Path
12
 
13
  from adam.config import ConfigManager
14
  from adam.executor import ToolCancelled, ToolContext, ToolExecutionError
 
15
  from adam.process_control import set_process_tree_paused, terminate_process_tree
16
 
17
 
@@ -29,7 +31,7 @@ def _latest_preview(folder: Path) -> Path | None:
29
  return None
30
 
31
 
32
- def train_flow(
33
  context: ToolContext,
34
  dataset_dir: str,
35
  model_name: str,
@@ -76,9 +78,11 @@ def train_flow(
76
  saved_resolution = int(metadata.get("resolution", 0) or 0)
77
  except (OSError, ValueError, TypeError, json.JSONDecodeError) as exc:
78
  raise ToolExecutionError("Choose a valid completed Flow Matching model to continue.") from exc
79
- if saved_resolution != int(resolution):
80
- raise ToolExecutionError(
81
- f"The selected Flow model is {saved_resolution}px; continuation must use the same resolution."
 
 
82
  )
83
  output.parent.mkdir(parents=True, exist_ok=True)
84
  command = [
@@ -159,3 +163,114 @@ def train_flow(
159
  "kind": "model", "name": safe_name, "path": str(output), "trainer": "flow",
160
  "dataset_path": str(dataset), "checkpoint": str(output), "epochs": int(epochs),
161
  }]}
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
8
  import sys
9
  import threading
10
  import time
11
+ import shutil
12
  from pathlib import Path
13
 
14
  from adam.config import ConfigManager
15
  from adam.executor import ToolCancelled, ToolContext, ToolExecutionError
16
+ from adam.progressive_training import parse_stages, stage_batch_settings, stage_summary
17
  from adam.process_control import set_process_tree_paused, terminate_process_tree
18
 
19
 
 
31
  return None
32
 
33
 
34
+ def _train_flow_stage(
35
  context: ToolContext,
36
  dataset_dir: str,
37
  model_name: str,
 
78
  saved_resolution = int(metadata.get("resolution", 0) or 0)
79
  except (OSError, ValueError, TypeError, json.JSONDecodeError) as exc:
80
  raise ToolExecutionError("Choose a valid completed Flow Matching model to continue.") from exc
81
+ if saved_resolution and saved_resolution != int(resolution):
82
+ context.log(
83
+ "Resolution-change fine-tune: loading "
84
+ f"{saved_resolution}px Flow weights for training at {int(resolution)}px. "
85
+ "Optimizer state will start fresh."
86
  )
87
  output.parent.mkdir(parents=True, exist_ok=True)
88
  command = [
 
163
  "kind": "model", "name": safe_name, "path": str(output), "trainer": "flow",
164
  "dataset_path": str(dataset), "checkpoint": str(output), "epochs": int(epochs),
165
  }]}
166
+
167
+
168
+ def _stage_context(context: ToolContext, *, stage_index: int, stage_count: int) -> ToolContext:
169
+ def report(percent: int, message: str, **details: object) -> None:
170
+ overall = round(((stage_index + max(0, min(percent, 100)) / 100) / stage_count) * 100)
171
+ context.progress(overall, f"Stage {stage_index + 1}/{stage_count} · {message}", **details)
172
+
173
+ return ToolContext(
174
+ root=context.root, job_id=context.job_id, tool=context.tool,
175
+ cancel_event=context.cancel_event, run_event=context.run_event,
176
+ progress_callback=report, log_callback=context.log_callback,
177
+ preview_callback=context.preview_callback, step_delay=context.step_delay,
178
+ )
179
+
180
+
181
+ def _stage_output(root: Path, stage_number: int, resolution: int) -> Path:
182
+ base = root / f"stage-{stage_number:02d}-{resolution}px"
183
+ if not base.exists():
184
+ return base
185
+ attempt = 2
186
+ while (candidate := root / f"{base.name}-retry-{attempt}").exists():
187
+ attempt += 1
188
+ return candidate
189
+
190
+
191
+ def train_flow(
192
+ context: ToolContext,
193
+ dataset_dir: str,
194
+ model_name: str,
195
+ epochs: int,
196
+ output_dir: str,
197
+ resume_from: str = "",
198
+ resolution: int = 256, batch_size: int = 8, learning_rate: float = 0.0002,
199
+ gradient_accumulation: int = 1, workers: int = 4, mixed_precision: str = "fp16",
200
+ save_every: int = 10, preview_every: int = 10, preview_steps: int = 10,
201
+ gradient_checkpointing: bool = False, preview_enabled: bool = True,
202
+ preview_prompt: str = "", preview_seed: int = 123456789,
203
+ progressive_stages: list[dict[str, object]] | None = None,
204
+ progressive_auto_batch: bool = True,
205
+ ) -> dict[str, object]:
206
+ """Train one Flow model or carry it through a saved resolution curriculum."""
207
+ if not progressive_stages:
208
+ return _train_flow_stage(
209
+ context, dataset_dir, model_name, epochs, output_dir, resume_from, resolution,
210
+ batch_size, learning_rate, gradient_accumulation, workers, mixed_precision,
211
+ save_every, preview_every, preview_steps, gradient_checkpointing,
212
+ preview_enabled, preview_prompt, preview_seed,
213
+ )
214
+ stages = parse_stages(progressive_stages, trainer="flow", total_epochs=epochs)
215
+ public_output = Path(output_dir).expanduser().resolve()
216
+ if public_output.exists():
217
+ raise ToolExecutionError("The chosen Flow Matching output folder already exists; ADAM will not overwrite it.")
218
+ stage_root = public_output.parent / f".{public_output.name}.progressive"
219
+ state_path = stage_root / "progressive_state.json"
220
+ stage_root.mkdir(parents=True, exist_ok=True)
221
+ try:
222
+ state = json.loads(state_path.read_text(encoding="utf-8"))
223
+ except (OSError, json.JSONDecodeError):
224
+ state = {"model_name": model_name, "stages": [], "completed": []}
225
+ completed = state.get("completed", []) if isinstance(state.get("completed"), list) else []
226
+ completed_by_index = {
227
+ int(item.get("index")): Path(str(item.get("output")))
228
+ for item in completed if isinstance(item, dict) and str(item.get("index", "")).isdigit()
229
+ }
230
+ final_resolution = stages[-1].resolution
231
+ prior_model = resume_from
232
+ context.log(
233
+ "Progressive Flow Matching schedule: " + stage_summary(stages) + ". "
234
+ + ("Auto batch caps are enabled." if progressive_auto_batch else "Using the same batch settings at every stage.")
235
+ )
236
+ for index, stage in enumerate(stages):
237
+ completed_output = completed_by_index.get(index)
238
+ if completed_output and completed_output.is_dir():
239
+ prior_model = str(completed_output)
240
+ context.log(f"Stage {index + 1}/{len(stages)} already completed; using its saved weights.")
241
+ continue
242
+ stage_output = _stage_output(stage_root, index + 1, stage.resolution)
243
+ stage_batch, stage_accumulation = stage_batch_settings(
244
+ trainer="flow", stage_resolution=stage.resolution, final_resolution=final_resolution,
245
+ final_batch_size=batch_size, base_accumulation=gradient_accumulation,
246
+ auto_batch=bool(progressive_auto_batch),
247
+ )
248
+ context.log(
249
+ f"Stage {index + 1}/{len(stages)}: {stage.resolution}px for {stage.epochs} epochs; "
250
+ f"batch {stage_batch}, gradient accumulation {stage_accumulation}."
251
+ )
252
+ _train_flow_stage(
253
+ _stage_context(context, stage_index=index, stage_count=len(stages)),
254
+ dataset_dir, model_name, stage.epochs, str(stage_output), prior_model,
255
+ stage.resolution, stage_batch, learning_rate, stage_accumulation, workers,
256
+ mixed_precision, min(save_every, stage.epochs), min(preview_every, stage.epochs),
257
+ preview_steps, gradient_checkpointing or stage.resolution >= 384,
258
+ preview_enabled, preview_prompt, preview_seed,
259
+ )
260
+ prior_model = str(stage_output)
261
+ completed.append({"index": index, "resolution": stage.resolution, "epochs": stage.epochs, "output": prior_model})
262
+ state.update({"stages": [{"resolution": item.resolution, "epochs": item.epochs} for item in stages], "completed": completed})
263
+ state_path.write_text(json.dumps(state, indent=2), encoding="utf-8")
264
+ if not prior_model or not Path(prior_model).is_dir():
265
+ raise ToolExecutionError("Progressive Flow Matching training did not produce a final stage model.")
266
+ shutil.move(prior_model, public_output)
267
+ context.progress(100, "Progressive Flow Matching training completed")
268
+ return {
269
+ "output_folder": str(public_output), "model_name": model_name,
270
+ "progressive_stages": [{"resolution": stage.resolution, "epochs": stage.epochs} for stage in stages],
271
+ "assets": [{
272
+ "kind": "model", "name": model_name, "path": str(public_output), "trainer": "flow",
273
+ "dataset_path": str(Path(dataset_dir).expanduser().resolve()), "checkpoint": str(public_output),
274
+ "epochs": sum(stage.epochs for stage in stages),
275
+ }],
276
+ }
adam/tools/flow_generator.py CHANGED
@@ -67,6 +67,8 @@ def generate_flow_images(
67
  seed: int,
68
  sampler: str,
69
  aspect_ratio: str,
 
 
70
  preview_interval: int = 0,
71
  smart_generation: bool = False,
72
  smart_wanted_results: int = 0,
@@ -142,6 +144,19 @@ def generate_flow_images(
142
  }
143
  if aspect_ratio not in allowed_aspects:
144
  raise ToolExecutionError("Choose one of the supported Flow aspect ratios.")
 
 
 
 
 
 
 
 
 
 
 
 
 
145
  if len(prompt) > 500:
146
  raise ToolExecutionError("The generation label must be 500 characters or shorter.")
147
  if not 0 <= int(preview_interval) <= step_count:
@@ -199,6 +214,8 @@ def generate_flow_images(
199
 
200
  try:
201
  settings = {"aspect_ratio": aspect_ratio}
 
 
202
  if preview_enabled and preview_supported:
203
  settings["preview_interval"] = int(preview_interval)
204
  settings["preview_callback"] = lambda payload, step=0, total_steps=step_count, current=index: publish_generation_preview(
@@ -262,6 +279,8 @@ def generate_flow_images(
262
  "steps": step_count,
263
  "sampler": method,
264
  "aspect_ratio": aspect_ratio,
 
 
265
  "preview_interval": int(preview_interval),
266
  "preview_supported": preview_supported,
267
  "images": ordered_images if smart_keep_rejected or not smart_enabled else selected_paths,
 
67
  seed: int,
68
  sampler: str,
69
  aspect_ratio: str,
70
+ width: int = 0,
71
+ height: int = 0,
72
  preview_interval: int = 0,
73
  smart_generation: bool = False,
74
  smart_wanted_results: int = 0,
 
144
  }
145
  if aspect_ratio not in allowed_aspects:
146
  raise ToolExecutionError("Choose one of the supported Flow aspect ratios.")
147
+ custom_width = int(width or 0)
148
+ custom_height = int(height or 0)
149
+ if bool(custom_width) != bool(custom_height):
150
+ raise ToolExecutionError("Set both Flow image width and height, or leave both unset.")
151
+ if custom_width and (
152
+ not 64 <= custom_width <= 2048
153
+ or not 64 <= custom_height <= 2048
154
+ or custom_width % 16
155
+ or custom_height % 16
156
+ ):
157
+ raise ToolExecutionError(
158
+ "Flow image width and height must each be between 64 and 2048 pixels and divisible by 16."
159
+ )
160
  if len(prompt) > 500:
161
  raise ToolExecutionError("The generation label must be 500 characters or shorter.")
162
  if not 0 <= int(preview_interval) <= step_count:
 
214
 
215
  try:
216
  settings = {"aspect_ratio": aspect_ratio}
217
+ if custom_width:
218
+ settings.update({"width": custom_width, "height": custom_height})
219
  if preview_enabled and preview_supported:
220
  settings["preview_interval"] = int(preview_interval)
221
  settings["preview_callback"] = lambda payload, step=0, total_steps=step_count, current=index: publish_generation_preview(
 
279
  "steps": step_count,
280
  "sampler": method,
281
  "aspect_ratio": aspect_ratio,
282
+ "width": custom_width,
283
+ "height": custom_height,
284
  "preview_interval": int(preview_interval),
285
  "preview_supported": preview_supported,
286
  "images": ordered_images if smart_keep_rejected or not smart_enabled else selected_paths,
adam/tools/lora_adapter.py CHANGED
@@ -257,7 +257,14 @@ def train_lora(
257
  "trigger_word": trigger,
258
  "epochs": int(epochs),
259
  "settings": settings,
260
- "training_overrides": training_overrides,
 
 
 
 
 
 
 
261
  }
262
  mp_context = multiprocessing.get_context("spawn")
263
  events = mp_context.Queue()
 
257
  "trigger_word": trigger,
258
  "epochs": int(epochs),
259
  "settings": settings,
260
+ # ``preview_prompt`` is a named argument, so Python removes it from
261
+ # ``training_overrides``. Explicitly include it here so a new ADAM
262
+ # run cannot inherit the connected trainer's last saved prompt.
263
+ "training_overrides": {
264
+ **training_overrides,
265
+ "preview_prompt": str(preview_prompt),
266
+ "preview_interval_epochs": max(1, int(preview_every)),
267
+ },
268
  }
269
  mp_context = multiprocessing.get_context("spawn")
270
  events = mp_context.Queue()
adam/tools/lora_generator.py CHANGED
@@ -3,11 +3,14 @@
3
  from __future__ import annotations
4
 
5
  import json
 
 
6
  import random
7
  import re
8
  import sys
9
  from datetime import datetime, timezone
10
  from pathlib import Path
 
11
 
12
  from adam.config import ConfigManager
13
  from adam.executor import ToolCancelled, ToolContext, ToolExecutionError
@@ -26,6 +29,114 @@ ASPECT_SIZES = {
26
  }
27
 
28
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
29
  def _safe_label(value: str) -> str:
30
  label = re.sub(r'[<>:"/\\|?*\x00-\x1f]+', " ", value.strip()[:96] or "LoRA")
31
  return re.sub(r"\s+", " ", label).strip(" .") or "LoRA"
@@ -95,7 +206,7 @@ def generate_lora_images(
95
  raise ToolExecutionError("Image count must be between 1 and 48.")
96
  if not 1 <= step_count <= 150:
97
  raise ToolExecutionError("LoRA generation steps must be between 1 and 150.")
98
- if sampler not in {"DPM++ 2M", "DPM++ SDE", "Euler", "Euler a", "DDIM"}:
99
  raise ToolExecutionError("Choose one of the supported LoRA samplers.")
100
  if aspect_ratio not in ASPECT_SIZES:
101
  raise ToolExecutionError("Choose one of the supported LoRA aspect ratios.")
@@ -111,7 +222,6 @@ def generate_lora_images(
111
  sys.path.insert(0, str(source_root))
112
  try:
113
  from loratrainer.models.generation_config import GenerationConfig
114
- from loratrainer.trainer.diffusers_sdxl_generator import DiffusersSDXLGeneratorBackend
115
  except Exception as exc:
116
  raise ToolExecutionError(f"Could not load the LoRA Trainer generator: {exc}") from exc
117
 
@@ -144,21 +254,10 @@ def generate_lora_images(
144
  prompt_weighting=bool(prompt_weighting),
145
  )
146
 
147
- def progress(update) -> None:
148
- context.checkpoint()
149
- current = max(0, int(getattr(update, "current_image", 0)))
150
- context.progress(min(99, max(1, round(current * 100 / count))), str(getattr(update, "message", "Generating with LoRA")))
151
- if int(preview_interval) > 0:
152
- payload = getattr(update, "preview", None) or getattr(update, "preview_path", None)
153
- if payload:
154
- publish_generation_preview(
155
- context, output, payload, image_index=max(0, current - 1), image_count=count,
156
- step=int(getattr(update, "step", getattr(update, "current_step", 0)) or 0),
157
- total_steps=step_count,
158
- )
159
-
160
  try:
161
- paths = DiffusersSDXLGeneratorBackend().generate(request, progress)
 
 
162
  except ToolCancelled:
163
  raise
164
  except Exception as exc:
 
3
  from __future__ import annotations
4
 
5
  import json
6
+ import multiprocessing
7
+ import queue
8
  import random
9
  import re
10
  import sys
11
  from datetime import datetime, timezone
12
  from pathlib import Path
13
+ from typing import Any
14
 
15
  from adam.config import ConfigManager
16
  from adam.executor import ToolCancelled, ToolContext, ToolExecutionError
 
29
  }
30
 
31
 
32
+ def _run_lora_generation_worker(payload: dict[str, Any], events: Any) -> None:
33
+ """Run SDXL inference outside ADAM's long-lived desktop process.
34
+
35
+ Diffusers owns CUDA allocations below Python's normal object lifecycle.
36
+ Keeping each LoRA batch in a spawned worker gives queued requests the same
37
+ clean CUDA context as the connected LoRA Trainer application.
38
+ """
39
+ try:
40
+ source_root = str(payload["source_root"])
41
+ if source_root not in sys.path:
42
+ sys.path.insert(0, source_root)
43
+ from loratrainer.models.generation_config import GenerationConfig
44
+ from loratrainer.trainer.diffusers_sdxl_generator import DiffusersSDXLGeneratorBackend
45
+
46
+ request = GenerationConfig(
47
+ base_model_path=Path(str(payload["base_model_path"])),
48
+ output_dir=Path(str(payload["output_dir"])),
49
+ positive_prompt=str(payload["positive_prompt"]),
50
+ negative_prompt=str(payload["negative_prompt"]),
51
+ lora_path=Path(str(payload["lora_path"])) if payload.get("lora_path") else None,
52
+ reference_image=Path(str(payload["reference_image"])) if payload.get("reference_image") else None,
53
+ width=int(payload["width"]), height=int(payload["height"]),
54
+ steps=int(payload["steps"]), cfg_scale=float(payload["cfg_scale"]),
55
+ seed=int(payload["seed"]), sampler=str(payload["sampler"]),
56
+ batch_count=int(payload["batch_count"]), lora_strength=float(payload["lora_strength"]),
57
+ denoise_strength=float(payload["denoise_strength"]),
58
+ prompt_weighting=bool(payload["prompt_weighting"]),
59
+ )
60
+
61
+ def progress(update: Any) -> None:
62
+ events.put({
63
+ "type": "progress", "current_image": int(getattr(update, "current_image", 0) or 0),
64
+ "message": str(getattr(update, "message", "Generating with LoRA") or "Generating with LoRA"),
65
+ "preview": getattr(update, "preview", None) or getattr(update, "preview_path", None),
66
+ "step": int(getattr(update, "step", getattr(update, "current_step", 0)) or 0),
67
+ })
68
+
69
+ paths = DiffusersSDXLGeneratorBackend().generate(request, progress)
70
+ events.put({"type": "result", "paths": [str(path) for path in paths]})
71
+ except Exception as exc:
72
+ events.put({"type": "error", "message": str(exc), "exception": type(exc).__name__})
73
+
74
+
75
+ def _generate_in_isolated_process(
76
+ context: ToolContext, source_root: Path, request: Any, output: Path,
77
+ count: int, step_count: int, preview_interval: int,
78
+ ) -> list[Path]:
79
+ """Generate a batch in a short-lived CUDA process and relay its progress."""
80
+ payload = {
81
+ "source_root": str(source_root), "base_model_path": str(request.base_model_path),
82
+ "output_dir": str(request.output_dir), "positive_prompt": request.positive_prompt,
83
+ "negative_prompt": request.negative_prompt, "lora_path": str(request.lora_path or ""),
84
+ "reference_image": str(request.reference_image or ""), "width": request.width,
85
+ "height": request.height, "steps": request.steps, "cfg_scale": request.cfg_scale,
86
+ "seed": request.seed, "sampler": request.sampler, "batch_count": request.batch_count,
87
+ "lora_strength": request.lora_strength, "denoise_strength": request.denoise_strength,
88
+ "prompt_weighting": request.prompt_weighting,
89
+ }
90
+ worker_context = multiprocessing.get_context("spawn")
91
+ events = worker_context.Queue()
92
+ process = worker_context.Process(target=_run_lora_generation_worker, args=(payload, events))
93
+ process.start()
94
+ result: list[Path] | None = None
95
+ error = ""
96
+ try:
97
+ while True:
98
+ if context.cancel_event.is_set():
99
+ process.terminate()
100
+ process.join(timeout=3)
101
+ raise ToolCancelled("Job cancelled by user.")
102
+ try:
103
+ event = events.get(timeout=0.1)
104
+ except queue.Empty:
105
+ if not process.is_alive():
106
+ break
107
+ continue
108
+ if event.get("type") == "progress":
109
+ current = max(0, int(event.get("current_image", 0)))
110
+ context.progress(
111
+ min(99, max(1, round(current * 100 / count))),
112
+ str(event.get("message", "Generating with LoRA")),
113
+ )
114
+ preview = event.get("preview")
115
+ if int(preview_interval) > 0 and preview:
116
+ publish_generation_preview(
117
+ context, output, preview, image_index=max(0, current - 1), image_count=count,
118
+ step=int(event.get("step", 0)), total_steps=step_count,
119
+ )
120
+ elif event.get("type") == "result":
121
+ result = [Path(str(path)) for path in event.get("paths", [])]
122
+ elif event.get("type") == "error":
123
+ error = str(event.get("message") or "Unknown worker error")
124
+ process.join(timeout=3)
125
+ finally:
126
+ if process.is_alive():
127
+ process.terminate()
128
+ process.join(timeout=3)
129
+ events.close()
130
+ events.join_thread()
131
+ if error:
132
+ raise ToolExecutionError(f"LoRA generation failed: {error}")
133
+ if result is None:
134
+ raise ToolExecutionError(
135
+ f"LoRA generation worker exited without returning images (exit code {process.exitcode})."
136
+ )
137
+ return result
138
+
139
+
140
  def _safe_label(value: str) -> str:
141
  label = re.sub(r'[<>:"/\\|?*\x00-\x1f]+', " ", value.strip()[:96] or "LoRA")
142
  return re.sub(r"\s+", " ", label).strip(" .") or "LoRA"
 
206
  raise ToolExecutionError("Image count must be between 1 and 48.")
207
  if not 1 <= step_count <= 150:
208
  raise ToolExecutionError("LoRA generation steps must be between 1 and 150.")
209
+ if sampler not in {"DPM++ 2M", "DPM++ 2M Karras", "DPM++ 2M SDE", "DPM++ 2M SDE Karras", "DPM++ SDE", "DPM++ SDE Karras", "Euler", "Euler a", "Heun", "LMS", "DDIM"}:
210
  raise ToolExecutionError("Choose one of the supported LoRA samplers.")
211
  if aspect_ratio not in ASPECT_SIZES:
212
  raise ToolExecutionError("Choose one of the supported LoRA aspect ratios.")
 
222
  sys.path.insert(0, str(source_root))
223
  try:
224
  from loratrainer.models.generation_config import GenerationConfig
 
225
  except Exception as exc:
226
  raise ToolExecutionError(f"Could not load the LoRA Trainer generator: {exc}") from exc
227
 
 
254
  prompt_weighting=bool(prompt_weighting),
255
  )
256
 
 
 
 
 
 
 
 
 
 
 
 
 
 
257
  try:
258
+ paths = _generate_in_isolated_process(
259
+ context, source_root, request, output, count, step_count, int(preview_interval),
260
+ )
261
  except ToolCancelled:
262
  raise
263
  except Exception as exc: