ADAM October 2026 source release: PixelRow, INRFlow, Wan Video, Oasis player and field guide
Browse filesUpdate 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
- .gitattributes +1 -0
- ADAM-source-2026-10-01.zip +3 -0
- ADAM.spec +1 -0
- README.md +591 -394
- SHA256SUMS.txt +1 -0
- adam/assets.py +72 -20
- adam/auto_training.py +109 -0
- adam/cnn_reviewer.py +172 -0
- adam/commands.py +1 -1
- adam/config.py +3 -0
- adam/dataset_registry.py +4 -3
- adam/generations.py +334 -14
- adam/intelligence.py +200 -0
- adam/job_manager.py +18 -3
- adam/model_plugins.py +3 -1
- adam/model_plugins_builtin/ddpm/manifest.py +3 -1
- adam/model_plugins_builtin/flow_matching/manifest.py +2 -0
- adam/model_plugins_builtin/inrflow/__init__.py +2 -0
- adam/model_plugins_builtin/inrflow/common.py +106 -0
- adam/model_plugins_builtin/inrflow/generator.py +305 -0
- adam/model_plugins_builtin/inrflow/manifest.py +394 -0
- adam/model_plugins_builtin/inrflow/model.py +400 -0
- adam/model_plugins_builtin/inrflow/trainer.py +523 -0
- adam/model_plugins_builtin/oasis/manifest.py +24 -9
- adam/model_plugins_builtin/pixelrow/__init__.py +2 -0
- adam/model_plugins_builtin/pixelrow/common.py +86 -0
- adam/model_plugins_builtin/pixelrow/generator.py +212 -0
- adam/model_plugins_builtin/pixelrow/manifest.py +301 -0
- adam/model_plugins_builtin/pixelrow/model.py +244 -0
- adam/model_plugins_builtin/pixelrow/trainer.py +411 -0
- adam/model_plugins_builtin/sdxl_lora/manifest.py +1 -1
- adam/model_plugins_builtin/wan_video/__init__.py +1 -0
- adam/model_plugins_builtin/wan_video/manifest.py +50 -0
- adam/nova.py +2 -0
- adam/oasis_dataset.py +94 -15
- adam/oasis_player.py +112 -0
- adam/ollama.py +56 -7
- adam/orion.py +85 -0
- adam/planner.py +254 -21
- adam/progressive_training.py +103 -0
- adam/recommendations.py +113 -15
- adam/remote_access.py +25 -4
- adam/remote_dashboard.py +22 -13
- adam/remote_v1.py +53 -35
- adam/tool_folders.py +5 -0
- adam/tools/ddpm_adapter.py +244 -28
- adam/tools/flow_adapter.py +119 -4
- adam/tools/flow_generator.py +19 -0
- adam/tools/lora_adapter.py +8 -1
- 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 |
-
|
| 12 |
-
|
| 13 |
-
|
| 14 |
-
|
| 15 |
-
|
| 16 |
-
|
| 17 |
-
|
| 18 |
-
|
| 19 |
-
|
| 20 |
-
|
| 21 |
-
|
| 22 |
-
|
| 23 |
-
|
| 24 |
-
|
| 25 |
-
|
| 26 |
-
|
| 27 |
-
|
| 28 |
-
|
| 29 |
-
|
| 30 |
-
|
| 31 |
-
|
| 32 |
-
|
| 33 |
-
|
| 34 |
-
|
| 35 |
-
|
| 36 |
-
|
| 37 |
-
|
| 38 |
-
|
| 39 |
-
|
| 40 |
-
|
| 41 |
-
|
| 42 |
-
|
| 43 |
-
|
| 44 |
-
|
| 45 |
-
|
| 46 |
-
|
| 47 |
-
|
| 48 |
-
|
| 49 |
-
|
| 50 |
-
|
| 51 |
-
|
| 52 |
-
|
| 53 |
-
|
| 54 |
-
|
| 55 |
-
|
| 56 |
-
|
| 57 |
-
|
| 58 |
-
|
| 59 |
-
|
| 60 |
-
|
| 61 |
-
|
| 62 |
-
|
| 63 |
-
|
| 64 |
-
|
| 65 |
-
|
| 66 |
-
|
| 67 |
-
|
| 68 |
-
|
| 69 |
-
|
| 70 |
-
|
| 71 |
-
|
| 72 |
-
|
| 73 |
-
|
| 74 |
-
|
| 75 |
-
|
| 76 |
-
|
| 77 |
-
|
| 78 |
-
|
| 79 |
-
|
| 80 |
-
|
| 81 |
-
|
| 82 |
-
|
| 83 |
-
|
| 84 |
-
|
| 85 |
-
|
| 86 |
-
|
| 87 |
-
|
| 88 |
-
|
| 89 |
-
|
| 90 |
-
|
| 91 |
-
|
| 92 |
-
|
| 93 |
-
|
| 94 |
-
|
| 95 |
-
|
| 96 |
-
|
| 97 |
-
|
| 98 |
-
|
| 99 |
-
|
| 100 |
-
|
| 101 |
-
|
| 102 |
-
|
| 103 |
-
|
| 104 |
-
|
| 105 |
-
|
| 106 |
-
|
| 107 |
-
|
| 108 |
-
|
| 109 |
-
|
| 110 |
-
|
| 111 |
-
|
| 112 |
-
|
| 113 |
-
|
| 114 |
-
|
| 115 |
-
|
| 116 |
-
|
| 117 |
-
|
| 118 |
-
|
| 119 |
-
|
| 120 |
-
|
| 121 |
-
|
| 122 |
-
|
| 123 |
-
|
| 124 |
-
|
| 125 |
-
|
| 126 |
-
|
| 127 |
-
|
| 128 |
-
|
| 129 |
-
|
| 130 |
-
|
| 131 |
-
|
| 132 |
-
|
| 133 |
-
|
| 134 |
-
|
| 135 |
-
|
| 136 |
-
|
| 137 |
-
and
|
| 138 |
-
|
| 139 |
-
|
| 140 |
-
|
| 141 |
-
|
| 142 |
-
|
| 143 |
-
|
| 144 |
-
|
| 145 |
-
|
| 146 |
-
|
| 147 |
-
|
| 148 |
-
in
|
| 149 |
-
|
| 150 |
-
|
| 151 |
-
|
| 152 |
-
|
| 153 |
-
job.
|
| 154 |
-
|
| 155 |
-
|
| 156 |
-
|
| 157 |
-
|
| 158 |
-
|
| 159 |
-
|
| 160 |
-
|
| 161 |
-
|
| 162 |
-
|
| 163 |
-
|
| 164 |
-
|
| 165 |
-
|
| 166 |
-
|
| 167 |
-
``
|
| 168 |
-
|
| 169 |
-
|
| 170 |
-
|
| 171 |
-
|
| 172 |
-
|
| 173 |
-
|
| 174 |
-
``
|
| 175 |
-
|
| 176 |
-
|
| 177 |
-
|
| 178 |
-
|
| 179 |
-
|
| 180 |
-
The
|
| 181 |
-
|
| 182 |
-
|
| 183 |
-
|
| 184 |
-
|
| 185 |
-
|
| 186 |
-
|
| 187 |
-
|
| 188 |
-
|
| 189 |
-
|
| 190 |
-
-
|
| 191 |
-
|
| 192 |
-
|
| 193 |
-
|
| 194 |
-
|
| 195 |
-
|
| 196 |
-
|
| 197 |
-
|
| 198 |
-
|
| 199 |
-
|
| 200 |
-
|
| 201 |
-
|
| 202 |
-
|
| 203 |
-
|
| 204 |
-
|
| 205 |
-
|
| 206 |
-
|
| 207 |
-
|
| 208 |
-
|
| 209 |
-
|
| 210 |
-
|
| 211 |
-
|
| 212 |
-
|
| 213 |
-
|
| 214 |
-
|
| 215 |
-
|
| 216 |
-
|
| 217 |
-
|
| 218 |
-
|
| 219 |
-
|
| 220 |
-
|
| 221 |
-
|
| 222 |
-
|
| 223 |
-
|
| 224 |
-
|
| 225 |
-
|
| 226 |
-
|
| 227 |
-
|
| 228 |
-
|
| 229 |
-
|
| 230 |
-
|
| 231 |
-
|
| 232 |
-
|
| 233 |
-
|
| 234 |
-
|
| 235 |
-
|
| 236 |
-
|
| 237 |
-
|
| 238 |
-
|
| 239 |
-
|
| 240 |
-
|
| 241 |
-
|
| 242 |
-
|
| 243 |
-
|
| 244 |
-
|
| 245 |
-
|
| 246 |
-
|
| 247 |
-
|
| 248 |
-
|
| 249 |
-
|
| 250 |
-
|
| 251 |
-
|
| 252 |
-
|
| 253 |
-
|
| 254 |
-
|
| 255 |
-
|
| 256 |
-
|
| 257 |
-
|
| 258 |
-
|
| 259 |
-
|
| 260 |
-
|
| 261 |
-
|
| 262 |
-
|
| 263 |
-
|
| 264 |
-
|
| 265 |
-
|
| 266 |
-
|
| 267 |
-
|
| 268 |
-
`
|
| 269 |
-
|
| 270 |
-
|
| 271 |
-
|
| 272 |
-
Training
|
| 273 |
-
|
| 274 |
-
|
| 275 |
-
|
| 276 |
-
|
| 277 |
-
|
| 278 |
-
|
| 279 |
-
|
| 280 |
-
|
| 281 |
-
|
| 282 |
-
|
| 283 |
-
|
| 284 |
-
|
| 285 |
-
|
| 286 |
-
|
| 287 |
-
|
| 288 |
-
|
| 289 |
-
|
| 290 |
-
|
| 291 |
-
|
| 292 |
-
|
| 293 |
-
|
| 294 |
-
|
| 295 |
-
|
| 296 |
-
|
| 297 |
-
|
| 298 |
-
|
| 299 |
-
|
| 300 |
-
|
| 301 |
-
|
| 302 |
-
|
| 303 |
-
|
| 304 |
-
|
| 305 |
-
|
| 306 |
-
|
| 307 |
-
|
| 308 |
-
|
| 309 |
-
|
| 310 |
-
|
| 311 |
-
|
| 312 |
-
|
| 313 |
-
|
| 314 |
-
|
| 315 |
-
|
| 316 |
-
|
| 317 |
-
|
| 318 |
-
|
| 319 |
-
|
| 320 |
-
|
| 321 |
-
|
| 322 |
-
|
| 323 |
-
|
| 324 |
-
|
| 325 |
-
|
| 326 |
-
|
| 327 |
-
|
| 328 |
-
|
| 329 |
-
|
| 330 |
-
|
| 331 |
-
|
| 332 |
-
|
| 333 |
-
|
| 334 |
-
|
| 335 |
-
|
| 336 |
-
|
| 337 |
-
|
| 338 |
-
|
| 339 |
-
|
| 340 |
-
``
|
| 341 |
-
|
| 342 |
-
|
| 343 |
-
|
| 344 |
-
|
| 345 |
-
|
| 346 |
-
|
| 347 |
-
|
| 348 |
-
|
| 349 |
-
|
| 350 |
-
|
| 351 |
-
|
| 352 |
-
|
| 353 |
-
|
| 354 |
-
|
| 355 |
-
|
| 356 |
-
|
| 357 |
-
|
| 358 |
-
|
| 359 |
-
|
| 360 |
-
|
| 361 |
-
|
| 362 |
-
|
| 363 |
-
|
| 364 |
-
|
| 365 |
-
|
| 366 |
-
|
| 367 |
-
|
| 368 |
-
|
| 369 |
-
|
| 370 |
-
|
| 371 |
-
|
| 372 |
-
|
| 373 |
-
|
| 374 |
-
|
| 375 |
-
|
| 376 |
-
|
| 377 |
-
|
| 378 |
-
|
| 379 |
-
|
| 380 |
-
|
| 381 |
-
|
| 382 |
-
|
| 383 |
-
|
| 384 |
-
|
| 385 |
-
|
| 386 |
-
|
| 387 |
-
|
| 388 |
-
|
| 389 |
-
|
| 390 |
-
|
| 391 |
-
|
| 392 |
-
|
| 393 |
-
|
| 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 |
+

|
| 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 |
+

|
| 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 |
+

|
| 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
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 |
-
|
|
|
|
| 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 |
-
|
|
|
|
|
|
|
| 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 |
-
|
| 78 |
-
|
| 79 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 =
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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"\
|
|
|
|
|
|
|
| 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 |
-
|
| 349 |
-
|
| 350 |
-
|
| 351 |
-
|
| 352 |
-
|
| 353 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 354 |
records.sort(key=lambda item: item.created_at or item.folder.name, reverse=True)
|
| 355 |
-
return records
|
| 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 `` `` 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,
|
| 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,
|
| 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(
|
|
|
|
|
|
|
| 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": "
|
|
|
|
|
|
|
| 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.
|
| 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", "
|
| 17 |
"estimate": "High; 256x144 with batch 2 is the conservative RTX 3060 starting point.",
|
| 18 |
},
|
| 19 |
}
|
| 20 |
|
| 21 |
TRAINING_SETTINGS = {
|
| 22 |
-
"
|
|
|
|
|
|
|
| 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": ["
|
| 27 |
"gradient_accumulation": {"label": "Gradient accumulation", "type": "int", "default": 1, "min": 1, "max": 16, "group": "Optimization"},
|
| 28 |
-
"
|
| 29 |
-
"
|
|
|
|
|
|
|
| 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":
|
| 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(
|
|
|
|
|
|
|
|
|
|
| 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(
|
| 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 |
-
|
| 149 |
-
|
| 150 |
-
|
| 151 |
-
|
| 152 |
-
|
| 153 |
-
|
| 154 |
-
|
| 155 |
-
|
| 156 |
-
|
| 157 |
-
seen_resolution
|
| 158 |
-
|
| 159 |
-
|
| 160 |
-
|
|
|
|
| 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 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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(
|
| 68 |
-
|
|
|
|
|
|
|
| 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
|
| 203 |
-
"job results
|
|
|
|
| 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
|
| 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 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 |
-
|
| 722 |
-
|
| 723 |
-
|
| 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 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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
|
| 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 |
-
|
| 96 |
-
|
| 97 |
-
|
| 98 |
-
|
| 99 |
-
|
| 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": "
|
| 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":
|
| 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 |
-
|
| 182 |
-
|
| 183 |
-
|
| 184 |
-
|
| 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 =
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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="
|
|
|
|
|
|
|
| 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 |
-
|
|
|
|
| 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"><</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 |
-
|
| 428 |
-
|
| 429 |
-
|
| 430 |
-
|
| 431 |
-
|
| 432 |
-
|
| 433 |
-
|
| 434 |
-
|
| 435 |
-
|
| 436 |
-
|
| 437 |
-
|
| 438 |
-
|
| 439 |
-
|
| 440 |
-
|
| 441 |
-
|
| 442 |
-
|
| 443 |
-
"
|
| 444 |
-
|
| 445 |
-
|
| 446 |
-
|
| 447 |
-
|
| 448 |
-
|
| 449 |
-
|
| 450 |
-
|
| 451 |
-
|
| 452 |
-
|
| 453 |
-
"
|
| 454 |
-
"
|
| 455 |
-
"
|
| 456 |
-
|
| 457 |
-
|
| 458 |
-
|
| 459 |
-
|
| 460 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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
|
| 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 =
|
| 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 (
|
| 175 |
raise ToolExecutionError("The saved DDPM model is incomplete and cannot be fine-tuned safely.")
|
| 176 |
-
pretrained_model =
|
| 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 |
-
|
| 194 |
-
|
| 195 |
-
|
| 196 |
-
|
| 197 |
-
output = output.with_name(f"{output.name}_finetuned_{timestamp}")
|
| 198 |
resume = None
|
| 199 |
context.log(
|
| 200 |
-
f"Changing DDPM
|
|
|
|
| 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
|
| 213 |
-
|
| 214 |
-
|
| 215 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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(
|
| 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(
|
|
|
|
|
|
|
|
|
|
| 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
|
| 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 |
-
|
| 81 |
-
|
|
|
|
|
|
|
| 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 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 =
|
|
|
|
|
|
|
| 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:
|