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

Update ADAM safety, UI, and model workflows (#1)

Browse files

- Update ADAM safety, UI, and model workflows (1ac97445998399d40a3662e6a2505c673e45aa2b)

This view is limited to 50 files because it contains too many changes.   See raw diff
Files changed (50) hide show
  1. .gitattributes +36 -44
  2. .gitignore +10 -16
  3. Launch ADAM.bat +14 -50
  4. README.md +365 -24
  5. adam/assets.py +76 -3
  6. adam/atlas.py +24 -4
  7. adam/commands.py +42 -17
  8. adam/config.py +9 -0
  9. adam/dataset_lab.py +127 -0
  10. adam/dataset_registry.py +484 -0
  11. adam/eve.py +13 -2
  12. adam/executor.py +19 -27
  13. adam/experiment_tracker.py +387 -0
  14. adam/generations.py +72 -0
  15. adam/image_preferences.py +313 -0
  16. adam/job_manager.py +311 -10
  17. adam/model_inspector/__init__.py +15 -0
  18. adam/model_inspector/base.py +146 -0
  19. adam/model_inspector/comparison.py +160 -0
  20. adam/model_inspector/ddpm.py +29 -0
  21. adam/model_inspector/detector.py +35 -0
  22. adam/model_inspector/flow_matching.py +25 -0
  23. adam/model_inspector/generic.py +349 -0
  24. adam/model_inspector/lora.py +23 -0
  25. adam/model_inspector/maskgit.py +24 -0
  26. adam/model_inspector/statistics.py +142 -0
  27. adam/model_plugin_backend.py +38 -0
  28. adam/model_plugins.py +494 -0
  29. adam/model_plugins_builtin/__init__.py +1 -0
  30. adam/model_plugins_builtin/ddpm/__init__.py +1 -0
  31. adam/model_plugins_builtin/ddpm/manifest.py +56 -0
  32. adam/model_plugins_builtin/flow_matching/__init__.py +1 -0
  33. adam/model_plugins_builtin/flow_matching/manifest.py +51 -0
  34. adam/model_plugins_builtin/model_template/__init__.py +1 -0
  35. adam/model_plugins_builtin/model_template/manifest.py +30 -0
  36. adam/model_plugins_builtin/oasis/__init__.py +1 -0
  37. adam/model_plugins_builtin/oasis/manifest.py +70 -0
  38. adam/model_plugins_builtin/sdxl_lora/__init__.py +1 -0
  39. adam/model_plugins_builtin/sdxl_lora/manifest.py +60 -0
  40. adam/model_profiles.py +93 -0
  41. adam/models.py +9 -0
  42. adam/oasis_dataset.py +215 -0
  43. adam/orion.py +17 -2
  44. adam/planner.py +303 -49
  45. adam/process_control.py +36 -0
  46. adam/recommendations.py +195 -0
  47. adam/registry.py +52 -0
  48. adam/remote_access.py +0 -0
  49. adam/remote_api.py +147 -0
  50. adam/remote_dashboard.py +141 -0
.gitattributes CHANGED
@@ -1,46 +1,38 @@
1
- *.7z filter=lfs diff=lfs merge=lfs -text
2
- *.arrow filter=lfs diff=lfs merge=lfs -text
3
- *.bin filter=lfs diff=lfs merge=lfs -text
4
- *.bz2 filter=lfs diff=lfs merge=lfs -text
5
- *.ckpt filter=lfs diff=lfs merge=lfs -text
6
- *.ftz filter=lfs diff=lfs merge=lfs -text
7
- *.gz filter=lfs diff=lfs merge=lfs -text
8
- *.h5 filter=lfs diff=lfs merge=lfs -text
9
- *.joblib filter=lfs diff=lfs merge=lfs -text
10
- *.lfs.* filter=lfs diff=lfs merge=lfs -text
11
- *.mlmodel filter=lfs diff=lfs merge=lfs -text
12
- *.model filter=lfs diff=lfs merge=lfs -text
13
- *.msgpack filter=lfs diff=lfs merge=lfs -text
14
- *.npy filter=lfs diff=lfs merge=lfs -text
15
- *.npz filter=lfs diff=lfs merge=lfs -text
16
- *.onnx filter=lfs diff=lfs merge=lfs -text
17
- *.ot filter=lfs diff=lfs merge=lfs -text
18
- *.parquet filter=lfs diff=lfs merge=lfs -text
19
- *.pb filter=lfs diff=lfs merge=lfs -text
20
- *.pickle filter=lfs diff=lfs merge=lfs -text
21
- *.pkl filter=lfs diff=lfs merge=lfs -text
22
- *.pt filter=lfs diff=lfs merge=lfs -text
23
- *.pth filter=lfs diff=lfs merge=lfs -text
24
- *.rar filter=lfs diff=lfs merge=lfs -text
25
- *.safetensors filter=lfs diff=lfs merge=lfs -text
26
- saved_model/**/* filter=lfs diff=lfs merge=lfs -text
27
- *.tar.* filter=lfs diff=lfs merge=lfs -text
28
- *.tar filter=lfs diff=lfs merge=lfs -text
29
- *.tflite filter=lfs diff=lfs merge=lfs -text
30
- *.tgz filter=lfs diff=lfs merge=lfs -text
31
- *.wasm filter=lfs diff=lfs merge=lfs -text
32
- *.xz filter=lfs diff=lfs merge=lfs -text
33
- *.zip filter=lfs diff=lfs merge=lfs -text
34
- *.zst filter=lfs diff=lfs merge=lfs -text
35
- *tfevents* filter=lfs diff=lfs merge=lfs -text
36
- adam/__pycache__/planner.cpython-311.pyc filter=lfs diff=lfs merge=lfs -text
37
- adam/ui/__pycache__/main_window.cpython-310.pyc filter=lfs diff=lfs merge=lfs -text
38
- adam/ui/__pycache__/main_window.cpython-311.pyc filter=lfs diff=lfs merge=lfs -text
39
- adam/ui/__pycache__/studio.cpython-311.pyc filter=lfs diff=lfs merge=lfs -text
40
  assets/adam_atom.ico filter=lfs diff=lfs merge=lfs -text
41
  assets/adam_atom.png filter=lfs diff=lfs merge=lfs -text
42
- build/ADAM/ADAM.exe filter=lfs diff=lfs merge=lfs -text
43
- build/ADAM/ADAM.pkg filter=lfs diff=lfs merge=lfs -text
44
- build/ADAM/PYZ-00.pyz filter=lfs diff=lfs merge=lfs -text
45
- build/ADAM/xref-ADAM.html filter=lfs diff=lfs merge=lfs -text
46
- tests/__pycache__/test_generations.cpython-311-pytest-8.4.2.pyc filter=lfs diff=lfs merge=lfs -text
 
1
+ *.7z filter=lfs diff=lfs merge=lfs -text
2
+ *.arrow filter=lfs diff=lfs merge=lfs -text
3
+ *.bin filter=lfs diff=lfs merge=lfs -text
4
+ *.bz2 filter=lfs diff=lfs merge=lfs -text
5
+ *.ckpt filter=lfs diff=lfs merge=lfs -text
6
+ *.ftz filter=lfs diff=lfs merge=lfs -text
7
+ *.gz filter=lfs diff=lfs merge=lfs -text
8
+ *.h5 filter=lfs diff=lfs merge=lfs -text
9
+ *.joblib filter=lfs diff=lfs merge=lfs -text
10
+ *.lfs.* filter=lfs diff=lfs merge=lfs -text
11
+ *.mlmodel filter=lfs diff=lfs merge=lfs -text
12
+ *.model filter=lfs diff=lfs merge=lfs -text
13
+ *.msgpack filter=lfs diff=lfs merge=lfs -text
14
+ *.npy filter=lfs diff=lfs merge=lfs -text
15
+ *.npz filter=lfs diff=lfs merge=lfs -text
16
+ *.onnx filter=lfs diff=lfs merge=lfs -text
17
+ *.ot filter=lfs diff=lfs merge=lfs -text
18
+ *.parquet filter=lfs diff=lfs merge=lfs -text
19
+ *.pb filter=lfs diff=lfs merge=lfs -text
20
+ *.pickle filter=lfs diff=lfs merge=lfs -text
21
+ *.pkl filter=lfs diff=lfs merge=lfs -text
22
+ *.pt filter=lfs diff=lfs merge=lfs -text
23
+ *.pth filter=lfs diff=lfs merge=lfs -text
24
+ *.rar filter=lfs diff=lfs merge=lfs -text
25
+ *.safetensors filter=lfs diff=lfs merge=lfs -text
26
+ saved_model/**/* filter=lfs diff=lfs merge=lfs -text
27
+ *.tar.* filter=lfs diff=lfs merge=lfs -text
28
+ *.tar filter=lfs diff=lfs merge=lfs -text
29
+ *.tflite filter=lfs diff=lfs merge=lfs -text
30
+ *.tgz filter=lfs diff=lfs merge=lfs -text
31
+ *.wasm filter=lfs diff=lfs merge=lfs -text
32
+ *.xz filter=lfs diff=lfs merge=lfs -text
33
+ *.zip filter=lfs diff=lfs merge=lfs -text
34
+ *.zst filter=lfs diff=lfs merge=lfs -text
35
+ *tfevents* 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
 
 
 
 
.gitignore CHANGED
@@ -1,28 +1,22 @@
1
- # Python and test caches
2
  __pycache__/
3
  *.py[cod]
4
  .pytest_cache/
5
  .venv/
6
  venv/
7
-
8
- # Local application state and user-connected tools
9
- config/settings.json
10
- config/settings.local.json
11
- config/external_tools.json
12
- !config/external_tools.json
13
- data/*
14
- !data/.gitkeep
15
- logs/
16
-
17
- # Generated datasets, models, and media
18
  ADAM_Datasets/
19
- artifacts/
20
  build/
21
  dist/
 
 
 
 
 
 
22
  *.safetensors
23
  *.ckpt
24
  *.pt
25
  *.pth
26
- *.bin
27
- *.onnx
28
- *.tmp
 
 
1
  __pycache__/
2
  *.py[cod]
3
  .pytest_cache/
4
  .venv/
5
  venv/
6
+ logs/*.log
7
+ data/
 
 
 
 
 
 
 
 
 
8
  ADAM_Datasets/
 
9
  build/
10
  dist/
11
+ hf-release-*/
12
+ artifacts/
13
+ config/settings.local.json
14
+ config/settings.json
15
+ *.tmp
16
+ *.sqlite3
17
  *.safetensors
18
  *.ckpt
19
  *.pt
20
  *.pth
21
+ LoRAModelsHere/
22
+ LoRA StableDiffusionModels Here/
 
Launch ADAM.bat CHANGED
@@ -1,61 +1,25 @@
1
  @echo off
2
  setlocal
3
  cd /d "%~dp0"
4
- python -c "import yt_dlp, cv2, numpy" 2>nul
5
  if errorlevel 1 (
6
- echo Installing the ADAM Video Dataset Collector requirements...
7
- python -m pip install -r "%~dp0requirements.txt"
8
- if errorlevel 1 goto :adam_dependency_error
9
- )
10
- python -c "import selenium, requests, PIL" 2>nul
11
- if errorlevel 1 (
12
- echo Installing the Dataset Collector requirements for ADAM...
13
- python -m pip install -r "D:\Users\PlayRobloxAllDay\Desktop\Programs\GoogleImageDatasetCollector\requirements.txt"
14
- if errorlevel 1 goto :dependency_error
15
- )
16
- python -c "import datasets, diffusers, transformers, accelerate, torch, torchvision" 2>nul
17
- if errorlevel 1 (
18
- echo Installing the DDPM training requirements for ADAM...
19
- python -m pip install -r "D:\Users\PlayRobloxAllDay\Desktop\Programs\DDPM\requirements.txt"
20
- if errorlevel 1 goto :ddpm_dependency_error
21
  )
 
 
22
  python main.py
23
  if errorlevel 1 (
24
  echo.
25
- echo ADAM could not start. Install the requirements with:
26
- echo python -m pip install -r requirements.txt
 
27
  echo.
28
  pause
 
29
  )
30
- endlocal
31
- exit /b
32
-
33
- :adam_dependency_error
34
- echo.
35
- echo ADAM could not install its Video Dataset Collector requirements.
36
- echo Run this command with the same Python used to start ADAM:
37
- echo python -m pip install -r "%~dp0requirements.txt"
38
- echo.
39
- pause
40
- endlocal
41
- exit /b
42
-
43
- :dependency_error
44
- echo.
45
- echo ADAM could not install the Dataset Collector requirements.
46
- echo Run this command and then launch ADAM again:
47
- echo python -m pip install -r "D:\Users\PlayRobloxAllDay\Desktop\Programs\GoogleImageDatasetCollector\requirements.txt"
48
- echo.
49
- pause
50
- endlocal
51
- exit /b
52
-
53
- :ddpm_dependency_error
54
- echo.
55
- echo ADAM could not install the DDPM training requirements.
56
- echo Run this command and then launch ADAM again:
57
- echo python -m pip install -r "D:\Users\PlayRobloxAllDay\Desktop\Programs\DDPM\requirements.txt"
58
- echo.
59
- pause
60
- endlocal
61
- exit /b
 
1
  @echo off
2
  setlocal
3
  cd /d "%~dp0"
4
+ python -c "import PySide6, psutil, PIL" 2>nul
5
  if errorlevel 1 (
6
+ echo ADAM's desktop requirements are missing from this Python environment.
7
+ echo Install them with:
8
+ echo python -m pip install -r "%~dp0requirements.txt"
9
+ echo.
10
+ pause
11
+ exit /b 1
 
 
 
 
 
 
 
 
 
12
  )
13
+ rem Optional trainers and collectors are connected through ADAM Settings.
14
+ rem Their dependencies belong to their own environments, not startup.
15
  python main.py
16
  if errorlevel 1 (
17
  echo.
18
+ echo ADAM could not start. Check the application log for details.
19
+ echo To install the application requirements, run:
20
+ echo python -m pip install -r "%~dp0requirements.txt"
21
  echo.
22
  pause
23
+ exit /b 1
24
  )
25
+ endlocal
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
README.md CHANGED
@@ -10,44 +10,385 @@ tags:
10
 
11
  # ADAM — AI Development and Automation Manager
12
 
13
- ADAM is a local Windows desktop hub for organizing AI image-development workflows. It helps you prepare and review datasets, connect your own training tools, plan runs with explicit approval, monitor jobs, and keep a history of your own assets.
 
 
14
 
15
- ## What is included
16
 
17
- - ADAM source code and the built-in tool registry.
18
- - A clean, empty external-tool registry.
19
- - No model weights, LoRAs, checkpoints, datasets, generated images, job history, personal paths, API keys, or local settings.
20
 
21
- ## Requirements
 
 
22
 
23
- - Windows 10/11
24
- - Python 3.10 or later
25
- - Optional: NVIDIA GPU for compatible training workflows
26
- - Optional: Ollama for local chat assistance
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
27
 
28
- ## Install and run
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
29
 
30
  ```powershell
31
- git clone https://huggingface.co/SyntheticMDProductions/AI_Development_Automation_Manager
32
- cd AI_Development_Automation_Manager
33
- python -m pip install -r requirements.txt
34
  python main.py
35
  ```
36
 
37
- On first launch, ADAM creates your personal `config/settings.json` automatically. In **Settings → Tool folders**, connect the local projects, datasets, and models that you own and want ADAM to manage. ADAM does not bundle or download model weights for you.
38
 
39
- ## Notes for users
 
40
 
41
- - Training and collection plans require approval before ADAM starts them.
42
- - You are responsible for the licenses, permissions, and rights for any datasets, models, and third-party tools you connect.
43
- - This repository is a downloadable desktop application. It is not a hosted Hugging Face Space or an inference model.
44
 
45
- ## Development
46
 
47
- ```powershell
48
- python -m pytest -q
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
49
  ```
50
 
51
- ## License
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
52
 
53
- ADAM is released under the [MIT License](LICENSE).
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
10
 
11
  # ADAM — AI Development and Automation Manager
12
 
13
+ ADAM is a local, safety-first desktop hub for orchestrating AI project tools.
14
+ It includes registered dataset, DDPM, and SDXL LoRA workflows with background
15
+ planning, approval gates, progress reporting, and persistent asset history.
16
 
17
+ ![ADAM command center](docs/screenshots/command-center.png)
18
 
19
+ Dataset preparation, captioning, and preview placeholders remain clearly marked
20
+ as demo tools. The connected Dataset Collector, DDPM trainer, and Local SDXL
21
+ LoRA Trainer use real adapters and never fall back to simulated training.
22
 
23
+ Existing program folders can be connected from **Settings → Tool folders**.
24
+ ADAM stores only the path and scans for likely entry points; it does not copy or
25
+ modify the external project. Folder assignments can also be pasted into chat:
26
 
27
+ ```text
28
+ DDPM Trainer: D:\AI\DDPM
29
+ Flow Matching Trainer: D:\AI\FlowMatchImageGenerator
30
+ ```
31
+
32
+ On a new computer, install `requirements.txt` in your chosen Python environment
33
+ before running `Launch ADAM.bat`. The launcher checks desktop dependencies and
34
+ does not install packages automatically or depend on the developer's personal
35
+ trainer folders. Install each optional trainer's dependencies according to that
36
+ tool's setup instructions before using its ADAM workflow.
37
+
38
+ Remote access is disabled by default. Devices with an access token can browse
39
+ datasets, edit captions/review marks and submit work. Only enable it for trusted
40
+ devices. The desktop **Allow remote job controls and approval changes** setting
41
+ also permits remote confirmation, stopping and retrying jobs. A remote browser
42
+ can enable training auto-approval only after that desktop permission is granted;
43
+ it can always turn auto-approval off. Saved token changes and disabled access
44
+ take effect for new requests without restarting the server.
45
+
46
+ Use private Tailscale access for connections beyond a trusted local network;
47
+ the built-in HTTP listener does not provide transport encryption by itself.
48
+ Phone URLs and QR codes contain the access token and should be treated as
49
+ credentials. Remote commands require JSON, have bounded request sizes and
50
+ connection counts, and reject cross-site browser submissions. These controls
51
+ do not sandbox installed Python plugins or connected trainers: install only
52
+ code you trust.
53
+
54
+ Detection does not automatically authorize training. A real training adapter
55
+ remains gated until its dataset, model name, run settings, and output location
56
+ are explicit.
57
+
58
+ ## Training agents
59
+
60
+ ![ADAM model creation settings](docs/screenshots/create-model.png)
61
+
62
+ ADAM's training lifecycle is divided into four explainable responsibilities:
63
+
64
+ - **EVE** reviews dataset membership and leaves uncertain images for the user.
65
+ - **ORION** reviews planned epochs, batch size, resolution, image exposures, and
66
+ estimated optimizer steps. He can require approval but never silently changes
67
+ the requested settings. In the Model Creation Assistant, **ORION: apply a
68
+ starting recipe** fills a conservative, editable draft from the image count
69
+ and selected resolution before a plan is built.
70
+ - **ATLAS** watches active training for non-finite loss, sustained critical GPU
71
+ temperature, critically low disk space, stalls, and large runtime overruns.
72
+ Critical conditions pause the trainer process tree so the user can inspect it.
73
+ - **NOVA** examines available post-training previews and samples for unreadable
74
+ files and exact-looking duplicate collapse. Her report explicitly separates
75
+ technical sample health from subjective or subject-quality review.
76
+
77
+ ORION, ATLAS, and NOVA reports are stored with each durable job record and are
78
+ shown in Current Plan, Active Job, and Jobs / History respectively. ATLAS's
79
+ default thresholds can be overridden in `config/settings.json` with the
80
+ `atlas_*` settings defined in `adam/config.py`.
81
+
82
+ Every new job passes through the shared preflight and ORION review before its
83
+ queue state is chosen. Desktop plans, Remote prompts, and Remote training forms
84
+ use the same review. Remote training auto-approval still applies to ordinary
85
+ plans, but an ORION warning leaves the job awaiting explicit approval. Reviewing
86
+ a plan does not change the requested training settings.
87
+
88
+ ## Real image collection
89
+
90
+ When a valid Dataset Collector folder is connected, the `dataset_collector`
91
+ registry entry uses ADAM's real visible-browser adapter. After plan approval it:
92
+
93
+ - opens Bing Images in a normal visible Chrome window;
94
+ - waits when consent/CAPTCHA/human-verification text is detected;
95
+ - resumes automatically after the user resolves the page;
96
+ - downloads valid images at least 256×256;
97
+ - removes exact duplicate downloads;
98
+ - writes a matching `.txt` caption beside every image; and
99
+ - records URLs, captions, sources, and dimensions in `metadata.csv`.
100
+
101
+ No CAPTCHA or website restriction is bypassed. Closing Chrome or stopping the
102
+ job ends collection safely. A new timestamped dataset folder is used rather
103
+ than overwriting an existing collection.
104
+
105
+ ADAM keeps an incomplete DDPM request in conversation memory. A follow-up such
106
+ as `dataset folder Mario, model name Mario V2, epoch count 100, output D:\Runs`
107
+ fills the pending fields and validates named datasets against the connected
108
+ collector. It will not start if the dataset cannot be found.
109
+
110
+ ## Showcase videos
111
+
112
+ The **Showcase Video** workspace creates a finished MP4 directly from completed
113
+ DDPM and Flow Matching models. Select and reorder the models, choose 12–24
114
+ images per model, a 3-, 4-, or 5-second image duration, shared steps and aspect
115
+ ratio, provider-compatible samplers, seed, and 720p or 1080p output. ADAM runs
116
+ the image batches sequentially and then renders a request-list interface that
117
+ tracks the active model, image number, trainer, steps, sampler, and aspect ratio.
118
+ LoRA models are intentionally excluded from this streamlined workflow.
119
+
120
+ When Ollama is reachable, messages that are not workflow commands receive a
121
+ short conversational answer. Ollama may explain or plan, but it still cannot
122
+ bypass the registry or confirmation gates.
123
+
124
+ ## Web search in Chat Mode
125
 
126
+ Chat Mode can give local Ollama current web context without an API key. Enable
127
+ it in **Settings → Planning model**, then ask naturally, for example:
128
+
129
+ ```text
130
+ Search the web for Dandy's World character ideas.
131
+ What are the latest Ollama release notes?
132
+ Look up a reference for a cyberpunk city character.
133
+ ```
134
+
135
+ ADAM sends only that search query to Bing's public results feed, reads the
136
+ result titles and snippets,
137
+ and passes up to five titles, snippets, and links to Ollama. It does not open
138
+ the result pages, download anything, or let web content run tools. Results are
139
+ untrusted reference material, so ADAM is instructed to cite the links and flag
140
+ uncertainty. Disable the setting to keep Chat Mode fully local.
141
+
142
+ When you explicitly ask ADAM to **read**, **open**, or **research** result links,
143
+ it can read up to three public HTML/text pages and give Ollama short extracts.
144
+ For example: `Search the web for Undertale character ideas and read the most
145
+ relevant links.` Direct links can be read with `Read https://example.com/ and
146
+ summarize it.` Private/local addresses, non-web protocols, oversized pages,
147
+ downloads, and more than three pages are blocked. This control can be disabled
148
+ in Settings.
149
+
150
+ Planning runs away from the interface thread, and conversational Ollama output
151
+ is streamed into the chat. ADAM validates training commands against a strict
152
+ schema and each registered trainer's declared capabilities before offering a
153
+ job.
154
+
155
+ In **Settings → Planning model**, **Chat response length** sets the maximum
156
+ number of generated tokens for a Chat Mode reply. Higher values allow longer
157
+ research summaries but use more time and GPU memory. The default is 1,024,
158
+ which gives Qwen3 enough room to reason and still produce a visible response.
159
+
160
+ ADAM stores friendly dataset/model names, paths, trainer types, epochs, and
161
+ resume checkpoints in `data/assets.json`. Requests such as:
162
+
163
+ ```text
164
+ From the Mario dataset, train it on a DDPM for 300 epochs.
165
+ With the Mario dataset, train it on a LoRA for 100 epochs.
166
+ Continue the Mario model from the DDPM for 50 epochs.
167
+ ```
168
+
169
+ are resolved to real paths before approval. Continuation is offered only when a
170
+ compatible checkpoint exists. New DDPM runs retain the latest resume checkpoint.
171
+
172
+ ## Run
173
 
174
  ```powershell
 
 
 
175
  python main.py
176
  ```
177
 
178
+ On Windows, you can also double-click `Launch ADAM.bat`.
179
 
180
+ The app requires Python 3.10+ and PySide6. Optional integrations use `psutil`
181
+ for system information and `pynvml` for NVIDIA GPU information.
182
 
183
+ ```powershell
184
+ python -m pip install -r requirements.txt
185
+ ```
186
 
187
+ Try:
188
 
189
+ - Click **Create a model…** in Trainer Mode for the guided Model Creation Assistant.
190
+ - `Adam, train a LoRA of Hatsune Miku`
191
+ - `Adam, collect a dataset of liminal spaces`
192
+ - `Adam, generate previews`
193
+ - `Adam, check GPU status`
194
+ - `From the Mario dataset, train it on a DDPM for 300 epochs`
195
+ - `With the Mario dataset, train it on a LoRA for 100 epochs`
196
+
197
+ Training and large collection plans are never started until you approve the
198
+ plan. All actions are recorded in `logs/adam.log`, while project artifacts live
199
+ under `data/projects/`.
200
+
201
+ The Model Creation Assistant can start from a built-in Character LoRA, Style
202
+ LoRA, DDPM, or Flow Matching preset. It can create a dataset or select a
203
+ registered one, recommends starting values, and saves personal presets. The
204
+ result still goes through ADAM's normal validated planner and approval gate.
205
+ Use **+ Add model** to build a multi-model training batch. Each wide model tab
206
+ keeps its own dataset, trainer, name, and settings; the minus button removes an
207
+ unwanted model, and tabs can be dragged to change the run order. ADAM validates
208
+ all models, presents one combined approval plan, and runs them sequentially so
209
+ only one training workflow uses the GPU at a time. A failed step stops the batch
210
+ before a later model starts.
211
+ Before approval, ADAM adds checks for connected tools, dataset contents, the
212
+ LoRA base model, and output-drive free space. Completed dataset and training
213
+ jobs also include a suggested next step.
214
+
215
+ ### Model Batch Builder
216
+
217
+ Use **Create model batch…** to paste one requested subject per line. ADAM turns
218
+ the list into editable model tabs, removes duplicate names, and lets the current
219
+ trainer recipe be applied to any multi-selection of models. The batch is saved
220
+ as a draft so it can be closed and resumed later.
221
+
222
+ For a review-first workflow, choose **Collect missing datasets first**. This
223
+ queues only sequential dataset collection and leaves training in the saved
224
+ draft. After collection, reopen the draft, use **Find collected datasets**, and
225
+ review each dataset in Training Studio. **Exclude rejected** moves rejected
226
+ images out of the training folder into a recoverable quarantine, and **Restore
227
+ excluded** reverses it. **Keep all images** marks the whole selected dataset as
228
+ accepted in one action, after which individual bad images can still be rejected.
229
+ Training remains locked until each model is explicitly
230
+ marked as reviewed and ready. If every linked dataset is acceptable as-is,
231
+ **Approve all datasets** marks the entire batch ready after one confirmation;
232
+ it does not inspect individual images or apply pending rejection decisions.
233
+
234
+ Completed Flow Matching models can be selected in **Fine-tune**. ADAM uses the
235
+ saved Flow model folder as the continuation source, locks the continuation to
236
+ the model's original resolution, and writes the fine-tuned result to a new
237
+ output folder. This continues the saved weights while starting a fresh optimizer
238
+ and learning-rate schedule; it does not overwrite the original model.
239
+
240
+ ## Training Studio
241
+
242
+ The **Training Studio** turns completed work into a reviewable experiment loop:
243
+
244
+ - **Datasets** provides an image gallery, keep/reject decisions, caption editing,
245
+ exact duplicate detection, and visually similar duplicate candidates.
246
+ - **Experiments** compares job settings and outcomes, opens outputs, marks a
247
+ preferred model, and converts successful settings into reusable recipes.
248
+ - **Checkpoint Lab** browses model checkpoints and output images, records
249
+ consistent prompt/seed evaluations, and sends preview requests through the
250
+ normal approval-aware planner.
251
+ - **Recipes** preserves training starting points and can import or export
252
+ portable JSON recipe files.
253
+
254
+ ### EVE AI Dataset Review
255
+
256
+ In Training Studio → Datasets, **EVE AI Review…** performs a local reference-
257
+ guided visual review. Add one or more good reference images and optional bad
258
+ references, then choose Keep and Reject confidence thresholds. EVE uses a small
259
+ DINOv2 vision model to divide the selected dataset into **Keep**, **Reject**, and
260
+ **Uncertain** galleries with confidence scores. The model is downloaded once on
261
+ first use and subsequent analysis stays local.
262
+
263
+ Nothing is applied automatically. Inspect both sides, double-click images for a
264
+ full view, and move selected results between the three groups before choosing
265
+ **Apply EVE review**. EVE's decisions remain ordinary Training Studio review
266
+ marks: they can be manually changed, and rejected files are not moved until
267
+ **Exclude rejected** is selected. The latest proposal is also saved under
268
+ `data/eve_reviews/` for auditing. Use **Select all in current group** (or
269
+ Ctrl/Shift selection) to move many images at once; EVE transfers only the
270
+ chosen thumbnails so manual sorting stays responsive on large datasets.
271
+
272
+ Training panels show elapsed time, a progress-based ETA, recent logs, and a
273
+ loss sparkline when the connected trainer reports `loss`. Preflight summaries
274
+ include clearly labelled workload, duration, VRAM, and disk estimates. These
275
+ estimates are planning hints rather than hardware guarantees.
276
+
277
+ Create a Model also supports live training previews with a configurable
278
+ epoch interval, prompt, and reproducible seed for each model tab. While a
279
+ training job is active, its newest 256×256 preview appears in the right sidebar
280
+ with the source epoch and next scheduled preview. The full-size trainer output
281
+ can be opened from the card. Built-in adapters may publish previews directly;
282
+ registered DDPM, Flow, LoRA, APVD, MaskGit, and other trainers can also
283
+ participate by writing conventionally named `preview`, `sample`, or `epoch`
284
+ images beneath their declared output folder.
285
+
286
+ ## Generations
287
+
288
+ The **Generations** workspace runs compatible registered image generators
289
+ without opening their separate desktop interfaces. The connected DDPM and Flow
290
+ Matching projects can generate from completed models with a reproducible seed,
291
+ sampler or ODE method, step count, image count, and aspect ratio. Generation
292
+ work uses the normal ADAM job queue, progress reporting, cancellation, and
293
+ logging.
294
+
295
+ Every completed batch is stored under `data/generations/` with its images and a
296
+ `generation.json` sidecar. The history gallery can open an image or batch folder
297
+ and restore the exact settings for another run. DDPM creative notes are stored
298
+ with a batch for organization; they are not presented as text conditioning for
299
+ an unconditional DDPM model.
300
+
301
+ Generation history opens as automatic model folders. Each folder uses the
302
+ registered model name and a recent generated image as its cover. Double-clicking
303
+ a folder filters history to that model and selects its provider and model in the
304
+ generation controls, so the next batch is generated into the same existing model
305
+ directory. This view does not move or rewrite older generation files.
306
+
307
+ **Generation Cycle…** selects multiple compatible completed models and queues
308
+ one generation step per model. Choose images per model, a shared prompt or
309
+ creative note, starting seed, slideshow duration, looping, fullscreen playback,
310
+ and an optional model/trainer label. When the cycle finishes, ADAM opens the
311
+ results as a local slideshow while preserving every ordinary generation record
312
+ in history.
313
+
314
+ If ADAM discovers a job interrupted by an unexpected shutdown, it offers to
315
+ open Jobs & History. The previous record remains intact and can be retried as a
316
+ new approval-gated job. Job logs can also be exported for troubleshooting.
317
+
318
+ ## Connect an existing tool
319
+
320
+ ADAM supports importable Python functions and command-line Python scripts.
321
+ For a no-code setup, open **Settings → External Tools → Add external tool**.
322
+ Choose the program folder, select its training entry script and important
323
+ configuration files, then review ADAM's static compatibility and safety report.
324
+ The report covers:
325
+
326
+ - detected command-line options and required inputs;
327
+ - likely dataset formats;
328
+ - output and checkpoint behavior;
329
+ - progress reporting;
330
+ - resume-training support; and
331
+ - potentially risky operations visible in the selected entry script.
332
+
333
+ The 1–10 rating measures how clearly the script fits ADAM's safe command-line
334
+ contract. It is not a guarantee that third-party code is harmless. ADAM does
335
+ not execute a script while scanning it, external tools cannot replace built-in
336
+ registry entries, and every external-tool run requires explicit approval.
337
+
338
+ After registration, a tool can be planned with a request such as:
339
+
340
+ ```text
341
+ Run APVD Model Trainer with dataset=D:\DreamData, epochs=20, output=D:\APVD\output
342
  ```
343
 
344
+ ADAM will ask for any required inputs that were omitted before it offers the
345
+ approval plan.
346
+
347
+ For manual registry configuration, edit the relevant item in
348
+ `config/tools.json`:
349
+
350
+ ```json
351
+ {
352
+ "backend": {
353
+ "type": "python",
354
+ "module": "my_tools.lora",
355
+ "function": "train"
356
+ },
357
+ "demo": false
358
+ }
359
+ ```
360
+
361
+ The function receives a `ToolContext` as its first argument and keyword
362
+ arguments from the approved plan. This keeps training code in one place: your
363
+ existing GUI and ADAM can both call the same backend.
364
+
365
+ For scripts:
366
 
367
+ ```json
368
+ {
369
+ "backend": {
370
+ "type": "script",
371
+ "path": "D:/AI/LoRATrainer/train.py"
372
+ },
373
+ "demo": false
374
+ }
375
+ ```
376
+
377
+ ADAM invokes scripts directly with the current Python interpreter, captures
378
+ stdout/stderr, and never drives another GUI with mouse clicks.
379
+
380
+ ## Safety model
381
+
382
+ - Plans are shown before execution.
383
+ - Long, destructive, or high-volume work requires confirmation.
384
+ - Unregistered tools cannot be invoked.
385
+ - External paths and arguments are validated before execution.
386
+ - The LLM may propose a plan, but only registered tools can execute it.
387
+ - Pause, resume, and cancel controls are available for active jobs.
388
+ - Every tool action and state transition is logged.
389
+
390
+ ## Tests
391
+
392
+ ```powershell
393
+ python -m pytest -q
394
+ ```
adam/assets.py CHANGED
@@ -17,6 +17,15 @@ def _normal(value: str) -> str:
17
  return re.sub(r"[^a-z0-9]+", " ", value.casefold()).strip()
18
 
19
 
 
 
 
 
 
 
 
 
 
20
  @dataclass(slots=True)
21
  class Asset:
22
  id: str
@@ -28,6 +37,7 @@ class Asset:
28
  checkpoint: str = ""
29
  epochs: int = 0
30
  created_at: str = ""
 
31
 
32
  @classmethod
33
  def from_dict(cls, payload: dict[str, Any]) -> "Asset":
@@ -41,6 +51,7 @@ class Asset:
41
  checkpoint=str(payload.get("checkpoint", "")),
42
  epochs=int(payload.get("epochs", 0) or 0),
43
  created_at=str(payload.get("created_at") or _now()),
 
44
  )
45
 
46
 
@@ -82,6 +93,7 @@ class AssetRegistry:
82
  dataset_id: str = "",
83
  checkpoint: str = "",
84
  epochs: int = 0,
 
85
  persist: bool = True,
86
  ) -> Asset:
87
  resolved = str(Path(path).expanduser().resolve())
@@ -100,6 +112,10 @@ class AssetRegistry:
100
  asset.checkpoint = checkpoint
101
  asset.epochs = int(epochs)
102
  asset.created_at = asset.created_at or _now()
 
 
 
 
103
  if existing is None:
104
  self.assets.insert(0, asset)
105
  if persist:
@@ -118,10 +134,14 @@ class AssetRegistry:
118
  key: item[key]
119
  for key in (
120
  "kind", "name", "path", "trainer", "dataset_id",
121
- "checkpoint", "epochs",
122
  )
123
  if key in item
124
  }
 
 
 
 
125
  dataset_path = str(item.get("dataset_path", ""))
126
  if item.get("kind") == "model" and dataset_path and Path(dataset_path).is_dir():
127
  dataset = self.register(
@@ -151,7 +171,18 @@ class AssetRegistry:
151
  folders = config.get("tool_folders", {})
152
  if not isinstance(folders, dict):
153
  return
 
154
  app_root = self.path.parent.parent
 
 
 
 
 
 
 
 
 
 
155
  external_lora_root = app_root / "LoRAModelsHere"
156
  if external_lora_root.is_dir():
157
  for path in external_lora_root.rglob("*.safetensors"):
@@ -194,6 +225,7 @@ class AssetRegistry:
194
  ("ddpm", "ddpm_trainer", "output"),
195
  ("lora", "lora_trainer", "output"),
196
  ("flow", "flow_trainer", "output_flow_models"),
 
197
  ):
198
  root = Path(str(folders.get(folder_name, ""))) / output_name
199
  if not root.is_dir():
@@ -219,6 +251,7 @@ class AssetRegistry:
219
  else -1,
220
  )
221
  elif trainer == "lora":
 
222
  checkpoints = sorted(
223
  (
224
  path for path in folder.glob("*.safetensors")
@@ -228,7 +261,12 @@ class AssetRegistry:
228
  )
229
  if checkpoints:
230
  name = checkpoints[-1].stem.removesuffix("_cancelled")
231
- else:
 
 
 
 
 
232
  checkpoints = []
233
  try:
234
  metadata = json.loads(
@@ -242,8 +280,24 @@ class AssetRegistry:
242
  dataset_path = flow_datasets.get(str(folder.resolve()), "")
243
  except (OSError, ValueError, TypeError, json.JSONDecodeError):
244
  continue
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
245
  checkpoint = (
246
- str(folder) if trainer == "flow" else str(checkpoints[-1]) if checkpoints else ""
247
  )
248
  dataset_id = ""
249
  if dataset_path and Path(dataset_path).is_dir():
@@ -261,8 +315,27 @@ class AssetRegistry:
261
  trainer=trainer,
262
  dataset_id=dataset_id,
263
  checkpoint=checkpoint,
 
264
  persist=False,
265
  )
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
266
  self.save()
267
 
268
  def _flow_dataset_paths(self) -> dict[str, str]:
 
17
  return re.sub(r"[^a-z0-9]+", " ", value.casefold()).strip()
18
 
19
 
20
+ def _friendly_name(value: str, fallback: str) -> str:
21
+ text = str(value or "").strip()
22
+ if not text:
23
+ return fallback
24
+ if re.search(r"^[A-Za-z]:[\\/]", text) or "/" in text or "\\" in text:
25
+ return Path(text).name or fallback
26
+ return text
27
+
28
+
29
  @dataclass(slots=True)
30
  class Asset:
31
  id: str
 
37
  checkpoint: str = ""
38
  epochs: int = 0
39
  created_at: str = ""
40
+ metadata: dict[str, Any] | None = None
41
 
42
  @classmethod
43
  def from_dict(cls, payload: dict[str, Any]) -> "Asset":
 
51
  checkpoint=str(payload.get("checkpoint", "")),
52
  epochs=int(payload.get("epochs", 0) or 0),
53
  created_at=str(payload.get("created_at") or _now()),
54
+ metadata=dict(payload.get("metadata") or {}),
55
  )
56
 
57
 
 
93
  dataset_id: str = "",
94
  checkpoint: str = "",
95
  epochs: int = 0,
96
+ metadata: dict[str, Any] | None = None,
97
  persist: bool = True,
98
  ) -> Asset:
99
  resolved = str(Path(path).expanduser().resolve())
 
112
  asset.checkpoint = checkpoint
113
  asset.epochs = int(epochs)
114
  asset.created_at = asset.created_at or _now()
115
+ if metadata:
116
+ current = dict(asset.metadata or {})
117
+ current.update(metadata)
118
+ asset.metadata = current
119
  if existing is None:
120
  self.assets.insert(0, asset)
121
  if persist:
 
134
  key: item[key]
135
  for key in (
136
  "kind", "name", "path", "trainer", "dataset_id",
137
+ "checkpoint", "epochs", "metadata",
138
  )
139
  if key in item
140
  }
141
+ if "trigger_word" in item:
142
+ metadata = dict(values.get("metadata") or {})
143
+ metadata["trigger_word"] = str(item.get("trigger_word") or "")
144
+ values["metadata"] = metadata
145
  dataset_path = str(item.get("dataset_path", ""))
146
  if item.get("kind") == "model" and dataset_path and Path(dataset_path).is_dir():
147
  dataset = self.register(
 
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"))
179
+ except (OSError, ValueError, TypeError, json.JSONDecodeError):
180
+ external = {}
181
+ for entry in external.get("tools", []) if isinstance(external, dict) else []:
182
+ if isinstance(entry, dict) and entry.get("id") == "external_oasis_game_trainer":
183
+ root = str(entry.get("backend", {}).get("root", ""))
184
+ if root:
185
+ folders["oasis_trainer"] = root
186
  external_lora_root = app_root / "LoRAModelsHere"
187
  if external_lora_root.is_dir():
188
  for path in external_lora_root.rglob("*.safetensors"):
 
225
  ("ddpm", "ddpm_trainer", "output"),
226
  ("lora", "lora_trainer", "output"),
227
  ("flow", "flow_trainer", "output_flow_models"),
228
+ ("oasis", "oasis_trainer", "output_action_flow_models"),
229
  ):
230
  root = Path(str(folders.get(folder_name, ""))) / output_name
231
  if not root.is_dir():
 
251
  else -1,
252
  )
253
  elif trainer == "lora":
254
+ trigger_word = ""
255
  checkpoints = sorted(
256
  (
257
  path for path in folder.glob("*.safetensors")
 
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:
272
  metadata = json.loads(
 
280
  dataset_path = flow_datasets.get(str(folder.resolve()), "")
281
  except (OSError, ValueError, TypeError, json.JSONDecodeError):
282
  continue
283
+ else:
284
+ checkpoints = []
285
+ try:
286
+ metadata = json.loads(
287
+ (folder / "action_flow_model_info.json").read_text(encoding="utf-8")
288
+ )
289
+ if metadata.get("model_type") != "action_conditioned_rectified_flow_video":
290
+ continue
291
+ if not (folder / "unet" / "config.json").is_file():
292
+ continue
293
+ name = _friendly_name(
294
+ str(metadata.get("model_name") or metadata.get("name") or name),
295
+ folder.name,
296
+ )
297
+ except (OSError, ValueError, TypeError, json.JSONDecodeError):
298
+ continue
299
  checkpoint = (
300
+ str(folder) if trainer in {"flow", "oasis"} else str(checkpoints[-1]) if checkpoints else ""
301
  )
302
  dataset_id = ""
303
  if dataset_path and Path(dataset_path).is_dir():
 
315
  trainer=trainer,
316
  dataset_id=dataset_id,
317
  checkpoint=checkpoint,
318
+ metadata=({"trigger_word": trigger_word or name} if trainer == "lora" else None),
319
  persist=False,
320
  )
321
+ try:
322
+ from adam.dataset_registry import DatasetRegistry
323
+
324
+ dataset_registry = DatasetRegistry(app_root, config)
325
+ valid_location_ids = {location.id for location in dataset_registry.known_locations()}
326
+ self.assets = [
327
+ item for item in self.assets
328
+ if not (
329
+ item.kind == "dataset"
330
+ and isinstance(item.metadata, dict)
331
+ and item.metadata.get("dataset_registry_source") in {"adam", "tool"}
332
+ and item.metadata.get("dataset_location_id")
333
+ and item.metadata.get("dataset_location_id") not in valid_location_ids
334
+ )
335
+ ]
336
+ dataset_registry.discover_into_assets(self, persist=False)
337
+ except Exception:
338
+ pass
339
  self.save()
340
 
341
  def _flow_dataset_paths(self) -> dict[str, str]:
adam/atlas.py CHANGED
@@ -2,7 +2,9 @@ from __future__ import annotations
2
 
3
  import math
4
  import re
 
5
  import time
 
6
  from dataclasses import dataclass
7
  from datetime import datetime
8
  from typing import Any
@@ -57,17 +59,35 @@ class AtlasSupervisor:
57
  return AtlasDecision("critical", f"GPU temperature remained at {temperature:.0f}°C. ATLAS paused the job.", "pause")
58
 
59
  free_disk = max(0.0, snapshot.disk_total_gb - snapshot.disk_used_gb)
60
- state["disk_samples"] = state["disk_samples"] + 1 if snapshot.disk_total_gb and free_disk <= self.critical_disk_gb else 0
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
61
  if state["disk_samples"] >= 2:
62
- return AtlasDecision("critical", f"Only {free_disk:.1f} GB remains on the output drive. ATLAS paused the job.", "pause")
63
 
64
  stalled_minutes = (moment - state["changed_at"]) / 60
65
  if stalled_minutes >= self.stall_minutes and snapshot.gpu_percent < 5:
66
  return AtlasDecision("warning", f"No recorded progress and little GPU activity for {stalled_minutes:.0f} minutes. Check the trainer.")
67
  if temperature is not None and temperature >= self.warning_temp:
68
  return AtlasDecision("warning", f"GPU temperature is elevated at {temperature:.0f}°C; ATLAS is watching it closely.")
69
- if snapshot.disk_total_gb and free_disk <= self.warning_disk_gb:
70
- return AtlasDecision("warning", f"Output drive space is getting low ({free_disk:.1f} GB free).")
 
 
71
  if snapshot.memory_percent >= 95:
72
  return AtlasDecision("warning", f"System memory usage is very high at {snapshot.memory_percent:.0f}%.")
73
 
 
2
 
3
  import math
4
  import re
5
+ import shutil
6
  import time
7
+ from pathlib import Path
8
  from dataclasses import dataclass
9
  from datetime import datetime
10
  from typing import Any
 
59
  return AtlasDecision("critical", f"GPU temperature remained at {temperature:.0f}°C. ATLAS paused the job.", "pause")
60
 
61
  free_disk = max(0.0, snapshot.disk_total_gb - snapshot.disk_used_gb)
62
+ disk_known = bool(snapshot.disk_total_gb)
63
+ disk_label = "monitored drive"
64
+ output = job.plan.steps[job.current_step].arguments.get("output_dir") or job.output_folder
65
+ disk_error = False
66
+ if output:
67
+ disk_label = "output drive"
68
+ try:
69
+ target = Path(str(output)).expanduser().resolve()
70
+ while not target.exists() and target != target.parent:
71
+ target = target.parent
72
+ usage = shutil.disk_usage(target)
73
+ free_disk = usage.free / (1024 ** 3)
74
+ disk_known = True
75
+ except (OSError, ValueError):
76
+ disk_known = False
77
+ disk_error = True
78
+ state["disk_samples"] = state["disk_samples"] + 1 if disk_known and free_disk <= self.critical_disk_gb else 0
79
  if state["disk_samples"] >= 2:
80
+ return AtlasDecision("critical", f"Only {free_disk:.1f} GB remains on the {disk_label}. ATLAS paused the job.", "pause")
81
 
82
  stalled_minutes = (moment - state["changed_at"]) / 60
83
  if stalled_minutes >= self.stall_minutes and snapshot.gpu_percent < 5:
84
  return AtlasDecision("warning", f"No recorded progress and little GPU activity for {stalled_minutes:.0f} minutes. Check the trainer.")
85
  if temperature is not None and temperature >= self.warning_temp:
86
  return AtlasDecision("warning", f"GPU temperature is elevated at {temperature:.0f}°C; ATLAS is watching it closely.")
87
+ if disk_error:
88
+ return AtlasDecision("warning", "ATLAS could not check the output drive's free space. Check that the output location is available.")
89
+ if disk_known and free_disk <= self.warning_disk_gb:
90
+ return AtlasDecision("warning", f"Space on the {disk_label} is getting low ({free_disk:.1f} GB free).")
91
  if snapshot.memory_percent >= 95:
92
  return AtlasDecision("warning", f"System memory usage is very high at {snapshot.memory_percent:.0f}%.")
93
 
adam/commands.py CHANGED
@@ -1,8 +1,11 @@
1
  from __future__ import annotations
2
 
3
  from dataclasses import dataclass
 
4
  from typing import Any
5
 
 
 
6
 
7
  class CommandValidationError(ValueError):
8
  pass
@@ -18,13 +21,14 @@ class TrainingCommand:
18
  output: str = "default output"
19
  resume_from: str = ""
20
  base_model: str = ""
 
21
  training_options: dict[str, Any] | None = None
22
 
23
  @classmethod
24
  def from_dict(cls, payload: dict[str, Any]) -> "TrainingCommand":
25
  allowed = {
26
  "action", "trainer", "dataset", "model_name", "epochs", "output",
27
- "resume_from", "base_model",
28
  "training_options",
29
  }
30
  unknown = set(payload) - allowed
@@ -42,42 +46,63 @@ class TrainingCommand:
42
  output=str(payload.get("output", "default output")).strip(),
43
  resume_from=str(payload.get("resume_from", "")).strip(),
44
  base_model=str(payload.get("base_model", "")).strip(),
 
 
 
 
 
45
  training_options=dict(payload.get("training_options") or {}),
46
  )
47
  except (TypeError, ValueError) as exc:
48
  raise CommandValidationError("Training command fields have invalid types.") from exc
49
  if command.action not in {"train", "resume_training"}:
50
  raise CommandValidationError("Training action must be train or resume_training.")
51
- if command.trainer not in {"ddpm", "lora", "flow"}:
52
- raise CommandValidationError("Trainer must be ddpm, flow, or lora.")
 
53
  if not command.dataset or not command.model_name:
54
  raise CommandValidationError("Dataset and model name are required.")
55
  if not 1 <= command.epochs <= 100_000:
56
  raise CommandValidationError("Epoch count must be between 1 and 100000.")
57
  if command.action == "resume_training" and not command.resume_from:
58
  raise CommandValidationError("Resume training requires an explicit checkpoint.")
 
 
 
 
59
  command._validate_options()
60
  return command
61
 
62
  def _validate_options(self) -> None:
63
  options = self.training_options or {}
64
- allowed = {
65
- "ddpm": {
66
- "resolution", "batch_size", "learning_rate", "gradient_accumulation_steps",
67
- "dataloader_num_workers", "mixed_precision", "save_every", "preview_steps",
68
- "training_intensity", "preview_enabled", "preview_every", "preview_prompt",
69
- "preview_seed",
70
- },
71
- "flow": {
72
- "resolution", "batch_size", "learning_rate", "gradient_accumulation",
73
- "workers", "mixed_precision", "save_every", "preview_every", "preview_steps",
74
- "gradient_checkpointing", "preview_enabled", "preview_prompt", "preview_seed",
75
- },
76
- "lora": {"preview_enabled", "preview_every", "preview_prompt", "preview_seed"},
77
- }[self.trainer]
 
 
 
78
  unknown = set(options) - allowed
79
  if unknown:
80
  raise CommandValidationError(f"Unsupported {self.trainer} training options: {', '.join(sorted(unknown))}")
 
 
 
 
 
 
 
 
81
  integer_ranges = {
82
  "resolution": (64, 512), "batch_size": (1, 64),
83
  "gradient_accumulation_steps": (1, 64), "gradient_accumulation": (1, 64),
 
1
  from __future__ import annotations
2
 
3
  from dataclasses import dataclass
4
+ from pathlib import Path
5
  from typing import Any
6
 
7
+ from adam.model_plugins import ModelPluginRegistry, validate_settings
8
+
9
 
10
  class CommandValidationError(ValueError):
11
  pass
 
21
  output: str = "default output"
22
  resume_from: str = ""
23
  base_model: str = ""
24
+ trigger_word: str = ""
25
  training_options: dict[str, Any] | None = None
26
 
27
  @classmethod
28
  def from_dict(cls, payload: dict[str, Any]) -> "TrainingCommand":
29
  allowed = {
30
  "action", "trainer", "dataset", "model_name", "epochs", "output",
31
+ "resume_from", "base_model", "trigger_word",
32
  "training_options",
33
  }
34
  unknown = set(payload) - allowed
 
46
  output=str(payload.get("output", "default output")).strip(),
47
  resume_from=str(payload.get("resume_from", "")).strip(),
48
  base_model=str(payload.get("base_model", "")).strip(),
49
+ trigger_word=str(
50
+ payload.get("trigger_word")
51
+ or (payload.get("training_options") or {}).get("trigger_word")
52
+ or ""
53
+ ).strip(),
54
  training_options=dict(payload.get("training_options") or {}),
55
  )
56
  except (TypeError, ValueError) as exc:
57
  raise CommandValidationError("Training command fields have invalid types.") from exc
58
  if command.action not in {"train", "resume_training"}:
59
  raise CommandValidationError("Training action must be train or resume_training.")
60
+ plugin_schema = ModelPluginRegistry(Path.cwd()).training_schema(command.trainer)
61
+ if command.trainer not in {"ddpm", "lora", "flow"} and not plugin_schema:
62
+ raise CommandValidationError("Trainer must be a discovered model plugin.")
63
  if not command.dataset or not command.model_name:
64
  raise CommandValidationError("Dataset and model name are required.")
65
  if not 1 <= command.epochs <= 100_000:
66
  raise CommandValidationError("Epoch count must be between 1 and 100000.")
67
  if command.action == "resume_training" and not command.resume_from:
68
  raise CommandValidationError("Resume training requires an explicit checkpoint.")
69
+ if command.trainer == "lora":
70
+ trigger = command.trigger_word or command.model_name
71
+ if len(trigger) > 128 or any(char in trigger for char in '<>:"/\\|?*\x00'):
72
+ raise CommandValidationError("LoRA trigger word must be short text without reserved characters.")
73
  command._validate_options()
74
  return command
75
 
76
  def _validate_options(self) -> None:
77
  options = self.training_options or {}
78
+ schema = ModelPluginRegistry(Path.cwd()).training_schema(self.trainer)
79
+ allowed = set(schema)
80
+ if not allowed:
81
+ allowed = {
82
+ "ddpm": {
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",
90
+ "workers", "mixed_precision", "save_every", "preview_every", "preview_steps",
91
+ "gradient_checkpointing", "preview_enabled", "preview_prompt", "preview_seed",
92
+ },
93
+ "lora": {"preview_enabled", "preview_every", "preview_prompt", "preview_seed", "trigger_word"},
94
+ }[self.trainer]
95
  unknown = set(options) - allowed
96
  if unknown:
97
  raise CommandValidationError(f"Unsupported {self.trainer} training options: {', '.join(sorted(unknown))}")
98
+ if schema:
99
+ errors = validate_settings(
100
+ {key: spec for key, spec in schema.items() if key in options},
101
+ options,
102
+ )
103
+ if errors:
104
+ raise CommandValidationError(" ".join(errors))
105
+ return
106
  integer_ranges = {
107
  "resolution": (64, 512), "batch_size": (1, 64),
108
  "gradient_accumulation_steps": (1, 64), "gradient_accumulation": (1, 64),
adam/config.py CHANGED
@@ -26,12 +26,21 @@ DEFAULT_SETTINGS: dict[str, Any] = {
26
  "demo_step_delay": 0.24,
27
  "max_dataset_images_without_confirmation": 100,
28
  "training_presets": {},
 
 
 
 
 
 
 
 
29
  "tool_folders": {
30
  "dataset_collector": "",
31
  "caption_generator": "",
32
  "lora_trainer": "",
33
  "ddpm_trainer": "",
34
  "flow_trainer": "",
 
35
  "preview_generator": "",
36
  },
37
  }
 
26
  "demo_step_delay": 0.24,
27
  "max_dataset_images_without_confirmation": 100,
28
  "training_presets": {},
29
+ "remote_access": {
30
+ "enabled": False,
31
+ "bind_address": "127.0.0.1",
32
+ "port": 8765,
33
+ "token": "",
34
+ "allow_job_control": False,
35
+ "auto_approve_training": False,
36
+ },
37
  "tool_folders": {
38
  "dataset_collector": "",
39
  "caption_generator": "",
40
  "lora_trainer": "",
41
  "ddpm_trainer": "",
42
  "flow_trainer": "",
43
+ "oasis_trainer": "",
44
  "preview_generator": "",
45
  },
46
  }
adam/dataset_lab.py ADDED
@@ -0,0 +1,127 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from __future__ import annotations
2
+
3
+ import hashlib
4
+ from dataclasses import asdict, dataclass, field
5
+ from pathlib import Path
6
+
7
+
8
+ IMAGE_EXTENSIONS = {".jpg", ".jpeg", ".png", ".webp", ".bmp"}
9
+ VIDEO_EXTENSIONS = {".mp4", ".mov", ".mkv", ".webm", ".avi"}
10
+ TEXT_EXTENSIONS = {".txt", ".caption", ".jsonl", ".json"}
11
+
12
+
13
+ @dataclass(slots=True)
14
+ class DatasetItem:
15
+ path: str
16
+ kind: str
17
+ size_bytes: int
18
+ width: int = 0
19
+ height: int = 0
20
+ caption_path: str = ""
21
+ duplicate_key: str = ""
22
+
23
+
24
+ @dataclass(slots=True)
25
+ class DatasetReport:
26
+ path: str
27
+ total_files: int = 0
28
+ image_count: int = 0
29
+ video_count: int = 0
30
+ text_count: int = 0
31
+ caption_count: int = 0
32
+ missing_caption_count: int = 0
33
+ duplicate_groups: int = 0
34
+ dimensions: dict[str, int] = field(default_factory=dict)
35
+ extensions: dict[str, int] = field(default_factory=dict)
36
+ items: list[DatasetItem] = field(default_factory=list)
37
+ warnings: list[str] = field(default_factory=list)
38
+
39
+ def to_dict(self) -> dict:
40
+ return asdict(self)
41
+
42
+
43
+ def _hash_file(path: Path) -> str:
44
+ digest = hashlib.sha1()
45
+ with path.open("rb") as handle:
46
+ for chunk in iter(lambda: handle.read(1024 * 1024), b""):
47
+ digest.update(chunk)
48
+ return digest.hexdigest()
49
+
50
+
51
+ def scan_dataset(path: str | Path, *, limit: int = 500) -> DatasetReport:
52
+ root = Path(path).expanduser().resolve()
53
+ report = DatasetReport(path=str(root))
54
+ if not root.is_dir():
55
+ report.warnings.append("Dataset folder does not exist.")
56
+ return report
57
+ hashes: dict[str, int] = {}
58
+ try:
59
+ files = [item for item in root.rglob("*") if item.is_file()]
60
+ except OSError as exc:
61
+ report.warnings.append(f"Dataset could not be scanned: {exc}")
62
+ return report
63
+ report.total_files = len(files)
64
+ for item in files:
65
+ suffix = item.suffix.casefold()
66
+ report.extensions[suffix or "(none)"] = report.extensions.get(suffix or "(none)", 0) + 1
67
+ if suffix in IMAGE_EXTENSIONS:
68
+ report.image_count += 1
69
+ if not any(item.with_suffix(ext).is_file() for ext in (".txt", ".caption")):
70
+ report.missing_caption_count += 1
71
+ elif suffix in VIDEO_EXTENSIONS:
72
+ report.video_count += 1
73
+ if suffix in TEXT_EXTENSIONS:
74
+ report.text_count += 1
75
+ if suffix in {".txt", ".caption"}:
76
+ report.caption_count += 1
77
+ for item in files[: max(1, limit)]:
78
+ suffix = item.suffix.casefold()
79
+ kind = "other"
80
+ width = height = 0
81
+ caption_path = ""
82
+ duplicate_key = ""
83
+ if suffix in IMAGE_EXTENSIONS:
84
+ kind = "image"
85
+ caption = next((item.with_suffix(ext) for ext in (".txt", ".caption") if item.with_suffix(ext).is_file()), None)
86
+ caption_path = str(caption) if caption else ""
87
+ try:
88
+ from PIL import Image
89
+
90
+ with Image.open(item) as image:
91
+ width, height = image.size
92
+ label = f"{width}x{height}"
93
+ report.dimensions[label] = report.dimensions.get(label, 0) + 1
94
+ except Exception:
95
+ pass
96
+ try:
97
+ duplicate_key = _hash_file(item)
98
+ hashes[duplicate_key] = hashes.get(duplicate_key, 0) + 1
99
+ except OSError:
100
+ duplicate_key = ""
101
+ elif suffix in VIDEO_EXTENSIONS:
102
+ kind = "video"
103
+ elif suffix in TEXT_EXTENSIONS:
104
+ kind = "text"
105
+ try:
106
+ size = item.stat().st_size
107
+ except OSError:
108
+ size = 0
109
+ report.items.append(
110
+ DatasetItem(
111
+ path=str(item),
112
+ kind=kind,
113
+ size_bytes=size,
114
+ width=width,
115
+ height=height,
116
+ caption_path=caption_path,
117
+ duplicate_key=duplicate_key,
118
+ )
119
+ )
120
+ if report.total_files > limit:
121
+ report.warnings.append(f"Showing first {limit:,} files; totals still include all files.")
122
+ report.duplicate_groups = sum(1 for count in hashes.values() if count > 1)
123
+ if report.image_count and report.missing_caption_count:
124
+ report.warnings.append(f"{report.missing_caption_count:,} sampled image(s) do not have sidecar captions.")
125
+ if report.duplicate_groups:
126
+ report.warnings.append(f"{report.duplicate_groups:,} duplicate image group(s) found in the sample.")
127
+ return report
adam/dataset_registry.py ADDED
@@ -0,0 +1,484 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from __future__ import annotations
2
+
3
+ import hashlib
4
+ import json
5
+ import threading
6
+ from dataclasses import asdict, dataclass, field
7
+ from datetime import datetime, timedelta, timezone
8
+ from pathlib import Path
9
+ from typing import Any, TYPE_CHECKING
10
+ from uuid import uuid4
11
+
12
+ from adam.dataset_lab import IMAGE_EXTENSIONS, TEXT_EXTENSIONS, VIDEO_EXTENSIONS
13
+
14
+ if TYPE_CHECKING:
15
+ from adam.assets import Asset, AssetRegistry
16
+
17
+
18
+ DATASET_MARKERS = {
19
+ "dataset_manifest.json",
20
+ "metadata",
21
+ "frames",
22
+ "videos",
23
+ "captions",
24
+ "actions.jsonl",
25
+ "actions.csv",
26
+ }
27
+ SCAN_LIMIT = 20_000
28
+ ASYNC_REFRESH_AFTER = timedelta(minutes=30)
29
+ _scan_lock = threading.Lock()
30
+ _active_scans: set[str] = set()
31
+
32
+
33
+ def _now() -> str:
34
+ return datetime.now(timezone.utc).isoformat()
35
+
36
+
37
+ def _path_key(path: str | Path) -> str:
38
+ resolved = str(Path(path).expanduser().resolve())
39
+ return hashlib.sha1(resolved.casefold().encode("utf-8")).hexdigest()[:16]
40
+
41
+
42
+ def _atomic_json(path: Path, payload: dict[str, Any]) -> None:
43
+ path.parent.mkdir(parents=True, exist_ok=True)
44
+ temporary = path.with_suffix(path.suffix + ".tmp")
45
+ temporary.write_text(json.dumps(payload, indent=2), encoding="utf-8")
46
+ temporary.replace(path)
47
+
48
+
49
+ def _date(value: str) -> datetime | None:
50
+ try:
51
+ return datetime.fromisoformat(value)
52
+ except (TypeError, ValueError):
53
+ return None
54
+
55
+
56
+ @dataclass(slots=True)
57
+ class DatasetLocation:
58
+ id: str
59
+ name: str
60
+ path: str
61
+ source: str = "user"
62
+ created_at: str = field(default_factory=_now)
63
+ last_seen_at: str = ""
64
+ exists: bool = True
65
+
66
+ @classmethod
67
+ def from_dict(cls, payload: dict[str, Any]) -> "DatasetLocation":
68
+ return cls(
69
+ id=str(payload.get("id") or _path_key(str(payload.get("path", "")))),
70
+ name=str(payload.get("name") or Path(str(payload.get("path", ""))).name or "Datasets"),
71
+ path=str(payload.get("path", "")),
72
+ source=str(payload.get("source") or "user"),
73
+ created_at=str(payload.get("created_at") or _now()),
74
+ last_seen_at=str(payload.get("last_seen_at") or ""),
75
+ exists=bool(payload.get("exists", True)),
76
+ )
77
+
78
+
79
+ @dataclass(slots=True)
80
+ class DatasetRecord:
81
+ id: str
82
+ name: str
83
+ path: str
84
+ source: str = "asset"
85
+ location_id: str = ""
86
+ favorite: bool = False
87
+ last_used_at: str = ""
88
+ discovered_at: str = field(default_factory=_now)
89
+ scanned_at: str = ""
90
+ exists: bool = True
91
+ item_count: int = 0
92
+ image_count: int = 0
93
+ video_count: int = 0
94
+ caption_count: int = 0
95
+ missing_caption_count: int = 0
96
+ sample_image: str = ""
97
+ dataset_format: str = "Unknown"
98
+ warnings: list[str] = field(default_factory=list)
99
+
100
+ @classmethod
101
+ def from_dict(cls, payload: dict[str, Any]) -> "DatasetRecord":
102
+ return cls(
103
+ id=str(payload.get("id") or _path_key(str(payload.get("path", "")))),
104
+ name=str(payload.get("name") or Path(str(payload.get("path", ""))).name or "Dataset"),
105
+ path=str(payload.get("path", "")),
106
+ source=str(payload.get("source") or "asset"),
107
+ location_id=str(payload.get("location_id") or ""),
108
+ favorite=bool(payload.get("favorite", False)),
109
+ last_used_at=str(payload.get("last_used_at") or ""),
110
+ discovered_at=str(payload.get("discovered_at") or _now()),
111
+ scanned_at=str(payload.get("scanned_at") or ""),
112
+ exists=bool(payload.get("exists", True)),
113
+ item_count=int(payload.get("item_count", 0) or 0),
114
+ image_count=int(payload.get("image_count", 0) or 0),
115
+ video_count=int(payload.get("video_count", 0) or 0),
116
+ caption_count=int(payload.get("caption_count", 0) or 0),
117
+ missing_caption_count=int(payload.get("missing_caption_count", 0) or 0),
118
+ sample_image=str(payload.get("sample_image") or ""),
119
+ dataset_format=str(payload.get("dataset_format") or "Unknown"),
120
+ warnings=[str(item) for item in payload.get("warnings", []) if str(item)],
121
+ )
122
+
123
+
124
+ class DatasetRegistry:
125
+ """Persistent ADAM-aware index of known dataset locations and datasets."""
126
+
127
+ def __init__(self, root: Path, config: Any | None = None) -> None:
128
+ self.root = root.resolve()
129
+ self.path = self.root / "data" / "dataset_registry.json"
130
+ self.config = config
131
+ self.locations: list[DatasetLocation] = []
132
+ self.datasets: dict[str, DatasetRecord] = {}
133
+ self.load()
134
+
135
+ def load(self) -> None:
136
+ try:
137
+ payload = json.loads(self.path.read_text(encoding="utf-8"))
138
+ except (OSError, ValueError, TypeError, json.JSONDecodeError):
139
+ payload = {}
140
+ self.locations = [
141
+ DatasetLocation.from_dict(item)
142
+ for item in payload.get("locations", [])
143
+ if isinstance(item, dict)
144
+ ]
145
+ self.datasets = {
146
+ str(key): DatasetRecord.from_dict(value)
147
+ for key, value in (payload.get("datasets", {}) or {}).items()
148
+ if isinstance(value, dict)
149
+ }
150
+
151
+ def save(self) -> None:
152
+ _atomic_json(
153
+ self.path,
154
+ {
155
+ "locations": [asdict(item) for item in self.locations],
156
+ "datasets": {key: asdict(value) for key, value in self.datasets.items()},
157
+ },
158
+ )
159
+
160
+ def register_location(self, path: str | Path, *, name: str = "", source: str = "user") -> DatasetLocation:
161
+ resolved = Path(path).expanduser().resolve()
162
+ if not resolved.is_dir():
163
+ raise ValueError("Choose an existing dataset folder.")
164
+ key = _path_key(resolved)
165
+ existing = next((item for item in self.locations if item.id == key), None)
166
+ if existing is None:
167
+ existing = DatasetLocation(
168
+ id=key,
169
+ name=name.strip() or resolved.name or str(resolved),
170
+ path=str(resolved),
171
+ source=source,
172
+ )
173
+ self.locations.insert(0, existing)
174
+ else:
175
+ existing.name = name.strip() or existing.name
176
+ existing.path = str(resolved)
177
+ existing.source = source or existing.source
178
+ existing.exists = True
179
+ existing.last_seen_at = _now()
180
+ self.save()
181
+ return existing
182
+
183
+ def remove_location(self, location_id: str) -> bool:
184
+ before = len(self.locations)
185
+ self.locations = [item for item in self.locations if item.id != location_id]
186
+ changed = len(self.locations) != before
187
+ if changed:
188
+ self.save()
189
+ return changed
190
+
191
+ def known_locations(self) -> list[DatasetLocation]:
192
+ locations = list(self.locations)
193
+ by_path = {Path(item.path).expanduser().resolve(): item for item in locations if item.path}
194
+ for path, name, source in self._automatic_location_candidates():
195
+ try:
196
+ resolved = path.expanduser().resolve()
197
+ except OSError:
198
+ continue
199
+ if resolved in by_path:
200
+ continue
201
+ locations.append(
202
+ DatasetLocation(
203
+ id=_path_key(resolved),
204
+ name=name or resolved.name or str(resolved),
205
+ path=str(resolved),
206
+ source=source,
207
+ exists=resolved.is_dir(),
208
+ last_seen_at=_now() if resolved.is_dir() else "",
209
+ )
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
219
+ asset = assets.register(
220
+ kind="dataset",
221
+ name=record.name,
222
+ path=record.path,
223
+ metadata={
224
+ "dataset_registry_source": record.source,
225
+ "dataset_location_id": record.location_id,
226
+ },
227
+ persist=False,
228
+ )
229
+ discovered.append(asset)
230
+ if persist and discovered:
231
+ assets.save()
232
+ return discovered
233
+
234
+ def discover(
235
+ self,
236
+ *,
237
+ asset_registry: "AssetRegistry | None" = None,
238
+ refresh_missing: bool = True,
239
+ ) -> list[DatasetRecord]:
240
+ self.load()
241
+ changed = False
242
+ locations = self.known_locations()
243
+ known_location_ids = {location.id for location in locations}
244
+ for key, record in list(self.datasets.items()):
245
+ if record.source in {"adam", "tool"} and record.location_id and record.location_id not in known_location_ids:
246
+ del self.datasets[key]
247
+ changed = True
248
+ for location in locations:
249
+ exists = Path(location.path).is_dir()
250
+ if location.source == "user":
251
+ stored = next((item for item in self.locations if item.id == location.id), None)
252
+ if stored:
253
+ stored.exists = exists
254
+ stored.last_seen_at = _now() if exists else stored.last_seen_at
255
+ changed = True
256
+ if not exists:
257
+ continue
258
+ for candidate in self._dataset_candidates(Path(location.path)):
259
+ record = self._cached_or_sampled(candidate, source=location.source, location_id=location.id)
260
+ self.datasets[record.id] = record
261
+ changed = True
262
+ if asset_registry is not None:
263
+ for asset in getattr(asset_registry, "assets", []):
264
+ if getattr(asset, "kind", "") != "dataset":
265
+ continue
266
+ record = self._cached_or_sampled(Path(asset.path), source="asset", location_id="")
267
+ record.name = asset.name or record.name
268
+ self.datasets[record.id] = record
269
+ changed = True
270
+ for record in self.datasets.values():
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
+
278
+ def sorted_records(self) -> list[DatasetRecord]:
279
+ records = list(self.datasets.values())
280
+ records.sort(
281
+ key=lambda item: (
282
+ not item.favorite,
283
+ not bool(item.last_used_at),
284
+ item.last_used_at or item.discovered_at,
285
+ item.name.casefold(),
286
+ ),
287
+ reverse=False,
288
+ )
289
+ favorites = sorted([item for item in records if item.favorite], key=lambda item: item.name.casefold())
290
+ recent = sorted(
291
+ [item for item in records if not item.favorite and item.last_used_at],
292
+ key=lambda item: item.last_used_at,
293
+ reverse=True,
294
+ )
295
+ others = sorted(
296
+ [item for item in records if not item.favorite and not item.last_used_at],
297
+ key=lambda item: item.discovered_at,
298
+ reverse=True,
299
+ )
300
+ return [*favorites, *recent, *others]
301
+
302
+ def record_for_path(self, path: str | Path) -> DatasetRecord:
303
+ key = _path_key(path)
304
+ record = self.datasets.get(key)
305
+ if record is None:
306
+ record = self._cached_or_sampled(Path(path), source="asset", location_id="")
307
+ self.datasets[key] = record
308
+ self.save()
309
+ return record
310
+
311
+ def favorite(self, path: str | Path, enabled: bool) -> DatasetRecord:
312
+ record = self.record_for_path(path)
313
+ record.favorite = bool(enabled)
314
+ self.save()
315
+ return record
316
+
317
+ def touch(self, path: str | Path) -> DatasetRecord:
318
+ record = self.record_for_path(path)
319
+ record.last_used_at = _now()
320
+ self.save()
321
+ return record
322
+
323
+ def refresh_async(self, path: str | Path, *, source: str = "asset", location_id: str = "") -> None:
324
+ resolved = str(Path(path).expanduser().resolve())
325
+ key = _path_key(resolved)
326
+ with _scan_lock:
327
+ if key in _active_scans:
328
+ return
329
+ _active_scans.add(key)
330
+
331
+ def worker() -> None:
332
+ try:
333
+ record = self._scan(Path(resolved), source=source, location_id=location_id, limit=SCAN_LIMIT)
334
+ fresh = DatasetRegistry(self.root, self.config)
335
+ current = fresh.datasets.get(record.id)
336
+ if current:
337
+ record.favorite = current.favorite
338
+ record.last_used_at = current.last_used_at
339
+ record.discovered_at = current.discovered_at
340
+ fresh.datasets[record.id] = record
341
+ fresh.save()
342
+ finally:
343
+ with _scan_lock:
344
+ _active_scans.discard(key)
345
+
346
+ threading.Thread(target=worker, name="ADAMDatasetScan", daemon=True).start()
347
+
348
+ def _automatic_location_candidates(self) -> list[tuple[Path, str, str]]:
349
+ candidates: list[tuple[Path, str, str]] = [(self.root / "ADAM_Datasets", "ADAM Datasets", "adam")]
350
+ folders = self.config.get("tool_folders", {}) if self.config is not None else {}
351
+ folders = folders if isinstance(folders, dict) else {}
352
+ collector_raw = str(folders.get("dataset_collector", "")).strip()
353
+ if collector_raw:
354
+ candidates.append((Path(collector_raw) / "Datasets", "Dataset Collector", "tool"))
355
+ for tool_id in ("ddpm_trainer", "lora_trainer", "flow_trainer", "oasis_trainer"):
356
+ raw_root = str(folders.get(tool_id, "")).strip()
357
+ if not raw_root:
358
+ continue
359
+ root = Path(raw_root)
360
+ label = tool_id.removesuffix("_trainer").upper()
361
+ for child in ("Datasets", "datasets", "OldDatasets"):
362
+ candidates.append((root / child, f"{label} {child}", "tool"))
363
+ external_tools = self.root / "config" / "external_tools.json"
364
+ try:
365
+ payload = json.loads(external_tools.read_text(encoding="utf-8"))
366
+ except (OSError, ValueError, TypeError, json.JSONDecodeError):
367
+ payload = {}
368
+ for entry in payload.get("tools", []) if isinstance(payload, dict) else []:
369
+ if not isinstance(entry, dict):
370
+ continue
371
+ backend = entry.get("backend", {})
372
+ if not isinstance(backend, dict):
373
+ continue
374
+ raw_root = str(backend.get("root", "")).strip()
375
+ if raw_root:
376
+ root = Path(raw_root)
377
+ candidates.append((root / "Datasets", f"{entry.get('name', 'External tool')} Datasets", "tool"))
378
+ candidates.append((root / "OldDatasets", f"{entry.get('name', 'External tool')} OldDatasets", "tool"))
379
+ return candidates
380
+
381
+ def _dataset_candidates(self, location: Path) -> list[Path]:
382
+ candidates: list[Path] = []
383
+ try:
384
+ children = [item for item in location.iterdir() if item.is_dir()]
385
+ except OSError:
386
+ children = []
387
+ for child in children[:1000]:
388
+ if self._looks_like_dataset(child):
389
+ candidates.append(child)
390
+ if candidates:
391
+ return candidates
392
+ if self._looks_like_dataset(location):
393
+ candidates.append(location)
394
+ return candidates
395
+
396
+ def _looks_like_dataset(self, folder: Path) -> bool:
397
+ if not folder.is_dir():
398
+ return False
399
+ try:
400
+ names = {item.name for item in folder.iterdir()}
401
+ except OSError:
402
+ return False
403
+ if names & DATASET_MARKERS:
404
+ return True
405
+ checked = 0
406
+ for item in folder.rglob("*"):
407
+ if checked >= 200:
408
+ break
409
+ checked += 1
410
+ if item.is_file() and item.suffix.casefold() in IMAGE_EXTENSIONS | VIDEO_EXTENSIONS:
411
+ return True
412
+ return False
413
+
414
+ def _cached_or_sampled(self, path: Path, *, source: str, location_id: str) -> DatasetRecord:
415
+ key = _path_key(path)
416
+ current = self.datasets.get(key)
417
+ if current is not None:
418
+ current.exists = path.is_dir()
419
+ current.source = current.source or source
420
+ current.location_id = current.location_id or location_id
421
+ return current
422
+ return self._scan(path, source=source, location_id=location_id, limit=800)
423
+
424
+ def _needs_refresh(self, record: DatasetRecord) -> bool:
425
+ if not record.scanned_at:
426
+ return True
427
+ scanned = _date(record.scanned_at)
428
+ return scanned is None or datetime.now(timezone.utc) - scanned > ASYNC_REFRESH_AFTER
429
+
430
+ def _scan(self, path: Path, *, source: str, location_id: str, limit: int) -> DatasetRecord:
431
+ resolved = path.expanduser().resolve()
432
+ record = DatasetRecord(
433
+ id=_path_key(resolved),
434
+ name=resolved.name or "Dataset",
435
+ path=str(resolved),
436
+ source=source,
437
+ location_id=location_id,
438
+ exists=resolved.is_dir(),
439
+ scanned_at=_now(),
440
+ )
441
+ if not resolved.is_dir():
442
+ record.warnings.append("Dataset folder is unavailable.")
443
+ return record
444
+ files_seen = 0
445
+ capped = False
446
+ try:
447
+ iterator = resolved.rglob("*")
448
+ for item in iterator:
449
+ if not item.is_file():
450
+ continue
451
+ files_seen += 1
452
+ suffix = item.suffix.casefold()
453
+ if suffix in IMAGE_EXTENSIONS:
454
+ record.image_count += 1
455
+ if not record.sample_image:
456
+ record.sample_image = str(item)
457
+ if not any(item.with_suffix(ext).is_file() for ext in (".txt", ".caption")):
458
+ record.missing_caption_count += 1
459
+ elif suffix in VIDEO_EXTENSIONS:
460
+ record.video_count += 1
461
+ if suffix in TEXT_EXTENSIONS and suffix in {".txt", ".caption"}:
462
+ record.caption_count += 1
463
+ if files_seen >= limit:
464
+ capped = True
465
+ break
466
+ except OSError as exc:
467
+ record.warnings.append(f"Dataset could not be scanned: {exc}")
468
+ record.item_count = record.image_count + record.video_count
469
+ record.dataset_format = self._format_label(resolved, record)
470
+ if capped:
471
+ record.warnings.append(f"Counts are sampled from the first {limit:,} files.")
472
+ return record
473
+
474
+ @staticmethod
475
+ def _format_label(folder: Path, record: DatasetRecord) -> str:
476
+ if (folder / "actions.jsonl").is_file() or (folder / "actions.csv").is_file():
477
+ return "Oasis action dataset"
478
+ if record.video_count:
479
+ return "Video dataset"
480
+ if record.image_count and record.caption_count:
481
+ return "Captioned image dataset"
482
+ if record.image_count:
483
+ return "Image dataset"
484
+ return "Dataset folder"
adam/eve.py CHANGED
@@ -87,8 +87,9 @@ def classify_eve_embeddings(
87
  class EveVisionModel:
88
  """Lazy local DINOv2 feature extractor used by EVE."""
89
 
90
- def __init__(self, model_id: str = EVE_MODEL_ID) -> None:
91
  self.model_id = model_id
 
92
  self._processor = None
93
  self._model = None
94
  self._device = "cpu"
@@ -106,10 +107,20 @@ class EveVisionModel:
106
  raise RuntimeError(
107
  "EVE needs PyTorch and Transformers. Launch ADAM with its normal Python environment."
108
  ) from exc
109
- self._device = "cuda" if torch.cuda.is_available() else "cpu"
110
  self._processor = AutoImageProcessor.from_pretrained(self.model_id, use_fast=True)
111
  self._model = AutoModel.from_pretrained(self.model_id).to(self._device).eval()
112
 
 
 
 
 
 
 
 
 
 
 
113
  def embed(self, paths: Sequence[str | Path], progress: Callable[[int, int], None] | None = None) -> list[list[float]]:
114
  self.load()
115
  import torch
 
87
  class EveVisionModel:
88
  """Lazy local DINOv2 feature extractor used by EVE."""
89
 
90
+ def __init__(self, model_id: str = EVE_MODEL_ID, *, prefer_gpu: bool = True) -> None:
91
  self.model_id = model_id
92
+ self.prefer_gpu = prefer_gpu
93
  self._processor = None
94
  self._model = None
95
  self._device = "cpu"
 
107
  raise RuntimeError(
108
  "EVE needs PyTorch and Transformers. Launch ADAM with its normal Python environment."
109
  ) from exc
110
+ self._device = "cuda" if self.prefer_gpu and torch.cuda.is_available() else "cpu"
111
  self._processor = AutoImageProcessor.from_pretrained(self.model_id, use_fast=True)
112
  self._model = AutoModel.from_pretrained(self.model_id).to(self._device).eval()
113
 
114
+ def unload(self) -> None:
115
+ self._processor = None
116
+ self._model = None
117
+ try:
118
+ import torch
119
+ if torch.cuda.is_available():
120
+ torch.cuda.empty_cache()
121
+ except ImportError:
122
+ pass
123
+
124
  def embed(self, paths: Sequence[str | Path], progress: Callable[[int, int], None] | None = None) -> list[list[float]]:
125
  self.load()
126
  import torch
adam/executor.py CHANGED
@@ -12,6 +12,7 @@ from dataclasses import dataclass
12
  from pathlib import Path
13
  from typing import Any, Callable
14
 
 
15
  from adam.registry import ToolRegistry, ToolSpec
16
 
17
 
@@ -23,7 +24,15 @@ class ToolCancelled(ToolExecutionError):
23
  pass
24
 
25
 
26
- ProgressCallback = Callable[[int, str], None]
 
 
 
 
 
 
 
 
27
  LogCallback = Callable[[str], None]
28
  PreviewCallback = Callable[[dict[str, Any]], None]
29
 
@@ -38,14 +47,16 @@ class ToolContext:
38
  progress_callback: ProgressCallback
39
  log_callback: LogCallback
40
  preview_callback: PreviewCallback = lambda _preview: None
 
 
41
  step_delay: float = 0.2
42
 
43
  def log(self, message: str) -> None:
44
  self.log_callback(message)
45
 
46
- def progress(self, percent: int, message: str) -> None:
47
  self.checkpoint()
48
- self.progress_callback(max(0, min(int(percent), 100)), message)
49
 
50
  def preview(
51
  self, path: str | Path, *, epoch: int = 0, next_epoch: int = 0,
@@ -104,6 +115,8 @@ class ToolExecutor:
104
  progress_callback: ProgressCallback,
105
  log_callback: LogCallback,
106
  preview_callback: PreviewCallback | None = None,
 
 
107
  ) -> dict[str, Any]:
108
  spec = self.registry.get(tool_id)
109
  self._validate_arguments(spec, arguments)
@@ -116,6 +129,8 @@ class ToolExecutor:
116
  progress_callback=progress_callback,
117
  log_callback=log_callback,
118
  preview_callback=preview_callback or (lambda _preview: None),
 
 
119
  step_delay=self.step_delay,
120
  )
121
  backend_type = str(spec.backend.get("type", "")).lower()
@@ -243,30 +258,7 @@ class ToolExecutor:
243
 
244
  def stop_process_tree() -> None:
245
  """Stop the script and any workers it launched."""
246
- if process_controller is not None:
247
- try:
248
- descendants = process_controller.children(recursive=True)
249
- for child in descendants:
250
- try:
251
- child.terminate()
252
- except Exception:
253
- pass
254
- process_controller.terminate()
255
- try:
256
- import psutil
257
-
258
- _gone, alive = psutil.wait_procs(descendants, timeout=2)
259
- for child in alive:
260
- try:
261
- child.kill()
262
- except Exception:
263
- pass
264
- except Exception:
265
- pass
266
- return
267
- except Exception:
268
- pass
269
- process.terminate()
270
 
271
  try:
272
  while True:
 
12
  from pathlib import Path
13
  from typing import Any, Callable
14
 
15
+ from adam.process_control import terminate_process_tree
16
  from adam.registry import ToolRegistry, ToolSpec
17
 
18
 
 
24
  pass
25
 
26
 
27
+ class ToolAdjustmentRequested(ToolExecutionError):
28
+ """A trainer stopped cleanly so a job can continue with new settings."""
29
+
30
+ def __init__(self, message: str, details: dict[str, Any] | None = None) -> None:
31
+ super().__init__(message)
32
+ self.details = details or {}
33
+
34
+
35
+ ProgressCallback = Callable[..., None]
36
  LogCallback = Callable[[str], None]
37
  PreviewCallback = Callable[[dict[str, Any]], None]
38
 
 
47
  progress_callback: ProgressCallback
48
  log_callback: LogCallback
49
  preview_callback: PreviewCallback = lambda _preview: None
50
+ adjustment_event: threading.Event | None = None
51
+ adjustment_request: dict[str, Any] | None = None
52
  step_delay: float = 0.2
53
 
54
  def log(self, message: str) -> None:
55
  self.log_callback(message)
56
 
57
+ def progress(self, percent: int, message: str, **details: Any) -> None:
58
  self.checkpoint()
59
+ self.progress_callback(max(0, min(int(percent), 100)), message, **details)
60
 
61
  def preview(
62
  self, path: str | Path, *, epoch: int = 0, next_epoch: int = 0,
 
115
  progress_callback: ProgressCallback,
116
  log_callback: LogCallback,
117
  preview_callback: PreviewCallback | None = None,
118
+ adjustment_event: threading.Event | None = None,
119
+ adjustment_request: dict[str, Any] | None = None,
120
  ) -> dict[str, Any]:
121
  spec = self.registry.get(tool_id)
122
  self._validate_arguments(spec, arguments)
 
129
  progress_callback=progress_callback,
130
  log_callback=log_callback,
131
  preview_callback=preview_callback or (lambda _preview: None),
132
+ adjustment_event=adjustment_event,
133
+ adjustment_request=adjustment_request,
134
  step_delay=self.step_delay,
135
  )
136
  backend_type = str(spec.backend.get("type", "")).lower()
 
258
 
259
  def stop_process_tree() -> None:
260
  """Stop the script and any workers it launched."""
261
+ terminate_process_tree(process, timeout=3)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
262
 
263
  try:
264
  while True:
adam/experiment_tracker.py ADDED
@@ -0,0 +1,387 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from __future__ import annotations
2
+
3
+ import json
4
+ import re
5
+ import sqlite3
6
+ from dataclasses import dataclass
7
+ from datetime import datetime, timezone
8
+ from pathlib import Path
9
+ from typing import Any
10
+
11
+ from adam.models import Job, JobStatus, SystemSnapshot
12
+
13
+
14
+ def _utc_now() -> str:
15
+ return datetime.now(timezone.utc).isoformat()
16
+
17
+
18
+ def _json(value: Any) -> str:
19
+ return json.dumps(value, sort_keys=True)
20
+
21
+
22
+ def _safe_json(value: str, fallback: Any) -> Any:
23
+ try:
24
+ parsed = json.loads(value or "")
25
+ except (TypeError, ValueError, json.JSONDecodeError):
26
+ return fallback
27
+ return parsed if isinstance(parsed, type(fallback)) else fallback
28
+
29
+
30
+ def _safe_int(value: Any, default: int = 0) -> int:
31
+ try:
32
+ if isinstance(value, bool):
33
+ return default
34
+ return int(value)
35
+ except (TypeError, ValueError):
36
+ return default
37
+
38
+
39
+ def _safe_float(value: Any, default: float = 0.0) -> float:
40
+ try:
41
+ if isinstance(value, bool):
42
+ return default
43
+ return float(value)
44
+ except (TypeError, ValueError):
45
+ return default
46
+
47
+
48
+ def _loss_from_logs(logs: list[str]) -> float | None:
49
+ for line in reversed(logs):
50
+ match = re.search(r"\bloss(?:\s*[:=]\s*|\s+)(-?\d+(?:\.\d+)?(?:e[+-]?\d+)?)", line, re.I)
51
+ if match:
52
+ try:
53
+ return float(match.group(1))
54
+ except ValueError:
55
+ return None
56
+ return None
57
+
58
+
59
+ def _duration_seconds(job: Job) -> int:
60
+ if not job.started_at:
61
+ return 0
62
+ try:
63
+ start = datetime.fromisoformat(job.started_at)
64
+ end = datetime.fromisoformat(job.ended_at) if job.ended_at else datetime.now(timezone.utc)
65
+ return max(0, int((end - start).total_seconds()))
66
+ except ValueError:
67
+ return 0
68
+
69
+
70
+ def _image_count(path: str) -> int:
71
+ folder = Path(path).expanduser()
72
+ if not folder.is_dir():
73
+ return 0
74
+ try:
75
+ return sum(
76
+ 1 for item in folder.rglob("*")
77
+ if item.is_file() and item.suffix.casefold() in {".png", ".jpg", ".jpeg", ".webp", ".bmp"}
78
+ )
79
+ except OSError:
80
+ return 0
81
+
82
+
83
+ @dataclass(slots=True)
84
+ class ExperimentRun:
85
+ id: str
86
+ job_id: str
87
+ timestamp: str
88
+ model_architecture: str
89
+ model_name: str
90
+ trigger_word: str
91
+ base_model: str
92
+ dataset_path: str
93
+ dataset_name: str
94
+ dataset_item_count: int
95
+ epochs: int
96
+ batch_size: int
97
+ learning_rate: float
98
+ optimizer: str
99
+ scheduler: str
100
+ resolution: int
101
+ seed: int
102
+ status: str
103
+ training_time_seconds: int
104
+ final_loss: float | None
105
+ output_folder: str
106
+ checkpoint_paths: list[str]
107
+ preview_images: list[str]
108
+ peak_vram_gb: float | None
109
+ hardware: dict[str, Any]
110
+ settings: dict[str, Any]
111
+ generation_settings: dict[str, Any]
112
+ notes: str = ""
113
+ quality_score: int | None = None
114
+
115
+ @classmethod
116
+ def from_row(cls, row: sqlite3.Row) -> "ExperimentRun":
117
+ payload = dict(row)
118
+ for key in ("checkpoint_paths", "preview_images"):
119
+ payload[key] = _safe_json(payload.get(key, "[]"), [])
120
+ for key in ("hardware", "settings", "generation_settings"):
121
+ payload[key] = _safe_json(payload.get(key, "{}"), {})
122
+ payload["quality_score"] = (
123
+ _safe_int(payload["quality_score"]) if payload.get("quality_score") is not None else None
124
+ )
125
+ return cls(**payload)
126
+
127
+
128
+ class ExperimentStore:
129
+ def __init__(self, root: Path) -> None:
130
+ self.root = root.resolve()
131
+ self.path = self.root / "data" / "experiments.sqlite3"
132
+ self.path.parent.mkdir(parents=True, exist_ok=True)
133
+ self._init_db()
134
+
135
+ def connect(self) -> sqlite3.Connection:
136
+ connection = sqlite3.connect(self.path)
137
+ connection.row_factory = sqlite3.Row
138
+ return connection
139
+
140
+ def _init_db(self) -> None:
141
+ with self.connect() as db:
142
+ db.execute(
143
+ """
144
+ CREATE TABLE IF NOT EXISTS experiments (
145
+ id TEXT PRIMARY KEY,
146
+ job_id TEXT UNIQUE NOT NULL,
147
+ timestamp TEXT NOT NULL,
148
+ model_architecture TEXT NOT NULL,
149
+ model_name TEXT NOT NULL,
150
+ trigger_word TEXT NOT NULL DEFAULT '',
151
+ base_model TEXT NOT NULL,
152
+ dataset_path TEXT NOT NULL,
153
+ dataset_name TEXT NOT NULL,
154
+ dataset_item_count INTEGER NOT NULL,
155
+ epochs INTEGER NOT NULL,
156
+ batch_size INTEGER NOT NULL,
157
+ learning_rate REAL NOT NULL,
158
+ optimizer TEXT NOT NULL,
159
+ scheduler TEXT NOT NULL,
160
+ resolution INTEGER NOT NULL,
161
+ seed INTEGER NOT NULL,
162
+ status TEXT NOT NULL,
163
+ training_time_seconds INTEGER NOT NULL,
164
+ final_loss REAL,
165
+ output_folder TEXT NOT NULL,
166
+ checkpoint_paths TEXT NOT NULL,
167
+ preview_images TEXT NOT NULL,
168
+ peak_vram_gb REAL,
169
+ hardware TEXT NOT NULL,
170
+ settings TEXT NOT NULL,
171
+ generation_settings TEXT NOT NULL,
172
+ notes TEXT NOT NULL DEFAULT '',
173
+ quality_score INTEGER
174
+ )
175
+ """
176
+ )
177
+ self._migrate_columns(db)
178
+
179
+ @staticmethod
180
+ def _migrate_columns(db: sqlite3.Connection) -> None:
181
+ existing = {row["name"] for row in db.execute("PRAGMA table_info(experiments)").fetchall()}
182
+ columns = {
183
+ "id": "TEXT PRIMARY KEY",
184
+ "job_id": "TEXT NOT NULL DEFAULT ''",
185
+ "timestamp": "TEXT NOT NULL DEFAULT ''",
186
+ "model_architecture": "TEXT NOT NULL DEFAULT ''",
187
+ "model_name": "TEXT NOT NULL DEFAULT ''",
188
+ "trigger_word": "TEXT NOT NULL DEFAULT ''",
189
+ "base_model": "TEXT NOT NULL DEFAULT ''",
190
+ "dataset_path": "TEXT NOT NULL DEFAULT ''",
191
+ "dataset_name": "TEXT NOT NULL DEFAULT ''",
192
+ "dataset_item_count": "INTEGER NOT NULL DEFAULT 0",
193
+ "epochs": "INTEGER NOT NULL DEFAULT 0",
194
+ "batch_size": "INTEGER NOT NULL DEFAULT 0",
195
+ "learning_rate": "REAL NOT NULL DEFAULT 0",
196
+ "optimizer": "TEXT NOT NULL DEFAULT ''",
197
+ "scheduler": "TEXT NOT NULL DEFAULT ''",
198
+ "resolution": "INTEGER NOT NULL DEFAULT 0",
199
+ "seed": "INTEGER NOT NULL DEFAULT 0",
200
+ "status": "TEXT NOT NULL DEFAULT ''",
201
+ "training_time_seconds": "INTEGER NOT NULL DEFAULT 0",
202
+ "final_loss": "REAL",
203
+ "output_folder": "TEXT NOT NULL DEFAULT ''",
204
+ "checkpoint_paths": "TEXT NOT NULL DEFAULT '[]'",
205
+ "preview_images": "TEXT NOT NULL DEFAULT '[]'",
206
+ "peak_vram_gb": "REAL",
207
+ "hardware": "TEXT NOT NULL DEFAULT '{}'",
208
+ "settings": "TEXT NOT NULL DEFAULT '{}'",
209
+ "generation_settings": "TEXT NOT NULL DEFAULT '{}'",
210
+ "notes": "TEXT NOT NULL DEFAULT ''",
211
+ "quality_score": "INTEGER",
212
+ }
213
+ for name, definition in columns.items():
214
+ if name not in existing and name != "id":
215
+ db.execute(f"ALTER TABLE experiments ADD COLUMN {name} {definition}")
216
+
217
+ def record_job(self, job: Job, snapshot: SystemSnapshot | None = None) -> ExperimentRun | None:
218
+ training_steps = [step for step in job.plan.steps if step.tool_id.endswith("_trainer")]
219
+ if not training_steps:
220
+ return None
221
+ step = training_steps[-1]
222
+ args = dict(step.arguments)
223
+ architecture = step.tool_id.removesuffix("_trainer")
224
+ dataset_path = str(args.get("dataset_dir", ""))
225
+ output_folder = str(job.output_folder or args.get("output_dir", ""))
226
+ preview_images = [job.preview_path] if job.preview_path else []
227
+ checkpoints = []
228
+ if output_folder:
229
+ folder = Path(output_folder)
230
+ if folder.is_dir():
231
+ try:
232
+ checkpoints = [
233
+ str(path)
234
+ for path in sorted(folder.rglob("*"))
235
+ if path.is_file() and path.suffix.casefold() in {".safetensors", ".ckpt", ".pt", ".bin"}
236
+ ][-10:]
237
+ discovered_previews = [
238
+ str(path)
239
+ for path in sorted(folder.rglob("*"))
240
+ if path.is_file()
241
+ and path.suffix.casefold() in {".png", ".jpg", ".jpeg", ".webp", ".bmp"}
242
+ and any(token in path.name.casefold() for token in ("preview", "sample", "epoch"))
243
+ ][-12:]
244
+ preview_images = list(dict.fromkeys([*preview_images, *discovered_previews]))
245
+ except OSError:
246
+ checkpoints = []
247
+ hardware = {}
248
+ peak_vram = None
249
+ if snapshot is not None:
250
+ hardware = {
251
+ "gpu_name": snapshot.gpu_name,
252
+ "gpu_percent": snapshot.gpu_percent,
253
+ "vram_used_gb": snapshot.vram_used_gb,
254
+ "vram_total_gb": snapshot.vram_total_gb,
255
+ "memory_used_gb": snapshot.memory_used_gb,
256
+ "memory_total_gb": snapshot.memory_total_gb,
257
+ "cpu_percent": snapshot.cpu_percent,
258
+ "gpu_temperature": snapshot.gpu_temperature,
259
+ }
260
+ peak_vram = snapshot.vram_used_gb or None
261
+ run = ExperimentRun(
262
+ id=f"EXP-{job.id}",
263
+ job_id=job.id,
264
+ timestamp=job.ended_at or _utc_now(),
265
+ model_architecture=architecture,
266
+ model_name=str(args.get("model_name", job.plan.project_name)),
267
+ trigger_word=str(args.get("trigger_word", "")),
268
+ base_model=str(args.get("base_model", args.get("base_model_path", ""))),
269
+ dataset_path=dataset_path,
270
+ dataset_name=Path(dataset_path).name if dataset_path else "",
271
+ dataset_item_count=_image_count(dataset_path),
272
+ epochs=_safe_int(args.get("epochs")),
273
+ batch_size=_safe_int(args.get("batch_size")),
274
+ learning_rate=_safe_float(args.get("learning_rate")),
275
+ optimizer=str(args.get("optimizer", "")),
276
+ scheduler=str(args.get("scheduler", args.get("sampler", ""))),
277
+ resolution=_safe_int(args.get("resolution")),
278
+ seed=_safe_int(args.get("seed", args.get("preview_seed", 0))),
279
+ status=job.status.value,
280
+ training_time_seconds=_duration_seconds(job),
281
+ final_loss=_loss_from_logs(job.logs),
282
+ output_folder=output_folder,
283
+ checkpoint_paths=checkpoints,
284
+ preview_images=preview_images,
285
+ peak_vram_gb=peak_vram,
286
+ hardware=hardware,
287
+ settings=args,
288
+ generation_settings={},
289
+ )
290
+ with self.connect() as db:
291
+ db.execute(
292
+ """
293
+ INSERT INTO experiments (
294
+ id, job_id, timestamp, model_architecture, model_name, trigger_word, base_model,
295
+ dataset_path, dataset_name, dataset_item_count, epochs, batch_size,
296
+ learning_rate, optimizer, scheduler, resolution, seed, status,
297
+ training_time_seconds, final_loss, output_folder, checkpoint_paths,
298
+ preview_images, peak_vram_gb, hardware, settings, generation_settings
299
+ ) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
300
+ ON CONFLICT(job_id) DO UPDATE SET
301
+ timestamp=excluded.timestamp,
302
+ trigger_word=excluded.trigger_word,
303
+ status=excluded.status,
304
+ training_time_seconds=excluded.training_time_seconds,
305
+ final_loss=excluded.final_loss,
306
+ output_folder=excluded.output_folder,
307
+ checkpoint_paths=excluded.checkpoint_paths,
308
+ preview_images=excluded.preview_images,
309
+ peak_vram_gb=excluded.peak_vram_gb,
310
+ hardware=excluded.hardware,
311
+ settings=excluded.settings
312
+ """,
313
+ (
314
+ run.id, run.job_id, run.timestamp, run.model_architecture, run.model_name,
315
+ run.trigger_word, run.base_model, run.dataset_path, run.dataset_name, run.dataset_item_count,
316
+ run.epochs, run.batch_size, run.learning_rate, run.optimizer, run.scheduler,
317
+ run.resolution, run.seed, run.status, run.training_time_seconds, run.final_loss,
318
+ run.output_folder, _json(run.checkpoint_paths), _json(run.preview_images),
319
+ run.peak_vram_gb, _json(run.hardware), _json(run.settings),
320
+ _json(run.generation_settings),
321
+ ),
322
+ )
323
+ return run
324
+
325
+ def list_runs(self, search: str = "", architecture: str = "", dataset: str = "", limit: int = 200) -> list[ExperimentRun]:
326
+ clauses = []
327
+ params: list[Any] = []
328
+ if search:
329
+ clauses.append("(id LIKE ? OR job_id LIKE ? OR model_name LIKE ? OR dataset_name LIKE ? OR dataset_path LIKE ? OR output_folder LIKE ? OR notes LIKE ? OR status LIKE ?)")
330
+ term = f"%{search}%"
331
+ params.extend([term, term, term, term, term, term, term, term])
332
+ if architecture:
333
+ clauses.append("model_architecture = ?")
334
+ params.append(architecture)
335
+ if dataset:
336
+ clauses.append("dataset_name LIKE ?")
337
+ params.append(f"%{dataset}%")
338
+ where = " WHERE " + " AND ".join(clauses) if clauses else ""
339
+ with self.connect() as db:
340
+ rows = db.execute(
341
+ "SELECT * FROM experiments" + where + " ORDER BY timestamp DESC LIMIT ?",
342
+ [*params, int(limit)],
343
+ ).fetchall()
344
+ return [ExperimentRun.from_row(row) for row in rows]
345
+
346
+ def get(self, run_id: str) -> ExperimentRun | None:
347
+ with self.connect() as db:
348
+ row = db.execute("SELECT * FROM experiments WHERE id = ?", (run_id,)).fetchone()
349
+ return ExperimentRun.from_row(row) if row else None
350
+
351
+ def update_notes(self, run_id: str, notes: str, quality_score: int | None) -> None:
352
+ with self.connect() as db:
353
+ db.execute(
354
+ "UPDATE experiments SET notes = ?, quality_score = ? WHERE id = ?",
355
+ (notes, quality_score, run_id),
356
+ )
357
+
358
+ def compare(self, run_ids: list[str]) -> list[dict[str, Any]]:
359
+ runs = [run for run_id in run_ids if (run := self.get(run_id)) is not None]
360
+ fields = [
361
+ "model_architecture", "epochs", "final_loss", "training_time_seconds",
362
+ "resolution", "batch_size", "learning_rate", "scheduler",
363
+ "peak_vram_gb", "dataset_name", "quality_score",
364
+ ]
365
+ rows = []
366
+ for field in fields:
367
+ values = {run.id: getattr(run, field) for run in runs}
368
+ comparable = {str(value) for value in values.values()}
369
+ rows.append({"field": field, "changed": len(comparable) > 1, **values})
370
+ return rows
371
+
372
+ def clone_request(self, run_id: str) -> str:
373
+ run = self.get(run_id)
374
+ if run is None:
375
+ return ""
376
+ options = {
377
+ key: value
378
+ for key, value in run.settings.items()
379
+ if key not in {"dataset_dir", "model_name", "epochs", "output_dir", "resume_from"}
380
+ }
381
+ return (
382
+ f"From the {run.dataset_name or run.dataset_path} dataset, train a "
383
+ f"{run.model_architecture.upper()} model for {run.epochs} epochs. "
384
+ f"Name the model {run.model_name} Clone. "
385
+ "[ADAM_TRAINING_OPTIONS:" + json.dumps(options, sort_keys=True) + "] "
386
+ "[ADAM_TRAINER:" + run.model_architecture + "]"
387
+ )
adam/generations.py CHANGED
@@ -219,6 +219,7 @@ def generation_metadata_path(folder: Path, timestamp: str, job_id: str) -> Path:
219
 
220
  @dataclass(frozen=True, slots=True)
221
  class GenerationRecord:
 
222
  folder: Path
223
  images: tuple[Path, ...]
224
  provider_id: str
@@ -231,6 +232,8 @@ class GenerationRecord:
231
  sampler: str
232
  aspect_ratio: str
233
  created_at: str
 
 
234
 
235
  @classmethod
236
  def from_metadata(cls, metadata_path: Path) -> "GenerationRecord | None":
@@ -246,6 +249,7 @@ class GenerationRecord:
246
  if not images:
247
  return None
248
  return cls(
 
249
  folder=folder,
250
  images=images,
251
  provider_id=str(payload.get("provider_id", "")),
@@ -258,9 +262,77 @@ class GenerationRecord:
258
  sampler=str(payload.get("sampler", "")),
259
  aspect_ratio=str(payload.get("aspect_ratio", "")),
260
  created_at=str(payload.get("created_at", "")),
 
 
 
 
 
 
261
  )
262
 
263
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
264
  def generation_tools(registry: ToolRegistry) -> list[ToolSpec]:
265
  return [
266
  tool
 
219
 
220
  @dataclass(frozen=True, slots=True)
221
  class GenerationRecord:
222
+ metadata_path: Path
223
  folder: Path
224
  images: tuple[Path, ...]
225
  provider_id: str
 
232
  sampler: str
233
  aspect_ratio: str
234
  created_at: str
235
+ smart_generation: dict[str, Any]
236
+ image_evaluations: dict[str, dict[str, Any]]
237
 
238
  @classmethod
239
  def from_metadata(cls, metadata_path: Path) -> "GenerationRecord | None":
 
249
  if not images:
250
  return None
251
  return cls(
252
+ metadata_path=metadata_path,
253
  folder=folder,
254
  images=images,
255
  provider_id=str(payload.get("provider_id", "")),
 
262
  sampler=str(payload.get("sampler", "")),
263
  aspect_ratio=str(payload.get("aspect_ratio", "")),
264
  created_at=str(payload.get("created_at", "")),
265
+ smart_generation=dict(payload.get("smart_generation") or {}),
266
+ image_evaluations={
267
+ str(Path(path).resolve()): dict(value)
268
+ for path, value in dict(payload.get("image_evaluations") or {}).items()
269
+ if isinstance(value, dict)
270
+ },
271
  )
272
 
273
 
274
+ @dataclass(frozen=True, slots=True)
275
+ class GenerationModelFolder:
276
+ """A model-centered view over existing generation batches."""
277
+
278
+ key: str
279
+ model_name: str
280
+ model_path: str
281
+ provider_id: str
282
+ provider_name: str
283
+ records: tuple[GenerationRecord, ...]
284
+ image_count: int
285
+ cover_image: Path | None
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()
292
+ if raw_path:
293
+ try:
294
+ return f"path:{Path(raw_path).expanduser().resolve()}".casefold()
295
+ except OSError:
296
+ return f"path:{raw_path}".casefold()
297
+ return f"name:{record.provider_id}:{record.model_name}".casefold()
298
+
299
+
300
+ def group_generation_records(
301
+ records: list[GenerationRecord],
302
+ ) -> list[GenerationModelFolder]:
303
+ """Build newest-first automatic model folders without changing files."""
304
+ grouped: dict[str, list[GenerationRecord]] = {}
305
+ for record in records:
306
+ grouped.setdefault(generation_model_key(record), []).append(record)
307
+ folders: list[GenerationModelFolder] = []
308
+ for key, model_records in grouped.items():
309
+ newest_first = sorted(
310
+ model_records,
311
+ key=lambda item: item.created_at or item.folder.name,
312
+ reverse=True,
313
+ )
314
+ latest = newest_first[0]
315
+ cover = next(
316
+ (path for record in newest_first for path in record.images if path.is_file()),
317
+ None,
318
+ )
319
+ folders.append(
320
+ GenerationModelFolder(
321
+ key=key,
322
+ model_name=latest.model_name,
323
+ model_path=latest.model_path,
324
+ provider_id=latest.provider_id,
325
+ provider_name=latest.provider_name,
326
+ records=tuple(newest_first),
327
+ image_count=sum(len(record.images) for record in newest_first),
328
+ cover_image=cover,
329
+ latest_at=latest.created_at,
330
+ )
331
+ )
332
+ folders.sort(key=lambda item: (item.latest_at, item.model_name.casefold()), reverse=True)
333
+ return folders
334
+
335
+
336
  def generation_tools(registry: ToolRegistry) -> list[ToolSpec]:
337
  return [
338
  tool
adam/image_preferences.py ADDED
@@ -0,0 +1,313 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from __future__ import annotations
2
+
3
+ from dataclasses import asdict, dataclass, field
4
+ from datetime import datetime, timezone
5
+ import hashlib
6
+ import json
7
+ from pathlib import Path
8
+ from typing import Any, Callable, Sequence
9
+
10
+ from adam.eve import EveVisionModel, classify_eve_embeddings
11
+
12
+
13
+ RATINGS = {"favorite", "keep", "unsure", "reject"}
14
+ POSITIVE_RATINGS = {"favorite", "keep"}
15
+ NEGATIVE_RATINGS = {"reject"}
16
+
17
+
18
+ def _now() -> str:
19
+ return datetime.now(timezone.utc).isoformat()
20
+
21
+
22
+ def preference_profile_id(provider_id: str, model_path: str) -> str:
23
+ key = f"{provider_id}\n{str(Path(model_path).expanduser().resolve())}"
24
+ return hashlib.sha1(key.encode("utf-8")).hexdigest()[:16]
25
+
26
+
27
+ def image_cache_id(path: str | Path) -> str:
28
+ resolved = str(Path(path).expanduser().resolve())
29
+ return hashlib.sha1(resolved.encode("utf-8")).hexdigest()
30
+
31
+
32
+ @dataclass(slots=True)
33
+ class GenerationRating:
34
+ image_path: str
35
+ rating: str
36
+ provider_id: str
37
+ model_name: str
38
+ model_path: str
39
+ seed: int = 0
40
+ sampler: str = ""
41
+ steps: int = 0
42
+ resolution: str = ""
43
+ generation_settings: dict[str, Any] = field(default_factory=dict)
44
+ generation_created_at: str = ""
45
+ rated_at: str = field(default_factory=_now)
46
+ embedding: list[float] | None = None
47
+
48
+ @classmethod
49
+ def from_dict(cls, payload: dict[str, Any]) -> "GenerationRating":
50
+ rating = str(payload.get("rating", "unsure")).casefold()
51
+ return cls(
52
+ image_path=str(Path(str(payload.get("image_path", ""))).expanduser().resolve()),
53
+ rating=rating if rating in RATINGS else "unsure",
54
+ provider_id=str(payload.get("provider_id", "")),
55
+ model_name=str(payload.get("model_name", "")),
56
+ model_path=str(payload.get("model_path", "")),
57
+ seed=int(payload.get("seed", 0) or 0),
58
+ sampler=str(payload.get("sampler", "")),
59
+ steps=int(payload.get("steps", 0) or 0),
60
+ resolution=str(payload.get("resolution", "")),
61
+ generation_settings=dict(payload.get("generation_settings") or {}),
62
+ generation_created_at=str(payload.get("generation_created_at", "")),
63
+ rated_at=str(payload.get("rated_at") or _now()),
64
+ embedding=[float(value) for value in payload["embedding"]]
65
+ if isinstance(payload.get("embedding"), list)
66
+ else None,
67
+ )
68
+
69
+
70
+ @dataclass(frozen=True, slots=True)
71
+ class PreferenceScore:
72
+ image_path: str
73
+ score: float | None
74
+ confidence: float
75
+ category: str
76
+ reason: str = ""
77
+
78
+
79
+ class PreferenceProfile:
80
+ def __init__(self, root: Path, provider_id: str, model_name: str, model_path: str) -> None:
81
+ self.root = root.resolve()
82
+ self.provider_id = provider_id
83
+ self.model_name = model_name
84
+ self.model_path = str(Path(model_path).expanduser().resolve()) if model_path else ""
85
+ self.id = preference_profile_id(provider_id, self.model_path)
86
+ self.path = self.root / "data" / "generation_preferences" / f"{self.id}.json"
87
+ self.keep_threshold = 0.70
88
+ self.reject_threshold = 0.35
89
+ self.ratings: dict[str, GenerationRating] = {}
90
+ self.load()
91
+
92
+ def load(self) -> None:
93
+ try:
94
+ payload = json.loads(self.path.read_text(encoding="utf-8"))
95
+ except (OSError, ValueError, TypeError, json.JSONDecodeError):
96
+ return
97
+ self.model_name = str(payload.get("model_name") or self.model_name)
98
+ self.provider_id = str(payload.get("provider_id") or self.provider_id)
99
+ self.model_path = str(payload.get("model_path") or self.model_path)
100
+ thresholds = payload.get("thresholds", {})
101
+ if isinstance(thresholds, dict):
102
+ self.keep_threshold = float(thresholds.get("keep", self.keep_threshold))
103
+ self.reject_threshold = float(thresholds.get("reject", self.reject_threshold))
104
+ ratings = payload.get("ratings", [])
105
+ if isinstance(ratings, list):
106
+ for item in ratings:
107
+ if isinstance(item, dict):
108
+ rating = GenerationRating.from_dict(item)
109
+ self.ratings[rating.image_path] = rating
110
+
111
+ def save(self) -> None:
112
+ self.path.parent.mkdir(parents=True, exist_ok=True)
113
+ temporary = self.path.with_suffix(".tmp")
114
+ temporary.write_text(
115
+ json.dumps(
116
+ {
117
+ "version": 1,
118
+ "profile_id": self.id,
119
+ "provider_id": self.provider_id,
120
+ "model_name": self.model_name,
121
+ "model_path": self.model_path,
122
+ "thresholds": {
123
+ "keep": self.keep_threshold,
124
+ "reject": self.reject_threshold,
125
+ },
126
+ "ratings": [asdict(item) for item in self.ratings.values()],
127
+ "updated_at": _now(),
128
+ },
129
+ indent=2,
130
+ ),
131
+ encoding="utf-8",
132
+ )
133
+ temporary.replace(self.path)
134
+
135
+ def set_rating(
136
+ self,
137
+ image_path: str | Path,
138
+ rating: str,
139
+ *,
140
+ seed: int = 0,
141
+ sampler: str = "",
142
+ steps: int = 0,
143
+ resolution: str = "",
144
+ generation_settings: dict[str, Any] | None = None,
145
+ generation_created_at: str = "",
146
+ embedding: Sequence[float] | None = None,
147
+ ) -> GenerationRating:
148
+ clean = rating.casefold().strip()
149
+ if clean not in RATINGS:
150
+ raise ValueError("Generation rating must be Favorite, Keep, Unsure, or Reject.")
151
+ resolved = str(Path(image_path).expanduser().resolve())
152
+ existing = self.ratings.get(resolved)
153
+ record = GenerationRating(
154
+ image_path=resolved,
155
+ rating=clean,
156
+ provider_id=self.provider_id,
157
+ model_name=self.model_name,
158
+ model_path=self.model_path,
159
+ seed=int(seed),
160
+ sampler=sampler,
161
+ steps=int(steps),
162
+ resolution=resolution,
163
+ generation_settings=dict(generation_settings or {}),
164
+ generation_created_at=generation_created_at,
165
+ rated_at=_now(),
166
+ embedding=[float(value) for value in embedding] if embedding is not None else (
167
+ existing.embedding if existing else None
168
+ ),
169
+ )
170
+ self.ratings[resolved] = record
171
+ self.save()
172
+ return record
173
+
174
+ def rating_for(self, image_path: str | Path) -> GenerationRating | None:
175
+ return self.ratings.get(str(Path(image_path).expanduser().resolve()))
176
+
177
+ def examples(self) -> tuple[list[GenerationRating], list[GenerationRating]]:
178
+ positive = [
179
+ item for item in self.ratings.values()
180
+ if item.rating in POSITIVE_RATINGS and Path(item.image_path).is_file()
181
+ ]
182
+ negative = [
183
+ item for item in self.ratings.values()
184
+ if item.rating in NEGATIVE_RATINGS and Path(item.image_path).is_file()
185
+ ]
186
+ return positive, negative
187
+
188
+ def has_signal(self) -> bool:
189
+ positive, _negative = self.examples()
190
+ return bool(positive)
191
+
192
+
193
+ class ImageEmbeddingCache:
194
+ def __init__(self, root: Path, model_id: str) -> None:
195
+ self.root = root.resolve()
196
+ self.model_id = model_id
197
+ self.folder = self.root / "data" / "image_embeddings" / hashlib.sha1(model_id.encode("utf-8")).hexdigest()[:12]
198
+
199
+ def get(self, path: str | Path) -> list[float] | None:
200
+ cache_path = self.folder / f"{image_cache_id(path)}.json"
201
+ try:
202
+ payload = json.loads(cache_path.read_text(encoding="utf-8"))
203
+ except (OSError, ValueError, TypeError, json.JSONDecodeError):
204
+ return None
205
+ source = Path(path).expanduser().resolve()
206
+ try:
207
+ stat = source.stat()
208
+ except OSError:
209
+ return None
210
+ if payload.get("path") != str(source) or payload.get("mtime") != stat.st_mtime:
211
+ return None
212
+ vector = payload.get("embedding")
213
+ return [float(value) for value in vector] if isinstance(vector, list) else None
214
+
215
+ def set(self, path: str | Path, embedding: Sequence[float]) -> None:
216
+ source = Path(path).expanduser().resolve()
217
+ try:
218
+ stat = source.stat()
219
+ except OSError:
220
+ return
221
+ self.folder.mkdir(parents=True, exist_ok=True)
222
+ cache_path = self.folder / f"{image_cache_id(source)}.json"
223
+ temporary = cache_path.with_suffix(".tmp")
224
+ temporary.write_text(
225
+ json.dumps(
226
+ {
227
+ "path": str(source),
228
+ "mtime": stat.st_mtime,
229
+ "model_id": self.model_id,
230
+ "embedding": [float(value) for value in embedding],
231
+ }
232
+ ),
233
+ encoding="utf-8",
234
+ )
235
+ temporary.replace(cache_path)
236
+
237
+
238
+ class GenerationPreferenceEvaluator:
239
+ """Shared EVE-backed scorer for generated images."""
240
+
241
+ def __init__(
242
+ self,
243
+ root: Path,
244
+ vision: EveVisionModel | None = None,
245
+ *,
246
+ embedder: Callable[[Sequence[str | Path]], list[list[float]]] | None = None,
247
+ ) -> None:
248
+ self.root = root.resolve()
249
+ self.vision = vision or EveVisionModel(prefer_gpu=False)
250
+ self.embedder = embedder
251
+ self.cache = ImageEmbeddingCache(self.root, self.vision.model_id)
252
+
253
+ def _embedding(self, path: str | Path) -> list[float]:
254
+ cached = self.cache.get(path)
255
+ if cached is not None:
256
+ return cached
257
+ vectors = self.embedder([path]) if self.embedder else self.vision.embed([path])
258
+ vector = [float(value) for value in vectors[0]]
259
+ self.cache.set(path, vector)
260
+ return vector
261
+
262
+ def score(
263
+ self,
264
+ profile: PreferenceProfile,
265
+ image_paths: Sequence[str | Path],
266
+ *,
267
+ keep_threshold: float | None = None,
268
+ reject_threshold: float | None = None,
269
+ ) -> list[PreferenceScore]:
270
+ positive, negative = profile.examples()
271
+ if not positive:
272
+ return [
273
+ PreferenceScore(str(Path(path).expanduser().resolve()), None, 0.0, "Needs Review", "No preference examples yet")
274
+ for path in image_paths
275
+ ]
276
+ positive_vectors = [item.embedding or self._embedding(item.image_path) for item in positive]
277
+ negative_vectors = [item.embedding or self._embedding(item.image_path) for item in negative]
278
+ image_vectors = [self._embedding(path) for path in image_paths]
279
+ keep = max(0.001, min(1.0, float(keep_threshold if keep_threshold is not None else profile.keep_threshold)))
280
+ reject = max(0.0, min(float(reject_threshold if reject_threshold is not None else profile.reject_threshold), keep - 0.001))
281
+ results = classify_eve_embeddings(
282
+ image_paths,
283
+ image_vectors,
284
+ positive_vectors,
285
+ negative_vectors,
286
+ keep_threshold=keep,
287
+ reject_threshold=reject,
288
+ )
289
+ categories = {"keep": "Strong Keep", "reject": "Likely Reject", "unreviewed": "Needs Review"}
290
+ return [
291
+ PreferenceScore(result.path, result.match_score, result.decision_confidence, categories[result.suggestion])
292
+ for result in results
293
+ ]
294
+
295
+
296
+ def score_generated_images(
297
+ root: Path,
298
+ *,
299
+ provider_id: str,
300
+ model_name: str,
301
+ model_path: str,
302
+ image_paths: Sequence[str | Path],
303
+ keep_threshold: float | None = None,
304
+ reject_threshold: float | None = None,
305
+ ) -> list[PreferenceScore]:
306
+ profile = PreferenceProfile(root, provider_id, model_name, model_path)
307
+ evaluator = GenerationPreferenceEvaluator(root)
308
+ return evaluator.score(
309
+ profile,
310
+ image_paths,
311
+ keep_threshold=keep_threshold,
312
+ reject_threshold=reject_threshold,
313
+ )
adam/job_manager.py CHANGED
@@ -5,17 +5,27 @@ import logging
5
  import math
6
  import re
7
  import threading
8
- from datetime import datetime, timezone
 
9
  from pathlib import Path
10
  from typing import Any
11
 
12
- from PySide6.QtCore import QObject, QThread, Signal
13
 
14
- from adam.executor import ToolCancelled, ToolExecutionError, ToolExecutor
15
  from adam.assets import AssetRegistry
16
  from adam.atlas import AtlasSupervisor
 
17
  from adam.models import ExecutionPlan, Job, JobStatus, StepStatus, utc_now
18
  from adam.nova import evaluate_job_output
 
 
 
 
 
 
 
 
19
 
20
 
21
  class JobWorker(QThread):
@@ -28,6 +38,8 @@ class JobWorker(QThread):
28
  self.cancel_event = threading.Event()
29
  self.run_event = threading.Event()
30
  self.run_event.set()
 
 
31
 
32
  def pause(self) -> None:
33
  self.run_event.clear()
@@ -39,6 +51,12 @@ class JobWorker(QThread):
39
  self.cancel_event.set()
40
  self.run_event.set()
41
 
 
 
 
 
 
 
42
  def run(self) -> None:
43
  total_steps = len(self.job.plan.steps)
44
  try:
@@ -54,15 +72,25 @@ class JobWorker(QThread):
54
  )
55
 
56
  preview_state = {"epoch": 0, "path": ""}
 
 
57
 
58
- def on_progress(percent: int, message: str, step_index: int = index) -> None:
59
  overall = int(((step_index + percent / 100) / total_steps) * 100)
 
 
 
 
 
 
 
60
  self.event.emit(
61
  {
62
  "type": "progress",
63
  "step_percent": percent,
64
  "overall": overall,
65
  "message": message,
 
66
  }
67
  )
68
  self._discover_external_preview(step, message, preview_state)
@@ -84,6 +112,8 @@ class JobWorker(QThread):
84
  progress_callback=on_progress,
85
  log_callback=on_log,
86
  preview_callback=on_preview,
 
 
87
  )
88
  self.event.emit(
89
  {
@@ -95,6 +125,8 @@ class JobWorker(QThread):
95
  self.event.emit({"type": "completed"})
96
  except ToolCancelled as exc:
97
  self.event.emit({"type": "cancelled", "message": str(exc)})
 
 
98
  except Exception as exc:
99
  self.event.emit(
100
  {
@@ -140,6 +172,61 @@ class JobWorker(QThread):
140
  "steps": int(step.arguments.get("preview_steps", 0) or 0),
141
  })
142
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
143
 
144
  class JobManager(QObject):
145
  job_created = Signal(object)
@@ -159,30 +246,57 @@ class JobManager(QObject):
159
  self.root = root
160
  self.executor = executor
161
  self.logger = logger
 
162
  self.jobs_path = root / "data" / "jobs.json"
163
  self.assets = AssetRegistry(root)
164
  self.atlas = AtlasSupervisor(config)
 
165
  self.jobs: list[Job] = []
166
  self._queue: list[str] = []
167
  self._worker: JobWorker | None = None
168
  self._active_job: Job | None = None
 
 
 
169
  self._load()
 
 
 
 
 
 
 
170
 
171
  @property
172
  def active_job(self) -> Job | None:
173
  return self._active_job
174
 
175
- def submit(self, plan: ExecutionPlan) -> Job:
 
 
 
 
 
 
 
 
 
 
 
176
  status = (
177
  JobStatus.AWAITING_CONFIRMATION
178
  if plan.requires_confirmation
179
- else JobStatus.QUEUED
180
  )
181
- job = Job(plan=plan, status=status)
182
  self.jobs.insert(0, job)
183
  self._append_log(job, f"Plan created: {plan.summary}")
184
  if plan.requires_confirmation:
185
  self._append_log(job, "Waiting for user confirmation.")
 
 
 
 
186
  else:
187
  self._queue.append(job.id)
188
  self._save()
@@ -196,9 +310,12 @@ class JobManager(QObject):
196
  job = self.get(job_id)
197
  if job.status != JobStatus.AWAITING_CONFIRMATION:
198
  return
199
- job.status = JobStatus.QUEUED
200
  self._append_log(job, "Plan approved by user.")
201
- self._queue.append(job.id)
 
 
 
202
  self._save()
203
  self.job_updated.emit(job)
204
  self._start_next()
@@ -236,10 +353,13 @@ class JobManager(QObject):
236
  if job is self._active_job and self._worker:
237
  self._append_log(job, "Cancellation requested.")
238
  self._worker.cancel()
 
 
239
  return
240
  if job.id in self._queue:
241
  self._queue.remove(job.id)
242
  if job.status in {
 
243
  JobStatus.QUEUED,
244
  JobStatus.AWAITING_CONFIRMATION,
245
  JobStatus.DRAFT,
@@ -250,6 +370,104 @@ class JobManager(QObject):
250
  self._save()
251
  self.job_updated.emit(job)
252
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
253
  def get(self, job_id: str) -> Job:
254
  for job in self.jobs:
255
  if job.id == job_id:
@@ -383,14 +601,30 @@ class JobManager(QObject):
383
  job.preview_total = 0
384
  job.preview_image_index = 0
385
  job.preview_image_count = 0
 
 
 
 
 
 
386
  self._append_log(job, f"Starting: {job.plan.steps[index].title}")
387
  elif event_type == "progress":
388
  job.progress = int(event["overall"])
 
 
 
 
 
 
389
  message = str(event["message"])
390
  if message and (not job.logs or message not in job.logs[-1]):
391
  self._append_log(job, message)
 
 
392
  elif event_type == "log":
393
  self._append_log(job, str(event["message"]))
 
 
394
  elif event_type == "preview":
395
  job.preview_path = str(event.get("path", "")) or None
396
  job.preview_epoch = int(event.get("epoch", 0) or 0)
@@ -407,6 +641,8 @@ class JobManager(QObject):
407
  label = "Denoising" if job.preview_kind == "generation" else "Training"
408
  position = f" step {job.preview_current}" if job.preview_current else f" epoch {job.preview_epoch}"
409
  self._append_log(job, f"{label} preview updated at{position}.")
 
 
410
  elif event_type == "step_finished":
411
  index = int(event["index"])
412
  job.plan.steps[index].status = StepStatus.FINISHED
@@ -432,8 +668,11 @@ class JobManager(QObject):
432
  elif event_type == "completed":
433
  job.status = JobStatus.FINISHED
434
  job.progress = 100
 
 
435
  job.ended_at = utc_now()
436
  self._append_log(job, "Job finished successfully.")
 
437
  self.notification.emit("Job complete", job.plan.project_name)
438
  elif event_type == "cancelled":
439
  job.status = JobStatus.CANCELLED
@@ -442,7 +681,31 @@ class JobManager(QObject):
442
  for step in job.plan.steps:
443
  if step.status == StepStatus.RUNNING:
444
  step.status = StepStatus.SKIPPED
 
445
  self.notification.emit("Job cancelled", job.plan.project_name)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
446
  elif event_type == "failed":
447
  job.status = JobStatus.FAILED
448
  job.ended_at = utc_now()
@@ -456,10 +719,35 @@ class JobManager(QObject):
456
  event.get("exception"),
457
  job.error,
458
  )
 
459
  self.notification.emit("Job failed", job.error)
460
  self._save()
461
  self.job_updated.emit(job)
462
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
463
  def _worker_finished(self) -> None:
464
  self._worker = None
465
  self._active_job = None
@@ -468,6 +756,7 @@ class JobManager(QObject):
468
 
469
  def supervise(self, snapshot: Any) -> None:
470
  """Let ATLAS inspect the active training run and apply critical pauses."""
 
471
  job = self._active_job
472
  if job is None or job.status != JobStatus.RUNNING:
473
  return
@@ -495,6 +784,14 @@ class JobManager(QObject):
495
  if decision.action == "pause" and job.status == JobStatus.RUNNING:
496
  self.pause(job.id)
497
 
 
 
 
 
 
 
 
 
498
  def _append_log(self, job: Job, message: str) -> None:
499
  timestamp = datetime.now().strftime("%H:%M:%S")
500
  line = f"[{timestamp}] {message}"
@@ -517,13 +814,17 @@ class JobManager(QObject):
517
  self.jobs = []
518
  return
519
  for job in self.jobs:
520
- if job.status in {JobStatus.RUNNING, JobStatus.PAUSED, JobStatus.QUEUED}:
521
  job.status = JobStatus.INTERRUPTED
522
  job.ended_at = utc_now()
523
  job.logs.append(
524
  "[startup] Previous session ended before this job. "
525
  "Review it before retrying."
526
  )
 
 
 
 
527
  self._save()
528
 
529
  def _save(self) -> None:
 
5
  import math
6
  import re
7
  import threading
8
+ import time
9
+ from datetime import datetime, timedelta, timezone
10
  from pathlib import Path
11
  from typing import Any
12
 
13
+ from PySide6.QtCore import QCoreApplication, QObject, QThread, QTimer, Signal
14
 
15
+ from adam.executor import ToolAdjustmentRequested, ToolCancelled, ToolExecutionError, ToolExecutor
16
  from adam.assets import AssetRegistry
17
  from adam.atlas import AtlasSupervisor
18
+ from adam.experiment_tracker import ExperimentStore
19
  from adam.models import ExecutionPlan, Job, JobStatus, StepStatus, utc_now
20
  from adam.nova import evaluate_job_output
21
+ from adam.training_assistant import append_preflight_summary
22
+
23
+
24
+ def _safe_int(value: Any) -> int:
25
+ try:
26
+ return max(0, int(value or 0))
27
+ except (TypeError, ValueError):
28
+ return 0
29
 
30
 
31
  class JobWorker(QThread):
 
38
  self.cancel_event = threading.Event()
39
  self.run_event = threading.Event()
40
  self.run_event.set()
41
+ self.adjustment_event = threading.Event()
42
+ self.adjustment_request: dict[str, Any] = {}
43
 
44
  def pause(self) -> None:
45
  self.run_event.clear()
 
51
  self.cancel_event.set()
52
  self.run_event.set()
53
 
54
+ def request_adjustment(self, updates: dict[str, Any]) -> None:
55
+ self.adjustment_request.clear()
56
+ self.adjustment_request.update(updates)
57
+ self.adjustment_event.set()
58
+ self.run_event.set()
59
+
60
  def run(self) -> None:
61
  total_steps = len(self.job.plan.steps)
62
  try:
 
72
  )
73
 
74
  preview_state = {"epoch": 0, "path": ""}
75
+ last_progress_emit = {"time": 0.0, "overall": -1, "message": ""}
76
+ progress_samples: list[dict[str, float]] = []
77
 
78
+ def on_progress(percent: int, message: str, step_index: int = index, **details: Any) -> None:
79
  overall = int(((step_index + percent / 100) / total_steps) * 100)
80
+ now = time.monotonic()
81
+ progress_eta = self._estimate_step_eta(details, progress_samples, now)
82
+ changed = overall != last_progress_emit["overall"] or message != last_progress_emit["message"]
83
+ terminal = percent >= 100 or overall >= 100
84
+ if not terminal and (not changed or now - last_progress_emit["time"] < 0.25):
85
+ return
86
+ last_progress_emit.update({"time": now, "overall": overall, "message": message})
87
  self.event.emit(
88
  {
89
  "type": "progress",
90
  "step_percent": percent,
91
  "overall": overall,
92
  "message": message,
93
+ **progress_eta,
94
  }
95
  )
96
  self._discover_external_preview(step, message, preview_state)
 
112
  progress_callback=on_progress,
113
  log_callback=on_log,
114
  preview_callback=on_preview,
115
+ adjustment_event=self.adjustment_event,
116
+ adjustment_request=self.adjustment_request,
117
  )
118
  self.event.emit(
119
  {
 
125
  self.event.emit({"type": "completed"})
126
  except ToolCancelled as exc:
127
  self.event.emit({"type": "cancelled", "message": str(exc)})
128
+ except ToolAdjustmentRequested as exc:
129
+ self.event.emit({"type": "adjustment_ready", "message": str(exc), **exc.details})
130
  except Exception as exc:
131
  self.event.emit(
132
  {
 
172
  "steps": int(step.arguments.get("preview_steps", 0) or 0),
173
  })
174
 
175
+ @staticmethod
176
+ def _estimate_step_eta(
177
+ details: dict[str, Any],
178
+ samples: list[dict[str, float]],
179
+ now: float,
180
+ ) -> dict[str, Any]:
181
+ """Estimate remaining runtime from real step cadence instead of percent alone."""
182
+ current = _safe_int(details.get("current_step", details.get("step", details.get("current"))))
183
+ total = _safe_int(details.get("total_steps", details.get("total")))
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] = {
190
+ "progress_current": current,
191
+ "progress_total": total,
192
+ "progress_unit": unit,
193
+ }
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
203
+
204
+ first = samples[0]
205
+ elapsed = now - first["time"]
206
+ completed = current - int(first["current"])
207
+ if completed <= 0 or elapsed <= 0:
208
+ return payload
209
+ lifetime_seconds_per_unit = elapsed / completed
210
+
211
+ recent = samples[-8:]
212
+ recent_first = recent[0]
213
+ recent_completed = current - int(recent_first["current"])
214
+ recent_elapsed = now - recent_first["time"]
215
+ recent_seconds_per_unit = (
216
+ recent_elapsed / recent_completed
217
+ if recent_completed > 0 and recent_elapsed > 0
218
+ else lifetime_seconds_per_unit
219
+ )
220
+ seconds_per_unit = (recent_seconds_per_unit * 0.65) + (lifetime_seconds_per_unit * 0.35)
221
+ remaining = max(0, int(round((total - current) * seconds_per_unit)))
222
+ if remaining:
223
+ payload["eta_seconds"] = remaining
224
+ payload["progress_rate"] = 1 / seconds_per_unit if seconds_per_unit > 0 else 0.0
225
+ payload["estimated_completion_at"] = (
226
+ datetime.now(timezone.utc) + timedelta(seconds=remaining)
227
+ ).isoformat()
228
+ return payload
229
+
230
 
231
  class JobManager(QObject):
232
  job_created = Signal(object)
 
246
  self.root = root
247
  self.executor = executor
248
  self.logger = logger
249
+ self.config = config if config is not None else {}
250
  self.jobs_path = root / "data" / "jobs.json"
251
  self.assets = AssetRegistry(root)
252
  self.atlas = AtlasSupervisor(config)
253
+ self.experiments = ExperimentStore(root)
254
  self.jobs: list[Job] = []
255
  self._queue: list[str] = []
256
  self._worker: JobWorker | None = None
257
  self._active_job: Job | None = None
258
+ self._last_snapshot = None
259
+ self._pending_update_job_ids: set[str] = set()
260
+ self._pending_update_timer_active = False
261
  self._load()
262
+ self._schedule_timer = QTimer(self)
263
+ self._schedule_timer.timeout.connect(self._release_due_scheduled)
264
+ if QCoreApplication.instance() is not None:
265
+ self._schedule_timer.start(15_000)
266
+ QTimer.singleShot(0, self._release_due_scheduled)
267
+ if self._queue:
268
+ QTimer.singleShot(0, self._start_next)
269
 
270
  @property
271
  def active_job(self) -> Job | None:
272
  return self._active_job
273
 
274
+ @property
275
+ def pending_jobs(self) -> list[Job]:
276
+ return [
277
+ job for job in self.jobs
278
+ if job.status in {JobStatus.AWAITING_CONFIRMATION, JobStatus.SCHEDULED, JobStatus.QUEUED}
279
+ ]
280
+
281
+ def submit(self, plan: ExecutionPlan, scheduled_for: str | None = None) -> Job:
282
+ # Every entry point must review a plan before its queue state is chosen.
283
+ # UI and Remote may prepare it earlier to keep filesystem work off Qt.
284
+ append_preflight_summary(plan, self.config)
285
+ is_future = self._is_future(scheduled_for)
286
  status = (
287
  JobStatus.AWAITING_CONFIRMATION
288
  if plan.requires_confirmation
289
+ else JobStatus.SCHEDULED if is_future else JobStatus.QUEUED
290
  )
291
+ job = Job(plan=plan, status=status, scheduled_for=scheduled_for if is_future else None)
292
  self.jobs.insert(0, job)
293
  self._append_log(job, f"Plan created: {plan.summary}")
294
  if plan.requires_confirmation:
295
  self._append_log(job, "Waiting for user confirmation.")
296
+ if is_future:
297
+ self._append_log(job, f"Requested start time: {self._display_time(scheduled_for)}.")
298
+ elif is_future:
299
+ self._append_log(job, f"Scheduled for {self._display_time(scheduled_for)}.")
300
  else:
301
  self._queue.append(job.id)
302
  self._save()
 
310
  job = self.get(job_id)
311
  if job.status != JobStatus.AWAITING_CONFIRMATION:
312
  return
313
+ job.status = JobStatus.SCHEDULED if self._is_future(job.scheduled_for) else JobStatus.QUEUED
314
  self._append_log(job, "Plan approved by user.")
315
+ if job.status == JobStatus.SCHEDULED:
316
+ self._append_log(job, f"Training will become eligible at {self._display_time(job.scheduled_for)}.")
317
+ else:
318
+ self._queue.append(job.id)
319
  self._save()
320
  self.job_updated.emit(job)
321
  self._start_next()
 
353
  if job is self._active_job and self._worker:
354
  self._append_log(job, "Cancellation requested.")
355
  self._worker.cancel()
356
+ self._save()
357
+ self.job_updated.emit(job)
358
  return
359
  if job.id in self._queue:
360
  self._queue.remove(job.id)
361
  if job.status in {
362
+ JobStatus.SCHEDULED,
363
  JobStatus.QUEUED,
364
  JobStatus.AWAITING_CONFIRMATION,
365
  JobStatus.DRAFT,
 
370
  self._save()
371
  self.job_updated.emit(job)
372
 
373
+ def request_training_adjustment(self, job_id: str, updates: dict[str, Any]) -> None:
374
+ """Apply safe DDPM settings after the current epoch and resume automatically."""
375
+ job = self.get(job_id)
376
+ if job is not self._active_job or job.status not in {JobStatus.RUNNING, JobStatus.PAUSED} or not self._worker:
377
+ raise ValueError("Only the active training job can be adjusted.")
378
+ if not (0 <= job.current_step < len(job.plan.steps)):
379
+ raise ValueError("The active training step is unavailable.")
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 \
386
+ or not 10 <= cleaned.get("training_intensity", 100) <= 100 \
387
+ or not 1 <= cleaned.get("gradient_accumulation_steps", 1) <= 64:
388
+ raise ValueError("The requested training settings are outside ADAM's safe range.")
389
+ previous = {key: step.arguments.get(key) for key in cleaned}
390
+ if all(previous[key] == value for key, value in cleaned.items()):
391
+ raise ValueError("Those settings are already active.")
392
+ job.status = JobStatus.RUNNING
393
+ self._append_log(job, f"Adjustment queued for the end of this epoch: {cleaned}.")
394
+ self._worker.request_adjustment(cleaned)
395
+ self._save()
396
+ self.job_updated.emit(job)
397
+
398
+ def safer_vram_retry(self, job_id: str) -> Job:
399
+ """Create a checkpoint-aware DDPM retry with a smaller physical batch."""
400
+ original = self.get(job_id)
401
+ if original.status != JobStatus.FAILED or not self._looks_like_vram_failure(original):
402
+ raise ValueError("This job did not fail with a recognizable VRAM error.")
403
+ plan = ExecutionPlan.from_dict(original.to_dict()["plan"])
404
+ start_index = max(0, min(original.current_step, len(plan.steps) - 1))
405
+ plan.steps = plan.steps[start_index:]
406
+ step = plan.steps[0]
407
+ old_batch = max(1, int(step.arguments.get("batch_size", 1)))
408
+ if old_batch <= 1:
409
+ raise ValueError("Batch size is already 1; lower resolution or enable other memory-saving options.")
410
+ original_epochs = max(1, int(step.arguments.get("epochs", 1)))
411
+ resume_note = self._prepare_ddpm_resume(step.arguments, step.tool_id)
412
+ remaining_epochs = max(1, int(step.arguments.get("epochs", original_epochs)))
413
+ completed_epochs = max(0, original_epochs - remaining_epochs) if resume_note else 0
414
+ new_batch = max(1, old_batch // 2)
415
+ old_accumulation = max(1, int(step.arguments.get("gradient_accumulation_steps", 1)))
416
+ step.arguments["batch_size"] = new_batch
417
+ step.arguments["gradient_accumulation_steps"] = min(64, old_accumulation * max(1, math.ceil(old_batch / new_batch)))
418
+ if completed_epochs:
419
+ step.arguments["completed_epochs"] = completed_epochs
420
+ for item in plan.steps:
421
+ item.status = StepStatus.PENDING
422
+ plan.id = original.plan.id + "-vram-retry"
423
+ plan.created_at = utc_now()
424
+ plan.requires_confirmation = True
425
+ plan.confirmation_reason = "VRAM recovery reduced the physical batch and preserved the effective batch with gradient accumulation."
426
+ retry = self.submit(plan)
427
+ self._append_log(retry, f"VRAM recovery changed batch {old_batch} → {new_batch} and gradient accumulation {old_accumulation} → {step.arguments['gradient_accumulation_steps']}.")
428
+ if resume_note:
429
+ self._append_log(retry, resume_note)
430
+ return retry
431
+
432
+ @staticmethod
433
+ def _looks_like_vram_failure(job: Job) -> bool:
434
+ text = "\n".join([job.error or "", *job.logs[-100:]]).lower()
435
+ return any(token in text for token in ("out of memory", "cuda oom", "cuda error: out of memory"))
436
+
437
+ @staticmethod
438
+ def _is_future(value: str | None) -> bool:
439
+ if not value:
440
+ return False
441
+ try:
442
+ scheduled = datetime.fromisoformat(value)
443
+ if scheduled.tzinfo is None:
444
+ scheduled = scheduled.astimezone()
445
+ return scheduled.astimezone(timezone.utc) > datetime.now(timezone.utc)
446
+ except (TypeError, ValueError):
447
+ return False
448
+
449
+ @staticmethod
450
+ def _display_time(value: str | None) -> str:
451
+ try:
452
+ return datetime.fromisoformat(str(value)).astimezone().strftime("%b %d at %I:%M %p")
453
+ except ValueError:
454
+ return str(value or "the requested time")
455
+
456
+ def _release_due_scheduled(self) -> None:
457
+ released: list[Job] = []
458
+ for job in reversed(self.jobs):
459
+ if job.status == JobStatus.SCHEDULED and not self._is_future(job.scheduled_for):
460
+ job.status = JobStatus.QUEUED
461
+ self._queue.append(job.id)
462
+ self._append_log(job, "Scheduled start time reached; waiting for the training slot.")
463
+ released.append(job)
464
+ if not released:
465
+ return
466
+ self._save()
467
+ for job in released:
468
+ self.job_updated.emit(job)
469
+ self._start_next()
470
+
471
  def get(self, job_id: str) -> Job:
472
  for job in self.jobs:
473
  if job.id == job_id:
 
601
  job.preview_total = 0
602
  job.preview_image_index = 0
603
  job.preview_image_count = 0
604
+ job.eta_seconds = None
605
+ job.estimated_completion_at = None
606
+ job.progress_current = 0
607
+ job.progress_total = 0
608
+ job.progress_rate = 0.0
609
+ job.progress_unit = "step"
610
  self._append_log(job, f"Starting: {job.plan.steps[index].title}")
611
  elif event_type == "progress":
612
  job.progress = int(event["overall"])
613
+ job.eta_seconds = _safe_int(event.get("eta_seconds")) or None
614
+ job.estimated_completion_at = str(event.get("estimated_completion_at") or "") or None
615
+ job.progress_current = _safe_int(event.get("progress_current"))
616
+ job.progress_total = _safe_int(event.get("progress_total"))
617
+ job.progress_rate = float(event.get("progress_rate", 0.0) or 0.0)
618
+ job.progress_unit = str(event.get("progress_unit", "step") or "step")
619
  message = str(event["message"])
620
  if message and (not job.logs or message not in job.logs[-1]):
621
  self._append_log(job, message)
622
+ self._schedule_job_update(job)
623
+ return
624
  elif event_type == "log":
625
  self._append_log(job, str(event["message"]))
626
+ self._schedule_job_update(job)
627
+ return
628
  elif event_type == "preview":
629
  job.preview_path = str(event.get("path", "")) or None
630
  job.preview_epoch = int(event.get("epoch", 0) or 0)
 
641
  label = "Denoising" if job.preview_kind == "generation" else "Training"
642
  position = f" step {job.preview_current}" if job.preview_current else f" epoch {job.preview_epoch}"
643
  self._append_log(job, f"{label} preview updated at{position}.")
644
+ self._schedule_job_update(job)
645
+ return
646
  elif event_type == "step_finished":
647
  index = int(event["index"])
648
  job.plan.steps[index].status = StepStatus.FINISHED
 
668
  elif event_type == "completed":
669
  job.status = JobStatus.FINISHED
670
  job.progress = 100
671
+ job.eta_seconds = 0
672
+ job.estimated_completion_at = utc_now()
673
  job.ended_at = utc_now()
674
  self._append_log(job, "Job finished successfully.")
675
+ self._record_experiment(job)
676
  self.notification.emit("Job complete", job.plan.project_name)
677
  elif event_type == "cancelled":
678
  job.status = JobStatus.CANCELLED
 
681
  for step in job.plan.steps:
682
  if step.status == StepStatus.RUNNING:
683
  step.status = StepStatus.SKIPPED
684
+ self._record_experiment(job)
685
  self.notification.emit("Job cancelled", job.plan.project_name)
686
+ elif event_type == "adjustment_ready":
687
+ index = max(0, min(job.current_step, len(job.plan.steps) - 1))
688
+ remaining_steps = job.plan.steps[index:]
689
+ step = remaining_steps[0]
690
+ updates = dict(event.get("updates") or {})
691
+ step.arguments.update(updates)
692
+ checkpoint = str(event.get("checkpoint", ""))
693
+ completed_epochs = max(0, int(event.get("completed_epochs", 0) or 0))
694
+ total_epochs = max(1, int(step.arguments.get("epochs", 1)))
695
+ step.arguments["epochs"] = max(1, total_epochs - completed_epochs)
696
+ step.arguments["completed_epochs"] = completed_epochs
697
+ if checkpoint:
698
+ step.arguments["resume_from"] = checkpoint
699
+ for pending in remaining_steps:
700
+ pending.status = StepStatus.PENDING
701
+ job.plan.steps = remaining_steps
702
+ job.current_step = -1
703
+ job.status = JobStatus.QUEUED
704
+ job.eta_seconds = None
705
+ job.estimated_completion_at = None
706
+ self._queue.insert(0, job.id)
707
+ self._append_log(job, f"Epoch {completed_epochs} checkpoint is complete. Restarting with {updates}.")
708
+ self.notification.emit("Training settings ready", "Restarting from the completed epoch checkpoint.")
709
  elif event_type == "failed":
710
  job.status = JobStatus.FAILED
711
  job.ended_at = utc_now()
 
719
  event.get("exception"),
720
  job.error,
721
  )
722
+ self._record_experiment(job)
723
  self.notification.emit("Job failed", job.error)
724
  self._save()
725
  self.job_updated.emit(job)
726
 
727
+ def _schedule_job_update(self, job: Job) -> None:
728
+ self._pending_update_job_ids.add(job.id)
729
+ if QCoreApplication.instance() is None:
730
+ self._flush_pending_job_updates()
731
+ return
732
+ if self._pending_update_timer_active:
733
+ return
734
+ self._pending_update_timer_active = True
735
+ QTimer.singleShot(300, self._flush_pending_job_updates)
736
+
737
+ def _flush_pending_job_updates(self) -> None:
738
+ if not self._pending_update_job_ids:
739
+ self._pending_update_timer_active = False
740
+ return
741
+ pending_ids = list(self._pending_update_job_ids)
742
+ self._pending_update_job_ids.clear()
743
+ self._pending_update_timer_active = False
744
+ self._save()
745
+ for job_id in pending_ids:
746
+ try:
747
+ self.job_updated.emit(self.get(job_id))
748
+ except KeyError:
749
+ continue
750
+
751
  def _worker_finished(self) -> None:
752
  self._worker = None
753
  self._active_job = None
 
756
 
757
  def supervise(self, snapshot: Any) -> None:
758
  """Let ATLAS inspect the active training run and apply critical pauses."""
759
+ self._last_snapshot = snapshot
760
  job = self._active_job
761
  if job is None or job.status != JobStatus.RUNNING:
762
  return
 
784
  if decision.action == "pause" and job.status == JobStatus.RUNNING:
785
  self.pause(job.id)
786
 
787
+ def _record_experiment(self, job: Job) -> None:
788
+ if not self._has_training(job):
789
+ return
790
+ try:
791
+ self.experiments.record_job(job, self._last_snapshot)
792
+ except Exception as exc:
793
+ self.logger.warning("Experiment tracking failed for %s: %s", job.id, exc)
794
+
795
  def _append_log(self, job: Job, message: str) -> None:
796
  timestamp = datetime.now().strftime("%H:%M:%S")
797
  line = f"[{timestamp}] {message}"
 
814
  self.jobs = []
815
  return
816
  for job in self.jobs:
817
+ if job.status in {JobStatus.RUNNING, JobStatus.PAUSED}:
818
  job.status = JobStatus.INTERRUPTED
819
  job.ended_at = utc_now()
820
  job.logs.append(
821
  "[startup] Previous session ended before this job. "
822
  "Review it before retrying."
823
  )
824
+ elif job.status == JobStatus.QUEUED:
825
+ self._queue.append(job.id)
826
+ if not any("Queued job restored" in line for line in job.logs[-5:]):
827
+ job.logs.append("[startup] Queued job restored and will run when ADAM is ready.")
828
  self._save()
829
 
830
  def _save(self) -> None:
adam/model_inspector/__init__.py ADDED
@@ -0,0 +1,15 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from __future__ import annotations
2
+
3
+ from .base import ModelComparison, ModelInspection, TensorComparison, TensorStats
4
+ from .comparison import compare_models
5
+ from .detector import inspect_model, inspector_for
6
+
7
+ __all__ = [
8
+ "ModelComparison",
9
+ "ModelInspection",
10
+ "TensorComparison",
11
+ "TensorStats",
12
+ "compare_models",
13
+ "inspect_model",
14
+ "inspector_for",
15
+ ]
adam/model_inspector/base.py ADDED
@@ -0,0 +1,146 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from __future__ import annotations
2
+
3
+ from dataclasses import dataclass, field
4
+ from pathlib import Path
5
+ from typing import Any, Callable, Iterable
6
+
7
+
8
+ MODEL_EXTENSIONS = {".safetensors", ".pt", ".pth", ".bin", ".ckpt"}
9
+ CONFIG_FILENAMES = {
10
+ "model_index.json",
11
+ "config.json",
12
+ "scheduler_config.json",
13
+ "flow_model_info.json",
14
+ "adapter_config.json",
15
+ }
16
+
17
+
18
+ @dataclass(slots=True)
19
+ class TensorStats:
20
+ name: str
21
+ shape: tuple[int, ...]
22
+ dtype: str
23
+ parameter_count: int
24
+ memory_bytes: int
25
+ minimum: float | None = None
26
+ maximum: float | None = None
27
+ mean: float | None = None
28
+ std: float | None = None
29
+ abs_mean: float | None = None
30
+ l2_norm: float | None = None
31
+ zero_percent: float | None = None
32
+ component: str = "other"
33
+ health: list[str] = field(default_factory=list)
34
+
35
+
36
+ @dataclass(slots=True)
37
+ class ModelInspection:
38
+ path: str
39
+ resolved_path: str
40
+ architecture: str
41
+ confidence: float
42
+ status: str
43
+ size_bytes: int
44
+ config_files: list[str]
45
+ resolution: int | None
46
+ epoch: int | None
47
+ step: int | None
48
+ tensor_count: int
49
+ total_parameters: int
50
+ trainable_parameters: int | None
51
+ parameter_memory_bytes: int
52
+ dtypes: dict[str, int]
53
+ components: dict[str, int]
54
+ largest_tensors: list[TensorStats]
55
+ tensors: list[TensorStats]
56
+ health: list[str]
57
+ messages: list[str]
58
+ lora: dict[str, Any] = field(default_factory=dict)
59
+ configs: dict[str, Any] = field(default_factory=dict)
60
+ histogram: dict[str, list[float]] = field(default_factory=dict)
61
+ tensor_size_distribution: list[tuple[str, int]] = field(default_factory=list)
62
+ checkpoints: list[str] = field(default_factory=list)
63
+ loss_history: list[tuple[int, float]] = field(default_factory=list)
64
+
65
+
66
+ @dataclass(slots=True)
67
+ class TensorComparison:
68
+ name: str
69
+ shape: tuple[int, ...]
70
+ component: str
71
+ mean_abs_difference: float | None
72
+ relative_difference: float | None
73
+ cosine_similarity: float | None
74
+ l2_distance: float | None
75
+ drift: float | None
76
+ change_score: float | None
77
+
78
+
79
+ @dataclass(slots=True)
80
+ class ModelComparison:
81
+ path_a: str
82
+ path_b: str
83
+ architecture_a: str
84
+ architecture_b: str
85
+ architecture_match: bool
86
+ config_differences: list[str]
87
+ resolution_difference: tuple[int | None, int | None] | None
88
+ parameter_count_difference: int
89
+ only_a: list[str]
90
+ only_b: list[str]
91
+ shape_mismatches: list[str]
92
+ tensor_comparisons: list[TensorComparison]
93
+ group_comparisons: dict[str, dict[str, float]]
94
+ messages: list[str]
95
+
96
+
97
+ ProgressCallback = Callable[[int, str], None]
98
+ CancelCallback = Callable[[], bool]
99
+
100
+
101
+ class InspectorError(RuntimeError):
102
+ pass
103
+
104
+
105
+ class BaseModelInspector:
106
+ architecture = "Generic / Unknown"
107
+
108
+ def inspect(
109
+ self,
110
+ path: str | Path,
111
+ *,
112
+ recorded_architecture: str = "",
113
+ run_settings: dict[str, Any] | None = None,
114
+ progress: ProgressCallback | None = None,
115
+ cancelled: CancelCallback | None = None,
116
+ ) -> ModelInspection:
117
+ raise NotImplementedError
118
+
119
+
120
+ def report(progress: ProgressCallback | None, value: int, message: str) -> None:
121
+ if progress:
122
+ progress(max(0, min(100, int(value))), message)
123
+
124
+
125
+ def is_cancelled(cancelled: CancelCallback | None) -> bool:
126
+ return bool(cancelled and cancelled())
127
+
128
+
129
+ def parameter_count(shape: Iterable[int]) -> int:
130
+ total = 1
131
+ for dim in shape:
132
+ total *= int(dim)
133
+ return int(total)
134
+
135
+
136
+ def dtype_size(dtype: str) -> int:
137
+ lowered = dtype.casefold()
138
+ if "float64" in lowered or "int64" in lowered:
139
+ return 8
140
+ if "float32" in lowered or "int32" in lowered:
141
+ return 4
142
+ if "float16" in lowered or "bfloat16" in lowered or "int16" in lowered:
143
+ return 2
144
+ if "int8" in lowered or "uint8" in lowered or "bool" in lowered:
145
+ return 1
146
+ return 4
adam/model_inspector/comparison.py ADDED
@@ -0,0 +1,160 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from __future__ import annotations
2
+
3
+ import json
4
+ from pathlib import Path
5
+ from typing import Any
6
+
7
+ from .base import ModelComparison, TensorComparison, is_cancelled, report
8
+ from .detector import inspect_model
9
+ from .generic import _extract_state_dict, _weight_files
10
+ from .statistics import component_for_name
11
+
12
+
13
+ def _iter_named_tensors(path: Path):
14
+ files = _weight_files(path)
15
+ for file in files:
16
+ if file.suffix.casefold() == ".safetensors":
17
+ from safetensors import safe_open
18
+
19
+ with safe_open(str(file), framework="pt", device="cpu") as handle:
20
+ for key in handle.keys():
21
+ yield key, handle.get_tensor(key)
22
+ else:
23
+ import torch
24
+
25
+ payload = torch.load(str(file), map_location="cpu", weights_only=False)
26
+ state = _extract_state_dict(payload)
27
+ for key, value in state.items():
28
+ if hasattr(value, "shape"):
29
+ yield str(key), value
30
+
31
+
32
+ def _tensor_map(path: Path) -> dict[str, Any]:
33
+ return {name: tensor for name, tensor in _iter_named_tensors(path)}
34
+
35
+
36
+ def _config_differences(configs_a: dict[str, Any], configs_b: dict[str, Any], *, limit: int = 40) -> list[str]:
37
+ diffs: list[str] = []
38
+ keys = sorted(set(configs_a) | set(configs_b))
39
+ for key in keys:
40
+ if key not in configs_a:
41
+ diffs.append(f"Only B has config {key}")
42
+ elif key not in configs_b:
43
+ diffs.append(f"Only A has config {key}")
44
+ elif json.dumps(configs_a[key], sort_keys=True, default=str) != json.dumps(configs_b[key], sort_keys=True, default=str):
45
+ diffs.append(f"Config differs: {key}")
46
+ if len(diffs) >= limit:
47
+ diffs.append("Additional config differences omitted.")
48
+ break
49
+ return diffs
50
+
51
+
52
+ def compare_models(
53
+ path_a: str | Path,
54
+ path_b: str | Path,
55
+ *,
56
+ arch_a: str = "",
57
+ arch_b: str = "",
58
+ settings_a: dict[str, Any] | None = None,
59
+ settings_b: dict[str, Any] | None = None,
60
+ progress=None,
61
+ cancelled=None,
62
+ ) -> ModelComparison:
63
+ report(progress, 2, "Inspecting first model")
64
+ summary_a = inspect_model(path_a, recorded_architecture=arch_a, run_settings=settings_a, progress=progress, cancelled=cancelled)
65
+ report(progress, 30, "Inspecting second model")
66
+ summary_b = inspect_model(path_b, recorded_architecture=arch_b, run_settings=settings_b, progress=progress, cancelled=cancelled)
67
+ report(progress, 55, "Loading comparable tensors")
68
+ tensors_a = _tensor_map(Path(summary_a.resolved_path))
69
+ if is_cancelled(cancelled):
70
+ raise RuntimeError("Comparison cancelled.")
71
+ tensors_b = _tensor_map(Path(summary_b.resolved_path))
72
+ names_a = set(tensors_a)
73
+ names_b = set(tensors_b)
74
+ common = sorted(names_a & names_b)
75
+ only_a = sorted(names_a - names_b)[:200]
76
+ only_b = sorted(names_b - names_a)[:200]
77
+ shape_mismatches = []
78
+ comparable = []
79
+ for name in common:
80
+ if tuple(tensors_a[name].shape) != tuple(tensors_b[name].shape):
81
+ shape_mismatches.append(name)
82
+ else:
83
+ comparable.append(name)
84
+ comparisons: list[TensorComparison] = []
85
+ import torch
86
+
87
+ for index, name in enumerate(comparable):
88
+ if is_cancelled(cancelled):
89
+ raise RuntimeError("Comparison cancelled.")
90
+ if index % 10 == 0:
91
+ report(progress, 58 + int(36 * index / max(1, len(comparable))), f"Comparing {index + 1} of {len(comparable)} tensors")
92
+ with torch.no_grad():
93
+ a = tensors_a[name].detach().to(device="cpu").float().reshape(-1)
94
+ b = tensors_b[name].detach().to(device="cpu").float().reshape(-1)
95
+ if a.numel() == 0:
96
+ continue
97
+ limit = 1_000_000
98
+ if a.numel() > limit:
99
+ stride = max(1, a.numel() // limit)
100
+ a = a[::stride][:limit]
101
+ b = b[::stride][:limit]
102
+ delta = b - a
103
+ mean_abs = delta.abs().mean().item()
104
+ base_abs = a.abs().mean().item()
105
+ relative = mean_abs / (base_abs + 1e-12)
106
+ l2 = torch.linalg.vector_norm(delta).item()
107
+ norm_a = torch.linalg.vector_norm(a).item()
108
+ norm_b = torch.linalg.vector_norm(b).item()
109
+ cosine = torch.nn.functional.cosine_similarity(a, b, dim=0).item() if norm_a and norm_b else None
110
+ drift = (norm_b - norm_a) / (norm_a + 1e-12) if norm_a else None
111
+ score = relative * 0.7 + (1 - cosine if cosine is not None else 0) * 0.3
112
+ comparisons.append(
113
+ TensorComparison(
114
+ name=name,
115
+ shape=tuple(int(dim) for dim in tensors_a[name].shape),
116
+ component=component_for_name(name),
117
+ mean_abs_difference=float(mean_abs),
118
+ relative_difference=float(relative),
119
+ cosine_similarity=float(cosine) if cosine is not None else None,
120
+ l2_distance=float(l2),
121
+ drift=float(drift) if drift is not None else None,
122
+ change_score=float(score),
123
+ )
124
+ )
125
+ comparisons.sort(key=lambda item: item.change_score or 0, reverse=True)
126
+ groups: dict[str, dict[str, float]] = {}
127
+ for item in comparisons:
128
+ group = groups.setdefault(item.component, {"tensors": 0, "mean_change_score": 0.0, "mean_abs_difference": 0.0})
129
+ group["tensors"] += 1
130
+ group["mean_change_score"] += item.change_score or 0
131
+ group["mean_abs_difference"] += item.mean_abs_difference or 0
132
+ for group in groups.values():
133
+ count = max(1, int(group["tensors"]))
134
+ group["mean_change_score"] /= count
135
+ group["mean_abs_difference"] /= count
136
+ messages = [
137
+ "Change Score is a statistical weight-change metric; it does not directly equal behavioral importance."
138
+ ]
139
+ if shape_mismatches:
140
+ messages.append("Some tensor comparisons are unavailable because tensor shapes differ.")
141
+ report(progress, 100, "Comparison complete")
142
+ return ModelComparison(
143
+ path_a=summary_a.resolved_path,
144
+ path_b=summary_b.resolved_path,
145
+ architecture_a=summary_a.architecture,
146
+ architecture_b=summary_b.architecture,
147
+ architecture_match=summary_a.architecture == summary_b.architecture,
148
+ config_differences=_config_differences(summary_a.configs, summary_b.configs),
149
+ resolution_difference=(
150
+ (summary_a.resolution, summary_b.resolution)
151
+ if summary_a.resolution != summary_b.resolution else None
152
+ ),
153
+ parameter_count_difference=summary_b.total_parameters - summary_a.total_parameters,
154
+ only_a=only_a,
155
+ only_b=only_b,
156
+ shape_mismatches=shape_mismatches[:200],
157
+ tensor_comparisons=comparisons[:200],
158
+ group_comparisons=groups,
159
+ messages=messages,
160
+ )
adam/model_inspector/ddpm.py ADDED
@@ -0,0 +1,29 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from __future__ import annotations
2
+
3
+ from pathlib import Path
4
+ from typing import Any
5
+
6
+ from .generic import GenericModelInspector
7
+
8
+
9
+ class DDPMInspector(GenericModelInspector):
10
+ architecture = "DDPM / Diffusers"
11
+
12
+ def inspect(self, path: str | Path, *, recorded_architecture: str = "", run_settings: dict[str, Any] | None = None, progress=None, cancelled=None):
13
+ summary = super().inspect(
14
+ path,
15
+ recorded_architecture=recorded_architecture or "ddpm",
16
+ run_settings=run_settings,
17
+ progress=progress,
18
+ cancelled=cancelled,
19
+ )
20
+ if summary.architecture == "Generic / Unknown":
21
+ summary.architecture = self.architecture
22
+ summary.messages.insert(0, "Model recognized as DDPM from ADAM trainer metadata")
23
+ expected = ("unet", "scheduler")
24
+ root = Path(summary.resolved_path)
25
+ existing = {part.name.casefold() for part in (root.iterdir() if root.is_dir() else [])}
26
+ for component in expected:
27
+ if root.is_dir() and component not in existing and not any(component in item.casefold() for item in summary.configs):
28
+ summary.health.append(f"Unusual: expected DDPM component not found: {component}")
29
+ return summary
adam/model_inspector/detector.py ADDED
@@ -0,0 +1,35 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from __future__ import annotations
2
+
3
+ from pathlib import Path
4
+ from typing import Any
5
+
6
+ from .ddpm import DDPMInspector
7
+ from .flow_matching import FlowMatchingInspector
8
+ from .generic import GenericModelInspector
9
+ from .lora import LoRAInspector
10
+ from .maskgit import MaskGITInspector
11
+
12
+
13
+ def inspector_for(path: str | Path, recorded_architecture: str = "", settings: dict[str, Any] | None = None):
14
+ text = " ".join((str(path), recorded_architecture, str(settings or {}))).casefold()
15
+ target = Path(path).expanduser()
16
+ if "lora" in text or target.suffix.casefold() == ".safetensors" and "adapter" in target.name.casefold():
17
+ return LoRAInspector()
18
+ if "maskgit" in text:
19
+ return MaskGITInspector()
20
+ if "flow" in text or (target / "flow_model_info.json").is_file():
21
+ return FlowMatchingInspector()
22
+ if "ddpm" in text or (target / "model_index.json").is_file() or (target / "scheduler").is_dir():
23
+ return DDPMInspector()
24
+ return GenericModelInspector()
25
+
26
+
27
+ def inspect_model(path: str | Path, *, recorded_architecture: str = "", run_settings: dict[str, Any] | None = None, progress=None, cancelled=None):
28
+ inspector = inspector_for(path, recorded_architecture, run_settings)
29
+ return inspector.inspect(
30
+ path,
31
+ recorded_architecture=recorded_architecture,
32
+ run_settings=run_settings,
33
+ progress=progress,
34
+ cancelled=cancelled,
35
+ )
adam/model_inspector/flow_matching.py ADDED
@@ -0,0 +1,25 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from __future__ import annotations
2
+
3
+ from pathlib import Path
4
+ from typing import Any
5
+
6
+ from .generic import GenericModelInspector
7
+
8
+
9
+ class FlowMatchingInspector(GenericModelInspector):
10
+ architecture = "Flow Matching"
11
+
12
+ def inspect(self, path: str | Path, *, recorded_architecture: str = "", run_settings: dict[str, Any] | None = None, progress=None, cancelled=None):
13
+ summary = super().inspect(
14
+ path,
15
+ recorded_architecture=recorded_architecture or "flow",
16
+ run_settings=run_settings,
17
+ progress=progress,
18
+ cancelled=cancelled,
19
+ )
20
+ summary.architecture = "Flow Matching" if summary.architecture == "Generic / Unknown" else summary.architecture
21
+ if not any("flow_model_info.json" in item for item in summary.config_files):
22
+ root = Path(summary.resolved_path)
23
+ if root.is_dir():
24
+ summary.health.append("Unusual: Flow Matching metadata file was not found")
25
+ return summary
adam/model_inspector/generic.py ADDED
@@ -0,0 +1,349 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from __future__ import annotations
2
+
3
+ import json
4
+ from collections import Counter
5
+ from pathlib import Path
6
+ from typing import Any, Iterator
7
+
8
+ from .base import BaseModelInspector, InspectorError, ModelInspection, TensorStats, is_cancelled, report
9
+ from .statistics import (
10
+ discover_checkpoint_paths,
11
+ discover_config_files,
12
+ step_from_name,
13
+ tensor_stats_from_torch,
14
+ )
15
+
16
+
17
+ def _folder_size(path: Path) -> int:
18
+ if path.is_file():
19
+ return path.stat().st_size
20
+ total = 0
21
+ try:
22
+ for item in path.rglob("*"):
23
+ if item.is_file():
24
+ total += item.stat().st_size
25
+ except OSError:
26
+ return total
27
+ return total
28
+
29
+
30
+ def _read_configs(files: list[Path], root: Path) -> dict[str, Any]:
31
+ configs: dict[str, Any] = {}
32
+ for file in files[:40]:
33
+ try:
34
+ key = str(file.relative_to(root if root.is_dir() else root.parent))
35
+ except ValueError:
36
+ key = file.name
37
+ try:
38
+ configs[key] = json.loads(file.read_text(encoding="utf-8"))
39
+ except (OSError, UnicodeDecodeError, json.JSONDecodeError):
40
+ configs[key] = "<unreadable>"
41
+ return configs
42
+
43
+
44
+ def _resolution_from_configs(configs: dict[str, Any], settings: dict[str, Any] | None) -> int | None:
45
+ for source in (settings or {}, *[value for value in configs.values() if isinstance(value, dict)]):
46
+ for key in ("resolution", "sample_size", "image_size", "size"):
47
+ value = source.get(key) if isinstance(source, dict) else None
48
+ if isinstance(value, int):
49
+ return value
50
+ if isinstance(value, (list, tuple)) and value and isinstance(value[0], int):
51
+ return int(value[0])
52
+ try:
53
+ if value:
54
+ return int(value)
55
+ except (TypeError, ValueError):
56
+ pass
57
+ return None
58
+
59
+
60
+ def _iter_safetensors(file: Path) -> Iterator[tuple[str, Any, dict[str, Any]]]:
61
+ from safetensors import safe_open
62
+
63
+ with safe_open(str(file), framework="pt", device="cpu") as handle:
64
+ metadata = handle.metadata() or {}
65
+ for key in handle.keys():
66
+ yield key, handle.get_tensor(key), metadata
67
+
68
+
69
+ def _extract_state_dict(payload: Any) -> dict[str, Any]:
70
+ try:
71
+ import torch
72
+ except Exception:
73
+ torch = None
74
+ if torch is not None and hasattr(payload, "shape"):
75
+ return {"tensor": payload}
76
+ if isinstance(payload, dict):
77
+ for key in ("state_dict", "model_state_dict", "model", "module", "unet", "network"):
78
+ value = payload.get(key)
79
+ if isinstance(value, dict) and any(hasattr(item, "shape") for item in value.values()):
80
+ return value
81
+ if any(hasattr(item, "shape") for item in payload.values()):
82
+ return payload
83
+ return {}
84
+
85
+
86
+ def _iter_torch_checkpoint(file: Path) -> Iterator[tuple[str, Any, dict[str, Any]]]:
87
+ import torch
88
+
89
+ try:
90
+ payload = torch.load(str(file), map_location="cpu", weights_only=True)
91
+ except TypeError:
92
+ payload = torch.load(str(file), map_location="cpu")
93
+ except Exception:
94
+ payload = torch.load(str(file), map_location="cpu", weights_only=False)
95
+ state = _extract_state_dict(payload)
96
+ metadata = {key: value for key, value in payload.items() if key not in state} if isinstance(payload, dict) else {}
97
+ for key, tensor in state.items():
98
+ if hasattr(tensor, "shape"):
99
+ yield str(key), tensor, metadata
100
+
101
+
102
+ def _weight_files(path: Path) -> list[Path]:
103
+ if path.is_file():
104
+ return [path]
105
+ ignored_names = {"optimizer.bin", "scheduler.bin", "scaler.pt"}
106
+ names = {
107
+ "diffusion_pytorch_model.safetensors",
108
+ "model.safetensors",
109
+ "pytorch_model.bin",
110
+ "adapter_model.safetensors",
111
+ "adapter_model.bin",
112
+ "checkpoint.pt",
113
+ "best_checkpoint.pt",
114
+ }
115
+ files: list[Path] = []
116
+ try:
117
+ for item in path.rglob("*"):
118
+ if item.is_file() and (item.name in names or item.suffix.casefold() in {".safetensors", ".pt", ".pth", ".bin", ".ckpt"}):
119
+ if item.name.casefold() not in ignored_names:
120
+ files.append(item)
121
+ except OSError:
122
+ return []
123
+ if (path / "model_index.json").is_file():
124
+ final_files = [
125
+ item for item in files
126
+ if not any(part.startswith("checkpoint-") for part in item.relative_to(path).parts)
127
+ ]
128
+ if final_files:
129
+ files = final_files
130
+ return sorted(files, key=lambda item: (0 if item.name in names else 1, str(item)))
131
+
132
+
133
+ class GenericModelInspector(BaseModelInspector):
134
+ architecture = "Generic / Unknown"
135
+
136
+ def inspect(
137
+ self,
138
+ path: str | Path,
139
+ *,
140
+ recorded_architecture: str = "",
141
+ run_settings: dict[str, Any] | None = None,
142
+ progress=None,
143
+ cancelled=None,
144
+ ) -> ModelInspection:
145
+ target = Path(path).expanduser()
146
+ if not target.exists():
147
+ raise InspectorError(f"Model path does not exist: {target}")
148
+ target = target.resolve()
149
+ report(progress, 3, "Finding model files")
150
+ config_files = discover_config_files(target)
151
+ configs = _read_configs(config_files, target)
152
+ files = _weight_files(target)
153
+ if not files:
154
+ message = "Model contains no readable tensor checkpoint"
155
+ return self._empty(target, recorded_architecture, run_settings, config_files, configs, message)
156
+
157
+ tensors: list[TensorStats] = []
158
+ dtypes: Counter[str] = Counter()
159
+ components: Counter[str] = Counter()
160
+ health: list[str] = []
161
+ messages: list[str] = []
162
+ metadata: dict[str, Any] = {}
163
+ for file_index, file in enumerate(files):
164
+ if is_cancelled(cancelled):
165
+ raise InspectorError("Inspection cancelled.")
166
+ report(progress, 8 + int(80 * file_index / max(1, len(files))), f"Reading {file.name}")
167
+ try:
168
+ if file.suffix.casefold() == ".safetensors":
169
+ iterator = _iter_safetensors(file)
170
+ else:
171
+ iterator = _iter_torch_checkpoint(file)
172
+ for name, tensor, file_metadata in iterator:
173
+ if is_cancelled(cancelled):
174
+ raise InspectorError("Inspection cancelled.")
175
+ prefix = file.parent.name if len(files) > 1 else ""
176
+ stat = tensor_stats_from_torch(f"{prefix}.{name}" if prefix and not name.startswith(prefix) else name, tensor)
177
+ tensors.append(stat)
178
+ dtypes[stat.dtype] += stat.parameter_count
179
+ components[stat.component] += stat.parameter_count
180
+ health.extend(f"{stat.name}: {item}" for item in stat.health)
181
+ if file_metadata:
182
+ metadata.update(file_metadata)
183
+ except Exception as exc:
184
+ health.append(f"{file.name}: unreadable checkpoint ({exc})")
185
+
186
+ if not tensors:
187
+ message = "Model contains no readable tensor checkpoint"
188
+ return self._empty(target, recorded_architecture, run_settings, config_files, configs, message, [*(health or []), message])
189
+
190
+ report(progress, 92, "Summarizing model")
191
+ total_parameters = sum(tensor.parameter_count for tensor in tensors)
192
+ parameter_memory = sum(tensor.memory_bytes for tensor in tensors)
193
+ largest = sorted(tensors, key=lambda item: item.parameter_count, reverse=True)[:20]
194
+ architecture, confidence, message = self._architecture_from_signals(
195
+ target, recorded_architecture, configs, [tensor.name for tensor in tensors]
196
+ )
197
+ messages.append(message)
198
+ duplicate_count = len(tensors) - len({tensor.name for tensor in tensors})
199
+ if duplicate_count:
200
+ health.append(f"Unusual: {duplicate_count} duplicate tensor names after folder merging")
201
+ checkpoints = [str(item) for item in discover_checkpoint_paths(target)]
202
+ return ModelInspection(
203
+ path=str(path),
204
+ resolved_path=str(target),
205
+ architecture=architecture,
206
+ confidence=confidence,
207
+ status="ok",
208
+ size_bytes=_folder_size(target),
209
+ config_files=[str(item) for item in config_files],
210
+ resolution=_resolution_from_configs(configs, run_settings),
211
+ epoch=self._number_from_metadata(metadata, "epoch"),
212
+ step=self._number_from_metadata(metadata, "step") or step_from_name(target.name),
213
+ tensor_count=len(tensors),
214
+ total_parameters=total_parameters,
215
+ trainable_parameters=self._trainable_parameters(tensors, architecture),
216
+ parameter_memory_bytes=parameter_memory,
217
+ dtypes=dict(dtypes),
218
+ components=dict(components),
219
+ largest_tensors=largest,
220
+ tensors=tensors,
221
+ health=health or ["No invalid tensor values found in sampled statistics."],
222
+ messages=messages,
223
+ lora=self._lora_info(tensors, configs),
224
+ configs=configs,
225
+ histogram=self._histogram(tensors),
226
+ tensor_size_distribution=[(tensor.name, tensor.parameter_count) for tensor in largest],
227
+ checkpoints=checkpoints,
228
+ loss_history=[],
229
+ )
230
+
231
+ def _empty(
232
+ self,
233
+ target: Path,
234
+ recorded_architecture: str,
235
+ run_settings: dict[str, Any] | None,
236
+ config_files: list[Path],
237
+ configs: dict[str, Any],
238
+ message: str,
239
+ health: list[str] | None = None,
240
+ ) -> ModelInspection:
241
+ architecture, confidence, detection_message = self._architecture_from_signals(target, recorded_architecture, configs, [])
242
+ return ModelInspection(
243
+ path=str(target),
244
+ resolved_path=str(target),
245
+ architecture=architecture,
246
+ confidence=confidence,
247
+ status="warning",
248
+ size_bytes=_folder_size(target),
249
+ config_files=[str(item) for item in config_files],
250
+ resolution=_resolution_from_configs(configs, run_settings),
251
+ epoch=None,
252
+ step=step_from_name(target.name),
253
+ tensor_count=0,
254
+ total_parameters=0,
255
+ trainable_parameters=None,
256
+ parameter_memory_bytes=0,
257
+ dtypes={},
258
+ components={},
259
+ largest_tensors=[],
260
+ tensors=[],
261
+ health=health or [message],
262
+ messages=[detection_message, message],
263
+ configs=configs,
264
+ checkpoints=[str(item) for item in discover_checkpoint_paths(target)],
265
+ )
266
+
267
+ @staticmethod
268
+ def _number_from_metadata(metadata: dict[str, Any], key: str) -> int | None:
269
+ for candidate in (key, f"global_{key}", f"current_{key}"):
270
+ try:
271
+ value = metadata.get(candidate)
272
+ if value is not None:
273
+ return int(value)
274
+ except (TypeError, ValueError):
275
+ pass
276
+ return None
277
+
278
+ @staticmethod
279
+ def _trainable_parameters(tensors: list[TensorStats], architecture: str) -> int | None:
280
+ if architecture == "LoRA":
281
+ return sum(tensor.parameter_count for tensor in tensors)
282
+ return None
283
+
284
+ @staticmethod
285
+ def _architecture_from_signals(
286
+ target: Path,
287
+ recorded_architecture: str,
288
+ configs: dict[str, Any],
289
+ tensor_names: list[str],
290
+ ) -> tuple[str, float, str]:
291
+ recorded = recorded_architecture.casefold()
292
+ joined_names = "\n".join(tensor_names).casefold()
293
+ config_text = json.dumps(configs, default=str).casefold()
294
+ folder_text = str(target).casefold()
295
+ signals = " ".join((joined_names, config_text, folder_text))
296
+ if "lora" in recorded or "lora" in signals or "adapter_config" in signals:
297
+ return "LoRA", 0.92, "Model recognized as LoRA"
298
+ if "maskgit" in recorded or "maskgit" in signals:
299
+ return "MaskGIT", 0.86, "Model recognized as MaskGIT"
300
+ if "flow" in recorded or "rectified_flow" in signals or "flow_model_info" in signals:
301
+ return "Flow Matching", 0.9, "Model recognized as Flow Matching"
302
+ if "ddpm" in recorded or "diffusers" in config_text or "unet" in signals or "scheduler_config" in signals:
303
+ return "DDPM / Diffusers", 0.88, "Model recognized as DDPM"
304
+ return "Generic / Unknown", 0.35, "Model type uncertain - using generic tensor inspection"
305
+
306
+ @staticmethod
307
+ def _lora_info(tensors: list[TensorStats], configs: dict[str, Any]) -> dict[str, Any]:
308
+ lora_tensors = [tensor for tensor in tensors if "lora" in tensor.name.casefold()]
309
+ if not lora_tensors:
310
+ return {}
311
+ down = [tensor for tensor in lora_tensors if any(token in tensor.name.casefold() for token in ("down", "lora_a"))]
312
+ up = [tensor for tensor in lora_tensors if any(token in tensor.name.casefold() for token in ("up", "lora_b"))]
313
+ ranks = sorted({tensor.shape[0] for tensor in down if tensor.shape})
314
+ alpha = None
315
+ targets: set[str] = set()
316
+ for config in configs.values():
317
+ if isinstance(config, dict):
318
+ alpha = config.get("lora_alpha", config.get("alpha", alpha))
319
+ modules = config.get("target_modules")
320
+ if isinstance(modules, list):
321
+ targets.update(str(item) for item in modules)
322
+ if not targets:
323
+ for tensor in lora_tensors:
324
+ parts = tensor.name.split(".")
325
+ if len(parts) > 2:
326
+ targets.add(parts[-3])
327
+ return {
328
+ "rank": ", ".join(str(item) for item in ranks[:8]) if ranks else "unknown",
329
+ "alpha": alpha if alpha is not None else "unknown",
330
+ "target_modules": sorted(targets)[:20],
331
+ "down_matrices": len(down),
332
+ "up_matrices": len(up),
333
+ "adapter_parameter_count": sum(tensor.parameter_count for tensor in lora_tensors),
334
+ "average_abs_mean": (
335
+ sum(tensor.abs_mean or 0 for tensor in lora_tensors) / max(1, len(lora_tensors))
336
+ ),
337
+ }
338
+
339
+ @staticmethod
340
+ def _histogram(tensors: list[TensorStats]) -> dict[str, list[float]]:
341
+ values = [tensor.abs_mean for tensor in tensors if tensor.abs_mean is not None]
342
+ if not values:
343
+ return {}
344
+ buckets = [0.0] * 10
345
+ high = max(values) or 1.0
346
+ for value in values:
347
+ index = min(9, int((value / high) * 10))
348
+ buckets[index] += 1
349
+ return {"abs_mean_bins": [round(high * index / 10, 6) for index in range(11)], "counts": buckets}
adam/model_inspector/lora.py ADDED
@@ -0,0 +1,23 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from __future__ import annotations
2
+
3
+ from pathlib import Path
4
+ from typing import Any
5
+
6
+ from .generic import GenericModelInspector
7
+
8
+
9
+ class LoRAInspector(GenericModelInspector):
10
+ architecture = "LoRA"
11
+
12
+ def inspect(self, path: str | Path, *, recorded_architecture: str = "", run_settings: dict[str, Any] | None = None, progress=None, cancelled=None):
13
+ summary = super().inspect(
14
+ path,
15
+ recorded_architecture=recorded_architecture or "lora",
16
+ run_settings=run_settings,
17
+ progress=progress,
18
+ cancelled=cancelled,
19
+ )
20
+ summary.architecture = "LoRA" if summary.architecture == "Generic / Unknown" else summary.architecture
21
+ if summary.tensors and not summary.lora:
22
+ summary.health.append("Worth inspecting: ADAM marked this as LoRA, but LoRA tensor naming was not obvious")
23
+ return summary
adam/model_inspector/maskgit.py ADDED
@@ -0,0 +1,24 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from __future__ import annotations
2
+
3
+ from pathlib import Path
4
+ from typing import Any
5
+
6
+ from .generic import GenericModelInspector
7
+
8
+
9
+ class MaskGITInspector(GenericModelInspector):
10
+ architecture = "MaskGIT"
11
+
12
+ def inspect(self, path: str | Path, *, recorded_architecture: str = "", run_settings: dict[str, Any] | None = None, progress=None, cancelled=None):
13
+ summary = super().inspect(
14
+ path,
15
+ recorded_architecture=recorded_architecture or "maskgit",
16
+ run_settings=run_settings,
17
+ progress=progress,
18
+ cancelled=cancelled,
19
+ )
20
+ summary.architecture = "MaskGIT" if summary.architecture == "Generic / Unknown" else summary.architecture
21
+ names = "\n".join(tensor.name for tensor in summary.tensors).casefold()
22
+ if summary.tensors and not any(token in names for token in ("attention", "attn", "transformer", "embed")):
23
+ summary.health.append("Worth inspecting: expected MaskGIT transformer or embedding tensors were not obvious")
24
+ return summary
adam/model_inspector/statistics.py ADDED
@@ -0,0 +1,142 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from __future__ import annotations
2
+
3
+ import math
4
+ from pathlib import Path
5
+ from typing import Any
6
+
7
+ from .base import CONFIG_FILENAMES, MODEL_EXTENSIONS, TensorStats, dtype_size, parameter_count
8
+
9
+
10
+ def bytes_label(size: int | float | None) -> str:
11
+ if size is None:
12
+ return "-"
13
+ value = float(size)
14
+ for unit in ("B", "KB", "MB", "GB", "TB"):
15
+ if abs(value) < 1024 or unit == "TB":
16
+ return f"{value:.1f} {unit}" if unit != "B" else f"{int(value)} B"
17
+ value /= 1024
18
+ return f"{value:.1f} TB"
19
+
20
+
21
+ def component_for_name(name: str) -> str:
22
+ lowered = name.casefold()
23
+ mapping = (
24
+ ("down_blocks", ("down_blocks", "down.", "downsample")),
25
+ ("mid_block", ("mid_block", "middle_block", "mid.")),
26
+ ("up_blocks", ("up_blocks", "up.", "upsample")),
27
+ ("attention", ("attn", "attention", "to_q", "to_k", "to_v", "query", "key", "value")),
28
+ ("embeddings", ("embed", "embedding", "position", "token")),
29
+ ("transformer blocks", ("transformer", "blocks.", "layers.", "encoder", "decoder")),
30
+ ("output layers", ("out.", "output", "proj_out", "lm_head", "conv_out")),
31
+ ("LoRA adapters", ("lora", "hada", "lokr", "adapter")),
32
+ ("normalization", ("norm", "bn", "ln", "group_norm", "layer_norm")),
33
+ )
34
+ for component, tokens in mapping:
35
+ if any(token in lowered for token in tokens):
36
+ return component
37
+ return name.split(".", 1)[0] if "." in name else "other"
38
+
39
+
40
+ def safe_number(value: Any) -> float | None:
41
+ try:
42
+ number = float(value)
43
+ except (TypeError, ValueError, OverflowError):
44
+ return None
45
+ return number if math.isfinite(number) else None
46
+
47
+
48
+ def tensor_stats_from_torch(name: str, tensor: Any, *, sample_limit: int = 1_000_000) -> TensorStats:
49
+ shape = tuple(int(dim) for dim in getattr(tensor, "shape", ()))
50
+ dtype = str(getattr(tensor, "dtype", "unknown")).replace("torch.", "")
51
+ count = parameter_count(shape)
52
+ stat = TensorStats(
53
+ name=name,
54
+ shape=shape,
55
+ dtype=dtype,
56
+ parameter_count=count,
57
+ memory_bytes=count * dtype_size(dtype),
58
+ component=component_for_name(name),
59
+ )
60
+ if count == 0:
61
+ stat.health.append("Empty tensor")
62
+ return stat
63
+ try:
64
+ import torch
65
+
66
+ with torch.no_grad():
67
+ values = tensor.detach().to(device="cpu")
68
+ if not values.is_floating_point() and not values.is_complex():
69
+ values = values.float()
70
+ else:
71
+ values = values.float()
72
+ flat = values.reshape(-1)
73
+ if flat.numel() > sample_limit:
74
+ stride = max(1, flat.numel() // sample_limit)
75
+ flat = flat[::stride][:sample_limit]
76
+ finite = torch.isfinite(flat)
77
+ if not bool(finite.all()):
78
+ if bool(torch.isnan(flat).any()):
79
+ stat.health.append("Invalid: NaN values found")
80
+ if bool(torch.isinf(flat).any()):
81
+ stat.health.append("Invalid: Inf values found")
82
+ flat = flat[finite]
83
+ if flat.numel() == 0:
84
+ return stat
85
+ stat.minimum = safe_number(flat.min().item())
86
+ stat.maximum = safe_number(flat.max().item())
87
+ stat.mean = safe_number(flat.mean().item())
88
+ stat.std = safe_number(flat.std(unbiased=False).item()) if flat.numel() > 1 else 0.0
89
+ stat.abs_mean = safe_number(flat.abs().mean().item())
90
+ stat.l2_norm = safe_number(torch.linalg.vector_norm(flat).item())
91
+ stat.zero_percent = safe_number((flat == 0).float().mean().item() * 100)
92
+ except Exception as exc:
93
+ stat.health.append(f"Statistics unavailable: {exc}")
94
+ if stat.abs_mean is not None and stat.abs_mean > 100:
95
+ stat.health.append("Unusual: very large average weight magnitude")
96
+ if stat.maximum is not None and stat.minimum is not None and max(abs(stat.maximum), abs(stat.minimum)) > 1_000:
97
+ stat.health.append("Unusual: very large absolute weight value")
98
+ return stat
99
+
100
+
101
+ def discover_config_files(path: Path) -> list[Path]:
102
+ root = path if path.is_dir() else path.parent
103
+ files: list[Path] = []
104
+ try:
105
+ for item in root.rglob("*"):
106
+ if item.is_file() and item.name in CONFIG_FILENAMES:
107
+ files.append(item)
108
+ except OSError:
109
+ return []
110
+ return sorted(files)
111
+
112
+
113
+ def discover_checkpoint_paths(path: Path) -> list[Path]:
114
+ root = path if path.is_dir() else path.parent
115
+ candidates: list[Path] = []
116
+ try:
117
+ for item in root.rglob("*"):
118
+ if item.is_file() and item.suffix.casefold() in MODEL_EXTENSIONS:
119
+ candidates.append(item)
120
+ elif item.is_dir() and item.name.startswith("checkpoint-"):
121
+ candidates.append(item)
122
+ except OSError:
123
+ return []
124
+ return sorted(candidates, key=lambda item: (step_from_name(item.name) or -1, str(item)))
125
+
126
+
127
+ def step_from_name(name: str) -> int | None:
128
+ import re
129
+
130
+ matches = re.findall(r"(?:step|checkpoint|epoch|e|s)[-_]?(\d+)", name, flags=re.I)
131
+ if not matches:
132
+ matches = re.findall(r"(\d+)", name)
133
+ if not matches:
134
+ return None
135
+ try:
136
+ return int(matches[-1])
137
+ except ValueError:
138
+ return None
139
+
140
+
141
+ def shape_label(shape: tuple[int, ...]) -> str:
142
+ return " x ".join(str(dim) for dim in shape) if shape else "scalar"
adam/model_plugin_backend.py ADDED
@@ -0,0 +1,38 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from __future__ import annotations
2
+
3
+ from typing import Any
4
+
5
+ from adam.executor import ToolContext, ToolExecutionError
6
+ from adam.model_plugins import ModelPluginRegistry, plugin_function, validate_settings
7
+
8
+
9
+ def _plugin_for_tool(context: ToolContext, mode: str):
10
+ registry = ModelPluginRegistry(context.root)
11
+ for plugin in registry.all():
12
+ if mode == "training" and plugin.trainer_id == context.tool.id:
13
+ return plugin
14
+ if mode == "generation" and plugin.generator_id == context.tool.id:
15
+ return plugin
16
+ raise ToolExecutionError(f"No model plugin owns {context.tool.id}.")
17
+
18
+
19
+ def train(context: ToolContext, **settings: Any) -> dict[str, Any]:
20
+ plugin = _plugin_for_tool(context, "training")
21
+ errors = validate_settings(plugin.training_settings, settings)
22
+ if errors:
23
+ raise ToolExecutionError(" ".join(errors))
24
+ function = plugin_function(plugin, "train")
25
+ if function is None:
26
+ raise ToolExecutionError(f"{plugin.name} does not implement train().")
27
+ return function(settings=settings, callbacks=context) or {}
28
+
29
+
30
+ def generate(context: ToolContext, **settings: Any) -> dict[str, Any]:
31
+ plugin = _plugin_for_tool(context, "generation")
32
+ errors = validate_settings(plugin.generation_settings, settings)
33
+ if errors:
34
+ raise ToolExecutionError(" ".join(errors))
35
+ function = plugin_function(plugin, "generate")
36
+ if function is None:
37
+ raise ToolExecutionError(f"{plugin.name} does not implement generate().")
38
+ return function(settings=settings, callbacks=context) or {}
adam/model_plugins.py ADDED
@@ -0,0 +1,494 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from __future__ import annotations
2
+
3
+ import importlib
4
+ import importlib.util
5
+ import json
6
+ import logging
7
+ import pkgutil
8
+ from dataclasses import dataclass, field
9
+ from pathlib import Path
10
+ from typing import Any, Callable
11
+
12
+
13
+ REQUIRED_INFO_FIELDS = {"name", "version", "category", "description"}
14
+ SUPPORTED_SETTING_TYPES = {
15
+ "int",
16
+ "float",
17
+ "bool",
18
+ "choice",
19
+ "text",
20
+ "multiline_text",
21
+ "path",
22
+ "folder",
23
+ "slider",
24
+ }
25
+
26
+
27
+ class ModelPluginError(RuntimeError):
28
+ pass
29
+
30
+
31
+ @dataclass(frozen=True, slots=True)
32
+ class ModelPlugin:
33
+ id: str
34
+ info: dict[str, Any]
35
+ training_settings: dict[str, dict[str, Any]] = field(default_factory=dict)
36
+ generation_settings: dict[str, dict[str, Any]] = field(default_factory=dict)
37
+ training_tool: dict[str, Any] = field(default_factory=dict)
38
+ generation_tool: dict[str, Any] = field(default_factory=dict)
39
+ module_name: str = ""
40
+ plugin_path: Path | None = None
41
+
42
+ @property
43
+ def name(self) -> str:
44
+ return str(self.info.get("name", self.id))
45
+
46
+ @property
47
+ def trainer_id(self) -> str:
48
+ return str(self.training_tool.get("id") or f"{self.id}_trainer")
49
+
50
+ @property
51
+ def generator_id(self) -> str:
52
+ return str(self.generation_tool.get("id") or f"{self.id}_generator")
53
+
54
+
55
+ class ModelPluginRegistry:
56
+ """Discovers model plugins and validates their setting schemas."""
57
+
58
+ def __init__(self, root: Path, logger: logging.Logger | None = None) -> None:
59
+ self.root = root.resolve()
60
+ self.logger = logger or logging.getLogger(__name__)
61
+ self.plugins: dict[str, ModelPlugin] = {}
62
+ self.errors: list[str] = []
63
+ self.discover()
64
+
65
+ def discover(self) -> None:
66
+ self.plugins = {}
67
+ self.errors = []
68
+ for module_name in self._candidate_modules():
69
+ try:
70
+ plugin = self._load_module_plugin(module_name)
71
+ except Exception as exc:
72
+ message = f"{module_name}: {exc}"
73
+ self.errors.append(message)
74
+ self.logger.warning("Model plugin failed to load: %s", message)
75
+ continue
76
+ if plugin.id in self.plugins:
77
+ self.errors.append(f"{module_name}: duplicate model plugin id {plugin.id}")
78
+ continue
79
+ self.plugins[plugin.id] = plugin
80
+
81
+ def _candidate_modules(self) -> list[str | Path]:
82
+ modules: list[str | Path] = []
83
+ try:
84
+ package = importlib.import_module("adam.model_plugins_builtin")
85
+ for item in pkgutil.iter_modules(package.__path__, package.__name__ + "."):
86
+ if not item.ispkg:
87
+ continue
88
+ modules.append(item.name + ".manifest")
89
+ except Exception as exc:
90
+ self.errors.append(f"adam.model_plugins_builtin: {exc}")
91
+
92
+ models_dir = self.root / "models"
93
+ if models_dir.is_dir():
94
+ for folder in sorted(models_dir.iterdir()):
95
+ manifest = folder / "manifest.py"
96
+ if not folder.is_dir() or not manifest.is_file():
97
+ continue
98
+ modules.append(manifest)
99
+ return modules
100
+
101
+ def _load_module_plugin(self, module_name: str | Path) -> ModelPlugin:
102
+ if isinstance(module_name, Path):
103
+ fallback_id = module_name.parent.name
104
+ unique_name = f"adam_user_model_{fallback_id}_{abs(hash(str(module_name.resolve())))}"
105
+ spec = importlib.util.spec_from_file_location(unique_name, module_name)
106
+ if spec is None or spec.loader is None:
107
+ raise ModelPluginError(f"Could not load manifest file: {module_name}")
108
+ module = importlib.util.module_from_spec(spec)
109
+ spec.loader.exec_module(module)
110
+ module_label = str(module_name)
111
+ else:
112
+ module = importlib.import_module(module_name)
113
+ fallback_id = module_name.split(".")[-2]
114
+ module_label = module_name
115
+ plugin_id = str(getattr(module, "PLUGIN_ID", "") or fallback_id)
116
+ info = dict(getattr(module, "MODEL_INFO", {}))
117
+ missing = REQUIRED_INFO_FIELDS - set(info)
118
+ if missing:
119
+ raise ModelPluginError(
120
+ "MODEL_INFO is missing " + ", ".join(sorted(missing))
121
+ )
122
+ training_settings = self._validate_schema(
123
+ dict(getattr(module, "TRAINING_SETTINGS", {})),
124
+ f"{plugin_id} training",
125
+ )
126
+ generation_settings = self._validate_schema(
127
+ dict(getattr(module, "GENERATION_SETTINGS", {})),
128
+ f"{plugin_id} generation",
129
+ )
130
+ plugin_path = Path(getattr(module, "__file__", "")).resolve().parent
131
+ return ModelPlugin(
132
+ id=plugin_id,
133
+ info=info,
134
+ training_settings=training_settings,
135
+ generation_settings=generation_settings,
136
+ training_tool=dict(getattr(module, "TRAINING_TOOL", {})),
137
+ generation_tool=dict(getattr(module, "GENERATION_TOOL", {})),
138
+ module_name=module_label,
139
+ plugin_path=plugin_path,
140
+ )
141
+
142
+ @staticmethod
143
+ def _validate_schema(
144
+ schema: dict[str, Any],
145
+ label: str,
146
+ ) -> dict[str, dict[str, Any]]:
147
+ clean: dict[str, dict[str, Any]] = {}
148
+ for key, raw in schema.items():
149
+ if not isinstance(raw, dict):
150
+ raise ModelPluginError(f"{label} setting {key} must be an object")
151
+ spec = dict(raw)
152
+ setting_type = str(spec.get("type", "text"))
153
+ if setting_type not in SUPPORTED_SETTING_TYPES:
154
+ raise ModelPluginError(
155
+ f"{label} setting {key} has unsupported type {setting_type}"
156
+ )
157
+ spec["type"] = setting_type
158
+ spec.setdefault("label", key.replace("_", " ").title())
159
+ spec.setdefault("group", "Basic")
160
+ if setting_type == "choice":
161
+ options = spec.get("options", [])
162
+ if not isinstance(options, (list, tuple)) or not options:
163
+ raise ModelPluginError(f"{label} setting {key} needs options")
164
+ spec["options"] = list(options)
165
+ spec.setdefault("default", spec["options"][0])
166
+ clean[str(key)] = spec
167
+ return clean
168
+
169
+ def get(self, plugin_id: str) -> ModelPlugin:
170
+ return self.plugins[plugin_id]
171
+
172
+ def all(self) -> list[ModelPlugin]:
173
+ return list(self.plugins.values())
174
+
175
+ def by_trainer(self, trainer: str) -> ModelPlugin | None:
176
+ return next((plugin for plugin in self.plugins.values() if plugin.id == trainer), None)
177
+
178
+ def training_schema(self, trainer: str) -> dict[str, dict[str, Any]]:
179
+ plugin = self.by_trainer(trainer)
180
+ return plugin.training_settings if plugin else {}
181
+
182
+ def generation_schema_for_tool(self, tool_id: str) -> dict[str, dict[str, Any]]:
183
+ for plugin in self.plugins.values():
184
+ if plugin.generator_id == tool_id:
185
+ return plugin.generation_settings
186
+ return {}
187
+
188
+ def training_tool_specs(self) -> list[dict[str, Any]]:
189
+ return [
190
+ self._tool_spec(plugin, mode="training")
191
+ for plugin in self.plugins.values()
192
+ if plugin.training_tool
193
+ ]
194
+
195
+ def generation_tool_specs(self) -> list[dict[str, Any]]:
196
+ return [
197
+ self._tool_spec(plugin, mode="generation")
198
+ for plugin in self.plugins.values()
199
+ if plugin.generation_tool
200
+ ]
201
+
202
+ @staticmethod
203
+ def _tool_spec(plugin: ModelPlugin, *, mode: str) -> dict[str, Any]:
204
+ tool = dict(plugin.training_tool if mode == "training" else plugin.generation_tool)
205
+ schema = plugin.training_settings if mode == "training" else plugin.generation_settings
206
+ core_arguments = (
207
+ ["dataset_dir", "model_name", "epochs", "output_dir", "resume_from"]
208
+ if mode == "training"
209
+ else [
210
+ "model_name", "model_path", "prompt", "image_count", "steps",
211
+ "seed", "sampler", "aspect_ratio",
212
+ ]
213
+ )
214
+ core_required = (
215
+ ["dataset_dir", "model_name", "epochs", "output_dir"]
216
+ if mode == "training"
217
+ else ["model_name", "model_path", "image_count", "steps", "seed"]
218
+ )
219
+ defaults = {
220
+ "id": plugin.trainer_id if mode == "training" else plugin.generator_id,
221
+ "name": f"{plugin.name} {'Trainer' if mode == 'training' else 'Generator'}",
222
+ "description": plugin.info.get("description", ""),
223
+ "category": "Training" if mode == "training" else "Output",
224
+ "entry_function": "train" if mode == "training" else "generate",
225
+ "arguments": [*core_arguments, *list(schema)],
226
+ "required_arguments": [
227
+ *core_required,
228
+ *[key for key, spec in schema.items() if bool(spec.get("required"))],
229
+ ],
230
+ "capabilities": (
231
+ ["fresh_training", "progress", "pause", "cancel"]
232
+ if mode == "training"
233
+ else ["image_generation", "progress", "cancel"]
234
+ ),
235
+ "requires_confirmation": mode == "training",
236
+ "enabled": True,
237
+ "demo": False,
238
+ }
239
+ defaults.update(tool)
240
+ defaults["arguments"] = list(defaults.get("arguments") or [*core_arguments, *list(schema)])
241
+ defaults["required_arguments"] = list(defaults.get("required_arguments") or [])
242
+ return defaults
243
+
244
+ def validate_settings(
245
+ self,
246
+ trainer: str,
247
+ values: dict[str, Any],
248
+ *,
249
+ mode: str = "training",
250
+ ) -> list[str]:
251
+ plugin = self.by_trainer(trainer)
252
+ if not plugin:
253
+ return [f"Unknown model plugin: {trainer}"]
254
+ schema = plugin.training_settings if mode == "training" else plugin.generation_settings
255
+ return validate_settings(schema, values)
256
+
257
+
258
+ def validate_settings(schema: dict[str, dict[str, Any]], values: dict[str, Any]) -> list[str]:
259
+ errors: list[str] = []
260
+ for key, spec in schema.items():
261
+ value = values.get(key, spec.get("default"))
262
+ label = str(spec.get("label", key))
263
+ if spec.get("required") and (value is None or str(value).strip() == ""):
264
+ errors.append(f"{label} is required.")
265
+ continue
266
+ if value in (None, "") and not spec.get("required"):
267
+ continue
268
+ setting_type = str(spec.get("type", "text"))
269
+ try:
270
+ if setting_type in {"int", "slider"}:
271
+ if isinstance(value, bool):
272
+ raise ValueError
273
+ numeric = int(value)
274
+ elif setting_type == "float":
275
+ if isinstance(value, bool):
276
+ raise ValueError
277
+ numeric = float(value)
278
+ else:
279
+ numeric = None
280
+ except (TypeError, ValueError):
281
+ errors.append(f"{label} must be a number.")
282
+ continue
283
+ if numeric is not None:
284
+ if "min" in spec and numeric < float(spec["min"]):
285
+ errors.append(f"{label} must be at least {spec['min']}.")
286
+ if "max" in spec and numeric > float(spec["max"]):
287
+ errors.append(f"{label} must be at most {spec['max']}.")
288
+ if setting_type == "choice" and "options" in spec and value not in spec["options"]:
289
+ errors.append(f"{label} must be one of: {', '.join(map(str, spec['options']))}.")
290
+ if setting_type == "path" and spec.get("must_exist") and not Path(str(value)).expanduser().is_file():
291
+ errors.append(f"{label} must point to an existing file.")
292
+ if setting_type == "folder" and spec.get("must_exist") and not Path(str(value)).expanduser().is_dir():
293
+ errors.append(f"{label} must point to an existing folder.")
294
+ return errors
295
+
296
+
297
+ def load_presets(root: Path, plugin_id: str, mode: str) -> dict[str, dict[str, Any]]:
298
+ path = root.resolve() / "config" / "model_presets.json"
299
+ try:
300
+ payload = json.loads(path.read_text(encoding="utf-8"))
301
+ except (OSError, json.JSONDecodeError):
302
+ return {}
303
+ presets = payload.get(plugin_id, {}).get(mode, {})
304
+ return dict(presets) if isinstance(presets, dict) else {}
305
+
306
+
307
+ def save_preset(
308
+ root: Path,
309
+ plugin_id: str,
310
+ mode: str,
311
+ name: str,
312
+ settings: dict[str, Any],
313
+ ) -> None:
314
+ path = root.resolve() / "config" / "model_presets.json"
315
+ try:
316
+ payload = json.loads(path.read_text(encoding="utf-8"))
317
+ except (OSError, json.JSONDecodeError):
318
+ payload = {}
319
+ payload.setdefault(plugin_id, {}).setdefault(mode, {})[name] = settings
320
+ temporary = path.with_suffix(".tmp")
321
+ temporary.write_text(json.dumps(payload, indent=2, sort_keys=True), encoding="utf-8")
322
+ temporary.replace(path)
323
+
324
+
325
+ def plugin_function(plugin: ModelPlugin, function_name: str) -> Callable[..., Any] | None:
326
+ if plugin.module_name.endswith("manifest.py"):
327
+ spec = importlib.util.spec_from_file_location(
328
+ f"adam_user_model_{plugin.id}_{abs(hash(plugin.module_name))}",
329
+ plugin.module_name,
330
+ )
331
+ if spec is None or spec.loader is None:
332
+ return None
333
+ module = importlib.util.module_from_spec(spec)
334
+ spec.loader.exec_module(module)
335
+ else:
336
+ module = importlib.import_module(plugin.module_name)
337
+ function = getattr(module, function_name, None)
338
+ return function if callable(function) else None
339
+
340
+
341
+ def safe_plugin_id(name: str) -> str:
342
+ cleaned = "".join(
343
+ character.lower() if character.isalnum() else "_"
344
+ for character in name.strip()
345
+ )
346
+ cleaned = "_".join(part for part in cleaned.split("_") if part)
347
+ return cleaned[:48] or "my_model"
348
+
349
+
350
+ def scaffold_model_plugin(
351
+ root: Path,
352
+ *,
353
+ plugin_id: str,
354
+ name: str,
355
+ architecture: str = "custom",
356
+ output_type: str = "image",
357
+ include_training: bool = True,
358
+ include_generation: bool = True,
359
+ ) -> Path:
360
+ """Create a simple user-editable model plugin folder."""
361
+ plugin_id = safe_plugin_id(plugin_id)
362
+ if plugin_id in {"ddpm", "flow", "lora", "model_template"}:
363
+ raise ModelPluginError("Choose a plugin id that does not conflict with a built-in model.")
364
+ folder = root.resolve() / "models" / plugin_id
365
+ if folder.exists():
366
+ raise ModelPluginError(f"A model plugin folder already exists: {folder}")
367
+ folder.mkdir(parents=True)
368
+ (folder / "__init__.py").write_text(
369
+ f'"""ADAM model plugin: {name}."""\n',
370
+ encoding="utf-8",
371
+ )
372
+ (folder / "manifest.py").write_text(
373
+ _manifest_template(
374
+ plugin_id=plugin_id,
375
+ name=name,
376
+ architecture=architecture,
377
+ output_type=output_type,
378
+ include_training=include_training,
379
+ include_generation=include_generation,
380
+ ),
381
+ encoding="utf-8",
382
+ )
383
+ (folder / "model.py").write_text(_model_template(), encoding="utf-8")
384
+ if include_training:
385
+ (folder / "trainer.py").write_text(_trainer_template(), encoding="utf-8")
386
+ if include_generation:
387
+ (folder / "generator.py").write_text(_generator_template(), encoding="utf-8")
388
+ return folder
389
+
390
+
391
+ def _manifest_template(
392
+ *,
393
+ plugin_id: str,
394
+ name: str,
395
+ architecture: str,
396
+ output_type: str,
397
+ include_training: bool,
398
+ include_generation: bool,
399
+ ) -> str:
400
+ plugin_id_json = json.dumps(plugin_id)
401
+ name_json = json.dumps(name)
402
+ architecture_json = json.dumps(architecture)
403
+ output_type_json = json.dumps(output_type)
404
+ training_tool = (
405
+ "{\n"
406
+ f' "id": "{plugin_id}_trainer",\n'
407
+ f' "name": {json.dumps(name + " Trainer")},\n'
408
+ f' "backend": {{"type": "python", "module": "models.{plugin_id}.trainer", "function": "train"}},\n'
409
+ "}"
410
+ if include_training else "{}"
411
+ )
412
+ generation_tool = (
413
+ "{\n"
414
+ f' "id": "{plugin_id}_generator",\n'
415
+ f' "name": {json.dumps(name + " Generator")},\n'
416
+ f' "model_trainers": ["{plugin_id}"],\n'
417
+ f' "backend": {{"type": "python", "module": "models.{plugin_id}.generator", "function": "generate"}},\n'
418
+ "}"
419
+ if include_generation else "{}"
420
+ )
421
+ return f'''PLUGIN_ID = {plugin_id_json}
422
+
423
+ MODEL_INFO = {{
424
+ "name": {name_json},
425
+ "version": "0.1",
426
+ "category": "Image Generation",
427
+ "description": {json.dumps("Describe what " + name + " trains or generates.")},
428
+ "architecture": {architecture_json},
429
+ "status": "experimental",
430
+ "output_type": {output_type_json},
431
+ }}
432
+
433
+ TRAINING_SETTINGS = {{
434
+ "resolution": {{"label": "Resolution", "type": "choice", "options": [64, 128, 256, 384, 512], "default": 256, "group": "Basic"}},
435
+ "batch_size": {{"label": "Batch size", "type": "int", "default": 1, "min": 1, "max": 64, "group": "Basic"}},
436
+ "learning_rate": {{"label": "Learning rate", "type": "float", "default": 0.0001, "min": 0.0000001, "max": 0.1, "decimals": 7, "group": "Optimization"}},
437
+ "mixed_precision": {{"label": "Precision", "type": "choice", "options": ["fp16", "no"], "default": "fp16", "group": "Optimization"}},
438
+ "preview_enabled": {{"label": "Generate previews while training", "type": "bool", "default": True, "group": "Preview"}},
439
+ "preview_every": {{"label": "Preview interval", "type": "int", "default": 5, "min": 1, "max": 100000, "group": "Preview"}},
440
+ "preview_prompt": {{"label": "Preview prompt", "type": "text", "default": "", "group": "Preview"}},
441
+ "preview_seed": {{"label": "Preview seed", "type": "int", "default": 123456789, "min": 0, "max": 2147483647, "group": "Preview"}},
442
+ }}
443
+
444
+ GENERATION_SETTINGS = {{
445
+ "prompt": {{"label": "Prompt", "type": "multiline_text", "default": "", "group": "Prompt"}},
446
+ "image_count": {{"label": "Images", "type": "int", "default": 1, "min": 1, "max": 48, "group": "Generation"}},
447
+ "steps": {{"label": "Steps", "type": "int", "default": 30, "min": 1, "max": 500, "group": "Generation"}},
448
+ "seed": {{"label": "Seed", "type": "int", "default": 0, "min": 0, "max": 2147483647, "group": "Generation"}},
449
+ }}
450
+
451
+ TRAINING_TOOL = {training_tool}
452
+
453
+ GENERATION_TOOL = {generation_tool}
454
+ '''
455
+
456
+
457
+ def _model_template() -> str:
458
+ return '''from __future__ import annotations
459
+
460
+ from typing import Any
461
+
462
+
463
+ def load_model(model_path: str, settings: dict[str, Any] | None = None) -> Any:
464
+ """Load your model or inference pipeline here."""
465
+ raise NotImplementedError("Add your model loading code.")
466
+ '''
467
+
468
+
469
+ def _trainer_template() -> str:
470
+ return '''from __future__ import annotations
471
+
472
+ from typing import Any
473
+
474
+
475
+ def train(context, **settings: Any) -> dict[str, Any]:
476
+ """Train the model and report progress back to ADAM."""
477
+ context.log("Replace this with real training code.")
478
+ context.progress(100, "Training placeholder complete")
479
+ return {}
480
+ '''
481
+
482
+
483
+ def _generator_template() -> str:
484
+ return '''from __future__ import annotations
485
+
486
+ from typing import Any
487
+
488
+
489
+ def generate(context, **settings: Any) -> dict[str, Any]:
490
+ """Generate outputs and report progress back to ADAM."""
491
+ context.log("Replace this with real generation code.")
492
+ context.progress(100, "Generation placeholder complete")
493
+ return {}
494
+ '''
adam/model_plugins_builtin/__init__.py ADDED
@@ -0,0 +1 @@
 
 
1
+ """Built-in model plugin manifests shipped with ADAM."""
adam/model_plugins_builtin/ddpm/__init__.py ADDED
@@ -0,0 +1 @@
 
 
1
+ """DDPM model plugin."""
adam/model_plugins_builtin/ddpm/manifest.py ADDED
@@ -0,0 +1,56 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ PLUGIN_ID = "ddpm"
2
+
3
+ MODEL_INFO = {
4
+ "name": "DDPM",
5
+ "version": "1.0",
6
+ "category": "Image Generation",
7
+ "description": "Denoising Diffusion Probabilistic Model image trainer and generator.",
8
+ "architecture": "diffusion",
9
+ "status": "stable",
10
+ "output_type": "image",
11
+ "capabilities": ["fresh_training", "resume_training", "image_generation", "smart_generation", "live_preview"],
12
+ "input_formats": ["image folder"],
13
+ "output_formats": ["diffusers pipeline", "checkpoint folder", "png preview"],
14
+ "hardware": {"recommended_vram_gb": 6, "recommended_system_ram_gb": 16},
15
+ "vram_behavior": {"scales_with": ["resolution", "batch_size"], "estimate": "Moderate; batch size should drop quickly above 256px."},
16
+ }
17
+
18
+ TRAINING_SETTINGS = {
19
+ "resolution": {"label": "Resolution", "type": "choice", "options": [64, 128, 256, 384, 512], "default": 128, "group": "Basic"},
20
+ "batch_size": {"label": "Batch size", "type": "int", "default": 1, "min": 1, "max": 64, "group": "Basic"},
21
+ "learning_rate": {"label": "Learning rate", "type": "float", "default": 0.0001, "min": 0.0000001, "max": 0.1, "decimals": 7, "step": 0.00005, "group": "Optimization"},
22
+ "gradient_accumulation_steps": {"label": "Gradient accumulation", "type": "int", "default": 1, "min": 1, "max": 64, "group": "Optimization"},
23
+ "dataloader_num_workers": {"label": "Loader workers", "type": "int", "default": 4, "min": 0, "max": 16, "group": "Dataset"},
24
+ "mixed_precision": {"label": "Precision", "type": "choice", "options": ["fp16", "no"], "default": "fp16", "group": "Optimization"},
25
+ "save_every": {"label": "Save every", "type": "int", "default": 10, "min": 1, "max": 1000, "group": "Checkpoints"},
26
+ "preview_steps": {"label": "Preview steps", "type": "int", "default": 50, "min": 1, "max": 500, "group": "Preview"},
27
+ "training_intensity": {"label": "Training intensity", "type": "slider", "default": 100, "min": 10, "max": 100, "group": "Advanced", "advanced": True},
28
+ "completed_epochs": {"label": "Completed epochs", "type": "int", "default": 0, "min": 0, "max": 100000, "group": "Internal", "advanced": True},
29
+ "preview_enabled": {"label": "Generate previews while training", "type": "bool", "default": True, "group": "Preview"},
30
+ "preview_every": {"label": "Preview interval", "type": "int", "default": 5, "min": 1, "max": 100000, "group": "Preview"},
31
+ "preview_prompt": {"label": "Preview prompt", "type": "text", "default": "", "group": "Preview"},
32
+ "preview_seed": {"label": "Preview seed", "type": "int", "default": 123456789, "min": 0, "max": 2147483647, "group": "Preview"},
33
+ }
34
+
35
+ GENERATION_SETTINGS = {
36
+ "prompt": {"label": "Creative note", "type": "multiline_text", "default": "", "group": "Generation"},
37
+ "image_count": {"label": "Images", "type": "int", "default": 1, "min": 1, "max": 48, "group": "Generation"},
38
+ "steps": {"label": "Sampling steps", "type": "int", "default": 50, "min": 5, "max": 500, "group": "Generation"},
39
+ "sampler": {"label": "Sampler", "type": "choice", "options": ["DDIM", "DDPM"], "default": "DDIM", "group": "Generation"},
40
+ "aspect_ratio": {"label": "Aspect ratio", "type": "choice", "options": ["1:1 (Square)", "16:9 (Widescreen)", "9:16 (Portrait)", "4:3 (Classic)", "3:4 (Portrait Classic)", "3:2 (Photo)", "2:3 (Portrait Photo)"], "default": "1:1 (Square)", "group": "Generation"},
41
+ "seed": {"label": "Seed", "type": "int", "default": 0, "min": 0, "max": 2147483647, "group": "Generation"},
42
+ "reference_image": {"label": "Reference image", "type": "path", "default": "", "group": "Reference"},
43
+ "reference_strength": {"label": "Reference strength", "type": "slider", "default": 65, "min": 0, "max": 100, "group": "Reference"},
44
+ "width": {"label": "Custom width", "type": "int", "default": 0, "min": 0, "max": 2048, "group": "Advanced", "advanced": True},
45
+ "height": {"label": "Custom height", "type": "int", "default": 0, "min": 0, "max": 2048, "group": "Advanced", "advanced": True},
46
+ "preview_interval": {"label": "Steps per preview", "type": "int", "default": 0, "min": 0, "max": 500, "group": "Preview"},
47
+ "smart_generation": {"label": "Smart Generation", "type": "bool", "default": False, "group": "Smart Generation"},
48
+ "smart_wanted_results": {"label": "Wanted results", "type": "int", "default": 8, "min": 1, "max": 48, "group": "Smart Generation"},
49
+ "smart_max_candidates": {"label": "Maximum candidates", "type": "int", "default": 32, "min": 1, "max": 256, "group": "Smart Generation"},
50
+ "smart_min_score": {"label": "Minimum score", "type": "float", "default": 0.7, "min": 0, "max": 1, "group": "Smart Generation"},
51
+ "smart_mode": {"label": "Selection mode", "type": "choice", "options": ["threshold", "top_n"], "default": "threshold", "group": "Smart Generation"},
52
+ "smart_keep_rejected": {"label": "Keep rejected candidates", "type": "bool", "default": True, "group": "Smart Generation"},
53
+ }
54
+
55
+ TRAINING_TOOL = {"id": "ddpm_trainer"}
56
+ GENERATION_TOOL = {"id": "ddpm_generator"}
adam/model_plugins_builtin/flow_matching/__init__.py ADDED
@@ -0,0 +1 @@
 
 
1
+ """Flow Matching model plugin."""
adam/model_plugins_builtin/flow_matching/manifest.py ADDED
@@ -0,0 +1,51 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ PLUGIN_ID = "flow"
2
+
3
+ MODEL_INFO = {
4
+ "name": "Flow Matching",
5
+ "version": "1.0",
6
+ "category": "Image Generation",
7
+ "description": "Rectified Flow image trainer and generator.",
8
+ "architecture": "rectified_flow",
9
+ "status": "stable",
10
+ "output_type": "image",
11
+ "capabilities": ["fresh_training", "resume_training", "image_generation", "smart_generation", "live_preview"],
12
+ "input_formats": ["image folder"],
13
+ "output_formats": ["diffusers unet folder", "flow metadata", "png preview"],
14
+ "hardware": {"recommended_vram_gb": 8, "recommended_system_ram_gb": 16},
15
+ "vram_behavior": {"scales_with": ["resolution", "batch_size"], "estimate": "Moderate to high; flow runs usually want smaller batches at 512px."},
16
+ }
17
+
18
+ TRAINING_SETTINGS = {
19
+ "resolution": {"label": "Resolution", "type": "choice", "options": [64, 128, 256, 384, 512], "default": 256, "group": "Basic"},
20
+ "batch_size": {"label": "Batch size", "type": "int", "default": 8, "min": 1, "max": 64, "group": "Basic"},
21
+ "learning_rate": {"label": "Learning rate", "type": "float", "default": 0.0002, "min": 0.0000001, "max": 0.1, "decimals": 7, "step": 0.00005, "group": "Optimization"},
22
+ "gradient_accumulation": {"label": "Gradient accumulation", "type": "int", "default": 1, "min": 1, "max": 64, "group": "Optimization"},
23
+ "workers": {"label": "Loader workers", "type": "int", "default": 4, "min": 0, "max": 16, "group": "Dataset"},
24
+ "mixed_precision": {"label": "Precision", "type": "choice", "options": ["fp16", "no"], "default": "fp16", "group": "Optimization"},
25
+ "save_every": {"label": "Save every", "type": "int", "default": 10, "min": 1, "max": 1000, "group": "Checkpoints"},
26
+ "preview_steps": {"label": "Preview steps", "type": "int", "default": 10, "min": 1, "max": 500, "group": "Preview"},
27
+ "gradient_checkpointing": {"label": "Gradient checkpointing", "type": "bool", "default": False, "group": "Advanced", "advanced": True},
28
+ "preview_enabled": {"label": "Generate previews while training", "type": "bool", "default": True, "group": "Preview"},
29
+ "preview_every": {"label": "Preview interval", "type": "int", "default": 5, "min": 1, "max": 100000, "group": "Preview"},
30
+ "preview_prompt": {"label": "Preview prompt", "type": "text", "default": "", "group": "Preview"},
31
+ "preview_seed": {"label": "Preview seed", "type": "int", "default": 123456789, "min": 0, "max": 2147483647, "group": "Preview"},
32
+ }
33
+
34
+ GENERATION_SETTINGS = {
35
+ "prompt": {"label": "Creative note", "type": "multiline_text", "default": "", "group": "Generation"},
36
+ "image_count": {"label": "Images", "type": "int", "default": 1, "min": 1, "max": 48, "group": "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
+ "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"},
43
+ "smart_wanted_results": {"label": "Wanted results", "type": "int", "default": 8, "min": 1, "max": 48, "group": "Smart Generation"},
44
+ "smart_max_candidates": {"label": "Maximum candidates", "type": "int", "default": 32, "min": 1, "max": 256, "group": "Smart Generation"},
45
+ "smart_min_score": {"label": "Minimum score", "type": "float", "default": 0.7, "min": 0, "max": 1, "group": "Smart Generation"},
46
+ "smart_mode": {"label": "Selection mode", "type": "choice", "options": ["threshold", "top_n"], "default": "threshold", "group": "Smart Generation"},
47
+ "smart_keep_rejected": {"label": "Keep rejected candidates", "type": "bool", "default": True, "group": "Smart Generation"},
48
+ }
49
+
50
+ TRAINING_TOOL = {"id": "flow_trainer"}
51
+ GENERATION_TOOL = {"id": "flow_generator"}
adam/model_plugins_builtin/model_template/__init__.py ADDED
@@ -0,0 +1 @@
 
 
1
+ """Copyable model plugin template."""
adam/model_plugins_builtin/model_template/manifest.py ADDED
@@ -0,0 +1,30 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ PLUGIN_ID = "model_template"
2
+
3
+ MODEL_INFO = {
4
+ "name": "Model Template",
5
+ "version": "0.1",
6
+ "category": "Template",
7
+ "description": "Example manifest for adding a new ADAM model plugin.",
8
+ "status": "example",
9
+ }
10
+
11
+ TRAINING_SETTINGS = {
12
+ "dataset_dir": {"label": "Dataset folder", "type": "folder", "required": True, "group": "Dataset"},
13
+ "epochs": {"label": "Epochs", "type": "int", "default": 10, "min": 1, "max": 100000, "group": "Basic"},
14
+ "learning_rate": {"label": "Learning rate", "type": "float", "default": 0.0001, "min": 0.0000001, "max": 0.1, "group": "Optimization"},
15
+ }
16
+
17
+ GENERATION_SETTINGS = {
18
+ "model_path": {"label": "Model file or folder", "type": "path", "required": True, "group": "Model Loading"},
19
+ "seed": {"label": "Seed", "type": "int", "default": 0, "min": 0, "max": 2147483647, "group": "Generation"},
20
+ }
21
+
22
+ # A real plugin can point to its own backend:
23
+ # TRAINING_TOOL = {
24
+ # "backend": {"type": "python", "module": "models.my_model.trainer", "function": "train"},
25
+ # }
26
+ # GENERATION_TOOL = {
27
+ # "backend": {"type": "python", "module": "models.my_model.generator", "function": "generate"},
28
+ # }
29
+ TRAINING_TOOL = {}
30
+ GENERATION_TOOL = {}
adam/model_plugins_builtin/oasis/__init__.py ADDED
@@ -0,0 +1 @@
 
 
1
+ """Oasis action-conditioned world model plugin."""
adam/model_plugins_builtin/oasis/manifest.py ADDED
@@ -0,0 +1,70 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ PLUGIN_ID = "oasis"
2
+
3
+ MODEL_INFO = {
4
+ "name": "Oasis Action World Model",
5
+ "version": "1.0",
6
+ "category": "Playable World Models",
7
+ "description": "Action-conditioned playable world model trainer for gameplay frame sequences.",
8
+ "architecture": "action_conditioned_rectified_flow_video",
9
+ "status": "experimental",
10
+ "output_type": "playable_world",
11
+ "capabilities": ["fresh_training", "resume_training", "playable_inference", "live_preview"],
12
+ "input_formats": ["Oasis action dataset folder", "semicolon-separated Oasis dataset folders"],
13
+ "output_formats": ["action_flow_model_info.json", "diffusers unet folder", "png preview"],
14
+ "hardware": {"recommended_vram_gb": 12, "recommended_system_ram_gb": 32},
15
+ "vram_behavior": {
16
+ "scales_with": ["resolution", "batch_size", "sequence_context"],
17
+ "estimate": "High; 256x144 with batch 2 is the conservative RTX 3060 starting point.",
18
+ },
19
+ }
20
+
21
+ TRAINING_SETTINGS = {
22
+ "resolution": {"label": "Resolution", "type": "choice", "options": ["128x72", "256x144", "384x216", "512x288"], "default": "256x144", "group": "Basic"},
23
+ "batch_size": {"label": "Batch size", "type": "int", "default": 2, "min": 1, "max": 16, "group": "Basic"},
24
+ "learning_rate": {"label": "Learning rate", "type": "float", "default": 0.00002, "min": 0.0000001, "max": 0.01, "decimals": 7, "step": 0.00001, "group": "Optimization"},
25
+ "workers": {"label": "Loader workers", "type": "int", "default": 2, "min": 0, "max": 8, "group": "Dataset"},
26
+ "mixed_precision": {"label": "Precision", "type": "choice", "options": ["fp32", "fp16", "no"], "default": "fp32", "group": "Optimization"},
27
+ "gradient_accumulation": {"label": "Gradient accumulation", "type": "int", "default": 1, "min": 1, "max": 16, "group": "Optimization"},
28
+ "frame_gap": {"label": "Prediction horizon", "type": "int", "default": 3, "min": 1, "max": 60, "group": "Sequence"},
29
+ "sequence_context": {"label": "Context length", "type": "int", "default": 1, "min": 1, "max": 32, "group": "Sequence"},
30
+ "action_aggregation": {"label": "Action aggregation", "type": "choice", "options": ["window", "mean", "last"], "default": "window", "group": "Sequence"},
31
+ "validation_split": {"label": "Validation split", "type": "float", "default": 0.1, "min": 0.01, "max": 0.5, "decimals": 3, "step": 0.01, "group": "Dataset"},
32
+ "validation_batches": {"label": "Validation batches", "type": "int", "default": 8, "min": 0, "max": 128, "group": "Dataset"},
33
+ "save_every": {"label": "Save every", "type": "int", "default": 5, "min": 1, "max": 1000, "group": "Checkpoints"},
34
+ "preview_enabled": {"label": "Generate previews while training", "type": "bool", "default": True, "group": "Preview"},
35
+ "preview_every": {"label": "Preview interval", "type": "int", "default": 5, "min": 1, "max": 100000, "group": "Preview"},
36
+ "preview_steps": {"label": "Preview steps", "type": "int", "default": 1, "min": 1, "max": 50, "group": "Preview"},
37
+ "seed": {"label": "Random seed", "type": "int", "default": 1234, "min": 0, "max": 2147483647, "group": "Reproducibility"},
38
+ "base_model": {"label": "Base video model", "type": "folder", "default": "", "group": "Checkpoints", "advanced": True},
39
+ "condition_noise": {"label": "Condition noise", "type": "float", "default": 0.03, "min": 0.0, "max": 0.5, "decimals": 4, "step": 0.01, "group": "Advanced", "advanced": True},
40
+ "temporal_loss_weight": {"label": "Temporal loss weight", "type": "float", "default": 0.1, "min": 0.0, "max": 10.0, "decimals": 3, "step": 0.05, "group": "Advanced", "advanced": True},
41
+ "motion_loss_weight": {"label": "Motion loss weight", "type": "float", "default": 2.0, "min": 0.0, "max": 10.0, "decimals": 3, "step": 0.25, "group": "Advanced", "advanced": True},
42
+ "action_input_scale": {"label": "Action input scale", "type": "float", "default": 8.0, "min": 0.1, "max": 32.0, "decimals": 3, "step": 0.5, "group": "Advanced", "advanced": True},
43
+ "neutral_action_dropout": {"label": "Neutral action dropout", "type": "float", "default": 0.15, "min": 0.0, "max": 0.9, "decimals": 3, "step": 0.05, "group": "Advanced", "advanced": True},
44
+ "action_contrast_weight": {"label": "Action contrast weight", "type": "float", "default": 0.35, "min": 0.0, "max": 5.0, "decimals": 3, "step": 0.05, "group": "Advanced", "advanced": True},
45
+ "action_contrast_margin": {"label": "Action contrast margin", "type": "float", "default": 0.02, "min": 0.0, "max": 1.0, "decimals": 4, "step": 0.005, "group": "Advanced", "advanced": True},
46
+ "gradient_checkpointing": {"label": "Gradient checkpointing", "type": "bool", "default": False, "group": "Advanced", "advanced": True},
47
+ "balance_actions": {"label": "Balance rare actions", "type": "bool", "default": False, "group": "Advanced", "advanced": True},
48
+ }
49
+
50
+ GENERATION_SETTINGS = {
51
+ "starting_frame": {"label": "Starting frame", "type": "path", "default": "", "group": "Player"},
52
+ "seed": {"label": "Seed", "type": "int", "default": 0, "min": 0, "max": 2147483647, "group": "Player"},
53
+ }
54
+
55
+ TRAINING_TOOL = {
56
+ "id": "oasis_trainer",
57
+ "name": "Oasis Action World Model Trainer",
58
+ "backend": {"type": "python", "module": "adam.tools.oasis_adapter", "function": "train_oasis"},
59
+ "capabilities": ["fresh_training", "resume_training", "progress", "pause", "cancel", "live_preview"],
60
+ }
61
+
62
+ GENERATION_TOOL = {
63
+ "id": "oasis_player",
64
+ "name": "Oasis Playable Inference",
65
+ "model_trainers": ["oasis"],
66
+ "arguments": ["model_name", "model_path", "starting_frame", "seed"],
67
+ "required_arguments": ["model_path"],
68
+ "capabilities": ["playable_inference", "keyboard_actions", "progress", "cancel"],
69
+ "backend": {"type": "python", "module": "adam.tools.oasis_adapter", "function": "launch_oasis_player"},
70
+ }
adam/model_plugins_builtin/sdxl_lora/__init__.py ADDED
@@ -0,0 +1 @@
 
 
1
+ """SDXL LoRA model plugin."""
adam/model_plugins_builtin/sdxl_lora/manifest.py ADDED
@@ -0,0 +1,60 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ PLUGIN_ID = "lora"
2
+
3
+ MODEL_INFO = {
4
+ "name": "SDXL LoRA",
5
+ "version": "1.0",
6
+ "category": "Image Generation",
7
+ "description": "Stable Diffusion XL LoRA adapter training and base-model plus adapter generation.",
8
+ "architecture": "sdxl_lora",
9
+ "status": "stable",
10
+ "output_type": "image",
11
+ "dependencies": ["diffusers", "safetensors"],
12
+ "capabilities": ["fresh_training", "resume_training", "lora_adapter", "image_generation", "reference_image"],
13
+ "input_formats": ["captioned image folder", "SDXL checkpoint"],
14
+ "output_formats": ["safetensors", "png generation"],
15
+ "hardware": {"recommended_vram_gb": 8, "recommended_system_ram_gb": 16},
16
+ "vram_behavior": {"scales_with": ["base_model_size", "resolution", "batch_size"], "estimate": "High; SDXL LoRA usually starts safely at batch 1 on 8-12 GB GPUs."},
17
+ }
18
+
19
+ TRAINING_SETTINGS = {
20
+ "base_model": {"label": "Base model", "type": "path", "default": "", "required": True, "must_exist": True, "group": "Basic"},
21
+ "trigger_word": {"label": "Trigger word", "type": "text", "default": "", "group": "LoRA"},
22
+ "resolution": {"label": "Resolution", "type": "choice", "options": [512, 768, 1024], "default": 1024, "group": "Basic"},
23
+ "rank": {"label": "Rank", "type": "int", "default": 16, "min": 1, "max": 256, "group": "LoRA"},
24
+ "alpha": {"label": "Alpha", "type": "int", "default": 16, "min": 1, "max": 256, "group": "LoRA"},
25
+ "learning_rate": {"label": "Learning rate", "type": "float", "default": 0.0001, "min": 0.0000001, "max": 0.01, "decimals": 7, "step": 0.00005, "group": "Optimization"},
26
+ "batch_size": {"label": "Batch size", "type": "int", "default": 1, "min": 1, "max": 16, "group": "Basic"},
27
+ "gradient_accumulation_steps": {"label": "Gradient accumulation", "type": "int", "default": 1, "min": 1, "max": 64, "group": "Optimization"},
28
+ "mixed_precision": {"label": "Precision", "type": "choice", "options": ["fp16", "bf16", "no"], "default": "fp16", "group": "Optimization"},
29
+ "caption_extension": {"label": "Caption extension", "type": "choice", "options": [".txt", ".caption"], "default": ".txt", "group": "Dataset"},
30
+ "save_every": {"label": "Save every", "type": "int", "default": 10, "min": 1, "max": 1000, "group": "Checkpoints"},
31
+ "optimizer": {"label": "Optimizer", "type": "choice", "options": ["AdamW", "AdamW8bit"], "default": "AdamW", "group": "Optimization", "advanced": True},
32
+ "gradient_clip_norm": {"label": "Gradient clipping", "type": "float", "default": 1.0, "min": 0.0, "max": 10.0, "decimals": 3, "step": 0.1, "group": "Advanced", "advanced": True},
33
+ "preview_enabled": {"label": "Generate previews while training", "type": "bool", "default": True, "group": "Preview"},
34
+ "preview_every": {"label": "Preview interval", "type": "int", "default": 5, "min": 1, "max": 100000, "group": "Preview"},
35
+ "preview_prompt": {"label": "Preview prompt", "type": "text", "default": "", "group": "Preview"},
36
+ "preview_seed": {"label": "Preview seed", "type": "int", "default": 123456789, "min": 0, "max": 2147483647, "group": "Preview"},
37
+ }
38
+
39
+ GENERATION_SETTINGS = {
40
+ "base_model_path": {"label": "Base model", "type": "path", "default": "", "required": True, "must_exist": True, "group": "Model Loading"},
41
+ "model_path": {"label": "LoRA adapter", "type": "path", "default": "", "required": True, "must_exist": True, "group": "Model Loading"},
42
+ "prompt": {"label": "Prompt", "type": "multiline_text", "default": "", "required": True, "group": "Prompt"},
43
+ "negative_prompt": {"label": "Negative prompt", "type": "multiline_text", "default": "", "group": "Prompt"},
44
+ "lora_strength": {"label": "LoRA strength", "type": "float", "default": 1.0, "min": 0.0, "max": 2.0, "decimals": 2, "step": 0.05, "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++ 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"},
52
+ "seed": {"label": "Seed", "type": "int", "default": 0, "min": 0, "max": 2147483647, "group": "Generation"},
53
+ "reference_image": {"label": "Reference image", "type": "path", "default": "", "group": "Reference"},
54
+ "denoise_strength": {"label": "Denoise strength", "type": "float", "default": 0.45, "min": 0.0, "max": 1.0, "decimals": 2, "step": 0.05, "group": "Reference"},
55
+ "prompt_weighting": {"label": "Use prompt weights", "type": "bool", "default": True, "group": "Advanced", "advanced": True},
56
+ "preview_interval": {"label": "Steps per preview", "type": "int", "default": 0, "min": 0, "max": 500, "group": "Preview"},
57
+ }
58
+
59
+ TRAINING_TOOL = {"id": "lora_trainer"}
60
+ GENERATION_TOOL = {"id": "lora_generator"}
adam/model_profiles.py ADDED
@@ -0,0 +1,93 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from __future__ import annotations
2
+
3
+ from dataclasses import asdict, dataclass, field
4
+ from typing import Any
5
+
6
+ from adam.model_plugins import ModelPlugin, ModelPluginRegistry
7
+
8
+
9
+ @dataclass(frozen=True, slots=True)
10
+ class ModelProfile:
11
+ """Normalized model profile built from ADAM's plugin manifests."""
12
+
13
+ id: str
14
+ name: str
15
+ category: str
16
+ architecture: str
17
+ version: str
18
+ description: str
19
+ status: str = "experimental"
20
+ output_type: str = "image"
21
+ training: dict[str, dict[str, Any]] = field(default_factory=dict)
22
+ generation: dict[str, dict[str, Any]] = field(default_factory=dict)
23
+ trainer_module: str = ""
24
+ generator_module: str = ""
25
+ trainer_tool: str = ""
26
+ generator_tool: str = ""
27
+ capabilities: list[str] = field(default_factory=list)
28
+ hardware: dict[str, Any] = field(default_factory=dict)
29
+ vram_behavior: dict[str, Any] = field(default_factory=dict)
30
+ input_formats: list[str] = field(default_factory=list)
31
+ output_formats: list[str] = field(default_factory=list)
32
+
33
+ def to_dict(self) -> dict[str, Any]:
34
+ return asdict(self)
35
+
36
+
37
+ def profile_from_plugin(plugin: ModelPlugin) -> ModelProfile:
38
+ info = dict(plugin.info)
39
+ training_tool = dict(plugin.training_tool)
40
+ generation_tool = dict(plugin.generation_tool)
41
+ capabilities = list(
42
+ dict.fromkeys(
43
+ [
44
+ *info.get("capabilities", []),
45
+ *training_tool.get("capabilities", []),
46
+ *generation_tool.get("capabilities", []),
47
+ ]
48
+ )
49
+ )
50
+ trainer_backend = dict(training_tool.get("backend", {}))
51
+ generator_backend = dict(generation_tool.get("backend", {}))
52
+ return ModelProfile(
53
+ id=plugin.id,
54
+ name=str(info.get("name", plugin.name)),
55
+ category=str(info.get("category", "")),
56
+ architecture=str(info.get("architecture", plugin.id)),
57
+ version=str(info.get("version", "")),
58
+ description=str(info.get("description", "")),
59
+ status=str(info.get("status", "experimental")),
60
+ output_type=str(info.get("output_type", "image")),
61
+ training=plugin.training_settings,
62
+ generation=plugin.generation_settings,
63
+ trainer_module=str(trainer_backend.get("module", "")),
64
+ generator_module=str(generator_backend.get("module", "")),
65
+ trainer_tool=plugin.trainer_id if plugin.training_settings else "",
66
+ generator_tool=plugin.generator_id if plugin.generation_settings else "",
67
+ capabilities=capabilities,
68
+ hardware=dict(info.get("hardware", {})),
69
+ vram_behavior=dict(info.get("vram_behavior", {})),
70
+ input_formats=list(info.get("input_formats", [])),
71
+ output_formats=list(info.get("output_formats", [])),
72
+ )
73
+
74
+
75
+ class ModelProfileRegistry:
76
+ """Read-only view over plugin manifests for UI and automation features."""
77
+
78
+ def __init__(self, plugins: ModelPluginRegistry) -> None:
79
+ self.plugins = plugins
80
+
81
+ def all(self) -> list[ModelProfile]:
82
+ return [
83
+ profile_from_plugin(plugin)
84
+ for plugin in self.plugins.all()
85
+ if plugin.info.get("category") != "Template"
86
+ ]
87
+
88
+ def get(self, profile_id: str) -> ModelProfile | None:
89
+ plugin = self.plugins.by_trainer(profile_id)
90
+ return profile_from_plugin(plugin) if plugin else None
91
+
92
+ def as_catalog(self) -> list[dict[str, Any]]:
93
+ return [profile.to_dict() for profile in self.all()]
adam/models.py CHANGED
@@ -14,6 +14,7 @@ def utc_now() -> str:
14
  class JobStatus(str, Enum):
15
  DRAFT = "Draft"
16
  AWAITING_CONFIRMATION = "Awaiting confirmation"
 
17
  QUEUED = "Queued"
18
  RUNNING = "Running"
19
  PAUSED = "Paused"
@@ -73,6 +74,7 @@ class Job:
73
  progress: int = 0
74
  current_step: int = -1
75
  created_at: str = field(default_factory=utc_now)
 
76
  started_at: str | None = None
77
  ended_at: str | None = None
78
  output_folder: str | None = None
@@ -89,8 +91,15 @@ class Job:
89
  preview_total: int = 0
90
  preview_image_index: int = 0
91
  preview_image_count: int = 0
 
 
 
 
 
 
92
  atlas_report: dict[str, Any] = field(default_factory=dict)
93
  nova_report: dict[str, Any] = field(default_factory=dict)
 
94
 
95
  def to_dict(self) -> dict[str, Any]:
96
  payload = asdict(self)
 
14
  class JobStatus(str, Enum):
15
  DRAFT = "Draft"
16
  AWAITING_CONFIRMATION = "Awaiting confirmation"
17
+ SCHEDULED = "Scheduled"
18
  QUEUED = "Queued"
19
  RUNNING = "Running"
20
  PAUSED = "Paused"
 
74
  progress: int = 0
75
  current_step: int = -1
76
  created_at: str = field(default_factory=utc_now)
77
+ scheduled_for: str | None = None
78
  started_at: str | None = None
79
  ended_at: str | None = None
80
  output_folder: str | None = None
 
91
  preview_total: int = 0
92
  preview_image_index: int = 0
93
  preview_image_count: int = 0
94
+ eta_seconds: int | None = None
95
+ estimated_completion_at: str | None = None
96
+ progress_current: int = 0
97
+ progress_total: int = 0
98
+ progress_rate: float = 0.0
99
+ progress_unit: str = "step"
100
  atlas_report: dict[str, Any] = field(default_factory=dict)
101
  nova_report: dict[str, Any] = field(default_factory=dict)
102
+ metadata: dict[str, Any] = field(default_factory=dict)
103
 
104
  def to_dict(self) -> dict[str, Any]:
105
  payload = asdict(self)
adam/oasis_dataset.py ADDED
@@ -0,0 +1,215 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from __future__ import annotations
2
+
3
+ import json
4
+ import re
5
+ from dataclasses import dataclass, field
6
+ from pathlib import Path
7
+ from typing import Any
8
+
9
+ from PIL import Image
10
+
11
+
12
+ BINARY_ACTIONS = {
13
+ "w", "a", "s", "d", "jump", "arrow_up", "arrow_left", "arrow_down",
14
+ "arrow_right", "enter", "shift", "ctrl", "alt", "tab", "escape",
15
+ "q", "e", "r", "f", "z", "x", "c", "v", "key_1", "key_2", "key_3",
16
+ "key_4", "mouse_left", "mouse_middle", "mouse_right",
17
+ }
18
+ CONTINUOUS_ACTIONS = {"mouse_dx", "mouse_dy", "zoom"}
19
+ DERIVED_ACTIONS = {
20
+ "move_x", "move_y", "right_mouse", "camera_active",
21
+ "camera_yaw_delta_degrees", "camera_pitch_delta_degrees",
22
+ "mouse_raw_dx", "mouse_raw_dy",
23
+ }
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",
30
+ }
31
+
32
+
33
+ @dataclass(slots=True)
34
+ class OasisDatasetReport:
35
+ dataset_folders: list[str] = field(default_factory=list)
36
+ frames: int = 0
37
+ metadata_rows: int = 0
38
+ valid_rows: int = 0
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
+
46
+ @property
47
+ def ok(self) -> bool:
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] = []
54
+ for entry in entries:
55
+ text = str(entry).strip().strip('"')
56
+ if not text:
57
+ continue
58
+ path = Path(text).expanduser()
59
+ if path not in directories:
60
+ directories.append(path)
61
+ return directories
62
+
63
+
64
+ def _numeric_frame_index(path: Path) -> int | None:
65
+ match = re.search(r"frame_(\d+)", path.stem, re.I) or re.search(r"(\d+)", path.stem)
66
+ return int(match.group(1)) if match else None
67
+
68
+
69
+ def validate_oasis_dataset(value: str | list[str] | tuple[str, ...], *, frame_gap: int = 1) -> OasisDatasetReport:
70
+ report = OasisDatasetReport()
71
+ frame_gap = max(1, int(frame_gap))
72
+ directories = dataset_directories(value)
73
+ if not directories:
74
+ report.errors.append("Select at least one Oasis action dataset folder.")
75
+ return report
76
+ seen_resolution: tuple[int, int] | None = None
77
+ action_counts = {name: 0 for name in sorted(REQUIRED_CANONICAL_ACTIONS | CONTINUOUS_ACTIONS)}
78
+ transition_total = 0
79
+ session_ids: set[str] = set()
80
+
81
+ for directory in directories:
82
+ resolved = directory.resolve()
83
+ report.dataset_folders.append(str(resolved))
84
+ if not directory.is_dir():
85
+ report.errors.append(f"Dataset folder does not exist: {directory}")
86
+ continue
87
+ frames_dir = directory / "frames"
88
+ actions_path = directory / "actions.jsonl"
89
+ if not frames_dir.is_dir():
90
+ report.errors.append(f"{directory} is missing a frames folder.")
91
+ continue
92
+ if not actions_path.is_file():
93
+ report.errors.append(f"{directory} is missing actions.jsonl.")
94
+ continue
95
+ frame_files = sorted(path for path in frames_dir.iterdir() if path.is_file() and path.suffix.casefold() in IMAGE_EXTENSIONS)
96
+ report.frames += len(frame_files)
97
+ if not frame_files:
98
+ report.errors.append(f"{directory} has no image frames.")
99
+ indexed_frames = [index for index in (_numeric_frame_index(path) for path in frame_files) if index is not None]
100
+ frame_paths_by_index: dict[int, list[Path]] = {}
101
+ for frame_file in frame_files:
102
+ index = _numeric_frame_index(frame_file)
103
+ if index is not None:
104
+ frame_paths_by_index.setdefault(index, []).append(frame_file)
105
+ if indexed_frames:
106
+ gaps = [
107
+ (left, right) for left, right in zip(indexed_frames, indexed_frames[1:])
108
+ if right != left + 1
109
+ ]
110
+ if gaps:
111
+ report.warnings.append(f"{directory} has frame ordering gaps such as {gaps[0][0]} to {gaps[0][1]}; invalid transitions will be skipped.")
112
+ rows_by_session: dict[str, list[dict[str, Any]]] = {}
113
+ seen_keys: set[tuple[str, int]] = set()
114
+ for line_number, line in enumerate(actions_path.read_text(encoding="utf-8").splitlines(), 1):
115
+ line = line.strip()
116
+ if not line:
117
+ continue
118
+ report.metadata_rows += 1
119
+ try:
120
+ row = json.loads(line)
121
+ except json.JSONDecodeError:
122
+ report.errors.append(f"{actions_path.name} line {line_number} is not valid JSON.")
123
+ continue
124
+ filename = str(row.get("filename", "")).strip()
125
+ if not filename:
126
+ report.errors.append(f"{actions_path.name} line {line_number} has no frame filename.")
127
+ continue
128
+ frame_path = frames_dir / filename
129
+ if not frame_path.is_file():
130
+ try:
131
+ frame_index = int(row.get("frame_index"))
132
+ except (TypeError, ValueError):
133
+ report.warnings.append(
134
+ f"{actions_path.name} line {line_number} points to missing frame {filename}; skipping row."
135
+ )
136
+ continue
137
+ candidates = frame_paths_by_index.get(frame_index, [])
138
+ if len(candidates) == 1:
139
+ frame_path = candidates[0]
140
+ report.warnings.append(
141
+ f"{actions_path.name} line {line_number} uses {frame_path.name} for missing legacy filename {filename}."
142
+ )
143
+ else:
144
+ report.warnings.append(
145
+ f"{actions_path.name} line {line_number} points to missing frame {filename}; skipping row."
146
+ )
147
+ continue
148
+ try:
149
+ with Image.open(frame_path) as image:
150
+ image.verify()
151
+ with Image.open(frame_path) as image:
152
+ size = image.size
153
+ except Exception as exc:
154
+ report.errors.append(f"Broken image file {frame_path.name}: {exc}")
155
+ continue
156
+ if seen_resolution is None:
157
+ seen_resolution = size
158
+ report.resolution = f"{size[0]}x{size[1]}"
159
+ elif size != seen_resolution:
160
+ report.errors.append(f"Inconsistent frame resolution: {frame_path.name} is {size[0]}x{size[1]}, expected {seen_resolution[0]}x{seen_resolution[1]}.")
161
+ unexpected = sorted(set(row) - SUPPORTED_ACTIONS - METADATA_FIELDS)
162
+ if unexpected:
163
+ report.errors.append(f"{actions_path.name} line {line_number} contains unsupported action field(s): {', '.join(unexpected[:6])}.")
164
+ missing = sorted(name for name in REQUIRED_CANONICAL_ACTIONS if name not in row)
165
+ if missing:
166
+ report.errors.append(f"{actions_path.name} line {line_number} is missing action label(s): {', '.join(missing)}.")
167
+ continue
168
+ try:
169
+ frame_index = int(row.get("frame_index"))
170
+ except (TypeError, ValueError):
171
+ report.errors.append(f"{actions_path.name} line {line_number} has an invalid frame_index.")
172
+ continue
173
+ session_id = str(row.get("session_id") or f"legacy-{directory.name}").strip()
174
+ key = (session_id, frame_index)
175
+ if key in seen_keys:
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))
182
+ except (TypeError, ValueError):
183
+ report.errors.append(f"{actions_path.name} line {line_number} has invalid {name} action value.")
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}")
190
+ report.valid_rows += 1
191
+ for session_id, rows in rows_by_session.items():
192
+ if not rows:
193
+ report.errors.append(f"{directory} has an empty sequence {session_id}.")
194
+ continue
195
+ rows.sort(key=lambda item: int(item["frame_index"]))
196
+ for left, right in zip(rows, rows[1:]):
197
+ if int(right["frame_index"]) != int(left["frame_index"]) + 1:
198
+ report.warnings.append(f"{directory} session {session_id} has an ordering gap at frame {left['frame_index']}; invalid transitions will be skipped.")
199
+ transition_total += sum(
200
+ 1 for left, right in zip(rows, rows[frame_gap:])
201
+ if int(right["frame_index"]) == int(left["frame_index"]) + frame_gap
202
+ )
203
+
204
+ report.sessions = len(session_ids)
205
+ report.action_counts = action_counts
206
+ report.valid_transitions = transition_total
207
+ if report.metadata_rows != report.frames:
208
+ report.warnings.append(f"Frame and label counts differ: {report.frames} frame files, {report.metadata_rows} action rows.")
209
+ if report.valid_rows < 2:
210
+ report.errors.append("The dataset needs at least two valid labelled frames.")
211
+ if report.valid_transitions < 1:
212
+ report.errors.append(f"No valid frame transitions were found for prediction horizon {frame_gap}.")
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
adam/orion.py CHANGED
@@ -10,6 +10,8 @@ IMAGE_EXTENSIONS = {".jpg", ".jpeg", ".png", ".webp", ".bmp"}
10
 
11
 
12
  def dataset_image_count(raw_path: object) -> int:
 
 
13
  path = Path(str(raw_path or "")).expanduser()
14
  if not path.is_dir():
15
  return 0
@@ -36,6 +38,19 @@ def _available_vram_gb() -> float | None:
36
  return None
37
 
38
 
 
 
 
 
 
 
 
 
 
 
 
 
 
39
  def recommend_training_settings(
40
  trainer: str,
41
  image_count: int,
@@ -46,7 +61,7 @@ def recommend_training_settings(
46
  """Return an explainable, conservative starting recipe for manual review."""
47
  trainer = str(trainer).casefold()
48
  images = max(10, int(image_count))
49
- resolution = max(64, min(512, int(resolution)))
50
  vram = _available_vram_gb() if vram_gb is None else vram_gb
51
  cpu_workers = max(2, min(8, (os.cpu_count() or 4) // 2))
52
 
@@ -131,7 +146,7 @@ def review_training_plan(plan: Any) -> dict[str, Any]:
131
  1,
132
  int(args.get("gradient_accumulation_steps", args.get("gradient_accumulation", 1)) or 1),
133
  )
134
- resolution = max(64, int(args.get("resolution", 256) or 256))
135
  exposures = images * epochs if images else 0
136
  optimizer_steps = math.ceil(images / batch / accumulation) * epochs if images else 0
137
  total_steps += optimizer_steps
 
10
 
11
 
12
  def dataset_image_count(raw_path: object) -> int:
13
+ if not str(raw_path or "").strip():
14
+ return 0
15
  path = Path(str(raw_path or "")).expanduser()
16
  if not path.is_dir():
17
  return 0
 
38
  return None
39
 
40
 
41
+ def _resolution_extent(value: object, fallback: int = 256) -> int:
42
+ raw = str(value or fallback).strip().lower()
43
+ if "x" in raw:
44
+ try:
45
+ return max(int(part.strip()) for part in raw.split("x", 1))
46
+ except ValueError:
47
+ return fallback
48
+ try:
49
+ return int(raw)
50
+ except ValueError:
51
+ return fallback
52
+
53
+
54
  def recommend_training_settings(
55
  trainer: str,
56
  image_count: int,
 
61
  """Return an explainable, conservative starting recipe for manual review."""
62
  trainer = str(trainer).casefold()
63
  images = max(10, int(image_count))
64
+ resolution = max(64, min(512, _resolution_extent(resolution)))
65
  vram = _available_vram_gb() if vram_gb is None else vram_gb
66
  cpu_workers = max(2, min(8, (os.cpu_count() or 4) // 2))
67
 
 
146
  1,
147
  int(args.get("gradient_accumulation_steps", args.get("gradient_accumulation", 1)) or 1),
148
  )
149
+ resolution = max(64, _resolution_extent(args.get("resolution", 256)))
150
  exposures = images * epochs if images else 0
151
  optimizer_steps = math.ceil(images / batch / accumulation) * epochs if images else 0
152
  total_steps += optimizer_steps
adam/planner.py CHANGED
@@ -43,6 +43,22 @@ def _project_name(subject: str, suffix: str) -> str:
43
  return f"{safe.title()} {suffix}".strip()[:64]
44
 
45
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
46
  def _collection_mode(request: str) -> str:
47
  """Return the user's requested stopping rule for internet collection."""
48
  return (
@@ -318,6 +334,10 @@ class Planner:
318
  def _deterministic_plan(self, request: str) -> ExecutionPlan | None:
319
  lowered = request.lower()
320
 
 
 
 
 
321
  youtube_plan = self._youtube_dataset_plan(request)
322
  if youtube_plan:
323
  return youtube_plan
@@ -556,10 +576,17 @@ class Planner:
556
  if not re.search(r"\b(train|fine[- ]?tune|retrain|continue|resume)\b", lowered):
557
  return None
558
  fine_tune_payload = self._fine_tune_payload(request)
 
559
  trainer = str(fine_tune_payload.get("trainer", "")) or (
 
 
560
  "lora" if re.search(r"\blora\b", lowered)
561
  else "ddpm" if re.search(r"\bddpm\b", lowered)
562
  else "flow" if re.search(r"\bflow(?:\s+matching)?\b", lowered)
 
 
 
 
563
  else ""
564
  )
565
  action = (
@@ -574,7 +601,7 @@ class Planner:
574
  model_query = ""
575
  resume_match = re.search(
576
  r"\b(?:fine[- ]?tune|retrain|continue|resume)\s+(?:the\s+)?(.+?)"
577
- r"(?:\s+model)?\s+(?:from|on|with)\s+(?:the\s+)?(?:ddpm|lora)\b",
578
  request,
579
  re.I,
580
  )
@@ -619,7 +646,7 @@ class Planner:
619
  if action == "resume_training":
620
  candidates: list[Asset] = []
621
  if model_query:
622
- candidates = self.assets.find("model", model_query, trainer=trainer)
623
  if not candidates:
624
  return ExecutionPlan(
625
  request=request,
@@ -642,7 +669,13 @@ class Planner:
642
  trainer = trainer or model.trainer
643
  ddpm_pipeline = trainer == "ddpm" and (Path(model.path) / "model_index.json").is_file()
644
  flow_model = trainer == "flow" and self._valid_flow_model(Path(model.path))
645
- if (not model.checkpoint or not Path(model.checkpoint).exists()) and not ddpm_pipeline and not flow_model:
 
 
 
 
 
 
646
  return ExecutionPlan(
647
  request=request,
648
  summary=(
@@ -653,10 +686,10 @@ class Planner:
653
  ),
654
  steps=[],
655
  project_name="Resume training",
656
- )
657
  dataset_mode = str(fine_tune_payload.get("dataset_mode", "original"))
658
  if dataset_mode == "existing":
659
- dataset = self._asset_dataset(str(fine_tune_payload.get("dataset_name", "")))
660
  else:
661
  dataset = self._dataset_for_model(model)
662
  if dataset_mode == "new":
@@ -677,21 +710,25 @@ class Planner:
677
  steps=[],
678
  project_name="Resume training",
679
  )
 
680
  command = TrainingCommand.from_dict(
681
  {
682
  "action": "resume_training",
683
  "trainer": trainer,
684
  "dataset": dataset.path,
685
- "model_name": model.name,
686
  "epochs": epochs,
687
  "output": (
688
- str(self._training_output(trainer, f"{model.name} Fine Tune") or model.path)
689
  if trainer == "flow" else model.path
690
  ),
691
  # The DDPM adapter can safely branch from a complete pipeline when
692
  # its exact Accelerate checkpoint has been cleaned up.
693
  "resume_from": model.checkpoint or model.path,
694
- "base_model": self._lora_base_model() if trainer == "lora" else "",
 
 
 
695
  "training_options": training_options,
696
  }
697
  )
@@ -731,7 +768,10 @@ class Planner:
731
  "model_name": model_name,
732
  "epochs": epochs,
733
  "output": str(output),
734
- "base_model": self._lora_base_model() if trainer == "lora" else "",
 
 
 
735
  "training_options": training_options,
736
  }
737
  )
@@ -787,23 +827,24 @@ class Planner:
787
  if dataset_dir.exists():
788
  dataset_dir = dataset_dir.with_name(f"{dataset_dir.name} {datetime.now().strftime('%Y%m%d_%H%M%S')}")
789
  image_count = max(10, min(int(payload.get("image_count", 60)), 5000))
 
790
  arguments: dict[str, Any] = {
791
- "dataset_dir": str(dataset_dir), "model_name": model.name,
792
  "epochs": epochs,
793
  "output_dir": (
794
- str(self._training_output(trainer, f"{model.name} Fine Tune") or model.path)
795
  if trainer == "flow" else model.path
796
  ),
797
  "resume_from": model.checkpoint or model.path, **training_options,
798
  }
799
  if trainer == "lora":
800
- base_model = self._lora_base_model()
801
  if not base_model or not Path(base_model).is_file():
802
  return ExecutionPlan(request=request, summary="Choose a valid SDXL base model in the LoRA app before fine-tuning.", steps=[], project_name="LoRA training")
803
  arguments["base_model"] = base_model
804
  return ExecutionPlan(
805
  request=request,
806
- summary=f"Collect {image_count} new images for {subject}, then continue {model.name} for {epochs} additional epochs.",
807
  steps=[
808
  PlanStep("dataset_collector", "Collect new fine-tune dataset", "Collect and save a reviewable dataset.", {"subject": subject, "image_count": image_count, "collection_mode": "target", "project_name": project, "output_dir": str(dataset_dir)}),
809
  PlanStep(f"{trainer}_trainer", f"Fine-tune {trainer.upper()} model", "Continue from the selected saved model using the newly collected dataset.", arguments),
@@ -884,7 +925,7 @@ class Planner:
884
 
885
  def _dataset_for_phrase(self, phrase: str) -> Asset | None:
886
  """Resolve a friendly dataset phrase, preferring the shortest clear folder match."""
887
- direct = self._asset_dataset(phrase)
888
  if direct:
889
  return direct
890
  wanted = re.sub(r"[^a-z0-9]+", " ", phrase.casefold()).strip()
@@ -900,6 +941,36 @@ class Planner:
900
  candidates.sort(key=lambda item: (len(item.name), item.name.casefold()))
901
  return candidates[0]
902
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
903
  def _plan_training_command(
904
  self,
905
  request: str,
@@ -917,25 +988,48 @@ class Planner:
917
  steps=[],
918
  project_name="Unsupported training request",
919
  )
920
- dataset_path = Path(command.dataset).expanduser()
921
- if not dataset_path.is_dir():
922
- raise PlanningError("The validated training dataset does not exist.")
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
923
  trainer_folder = self._configured_tool_folder(tool_id)
924
- if not trainer_folder:
 
 
925
  raise PlanningError(f"The {spec.name} folder is not connected.")
926
- output_folder = "output_flow_models" if command.trainer == "flow" else "output"
927
- output_root = (Path(trainer_folder) / output_folder).resolve()
928
  output_path = Path(command.output).expanduser().resolve()
929
- try:
930
- output_path.relative_to(output_root)
931
- except ValueError as exc:
932
- raise PlanningError(
933
- f"{spec.name} outputs must stay inside {output_root}."
934
- ) from exc
 
 
 
 
 
 
935
  if command.resume_from and not Path(command.resume_from).exists():
936
  raise PlanningError("The validated resume checkpoint does not exist.")
937
  arguments: dict[str, Any] = {
938
- "dataset_dir": str(dataset_path.resolve()),
939
  "model_name": command.model_name,
940
  "epochs": command.epochs,
941
  "output_dir": str(output_path),
@@ -944,31 +1038,37 @@ class Planner:
944
  if command.resume_from:
945
  arguments["resume_from"] = command.resume_from
946
  if command.trainer == "lora":
947
- if not command.base_model or not Path(command.base_model).is_file():
 
948
  return ExecutionPlan(
949
  request=request,
950
  summary=(
951
- "I found the LoRA dataset, but the connected LoRA trainer has no "
952
- "valid SDXL base model selected. Choose one in the LoRA app first."
953
  ),
954
  steps=[],
955
  project_name="LoRA training",
956
  )
957
- arguments["base_model"] = command.base_model
 
 
 
 
 
958
  verb = "Continue" if command.action == "resume_training" else "Train"
959
  epoch_kind = "additional epochs" if command.action == "resume_training" else "epochs"
960
  return ExecutionPlan(
961
  request=request,
962
  summary=(
963
- f"{verb} {command.model_name} with the registered {command.trainer.upper()} "
964
- f"trainer for {command.epochs} {epoch_kind}. Dataset: {command.dataset}. "
965
  f"Output: {command.output}."
966
  + (f" Training options: {command.training_options}." if command.training_options else "")
967
  ),
968
  steps=[
969
  PlanStep(
970
  tool_id,
971
- f"{verb} {command.trainer.upper()} model",
972
  "Launch the connected trainer with validated paths and stream progress.",
973
  arguments,
974
  )
@@ -987,7 +1087,7 @@ class Planner:
987
  # causing a valid Flow Matching request to fall back to its old
988
  # clarification screen instead of creating a training plan.
989
  explicit_path = re.search(
990
- r"\bfrom\s+(?:the\s+)?([A-Za-z]:[\\/].+?)\s+dataset\s*"
991
  r"(?=[,.;]?\s*(?:train|continue|resume|name|call|save|output|put)\b)",
992
  request,
993
  re.I,
@@ -996,6 +1096,7 @@ class Planner:
996
  return explicit_path.group(1).strip()
997
  patterns = (
998
  r"\btrain\s+(?:the\s+)?(.+?)\s+dataset\s+(?:on|with|for)\b",
 
999
  r"\bfrom\s+(?:the\s+)?(.+?)\s+dataset\b",
1000
  r"\bwith\s+(?:the\s+)?(.+?)\s+dataset\b",
1001
  r"\b(?:the\s+)?(.+?)\s+dataset\s*,?\s+(?:train|use)\b",
@@ -1015,6 +1116,9 @@ class Planner:
1015
  # metadata before looking for the user-facing name.
1016
  request = re.sub(r"\s*\[ADAM_TRAINING_OPTIONS:\{.*?\}\]", "", request, flags=re.I | re.S)
1017
  match = re.search(r"\b(?:name|call)\s+(?:the\s+)?model\s+(.+?)(?:[,\[\{]|$)", request, re.I)
 
 
 
1018
  return _clean_subject(match.group(1)) if match else ""
1019
 
1020
  def _asset_dataset(self, name: str) -> Asset | None:
@@ -1047,19 +1151,68 @@ class Planner:
1047
  return ranked[0]
1048
  return None
1049
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1050
  def _training_output(self, trainer: str, model_name: str) -> Path | None:
1051
  folder = self._configured_tool_folder(f"{trainer}_trainer")
1052
- if not folder:
 
1053
  return None
1054
  safe = re.sub(r"[^A-Za-z0-9._-]+", "_", model_name).strip("._") or "model"
1055
- output_root = "output_flow_models" if trainer == "flow" else "output"
1056
- candidate = (Path(folder) / output_root / safe).resolve()
 
 
 
1057
  if candidate.exists():
1058
  candidate = candidate.with_name(
1059
  f"{candidate.name}_{datetime.now().strftime('%Y%m%d_%H%M%S')}"
1060
  )
1061
  return candidate
1062
 
 
 
 
 
 
 
 
1063
  @staticmethod
1064
  def _valid_flow_model(folder: Path) -> bool:
1065
  try:
@@ -1145,12 +1298,22 @@ class Planner:
1145
  @staticmethod
1146
  def _training_options_from_request(request: str) -> dict[str, Any]:
1147
  match = re.search(r"\[ADAM_TRAINING_OPTIONS:(\{.*?\})\]", request, re.S)
1148
- if not match:
1149
- return {}
1150
- try:
1151
- options = json.loads(match.group(1))
1152
- except json.JSONDecodeError as exc:
1153
- raise PlanningError("Training options could not be read safely.") from exc
 
 
 
 
 
 
 
 
 
 
1154
  if not isinstance(options, dict):
1155
  raise PlanningError("Training options must be a settings object.")
1156
  return options
@@ -1475,6 +1638,82 @@ class Planner:
1475
  project_name="DDPM training",
1476
  )
1477
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1478
  @staticmethod
1479
  def _missing_ddpm_message(fields: dict[str, Any]) -> str:
1480
  missing = [
@@ -1539,17 +1778,31 @@ class Planner:
1539
  def _configured_tool_folder(self, tool_id: str) -> str:
1540
  folders = self.config.get("tool_folders", {})
1541
  if not isinstance(folders, dict):
1542
- return ""
1543
  raw_path = str(folders.get(tool_id, "")).strip()
1544
- return raw_path if raw_path and Path(raw_path).is_dir() else ""
 
 
 
 
 
 
 
 
 
 
 
 
 
1545
 
1546
  def _lora_plan(self, request: str, subject: str) -> ExecutionPlan:
1547
  requested_name = self._model_name_from_request(request)
1548
  model_name = requested_name or subject
 
1549
  project = _project_name(model_name, "LoRA")
1550
  collector_root = self._configured_tool_folder("dataset_collector")
1551
  trainer_root = self._configured_tool_folder("lora_trainer")
1552
- base_model = self._lora_base_model()
1553
  if not collector_root or not trainer_root:
1554
  return ExecutionPlan(
1555
  request=request,
@@ -1564,8 +1817,8 @@ class Planner:
1564
  return ExecutionPlan(
1565
  request=request,
1566
  summary=(
1567
- "Select a valid SDXL base model in the connected LoRA app first. "
1568
- "ADAM will reuse that reviewed setting."
1569
  ),
1570
  steps=[],
1571
  project_name="LoRA training",
@@ -1610,6 +1863,7 @@ class Planner:
1610
  "epochs": epochs,
1611
  "output_dir": str(output_dir),
1612
  "base_model": base_model,
 
1613
  },
1614
  ),
1615
  ]
 
43
  return f"{safe.title()} {suffix}".strip()[:64]
44
 
45
 
46
+ 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())
53
+
54
+
55
+ def _friendly_model_name(asset: Asset) -> str:
56
+ name = str(asset.name or "").strip()
57
+ if re.search(r"^[A-Za-z]:[\\/]", name) or "/" in name or "\\" in name:
58
+ return Path(asset.path).name
59
+ return name or Path(asset.path).name
60
+
61
+
62
  def _collection_mode(request: str) -> str:
63
  """Return the user's requested stopping rule for internet collection."""
64
  return (
 
334
  def _deterministic_plan(self, request: str) -> ExecutionPlan | None:
335
  lowered = request.lower()
336
 
337
+ oasis_player = self._oasis_player_plan(request)
338
+ if oasis_player:
339
+ return oasis_player
340
+
341
  youtube_plan = self._youtube_dataset_plan(request)
342
  if youtube_plan:
343
  return youtube_plan
 
576
  if not re.search(r"\b(train|fine[- ]?tune|retrain|continue|resume)\b", lowered):
577
  return None
578
  fine_tune_payload = self._fine_tune_payload(request)
579
+ marker = re.search(r"\[ADAM_TRAINER:([A-Za-z0-9_ -]+)\]", request, re.I)
580
  trainer = str(fine_tune_payload.get("trainer", "")) or (
581
+ marker.group(1).strip().casefold().replace(" ", "_") if marker else ""
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",
588
+ lowered,
589
+ )
590
  else ""
591
  )
592
  action = (
 
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
  )
 
646
  if action == "resume_training":
647
  candidates: list[Asset] = []
648
  if model_query:
649
+ candidates = self._model_candidates(model_query, trainer=trainer)
650
  if not candidates:
651
  return ExecutionPlan(
652
  request=request,
 
669
  trainer = trainer or model.trainer
670
  ddpm_pipeline = trainer == "ddpm" and (Path(model.path) / "model_index.json").is_file()
671
  flow_model = trainer == "flow" and self._valid_flow_model(Path(model.path))
672
+ oasis_model = trainer == "oasis" and self._valid_oasis_model(Path(model.path))
673
+ if (
674
+ (not model.checkpoint or not Path(model.checkpoint).exists())
675
+ and not ddpm_pipeline
676
+ and not flow_model
677
+ and not oasis_model
678
+ ):
679
  return ExecutionPlan(
680
  request=request,
681
  summary=(
 
686
  ),
687
  steps=[],
688
  project_name="Resume training",
689
+ )
690
  dataset_mode = str(fine_tune_payload.get("dataset_mode", "original"))
691
  if dataset_mode == "existing":
692
+ dataset = self._resolve_dataset_asset(str(fine_tune_payload.get("dataset_name", "")))
693
  else:
694
  dataset = self._dataset_for_model(model)
695
  if dataset_mode == "new":
 
710
  steps=[],
711
  project_name="Resume training",
712
  )
713
+ resumed_model_name = _friendly_model_name(model)
714
  command = TrainingCommand.from_dict(
715
  {
716
  "action": "resume_training",
717
  "trainer": trainer,
718
  "dataset": dataset.path,
719
+ "model_name": resumed_model_name,
720
  "epochs": epochs,
721
  "output": (
722
+ str(self._training_output(trainer, f"{resumed_model_name} Fine Tune") or model.path)
723
  if trainer == "flow" else model.path
724
  ),
725
  # The DDPM adapter can safely branch from a complete pipeline when
726
  # its exact Accelerate checkpoint has been cleaned up.
727
  "resume_from": model.checkpoint or model.path,
728
+ "base_model": (
729
+ str(training_options.get("base_model") or self._lora_base_model())
730
+ if trainer == "lora" else ""
731
+ ),
732
  "training_options": training_options,
733
  }
734
  )
 
768
  "model_name": model_name,
769
  "epochs": epochs,
770
  "output": str(output),
771
+ "base_model": (
772
+ str(training_options.get("base_model") or self._lora_base_model())
773
+ if trainer == "lora" else ""
774
+ ),
775
  "training_options": training_options,
776
  }
777
  )
 
827
  if dataset_dir.exists():
828
  dataset_dir = dataset_dir.with_name(f"{dataset_dir.name} {datetime.now().strftime('%Y%m%d_%H%M%S')}")
829
  image_count = max(10, min(int(payload.get("image_count", 60)), 5000))
830
+ model_name = _friendly_model_name(model)
831
  arguments: dict[str, Any] = {
832
+ "dataset_dir": str(dataset_dir), "model_name": model_name,
833
  "epochs": epochs,
834
  "output_dir": (
835
+ str(self._training_output(trainer, f"{model_name} Fine Tune") or model.path)
836
  if trainer == "flow" else model.path
837
  ),
838
  "resume_from": model.checkpoint or model.path, **training_options,
839
  }
840
  if trainer == "lora":
841
+ base_model = str(training_options.get("base_model") or self._lora_base_model())
842
  if not base_model or not Path(base_model).is_file():
843
  return ExecutionPlan(request=request, summary="Choose a valid SDXL base model in the LoRA app before fine-tuning.", steps=[], project_name="LoRA training")
844
  arguments["base_model"] = base_model
845
  return ExecutionPlan(
846
  request=request,
847
+ summary=f"Collect {image_count} new images for {subject}, then continue {model_name} for {epochs} additional epochs.",
848
  steps=[
849
  PlanStep("dataset_collector", "Collect new fine-tune dataset", "Collect and save a reviewable dataset.", {"subject": subject, "image_count": image_count, "collection_mode": "target", "project_name": project, "output_dir": str(dataset_dir)}),
850
  PlanStep(f"{trainer}_trainer", f"Fine-tune {trainer.upper()} model", "Continue from the selected saved model using the newly collected dataset.", arguments),
 
925
 
926
  def _dataset_for_phrase(self, phrase: str) -> Asset | None:
927
  """Resolve a friendly dataset phrase, preferring the shortest clear folder match."""
928
+ direct = self._resolve_dataset_asset(phrase)
929
  if direct:
930
  return direct
931
  wanted = re.sub(r"[^a-z0-9]+", " ", phrase.casefold()).strip()
 
941
  candidates.sort(key=lambda item: (len(item.name), item.name.casefold()))
942
  return candidates[0]
943
 
944
+ def _model_candidates(self, query: str, *, trainer: str = "") -> list[Asset]:
945
+ raw_query = str(query).strip().strip('"').replace("\\_", "_")
946
+ path = Path(raw_query).expanduser()
947
+ if path.exists():
948
+ resolved = path.resolve()
949
+ matches = [
950
+ asset for asset in self.assets.assets
951
+ if asset.kind == "model"
952
+ and (not trainer or asset.trainer == trainer)
953
+ and Path(asset.path).expanduser().resolve() == resolved
954
+ ]
955
+ if matches:
956
+ return matches
957
+ path_like = re.search(r"^[A-Za-z]:[\\/]", raw_query) or "/" in raw_query or "\\" in raw_query
958
+ if path_like and path.name:
959
+ matches = self.assets.find("model", path.name, trainer=trainer)
960
+ existing = [asset for asset in matches if Path(asset.path).expanduser().exists()]
961
+ if trainer == "oasis":
962
+ valid = [
963
+ asset for asset in existing
964
+ if self._valid_oasis_model(Path(asset.path).expanduser())
965
+ ]
966
+ if valid:
967
+ return valid
968
+ if existing:
969
+ return existing
970
+ if matches:
971
+ return matches
972
+ return self.assets.find("model", raw_query, trainer=trainer)
973
+
974
  def _plan_training_command(
975
  self,
976
  request: str,
 
988
  steps=[],
989
  project_name="Unsupported training request",
990
  )
991
+ if command.trainer == "oasis":
992
+ dataset_paths = self._oasis_dataset_paths(command.dataset)
993
+ if not dataset_paths:
994
+ raise PlanningError(
995
+ "The validated Oasis dataset does not exist or contains no action dataset folders."
996
+ )
997
+ dataset_path = dataset_paths[0]
998
+ resolved_dataset_paths = [str(path.resolve()) for path in dataset_paths]
999
+ dataset_argument = (
1000
+ resolved_dataset_paths[0]
1001
+ if len(resolved_dataset_paths) == 1
1002
+ else resolved_dataset_paths
1003
+ )
1004
+ dataset_label = ";".join(resolved_dataset_paths)
1005
+ else:
1006
+ dataset_path = Path(command.dataset).expanduser()
1007
+ if not dataset_path.is_dir():
1008
+ raise PlanningError("The validated training dataset does not exist.")
1009
+ dataset_argument = str(dataset_path.resolve())
1010
+ dataset_label = dataset_argument
1011
  trainer_folder = self._configured_tool_folder(tool_id)
1012
+ plugin = self.registry.model_plugins.by_trainer(command.trainer)
1013
+ custom_plugin = plugin is not None and command.trainer not in {"ddpm", "flow", "lora", "oasis"}
1014
+ if not trainer_folder and not custom_plugin:
1015
  raise PlanningError(f"The {spec.name} folder is not connected.")
 
 
1016
  output_path = Path(command.output).expanduser().resolve()
1017
+ if custom_plugin:
1018
+ output_root = (self.root / "data" / "model_plugin_outputs" / command.trainer).resolve()
1019
+ output_root.mkdir(parents=True, exist_ok=True)
1020
+ else:
1021
+ output_folder = self._trainer_output_folder(command.trainer)
1022
+ output_root = (Path(trainer_folder) / output_folder).resolve()
1023
+ try:
1024
+ output_path.relative_to(output_root)
1025
+ except ValueError as exc:
1026
+ raise PlanningError(
1027
+ f"{spec.name} outputs must stay inside {output_root}."
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,
1034
  "epochs": command.epochs,
1035
  "output_dir": str(output_path),
 
1038
  if command.resume_from:
1039
  arguments["resume_from"] = command.resume_from
1040
  if command.trainer == "lora":
1041
+ base_model = command.base_model or str((command.training_options or {}).get("base_model", ""))
1042
+ if not base_model or not Path(base_model).is_file():
1043
  return ExecutionPlan(
1044
  request=request,
1045
  summary=(
1046
+ "I found the LoRA dataset, but no valid SDXL base model is selected. "
1047
+ "Choose one in the generated LoRA settings before training."
1048
  ),
1049
  steps=[],
1050
  project_name="LoRA training",
1051
  )
1052
+ arguments["base_model"] = base_model
1053
+ arguments["trigger_word"] = (
1054
+ command.trigger_word
1055
+ or str((command.training_options or {}).get("trigger_word") or "")
1056
+ or command.model_name
1057
+ )
1058
  verb = "Continue" if command.action == "resume_training" else "Train"
1059
  epoch_kind = "additional epochs" if command.action == "resume_training" else "epochs"
1060
  return ExecutionPlan(
1061
  request=request,
1062
  summary=(
1063
+ f"{verb} {command.model_name} with the registered {_trainer_label(command.trainer)} "
1064
+ f"trainer for {command.epochs} {epoch_kind}. Dataset: {dataset_label}. "
1065
  f"Output: {command.output}."
1066
  + (f" Training options: {command.training_options}." if command.training_options else "")
1067
  ),
1068
  steps=[
1069
  PlanStep(
1070
  tool_id,
1071
+ f"{verb} {_trainer_label(command.trainer)} model",
1072
  "Launch the connected trainer with validated paths and stream progress.",
1073
  arguments,
1074
  )
 
1087
  # causing a valid Flow Matching request to fall back to its old
1088
  # clarification screen instead of creating a training plan.
1089
  explicit_path = re.search(
1090
+ r"\b(?:from|use)\s+(?:the\s+)?([A-Za-z]:[\\/].+?)\s+dataset\s*"
1091
  r"(?=[,.;]?\s*(?:train|continue|resume|name|call|save|output|put)\b)",
1092
  request,
1093
  re.I,
 
1096
  return explicit_path.group(1).strip()
1097
  patterns = (
1098
  r"\btrain\s+(?:the\s+)?(.+?)\s+dataset\s+(?:on|with|for)\b",
1099
+ r"\buse\s+(?:the\s+)?(.+?)\s+dataset\b",
1100
  r"\bfrom\s+(?:the\s+)?(.+?)\s+dataset\b",
1101
  r"\bwith\s+(?:the\s+)?(.+?)\s+dataset\b",
1102
  r"\b(?:the\s+)?(.+?)\s+dataset\s*,?\s+(?:train|use)\b",
 
1116
  # metadata before looking for the user-facing name.
1117
  request = re.sub(r"\s*\[ADAM_TRAINING_OPTIONS:\{.*?\}\]", "", request, flags=re.I | re.S)
1118
  match = re.search(r"\b(?:name|call)\s+(?:the\s+)?model\s+(.+?)(?:[,\[\{]|$)", request, re.I)
1119
+ if match:
1120
+ return _clean_subject(match.group(1))
1121
+ match = re.search(r"\b(?:model|checkpoint)\s+called\s+(.+?)(?:[,\[\{]|$)", request, re.I)
1122
  return _clean_subject(match.group(1)) if match else ""
1123
 
1124
  def _asset_dataset(self, name: str) -> Asset | None:
 
1151
  return ranked[0]
1152
  return None
1153
 
1154
+ def _resolve_dataset_asset(self, value: str) -> Asset | None:
1155
+ path = self._resolve_dataset(value)
1156
+ if path:
1157
+ return self.assets.register(
1158
+ kind="dataset",
1159
+ name=path.name,
1160
+ path=str(path),
1161
+ persist=False,
1162
+ )
1163
+ return None
1164
+
1165
+ @staticmethod
1166
+ def _is_oasis_dataset_folder(path: Path) -> bool:
1167
+ return path.is_dir() and (path / "frames").is_dir() and (path / "actions.jsonl").is_file()
1168
+
1169
+ def _oasis_dataset_paths(self, value: str) -> list[Path]:
1170
+ paths: list[Path] = []
1171
+ for raw in str(value or "").split(";"):
1172
+ text = raw.strip().strip('"')
1173
+ if not text:
1174
+ continue
1175
+ path = Path(text).expanduser()
1176
+ if self._is_oasis_dataset_folder(path):
1177
+ paths.append(path.resolve())
1178
+ continue
1179
+ if path.is_dir():
1180
+ children = [
1181
+ child.resolve()
1182
+ for child in sorted(path.rglob("*"), key=lambda item: str(item).casefold())
1183
+ if self._is_oasis_dataset_folder(child)
1184
+ ]
1185
+ paths.extend(children)
1186
+ unique: list[Path] = []
1187
+ for path in paths:
1188
+ if path not in unique:
1189
+ unique.append(path)
1190
+ return unique
1191
+
1192
  def _training_output(self, trainer: str, model_name: str) -> Path | None:
1193
  folder = self._configured_tool_folder(f"{trainer}_trainer")
1194
+ plugin = self.registry.model_plugins.by_trainer(trainer)
1195
+ if not folder and plugin is None:
1196
  return None
1197
  safe = re.sub(r"[^A-Za-z0-9._-]+", "_", model_name).strip("._") or "model"
1198
+ if folder:
1199
+ output_root = self._trainer_output_folder(trainer)
1200
+ candidate = (Path(folder) / output_root / safe).resolve()
1201
+ else:
1202
+ candidate = (self.root / "data" / "model_plugin_outputs" / trainer / safe).resolve()
1203
  if candidate.exists():
1204
  candidate = candidate.with_name(
1205
  f"{candidate.name}_{datetime.now().strftime('%Y%m%d_%H%M%S')}"
1206
  )
1207
  return candidate
1208
 
1209
+ @staticmethod
1210
+ def _trainer_output_folder(trainer: str) -> str:
1211
+ return {
1212
+ "flow": "output_flow_models",
1213
+ "oasis": "output_action_flow_models",
1214
+ }.get(trainer, "output")
1215
+
1216
  @staticmethod
1217
  def _valid_flow_model(folder: Path) -> bool:
1218
  try:
 
1298
  @staticmethod
1299
  def _training_options_from_request(request: str) -> dict[str, Any]:
1300
  match = re.search(r"\[ADAM_TRAINING_OPTIONS:(\{.*?\})\]", request, re.S)
1301
+ options: dict[str, Any] = {}
1302
+ if match:
1303
+ try:
1304
+ parsed = json.loads(match.group(1))
1305
+ except json.JSONDecodeError as exc:
1306
+ raise PlanningError("Training options could not be read safely.") from exc
1307
+ if not isinstance(parsed, dict):
1308
+ raise PlanningError("Training options must be a settings object.")
1309
+ options = parsed
1310
+ trigger_match = re.search(
1311
+ r"\b(?:trigger\s+word|concept\s+token)(?:\s+of|\s*=|\s*:)?\s*['\"\u201c\u201d]?([A-Za-z0-9_.-]{1,128})",
1312
+ request,
1313
+ re.I,
1314
+ )
1315
+ if trigger_match and "trigger_word" not in options:
1316
+ options["trigger_word"] = trigger_match.group(1).strip()
1317
  if not isinstance(options, dict):
1318
  raise PlanningError("Training options must be a settings object.")
1319
  return options
 
1638
  project_name="DDPM training",
1639
  )
1640
 
1641
+ def _oasis_player_plan(self, request: str) -> ExecutionPlan | None:
1642
+ lowered = request.casefold()
1643
+ if not re.search(r"\b(launch|start|play|open|run)\b", lowered):
1644
+ return None
1645
+ if not re.search(r"\b(oasis|playable\s+ai\s+game|action\s+player|world\s+model)\b", lowered):
1646
+ return None
1647
+ self.assets.discover(self.config)
1648
+ path_match = re.search(r"([A-Za-z]:[\\/][^,\n]+)", request)
1649
+ model_path = Path(path_match.group(1).strip().strip("\"'")) if path_match else None
1650
+ model_name = ""
1651
+ if model_path is not None and not self._valid_oasis_model(model_path):
1652
+ return ExecutionPlan(
1653
+ request=request,
1654
+ summary="I need a valid Oasis action model folder to launch the playable window.",
1655
+ steps=[],
1656
+ project_name="Oasis player",
1657
+ )
1658
+ if model_path is None:
1659
+ query_text = re.sub(r"\b(?:with\s+)?seed\s+\d+\b", " ", request, flags=re.I)
1660
+ query_text = re.sub(r"\b(?:starting|reference)\s+frame\s*(?:is|:|=)?\s*[A-Za-z]:[\\/][^,\n]+", " ", query_text, flags=re.I)
1661
+ query = _clean_subject(
1662
+ re.sub(r"\b(launch|start|play|open|run|oasis|action player|world model)\b", " ", query_text, flags=re.I)
1663
+ )
1664
+ candidates = self.assets.find("model", query, trainer="oasis") if query else [
1665
+ asset for asset in self.assets.assets if asset.kind == "model" and asset.trainer == "oasis"
1666
+ ]
1667
+ candidates = [asset for asset in candidates if Path(asset.path).is_dir()]
1668
+ if len(candidates) != 1:
1669
+ examples = ", ".join(asset.name for asset in candidates[:4])
1670
+ return ExecutionPlan(
1671
+ request=request,
1672
+ summary=(
1673
+ "I need one Oasis checkpoint folder to launch."
1674
+ + (f" Matching models: {examples}." if examples else "")
1675
+ ),
1676
+ steps=[],
1677
+ project_name="Oasis player",
1678
+ )
1679
+ model_name = candidates[0].name
1680
+ model_path = Path(candidates[0].path)
1681
+ starting_frame = ""
1682
+ start_match = re.search(r"\b(?:starting|reference)\s+frame\s*(?:is|:|=)?\s*([A-Za-z]:[\\/][^,\n]+)", request, re.I)
1683
+ if start_match:
1684
+ starting_frame = start_match.group(1).strip().strip("\"'")
1685
+ seed_match = re.search(r"\bseed\s+(\d+)", request, re.I)
1686
+ return ExecutionPlan(
1687
+ request=request,
1688
+ summary=f"Launch the Oasis playable window for {model_name or model_path.name}.",
1689
+ steps=[
1690
+ PlanStep(
1691
+ "oasis_player",
1692
+ "Launch Oasis player",
1693
+ "Open the existing Oasis playable inference window in its own process.",
1694
+ {
1695
+ "model_name": model_name or model_path.name,
1696
+ "model_path": str(model_path.resolve()),
1697
+ "starting_frame": starting_frame,
1698
+ "seed": int(seed_match.group(1)) if seed_match else 0,
1699
+ },
1700
+ )
1701
+ ],
1702
+ requires_confirmation=False,
1703
+ project_name="Oasis player",
1704
+ )
1705
+
1706
+ @staticmethod
1707
+ def _valid_oasis_model(folder: Path) -> bool:
1708
+ try:
1709
+ metadata = json.loads((folder / "action_flow_model_info.json").read_text(encoding="utf-8"))
1710
+ return (
1711
+ metadata.get("model_type") == "action_conditioned_rectified_flow_video"
1712
+ and (folder / "unet" / "config.json").is_file()
1713
+ )
1714
+ except (OSError, ValueError, TypeError, json.JSONDecodeError):
1715
+ return False
1716
+
1717
  @staticmethod
1718
  def _missing_ddpm_message(fields: dict[str, Any]) -> str:
1719
  missing = [
 
1778
  def _configured_tool_folder(self, tool_id: str) -> str:
1779
  folders = self.config.get("tool_folders", {})
1780
  if not isinstance(folders, dict):
1781
+ folders = {}
1782
  raw_path = str(folders.get(tool_id, "")).strip()
1783
+ if raw_path and Path(raw_path).is_dir():
1784
+ return raw_path
1785
+ if tool_id == "oasis_trainer":
1786
+ external = self.root / "config" / "external_tools.json"
1787
+ try:
1788
+ payload = json.loads(external.read_text(encoding="utf-8"))
1789
+ except (OSError, ValueError, TypeError, json.JSONDecodeError):
1790
+ return ""
1791
+ for entry in payload.get("tools", []):
1792
+ if isinstance(entry, dict) and entry.get("id") == "external_oasis_game_trainer":
1793
+ candidate = Path(str(entry.get("backend", {}).get("root", ""))).expanduser()
1794
+ if candidate.is_dir():
1795
+ return str(candidate)
1796
+ return ""
1797
 
1798
  def _lora_plan(self, request: str, subject: str) -> ExecutionPlan:
1799
  requested_name = self._model_name_from_request(request)
1800
  model_name = requested_name or subject
1801
+ training_options = self._training_options_from_request(request)
1802
  project = _project_name(model_name, "LoRA")
1803
  collector_root = self._configured_tool_folder("dataset_collector")
1804
  trainer_root = self._configured_tool_folder("lora_trainer")
1805
+ base_model = str(training_options.get("base_model") or self._lora_base_model())
1806
  if not collector_root or not trainer_root:
1807
  return ExecutionPlan(
1808
  request=request,
 
1817
  return ExecutionPlan(
1818
  request=request,
1819
  summary=(
1820
+ "Choose a valid SDXL base model in the generated LoRA settings "
1821
+ "before training."
1822
  ),
1823
  steps=[],
1824
  project_name="LoRA training",
 
1863
  "epochs": epochs,
1864
  "output_dir": str(output_dir),
1865
  "base_model": base_model,
1866
+ **training_options,
1867
  },
1868
  ),
1869
  ]
adam/process_control.py CHANGED
@@ -1,6 +1,7 @@
1
  from __future__ import annotations
2
 
3
  import subprocess
 
4
 
5
 
6
  def set_process_tree_paused(process: subprocess.Popen, paused: bool) -> bool:
@@ -19,3 +20,38 @@ def set_process_tree_paused(process: subprocess.Popen, paused: bool) -> bool:
19
  return True
20
  except Exception:
21
  return False
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
  from __future__ import annotations
2
 
3
  import subprocess
4
+ from typing import Any
5
 
6
 
7
  def set_process_tree_paused(process: subprocess.Popen, paused: bool) -> bool:
 
20
  return True
21
  except Exception:
22
  return False
23
+
24
+
25
+ def terminate_process_tree(process: Any, *, timeout: float = 3.0) -> None:
26
+ """Terminate a process and any children it launched."""
27
+ try:
28
+ import psutil
29
+
30
+ parent = psutil.Process(process.pid)
31
+ children = parent.children(recursive=True)
32
+ targets = [*children, parent]
33
+ for target in targets:
34
+ try:
35
+ target.terminate()
36
+ except (psutil.NoSuchProcess, psutil.AccessDenied):
37
+ continue
38
+ _gone, alive = psutil.wait_procs(targets, timeout=timeout)
39
+ for target in alive:
40
+ try:
41
+ target.kill()
42
+ except (psutil.NoSuchProcess, psutil.AccessDenied):
43
+ continue
44
+ return
45
+ except Exception:
46
+ pass
47
+ try:
48
+ process.terminate()
49
+ except Exception:
50
+ return
51
+ try:
52
+ process.wait(timeout=timeout)
53
+ except Exception:
54
+ try:
55
+ process.kill()
56
+ except Exception:
57
+ pass
adam/recommendations.py ADDED
@@ -0,0 +1,195 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from __future__ import annotations
2
+
3
+ import math
4
+ import os
5
+ from dataclasses import asdict, dataclass, field
6
+ from typing import Any
7
+
8
+ from adam.model_profiles import ModelProfile
9
+ from adam.models import SystemSnapshot
10
+
11
+
12
+ @dataclass(slots=True)
13
+ class SettingsRecommendation:
14
+ profile_id: str
15
+ epochs: int
16
+ settings: dict[str, Any] = field(default_factory=dict)
17
+ reasons: list[str] = field(default_factory=list)
18
+ warnings: list[str] = field(default_factory=list)
19
+ summary: str = ""
20
+ estimated_vram_gb: float | None = None
21
+ risk_level: str = "normal"
22
+ confidence: str = "conservative"
23
+
24
+ def to_dict(self) -> dict[str, Any]:
25
+ return asdict(self)
26
+
27
+
28
+ def _field_default(profile: ModelProfile, key: str, fallback: Any) -> Any:
29
+ return profile.training.get(key, {}).get("default", fallback)
30
+
31
+
32
+ def _clamp_to_schema(profile: ModelProfile, key: str, value: Any) -> Any:
33
+ spec = profile.training.get(key, {})
34
+ kind = str(spec.get("type", "text"))
35
+ try:
36
+ if kind in {"int", "slider"}:
37
+ numeric = int(value)
38
+ return max(int(spec.get("min", numeric)), min(numeric, int(spec.get("max", numeric))))
39
+ if kind == "float":
40
+ numeric = float(value)
41
+ return max(float(spec.get("min", numeric)), min(numeric, float(spec.get("max", numeric))))
42
+ except (TypeError, ValueError):
43
+ return spec.get("default", value)
44
+ if kind == "choice":
45
+ options = list(spec.get("options", []))
46
+ return value if value in options else (options[0] if options else value)
47
+ return value
48
+
49
+
50
+ def estimate_vram_gb(profile: ModelProfile, resolution: int | str, batch_size: int, base_model_gb: float = 0.0) -> float:
51
+ """Broad VRAM estimate used only for warnings and conservative defaults."""
52
+ architecture = profile.architecture.casefold()
53
+ pixels = (max(64, resolution) / 512) ** 2
54
+ if profile.id == "lora" or "lora" in architecture:
55
+ base = max(6.0, base_model_gb * 1.8)
56
+ return base + pixels * max(1, batch_size) * 1.2
57
+ if profile.id == "oasis" or "action_conditioned" in architecture:
58
+ width, height = (resolution, resolution)
59
+ if isinstance(resolution, str) and "x" in resolution:
60
+ try:
61
+ width, height = (int(part) for part in resolution.lower().split("x", 1))
62
+ except ValueError:
63
+ width, height = (256, 144)
64
+ pixels = (max(width, height) / 512) ** 2
65
+ return 4.0 + pixels * max(1, batch_size) * 2.4
66
+ if "flow" in architecture:
67
+ return 2.8 + pixels * max(1, batch_size) * 1.0
68
+ if "diffusion" in architecture:
69
+ return 2.2 + pixels * max(1, batch_size) * 0.9
70
+ return 3.0 + pixels * max(1, batch_size) * 0.8
71
+
72
+
73
+ 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,
80
+ ) -> SettingsRecommendation:
81
+ images = max(10, int(dataset_items or 10))
82
+ raw_resolution = resolution or _field_default(profile, "resolution", 256) or 256
83
+ if isinstance(raw_resolution, str) and "x" in raw_resolution:
84
+ resolution = int(raw_resolution.lower().split("x", 1)[0])
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)
91
+ if snapshot and snapshot.vram_total_gb
92
+ else vram_total
93
+ )
94
+ architecture = profile.architecture.casefold()
95
+ target_exposures = 80_000 if profile.id == "lora" else 180_000 if "diffusion" in architecture else 120_000
96
+ max_epochs = 220 if profile.id == "lora" else 600 if "diffusion" in architecture else 300
97
+ epochs = max(10 if profile.id == "lora" else 25, min(max_epochs, round(target_exposures / images)))
98
+ reasons.append(
99
+ f"Epochs target roughly {target_exposures:,} image exposures, then clamp to the profile's safe range."
100
+ )
101
+
102
+ batch_defaults = {
103
+ 64: 16,
104
+ 128: 12,
105
+ 256: 4,
106
+ 384: 2,
107
+ 512: 1,
108
+ 768: 1,
109
+ 1024: 1,
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":
116
+ batch_defaults.update({512: 2, 768: 1, 1024: 1})
117
+ nearest = min(batch_defaults, key=lambda size: abs(size - resolution))
118
+ batch_size = batch_defaults[nearest]
119
+ reasons.append(f"Batch starts from the closest resolution preset ({nearest}px).")
120
+ if available_vram is not None and available_vram < 8:
121
+ batch_size = max(1, batch_size // 2)
122
+ reasons.append("Available VRAM is below 8 GB, so batch size is reduced conservatively.")
123
+
124
+ settings: dict[str, Any] = {}
125
+ for key in ("resolution", "batch_size"):
126
+ if key in profile.training:
127
+ value = batch_size
128
+ if key == "resolution":
129
+ value = raw_resolution if profile.id == "oasis" else resolution
130
+ settings[key] = _clamp_to_schema(
131
+ profile,
132
+ key,
133
+ value,
134
+ )
135
+ if "learning_rate" in profile.training:
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"):
143
+ if key in profile.training:
144
+ settings[key] = _clamp_to_schema(profile, key, workers)
145
+ for key, value in {
146
+ "gradient_accumulation_steps": 1,
147
+ "gradient_accumulation": 1,
148
+ "mixed_precision": "fp32" if profile.id == "oasis" else "fp16",
149
+ "save_every": max(5, min(25, max(1, epochs // 10))),
150
+ "preview_every": max(5, min(50, max(1, epochs // 10))),
151
+ "training_intensity": 100,
152
+ "gradient_checkpointing": resolution >= 384 or (available_vram is not None and available_vram < 8),
153
+ "rank": 16,
154
+ "alpha": 16,
155
+ "frame_gap": 3,
156
+ "sequence_context": 1,
157
+ "preview_steps": 1 if profile.id == "oasis" else 50 if profile.id == "ddpm" else 10,
158
+ }.items():
159
+ if key in profile.training:
160
+ settings[key] = _clamp_to_schema(profile, key, value)
161
+
162
+ estimated = estimate_vram_gb(profile, resolution, int(settings.get("batch_size", batch_size)), base_model_gb)
163
+ warnings: list[str] = []
164
+ if available_vram is not None and estimated > available_vram * 0.9:
165
+ warnings.append(
166
+ f"Estimated VRAM need is about {estimated:.1f} GB, above the conservative {available_vram * 0.9:.1f} GB working limit."
167
+ )
168
+ if "batch_size" in settings and int(settings["batch_size"]) > 1:
169
+ settings["batch_size"] = max(1, int(settings["batch_size"]) // 2)
170
+ estimated = estimate_vram_gb(profile, resolution, int(settings["batch_size"]), base_model_gb)
171
+ warnings.append(f"Batch size was reduced to {settings['batch_size']} for a safer first run.")
172
+ reasons.append("The initial VRAM estimate was high, so ADAM reduced the batch before applying the recipe.")
173
+ if images < 20:
174
+ warnings.append("Dataset is very small; expect overfitting unless this is just a smoke test.")
175
+ reasons.append("Very small datasets get a warning because quality usually depends more on data cleanup than long training.")
176
+ risk_level = "risky" if warnings else "normal"
177
+
178
+ memory_note = (
179
+ f" using about {available_vram:.1f} GB available VRAM" if available_vram is not None else " without detected VRAM"
180
+ )
181
+ summary = (
182
+ f"Recommended {epochs:,} epochs for {images:,} item(s), "
183
+ f"batch {settings.get('batch_size', batch_size)} at {resolution}px{memory_note}. "
184
+ "Treat this as a starting recipe, not a guarantee."
185
+ )
186
+ return SettingsRecommendation(
187
+ profile_id=profile.id,
188
+ epochs=epochs,
189
+ settings=settings,
190
+ reasons=reasons,
191
+ warnings=warnings,
192
+ summary=summary,
193
+ estimated_vram_gb=estimated,
194
+ risk_level=risk_level,
195
+ )
adam/registry.py CHANGED
@@ -1,10 +1,13 @@
1
  from __future__ import annotations
2
 
3
  import json
 
4
  from dataclasses import dataclass, field
5
  from pathlib import Path
6
  from typing import Any
7
 
 
 
8
 
9
  class RegistryError(RuntimeError):
10
  pass
@@ -60,6 +63,7 @@ class ToolRegistry:
60
  self.root = root.resolve()
61
  self.path = self.root / "config" / "tools.json"
62
  self._tools: dict[str, ToolSpec] = {}
 
63
  self.load()
64
 
65
  def load(self) -> None:
@@ -112,6 +116,26 @@ class ToolRegistry:
112
  safe_entry["demo"] = False
113
  spec = ToolSpec.from_dict(safe_entry)
114
  loaded[spec.id] = spec
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
115
  self._tools = loaded
116
 
117
  def get(self, tool_id: str, *, require_enabled: bool = True) -> ToolSpec:
@@ -144,3 +168,31 @@ class ToolRegistry:
144
  }
145
  for tool in self.enabled()
146
  ]
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
  from __future__ import annotations
2
 
3
  import json
4
+ import logging
5
  from dataclasses import dataclass, field
6
  from pathlib import Path
7
  from typing import Any
8
 
9
+ from adam.model_plugins import ModelPluginRegistry
10
+
11
 
12
  class RegistryError(RuntimeError):
13
  pass
 
63
  self.root = root.resolve()
64
  self.path = self.root / "config" / "tools.json"
65
  self._tools: dict[str, ToolSpec] = {}
66
+ self.model_plugins = ModelPluginRegistry(self.root, logging.getLogger(__name__))
67
  self.load()
68
 
69
  def load(self) -> None:
 
116
  safe_entry["demo"] = False
117
  spec = ToolSpec.from_dict(safe_entry)
118
  loaded[spec.id] = spec
119
+ for entry in self.model_plugins.training_tool_specs() + self.model_plugins.generation_tool_specs():
120
+ if not isinstance(entry, dict):
121
+ continue
122
+ safe_entry = dict(entry)
123
+ backend = dict(safe_entry.get("backend", {}))
124
+ if not backend:
125
+ safe_entry["backend"] = {
126
+ "type": "python",
127
+ "module": "adam.model_plugin_backend",
128
+ "function": "train" if safe_entry.get("category") == "Training" else "generate",
129
+ }
130
+ try:
131
+ spec = ToolSpec.from_dict(safe_entry)
132
+ except RegistryError as exc:
133
+ self.model_plugins.errors.append(f"{safe_entry.get('id', 'unknown')}: {exc}")
134
+ continue
135
+ if spec.id in loaded:
136
+ loaded[spec.id] = _merge_tool_specs(loaded[spec.id], spec)
137
+ else:
138
+ loaded[spec.id] = spec
139
  self._tools = loaded
140
 
141
  def get(self, tool_id: str, *, require_enabled: bool = True) -> ToolSpec:
 
168
  }
169
  for tool in self.enabled()
170
  ]
171
+
172
+
173
+ def _merge_tool_specs(existing: ToolSpec, plugin: ToolSpec) -> ToolSpec:
174
+ """Keep the existing backend while accepting plugin-declared schema arguments."""
175
+ arguments = tuple(dict.fromkeys([*existing.arguments, *plugin.arguments]))
176
+ required_arguments = existing.required_arguments or plugin.required_arguments
177
+ capabilities = tuple(dict.fromkeys([*existing.capabilities, *plugin.capabilities]))
178
+ model_trainers = tuple(
179
+ dict.fromkeys([*existing.model_trainers, *plugin.model_trainers])
180
+ )
181
+ generation_options = dict(existing.generation_options)
182
+ generation_options.update(plugin.generation_options)
183
+ return ToolSpec(
184
+ id=existing.id,
185
+ name=existing.name,
186
+ description=existing.description,
187
+ category=existing.category,
188
+ entry_function=existing.entry_function,
189
+ arguments=arguments,
190
+ required_arguments=required_arguments,
191
+ capabilities=capabilities,
192
+ model_trainers=model_trainers,
193
+ generation_options=generation_options,
194
+ requires_confirmation=existing.requires_confirmation,
195
+ enabled=existing.enabled,
196
+ demo=existing.demo,
197
+ backend=existing.backend,
198
+ )
adam/remote_access.py ADDED
The diff for this file is too large to render. See raw diff
 
adam/remote_api.py ADDED
@@ -0,0 +1,147 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from __future__ import annotations
2
+
3
+ import json
4
+ import math
5
+ from dataclasses import dataclass
6
+ from typing import Any
7
+
8
+
9
+ class RemoteApiError(ValueError):
10
+ def __init__(self, message: str, *, status: int = 400) -> None:
11
+ super().__init__(message)
12
+ self.status = status
13
+
14
+
15
+ @dataclass(frozen=True, slots=True)
16
+ class RemoteResponse:
17
+ status: int
18
+ body: bytes
19
+ content_type: str = "application/json"
20
+ headers: dict[str, str] | None = None
21
+
22
+
23
+ def json_response(payload: dict[str, Any], *, status: int = 200) -> RemoteResponse:
24
+ return RemoteResponse(
25
+ status=status,
26
+ body=json.dumps(payload, separators=(",", ":")).encode("utf-8"),
27
+ content_type="application/json",
28
+ )
29
+
30
+
31
+ def error_response(message: str, *, status: int = 400) -> RemoteResponse:
32
+ return json_response({"ok": False, "error": message}, status=status)
33
+
34
+
35
+ def media_response(body: bytes, content_type: str, *, cache_seconds: int = 86400) -> RemoteResponse:
36
+ return RemoteResponse(
37
+ status=200,
38
+ body=body,
39
+ content_type=content_type,
40
+ headers={"Cache-Control": f"private, max-age={max(0, int(cache_seconds))}"},
41
+ )
42
+
43
+
44
+ def bounded_text(value: Any, *, max_length: int, label: str, required: bool = False) -> str:
45
+ if value is None:
46
+ value = ""
47
+ if not isinstance(value, str):
48
+ value = str(value)
49
+ text = value.strip()
50
+ if required and not text:
51
+ raise RemoteApiError(f"{label} is required.")
52
+ if len(text) > max_length:
53
+ raise RemoteApiError(f"{label} must be {max_length} characters or shorter.")
54
+ return text
55
+
56
+
57
+ def bounded_int(
58
+ value: Any,
59
+ *,
60
+ minimum: int,
61
+ maximum: int,
62
+ default: int,
63
+ label: str,
64
+ ) -> int:
65
+ if value in (None, ""):
66
+ return default
67
+ try:
68
+ if isinstance(value, bool):
69
+ raise ValueError
70
+ number = int(value)
71
+ except (TypeError, ValueError) as exc:
72
+ raise RemoteApiError(f"{label} must be a whole number.") from exc
73
+ if number < minimum or number > maximum:
74
+ raise RemoteApiError(f"{label} must be between {minimum} and {maximum}.")
75
+ return number
76
+
77
+
78
+ def bounded_float(
79
+ value: Any,
80
+ *,
81
+ minimum: float,
82
+ maximum: float,
83
+ default: float,
84
+ label: str,
85
+ ) -> float:
86
+ if value in (None, ""):
87
+ return default
88
+ try:
89
+ if isinstance(value, bool):
90
+ raise ValueError
91
+ number = float(value)
92
+ except (TypeError, ValueError) as exc:
93
+ raise RemoteApiError(f"{label} must be a number.") from exc
94
+ if not math.isfinite(number) or number < minimum or number > maximum:
95
+ raise RemoteApiError(f"{label} must be between {minimum:g} and {maximum:g}.")
96
+ return number
97
+
98
+
99
+ def parse_pagination(query: dict[str, list[str]], *, default_size: int = 24, max_size: int = 80) -> dict[str, int]:
100
+ page = bounded_int(
101
+ (query.get("page") or ["1"])[0],
102
+ minimum=1,
103
+ maximum=1_000_000,
104
+ default=1,
105
+ label="Page",
106
+ )
107
+ page_size = bounded_int(
108
+ (query.get("page_size") or [str(default_size)])[0],
109
+ minimum=1,
110
+ maximum=max_size,
111
+ default=default_size,
112
+ label="Page size",
113
+ )
114
+ return {
115
+ "page": page,
116
+ "page_size": page_size,
117
+ "offset": (page - 1) * page_size,
118
+ "limit": page_size,
119
+ }
120
+
121
+
122
+ def coerce_json_object(payload: Any) -> dict[str, Any]:
123
+ if not isinstance(payload, dict):
124
+ raise RemoteApiError("Send a JSON object.")
125
+ return payload
126
+
127
+
128
+ def sanitized_arguments(arguments: dict[str, Any]) -> dict[str, Any]:
129
+ """Return client-safe arguments without absolute filesystem paths."""
130
+ hidden = {
131
+ "dataset_dir",
132
+ "output_dir",
133
+ "model_path",
134
+ "base_model",
135
+ "base_model_path",
136
+ "resume_from",
137
+ "reference_image",
138
+ }
139
+ clean: dict[str, Any] = {}
140
+ for key, value in arguments.items():
141
+ if key in hidden:
142
+ text = str(value or "")
143
+ clean[f"{key}_name"] = text.replace("\\", "/").rstrip("/").rsplit("/", 1)[-1] if text else ""
144
+ continue
145
+ if isinstance(value, (str, int, float, bool)) or value is None:
146
+ clean[key] = value
147
+ return clean
adam/remote_dashboard.py ADDED
@@ -0,0 +1,141 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from __future__ import annotations
2
+
3
+
4
+ def remote_dashboard_app_html() -> str:
5
+ return """<!doctype html>
6
+ <html lang="en">
7
+ <head>
8
+ <meta charset="utf-8">
9
+ <meta name="viewport" content="width=device-width, initial-scale=1, viewport-fit=cover">
10
+ <title>ADAM Remote</title>
11
+ <style>
12
+ :root{color-scheme:dark;--bg:#050911;--panel:#111823;--panel2:#171f2b;--line:#26374b;--text:#f4f8ff;--muted:#94a2b1;--blue:#3988ff;--blue2:#74bdff;--gold:#f5c84d;--warn:#f8c35d;--bad:#ff6d86}
13
+ *{box-sizing:border-box}body{margin:0;background:radial-gradient(circle at 50% -20%,#102747 0,#07101c 34%,#02060b 100%);color:var(--text);font:15px/1.38 system-ui,-apple-system,BlinkMacSystemFont,"Segoe UI",sans-serif}button,input,select,textarea{font:inherit}button{border:1px solid #32475d;border-radius:8px;background:#1a2431;color:var(--text);padding:10px 12px;font-weight:750;min-height:42px}button.primary{background:linear-gradient(180deg,#438dff,#2467df);color:#fff;border-color:#4a97ff;box-shadow:0 8px 22px rgba(36,103,223,.28)}button.ghost{background:#121923;color:#c5e4ff;border-color:#2d76c9}button.warn{background:#352816;color:#ffe3a3}button.bad{background:#391d29;color:#ffd2db}button.icon{width:40px;min-width:40px;padding:0;color:var(--gold);font-size:22px;background:transparent;border:0}button:disabled{opacity:.46}input,select,textarea{width:100%;border:1px solid var(--line);border-radius:8px;background:#101722;color:var(--text);min-height:44px;padding:10px 12px}textarea{resize:vertical;min-height:86px}.app{min-height:100vh;padding:12px 14px 88px}.top,.titlebar{display:grid;grid-template-columns:42px 1fr 42px;align-items:center;margin:2px 0 14px}.brand{text-align:center}.brand b{display:block;font-size:22px}.brand span,.muted{color:var(--muted);font-size:13px}.titlebar h1{margin:0;text-align:center;font-size:22px;letter-spacing:0}.pill{border:1px solid var(--line);border-radius:999px;padding:6px 10px;color:var(--blue2);background:#0d1721;font-size:12px}.tabs{position:fixed;left:0;right:0;bottom:0;display:grid;grid-template-columns:repeat(5,1fr);background:rgba(4,8,13,.97);border-top:1px solid var(--line);z-index:20}.tabs button{border-radius:0;border:0;border-left:1px solid #111b25;background:transparent;color:var(--muted);font-size:12px;padding:9px 4px}.tabs button b{display:block;font-size:20px;line-height:1}.tabs button.active{color:var(--blue2);background:#0b1421}.view{display:none}.view.active{display:block}.grid{display:grid;gap:10px}.two{grid-template-columns:1fr 1fr}.card{background:rgba(17,24,35,.92);border:1px solid var(--line);border-radius:8px;padding:12px}.panel{background:transparent;border:0;padding:0}.section-row{display:flex;justify-content:space-between;align-items:center;margin:16px 4px 8px}.section{font-size:17px;color:var(--text);font-weight:850;margin:0}.view-all{color:var(--blue2);font-size:13px;font-weight:800}.row{display:flex;gap:8px;align-items:center;flex-wrap:wrap}.row>*{flex:1 1 auto}.status{min-height:20px;color:var(--muted);font-size:13px;white-space:pre-wrap}.progress{height:9px;border-radius:999px;background:#05080c;overflow:hidden;margin-top:10px}.bar{height:100%;width:0;background:linear-gradient(90deg,var(--blue),var(--blue2))}.dataset-list{display:grid;gap:8px}.dataset-card{display:grid;grid-template-columns:82px 1fr 36px;gap:10px;align-items:center;border:1px solid rgba(88,112,140,.35);border-radius:8px;background:linear-gradient(180deg,#171f2b,#111923);padding:8px;width:100%;box-shadow:0 10px 24px rgba(0,0,0,.18)}.dataset-card.compact{grid-template-columns:54px 1fr 22px}.dataset-card.unavailable{opacity:.62}.dataset-card img{width:82px;height:74px;object-fit:cover;border-radius:7px;background:#05080c}.dataset-card.compact img{width:54px;height:54px}.dataset-card b{font-size:16px;overflow-wrap:anywhere}.dataset-actions{grid-column:2/4;display:grid;grid-template-columns:1fr 1.15fr;gap:8px}.dataset-actions button{min-height:36px;padding:7px 10px}.meta{color:var(--muted);font-size:12px}.search-wrap{position:relative}.search-wrap span{position:absolute;left:12px;top:11px;color:var(--muted)}.search-wrap input{padding-left:38px}.chips{display:flex;gap:6px;overflow:auto;padding-bottom:2px}.chip{white-space:nowrap;border:1px solid var(--line);border-radius:999px;background:#0b1219;color:var(--muted);padding:7px 10px}.chip.active{color:#fff;background:#173b63;border-color:#3f8fff}.thumbs{display:grid;grid-template-columns:repeat(auto-fill,minmax(92px,1fr));gap:8px}.thumb{padding:0;overflow:hidden;text-align:left;background:#0d151e}.thumb img{display:block;width:100%;aspect-ratio:1;object-fit:cover;background:#05080c}.thumb div{padding:6px;color:var(--muted);font-size:12px;overflow-wrap:anywhere}.preview{display:block;width:100%;max-height:54vh;object-fit:contain;border-radius:8px;background:#05080c}.form{display:grid;gap:18px}.form-row{display:grid;grid-template-columns:112px 1fr;gap:10px;align-items:center}.form-row.stack{grid-template-columns:1fr}.form-row label,.field label{color:var(--text);font-size:15px;font-weight:780}.hint{color:var(--muted);font-size:12px;margin-top:5px}.toggle{display:grid;grid-template-columns:1fr 58px;gap:10px;align-items:center}.toggle input{width:58px;min-height:32px;accent-color:var(--blue)}.settings-grid{display:grid;gap:12px}.details summary{cursor:pointer;color:var(--blue2);font-weight:800;padding:8px 0}.queue-item{border:1px solid var(--line);border-radius:8px;background:#0d151e;padding:10px}.bad-text{color:var(--bad)}pre{overflow:auto}.big-action{width:100%;min-height:58px;font-size:22px}.selected-note{margin-top:6px;color:var(--muted);font-size:13px}
14
+ @media(min-width:760px){.app{max-width:1020px;margin:0 auto}.wide{grid-template-columns:1.05fr .95fr}.settings-grid{grid-template-columns:1fr 1fr}.tabs{left:50%;transform:translateX(-50%);max-width:1020px;border-left:1px solid var(--line);border-right:1px solid var(--line)}}
15
+ </style>
16
+ </head>
17
+ <body>
18
+ <main class="app">
19
+ <div class="top"><button id="refresh" class="ghost">Refresh</button><div class="brand"><b>ADAM Remote</b><span id="connection">Connecting</span></div><span></span></div>
20
+
21
+ <section id="home" class="view active">
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>
29
+
30
+ <section id="datasetsView" class="view">
31
+ <div class="titlebar"><span></span><h1>Datasets</h1><button id="filterDatasets" class="icon" title="Filter">+</button></div>
32
+ <div class="search-wrap"><span>Q</span><input id="datasetSearch" placeholder="Search datasets..."></div>
33
+ <div id="locationChips" class="chips" style="margin-top:10px"></div>
34
+ <div class="section-row"><div class="section">Favorites</div><div class="view-all">View All</div></div><div id="favoriteDatasets" class="dataset-list"></div>
35
+ <div class="section-row"><div class="section">Recently Used</div><div class="view-all">View All</div></div><div id="recentDatasets" class="dataset-list"></div>
36
+ <div class="section-row"><div class="section">Locations</div></div><div id="locationsListInline" class="dataset-list"></div>
37
+ <div class="section-row"><div class="section">All Remembered</div></div><div id="allDatasets" class="dataset-list"></div>
38
+ <article id="datasetDetail" class="card" style="display:none;margin-top:12px">
39
+ <div class="section">Preview</div>
40
+ <div id="datasetTitle"></div><div id="datasetStats" class="meta"></div>
41
+ <div class="row" style="margin:10px 0"><button id="useDataset" class="primary">Use Dataset</button><button id="selectDatasetForCreate" class="ghost">Select For Create</button></div>
42
+ <div id="datasetGrid" class="thumbs"></div>
43
+ <div class="row" style="margin-top:10px"><button id="prevPage">Previous</button><span id="pageLabel" class="pill">Page 1</span><button id="nextPage">Next</button></div>
44
+ </article>
45
+ <article id="imageDetail" class="card" style="display:none;margin-top:12px"><img id="detailImage" class="preview" alt="Dataset image"><div id="detailName"></div><div id="detailDims" class="meta"></div><div class="field"><label>Caption</label><textarea id="captionEditor"></textarea></div><div class="row"><button id="saveCaption">Save</button><button id="keepImage">Keep</button><button class="warn" id="rejectImage">Reject</button><button id="unreviewImage">Unreview</button></div><div id="imageStatus" class="status"></div></article>
46
+ </section>
47
+
48
+ <section id="createView" class="view">
49
+ <div class="titlebar"><button id="createBack" class="icon" title="Back">&lt;</button><h1>Create Model</h1><button id="createHelp" class="icon" title="Help">?</button></div>
50
+ <article class="panel">
51
+ <div class="form">
52
+ <div class="form-row"><label>Preset</label><div><select id="preset"></select><div class="hint">Save time with a preset configuration.</div></div></div>
53
+ <div class="form-row"><label>Trainer</label><select id="trainer"></select></div>
54
+ <div class="form-row stack"><label>Model Name</label><input id="modelName" placeholder="Minecraft Oasis V3"></div>
55
+ <div class="form-row stack"><label>Dataset</label><button id="chooseDataset" class="ghost">Choose Dataset</button><select id="trainDataset" style="display:none"></select><div id="selectedDatasetNote" class="selected-note">No dataset selected</div></div>
56
+ <div id="baseModelWrap" class="form-row stack"><label>Base model</label><select id="baseModel"></select></div>
57
+ <div id="basicSettings" class="settings-grid"></div>
58
+ <details class="details"><summary>Advanced Settings</summary><div id="advancedSettings" class="settings-grid"></div></details>
59
+ <button id="startTraining" class="primary big-action">Start Training</button>
60
+ <button id="askInstead" class="ghost big-action">Ask ADAM Instead</button>
61
+ <button id="reviewPlan" class="ghost">Review Plan</button>
62
+ <pre id="trainingReview" class="status"></pre>
63
+ </div>
64
+ </article>
65
+ </section>
66
+
67
+ <section id="jobsView" class="view"><article class="card"><div class="section">Jobs</div><div id="queues" class="grid"></div></article></section>
68
+ <section id="settingsView" class="view"><div class="grid"><article class="card"><div class="section">Remote Control</div><label class="toggle"><span>Auto-approve remote training</span><input id="autoApproveTraining" type="checkbox"></label><label class="toggle"><span>Keep screen updated</span><input id="keepAwake" type="checkbox"></label><div id="settingsStatus" class="status"></div></article><article class="card"><div class="section">Remembered Locations</div><div id="locationsList" class="dataset-list"></div></article><article class="card"><div class="section">System</div><div id="system" class="grid two"></div></article></div></section>
69
+ </main>
70
+ <nav class="tabs"><button class="active" data-view="home"><b>^</b>Home</button><button data-view="datasetsView"><b>O</b>Datasets</button><button data-view="createView"><b>+</b>Create</button><button data-view="jobsView"><b>/</b>Jobs</button><button data-view="settingsView"><b>*</b>Settings</button></nav>
71
+ <script>
72
+ (function(){
73
+ var queryString=window.location.search||"";
74
+ var state={datasets:[],locations:[],trainingSchema:{trainers:[],base_models:[],presets:[]},selectedDataset:"",datasetPage:1,selectedItem:null,activeJobId:"",refreshMs:3000,locationFilter:""};
75
+ var timer=null;
76
+ function $(id){return document.getElementById(id)}
77
+ function list(v){return Array.isArray(v)?v:[]}
78
+ function clear(n){if(!n)return;while(n.firstChild)n.removeChild(n.firstChild)}
79
+ function text(id,v){var n=$(id);if(n)n.textContent=v==null?"":String(v)}
80
+ function authUrl(path,extra){var token=queryString.replace(/^\\?/,"");var url=path;if(token)url+=(url.indexOf("?")>=0?"&":"?")+token;if(extra)url+=(url.indexOf("?")>=0?"&":"?")+extra;return url}
81
+ function errorMessage(e){return e&&e.message?e.message:String(e||"Remote request failed.")}
82
+ function requestJson(path,method,payload){return new Promise(function(resolve,reject){var x=new XMLHttpRequest();x.open(method||"GET",authUrl(path),true);x.timeout=25000;x.setRequestHeader("Accept","application/json");if(payload!==undefined)x.setRequestHeader("Content-Type","application/json");x.onreadystatechange=function(){if(x.readyState!==4)return;var p={};try{p=x.responseText?JSON.parse(x.responseText):{}}catch(_e){reject(new Error("ADAM returned an unreadable response."));return}if(x.status<200||x.status>=300){reject(new Error(p.error||"Remote request failed ("+x.status+")."));return}resolve(p)};x.onerror=function(){reject(new Error("Connection failed. Check that ADAM Remote is still running."))};x.ontimeout=function(){reject(new Error("Connection timed out. The PC may be busy."))};x.send(payload===undefined?null:JSON.stringify(payload))})}
83
+ function getJson(path){return requestJson(path,"GET")}function postJson(path,payload){return requestJson(path,"POST",payload||{})}
84
+ function option(select,label,value){var o=document.createElement("option");o.textContent=label||"";o.value=value||"";select.appendChild(o)}
85
+ function appendText(parent,tag,value,cls){var node=document.createElement(tag);node.textContent=value==null?"":String(value);if(cls)node.className=cls;parent.appendChild(node);return node}
86
+ function fieldId(key){return "setting_"+key.replace(/[^A-Za-z0-9_]/g,"_")}
87
+ function activeTrainer(){var id=$("trainer").value;return list(state.trainingSchema.trainers).filter(function(t){return t.id===id})[0]||{settings:{}}}
88
+ function schemaFor(key){return activeTrainer().settings[key]||{}}
89
+ function closestButton(n){while(n&&n!==document.body){if(n.tagName&&n.tagName.toLowerCase()==="button")return n;n=n.parentNode}return null}
90
+ function switchView(id){var nodes=document.querySelectorAll(".view,.tabs button");for(var i=0;i<nodes.length;i++)nodes[i].classList.remove("active");if($(id))$(id).classList.add("active");var tab=document.querySelector('[data-view="'+id+'"]');if(tab)tab.classList.add("active")}
91
+ document.body.addEventListener("click",function(ev){var b=closestButton(ev.target);if(!b)return;var view=b.getAttribute("data-view");if(view)switchView(view);var action=b.getAttribute("data-action");if(action&&state.activeJobId){postJson("/api/job",{job_id:state.activeJobId,action:action}).then(load).catch(function(e){text("promptStatus",errorMessage(e))})}});
92
+ window.addEventListener("error",function(e){text("connection","Phone app error: "+(e.message||"unknown"))});
93
+ window.addEventListener("unhandledrejection",function(e){text("connection","Remote request error: "+errorMessage(e.reason))});
94
+
95
+ function img(src,alt,note){var im=document.createElement("img");im.alt=alt||"";if(src)im.src=authUrl(src);im.onerror=function(){im.style.display="none";if(note)note.textContent="Preview unavailable"};return im}
96
+ function renderStatus(p){p=p||{};text("connection",(p.app||"ADAM")+" online - "+(p.scope||"remote"));var perms=p.permissions||{};var job=p.active_job||null;state.activeJobId=job&&job.id?job.id:"";if(job){text("activeJob",(job.project||"Active job")+" - "+(job.status||"")+" - "+String(job.progress||0)+"% - Time left "+((job.timing||{}).remaining_label||""));text("timeLeft",(job.timing||{}).remaining_label||"-");text("finishTime",(job.timing||{}).finish_label||"-");$("activeBar").style.width=String(job.progress||0)+"%"}else{text("activeJob","No active job.");text("timeLeft","-");text("finishTime","-");$("activeBar").style.width="0%"}var buttons=document.querySelectorAll("[data-action]");for(var i=0;i<buttons.length;i++)buttons[i].disabled=!job||!perms.job_control;$("autoApproveTraining").checked=!!perms.auto_approve_training;$("autoApproveTraining").disabled=!perms.job_control&&!perms.auto_approve_training;renderPreview(p.preview||{});renderLatest(p.latest_generation||{});renderSystem(p.system||{});renderQueues(p)}
97
+ function renderPreview(info){var p=$("preview");if(info.available){p.style.display="block";p.src=authUrl("/api/preview","t="+Date.now());text("previewNote",(info.kind||"preview")+" "+(info.current||info.epoch||"")+"/"+(info.total||""))}else{p.style.display="none";p.removeAttribute("src");text("previewNote","Waiting for a preview.")}}
98
+ function renderLatest(latest){var root=$("latestGeneration");clear(root);var images=list(latest.images).slice(0,8);if(!latest.available||!images.length){var d=document.createElement("div");d.className="status";d.textContent="Finished generated images will appear here.";root.appendChild(d);return}images.forEach(function(item){var b=document.createElement("button");b.className="thumb";b.type="button";var n=document.createElement("div");n.textContent=latest.model_name||"Generated image";b.appendChild(img(item.url,"Generated image",n));b.appendChild(n);b.onclick=function(){window.open(authUrl(item.url,"t="+Date.now()),"_blank","noopener,noreferrer")};root.appendChild(b)})}
99
+ function renderSystem(sys){var root=$("system");clear(root);[["CPU",sys.cpu_percent],["RAM",sys.memory_percent],["GPU",sys.gpu_percent],["VRAM",sys.vram_percent]].forEach(function(m){var d=document.createElement("div");d.className="card";appendText(d,"div",m[0],"meta");appendText(d,"b",m[1]==null?"-":m[1]+"%");root.appendChild(d)})}
100
+ function renderQueues(p){var root=$("queues");clear(root);var jobs=list(p.queue).concat(list(p.completed_jobs),list(p.failed_jobs)).slice(0,40);if(!jobs.length){text("queues","No jobs yet.");return}jobs.forEach(function(job){var d=document.createElement("div");d.className="queue-item";appendText(d,"b",job.project||"ADAM Job");appendText(d,"div",(job.status||"")+" - "+String(job.progress||0)+"%","meta");if(job.current_step_title){var s=document.createElement("div");s.className="meta";s.textContent=job.current_step_title;d.appendChild(s)}if(job.error){var e=document.createElement("div");e.className="bad-text";e.textContent=job.error;d.appendChild(e)}root.appendChild(d)})}
101
+
102
+ function load(){return getJson("/api/status").then(renderStatus).catch(function(e){text("connection","Offline: "+errorMessage(e))})}
103
+ function loadData(){return Promise.all([
104
+ getJson("/api/v1/datasets").then(function(p){state.datasets=list(p.datasets);renderDatasets();fillDatasets()}),
105
+ getJson("/api/v1/datasets/locations").then(function(p){state.locations=list(p.locations);renderLocations()}),
106
+ getJson("/api/v1/training/schema").then(function(p){state.trainingSchema=p||{trainers:[]};fillCreate()})
107
+ ]).catch(function(e){text("trainingReview",errorMessage(e))})}
108
+ function datasetMatches(d){var q=($("datasetSearch").value||"").toLowerCase();if(q&&(d.name||"").toLowerCase().indexOf(q)<0)return false;if(state.locationFilter&&d.location_id!==state.locationFilter)return false;return true}
109
+ function datasetImage(d){return d.thumbnail_url?img(d.thumbnail_url,"Dataset thumbnail"):document.createElement("span")}
110
+ function useDataset(id){if(!id)return;postJson("/api/v1/datasets/"+encodeURIComponent(id)+"/use",{}).then(function(p){var d=p.dataset||{};state.selectedDataset=d.id||id;fillDatasets();if(!$("modelName").value)$("modelName").value=d.name||"";updateSelectedDatasetNote();switchView("createView");loadData()}).catch(function(e){text("datasetStats",errorMessage(e))})}
111
+ function datasetCard(d,compact){var row=document.createElement("div");row.className="dataset-card"+(compact?" compact":"")+(d.available?"":" unavailable");var image=datasetImage(d);var body=document.createElement("div");var name=document.createElement("b");name.textContent=d.name||"Dataset";body.appendChild(name);var meta=document.createElement("div");meta.className="meta";meta.textContent=(d.available?String(d.image_count||d.item_count||0)+" images":"Unavailable")+" - "+(d.dataset_format||"Dataset");body.appendChild(meta);var star=document.createElement("button");star.type="button";star.className="icon";star.textContent=d.favorite?"*":"+";star.title=d.favorite?"Remove favorite":"Favorite";star.onclick=function(ev){ev.stopPropagation();postJson("/api/v1/datasets/"+encodeURIComponent(d.id)+"/favorite",{favorite:!d.favorite}).then(loadData)};row.appendChild(image);row.appendChild(body);row.appendChild(star);if(compact){row.onclick=function(){useDataset(d.id)};return row}var actions=document.createElement("div");actions.className="dataset-actions";var preview=document.createElement("button");preview.type="button";preview.className="ghost";preview.textContent="Preview";preview.onclick=function(){openDataset(d.id,1)};var use=document.createElement("button");use.type="button";use.className="primary";use.textContent="Use Dataset";use.onclick=function(){useDataset(d.id)};actions.appendChild(preview);actions.appendChild(use);row.appendChild(actions);return row}
112
+ function fillList(id,items,empty,compact){var root=$(id);clear(root);items=items.filter(datasetMatches);if(!items.length){var e=document.createElement("div");e.className="status";e.textContent=empty;root.appendChild(e);return}items.forEach(function(d){root.appendChild(datasetCard(d,compact))})}
113
+ function renderDatasets(){var all=state.datasets;fillList("favoriteDatasets",all.filter(function(d){return d.favorite}),"No favorites yet.",false);fillList("recentDatasets",all.filter(function(d){return d.last_used_at}),"Recently used datasets will appear here.",true);fillList("allDatasets",all,"No remembered datasets found. Add locations from desktop Training Studio.",false)}
114
+ function renderLocations(){var chips=$("locationChips");clear(chips);var all=document.createElement("button");all.className="chip"+(state.locationFilter?"":" active");all.textContent="All locations";all.onclick=function(){state.locationFilter="";renderLocations();renderDatasets()};chips.appendChild(all);var listRoot=$("locationsList");var inlineRoot=$("locationsListInline");clear(listRoot);clear(inlineRoot);state.locations.forEach(function(loc){var c=document.createElement("button");c.className="chip"+(state.locationFilter===loc.id?" active":"");c.textContent=loc.name+(loc.available?"":" (unavailable)");c.onclick=function(){state.locationFilter=loc.id;renderLocations();renderDatasets()};chips.appendChild(c);var item=document.createElement("div");item.className="queue-item";appendText(item,"b",loc.name);appendText(item,"div",(loc.available?"available":"unavailable")+" - "+loc.source,"meta");listRoot.appendChild(item.cloneNode(true));inlineRoot.appendChild(item)})}
115
+ function fillDatasets(){var current=$("trainDataset")?$("trainDataset").value:"";if(!$("trainDataset"))return;clear($("trainDataset"));state.datasets.filter(function(d){return d.available}).forEach(function(d){option($("trainDataset"),d.name+" ("+String(d.image_count||d.item_count||0)+")",d.id)});if(current)$("trainDataset").value=current;if(state.selectedDataset)$("trainDataset").value=state.selectedDataset;updateSelectedDatasetNote()}
116
+ function updateSelectedDatasetNote(){var id=state.selectedDataset||($("trainDataset")&&$("trainDataset").value)||"";var d=state.datasets.filter(function(item){return item.id===id})[0];text("selectedDatasetNote",d?d.name+" - "+String(d.image_count||d.item_count||0)+" images":"No dataset selected")}
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();
138
+ }());
139
+ </script>
140
+ </body>
141
+ </html>"""