diff --git a/.gitattributes b/.gitattributes index f63580a03eb89f005fa648c94268fefdb9dd8753..bd6c7be67da881cf7ebf22b4c719020a180f90b0 100644 --- a/.gitattributes +++ b/.gitattributes @@ -1,46 +1,38 @@ -*.7z filter=lfs diff=lfs merge=lfs -text -*.arrow filter=lfs diff=lfs merge=lfs -text -*.bin filter=lfs diff=lfs merge=lfs -text -*.bz2 filter=lfs diff=lfs merge=lfs -text -*.ckpt filter=lfs diff=lfs merge=lfs -text -*.ftz filter=lfs diff=lfs merge=lfs -text -*.gz filter=lfs diff=lfs merge=lfs -text -*.h5 filter=lfs diff=lfs merge=lfs -text -*.joblib filter=lfs diff=lfs merge=lfs -text -*.lfs.* filter=lfs diff=lfs merge=lfs -text -*.mlmodel filter=lfs diff=lfs merge=lfs -text -*.model filter=lfs diff=lfs merge=lfs -text -*.msgpack filter=lfs diff=lfs merge=lfs -text -*.npy filter=lfs diff=lfs merge=lfs -text -*.npz filter=lfs diff=lfs merge=lfs -text -*.onnx filter=lfs diff=lfs merge=lfs -text -*.ot filter=lfs diff=lfs merge=lfs -text -*.parquet filter=lfs diff=lfs merge=lfs -text -*.pb filter=lfs diff=lfs merge=lfs -text -*.pickle filter=lfs diff=lfs merge=lfs -text -*.pkl filter=lfs diff=lfs merge=lfs -text -*.pt filter=lfs diff=lfs merge=lfs -text -*.pth filter=lfs diff=lfs merge=lfs -text -*.rar filter=lfs diff=lfs merge=lfs -text -*.safetensors filter=lfs diff=lfs merge=lfs -text -saved_model/**/* filter=lfs diff=lfs merge=lfs -text -*.tar.* filter=lfs diff=lfs merge=lfs -text -*.tar filter=lfs diff=lfs merge=lfs -text -*.tflite filter=lfs diff=lfs merge=lfs -text -*.tgz filter=lfs diff=lfs merge=lfs -text -*.wasm filter=lfs diff=lfs merge=lfs -text -*.xz filter=lfs diff=lfs merge=lfs -text -*.zip filter=lfs diff=lfs merge=lfs -text -*.zst filter=lfs diff=lfs merge=lfs -text -*tfevents* filter=lfs diff=lfs merge=lfs -text -adam/__pycache__/planner.cpython-311.pyc filter=lfs diff=lfs merge=lfs -text -adam/ui/__pycache__/main_window.cpython-310.pyc filter=lfs diff=lfs merge=lfs -text -adam/ui/__pycache__/main_window.cpython-311.pyc filter=lfs diff=lfs merge=lfs -text -adam/ui/__pycache__/studio.cpython-311.pyc filter=lfs diff=lfs merge=lfs -text +*.7z filter=lfs diff=lfs merge=lfs -text +*.arrow filter=lfs diff=lfs merge=lfs -text +*.bin filter=lfs diff=lfs merge=lfs -text +*.bz2 filter=lfs diff=lfs merge=lfs -text +*.ckpt filter=lfs diff=lfs merge=lfs -text +*.ftz filter=lfs diff=lfs merge=lfs -text +*.gz filter=lfs diff=lfs merge=lfs -text +*.h5 filter=lfs diff=lfs merge=lfs -text +*.joblib filter=lfs diff=lfs merge=lfs -text +*.lfs.* filter=lfs diff=lfs merge=lfs -text +*.mlmodel filter=lfs diff=lfs merge=lfs -text +*.model filter=lfs diff=lfs merge=lfs -text +*.msgpack filter=lfs diff=lfs merge=lfs -text +*.npy filter=lfs diff=lfs merge=lfs -text +*.npz filter=lfs diff=lfs merge=lfs -text +*.onnx filter=lfs diff=lfs merge=lfs -text +*.ot filter=lfs diff=lfs merge=lfs -text +*.parquet filter=lfs diff=lfs merge=lfs -text +*.pb filter=lfs diff=lfs merge=lfs -text +*.pickle filter=lfs diff=lfs merge=lfs -text +*.pkl filter=lfs diff=lfs merge=lfs -text +*.pt filter=lfs diff=lfs merge=lfs -text +*.pth filter=lfs diff=lfs merge=lfs -text +*.rar filter=lfs diff=lfs merge=lfs -text +*.safetensors filter=lfs diff=lfs merge=lfs -text +saved_model/**/* filter=lfs diff=lfs merge=lfs -text +*.tar.* filter=lfs diff=lfs merge=lfs -text +*.tar filter=lfs diff=lfs merge=lfs -text +*.tflite filter=lfs diff=lfs merge=lfs -text +*.tgz filter=lfs diff=lfs merge=lfs -text +*.wasm filter=lfs diff=lfs merge=lfs -text +*.xz filter=lfs diff=lfs merge=lfs -text +*.zip filter=lfs diff=lfs merge=lfs -text +*.zst filter=lfs diff=lfs merge=lfs -text +*tfevents* filter=lfs diff=lfs merge=lfs -text assets/adam_atom.ico filter=lfs diff=lfs merge=lfs -text assets/adam_atom.png filter=lfs diff=lfs merge=lfs -text -build/ADAM/ADAM.exe filter=lfs diff=lfs merge=lfs -text -build/ADAM/ADAM.pkg filter=lfs diff=lfs merge=lfs -text -build/ADAM/PYZ-00.pyz filter=lfs diff=lfs merge=lfs -text -build/ADAM/xref-ADAM.html filter=lfs diff=lfs merge=lfs -text -tests/__pycache__/test_generations.cpython-311-pytest-8.4.2.pyc filter=lfs diff=lfs merge=lfs -text +docs/screenshots/command-center.png filter=lfs diff=lfs merge=lfs -text diff --git a/.gitignore b/.gitignore index 1a1c86423190de659784eefbe0fdc0d119d69bda..85d90917a26c78e56e6555ac7998db5d05e70acc 100644 --- a/.gitignore +++ b/.gitignore @@ -1,28 +1,22 @@ -# Python and test caches __pycache__/ *.py[cod] .pytest_cache/ .venv/ venv/ - -# Local application state and user-connected tools -config/settings.json -config/settings.local.json -config/external_tools.json -!config/external_tools.json -data/* -!data/.gitkeep -logs/ - -# Generated datasets, models, and media +logs/*.log +data/ ADAM_Datasets/ -artifacts/ build/ dist/ +hf-release-*/ +artifacts/ +config/settings.local.json +config/settings.json +*.tmp +*.sqlite3 *.safetensors *.ckpt *.pt *.pth -*.bin -*.onnx -*.tmp +LoRAModelsHere/ +LoRA StableDiffusionModels Here/ diff --git a/Launch ADAM.bat b/Launch ADAM.bat index 0ac66162a0a23d62585a11e053f49aceae7dc433..5b3b4637ead7ee083044cdc3f3a3af524dd8201c 100644 --- a/Launch ADAM.bat +++ b/Launch ADAM.bat @@ -1,61 +1,25 @@ @echo off setlocal cd /d "%~dp0" -python -c "import yt_dlp, cv2, numpy" 2>nul +python -c "import PySide6, psutil, PIL" 2>nul if errorlevel 1 ( - echo Installing the ADAM Video Dataset Collector requirements... - python -m pip install -r "%~dp0requirements.txt" - if errorlevel 1 goto :adam_dependency_error -) -python -c "import selenium, requests, PIL" 2>nul -if errorlevel 1 ( - echo Installing the Dataset Collector requirements for ADAM... - python -m pip install -r "D:\Users\PlayRobloxAllDay\Desktop\Programs\GoogleImageDatasetCollector\requirements.txt" - if errorlevel 1 goto :dependency_error -) -python -c "import datasets, diffusers, transformers, accelerate, torch, torchvision" 2>nul -if errorlevel 1 ( - echo Installing the DDPM training requirements for ADAM... - python -m pip install -r "D:\Users\PlayRobloxAllDay\Desktop\Programs\DDPM\requirements.txt" - if errorlevel 1 goto :ddpm_dependency_error + echo ADAM's desktop requirements are missing from this Python environment. + echo Install them with: + echo python -m pip install -r "%~dp0requirements.txt" + echo. + pause + exit /b 1 ) +rem Optional trainers and collectors are connected through ADAM Settings. +rem Their dependencies belong to their own environments, not startup. python main.py if errorlevel 1 ( echo. - echo ADAM could not start. Install the requirements with: - echo python -m pip install -r requirements.txt + echo ADAM could not start. Check the application log for details. + echo To install the application requirements, run: + echo python -m pip install -r "%~dp0requirements.txt" echo. pause + exit /b 1 ) -endlocal -exit /b - -:adam_dependency_error -echo. -echo ADAM could not install its Video Dataset Collector requirements. -echo Run this command with the same Python used to start ADAM: -echo python -m pip install -r "%~dp0requirements.txt" -echo. -pause -endlocal -exit /b - -:dependency_error -echo. -echo ADAM could not install the Dataset Collector requirements. -echo Run this command and then launch ADAM again: -echo python -m pip install -r "D:\Users\PlayRobloxAllDay\Desktop\Programs\GoogleImageDatasetCollector\requirements.txt" -echo. -pause -endlocal -exit /b - -:ddpm_dependency_error -echo. -echo ADAM could not install the DDPM training requirements. -echo Run this command and then launch ADAM again: -echo python -m pip install -r "D:\Users\PlayRobloxAllDay\Desktop\Programs\DDPM\requirements.txt" -echo. -pause -endlocal -exit /b +endlocal diff --git a/README.md b/README.md index e9adf6a702608dc64b434f12621918f828de3acf..b91e56a48f6ee02d6dd15ba2e8cb597c7100a936 100644 --- a/README.md +++ b/README.md @@ -10,44 +10,385 @@ tags: # ADAM — AI Development and Automation Manager -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. +ADAM is a local, safety-first desktop hub for orchestrating AI project tools. +It includes registered dataset, DDPM, and SDXL LoRA workflows with background +planning, approval gates, progress reporting, and persistent asset history. -## What is included +![ADAM command center](docs/screenshots/command-center.png) -- ADAM source code and the built-in tool registry. -- A clean, empty external-tool registry. -- No model weights, LoRAs, checkpoints, datasets, generated images, job history, personal paths, API keys, or local settings. +Dataset preparation, captioning, and preview placeholders remain clearly marked +as demo tools. The connected Dataset Collector, DDPM trainer, and Local SDXL +LoRA Trainer use real adapters and never fall back to simulated training. -## Requirements +Existing program folders can be connected from **Settings → Tool folders**. +ADAM stores only the path and scans for likely entry points; it does not copy or +modify the external project. Folder assignments can also be pasted into chat: -- Windows 10/11 -- Python 3.10 or later -- Optional: NVIDIA GPU for compatible training workflows -- Optional: Ollama for local chat assistance +```text +DDPM Trainer: D:\AI\DDPM +Flow Matching Trainer: D:\AI\FlowMatchImageGenerator +``` + +On a new computer, install `requirements.txt` in your chosen Python environment +before running `Launch ADAM.bat`. The launcher checks desktop dependencies and +does not install packages automatically or depend on the developer's personal +trainer folders. Install each optional trainer's dependencies according to that +tool's setup instructions before using its ADAM workflow. + +Remote access is disabled by default. Devices with an access token can browse +datasets, edit captions/review marks and submit work. Only enable it for trusted +devices. The desktop **Allow remote job controls and approval changes** setting +also permits remote confirmation, stopping and retrying jobs. A remote browser +can enable training auto-approval only after that desktop permission is granted; +it can always turn auto-approval off. Saved token changes and disabled access +take effect for new requests without restarting the server. + +Use private Tailscale access for connections beyond a trusted local network; +the built-in HTTP listener does not provide transport encryption by itself. +Phone URLs and QR codes contain the access token and should be treated as +credentials. Remote commands require JSON, have bounded request sizes and +connection counts, and reject cross-site browser submissions. These controls +do not sandbox installed Python plugins or connected trainers: install only +code you trust. + +Detection does not automatically authorize training. A real training adapter +remains gated until its dataset, model name, run settings, and output location +are explicit. + +## Training agents + +![ADAM model creation settings](docs/screenshots/create-model.png) + +ADAM's training lifecycle is divided into four explainable responsibilities: + +- **EVE** reviews dataset membership and leaves uncertain images for the user. +- **ORION** reviews planned epochs, batch size, resolution, image exposures, and + estimated optimizer steps. He can require approval but never silently changes + the requested settings. In the Model Creation Assistant, **ORION: apply a + starting recipe** fills a conservative, editable draft from the image count + and selected resolution before a plan is built. +- **ATLAS** watches active training for non-finite loss, sustained critical GPU + temperature, critically low disk space, stalls, and large runtime overruns. + Critical conditions pause the trainer process tree so the user can inspect it. +- **NOVA** examines available post-training previews and samples for unreadable + files and exact-looking duplicate collapse. Her report explicitly separates + technical sample health from subjective or subject-quality review. + +ORION, ATLAS, and NOVA reports are stored with each durable job record and are +shown in Current Plan, Active Job, and Jobs / History respectively. ATLAS's +default thresholds can be overridden in `config/settings.json` with the +`atlas_*` settings defined in `adam/config.py`. + +Every new job passes through the shared preflight and ORION review before its +queue state is chosen. Desktop plans, Remote prompts, and Remote training forms +use the same review. Remote training auto-approval still applies to ordinary +plans, but an ORION warning leaves the job awaiting explicit approval. Reviewing +a plan does not change the requested training settings. + +## Real image collection + +When a valid Dataset Collector folder is connected, the `dataset_collector` +registry entry uses ADAM's real visible-browser adapter. After plan approval it: + +- opens Bing Images in a normal visible Chrome window; +- waits when consent/CAPTCHA/human-verification text is detected; +- resumes automatically after the user resolves the page; +- downloads valid images at least 256×256; +- removes exact duplicate downloads; +- writes a matching `.txt` caption beside every image; and +- records URLs, captions, sources, and dimensions in `metadata.csv`. + +No CAPTCHA or website restriction is bypassed. Closing Chrome or stopping the +job ends collection safely. A new timestamped dataset folder is used rather +than overwriting an existing collection. + +ADAM keeps an incomplete DDPM request in conversation memory. A follow-up such +as `dataset folder Mario, model name Mario V2, epoch count 100, output D:\Runs` +fills the pending fields and validates named datasets against the connected +collector. It will not start if the dataset cannot be found. + +## Showcase videos + +The **Showcase Video** workspace creates a finished MP4 directly from completed +DDPM and Flow Matching models. Select and reorder the models, choose 12–24 +images per model, a 3-, 4-, or 5-second image duration, shared steps and aspect +ratio, provider-compatible samplers, seed, and 720p or 1080p output. ADAM runs +the image batches sequentially and then renders a request-list interface that +tracks the active model, image number, trainer, steps, sampler, and aspect ratio. +LoRA models are intentionally excluded from this streamlined workflow. + +When Ollama is reachable, messages that are not workflow commands receive a +short conversational answer. Ollama may explain or plan, but it still cannot +bypass the registry or confirmation gates. + +## Web search in Chat Mode -## Install and run +Chat Mode can give local Ollama current web context without an API key. Enable +it in **Settings → Planning model**, then ask naturally, for example: + +```text +Search the web for Dandy's World character ideas. +What are the latest Ollama release notes? +Look up a reference for a cyberpunk city character. +``` + +ADAM sends only that search query to Bing's public results feed, reads the +result titles and snippets, +and passes up to five titles, snippets, and links to Ollama. It does not open +the result pages, download anything, or let web content run tools. Results are +untrusted reference material, so ADAM is instructed to cite the links and flag +uncertainty. Disable the setting to keep Chat Mode fully local. + +When you explicitly ask ADAM to **read**, **open**, or **research** result links, +it can read up to three public HTML/text pages and give Ollama short extracts. +For example: `Search the web for Undertale character ideas and read the most +relevant links.` Direct links can be read with `Read https://example.com/ and +summarize it.` Private/local addresses, non-web protocols, oversized pages, +downloads, and more than three pages are blocked. This control can be disabled +in Settings. + +Planning runs away from the interface thread, and conversational Ollama output +is streamed into the chat. ADAM validates training commands against a strict +schema and each registered trainer's declared capabilities before offering a +job. + +In **Settings → Planning model**, **Chat response length** sets the maximum +number of generated tokens for a Chat Mode reply. Higher values allow longer +research summaries but use more time and GPU memory. The default is 1,024, +which gives Qwen3 enough room to reason and still produce a visible response. + +ADAM stores friendly dataset/model names, paths, trainer types, epochs, and +resume checkpoints in `data/assets.json`. Requests such as: + +```text +From the Mario dataset, train it on a DDPM for 300 epochs. +With the Mario dataset, train it on a LoRA for 100 epochs. +Continue the Mario model from the DDPM for 50 epochs. +``` + +are resolved to real paths before approval. Continuation is offered only when a +compatible checkpoint exists. New DDPM runs retain the latest resume checkpoint. + +## Run ```powershell -git clone https://huggingface.co/SyntheticMDProductions/AI_Development_Automation_Manager -cd AI_Development_Automation_Manager -python -m pip install -r requirements.txt python main.py ``` -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. +On Windows, you can also double-click `Launch ADAM.bat`. -## Notes for users +The app requires Python 3.10+ and PySide6. Optional integrations use `psutil` +for system information and `pynvml` for NVIDIA GPU information. -- Training and collection plans require approval before ADAM starts them. -- You are responsible for the licenses, permissions, and rights for any datasets, models, and third-party tools you connect. -- This repository is a downloadable desktop application. It is not a hosted Hugging Face Space or an inference model. +```powershell +python -m pip install -r requirements.txt +``` -## Development +Try: -```powershell -python -m pytest -q +- Click **Create a model…** in Trainer Mode for the guided Model Creation Assistant. +- `Adam, train a LoRA of Hatsune Miku` +- `Adam, collect a dataset of liminal spaces` +- `Adam, generate previews` +- `Adam, check GPU status` +- `From the Mario dataset, train it on a DDPM for 300 epochs` +- `With the Mario dataset, train it on a LoRA for 100 epochs` + +Training and large collection plans are never started until you approve the +plan. All actions are recorded in `logs/adam.log`, while project artifacts live +under `data/projects/`. + +The Model Creation Assistant can start from a built-in Character LoRA, Style +LoRA, DDPM, or Flow Matching preset. It can create a dataset or select a +registered one, recommends starting values, and saves personal presets. The +result still goes through ADAM's normal validated planner and approval gate. +Use **+ Add model** to build a multi-model training batch. Each wide model tab +keeps its own dataset, trainer, name, and settings; the minus button removes an +unwanted model, and tabs can be dragged to change the run order. ADAM validates +all models, presents one combined approval plan, and runs them sequentially so +only one training workflow uses the GPU at a time. A failed step stops the batch +before a later model starts. +Before approval, ADAM adds checks for connected tools, dataset contents, the +LoRA base model, and output-drive free space. Completed dataset and training +jobs also include a suggested next step. + +### Model Batch Builder + +Use **Create model batch…** to paste one requested subject per line. ADAM turns +the list into editable model tabs, removes duplicate names, and lets the current +trainer recipe be applied to any multi-selection of models. The batch is saved +as a draft so it can be closed and resumed later. + +For a review-first workflow, choose **Collect missing datasets first**. This +queues only sequential dataset collection and leaves training in the saved +draft. After collection, reopen the draft, use **Find collected datasets**, and +review each dataset in Training Studio. **Exclude rejected** moves rejected +images out of the training folder into a recoverable quarantine, and **Restore +excluded** reverses it. **Keep all images** marks the whole selected dataset as +accepted in one action, after which individual bad images can still be rejected. +Training remains locked until each model is explicitly +marked as reviewed and ready. If every linked dataset is acceptable as-is, +**Approve all datasets** marks the entire batch ready after one confirmation; +it does not inspect individual images or apply pending rejection decisions. + +Completed Flow Matching models can be selected in **Fine-tune**. ADAM uses the +saved Flow model folder as the continuation source, locks the continuation to +the model's original resolution, and writes the fine-tuned result to a new +output folder. This continues the saved weights while starting a fresh optimizer +and learning-rate schedule; it does not overwrite the original model. + +## Training Studio + +The **Training Studio** turns completed work into a reviewable experiment loop: + +- **Datasets** provides an image gallery, keep/reject decisions, caption editing, + exact duplicate detection, and visually similar duplicate candidates. +- **Experiments** compares job settings and outcomes, opens outputs, marks a + preferred model, and converts successful settings into reusable recipes. +- **Checkpoint Lab** browses model checkpoints and output images, records + consistent prompt/seed evaluations, and sends preview requests through the + normal approval-aware planner. +- **Recipes** preserves training starting points and can import or export + portable JSON recipe files. + +### EVE AI Dataset Review + +In Training Studio → Datasets, **EVE AI Review…** performs a local reference- +guided visual review. Add one or more good reference images and optional bad +references, then choose Keep and Reject confidence thresholds. EVE uses a small +DINOv2 vision model to divide the selected dataset into **Keep**, **Reject**, and +**Uncertain** galleries with confidence scores. The model is downloaded once on +first use and subsequent analysis stays local. + +Nothing is applied automatically. Inspect both sides, double-click images for a +full view, and move selected results between the three groups before choosing +**Apply EVE review**. EVE's decisions remain ordinary Training Studio review +marks: they can be manually changed, and rejected files are not moved until +**Exclude rejected** is selected. The latest proposal is also saved under +`data/eve_reviews/` for auditing. Use **Select all in current group** (or +Ctrl/Shift selection) to move many images at once; EVE transfers only the +chosen thumbnails so manual sorting stays responsive on large datasets. + +Training panels show elapsed time, a progress-based ETA, recent logs, and a +loss sparkline when the connected trainer reports `loss`. Preflight summaries +include clearly labelled workload, duration, VRAM, and disk estimates. These +estimates are planning hints rather than hardware guarantees. + +Create a Model also supports live training previews with a configurable +epoch interval, prompt, and reproducible seed for each model tab. While a +training job is active, its newest 256×256 preview appears in the right sidebar +with the source epoch and next scheduled preview. The full-size trainer output +can be opened from the card. Built-in adapters may publish previews directly; +registered DDPM, Flow, LoRA, APVD, MaskGit, and other trainers can also +participate by writing conventionally named `preview`, `sample`, or `epoch` +images beneath their declared output folder. + +## Generations + +The **Generations** workspace runs compatible registered image generators +without opening their separate desktop interfaces. The connected DDPM and Flow +Matching projects can generate from completed models with a reproducible seed, +sampler or ODE method, step count, image count, and aspect ratio. Generation +work uses the normal ADAM job queue, progress reporting, cancellation, and +logging. + +Every completed batch is stored under `data/generations/` with its images and a +`generation.json` sidecar. The history gallery can open an image or batch folder +and restore the exact settings for another run. DDPM creative notes are stored +with a batch for organization; they are not presented as text conditioning for +an unconditional DDPM model. + +Generation history opens as automatic model folders. Each folder uses the +registered model name and a recent generated image as its cover. Double-clicking +a folder filters history to that model and selects its provider and model in the +generation controls, so the next batch is generated into the same existing model +directory. This view does not move or rewrite older generation files. + +**Generation Cycle…** selects multiple compatible completed models and queues +one generation step per model. Choose images per model, a shared prompt or +creative note, starting seed, slideshow duration, looping, fullscreen playback, +and an optional model/trainer label. When the cycle finishes, ADAM opens the +results as a local slideshow while preserving every ordinary generation record +in history. + +If ADAM discovers a job interrupted by an unexpected shutdown, it offers to +open Jobs & History. The previous record remains intact and can be retried as a +new approval-gated job. Job logs can also be exported for troubleshooting. + +## Connect an existing tool + +ADAM supports importable Python functions and command-line Python scripts. +For a no-code setup, open **Settings → External Tools → Add external tool**. +Choose the program folder, select its training entry script and important +configuration files, then review ADAM's static compatibility and safety report. +The report covers: + +- detected command-line options and required inputs; +- likely dataset formats; +- output and checkpoint behavior; +- progress reporting; +- resume-training support; and +- potentially risky operations visible in the selected entry script. + +The 1–10 rating measures how clearly the script fits ADAM's safe command-line +contract. It is not a guarantee that third-party code is harmless. ADAM does +not execute a script while scanning it, external tools cannot replace built-in +registry entries, and every external-tool run requires explicit approval. + +After registration, a tool can be planned with a request such as: + +```text +Run APVD Model Trainer with dataset=D:\DreamData, epochs=20, output=D:\APVD\output ``` -## License +ADAM will ask for any required inputs that were omitted before it offers the +approval plan. + +For manual registry configuration, edit the relevant item in +`config/tools.json`: + +```json +{ + "backend": { + "type": "python", + "module": "my_tools.lora", + "function": "train" + }, + "demo": false +} +``` + +The function receives a `ToolContext` as its first argument and keyword +arguments from the approved plan. This keeps training code in one place: your +existing GUI and ADAM can both call the same backend. + +For scripts: -ADAM is released under the [MIT License](LICENSE). +```json +{ + "backend": { + "type": "script", + "path": "D:/AI/LoRATrainer/train.py" + }, + "demo": false +} +``` + +ADAM invokes scripts directly with the current Python interpreter, captures +stdout/stderr, and never drives another GUI with mouse clicks. + +## Safety model + +- Plans are shown before execution. +- Long, destructive, or high-volume work requires confirmation. +- Unregistered tools cannot be invoked. +- External paths and arguments are validated before execution. +- The LLM may propose a plan, but only registered tools can execute it. +- Pause, resume, and cancel controls are available for active jobs. +- Every tool action and state transition is logged. + +## Tests + +```powershell +python -m pytest -q +``` diff --git a/adam/assets.py b/adam/assets.py index 0a45ff2a1019832e86de22af11263b8cf0b7bcc2..cd07b48c73a02c6dc57d3fe59831a13c35a9f882 100644 --- a/adam/assets.py +++ b/adam/assets.py @@ -17,6 +17,15 @@ def _normal(value: str) -> str: return re.sub(r"[^a-z0-9]+", " ", value.casefold()).strip() +def _friendly_name(value: str, fallback: str) -> str: + text = str(value or "").strip() + if not text: + return fallback + if re.search(r"^[A-Za-z]:[\\/]", text) or "/" in text or "\\" in text: + return Path(text).name or fallback + return text + + @dataclass(slots=True) class Asset: id: str @@ -28,6 +37,7 @@ class Asset: checkpoint: str = "" epochs: int = 0 created_at: str = "" + metadata: dict[str, Any] | None = None @classmethod def from_dict(cls, payload: dict[str, Any]) -> "Asset": @@ -41,6 +51,7 @@ class Asset: checkpoint=str(payload.get("checkpoint", "")), epochs=int(payload.get("epochs", 0) or 0), created_at=str(payload.get("created_at") or _now()), + metadata=dict(payload.get("metadata") or {}), ) @@ -82,6 +93,7 @@ class AssetRegistry: dataset_id: str = "", checkpoint: str = "", epochs: int = 0, + metadata: dict[str, Any] | None = None, persist: bool = True, ) -> Asset: resolved = str(Path(path).expanduser().resolve()) @@ -100,6 +112,10 @@ class AssetRegistry: asset.checkpoint = checkpoint asset.epochs = int(epochs) asset.created_at = asset.created_at or _now() + if metadata: + current = dict(asset.metadata or {}) + current.update(metadata) + asset.metadata = current if existing is None: self.assets.insert(0, asset) if persist: @@ -118,10 +134,14 @@ class AssetRegistry: key: item[key] for key in ( "kind", "name", "path", "trainer", "dataset_id", - "checkpoint", "epochs", + "checkpoint", "epochs", "metadata", ) if key in item } + if "trigger_word" in item: + metadata = dict(values.get("metadata") or {}) + metadata["trigger_word"] = str(item.get("trigger_word") or "") + values["metadata"] = metadata dataset_path = str(item.get("dataset_path", "")) if item.get("kind") == "model" and dataset_path and Path(dataset_path).is_dir(): dataset = self.register( @@ -151,7 +171,18 @@ class AssetRegistry: folders = config.get("tool_folders", {}) if not isinstance(folders, dict): return + folders = dict(folders) app_root = self.path.parent.parent + if not folders.get("oasis_trainer"): + try: + external = json.loads((app_root / "config" / "external_tools.json").read_text(encoding="utf-8")) + except (OSError, ValueError, TypeError, json.JSONDecodeError): + external = {} + for entry in external.get("tools", []) if isinstance(external, dict) else []: + if isinstance(entry, dict) and entry.get("id") == "external_oasis_game_trainer": + root = str(entry.get("backend", {}).get("root", "")) + if root: + folders["oasis_trainer"] = root external_lora_root = app_root / "LoRAModelsHere" if external_lora_root.is_dir(): for path in external_lora_root.rglob("*.safetensors"): @@ -194,6 +225,7 @@ class AssetRegistry: ("ddpm", "ddpm_trainer", "output"), ("lora", "lora_trainer", "output"), ("flow", "flow_trainer", "output_flow_models"), + ("oasis", "oasis_trainer", "output_action_flow_models"), ): root = Path(str(folders.get(folder_name, ""))) / output_name if not root.is_dir(): @@ -219,6 +251,7 @@ class AssetRegistry: else -1, ) elif trainer == "lora": + trigger_word = "" checkpoints = sorted( ( path for path in folder.glob("*.safetensors") @@ -228,7 +261,12 @@ class AssetRegistry: ) if checkpoints: name = checkpoints[-1].stem.removesuffix("_cancelled") - else: + try: + metadata = json.loads((folder / "model_info.json").read_text(encoding="utf-8")) + trigger_word = str(metadata.get("trigger_word") or "") + except (OSError, ValueError, TypeError, json.JSONDecodeError): + trigger_word = "" + elif trainer == "flow": checkpoints = [] try: metadata = json.loads( @@ -242,8 +280,24 @@ class AssetRegistry: dataset_path = flow_datasets.get(str(folder.resolve()), "") except (OSError, ValueError, TypeError, json.JSONDecodeError): continue + else: + checkpoints = [] + try: + metadata = json.loads( + (folder / "action_flow_model_info.json").read_text(encoding="utf-8") + ) + if metadata.get("model_type") != "action_conditioned_rectified_flow_video": + continue + if not (folder / "unet" / "config.json").is_file(): + continue + name = _friendly_name( + str(metadata.get("model_name") or metadata.get("name") or name), + folder.name, + ) + except (OSError, ValueError, TypeError, json.JSONDecodeError): + continue checkpoint = ( - str(folder) if trainer == "flow" else str(checkpoints[-1]) if checkpoints else "" + str(folder) if trainer in {"flow", "oasis"} else str(checkpoints[-1]) if checkpoints else "" ) dataset_id = "" if dataset_path and Path(dataset_path).is_dir(): @@ -261,8 +315,27 @@ class AssetRegistry: trainer=trainer, dataset_id=dataset_id, checkpoint=checkpoint, + metadata=({"trigger_word": trigger_word or name} if trainer == "lora" else None), persist=False, ) + try: + from adam.dataset_registry import DatasetRegistry + + dataset_registry = DatasetRegistry(app_root, config) + valid_location_ids = {location.id for location in dataset_registry.known_locations()} + self.assets = [ + item for item in self.assets + if not ( + item.kind == "dataset" + and isinstance(item.metadata, dict) + and item.metadata.get("dataset_registry_source") in {"adam", "tool"} + and item.metadata.get("dataset_location_id") + and item.metadata.get("dataset_location_id") not in valid_location_ids + ) + ] + dataset_registry.discover_into_assets(self, persist=False) + except Exception: + pass self.save() def _flow_dataset_paths(self) -> dict[str, str]: diff --git a/adam/atlas.py b/adam/atlas.py index 2d37da60e1d62ab38a947d612c5454d2f807109f..ee5398fa204d09c73a1c672ff7b804e66afeebbf 100644 --- a/adam/atlas.py +++ b/adam/atlas.py @@ -2,7 +2,9 @@ from __future__ import annotations import math import re +import shutil import time +from pathlib import Path from dataclasses import dataclass from datetime import datetime from typing import Any @@ -57,17 +59,35 @@ class AtlasSupervisor: return AtlasDecision("critical", f"GPU temperature remained at {temperature:.0f}°C. ATLAS paused the job.", "pause") free_disk = max(0.0, snapshot.disk_total_gb - snapshot.disk_used_gb) - state["disk_samples"] = state["disk_samples"] + 1 if snapshot.disk_total_gb and free_disk <= self.critical_disk_gb else 0 + disk_known = bool(snapshot.disk_total_gb) + disk_label = "monitored drive" + output = job.plan.steps[job.current_step].arguments.get("output_dir") or job.output_folder + disk_error = False + if output: + disk_label = "output drive" + try: + target = Path(str(output)).expanduser().resolve() + while not target.exists() and target != target.parent: + target = target.parent + usage = shutil.disk_usage(target) + free_disk = usage.free / (1024 ** 3) + disk_known = True + except (OSError, ValueError): + disk_known = False + disk_error = True + state["disk_samples"] = state["disk_samples"] + 1 if disk_known and free_disk <= self.critical_disk_gb else 0 if state["disk_samples"] >= 2: - return AtlasDecision("critical", f"Only {free_disk:.1f} GB remains on the output drive. ATLAS paused the job.", "pause") + return AtlasDecision("critical", f"Only {free_disk:.1f} GB remains on the {disk_label}. ATLAS paused the job.", "pause") stalled_minutes = (moment - state["changed_at"]) / 60 if stalled_minutes >= self.stall_minutes and snapshot.gpu_percent < 5: return AtlasDecision("warning", f"No recorded progress and little GPU activity for {stalled_minutes:.0f} minutes. Check the trainer.") if temperature is not None and temperature >= self.warning_temp: return AtlasDecision("warning", f"GPU temperature is elevated at {temperature:.0f}°C; ATLAS is watching it closely.") - if snapshot.disk_total_gb and free_disk <= self.warning_disk_gb: - return AtlasDecision("warning", f"Output drive space is getting low ({free_disk:.1f} GB free).") + if disk_error: + return AtlasDecision("warning", "ATLAS could not check the output drive's free space. Check that the output location is available.") + if disk_known and free_disk <= self.warning_disk_gb: + return AtlasDecision("warning", f"Space on the {disk_label} is getting low ({free_disk:.1f} GB free).") if snapshot.memory_percent >= 95: return AtlasDecision("warning", f"System memory usage is very high at {snapshot.memory_percent:.0f}%.") diff --git a/adam/commands.py b/adam/commands.py index d7fa4a243f4998cee1d06ce9d25aadaabab627e1..50e42f4234e0844fb389a04b2d5895079a9d98f6 100644 --- a/adam/commands.py +++ b/adam/commands.py @@ -1,8 +1,11 @@ from __future__ import annotations from dataclasses import dataclass +from pathlib import Path from typing import Any +from adam.model_plugins import ModelPluginRegistry, validate_settings + class CommandValidationError(ValueError): pass @@ -18,13 +21,14 @@ class TrainingCommand: output: str = "default output" resume_from: str = "" base_model: str = "" + trigger_word: str = "" training_options: dict[str, Any] | None = None @classmethod def from_dict(cls, payload: dict[str, Any]) -> "TrainingCommand": allowed = { "action", "trainer", "dataset", "model_name", "epochs", "output", - "resume_from", "base_model", + "resume_from", "base_model", "trigger_word", "training_options", } unknown = set(payload) - allowed @@ -42,42 +46,63 @@ class TrainingCommand: output=str(payload.get("output", "default output")).strip(), resume_from=str(payload.get("resume_from", "")).strip(), base_model=str(payload.get("base_model", "")).strip(), + trigger_word=str( + payload.get("trigger_word") + or (payload.get("training_options") or {}).get("trigger_word") + or "" + ).strip(), training_options=dict(payload.get("training_options") or {}), ) except (TypeError, ValueError) as exc: raise CommandValidationError("Training command fields have invalid types.") from exc if command.action not in {"train", "resume_training"}: raise CommandValidationError("Training action must be train or resume_training.") - if command.trainer not in {"ddpm", "lora", "flow"}: - raise CommandValidationError("Trainer must be ddpm, flow, or lora.") + plugin_schema = ModelPluginRegistry(Path.cwd()).training_schema(command.trainer) + if command.trainer not in {"ddpm", "lora", "flow"} and not plugin_schema: + raise CommandValidationError("Trainer must be a discovered model plugin.") if not command.dataset or not command.model_name: raise CommandValidationError("Dataset and model name are required.") if not 1 <= command.epochs <= 100_000: raise CommandValidationError("Epoch count must be between 1 and 100000.") if command.action == "resume_training" and not command.resume_from: raise CommandValidationError("Resume training requires an explicit checkpoint.") + if command.trainer == "lora": + trigger = command.trigger_word or command.model_name + if len(trigger) > 128 or any(char in trigger for char in '<>:"/\\|?*\x00'): + raise CommandValidationError("LoRA trigger word must be short text without reserved characters.") command._validate_options() return command def _validate_options(self) -> None: options = self.training_options or {} - allowed = { - "ddpm": { - "resolution", "batch_size", "learning_rate", "gradient_accumulation_steps", - "dataloader_num_workers", "mixed_precision", "save_every", "preview_steps", - "training_intensity", "preview_enabled", "preview_every", "preview_prompt", - "preview_seed", - }, - "flow": { - "resolution", "batch_size", "learning_rate", "gradient_accumulation", - "workers", "mixed_precision", "save_every", "preview_every", "preview_steps", - "gradient_checkpointing", "preview_enabled", "preview_prompt", "preview_seed", - }, - "lora": {"preview_enabled", "preview_every", "preview_prompt", "preview_seed"}, - }[self.trainer] + schema = ModelPluginRegistry(Path.cwd()).training_schema(self.trainer) + allowed = set(schema) + if not allowed: + allowed = { + "ddpm": { + "resolution", "batch_size", "learning_rate", "gradient_accumulation_steps", + "dataloader_num_workers", "mixed_precision", "save_every", "preview_steps", + "training_intensity", "preview_enabled", "preview_every", "preview_prompt", + "preview_seed", + }, + "flow": { + "resolution", "batch_size", "learning_rate", "gradient_accumulation", + "workers", "mixed_precision", "save_every", "preview_every", "preview_steps", + "gradient_checkpointing", "preview_enabled", "preview_prompt", "preview_seed", + }, + "lora": {"preview_enabled", "preview_every", "preview_prompt", "preview_seed", "trigger_word"}, + }[self.trainer] unknown = set(options) - allowed if unknown: raise CommandValidationError(f"Unsupported {self.trainer} training options: {', '.join(sorted(unknown))}") + if schema: + errors = validate_settings( + {key: spec for key, spec in schema.items() if key in options}, + options, + ) + if errors: + raise CommandValidationError(" ".join(errors)) + return integer_ranges = { "resolution": (64, 512), "batch_size": (1, 64), "gradient_accumulation_steps": (1, 64), "gradient_accumulation": (1, 64), diff --git a/adam/config.py b/adam/config.py index 7866f3d0ab20476d365371c4736733d60701485c..5c05541fdcd9e6ec6a2d233041c81d366531fb61 100644 --- a/adam/config.py +++ b/adam/config.py @@ -26,12 +26,21 @@ DEFAULT_SETTINGS: dict[str, Any] = { "demo_step_delay": 0.24, "max_dataset_images_without_confirmation": 100, "training_presets": {}, + "remote_access": { + "enabled": False, + "bind_address": "127.0.0.1", + "port": 8765, + "token": "", + "allow_job_control": False, + "auto_approve_training": False, + }, "tool_folders": { "dataset_collector": "", "caption_generator": "", "lora_trainer": "", "ddpm_trainer": "", "flow_trainer": "", + "oasis_trainer": "", "preview_generator": "", }, } diff --git a/adam/dataset_lab.py b/adam/dataset_lab.py new file mode 100644 index 0000000000000000000000000000000000000000..276c97f676c5e84c9f31034bf3a1e7d4daf60225 --- /dev/null +++ b/adam/dataset_lab.py @@ -0,0 +1,127 @@ +from __future__ import annotations + +import hashlib +from dataclasses import asdict, dataclass, field +from pathlib import Path + + +IMAGE_EXTENSIONS = {".jpg", ".jpeg", ".png", ".webp", ".bmp"} +VIDEO_EXTENSIONS = {".mp4", ".mov", ".mkv", ".webm", ".avi"} +TEXT_EXTENSIONS = {".txt", ".caption", ".jsonl", ".json"} + + +@dataclass(slots=True) +class DatasetItem: + path: str + kind: str + size_bytes: int + width: int = 0 + height: int = 0 + caption_path: str = "" + duplicate_key: str = "" + + +@dataclass(slots=True) +class DatasetReport: + path: str + total_files: int = 0 + image_count: int = 0 + video_count: int = 0 + text_count: int = 0 + caption_count: int = 0 + missing_caption_count: int = 0 + duplicate_groups: int = 0 + dimensions: dict[str, int] = field(default_factory=dict) + extensions: dict[str, int] = field(default_factory=dict) + items: list[DatasetItem] = field(default_factory=list) + warnings: list[str] = field(default_factory=list) + + def to_dict(self) -> dict: + return asdict(self) + + +def _hash_file(path: Path) -> str: + digest = hashlib.sha1() + with path.open("rb") as handle: + for chunk in iter(lambda: handle.read(1024 * 1024), b""): + digest.update(chunk) + return digest.hexdigest() + + +def scan_dataset(path: str | Path, *, limit: int = 500) -> DatasetReport: + root = Path(path).expanduser().resolve() + report = DatasetReport(path=str(root)) + if not root.is_dir(): + report.warnings.append("Dataset folder does not exist.") + return report + hashes: dict[str, int] = {} + try: + files = [item for item in root.rglob("*") if item.is_file()] + except OSError as exc: + report.warnings.append(f"Dataset could not be scanned: {exc}") + return report + report.total_files = len(files) + for item in files: + suffix = item.suffix.casefold() + report.extensions[suffix or "(none)"] = report.extensions.get(suffix or "(none)", 0) + 1 + if suffix in IMAGE_EXTENSIONS: + report.image_count += 1 + if not any(item.with_suffix(ext).is_file() for ext in (".txt", ".caption")): + report.missing_caption_count += 1 + elif suffix in VIDEO_EXTENSIONS: + report.video_count += 1 + if suffix in TEXT_EXTENSIONS: + report.text_count += 1 + if suffix in {".txt", ".caption"}: + report.caption_count += 1 + for item in files[: max(1, limit)]: + suffix = item.suffix.casefold() + kind = "other" + width = height = 0 + caption_path = "" + duplicate_key = "" + if suffix in IMAGE_EXTENSIONS: + kind = "image" + caption = next((item.with_suffix(ext) for ext in (".txt", ".caption") if item.with_suffix(ext).is_file()), None) + caption_path = str(caption) if caption else "" + try: + from PIL import Image + + with Image.open(item) as image: + width, height = image.size + label = f"{width}x{height}" + report.dimensions[label] = report.dimensions.get(label, 0) + 1 + except Exception: + pass + try: + duplicate_key = _hash_file(item) + hashes[duplicate_key] = hashes.get(duplicate_key, 0) + 1 + except OSError: + duplicate_key = "" + elif suffix in VIDEO_EXTENSIONS: + kind = "video" + elif suffix in TEXT_EXTENSIONS: + kind = "text" + try: + size = item.stat().st_size + except OSError: + size = 0 + report.items.append( + DatasetItem( + path=str(item), + kind=kind, + size_bytes=size, + width=width, + height=height, + caption_path=caption_path, + duplicate_key=duplicate_key, + ) + ) + if report.total_files > limit: + report.warnings.append(f"Showing first {limit:,} files; totals still include all files.") + report.duplicate_groups = sum(1 for count in hashes.values() if count > 1) + if report.image_count and report.missing_caption_count: + report.warnings.append(f"{report.missing_caption_count:,} sampled image(s) do not have sidecar captions.") + if report.duplicate_groups: + report.warnings.append(f"{report.duplicate_groups:,} duplicate image group(s) found in the sample.") + return report diff --git a/adam/dataset_registry.py b/adam/dataset_registry.py new file mode 100644 index 0000000000000000000000000000000000000000..e995a6dbd5b6a3160efd9cb7619777646b4c9d69 --- /dev/null +++ b/adam/dataset_registry.py @@ -0,0 +1,484 @@ +from __future__ import annotations + +import hashlib +import json +import threading +from dataclasses import asdict, dataclass, field +from datetime import datetime, timedelta, timezone +from pathlib import Path +from typing import Any, TYPE_CHECKING +from uuid import uuid4 + +from adam.dataset_lab import IMAGE_EXTENSIONS, TEXT_EXTENSIONS, VIDEO_EXTENSIONS + +if TYPE_CHECKING: + from adam.assets import Asset, AssetRegistry + + +DATASET_MARKERS = { + "dataset_manifest.json", + "metadata", + "frames", + "videos", + "captions", + "actions.jsonl", + "actions.csv", +} +SCAN_LIMIT = 20_000 +ASYNC_REFRESH_AFTER = timedelta(minutes=30) +_scan_lock = threading.Lock() +_active_scans: set[str] = set() + + +def _now() -> str: + return datetime.now(timezone.utc).isoformat() + + +def _path_key(path: str | Path) -> str: + resolved = str(Path(path).expanduser().resolve()) + return hashlib.sha1(resolved.casefold().encode("utf-8")).hexdigest()[:16] + + +def _atomic_json(path: Path, payload: dict[str, Any]) -> None: + path.parent.mkdir(parents=True, exist_ok=True) + temporary = path.with_suffix(path.suffix + ".tmp") + temporary.write_text(json.dumps(payload, indent=2), encoding="utf-8") + temporary.replace(path) + + +def _date(value: str) -> datetime | None: + try: + return datetime.fromisoformat(value) + except (TypeError, ValueError): + return None + + +@dataclass(slots=True) +class DatasetLocation: + id: str + name: str + path: str + source: str = "user" + created_at: str = field(default_factory=_now) + last_seen_at: str = "" + exists: bool = True + + @classmethod + def from_dict(cls, payload: dict[str, Any]) -> "DatasetLocation": + return cls( + id=str(payload.get("id") or _path_key(str(payload.get("path", "")))), + name=str(payload.get("name") or Path(str(payload.get("path", ""))).name or "Datasets"), + path=str(payload.get("path", "")), + source=str(payload.get("source") or "user"), + created_at=str(payload.get("created_at") or _now()), + last_seen_at=str(payload.get("last_seen_at") or ""), + exists=bool(payload.get("exists", True)), + ) + + +@dataclass(slots=True) +class DatasetRecord: + id: str + name: str + path: str + source: str = "asset" + location_id: str = "" + favorite: bool = False + last_used_at: str = "" + discovered_at: str = field(default_factory=_now) + scanned_at: str = "" + exists: bool = True + item_count: int = 0 + image_count: int = 0 + video_count: int = 0 + caption_count: int = 0 + missing_caption_count: int = 0 + sample_image: str = "" + dataset_format: str = "Unknown" + warnings: list[str] = field(default_factory=list) + + @classmethod + def from_dict(cls, payload: dict[str, Any]) -> "DatasetRecord": + return cls( + id=str(payload.get("id") or _path_key(str(payload.get("path", "")))), + name=str(payload.get("name") or Path(str(payload.get("path", ""))).name or "Dataset"), + path=str(payload.get("path", "")), + source=str(payload.get("source") or "asset"), + location_id=str(payload.get("location_id") or ""), + favorite=bool(payload.get("favorite", False)), + last_used_at=str(payload.get("last_used_at") or ""), + discovered_at=str(payload.get("discovered_at") or _now()), + scanned_at=str(payload.get("scanned_at") or ""), + exists=bool(payload.get("exists", True)), + item_count=int(payload.get("item_count", 0) or 0), + image_count=int(payload.get("image_count", 0) or 0), + video_count=int(payload.get("video_count", 0) or 0), + caption_count=int(payload.get("caption_count", 0) or 0), + missing_caption_count=int(payload.get("missing_caption_count", 0) or 0), + sample_image=str(payload.get("sample_image") or ""), + dataset_format=str(payload.get("dataset_format") or "Unknown"), + warnings=[str(item) for item in payload.get("warnings", []) if str(item)], + ) + + +class DatasetRegistry: + """Persistent ADAM-aware index of known dataset locations and datasets.""" + + def __init__(self, root: Path, config: Any | None = None) -> None: + self.root = root.resolve() + self.path = self.root / "data" / "dataset_registry.json" + self.config = config + self.locations: list[DatasetLocation] = [] + self.datasets: dict[str, DatasetRecord] = {} + self.load() + + def load(self) -> None: + try: + payload = json.loads(self.path.read_text(encoding="utf-8")) + except (OSError, ValueError, TypeError, json.JSONDecodeError): + payload = {} + self.locations = [ + DatasetLocation.from_dict(item) + for item in payload.get("locations", []) + if isinstance(item, dict) + ] + self.datasets = { + str(key): DatasetRecord.from_dict(value) + for key, value in (payload.get("datasets", {}) or {}).items() + if isinstance(value, dict) + } + + def save(self) -> None: + _atomic_json( + self.path, + { + "locations": [asdict(item) for item in self.locations], + "datasets": {key: asdict(value) for key, value in self.datasets.items()}, + }, + ) + + def register_location(self, path: str | Path, *, name: str = "", source: str = "user") -> DatasetLocation: + resolved = Path(path).expanduser().resolve() + if not resolved.is_dir(): + raise ValueError("Choose an existing dataset folder.") + key = _path_key(resolved) + existing = next((item for item in self.locations if item.id == key), None) + if existing is None: + existing = DatasetLocation( + id=key, + name=name.strip() or resolved.name or str(resolved), + path=str(resolved), + source=source, + ) + self.locations.insert(0, existing) + else: + existing.name = name.strip() or existing.name + existing.path = str(resolved) + existing.source = source or existing.source + existing.exists = True + existing.last_seen_at = _now() + self.save() + return existing + + def remove_location(self, location_id: str) -> bool: + before = len(self.locations) + self.locations = [item for item in self.locations if item.id != location_id] + changed = len(self.locations) != before + if changed: + self.save() + return changed + + def known_locations(self) -> list[DatasetLocation]: + locations = list(self.locations) + by_path = {Path(item.path).expanduser().resolve(): item for item in locations if item.path} + for path, name, source in self._automatic_location_candidates(): + try: + resolved = path.expanduser().resolve() + except OSError: + continue + if resolved in by_path: + continue + locations.append( + DatasetLocation( + id=_path_key(resolved), + name=name or resolved.name or str(resolved), + path=str(resolved), + source=source, + exists=resolved.is_dir(), + last_seen_at=_now() if resolved.is_dir() else "", + ) + ) + return locations + + def discover_into_assets(self, assets: "AssetRegistry", *, persist: bool = False) -> list["Asset"]: + discovered: list[Asset] = [] + records = self.discover(asset_registry=assets, refresh_missing=False) + for record in records: + if not record.exists: + continue + asset = assets.register( + kind="dataset", + name=record.name, + path=record.path, + metadata={ + "dataset_registry_source": record.source, + "dataset_location_id": record.location_id, + }, + persist=False, + ) + discovered.append(asset) + if persist and discovered: + assets.save() + return discovered + + def discover( + self, + *, + asset_registry: "AssetRegistry | None" = None, + refresh_missing: bool = True, + ) -> list[DatasetRecord]: + self.load() + changed = False + locations = self.known_locations() + known_location_ids = {location.id for location in locations} + for key, record in list(self.datasets.items()): + if record.source in {"adam", "tool"} and record.location_id and record.location_id not in known_location_ids: + del self.datasets[key] + changed = True + for location in locations: + exists = Path(location.path).is_dir() + if location.source == "user": + stored = next((item for item in self.locations if item.id == location.id), None) + if stored: + stored.exists = exists + stored.last_seen_at = _now() if exists else stored.last_seen_at + changed = True + if not exists: + continue + for candidate in self._dataset_candidates(Path(location.path)): + record = self._cached_or_sampled(candidate, source=location.source, location_id=location.id) + self.datasets[record.id] = record + changed = True + if asset_registry is not None: + for asset in getattr(asset_registry, "assets", []): + if getattr(asset, "kind", "") != "dataset": + continue + record = self._cached_or_sampled(Path(asset.path), source="asset", location_id="") + record.name = asset.name or record.name + self.datasets[record.id] = record + changed = True + for record in self.datasets.values(): + record.exists = Path(record.path).is_dir() + if refresh_missing and record.exists and self._needs_refresh(record): + self.refresh_async(record.path, source=record.source, location_id=record.location_id) + if changed: + self.save() + return self.sorted_records() + + def sorted_records(self) -> list[DatasetRecord]: + records = list(self.datasets.values()) + records.sort( + key=lambda item: ( + not item.favorite, + not bool(item.last_used_at), + item.last_used_at or item.discovered_at, + item.name.casefold(), + ), + reverse=False, + ) + favorites = sorted([item for item in records if item.favorite], key=lambda item: item.name.casefold()) + recent = sorted( + [item for item in records if not item.favorite and item.last_used_at], + key=lambda item: item.last_used_at, + reverse=True, + ) + others = sorted( + [item for item in records if not item.favorite and not item.last_used_at], + key=lambda item: item.discovered_at, + reverse=True, + ) + return [*favorites, *recent, *others] + + def record_for_path(self, path: str | Path) -> DatasetRecord: + key = _path_key(path) + record = self.datasets.get(key) + if record is None: + record = self._cached_or_sampled(Path(path), source="asset", location_id="") + self.datasets[key] = record + self.save() + return record + + def favorite(self, path: str | Path, enabled: bool) -> DatasetRecord: + record = self.record_for_path(path) + record.favorite = bool(enabled) + self.save() + return record + + def touch(self, path: str | Path) -> DatasetRecord: + record = self.record_for_path(path) + record.last_used_at = _now() + self.save() + return record + + def refresh_async(self, path: str | Path, *, source: str = "asset", location_id: str = "") -> None: + resolved = str(Path(path).expanduser().resolve()) + key = _path_key(resolved) + with _scan_lock: + if key in _active_scans: + return + _active_scans.add(key) + + def worker() -> None: + try: + record = self._scan(Path(resolved), source=source, location_id=location_id, limit=SCAN_LIMIT) + fresh = DatasetRegistry(self.root, self.config) + current = fresh.datasets.get(record.id) + if current: + record.favorite = current.favorite + record.last_used_at = current.last_used_at + record.discovered_at = current.discovered_at + fresh.datasets[record.id] = record + fresh.save() + finally: + with _scan_lock: + _active_scans.discard(key) + + threading.Thread(target=worker, name="ADAMDatasetScan", daemon=True).start() + + def _automatic_location_candidates(self) -> list[tuple[Path, str, str]]: + candidates: list[tuple[Path, str, str]] = [(self.root / "ADAM_Datasets", "ADAM Datasets", "adam")] + folders = self.config.get("tool_folders", {}) if self.config is not None else {} + folders = folders if isinstance(folders, dict) else {} + collector_raw = str(folders.get("dataset_collector", "")).strip() + if collector_raw: + candidates.append((Path(collector_raw) / "Datasets", "Dataset Collector", "tool")) + for tool_id in ("ddpm_trainer", "lora_trainer", "flow_trainer", "oasis_trainer"): + raw_root = str(folders.get(tool_id, "")).strip() + if not raw_root: + continue + root = Path(raw_root) + label = tool_id.removesuffix("_trainer").upper() + for child in ("Datasets", "datasets", "OldDatasets"): + candidates.append((root / child, f"{label} {child}", "tool")) + external_tools = self.root / "config" / "external_tools.json" + try: + payload = json.loads(external_tools.read_text(encoding="utf-8")) + except (OSError, ValueError, TypeError, json.JSONDecodeError): + payload = {} + for entry in payload.get("tools", []) if isinstance(payload, dict) else []: + if not isinstance(entry, dict): + continue + backend = entry.get("backend", {}) + if not isinstance(backend, dict): + continue + raw_root = str(backend.get("root", "")).strip() + if raw_root: + root = Path(raw_root) + candidates.append((root / "Datasets", f"{entry.get('name', 'External tool')} Datasets", "tool")) + candidates.append((root / "OldDatasets", f"{entry.get('name', 'External tool')} OldDatasets", "tool")) + return candidates + + def _dataset_candidates(self, location: Path) -> list[Path]: + candidates: list[Path] = [] + try: + children = [item for item in location.iterdir() if item.is_dir()] + except OSError: + children = [] + for child in children[:1000]: + if self._looks_like_dataset(child): + candidates.append(child) + if candidates: + return candidates + if self._looks_like_dataset(location): + candidates.append(location) + return candidates + + def _looks_like_dataset(self, folder: Path) -> bool: + if not folder.is_dir(): + return False + try: + names = {item.name for item in folder.iterdir()} + except OSError: + return False + if names & DATASET_MARKERS: + return True + checked = 0 + for item in folder.rglob("*"): + if checked >= 200: + break + checked += 1 + if item.is_file() and item.suffix.casefold() in IMAGE_EXTENSIONS | VIDEO_EXTENSIONS: + return True + return False + + def _cached_or_sampled(self, path: Path, *, source: str, location_id: str) -> DatasetRecord: + key = _path_key(path) + current = self.datasets.get(key) + if current is not None: + current.exists = path.is_dir() + current.source = current.source or source + current.location_id = current.location_id or location_id + return current + return self._scan(path, source=source, location_id=location_id, limit=800) + + def _needs_refresh(self, record: DatasetRecord) -> bool: + if not record.scanned_at: + return True + scanned = _date(record.scanned_at) + return scanned is None or datetime.now(timezone.utc) - scanned > ASYNC_REFRESH_AFTER + + def _scan(self, path: Path, *, source: str, location_id: str, limit: int) -> DatasetRecord: + resolved = path.expanduser().resolve() + record = DatasetRecord( + id=_path_key(resolved), + name=resolved.name or "Dataset", + path=str(resolved), + source=source, + location_id=location_id, + exists=resolved.is_dir(), + scanned_at=_now(), + ) + if not resolved.is_dir(): + record.warnings.append("Dataset folder is unavailable.") + return record + files_seen = 0 + capped = False + try: + iterator = resolved.rglob("*") + for item in iterator: + if not item.is_file(): + continue + files_seen += 1 + suffix = item.suffix.casefold() + if suffix in IMAGE_EXTENSIONS: + record.image_count += 1 + if not record.sample_image: + record.sample_image = str(item) + if not any(item.with_suffix(ext).is_file() for ext in (".txt", ".caption")): + record.missing_caption_count += 1 + elif suffix in VIDEO_EXTENSIONS: + record.video_count += 1 + if suffix in TEXT_EXTENSIONS and suffix in {".txt", ".caption"}: + record.caption_count += 1 + if files_seen >= limit: + capped = True + break + except OSError as exc: + record.warnings.append(f"Dataset could not be scanned: {exc}") + record.item_count = record.image_count + record.video_count + record.dataset_format = self._format_label(resolved, record) + if capped: + record.warnings.append(f"Counts are sampled from the first {limit:,} files.") + return record + + @staticmethod + def _format_label(folder: Path, record: DatasetRecord) -> str: + if (folder / "actions.jsonl").is_file() or (folder / "actions.csv").is_file(): + return "Oasis action dataset" + if record.video_count: + return "Video dataset" + if record.image_count and record.caption_count: + return "Captioned image dataset" + if record.image_count: + return "Image dataset" + return "Dataset folder" diff --git a/adam/eve.py b/adam/eve.py index 010f50594b7aeff7ab3f27649cee496c434d6602..5a56935ea251f6a001e10e28ea9de65c2d33dd16 100644 --- a/adam/eve.py +++ b/adam/eve.py @@ -87,8 +87,9 @@ def classify_eve_embeddings( class EveVisionModel: """Lazy local DINOv2 feature extractor used by EVE.""" - def __init__(self, model_id: str = EVE_MODEL_ID) -> None: + def __init__(self, model_id: str = EVE_MODEL_ID, *, prefer_gpu: bool = True) -> None: self.model_id = model_id + self.prefer_gpu = prefer_gpu self._processor = None self._model = None self._device = "cpu" @@ -106,10 +107,20 @@ class EveVisionModel: raise RuntimeError( "EVE needs PyTorch and Transformers. Launch ADAM with its normal Python environment." ) from exc - self._device = "cuda" if torch.cuda.is_available() else "cpu" + self._device = "cuda" if self.prefer_gpu and torch.cuda.is_available() else "cpu" self._processor = AutoImageProcessor.from_pretrained(self.model_id, use_fast=True) self._model = AutoModel.from_pretrained(self.model_id).to(self._device).eval() + def unload(self) -> None: + self._processor = None + self._model = None + try: + import torch + if torch.cuda.is_available(): + torch.cuda.empty_cache() + except ImportError: + pass + def embed(self, paths: Sequence[str | Path], progress: Callable[[int, int], None] | None = None) -> list[list[float]]: self.load() import torch diff --git a/adam/executor.py b/adam/executor.py index 09b20ec898bb8aae6fa39a98538c530362b645ed..36da82bf8208326aa3e9181a4f249e56e73fe3c1 100644 --- a/adam/executor.py +++ b/adam/executor.py @@ -12,6 +12,7 @@ from dataclasses import dataclass from pathlib import Path from typing import Any, Callable +from adam.process_control import terminate_process_tree from adam.registry import ToolRegistry, ToolSpec @@ -23,7 +24,15 @@ class ToolCancelled(ToolExecutionError): pass -ProgressCallback = Callable[[int, str], None] +class ToolAdjustmentRequested(ToolExecutionError): + """A trainer stopped cleanly so a job can continue with new settings.""" + + def __init__(self, message: str, details: dict[str, Any] | None = None) -> None: + super().__init__(message) + self.details = details or {} + + +ProgressCallback = Callable[..., None] LogCallback = Callable[[str], None] PreviewCallback = Callable[[dict[str, Any]], None] @@ -38,14 +47,16 @@ class ToolContext: progress_callback: ProgressCallback log_callback: LogCallback preview_callback: PreviewCallback = lambda _preview: None + adjustment_event: threading.Event | None = None + adjustment_request: dict[str, Any] | None = None step_delay: float = 0.2 def log(self, message: str) -> None: self.log_callback(message) - def progress(self, percent: int, message: str) -> None: + def progress(self, percent: int, message: str, **details: Any) -> None: self.checkpoint() - self.progress_callback(max(0, min(int(percent), 100)), message) + self.progress_callback(max(0, min(int(percent), 100)), message, **details) def preview( self, path: str | Path, *, epoch: int = 0, next_epoch: int = 0, @@ -104,6 +115,8 @@ class ToolExecutor: progress_callback: ProgressCallback, log_callback: LogCallback, preview_callback: PreviewCallback | None = None, + adjustment_event: threading.Event | None = None, + adjustment_request: dict[str, Any] | None = None, ) -> dict[str, Any]: spec = self.registry.get(tool_id) self._validate_arguments(spec, arguments) @@ -116,6 +129,8 @@ class ToolExecutor: progress_callback=progress_callback, log_callback=log_callback, preview_callback=preview_callback or (lambda _preview: None), + adjustment_event=adjustment_event, + adjustment_request=adjustment_request, step_delay=self.step_delay, ) backend_type = str(spec.backend.get("type", "")).lower() @@ -243,30 +258,7 @@ class ToolExecutor: def stop_process_tree() -> None: """Stop the script and any workers it launched.""" - if process_controller is not None: - try: - descendants = process_controller.children(recursive=True) - for child in descendants: - try: - child.terminate() - except Exception: - pass - process_controller.terminate() - try: - import psutil - - _gone, alive = psutil.wait_procs(descendants, timeout=2) - for child in alive: - try: - child.kill() - except Exception: - pass - except Exception: - pass - return - except Exception: - pass - process.terminate() + terminate_process_tree(process, timeout=3) try: while True: diff --git a/adam/experiment_tracker.py b/adam/experiment_tracker.py new file mode 100644 index 0000000000000000000000000000000000000000..d533dd05ed9c01c9b30768e9990e5b91928b7580 --- /dev/null +++ b/adam/experiment_tracker.py @@ -0,0 +1,387 @@ +from __future__ import annotations + +import json +import re +import sqlite3 +from dataclasses import dataclass +from datetime import datetime, timezone +from pathlib import Path +from typing import Any + +from adam.models import Job, JobStatus, SystemSnapshot + + +def _utc_now() -> str: + return datetime.now(timezone.utc).isoformat() + + +def _json(value: Any) -> str: + return json.dumps(value, sort_keys=True) + + +def _safe_json(value: str, fallback: Any) -> Any: + try: + parsed = json.loads(value or "") + except (TypeError, ValueError, json.JSONDecodeError): + return fallback + return parsed if isinstance(parsed, type(fallback)) else fallback + + +def _safe_int(value: Any, default: int = 0) -> int: + try: + if isinstance(value, bool): + return default + return int(value) + except (TypeError, ValueError): + return default + + +def _safe_float(value: Any, default: float = 0.0) -> float: + try: + if isinstance(value, bool): + return default + return float(value) + except (TypeError, ValueError): + return default + + +def _loss_from_logs(logs: list[str]) -> float | None: + for line in reversed(logs): + match = re.search(r"\bloss(?:\s*[:=]\s*|\s+)(-?\d+(?:\.\d+)?(?:e[+-]?\d+)?)", line, re.I) + if match: + try: + return float(match.group(1)) + except ValueError: + return None + return None + + +def _duration_seconds(job: Job) -> int: + if not job.started_at: + return 0 + try: + start = datetime.fromisoformat(job.started_at) + end = datetime.fromisoformat(job.ended_at) if job.ended_at else datetime.now(timezone.utc) + return max(0, int((end - start).total_seconds())) + except ValueError: + return 0 + + +def _image_count(path: str) -> int: + folder = Path(path).expanduser() + if not folder.is_dir(): + return 0 + try: + return sum( + 1 for item in folder.rglob("*") + if item.is_file() and item.suffix.casefold() in {".png", ".jpg", ".jpeg", ".webp", ".bmp"} + ) + except OSError: + return 0 + + +@dataclass(slots=True) +class ExperimentRun: + id: str + job_id: str + timestamp: str + model_architecture: str + model_name: str + trigger_word: str + base_model: str + dataset_path: str + dataset_name: str + dataset_item_count: int + epochs: int + batch_size: int + learning_rate: float + optimizer: str + scheduler: str + resolution: int + seed: int + status: str + training_time_seconds: int + final_loss: float | None + output_folder: str + checkpoint_paths: list[str] + preview_images: list[str] + peak_vram_gb: float | None + hardware: dict[str, Any] + settings: dict[str, Any] + generation_settings: dict[str, Any] + notes: str = "" + quality_score: int | None = None + + @classmethod + def from_row(cls, row: sqlite3.Row) -> "ExperimentRun": + payload = dict(row) + for key in ("checkpoint_paths", "preview_images"): + payload[key] = _safe_json(payload.get(key, "[]"), []) + for key in ("hardware", "settings", "generation_settings"): + payload[key] = _safe_json(payload.get(key, "{}"), {}) + payload["quality_score"] = ( + _safe_int(payload["quality_score"]) if payload.get("quality_score") is not None else None + ) + return cls(**payload) + + +class ExperimentStore: + def __init__(self, root: Path) -> None: + self.root = root.resolve() + self.path = self.root / "data" / "experiments.sqlite3" + self.path.parent.mkdir(parents=True, exist_ok=True) + self._init_db() + + def connect(self) -> sqlite3.Connection: + connection = sqlite3.connect(self.path) + connection.row_factory = sqlite3.Row + return connection + + def _init_db(self) -> None: + with self.connect() as db: + db.execute( + """ + CREATE TABLE IF NOT EXISTS experiments ( + id TEXT PRIMARY KEY, + job_id TEXT UNIQUE NOT NULL, + timestamp TEXT NOT NULL, + model_architecture TEXT NOT NULL, + model_name TEXT NOT NULL, + trigger_word TEXT NOT NULL DEFAULT '', + base_model TEXT NOT NULL, + dataset_path TEXT NOT NULL, + dataset_name TEXT NOT NULL, + dataset_item_count INTEGER NOT NULL, + epochs INTEGER NOT NULL, + batch_size INTEGER NOT NULL, + learning_rate REAL NOT NULL, + optimizer TEXT NOT NULL, + scheduler TEXT NOT NULL, + resolution INTEGER NOT NULL, + seed INTEGER NOT NULL, + status TEXT NOT NULL, + training_time_seconds INTEGER NOT NULL, + final_loss REAL, + output_folder TEXT NOT NULL, + checkpoint_paths TEXT NOT NULL, + preview_images TEXT NOT NULL, + peak_vram_gb REAL, + hardware TEXT NOT NULL, + settings TEXT NOT NULL, + generation_settings TEXT NOT NULL, + notes TEXT NOT NULL DEFAULT '', + quality_score INTEGER + ) + """ + ) + self._migrate_columns(db) + + @staticmethod + def _migrate_columns(db: sqlite3.Connection) -> None: + existing = {row["name"] for row in db.execute("PRAGMA table_info(experiments)").fetchall()} + columns = { + "id": "TEXT PRIMARY KEY", + "job_id": "TEXT NOT NULL DEFAULT ''", + "timestamp": "TEXT NOT NULL DEFAULT ''", + "model_architecture": "TEXT NOT NULL DEFAULT ''", + "model_name": "TEXT NOT NULL DEFAULT ''", + "trigger_word": "TEXT NOT NULL DEFAULT ''", + "base_model": "TEXT NOT NULL DEFAULT ''", + "dataset_path": "TEXT NOT NULL DEFAULT ''", + "dataset_name": "TEXT NOT NULL DEFAULT ''", + "dataset_item_count": "INTEGER NOT NULL DEFAULT 0", + "epochs": "INTEGER NOT NULL DEFAULT 0", + "batch_size": "INTEGER NOT NULL DEFAULT 0", + "learning_rate": "REAL NOT NULL DEFAULT 0", + "optimizer": "TEXT NOT NULL DEFAULT ''", + "scheduler": "TEXT NOT NULL DEFAULT ''", + "resolution": "INTEGER NOT NULL DEFAULT 0", + "seed": "INTEGER NOT NULL DEFAULT 0", + "status": "TEXT NOT NULL DEFAULT ''", + "training_time_seconds": "INTEGER NOT NULL DEFAULT 0", + "final_loss": "REAL", + "output_folder": "TEXT NOT NULL DEFAULT ''", + "checkpoint_paths": "TEXT NOT NULL DEFAULT '[]'", + "preview_images": "TEXT NOT NULL DEFAULT '[]'", + "peak_vram_gb": "REAL", + "hardware": "TEXT NOT NULL DEFAULT '{}'", + "settings": "TEXT NOT NULL DEFAULT '{}'", + "generation_settings": "TEXT NOT NULL DEFAULT '{}'", + "notes": "TEXT NOT NULL DEFAULT ''", + "quality_score": "INTEGER", + } + for name, definition in columns.items(): + if name not in existing and name != "id": + db.execute(f"ALTER TABLE experiments ADD COLUMN {name} {definition}") + + def record_job(self, job: Job, snapshot: SystemSnapshot | None = None) -> ExperimentRun | None: + training_steps = [step for step in job.plan.steps if step.tool_id.endswith("_trainer")] + if not training_steps: + return None + step = training_steps[-1] + args = dict(step.arguments) + architecture = step.tool_id.removesuffix("_trainer") + dataset_path = str(args.get("dataset_dir", "")) + output_folder = str(job.output_folder or args.get("output_dir", "")) + preview_images = [job.preview_path] if job.preview_path else [] + checkpoints = [] + if output_folder: + folder = Path(output_folder) + if folder.is_dir(): + try: + checkpoints = [ + str(path) + for path in sorted(folder.rglob("*")) + if path.is_file() and path.suffix.casefold() in {".safetensors", ".ckpt", ".pt", ".bin"} + ][-10:] + discovered_previews = [ + str(path) + for path in sorted(folder.rglob("*")) + if path.is_file() + and path.suffix.casefold() in {".png", ".jpg", ".jpeg", ".webp", ".bmp"} + and any(token in path.name.casefold() for token in ("preview", "sample", "epoch")) + ][-12:] + preview_images = list(dict.fromkeys([*preview_images, *discovered_previews])) + except OSError: + checkpoints = [] + hardware = {} + peak_vram = None + if snapshot is not None: + hardware = { + "gpu_name": snapshot.gpu_name, + "gpu_percent": snapshot.gpu_percent, + "vram_used_gb": snapshot.vram_used_gb, + "vram_total_gb": snapshot.vram_total_gb, + "memory_used_gb": snapshot.memory_used_gb, + "memory_total_gb": snapshot.memory_total_gb, + "cpu_percent": snapshot.cpu_percent, + "gpu_temperature": snapshot.gpu_temperature, + } + peak_vram = snapshot.vram_used_gb or None + run = ExperimentRun( + id=f"EXP-{job.id}", + job_id=job.id, + timestamp=job.ended_at or _utc_now(), + model_architecture=architecture, + model_name=str(args.get("model_name", job.plan.project_name)), + trigger_word=str(args.get("trigger_word", "")), + base_model=str(args.get("base_model", args.get("base_model_path", ""))), + dataset_path=dataset_path, + dataset_name=Path(dataset_path).name if dataset_path else "", + dataset_item_count=_image_count(dataset_path), + epochs=_safe_int(args.get("epochs")), + batch_size=_safe_int(args.get("batch_size")), + learning_rate=_safe_float(args.get("learning_rate")), + optimizer=str(args.get("optimizer", "")), + scheduler=str(args.get("scheduler", args.get("sampler", ""))), + resolution=_safe_int(args.get("resolution")), + seed=_safe_int(args.get("seed", args.get("preview_seed", 0))), + status=job.status.value, + training_time_seconds=_duration_seconds(job), + final_loss=_loss_from_logs(job.logs), + output_folder=output_folder, + checkpoint_paths=checkpoints, + preview_images=preview_images, + peak_vram_gb=peak_vram, + hardware=hardware, + settings=args, + generation_settings={}, + ) + with self.connect() as db: + db.execute( + """ + INSERT INTO experiments ( + id, job_id, timestamp, model_architecture, model_name, trigger_word, base_model, + dataset_path, dataset_name, dataset_item_count, epochs, batch_size, + learning_rate, optimizer, scheduler, resolution, seed, status, + training_time_seconds, final_loss, output_folder, checkpoint_paths, + preview_images, peak_vram_gb, hardware, settings, generation_settings + ) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) + ON CONFLICT(job_id) DO UPDATE SET + timestamp=excluded.timestamp, + trigger_word=excluded.trigger_word, + status=excluded.status, + training_time_seconds=excluded.training_time_seconds, + final_loss=excluded.final_loss, + output_folder=excluded.output_folder, + checkpoint_paths=excluded.checkpoint_paths, + preview_images=excluded.preview_images, + peak_vram_gb=excluded.peak_vram_gb, + hardware=excluded.hardware, + settings=excluded.settings + """, + ( + run.id, run.job_id, run.timestamp, run.model_architecture, run.model_name, + run.trigger_word, run.base_model, run.dataset_path, run.dataset_name, run.dataset_item_count, + run.epochs, run.batch_size, run.learning_rate, run.optimizer, run.scheduler, + run.resolution, run.seed, run.status, run.training_time_seconds, run.final_loss, + run.output_folder, _json(run.checkpoint_paths), _json(run.preview_images), + run.peak_vram_gb, _json(run.hardware), _json(run.settings), + _json(run.generation_settings), + ), + ) + return run + + def list_runs(self, search: str = "", architecture: str = "", dataset: str = "", limit: int = 200) -> list[ExperimentRun]: + clauses = [] + params: list[Any] = [] + if search: + 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 ?)") + term = f"%{search}%" + params.extend([term, term, term, term, term, term, term, term]) + if architecture: + clauses.append("model_architecture = ?") + params.append(architecture) + if dataset: + clauses.append("dataset_name LIKE ?") + params.append(f"%{dataset}%") + where = " WHERE " + " AND ".join(clauses) if clauses else "" + with self.connect() as db: + rows = db.execute( + "SELECT * FROM experiments" + where + " ORDER BY timestamp DESC LIMIT ?", + [*params, int(limit)], + ).fetchall() + return [ExperimentRun.from_row(row) for row in rows] + + def get(self, run_id: str) -> ExperimentRun | None: + with self.connect() as db: + row = db.execute("SELECT * FROM experiments WHERE id = ?", (run_id,)).fetchone() + return ExperimentRun.from_row(row) if row else None + + def update_notes(self, run_id: str, notes: str, quality_score: int | None) -> None: + with self.connect() as db: + db.execute( + "UPDATE experiments SET notes = ?, quality_score = ? WHERE id = ?", + (notes, quality_score, run_id), + ) + + def compare(self, run_ids: list[str]) -> list[dict[str, Any]]: + runs = [run for run_id in run_ids if (run := self.get(run_id)) is not None] + fields = [ + "model_architecture", "epochs", "final_loss", "training_time_seconds", + "resolution", "batch_size", "learning_rate", "scheduler", + "peak_vram_gb", "dataset_name", "quality_score", + ] + rows = [] + for field in fields: + values = {run.id: getattr(run, field) for run in runs} + comparable = {str(value) for value in values.values()} + rows.append({"field": field, "changed": len(comparable) > 1, **values}) + return rows + + def clone_request(self, run_id: str) -> str: + run = self.get(run_id) + if run is None: + return "" + options = { + key: value + for key, value in run.settings.items() + if key not in {"dataset_dir", "model_name", "epochs", "output_dir", "resume_from"} + } + return ( + f"From the {run.dataset_name or run.dataset_path} dataset, train a " + f"{run.model_architecture.upper()} model for {run.epochs} epochs. " + f"Name the model {run.model_name} Clone. " + "[ADAM_TRAINING_OPTIONS:" + json.dumps(options, sort_keys=True) + "] " + "[ADAM_TRAINER:" + run.model_architecture + "]" + ) diff --git a/adam/generations.py b/adam/generations.py index 1d102c64a69ecb987588a0c06e5e91a73c135957..9c3f51b08623e75bfe5409f4ed87a8c120b45213 100644 --- a/adam/generations.py +++ b/adam/generations.py @@ -219,6 +219,7 @@ def generation_metadata_path(folder: Path, timestamp: str, job_id: str) -> Path: @dataclass(frozen=True, slots=True) class GenerationRecord: + metadata_path: Path folder: Path images: tuple[Path, ...] provider_id: str @@ -231,6 +232,8 @@ class GenerationRecord: sampler: str aspect_ratio: str created_at: str + smart_generation: dict[str, Any] + image_evaluations: dict[str, dict[str, Any]] @classmethod def from_metadata(cls, metadata_path: Path) -> "GenerationRecord | None": @@ -246,6 +249,7 @@ class GenerationRecord: if not images: return None return cls( + metadata_path=metadata_path, folder=folder, images=images, provider_id=str(payload.get("provider_id", "")), @@ -258,9 +262,77 @@ class GenerationRecord: sampler=str(payload.get("sampler", "")), aspect_ratio=str(payload.get("aspect_ratio", "")), created_at=str(payload.get("created_at", "")), + smart_generation=dict(payload.get("smart_generation") or {}), + image_evaluations={ + str(Path(path).resolve()): dict(value) + for path, value in dict(payload.get("image_evaluations") or {}).items() + if isinstance(value, dict) + }, ) +@dataclass(frozen=True, slots=True) +class GenerationModelFolder: + """A model-centered view over existing generation batches.""" + + key: str + model_name: str + model_path: str + provider_id: str + provider_name: str + records: tuple[GenerationRecord, ...] + image_count: int + cover_image: Path | None + latest_at: str + + +def generation_model_key(record: GenerationRecord) -> str: + """Keep renamed or duplicated display names separated by model identity.""" + raw_path = str(record.model_path or "").strip() + if raw_path: + try: + return f"path:{Path(raw_path).expanduser().resolve()}".casefold() + except OSError: + return f"path:{raw_path}".casefold() + return f"name:{record.provider_id}:{record.model_name}".casefold() + + +def group_generation_records( + records: list[GenerationRecord], +) -> list[GenerationModelFolder]: + """Build newest-first automatic model folders without changing files.""" + grouped: dict[str, list[GenerationRecord]] = {} + for record in records: + grouped.setdefault(generation_model_key(record), []).append(record) + folders: list[GenerationModelFolder] = [] + for key, model_records in grouped.items(): + newest_first = sorted( + model_records, + key=lambda item: item.created_at or item.folder.name, + reverse=True, + ) + latest = newest_first[0] + cover = next( + (path for record in newest_first for path in record.images if path.is_file()), + None, + ) + folders.append( + GenerationModelFolder( + key=key, + model_name=latest.model_name, + model_path=latest.model_path, + provider_id=latest.provider_id, + provider_name=latest.provider_name, + records=tuple(newest_first), + image_count=sum(len(record.images) for record in newest_first), + cover_image=cover, + latest_at=latest.created_at, + ) + ) + folders.sort(key=lambda item: (item.latest_at, item.model_name.casefold()), reverse=True) + return folders + + def generation_tools(registry: ToolRegistry) -> list[ToolSpec]: return [ tool diff --git a/adam/image_preferences.py b/adam/image_preferences.py new file mode 100644 index 0000000000000000000000000000000000000000..ca04363767abac339b7771c5ff392d19f8bde190 --- /dev/null +++ b/adam/image_preferences.py @@ -0,0 +1,313 @@ +from __future__ import annotations + +from dataclasses import asdict, dataclass, field +from datetime import datetime, timezone +import hashlib +import json +from pathlib import Path +from typing import Any, Callable, Sequence + +from adam.eve import EveVisionModel, classify_eve_embeddings + + +RATINGS = {"favorite", "keep", "unsure", "reject"} +POSITIVE_RATINGS = {"favorite", "keep"} +NEGATIVE_RATINGS = {"reject"} + + +def _now() -> str: + return datetime.now(timezone.utc).isoformat() + + +def preference_profile_id(provider_id: str, model_path: str) -> str: + key = f"{provider_id}\n{str(Path(model_path).expanduser().resolve())}" + return hashlib.sha1(key.encode("utf-8")).hexdigest()[:16] + + +def image_cache_id(path: str | Path) -> str: + resolved = str(Path(path).expanduser().resolve()) + return hashlib.sha1(resolved.encode("utf-8")).hexdigest() + + +@dataclass(slots=True) +class GenerationRating: + image_path: str + rating: str + provider_id: str + model_name: str + model_path: str + seed: int = 0 + sampler: str = "" + steps: int = 0 + resolution: str = "" + generation_settings: dict[str, Any] = field(default_factory=dict) + generation_created_at: str = "" + rated_at: str = field(default_factory=_now) + embedding: list[float] | None = None + + @classmethod + def from_dict(cls, payload: dict[str, Any]) -> "GenerationRating": + rating = str(payload.get("rating", "unsure")).casefold() + return cls( + image_path=str(Path(str(payload.get("image_path", ""))).expanduser().resolve()), + rating=rating if rating in RATINGS else "unsure", + provider_id=str(payload.get("provider_id", "")), + model_name=str(payload.get("model_name", "")), + model_path=str(payload.get("model_path", "")), + seed=int(payload.get("seed", 0) or 0), + sampler=str(payload.get("sampler", "")), + steps=int(payload.get("steps", 0) or 0), + resolution=str(payload.get("resolution", "")), + generation_settings=dict(payload.get("generation_settings") or {}), + generation_created_at=str(payload.get("generation_created_at", "")), + rated_at=str(payload.get("rated_at") or _now()), + embedding=[float(value) for value in payload["embedding"]] + if isinstance(payload.get("embedding"), list) + else None, + ) + + +@dataclass(frozen=True, slots=True) +class PreferenceScore: + image_path: str + score: float | None + confidence: float + category: str + reason: str = "" + + +class PreferenceProfile: + def __init__(self, root: Path, provider_id: str, model_name: str, model_path: str) -> None: + self.root = root.resolve() + self.provider_id = provider_id + self.model_name = model_name + self.model_path = str(Path(model_path).expanduser().resolve()) if model_path else "" + self.id = preference_profile_id(provider_id, self.model_path) + self.path = self.root / "data" / "generation_preferences" / f"{self.id}.json" + self.keep_threshold = 0.70 + self.reject_threshold = 0.35 + self.ratings: dict[str, GenerationRating] = {} + self.load() + + def load(self) -> None: + try: + payload = json.loads(self.path.read_text(encoding="utf-8")) + except (OSError, ValueError, TypeError, json.JSONDecodeError): + return + self.model_name = str(payload.get("model_name") or self.model_name) + self.provider_id = str(payload.get("provider_id") or self.provider_id) + self.model_path = str(payload.get("model_path") or self.model_path) + thresholds = payload.get("thresholds", {}) + if isinstance(thresholds, dict): + self.keep_threshold = float(thresholds.get("keep", self.keep_threshold)) + self.reject_threshold = float(thresholds.get("reject", self.reject_threshold)) + ratings = payload.get("ratings", []) + if isinstance(ratings, list): + for item in ratings: + if isinstance(item, dict): + rating = GenerationRating.from_dict(item) + self.ratings[rating.image_path] = rating + + def save(self) -> None: + self.path.parent.mkdir(parents=True, exist_ok=True) + temporary = self.path.with_suffix(".tmp") + temporary.write_text( + json.dumps( + { + "version": 1, + "profile_id": self.id, + "provider_id": self.provider_id, + "model_name": self.model_name, + "model_path": self.model_path, + "thresholds": { + "keep": self.keep_threshold, + "reject": self.reject_threshold, + }, + "ratings": [asdict(item) for item in self.ratings.values()], + "updated_at": _now(), + }, + indent=2, + ), + encoding="utf-8", + ) + temporary.replace(self.path) + + def set_rating( + self, + image_path: str | Path, + rating: str, + *, + seed: int = 0, + sampler: str = "", + steps: int = 0, + resolution: str = "", + generation_settings: dict[str, Any] | None = None, + generation_created_at: str = "", + embedding: Sequence[float] | None = None, + ) -> GenerationRating: + clean = rating.casefold().strip() + if clean not in RATINGS: + raise ValueError("Generation rating must be Favorite, Keep, Unsure, or Reject.") + resolved = str(Path(image_path).expanduser().resolve()) + existing = self.ratings.get(resolved) + record = GenerationRating( + image_path=resolved, + rating=clean, + provider_id=self.provider_id, + model_name=self.model_name, + model_path=self.model_path, + seed=int(seed), + sampler=sampler, + steps=int(steps), + resolution=resolution, + generation_settings=dict(generation_settings or {}), + generation_created_at=generation_created_at, + rated_at=_now(), + embedding=[float(value) for value in embedding] if embedding is not None else ( + existing.embedding if existing else None + ), + ) + self.ratings[resolved] = record + self.save() + return record + + def rating_for(self, image_path: str | Path) -> GenerationRating | None: + return self.ratings.get(str(Path(image_path).expanduser().resolve())) + + def examples(self) -> tuple[list[GenerationRating], list[GenerationRating]]: + positive = [ + item for item in self.ratings.values() + if item.rating in POSITIVE_RATINGS and Path(item.image_path).is_file() + ] + negative = [ + item for item in self.ratings.values() + if item.rating in NEGATIVE_RATINGS and Path(item.image_path).is_file() + ] + return positive, negative + + def has_signal(self) -> bool: + positive, _negative = self.examples() + return bool(positive) + + +class ImageEmbeddingCache: + def __init__(self, root: Path, model_id: str) -> None: + self.root = root.resolve() + self.model_id = model_id + self.folder = self.root / "data" / "image_embeddings" / hashlib.sha1(model_id.encode("utf-8")).hexdigest()[:12] + + def get(self, path: str | Path) -> list[float] | None: + cache_path = self.folder / f"{image_cache_id(path)}.json" + try: + payload = json.loads(cache_path.read_text(encoding="utf-8")) + except (OSError, ValueError, TypeError, json.JSONDecodeError): + return None + source = Path(path).expanduser().resolve() + try: + stat = source.stat() + except OSError: + return None + if payload.get("path") != str(source) or payload.get("mtime") != stat.st_mtime: + return None + vector = payload.get("embedding") + return [float(value) for value in vector] if isinstance(vector, list) else None + + def set(self, path: str | Path, embedding: Sequence[float]) -> None: + source = Path(path).expanduser().resolve() + try: + stat = source.stat() + except OSError: + return + self.folder.mkdir(parents=True, exist_ok=True) + cache_path = self.folder / f"{image_cache_id(source)}.json" + temporary = cache_path.with_suffix(".tmp") + temporary.write_text( + json.dumps( + { + "path": str(source), + "mtime": stat.st_mtime, + "model_id": self.model_id, + "embedding": [float(value) for value in embedding], + } + ), + encoding="utf-8", + ) + temporary.replace(cache_path) + + +class GenerationPreferenceEvaluator: + """Shared EVE-backed scorer for generated images.""" + + def __init__( + self, + root: Path, + vision: EveVisionModel | None = None, + *, + embedder: Callable[[Sequence[str | Path]], list[list[float]]] | None = None, + ) -> None: + self.root = root.resolve() + self.vision = vision or EveVisionModel(prefer_gpu=False) + self.embedder = embedder + self.cache = ImageEmbeddingCache(self.root, self.vision.model_id) + + def _embedding(self, path: str | Path) -> list[float]: + cached = self.cache.get(path) + if cached is not None: + return cached + vectors = self.embedder([path]) if self.embedder else self.vision.embed([path]) + vector = [float(value) for value in vectors[0]] + self.cache.set(path, vector) + return vector + + def score( + self, + profile: PreferenceProfile, + image_paths: Sequence[str | Path], + *, + keep_threshold: float | None = None, + reject_threshold: float | None = None, + ) -> list[PreferenceScore]: + positive, negative = profile.examples() + if not positive: + return [ + PreferenceScore(str(Path(path).expanduser().resolve()), None, 0.0, "Needs Review", "No preference examples yet") + for path in image_paths + ] + positive_vectors = [item.embedding or self._embedding(item.image_path) for item in positive] + negative_vectors = [item.embedding or self._embedding(item.image_path) for item in negative] + image_vectors = [self._embedding(path) for path in image_paths] + keep = max(0.001, min(1.0, float(keep_threshold if keep_threshold is not None else profile.keep_threshold))) + reject = max(0.0, min(float(reject_threshold if reject_threshold is not None else profile.reject_threshold), keep - 0.001)) + results = classify_eve_embeddings( + image_paths, + image_vectors, + positive_vectors, + negative_vectors, + keep_threshold=keep, + reject_threshold=reject, + ) + categories = {"keep": "Strong Keep", "reject": "Likely Reject", "unreviewed": "Needs Review"} + return [ + PreferenceScore(result.path, result.match_score, result.decision_confidence, categories[result.suggestion]) + for result in results + ] + + +def score_generated_images( + root: Path, + *, + provider_id: str, + model_name: str, + model_path: str, + image_paths: Sequence[str | Path], + keep_threshold: float | None = None, + reject_threshold: float | None = None, +) -> list[PreferenceScore]: + profile = PreferenceProfile(root, provider_id, model_name, model_path) + evaluator = GenerationPreferenceEvaluator(root) + return evaluator.score( + profile, + image_paths, + keep_threshold=keep_threshold, + reject_threshold=reject_threshold, + ) diff --git a/adam/job_manager.py b/adam/job_manager.py index 450f524d10a9b71717cef5be8397ab827e100945..8dff72d1ea6d8a071704f657fd79ff56de289c07 100644 --- a/adam/job_manager.py +++ b/adam/job_manager.py @@ -5,17 +5,27 @@ import logging import math import re import threading -from datetime import datetime, timezone +import time +from datetime import datetime, timedelta, timezone from pathlib import Path from typing import Any -from PySide6.QtCore import QObject, QThread, Signal +from PySide6.QtCore import QCoreApplication, QObject, QThread, QTimer, Signal -from adam.executor import ToolCancelled, ToolExecutionError, ToolExecutor +from adam.executor import ToolAdjustmentRequested, ToolCancelled, ToolExecutionError, ToolExecutor from adam.assets import AssetRegistry from adam.atlas import AtlasSupervisor +from adam.experiment_tracker import ExperimentStore from adam.models import ExecutionPlan, Job, JobStatus, StepStatus, utc_now from adam.nova import evaluate_job_output +from adam.training_assistant import append_preflight_summary + + +def _safe_int(value: Any) -> int: + try: + return max(0, int(value or 0)) + except (TypeError, ValueError): + return 0 class JobWorker(QThread): @@ -28,6 +38,8 @@ class JobWorker(QThread): self.cancel_event = threading.Event() self.run_event = threading.Event() self.run_event.set() + self.adjustment_event = threading.Event() + self.adjustment_request: dict[str, Any] = {} def pause(self) -> None: self.run_event.clear() @@ -39,6 +51,12 @@ class JobWorker(QThread): self.cancel_event.set() self.run_event.set() + def request_adjustment(self, updates: dict[str, Any]) -> None: + self.adjustment_request.clear() + self.adjustment_request.update(updates) + self.adjustment_event.set() + self.run_event.set() + def run(self) -> None: total_steps = len(self.job.plan.steps) try: @@ -54,15 +72,25 @@ class JobWorker(QThread): ) preview_state = {"epoch": 0, "path": ""} + last_progress_emit = {"time": 0.0, "overall": -1, "message": ""} + progress_samples: list[dict[str, float]] = [] - def on_progress(percent: int, message: str, step_index: int = index) -> None: + def on_progress(percent: int, message: str, step_index: int = index, **details: Any) -> None: overall = int(((step_index + percent / 100) / total_steps) * 100) + now = time.monotonic() + progress_eta = self._estimate_step_eta(details, progress_samples, now) + changed = overall != last_progress_emit["overall"] or message != last_progress_emit["message"] + terminal = percent >= 100 or overall >= 100 + if not terminal and (not changed or now - last_progress_emit["time"] < 0.25): + return + last_progress_emit.update({"time": now, "overall": overall, "message": message}) self.event.emit( { "type": "progress", "step_percent": percent, "overall": overall, "message": message, + **progress_eta, } ) self._discover_external_preview(step, message, preview_state) @@ -84,6 +112,8 @@ class JobWorker(QThread): progress_callback=on_progress, log_callback=on_log, preview_callback=on_preview, + adjustment_event=self.adjustment_event, + adjustment_request=self.adjustment_request, ) self.event.emit( { @@ -95,6 +125,8 @@ class JobWorker(QThread): self.event.emit({"type": "completed"}) except ToolCancelled as exc: self.event.emit({"type": "cancelled", "message": str(exc)}) + except ToolAdjustmentRequested as exc: + self.event.emit({"type": "adjustment_ready", "message": str(exc), **exc.details}) except Exception as exc: self.event.emit( { @@ -140,6 +172,61 @@ class JobWorker(QThread): "steps": int(step.arguments.get("preview_steps", 0) or 0), }) + @staticmethod + def _estimate_step_eta( + details: dict[str, Any], + samples: list[dict[str, float]], + now: float, + ) -> dict[str, Any]: + """Estimate remaining runtime from real step cadence instead of percent alone.""" + current = _safe_int(details.get("current_step", details.get("step", details.get("current")))) + total = _safe_int(details.get("total_steps", details.get("total"))) + unit = str(details.get("unit", "step") or "step") + epoch = _safe_int(details.get("epoch")) + total_epochs = _safe_int(details.get("total_epochs")) + if (not current or not total) and epoch and total_epochs: + current, total, unit = epoch, total_epochs, "epoch" + payload: dict[str, Any] = { + "progress_current": current, + "progress_total": total, + "progress_unit": unit, + } + if not current or not total or current >= total: + return payload + last = samples[-1] if samples else None + if last and current <= last["current"]: + return payload + samples.append({"time": now, "current": float(current)}) + del samples[:-25] + if len(samples) < 2: + return payload + + first = samples[0] + elapsed = now - first["time"] + completed = current - int(first["current"]) + if completed <= 0 or elapsed <= 0: + return payload + lifetime_seconds_per_unit = elapsed / completed + + recent = samples[-8:] + recent_first = recent[0] + recent_completed = current - int(recent_first["current"]) + recent_elapsed = now - recent_first["time"] + recent_seconds_per_unit = ( + recent_elapsed / recent_completed + if recent_completed > 0 and recent_elapsed > 0 + else lifetime_seconds_per_unit + ) + seconds_per_unit = (recent_seconds_per_unit * 0.65) + (lifetime_seconds_per_unit * 0.35) + remaining = max(0, int(round((total - current) * seconds_per_unit))) + if remaining: + payload["eta_seconds"] = remaining + payload["progress_rate"] = 1 / seconds_per_unit if seconds_per_unit > 0 else 0.0 + payload["estimated_completion_at"] = ( + datetime.now(timezone.utc) + timedelta(seconds=remaining) + ).isoformat() + return payload + class JobManager(QObject): job_created = Signal(object) @@ -159,30 +246,57 @@ class JobManager(QObject): self.root = root self.executor = executor self.logger = logger + self.config = config if config is not None else {} self.jobs_path = root / "data" / "jobs.json" self.assets = AssetRegistry(root) self.atlas = AtlasSupervisor(config) + self.experiments = ExperimentStore(root) self.jobs: list[Job] = [] self._queue: list[str] = [] self._worker: JobWorker | None = None self._active_job: Job | None = None + self._last_snapshot = None + self._pending_update_job_ids: set[str] = set() + self._pending_update_timer_active = False self._load() + self._schedule_timer = QTimer(self) + self._schedule_timer.timeout.connect(self._release_due_scheduled) + if QCoreApplication.instance() is not None: + self._schedule_timer.start(15_000) + QTimer.singleShot(0, self._release_due_scheduled) + if self._queue: + QTimer.singleShot(0, self._start_next) @property def active_job(self) -> Job | None: return self._active_job - def submit(self, plan: ExecutionPlan) -> Job: + @property + def pending_jobs(self) -> list[Job]: + return [ + job for job in self.jobs + if job.status in {JobStatus.AWAITING_CONFIRMATION, JobStatus.SCHEDULED, JobStatus.QUEUED} + ] + + def submit(self, plan: ExecutionPlan, scheduled_for: str | None = None) -> Job: + # Every entry point must review a plan before its queue state is chosen. + # UI and Remote may prepare it earlier to keep filesystem work off Qt. + append_preflight_summary(plan, self.config) + is_future = self._is_future(scheduled_for) status = ( JobStatus.AWAITING_CONFIRMATION if plan.requires_confirmation - else JobStatus.QUEUED + else JobStatus.SCHEDULED if is_future else JobStatus.QUEUED ) - job = Job(plan=plan, status=status) + job = Job(plan=plan, status=status, scheduled_for=scheduled_for if is_future else None) self.jobs.insert(0, job) self._append_log(job, f"Plan created: {plan.summary}") if plan.requires_confirmation: self._append_log(job, "Waiting for user confirmation.") + if is_future: + self._append_log(job, f"Requested start time: {self._display_time(scheduled_for)}.") + elif is_future: + self._append_log(job, f"Scheduled for {self._display_time(scheduled_for)}.") else: self._queue.append(job.id) self._save() @@ -196,9 +310,12 @@ class JobManager(QObject): job = self.get(job_id) if job.status != JobStatus.AWAITING_CONFIRMATION: return - job.status = JobStatus.QUEUED + job.status = JobStatus.SCHEDULED if self._is_future(job.scheduled_for) else JobStatus.QUEUED self._append_log(job, "Plan approved by user.") - self._queue.append(job.id) + if job.status == JobStatus.SCHEDULED: + self._append_log(job, f"Training will become eligible at {self._display_time(job.scheduled_for)}.") + else: + self._queue.append(job.id) self._save() self.job_updated.emit(job) self._start_next() @@ -236,10 +353,13 @@ class JobManager(QObject): if job is self._active_job and self._worker: self._append_log(job, "Cancellation requested.") self._worker.cancel() + self._save() + self.job_updated.emit(job) return if job.id in self._queue: self._queue.remove(job.id) if job.status in { + JobStatus.SCHEDULED, JobStatus.QUEUED, JobStatus.AWAITING_CONFIRMATION, JobStatus.DRAFT, @@ -250,6 +370,104 @@ class JobManager(QObject): self._save() self.job_updated.emit(job) + def request_training_adjustment(self, job_id: str, updates: dict[str, Any]) -> None: + """Apply safe DDPM settings after the current epoch and resume automatically.""" + job = self.get(job_id) + if job is not self._active_job or job.status not in {JobStatus.RUNNING, JobStatus.PAUSED} or not self._worker: + raise ValueError("Only the active training job can be adjusted.") + if not (0 <= job.current_step < len(job.plan.steps)): + raise ValueError("The active training step is unavailable.") + step = job.plan.steps[job.current_step] + if step.tool_id != "ddpm_trainer": + raise ValueError("Safe epoch-boundary adjustment currently supports DDPM training.") + allowed = {"batch_size", "training_intensity", "gradient_accumulation_steps"} + cleaned = {key: int(value) for key, value in updates.items() if key in allowed} + if not cleaned or not 1 <= cleaned.get("batch_size", 1) <= 64 \ + or not 10 <= cleaned.get("training_intensity", 100) <= 100 \ + or not 1 <= cleaned.get("gradient_accumulation_steps", 1) <= 64: + raise ValueError("The requested training settings are outside ADAM's safe range.") + previous = {key: step.arguments.get(key) for key in cleaned} + if all(previous[key] == value for key, value in cleaned.items()): + raise ValueError("Those settings are already active.") + job.status = JobStatus.RUNNING + self._append_log(job, f"Adjustment queued for the end of this epoch: {cleaned}.") + self._worker.request_adjustment(cleaned) + self._save() + self.job_updated.emit(job) + + def safer_vram_retry(self, job_id: str) -> Job: + """Create a checkpoint-aware DDPM retry with a smaller physical batch.""" + original = self.get(job_id) + if original.status != JobStatus.FAILED or not self._looks_like_vram_failure(original): + raise ValueError("This job did not fail with a recognizable VRAM error.") + plan = ExecutionPlan.from_dict(original.to_dict()["plan"]) + start_index = max(0, min(original.current_step, len(plan.steps) - 1)) + plan.steps = plan.steps[start_index:] + step = plan.steps[0] + old_batch = max(1, int(step.arguments.get("batch_size", 1))) + if old_batch <= 1: + raise ValueError("Batch size is already 1; lower resolution or enable other memory-saving options.") + original_epochs = max(1, int(step.arguments.get("epochs", 1))) + resume_note = self._prepare_ddpm_resume(step.arguments, step.tool_id) + remaining_epochs = max(1, int(step.arguments.get("epochs", original_epochs))) + completed_epochs = max(0, original_epochs - remaining_epochs) if resume_note else 0 + new_batch = max(1, old_batch // 2) + old_accumulation = max(1, int(step.arguments.get("gradient_accumulation_steps", 1))) + step.arguments["batch_size"] = new_batch + step.arguments["gradient_accumulation_steps"] = min(64, old_accumulation * max(1, math.ceil(old_batch / new_batch))) + if completed_epochs: + step.arguments["completed_epochs"] = completed_epochs + for item in plan.steps: + item.status = StepStatus.PENDING + plan.id = original.plan.id + "-vram-retry" + plan.created_at = utc_now() + plan.requires_confirmation = True + plan.confirmation_reason = "VRAM recovery reduced the physical batch and preserved the effective batch with gradient accumulation." + retry = self.submit(plan) + self._append_log(retry, f"VRAM recovery changed batch {old_batch} → {new_batch} and gradient accumulation {old_accumulation} → {step.arguments['gradient_accumulation_steps']}.") + if resume_note: + self._append_log(retry, resume_note) + return retry + + @staticmethod + def _looks_like_vram_failure(job: Job) -> bool: + text = "\n".join([job.error or "", *job.logs[-100:]]).lower() + return any(token in text for token in ("out of memory", "cuda oom", "cuda error: out of memory")) + + @staticmethod + def _is_future(value: str | None) -> bool: + if not value: + return False + try: + scheduled = datetime.fromisoformat(value) + if scheduled.tzinfo is None: + scheduled = scheduled.astimezone() + return scheduled.astimezone(timezone.utc) > datetime.now(timezone.utc) + except (TypeError, ValueError): + return False + + @staticmethod + def _display_time(value: str | None) -> str: + try: + return datetime.fromisoformat(str(value)).astimezone().strftime("%b %d at %I:%M %p") + except ValueError: + return str(value or "the requested time") + + def _release_due_scheduled(self) -> None: + released: list[Job] = [] + for job in reversed(self.jobs): + if job.status == JobStatus.SCHEDULED and not self._is_future(job.scheduled_for): + job.status = JobStatus.QUEUED + self._queue.append(job.id) + self._append_log(job, "Scheduled start time reached; waiting for the training slot.") + released.append(job) + if not released: + return + self._save() + for job in released: + self.job_updated.emit(job) + self._start_next() + def get(self, job_id: str) -> Job: for job in self.jobs: if job.id == job_id: @@ -383,14 +601,30 @@ class JobManager(QObject): job.preview_total = 0 job.preview_image_index = 0 job.preview_image_count = 0 + job.eta_seconds = None + job.estimated_completion_at = None + job.progress_current = 0 + job.progress_total = 0 + job.progress_rate = 0.0 + job.progress_unit = "step" self._append_log(job, f"Starting: {job.plan.steps[index].title}") elif event_type == "progress": job.progress = int(event["overall"]) + job.eta_seconds = _safe_int(event.get("eta_seconds")) or None + job.estimated_completion_at = str(event.get("estimated_completion_at") or "") or None + job.progress_current = _safe_int(event.get("progress_current")) + job.progress_total = _safe_int(event.get("progress_total")) + job.progress_rate = float(event.get("progress_rate", 0.0) or 0.0) + job.progress_unit = str(event.get("progress_unit", "step") or "step") message = str(event["message"]) if message and (not job.logs or message not in job.logs[-1]): self._append_log(job, message) + self._schedule_job_update(job) + return elif event_type == "log": self._append_log(job, str(event["message"])) + self._schedule_job_update(job) + return elif event_type == "preview": job.preview_path = str(event.get("path", "")) or None job.preview_epoch = int(event.get("epoch", 0) or 0) @@ -407,6 +641,8 @@ class JobManager(QObject): label = "Denoising" if job.preview_kind == "generation" else "Training" position = f" step {job.preview_current}" if job.preview_current else f" epoch {job.preview_epoch}" self._append_log(job, f"{label} preview updated at{position}.") + self._schedule_job_update(job) + return elif event_type == "step_finished": index = int(event["index"]) job.plan.steps[index].status = StepStatus.FINISHED @@ -432,8 +668,11 @@ class JobManager(QObject): elif event_type == "completed": job.status = JobStatus.FINISHED job.progress = 100 + job.eta_seconds = 0 + job.estimated_completion_at = utc_now() job.ended_at = utc_now() self._append_log(job, "Job finished successfully.") + self._record_experiment(job) self.notification.emit("Job complete", job.plan.project_name) elif event_type == "cancelled": job.status = JobStatus.CANCELLED @@ -442,7 +681,31 @@ class JobManager(QObject): for step in job.plan.steps: if step.status == StepStatus.RUNNING: step.status = StepStatus.SKIPPED + self._record_experiment(job) self.notification.emit("Job cancelled", job.plan.project_name) + elif event_type == "adjustment_ready": + index = max(0, min(job.current_step, len(job.plan.steps) - 1)) + remaining_steps = job.plan.steps[index:] + step = remaining_steps[0] + updates = dict(event.get("updates") or {}) + step.arguments.update(updates) + checkpoint = str(event.get("checkpoint", "")) + completed_epochs = max(0, int(event.get("completed_epochs", 0) or 0)) + total_epochs = max(1, int(step.arguments.get("epochs", 1))) + step.arguments["epochs"] = max(1, total_epochs - completed_epochs) + step.arguments["completed_epochs"] = completed_epochs + if checkpoint: + step.arguments["resume_from"] = checkpoint + for pending in remaining_steps: + pending.status = StepStatus.PENDING + job.plan.steps = remaining_steps + job.current_step = -1 + job.status = JobStatus.QUEUED + job.eta_seconds = None + job.estimated_completion_at = None + self._queue.insert(0, job.id) + self._append_log(job, f"Epoch {completed_epochs} checkpoint is complete. Restarting with {updates}.") + self.notification.emit("Training settings ready", "Restarting from the completed epoch checkpoint.") elif event_type == "failed": job.status = JobStatus.FAILED job.ended_at = utc_now() @@ -456,10 +719,35 @@ class JobManager(QObject): event.get("exception"), job.error, ) + self._record_experiment(job) self.notification.emit("Job failed", job.error) self._save() self.job_updated.emit(job) + def _schedule_job_update(self, job: Job) -> None: + self._pending_update_job_ids.add(job.id) + if QCoreApplication.instance() is None: + self._flush_pending_job_updates() + return + if self._pending_update_timer_active: + return + self._pending_update_timer_active = True + QTimer.singleShot(300, self._flush_pending_job_updates) + + def _flush_pending_job_updates(self) -> None: + if not self._pending_update_job_ids: + self._pending_update_timer_active = False + return + pending_ids = list(self._pending_update_job_ids) + self._pending_update_job_ids.clear() + self._pending_update_timer_active = False + self._save() + for job_id in pending_ids: + try: + self.job_updated.emit(self.get(job_id)) + except KeyError: + continue + def _worker_finished(self) -> None: self._worker = None self._active_job = None @@ -468,6 +756,7 @@ class JobManager(QObject): def supervise(self, snapshot: Any) -> None: """Let ATLAS inspect the active training run and apply critical pauses.""" + self._last_snapshot = snapshot job = self._active_job if job is None or job.status != JobStatus.RUNNING: return @@ -495,6 +784,14 @@ class JobManager(QObject): if decision.action == "pause" and job.status == JobStatus.RUNNING: self.pause(job.id) + def _record_experiment(self, job: Job) -> None: + if not self._has_training(job): + return + try: + self.experiments.record_job(job, self._last_snapshot) + except Exception as exc: + self.logger.warning("Experiment tracking failed for %s: %s", job.id, exc) + def _append_log(self, job: Job, message: str) -> None: timestamp = datetime.now().strftime("%H:%M:%S") line = f"[{timestamp}] {message}" @@ -517,13 +814,17 @@ class JobManager(QObject): self.jobs = [] return for job in self.jobs: - if job.status in {JobStatus.RUNNING, JobStatus.PAUSED, JobStatus.QUEUED}: + if job.status in {JobStatus.RUNNING, JobStatus.PAUSED}: job.status = JobStatus.INTERRUPTED job.ended_at = utc_now() job.logs.append( "[startup] Previous session ended before this job. " "Review it before retrying." ) + elif job.status == JobStatus.QUEUED: + self._queue.append(job.id) + if not any("Queued job restored" in line for line in job.logs[-5:]): + job.logs.append("[startup] Queued job restored and will run when ADAM is ready.") self._save() def _save(self) -> None: diff --git a/adam/model_inspector/__init__.py b/adam/model_inspector/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..95bef929d41338fcd1160dc5576523b1c2930437 --- /dev/null +++ b/adam/model_inspector/__init__.py @@ -0,0 +1,15 @@ +from __future__ import annotations + +from .base import ModelComparison, ModelInspection, TensorComparison, TensorStats +from .comparison import compare_models +from .detector import inspect_model, inspector_for + +__all__ = [ + "ModelComparison", + "ModelInspection", + "TensorComparison", + "TensorStats", + "compare_models", + "inspect_model", + "inspector_for", +] diff --git a/adam/model_inspector/base.py b/adam/model_inspector/base.py new file mode 100644 index 0000000000000000000000000000000000000000..8b095c12090e7a0550fef81c2e44961204f166f8 --- /dev/null +++ b/adam/model_inspector/base.py @@ -0,0 +1,146 @@ +from __future__ import annotations + +from dataclasses import dataclass, field +from pathlib import Path +from typing import Any, Callable, Iterable + + +MODEL_EXTENSIONS = {".safetensors", ".pt", ".pth", ".bin", ".ckpt"} +CONFIG_FILENAMES = { + "model_index.json", + "config.json", + "scheduler_config.json", + "flow_model_info.json", + "adapter_config.json", +} + + +@dataclass(slots=True) +class TensorStats: + name: str + shape: tuple[int, ...] + dtype: str + parameter_count: int + memory_bytes: int + minimum: float | None = None + maximum: float | None = None + mean: float | None = None + std: float | None = None + abs_mean: float | None = None + l2_norm: float | None = None + zero_percent: float | None = None + component: str = "other" + health: list[str] = field(default_factory=list) + + +@dataclass(slots=True) +class ModelInspection: + path: str + resolved_path: str + architecture: str + confidence: float + status: str + size_bytes: int + config_files: list[str] + resolution: int | None + epoch: int | None + step: int | None + tensor_count: int + total_parameters: int + trainable_parameters: int | None + parameter_memory_bytes: int + dtypes: dict[str, int] + components: dict[str, int] + largest_tensors: list[TensorStats] + tensors: list[TensorStats] + health: list[str] + messages: list[str] + lora: dict[str, Any] = field(default_factory=dict) + configs: dict[str, Any] = field(default_factory=dict) + histogram: dict[str, list[float]] = field(default_factory=dict) + tensor_size_distribution: list[tuple[str, int]] = field(default_factory=list) + checkpoints: list[str] = field(default_factory=list) + loss_history: list[tuple[int, float]] = field(default_factory=list) + + +@dataclass(slots=True) +class TensorComparison: + name: str + shape: tuple[int, ...] + component: str + mean_abs_difference: float | None + relative_difference: float | None + cosine_similarity: float | None + l2_distance: float | None + drift: float | None + change_score: float | None + + +@dataclass(slots=True) +class ModelComparison: + path_a: str + path_b: str + architecture_a: str + architecture_b: str + architecture_match: bool + config_differences: list[str] + resolution_difference: tuple[int | None, int | None] | None + parameter_count_difference: int + only_a: list[str] + only_b: list[str] + shape_mismatches: list[str] + tensor_comparisons: list[TensorComparison] + group_comparisons: dict[str, dict[str, float]] + messages: list[str] + + +ProgressCallback = Callable[[int, str], None] +CancelCallback = Callable[[], bool] + + +class InspectorError(RuntimeError): + pass + + +class BaseModelInspector: + architecture = "Generic / Unknown" + + def inspect( + self, + path: str | Path, + *, + recorded_architecture: str = "", + run_settings: dict[str, Any] | None = None, + progress: ProgressCallback | None = None, + cancelled: CancelCallback | None = None, + ) -> ModelInspection: + raise NotImplementedError + + +def report(progress: ProgressCallback | None, value: int, message: str) -> None: + if progress: + progress(max(0, min(100, int(value))), message) + + +def is_cancelled(cancelled: CancelCallback | None) -> bool: + return bool(cancelled and cancelled()) + + +def parameter_count(shape: Iterable[int]) -> int: + total = 1 + for dim in shape: + total *= int(dim) + return int(total) + + +def dtype_size(dtype: str) -> int: + lowered = dtype.casefold() + if "float64" in lowered or "int64" in lowered: + return 8 + if "float32" in lowered or "int32" in lowered: + return 4 + if "float16" in lowered or "bfloat16" in lowered or "int16" in lowered: + return 2 + if "int8" in lowered or "uint8" in lowered or "bool" in lowered: + return 1 + return 4 diff --git a/adam/model_inspector/comparison.py b/adam/model_inspector/comparison.py new file mode 100644 index 0000000000000000000000000000000000000000..51d17b3f05e43d6adbd83085d282dd78e56cce4a --- /dev/null +++ b/adam/model_inspector/comparison.py @@ -0,0 +1,160 @@ +from __future__ import annotations + +import json +from pathlib import Path +from typing import Any + +from .base import ModelComparison, TensorComparison, is_cancelled, report +from .detector import inspect_model +from .generic import _extract_state_dict, _weight_files +from .statistics import component_for_name + + +def _iter_named_tensors(path: Path): + files = _weight_files(path) + for file in files: + if file.suffix.casefold() == ".safetensors": + from safetensors import safe_open + + with safe_open(str(file), framework="pt", device="cpu") as handle: + for key in handle.keys(): + yield key, handle.get_tensor(key) + else: + import torch + + payload = torch.load(str(file), map_location="cpu", weights_only=False) + state = _extract_state_dict(payload) + for key, value in state.items(): + if hasattr(value, "shape"): + yield str(key), value + + +def _tensor_map(path: Path) -> dict[str, Any]: + return {name: tensor for name, tensor in _iter_named_tensors(path)} + + +def _config_differences(configs_a: dict[str, Any], configs_b: dict[str, Any], *, limit: int = 40) -> list[str]: + diffs: list[str] = [] + keys = sorted(set(configs_a) | set(configs_b)) + for key in keys: + if key not in configs_a: + diffs.append(f"Only B has config {key}") + elif key not in configs_b: + diffs.append(f"Only A has config {key}") + elif json.dumps(configs_a[key], sort_keys=True, default=str) != json.dumps(configs_b[key], sort_keys=True, default=str): + diffs.append(f"Config differs: {key}") + if len(diffs) >= limit: + diffs.append("Additional config differences omitted.") + break + return diffs + + +def compare_models( + path_a: str | Path, + path_b: str | Path, + *, + arch_a: str = "", + arch_b: str = "", + settings_a: dict[str, Any] | None = None, + settings_b: dict[str, Any] | None = None, + progress=None, + cancelled=None, +) -> ModelComparison: + report(progress, 2, "Inspecting first model") + summary_a = inspect_model(path_a, recorded_architecture=arch_a, run_settings=settings_a, progress=progress, cancelled=cancelled) + report(progress, 30, "Inspecting second model") + summary_b = inspect_model(path_b, recorded_architecture=arch_b, run_settings=settings_b, progress=progress, cancelled=cancelled) + report(progress, 55, "Loading comparable tensors") + tensors_a = _tensor_map(Path(summary_a.resolved_path)) + if is_cancelled(cancelled): + raise RuntimeError("Comparison cancelled.") + tensors_b = _tensor_map(Path(summary_b.resolved_path)) + names_a = set(tensors_a) + names_b = set(tensors_b) + common = sorted(names_a & names_b) + only_a = sorted(names_a - names_b)[:200] + only_b = sorted(names_b - names_a)[:200] + shape_mismatches = [] + comparable = [] + for name in common: + if tuple(tensors_a[name].shape) != tuple(tensors_b[name].shape): + shape_mismatches.append(name) + else: + comparable.append(name) + comparisons: list[TensorComparison] = [] + import torch + + for index, name in enumerate(comparable): + if is_cancelled(cancelled): + raise RuntimeError("Comparison cancelled.") + if index % 10 == 0: + report(progress, 58 + int(36 * index / max(1, len(comparable))), f"Comparing {index + 1} of {len(comparable)} tensors") + with torch.no_grad(): + a = tensors_a[name].detach().to(device="cpu").float().reshape(-1) + b = tensors_b[name].detach().to(device="cpu").float().reshape(-1) + if a.numel() == 0: + continue + limit = 1_000_000 + if a.numel() > limit: + stride = max(1, a.numel() // limit) + a = a[::stride][:limit] + b = b[::stride][:limit] + delta = b - a + mean_abs = delta.abs().mean().item() + base_abs = a.abs().mean().item() + relative = mean_abs / (base_abs + 1e-12) + l2 = torch.linalg.vector_norm(delta).item() + norm_a = torch.linalg.vector_norm(a).item() + norm_b = torch.linalg.vector_norm(b).item() + cosine = torch.nn.functional.cosine_similarity(a, b, dim=0).item() if norm_a and norm_b else None + drift = (norm_b - norm_a) / (norm_a + 1e-12) if norm_a else None + score = relative * 0.7 + (1 - cosine if cosine is not None else 0) * 0.3 + comparisons.append( + TensorComparison( + name=name, + shape=tuple(int(dim) for dim in tensors_a[name].shape), + component=component_for_name(name), + mean_abs_difference=float(mean_abs), + relative_difference=float(relative), + cosine_similarity=float(cosine) if cosine is not None else None, + l2_distance=float(l2), + drift=float(drift) if drift is not None else None, + change_score=float(score), + ) + ) + comparisons.sort(key=lambda item: item.change_score or 0, reverse=True) + groups: dict[str, dict[str, float]] = {} + for item in comparisons: + group = groups.setdefault(item.component, {"tensors": 0, "mean_change_score": 0.0, "mean_abs_difference": 0.0}) + group["tensors"] += 1 + group["mean_change_score"] += item.change_score or 0 + group["mean_abs_difference"] += item.mean_abs_difference or 0 + for group in groups.values(): + count = max(1, int(group["tensors"])) + group["mean_change_score"] /= count + group["mean_abs_difference"] /= count + messages = [ + "Change Score is a statistical weight-change metric; it does not directly equal behavioral importance." + ] + if shape_mismatches: + messages.append("Some tensor comparisons are unavailable because tensor shapes differ.") + report(progress, 100, "Comparison complete") + return ModelComparison( + path_a=summary_a.resolved_path, + path_b=summary_b.resolved_path, + architecture_a=summary_a.architecture, + architecture_b=summary_b.architecture, + architecture_match=summary_a.architecture == summary_b.architecture, + config_differences=_config_differences(summary_a.configs, summary_b.configs), + resolution_difference=( + (summary_a.resolution, summary_b.resolution) + if summary_a.resolution != summary_b.resolution else None + ), + parameter_count_difference=summary_b.total_parameters - summary_a.total_parameters, + only_a=only_a, + only_b=only_b, + shape_mismatches=shape_mismatches[:200], + tensor_comparisons=comparisons[:200], + group_comparisons=groups, + messages=messages, + ) diff --git a/adam/model_inspector/ddpm.py b/adam/model_inspector/ddpm.py new file mode 100644 index 0000000000000000000000000000000000000000..ddccad125ce9d303b80cd98541bf53f776a1aead --- /dev/null +++ b/adam/model_inspector/ddpm.py @@ -0,0 +1,29 @@ +from __future__ import annotations + +from pathlib import Path +from typing import Any + +from .generic import GenericModelInspector + + +class DDPMInspector(GenericModelInspector): + architecture = "DDPM / Diffusers" + + def inspect(self, path: str | Path, *, recorded_architecture: str = "", run_settings: dict[str, Any] | None = None, progress=None, cancelled=None): + summary = super().inspect( + path, + recorded_architecture=recorded_architecture or "ddpm", + run_settings=run_settings, + progress=progress, + cancelled=cancelled, + ) + if summary.architecture == "Generic / Unknown": + summary.architecture = self.architecture + summary.messages.insert(0, "Model recognized as DDPM from ADAM trainer metadata") + expected = ("unet", "scheduler") + root = Path(summary.resolved_path) + existing = {part.name.casefold() for part in (root.iterdir() if root.is_dir() else [])} + for component in expected: + if root.is_dir() and component not in existing and not any(component in item.casefold() for item in summary.configs): + summary.health.append(f"Unusual: expected DDPM component not found: {component}") + return summary diff --git a/adam/model_inspector/detector.py b/adam/model_inspector/detector.py new file mode 100644 index 0000000000000000000000000000000000000000..b2d102b433dcbe9b3b95810989706943bbb8a0e9 --- /dev/null +++ b/adam/model_inspector/detector.py @@ -0,0 +1,35 @@ +from __future__ import annotations + +from pathlib import Path +from typing import Any + +from .ddpm import DDPMInspector +from .flow_matching import FlowMatchingInspector +from .generic import GenericModelInspector +from .lora import LoRAInspector +from .maskgit import MaskGITInspector + + +def inspector_for(path: str | Path, recorded_architecture: str = "", settings: dict[str, Any] | None = None): + text = " ".join((str(path), recorded_architecture, str(settings or {}))).casefold() + target = Path(path).expanduser() + if "lora" in text or target.suffix.casefold() == ".safetensors" and "adapter" in target.name.casefold(): + return LoRAInspector() + if "maskgit" in text: + return MaskGITInspector() + if "flow" in text or (target / "flow_model_info.json").is_file(): + return FlowMatchingInspector() + if "ddpm" in text or (target / "model_index.json").is_file() or (target / "scheduler").is_dir(): + return DDPMInspector() + return GenericModelInspector() + + +def inspect_model(path: str | Path, *, recorded_architecture: str = "", run_settings: dict[str, Any] | None = None, progress=None, cancelled=None): + inspector = inspector_for(path, recorded_architecture, run_settings) + return inspector.inspect( + path, + recorded_architecture=recorded_architecture, + run_settings=run_settings, + progress=progress, + cancelled=cancelled, + ) diff --git a/adam/model_inspector/flow_matching.py b/adam/model_inspector/flow_matching.py new file mode 100644 index 0000000000000000000000000000000000000000..48d2262013a404f26316a55d4c310c9fbd19f553 --- /dev/null +++ b/adam/model_inspector/flow_matching.py @@ -0,0 +1,25 @@ +from __future__ import annotations + +from pathlib import Path +from typing import Any + +from .generic import GenericModelInspector + + +class FlowMatchingInspector(GenericModelInspector): + architecture = "Flow Matching" + + def inspect(self, path: str | Path, *, recorded_architecture: str = "", run_settings: dict[str, Any] | None = None, progress=None, cancelled=None): + summary = super().inspect( + path, + recorded_architecture=recorded_architecture or "flow", + run_settings=run_settings, + progress=progress, + cancelled=cancelled, + ) + summary.architecture = "Flow Matching" if summary.architecture == "Generic / Unknown" else summary.architecture + if not any("flow_model_info.json" in item for item in summary.config_files): + root = Path(summary.resolved_path) + if root.is_dir(): + summary.health.append("Unusual: Flow Matching metadata file was not found") + return summary diff --git a/adam/model_inspector/generic.py b/adam/model_inspector/generic.py new file mode 100644 index 0000000000000000000000000000000000000000..0f88b1c343822d6219906a491cfcf61856110280 --- /dev/null +++ b/adam/model_inspector/generic.py @@ -0,0 +1,349 @@ +from __future__ import annotations + +import json +from collections import Counter +from pathlib import Path +from typing import Any, Iterator + +from .base import BaseModelInspector, InspectorError, ModelInspection, TensorStats, is_cancelled, report +from .statistics import ( + discover_checkpoint_paths, + discover_config_files, + step_from_name, + tensor_stats_from_torch, +) + + +def _folder_size(path: Path) -> int: + if path.is_file(): + return path.stat().st_size + total = 0 + try: + for item in path.rglob("*"): + if item.is_file(): + total += item.stat().st_size + except OSError: + return total + return total + + +def _read_configs(files: list[Path], root: Path) -> dict[str, Any]: + configs: dict[str, Any] = {} + for file in files[:40]: + try: + key = str(file.relative_to(root if root.is_dir() else root.parent)) + except ValueError: + key = file.name + try: + configs[key] = json.loads(file.read_text(encoding="utf-8")) + except (OSError, UnicodeDecodeError, json.JSONDecodeError): + configs[key] = "" + return configs + + +def _resolution_from_configs(configs: dict[str, Any], settings: dict[str, Any] | None) -> int | None: + for source in (settings or {}, *[value for value in configs.values() if isinstance(value, dict)]): + for key in ("resolution", "sample_size", "image_size", "size"): + value = source.get(key) if isinstance(source, dict) else None + if isinstance(value, int): + return value + if isinstance(value, (list, tuple)) and value and isinstance(value[0], int): + return int(value[0]) + try: + if value: + return int(value) + except (TypeError, ValueError): + pass + return None + + +def _iter_safetensors(file: Path) -> Iterator[tuple[str, Any, dict[str, Any]]]: + from safetensors import safe_open + + with safe_open(str(file), framework="pt", device="cpu") as handle: + metadata = handle.metadata() or {} + for key in handle.keys(): + yield key, handle.get_tensor(key), metadata + + +def _extract_state_dict(payload: Any) -> dict[str, Any]: + try: + import torch + except Exception: + torch = None + if torch is not None and hasattr(payload, "shape"): + return {"tensor": payload} + if isinstance(payload, dict): + for key in ("state_dict", "model_state_dict", "model", "module", "unet", "network"): + value = payload.get(key) + if isinstance(value, dict) and any(hasattr(item, "shape") for item in value.values()): + return value + if any(hasattr(item, "shape") for item in payload.values()): + return payload + return {} + + +def _iter_torch_checkpoint(file: Path) -> Iterator[tuple[str, Any, dict[str, Any]]]: + import torch + + try: + payload = torch.load(str(file), map_location="cpu", weights_only=True) + except TypeError: + payload = torch.load(str(file), map_location="cpu") + except Exception: + payload = torch.load(str(file), map_location="cpu", weights_only=False) + state = _extract_state_dict(payload) + metadata = {key: value for key, value in payload.items() if key not in state} if isinstance(payload, dict) else {} + for key, tensor in state.items(): + if hasattr(tensor, "shape"): + yield str(key), tensor, metadata + + +def _weight_files(path: Path) -> list[Path]: + if path.is_file(): + return [path] + ignored_names = {"optimizer.bin", "scheduler.bin", "scaler.pt"} + names = { + "diffusion_pytorch_model.safetensors", + "model.safetensors", + "pytorch_model.bin", + "adapter_model.safetensors", + "adapter_model.bin", + "checkpoint.pt", + "best_checkpoint.pt", + } + files: list[Path] = [] + try: + for item in path.rglob("*"): + if item.is_file() and (item.name in names or item.suffix.casefold() in {".safetensors", ".pt", ".pth", ".bin", ".ckpt"}): + if item.name.casefold() not in ignored_names: + files.append(item) + except OSError: + return [] + if (path / "model_index.json").is_file(): + final_files = [ + item for item in files + if not any(part.startswith("checkpoint-") for part in item.relative_to(path).parts) + ] + if final_files: + files = final_files + return sorted(files, key=lambda item: (0 if item.name in names else 1, str(item))) + + +class GenericModelInspector(BaseModelInspector): + architecture = "Generic / Unknown" + + def inspect( + self, + path: str | Path, + *, + recorded_architecture: str = "", + run_settings: dict[str, Any] | None = None, + progress=None, + cancelled=None, + ) -> ModelInspection: + target = Path(path).expanduser() + if not target.exists(): + raise InspectorError(f"Model path does not exist: {target}") + target = target.resolve() + report(progress, 3, "Finding model files") + config_files = discover_config_files(target) + configs = _read_configs(config_files, target) + files = _weight_files(target) + if not files: + message = "Model contains no readable tensor checkpoint" + return self._empty(target, recorded_architecture, run_settings, config_files, configs, message) + + tensors: list[TensorStats] = [] + dtypes: Counter[str] = Counter() + components: Counter[str] = Counter() + health: list[str] = [] + messages: list[str] = [] + metadata: dict[str, Any] = {} + for file_index, file in enumerate(files): + if is_cancelled(cancelled): + raise InspectorError("Inspection cancelled.") + report(progress, 8 + int(80 * file_index / max(1, len(files))), f"Reading {file.name}") + try: + if file.suffix.casefold() == ".safetensors": + iterator = _iter_safetensors(file) + else: + iterator = _iter_torch_checkpoint(file) + for name, tensor, file_metadata in iterator: + if is_cancelled(cancelled): + raise InspectorError("Inspection cancelled.") + prefix = file.parent.name if len(files) > 1 else "" + stat = tensor_stats_from_torch(f"{prefix}.{name}" if prefix and not name.startswith(prefix) else name, tensor) + tensors.append(stat) + dtypes[stat.dtype] += stat.parameter_count + components[stat.component] += stat.parameter_count + health.extend(f"{stat.name}: {item}" for item in stat.health) + if file_metadata: + metadata.update(file_metadata) + except Exception as exc: + health.append(f"{file.name}: unreadable checkpoint ({exc})") + + if not tensors: + message = "Model contains no readable tensor checkpoint" + return self._empty(target, recorded_architecture, run_settings, config_files, configs, message, [*(health or []), message]) + + report(progress, 92, "Summarizing model") + total_parameters = sum(tensor.parameter_count for tensor in tensors) + parameter_memory = sum(tensor.memory_bytes for tensor in tensors) + largest = sorted(tensors, key=lambda item: item.parameter_count, reverse=True)[:20] + architecture, confidence, message = self._architecture_from_signals( + target, recorded_architecture, configs, [tensor.name for tensor in tensors] + ) + messages.append(message) + duplicate_count = len(tensors) - len({tensor.name for tensor in tensors}) + if duplicate_count: + health.append(f"Unusual: {duplicate_count} duplicate tensor names after folder merging") + checkpoints = [str(item) for item in discover_checkpoint_paths(target)] + return ModelInspection( + path=str(path), + resolved_path=str(target), + architecture=architecture, + confidence=confidence, + status="ok", + size_bytes=_folder_size(target), + config_files=[str(item) for item in config_files], + resolution=_resolution_from_configs(configs, run_settings), + epoch=self._number_from_metadata(metadata, "epoch"), + step=self._number_from_metadata(metadata, "step") or step_from_name(target.name), + tensor_count=len(tensors), + total_parameters=total_parameters, + trainable_parameters=self._trainable_parameters(tensors, architecture), + parameter_memory_bytes=parameter_memory, + dtypes=dict(dtypes), + components=dict(components), + largest_tensors=largest, + tensors=tensors, + health=health or ["No invalid tensor values found in sampled statistics."], + messages=messages, + lora=self._lora_info(tensors, configs), + configs=configs, + histogram=self._histogram(tensors), + tensor_size_distribution=[(tensor.name, tensor.parameter_count) for tensor in largest], + checkpoints=checkpoints, + loss_history=[], + ) + + def _empty( + self, + target: Path, + recorded_architecture: str, + run_settings: dict[str, Any] | None, + config_files: list[Path], + configs: dict[str, Any], + message: str, + health: list[str] | None = None, + ) -> ModelInspection: + architecture, confidence, detection_message = self._architecture_from_signals(target, recorded_architecture, configs, []) + return ModelInspection( + path=str(target), + resolved_path=str(target), + architecture=architecture, + confidence=confidence, + status="warning", + size_bytes=_folder_size(target), + config_files=[str(item) for item in config_files], + resolution=_resolution_from_configs(configs, run_settings), + epoch=None, + step=step_from_name(target.name), + tensor_count=0, + total_parameters=0, + trainable_parameters=None, + parameter_memory_bytes=0, + dtypes={}, + components={}, + largest_tensors=[], + tensors=[], + health=health or [message], + messages=[detection_message, message], + configs=configs, + checkpoints=[str(item) for item in discover_checkpoint_paths(target)], + ) + + @staticmethod + def _number_from_metadata(metadata: dict[str, Any], key: str) -> int | None: + for candidate in (key, f"global_{key}", f"current_{key}"): + try: + value = metadata.get(candidate) + if value is not None: + return int(value) + except (TypeError, ValueError): + pass + return None + + @staticmethod + def _trainable_parameters(tensors: list[TensorStats], architecture: str) -> int | None: + if architecture == "LoRA": + return sum(tensor.parameter_count for tensor in tensors) + return None + + @staticmethod + def _architecture_from_signals( + target: Path, + recorded_architecture: str, + configs: dict[str, Any], + tensor_names: list[str], + ) -> tuple[str, float, str]: + recorded = recorded_architecture.casefold() + joined_names = "\n".join(tensor_names).casefold() + config_text = json.dumps(configs, default=str).casefold() + folder_text = str(target).casefold() + signals = " ".join((joined_names, config_text, folder_text)) + if "lora" in recorded or "lora" in signals or "adapter_config" in signals: + return "LoRA", 0.92, "Model recognized as LoRA" + if "maskgit" in recorded or "maskgit" in signals: + return "MaskGIT", 0.86, "Model recognized as MaskGIT" + if "flow" in recorded or "rectified_flow" in signals or "flow_model_info" in signals: + return "Flow Matching", 0.9, "Model recognized as Flow Matching" + if "ddpm" in recorded or "diffusers" in config_text or "unet" in signals or "scheduler_config" in signals: + return "DDPM / Diffusers", 0.88, "Model recognized as DDPM" + return "Generic / Unknown", 0.35, "Model type uncertain - using generic tensor inspection" + + @staticmethod + def _lora_info(tensors: list[TensorStats], configs: dict[str, Any]) -> dict[str, Any]: + lora_tensors = [tensor for tensor in tensors if "lora" in tensor.name.casefold()] + if not lora_tensors: + return {} + down = [tensor for tensor in lora_tensors if any(token in tensor.name.casefold() for token in ("down", "lora_a"))] + up = [tensor for tensor in lora_tensors if any(token in tensor.name.casefold() for token in ("up", "lora_b"))] + ranks = sorted({tensor.shape[0] for tensor in down if tensor.shape}) + alpha = None + targets: set[str] = set() + for config in configs.values(): + if isinstance(config, dict): + alpha = config.get("lora_alpha", config.get("alpha", alpha)) + modules = config.get("target_modules") + if isinstance(modules, list): + targets.update(str(item) for item in modules) + if not targets: + for tensor in lora_tensors: + parts = tensor.name.split(".") + if len(parts) > 2: + targets.add(parts[-3]) + return { + "rank": ", ".join(str(item) for item in ranks[:8]) if ranks else "unknown", + "alpha": alpha if alpha is not None else "unknown", + "target_modules": sorted(targets)[:20], + "down_matrices": len(down), + "up_matrices": len(up), + "adapter_parameter_count": sum(tensor.parameter_count for tensor in lora_tensors), + "average_abs_mean": ( + sum(tensor.abs_mean or 0 for tensor in lora_tensors) / max(1, len(lora_tensors)) + ), + } + + @staticmethod + def _histogram(tensors: list[TensorStats]) -> dict[str, list[float]]: + values = [tensor.abs_mean for tensor in tensors if tensor.abs_mean is not None] + if not values: + return {} + buckets = [0.0] * 10 + high = max(values) or 1.0 + for value in values: + index = min(9, int((value / high) * 10)) + buckets[index] += 1 + return {"abs_mean_bins": [round(high * index / 10, 6) for index in range(11)], "counts": buckets} diff --git a/adam/model_inspector/lora.py b/adam/model_inspector/lora.py new file mode 100644 index 0000000000000000000000000000000000000000..bbd76baba87d2f103026d4ea7a21c74d4f845b21 --- /dev/null +++ b/adam/model_inspector/lora.py @@ -0,0 +1,23 @@ +from __future__ import annotations + +from pathlib import Path +from typing import Any + +from .generic import GenericModelInspector + + +class LoRAInspector(GenericModelInspector): + architecture = "LoRA" + + def inspect(self, path: str | Path, *, recorded_architecture: str = "", run_settings: dict[str, Any] | None = None, progress=None, cancelled=None): + summary = super().inspect( + path, + recorded_architecture=recorded_architecture or "lora", + run_settings=run_settings, + progress=progress, + cancelled=cancelled, + ) + summary.architecture = "LoRA" if summary.architecture == "Generic / Unknown" else summary.architecture + if summary.tensors and not summary.lora: + summary.health.append("Worth inspecting: ADAM marked this as LoRA, but LoRA tensor naming was not obvious") + return summary diff --git a/adam/model_inspector/maskgit.py b/adam/model_inspector/maskgit.py new file mode 100644 index 0000000000000000000000000000000000000000..e8428feeb470c4112e7957db97eb2a3daf41dd70 --- /dev/null +++ b/adam/model_inspector/maskgit.py @@ -0,0 +1,24 @@ +from __future__ import annotations + +from pathlib import Path +from typing import Any + +from .generic import GenericModelInspector + + +class MaskGITInspector(GenericModelInspector): + architecture = "MaskGIT" + + def inspect(self, path: str | Path, *, recorded_architecture: str = "", run_settings: dict[str, Any] | None = None, progress=None, cancelled=None): + summary = super().inspect( + path, + recorded_architecture=recorded_architecture or "maskgit", + run_settings=run_settings, + progress=progress, + cancelled=cancelled, + ) + summary.architecture = "MaskGIT" if summary.architecture == "Generic / Unknown" else summary.architecture + names = "\n".join(tensor.name for tensor in summary.tensors).casefold() + if summary.tensors and not any(token in names for token in ("attention", "attn", "transformer", "embed")): + summary.health.append("Worth inspecting: expected MaskGIT transformer or embedding tensors were not obvious") + return summary diff --git a/adam/model_inspector/statistics.py b/adam/model_inspector/statistics.py new file mode 100644 index 0000000000000000000000000000000000000000..d48b69c29530df3f5c3ebb45f232f6c6002e5f15 --- /dev/null +++ b/adam/model_inspector/statistics.py @@ -0,0 +1,142 @@ +from __future__ import annotations + +import math +from pathlib import Path +from typing import Any + +from .base import CONFIG_FILENAMES, MODEL_EXTENSIONS, TensorStats, dtype_size, parameter_count + + +def bytes_label(size: int | float | None) -> str: + if size is None: + return "-" + value = float(size) + for unit in ("B", "KB", "MB", "GB", "TB"): + if abs(value) < 1024 or unit == "TB": + return f"{value:.1f} {unit}" if unit != "B" else f"{int(value)} B" + value /= 1024 + return f"{value:.1f} TB" + + +def component_for_name(name: str) -> str: + lowered = name.casefold() + mapping = ( + ("down_blocks", ("down_blocks", "down.", "downsample")), + ("mid_block", ("mid_block", "middle_block", "mid.")), + ("up_blocks", ("up_blocks", "up.", "upsample")), + ("attention", ("attn", "attention", "to_q", "to_k", "to_v", "query", "key", "value")), + ("embeddings", ("embed", "embedding", "position", "token")), + ("transformer blocks", ("transformer", "blocks.", "layers.", "encoder", "decoder")), + ("output layers", ("out.", "output", "proj_out", "lm_head", "conv_out")), + ("LoRA adapters", ("lora", "hada", "lokr", "adapter")), + ("normalization", ("norm", "bn", "ln", "group_norm", "layer_norm")), + ) + for component, tokens in mapping: + if any(token in lowered for token in tokens): + return component + return name.split(".", 1)[0] if "." in name else "other" + + +def safe_number(value: Any) -> float | None: + try: + number = float(value) + except (TypeError, ValueError, OverflowError): + return None + return number if math.isfinite(number) else None + + +def tensor_stats_from_torch(name: str, tensor: Any, *, sample_limit: int = 1_000_000) -> TensorStats: + shape = tuple(int(dim) for dim in getattr(tensor, "shape", ())) + dtype = str(getattr(tensor, "dtype", "unknown")).replace("torch.", "") + count = parameter_count(shape) + stat = TensorStats( + name=name, + shape=shape, + dtype=dtype, + parameter_count=count, + memory_bytes=count * dtype_size(dtype), + component=component_for_name(name), + ) + if count == 0: + stat.health.append("Empty tensor") + return stat + try: + import torch + + with torch.no_grad(): + values = tensor.detach().to(device="cpu") + if not values.is_floating_point() and not values.is_complex(): + values = values.float() + else: + values = values.float() + flat = values.reshape(-1) + if flat.numel() > sample_limit: + stride = max(1, flat.numel() // sample_limit) + flat = flat[::stride][:sample_limit] + finite = torch.isfinite(flat) + if not bool(finite.all()): + if bool(torch.isnan(flat).any()): + stat.health.append("Invalid: NaN values found") + if bool(torch.isinf(flat).any()): + stat.health.append("Invalid: Inf values found") + flat = flat[finite] + if flat.numel() == 0: + return stat + stat.minimum = safe_number(flat.min().item()) + stat.maximum = safe_number(flat.max().item()) + stat.mean = safe_number(flat.mean().item()) + stat.std = safe_number(flat.std(unbiased=False).item()) if flat.numel() > 1 else 0.0 + stat.abs_mean = safe_number(flat.abs().mean().item()) + stat.l2_norm = safe_number(torch.linalg.vector_norm(flat).item()) + stat.zero_percent = safe_number((flat == 0).float().mean().item() * 100) + except Exception as exc: + stat.health.append(f"Statistics unavailable: {exc}") + if stat.abs_mean is not None and stat.abs_mean > 100: + stat.health.append("Unusual: very large average weight magnitude") + if stat.maximum is not None and stat.minimum is not None and max(abs(stat.maximum), abs(stat.minimum)) > 1_000: + stat.health.append("Unusual: very large absolute weight value") + return stat + + +def discover_config_files(path: Path) -> list[Path]: + root = path if path.is_dir() else path.parent + files: list[Path] = [] + try: + for item in root.rglob("*"): + if item.is_file() and item.name in CONFIG_FILENAMES: + files.append(item) + except OSError: + return [] + return sorted(files) + + +def discover_checkpoint_paths(path: Path) -> list[Path]: + root = path if path.is_dir() else path.parent + candidates: list[Path] = [] + try: + for item in root.rglob("*"): + if item.is_file() and item.suffix.casefold() in MODEL_EXTENSIONS: + candidates.append(item) + elif item.is_dir() and item.name.startswith("checkpoint-"): + candidates.append(item) + except OSError: + return [] + return sorted(candidates, key=lambda item: (step_from_name(item.name) or -1, str(item))) + + +def step_from_name(name: str) -> int | None: + import re + + matches = re.findall(r"(?:step|checkpoint|epoch|e|s)[-_]?(\d+)", name, flags=re.I) + if not matches: + matches = re.findall(r"(\d+)", name) + if not matches: + return None + try: + return int(matches[-1]) + except ValueError: + return None + + +def shape_label(shape: tuple[int, ...]) -> str: + return " x ".join(str(dim) for dim in shape) if shape else "scalar" diff --git a/adam/model_plugin_backend.py b/adam/model_plugin_backend.py new file mode 100644 index 0000000000000000000000000000000000000000..ecfdb86418e97712a313185c1c65a0827e11afd7 --- /dev/null +++ b/adam/model_plugin_backend.py @@ -0,0 +1,38 @@ +from __future__ import annotations + +from typing import Any + +from adam.executor import ToolContext, ToolExecutionError +from adam.model_plugins import ModelPluginRegistry, plugin_function, validate_settings + + +def _plugin_for_tool(context: ToolContext, mode: str): + registry = ModelPluginRegistry(context.root) + for plugin in registry.all(): + if mode == "training" and plugin.trainer_id == context.tool.id: + return plugin + if mode == "generation" and plugin.generator_id == context.tool.id: + return plugin + raise ToolExecutionError(f"No model plugin owns {context.tool.id}.") + + +def train(context: ToolContext, **settings: Any) -> dict[str, Any]: + plugin = _plugin_for_tool(context, "training") + errors = validate_settings(plugin.training_settings, settings) + if errors: + raise ToolExecutionError(" ".join(errors)) + function = plugin_function(plugin, "train") + if function is None: + raise ToolExecutionError(f"{plugin.name} does not implement train().") + return function(settings=settings, callbacks=context) or {} + + +def generate(context: ToolContext, **settings: Any) -> dict[str, Any]: + plugin = _plugin_for_tool(context, "generation") + errors = validate_settings(plugin.generation_settings, settings) + if errors: + raise ToolExecutionError(" ".join(errors)) + function = plugin_function(plugin, "generate") + if function is None: + raise ToolExecutionError(f"{plugin.name} does not implement generate().") + return function(settings=settings, callbacks=context) or {} diff --git a/adam/model_plugins.py b/adam/model_plugins.py new file mode 100644 index 0000000000000000000000000000000000000000..ac477fff3e91a069e3e811e6db3dd7abf914b03b --- /dev/null +++ b/adam/model_plugins.py @@ -0,0 +1,494 @@ +from __future__ import annotations + +import importlib +import importlib.util +import json +import logging +import pkgutil +from dataclasses import dataclass, field +from pathlib import Path +from typing import Any, Callable + + +REQUIRED_INFO_FIELDS = {"name", "version", "category", "description"} +SUPPORTED_SETTING_TYPES = { + "int", + "float", + "bool", + "choice", + "text", + "multiline_text", + "path", + "folder", + "slider", +} + + +class ModelPluginError(RuntimeError): + pass + + +@dataclass(frozen=True, slots=True) +class ModelPlugin: + id: str + info: dict[str, Any] + training_settings: dict[str, dict[str, Any]] = field(default_factory=dict) + generation_settings: dict[str, dict[str, Any]] = field(default_factory=dict) + training_tool: dict[str, Any] = field(default_factory=dict) + generation_tool: dict[str, Any] = field(default_factory=dict) + module_name: str = "" + plugin_path: Path | None = None + + @property + def name(self) -> str: + return str(self.info.get("name", self.id)) + + @property + def trainer_id(self) -> str: + return str(self.training_tool.get("id") or f"{self.id}_trainer") + + @property + def generator_id(self) -> str: + return str(self.generation_tool.get("id") or f"{self.id}_generator") + + +class ModelPluginRegistry: + """Discovers model plugins and validates their setting schemas.""" + + def __init__(self, root: Path, logger: logging.Logger | None = None) -> None: + self.root = root.resolve() + self.logger = logger or logging.getLogger(__name__) + self.plugins: dict[str, ModelPlugin] = {} + self.errors: list[str] = [] + self.discover() + + def discover(self) -> None: + self.plugins = {} + self.errors = [] + for module_name in self._candidate_modules(): + try: + plugin = self._load_module_plugin(module_name) + except Exception as exc: + message = f"{module_name}: {exc}" + self.errors.append(message) + self.logger.warning("Model plugin failed to load: %s", message) + continue + if plugin.id in self.plugins: + self.errors.append(f"{module_name}: duplicate model plugin id {plugin.id}") + continue + self.plugins[plugin.id] = plugin + + def _candidate_modules(self) -> list[str | Path]: + modules: list[str | Path] = [] + try: + package = importlib.import_module("adam.model_plugins_builtin") + for item in pkgutil.iter_modules(package.__path__, package.__name__ + "."): + if not item.ispkg: + continue + modules.append(item.name + ".manifest") + except Exception as exc: + self.errors.append(f"adam.model_plugins_builtin: {exc}") + + models_dir = self.root / "models" + if models_dir.is_dir(): + for folder in sorted(models_dir.iterdir()): + manifest = folder / "manifest.py" + if not folder.is_dir() or not manifest.is_file(): + continue + modules.append(manifest) + return modules + + def _load_module_plugin(self, module_name: str | Path) -> ModelPlugin: + if isinstance(module_name, Path): + fallback_id = module_name.parent.name + unique_name = f"adam_user_model_{fallback_id}_{abs(hash(str(module_name.resolve())))}" + spec = importlib.util.spec_from_file_location(unique_name, module_name) + if spec is None or spec.loader is None: + raise ModelPluginError(f"Could not load manifest file: {module_name}") + module = importlib.util.module_from_spec(spec) + spec.loader.exec_module(module) + module_label = str(module_name) + else: + module = importlib.import_module(module_name) + fallback_id = module_name.split(".")[-2] + module_label = module_name + plugin_id = str(getattr(module, "PLUGIN_ID", "") or fallback_id) + info = dict(getattr(module, "MODEL_INFO", {})) + missing = REQUIRED_INFO_FIELDS - set(info) + if missing: + raise ModelPluginError( + "MODEL_INFO is missing " + ", ".join(sorted(missing)) + ) + training_settings = self._validate_schema( + dict(getattr(module, "TRAINING_SETTINGS", {})), + f"{plugin_id} training", + ) + generation_settings = self._validate_schema( + dict(getattr(module, "GENERATION_SETTINGS", {})), + f"{plugin_id} generation", + ) + plugin_path = Path(getattr(module, "__file__", "")).resolve().parent + return ModelPlugin( + id=plugin_id, + info=info, + training_settings=training_settings, + generation_settings=generation_settings, + training_tool=dict(getattr(module, "TRAINING_TOOL", {})), + generation_tool=dict(getattr(module, "GENERATION_TOOL", {})), + module_name=module_label, + plugin_path=plugin_path, + ) + + @staticmethod + def _validate_schema( + schema: dict[str, Any], + label: str, + ) -> dict[str, dict[str, Any]]: + clean: dict[str, dict[str, Any]] = {} + for key, raw in schema.items(): + if not isinstance(raw, dict): + raise ModelPluginError(f"{label} setting {key} must be an object") + spec = dict(raw) + setting_type = str(spec.get("type", "text")) + if setting_type not in SUPPORTED_SETTING_TYPES: + raise ModelPluginError( + f"{label} setting {key} has unsupported type {setting_type}" + ) + spec["type"] = setting_type + spec.setdefault("label", key.replace("_", " ").title()) + spec.setdefault("group", "Basic") + if setting_type == "choice": + options = spec.get("options", []) + if not isinstance(options, (list, tuple)) or not options: + raise ModelPluginError(f"{label} setting {key} needs options") + spec["options"] = list(options) + spec.setdefault("default", spec["options"][0]) + clean[str(key)] = spec + return clean + + def get(self, plugin_id: str) -> ModelPlugin: + return self.plugins[plugin_id] + + def all(self) -> list[ModelPlugin]: + return list(self.plugins.values()) + + def by_trainer(self, trainer: str) -> ModelPlugin | None: + return next((plugin for plugin in self.plugins.values() if plugin.id == trainer), None) + + def training_schema(self, trainer: str) -> dict[str, dict[str, Any]]: + plugin = self.by_trainer(trainer) + return plugin.training_settings if plugin else {} + + def generation_schema_for_tool(self, tool_id: str) -> dict[str, dict[str, Any]]: + for plugin in self.plugins.values(): + if plugin.generator_id == tool_id: + return plugin.generation_settings + return {} + + def training_tool_specs(self) -> list[dict[str, Any]]: + return [ + self._tool_spec(plugin, mode="training") + for plugin in self.plugins.values() + if plugin.training_tool + ] + + def generation_tool_specs(self) -> list[dict[str, Any]]: + return [ + self._tool_spec(plugin, mode="generation") + for plugin in self.plugins.values() + if plugin.generation_tool + ] + + @staticmethod + def _tool_spec(plugin: ModelPlugin, *, mode: str) -> dict[str, Any]: + tool = dict(plugin.training_tool if mode == "training" else plugin.generation_tool) + schema = plugin.training_settings if mode == "training" else plugin.generation_settings + core_arguments = ( + ["dataset_dir", "model_name", "epochs", "output_dir", "resume_from"] + if mode == "training" + else [ + "model_name", "model_path", "prompt", "image_count", "steps", + "seed", "sampler", "aspect_ratio", + ] + ) + core_required = ( + ["dataset_dir", "model_name", "epochs", "output_dir"] + if mode == "training" + else ["model_name", "model_path", "image_count", "steps", "seed"] + ) + defaults = { + "id": plugin.trainer_id if mode == "training" else plugin.generator_id, + "name": f"{plugin.name} {'Trainer' if mode == 'training' else 'Generator'}", + "description": plugin.info.get("description", ""), + "category": "Training" if mode == "training" else "Output", + "entry_function": "train" if mode == "training" else "generate", + "arguments": [*core_arguments, *list(schema)], + "required_arguments": [ + *core_required, + *[key for key, spec in schema.items() if bool(spec.get("required"))], + ], + "capabilities": ( + ["fresh_training", "progress", "pause", "cancel"] + if mode == "training" + else ["image_generation", "progress", "cancel"] + ), + "requires_confirmation": mode == "training", + "enabled": True, + "demo": False, + } + defaults.update(tool) + defaults["arguments"] = list(defaults.get("arguments") or [*core_arguments, *list(schema)]) + defaults["required_arguments"] = list(defaults.get("required_arguments") or []) + return defaults + + def validate_settings( + self, + trainer: str, + values: dict[str, Any], + *, + mode: str = "training", + ) -> list[str]: + plugin = self.by_trainer(trainer) + if not plugin: + return [f"Unknown model plugin: {trainer}"] + schema = plugin.training_settings if mode == "training" else plugin.generation_settings + return validate_settings(schema, values) + + +def validate_settings(schema: dict[str, dict[str, Any]], values: dict[str, Any]) -> list[str]: + errors: list[str] = [] + for key, spec in schema.items(): + value = values.get(key, spec.get("default")) + label = str(spec.get("label", key)) + if spec.get("required") and (value is None or str(value).strip() == ""): + errors.append(f"{label} is required.") + continue + if value in (None, "") and not spec.get("required"): + continue + setting_type = str(spec.get("type", "text")) + try: + if setting_type in {"int", "slider"}: + if isinstance(value, bool): + raise ValueError + numeric = int(value) + elif setting_type == "float": + if isinstance(value, bool): + raise ValueError + numeric = float(value) + else: + numeric = None + except (TypeError, ValueError): + errors.append(f"{label} must be a number.") + continue + if numeric is not None: + if "min" in spec and numeric < float(spec["min"]): + errors.append(f"{label} must be at least {spec['min']}.") + if "max" in spec and numeric > float(spec["max"]): + errors.append(f"{label} must be at most {spec['max']}.") + if setting_type == "choice" and "options" in spec and value not in spec["options"]: + errors.append(f"{label} must be one of: {', '.join(map(str, spec['options']))}.") + if setting_type == "path" and spec.get("must_exist") and not Path(str(value)).expanduser().is_file(): + errors.append(f"{label} must point to an existing file.") + if setting_type == "folder" and spec.get("must_exist") and not Path(str(value)).expanduser().is_dir(): + errors.append(f"{label} must point to an existing folder.") + return errors + + +def load_presets(root: Path, plugin_id: str, mode: str) -> dict[str, dict[str, Any]]: + path = root.resolve() / "config" / "model_presets.json" + try: + payload = json.loads(path.read_text(encoding="utf-8")) + except (OSError, json.JSONDecodeError): + return {} + presets = payload.get(plugin_id, {}).get(mode, {}) + return dict(presets) if isinstance(presets, dict) else {} + + +def save_preset( + root: Path, + plugin_id: str, + mode: str, + name: str, + settings: dict[str, Any], +) -> None: + path = root.resolve() / "config" / "model_presets.json" + try: + payload = json.loads(path.read_text(encoding="utf-8")) + except (OSError, json.JSONDecodeError): + payload = {} + payload.setdefault(plugin_id, {}).setdefault(mode, {})[name] = settings + temporary = path.with_suffix(".tmp") + temporary.write_text(json.dumps(payload, indent=2, sort_keys=True), encoding="utf-8") + temporary.replace(path) + + +def plugin_function(plugin: ModelPlugin, function_name: str) -> Callable[..., Any] | None: + if plugin.module_name.endswith("manifest.py"): + spec = importlib.util.spec_from_file_location( + f"adam_user_model_{plugin.id}_{abs(hash(plugin.module_name))}", + plugin.module_name, + ) + if spec is None or spec.loader is None: + return None + module = importlib.util.module_from_spec(spec) + spec.loader.exec_module(module) + else: + module = importlib.import_module(plugin.module_name) + function = getattr(module, function_name, None) + return function if callable(function) else None + + +def safe_plugin_id(name: str) -> str: + cleaned = "".join( + character.lower() if character.isalnum() else "_" + for character in name.strip() + ) + cleaned = "_".join(part for part in cleaned.split("_") if part) + return cleaned[:48] or "my_model" + + +def scaffold_model_plugin( + root: Path, + *, + plugin_id: str, + name: str, + architecture: str = "custom", + output_type: str = "image", + include_training: bool = True, + include_generation: bool = True, +) -> Path: + """Create a simple user-editable model plugin folder.""" + plugin_id = safe_plugin_id(plugin_id) + if plugin_id in {"ddpm", "flow", "lora", "model_template"}: + raise ModelPluginError("Choose a plugin id that does not conflict with a built-in model.") + folder = root.resolve() / "models" / plugin_id + if folder.exists(): + raise ModelPluginError(f"A model plugin folder already exists: {folder}") + folder.mkdir(parents=True) + (folder / "__init__.py").write_text( + f'"""ADAM model plugin: {name}."""\n', + encoding="utf-8", + ) + (folder / "manifest.py").write_text( + _manifest_template( + plugin_id=plugin_id, + name=name, + architecture=architecture, + output_type=output_type, + include_training=include_training, + include_generation=include_generation, + ), + encoding="utf-8", + ) + (folder / "model.py").write_text(_model_template(), encoding="utf-8") + if include_training: + (folder / "trainer.py").write_text(_trainer_template(), encoding="utf-8") + if include_generation: + (folder / "generator.py").write_text(_generator_template(), encoding="utf-8") + return folder + + +def _manifest_template( + *, + plugin_id: str, + name: str, + architecture: str, + output_type: str, + include_training: bool, + include_generation: bool, +) -> str: + plugin_id_json = json.dumps(plugin_id) + name_json = json.dumps(name) + architecture_json = json.dumps(architecture) + output_type_json = json.dumps(output_type) + training_tool = ( + "{\n" + f' "id": "{plugin_id}_trainer",\n' + f' "name": {json.dumps(name + " Trainer")},\n' + f' "backend": {{"type": "python", "module": "models.{plugin_id}.trainer", "function": "train"}},\n' + "}" + if include_training else "{}" + ) + generation_tool = ( + "{\n" + f' "id": "{plugin_id}_generator",\n' + f' "name": {json.dumps(name + " Generator")},\n' + f' "model_trainers": ["{plugin_id}"],\n' + f' "backend": {{"type": "python", "module": "models.{plugin_id}.generator", "function": "generate"}},\n' + "}" + if include_generation else "{}" + ) + return f'''PLUGIN_ID = {plugin_id_json} + +MODEL_INFO = {{ + "name": {name_json}, + "version": "0.1", + "category": "Image Generation", + "description": {json.dumps("Describe what " + name + " trains or generates.")}, + "architecture": {architecture_json}, + "status": "experimental", + "output_type": {output_type_json}, +}} + +TRAINING_SETTINGS = {{ + "resolution": {{"label": "Resolution", "type": "choice", "options": [64, 128, 256, 384, 512], "default": 256, "group": "Basic"}}, + "batch_size": {{"label": "Batch size", "type": "int", "default": 1, "min": 1, "max": 64, "group": "Basic"}}, + "learning_rate": {{"label": "Learning rate", "type": "float", "default": 0.0001, "min": 0.0000001, "max": 0.1, "decimals": 7, "group": "Optimization"}}, + "mixed_precision": {{"label": "Precision", "type": "choice", "options": ["fp16", "no"], "default": "fp16", "group": "Optimization"}}, + "preview_enabled": {{"label": "Generate previews while training", "type": "bool", "default": True, "group": "Preview"}}, + "preview_every": {{"label": "Preview interval", "type": "int", "default": 5, "min": 1, "max": 100000, "group": "Preview"}}, + "preview_prompt": {{"label": "Preview prompt", "type": "text", "default": "", "group": "Preview"}}, + "preview_seed": {{"label": "Preview seed", "type": "int", "default": 123456789, "min": 0, "max": 2147483647, "group": "Preview"}}, +}} + +GENERATION_SETTINGS = {{ + "prompt": {{"label": "Prompt", "type": "multiline_text", "default": "", "group": "Prompt"}}, + "image_count": {{"label": "Images", "type": "int", "default": 1, "min": 1, "max": 48, "group": "Generation"}}, + "steps": {{"label": "Steps", "type": "int", "default": 30, "min": 1, "max": 500, "group": "Generation"}}, + "seed": {{"label": "Seed", "type": "int", "default": 0, "min": 0, "max": 2147483647, "group": "Generation"}}, +}} + +TRAINING_TOOL = {training_tool} + +GENERATION_TOOL = {generation_tool} +''' + + +def _model_template() -> str: + return '''from __future__ import annotations + +from typing import Any + + +def load_model(model_path: str, settings: dict[str, Any] | None = None) -> Any: + """Load your model or inference pipeline here.""" + raise NotImplementedError("Add your model loading code.") +''' + + +def _trainer_template() -> str: + return '''from __future__ import annotations + +from typing import Any + + +def train(context, **settings: Any) -> dict[str, Any]: + """Train the model and report progress back to ADAM.""" + context.log("Replace this with real training code.") + context.progress(100, "Training placeholder complete") + return {} +''' + + +def _generator_template() -> str: + return '''from __future__ import annotations + +from typing import Any + + +def generate(context, **settings: Any) -> dict[str, Any]: + """Generate outputs and report progress back to ADAM.""" + context.log("Replace this with real generation code.") + context.progress(100, "Generation placeholder complete") + return {} +''' diff --git a/adam/model_plugins_builtin/__init__.py b/adam/model_plugins_builtin/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..d7bf666d97bd03851145e6e057063175d2675b2a --- /dev/null +++ b/adam/model_plugins_builtin/__init__.py @@ -0,0 +1 @@ +"""Built-in model plugin manifests shipped with ADAM.""" diff --git a/adam/model_plugins_builtin/ddpm/__init__.py b/adam/model_plugins_builtin/ddpm/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..f834fd1ce9571dfc141eeb4e86267e7edb72e00a --- /dev/null +++ b/adam/model_plugins_builtin/ddpm/__init__.py @@ -0,0 +1 @@ +"""DDPM model plugin.""" diff --git a/adam/model_plugins_builtin/ddpm/manifest.py b/adam/model_plugins_builtin/ddpm/manifest.py new file mode 100644 index 0000000000000000000000000000000000000000..92e333ea342718b4a552753af032b3a418cef2d1 --- /dev/null +++ b/adam/model_plugins_builtin/ddpm/manifest.py @@ -0,0 +1,56 @@ +PLUGIN_ID = "ddpm" + +MODEL_INFO = { + "name": "DDPM", + "version": "1.0", + "category": "Image Generation", + "description": "Denoising Diffusion Probabilistic Model image trainer and generator.", + "architecture": "diffusion", + "status": "stable", + "output_type": "image", + "capabilities": ["fresh_training", "resume_training", "image_generation", "smart_generation", "live_preview"], + "input_formats": ["image folder"], + "output_formats": ["diffusers pipeline", "checkpoint folder", "png preview"], + "hardware": {"recommended_vram_gb": 6, "recommended_system_ram_gb": 16}, + "vram_behavior": {"scales_with": ["resolution", "batch_size"], "estimate": "Moderate; batch size should drop quickly above 256px."}, +} + +TRAINING_SETTINGS = { + "resolution": {"label": "Resolution", "type": "choice", "options": [64, 128, 256, 384, 512], "default": 128, "group": "Basic"}, + "batch_size": {"label": "Batch size", "type": "int", "default": 1, "min": 1, "max": 64, "group": "Basic"}, + "learning_rate": {"label": "Learning rate", "type": "float", "default": 0.0001, "min": 0.0000001, "max": 0.1, "decimals": 7, "step": 0.00005, "group": "Optimization"}, + "gradient_accumulation_steps": {"label": "Gradient accumulation", "type": "int", "default": 1, "min": 1, "max": 64, "group": "Optimization"}, + "dataloader_num_workers": {"label": "Loader workers", "type": "int", "default": 4, "min": 0, "max": 16, "group": "Dataset"}, + "mixed_precision": {"label": "Precision", "type": "choice", "options": ["fp16", "no"], "default": "fp16", "group": "Optimization"}, + "save_every": {"label": "Save every", "type": "int", "default": 10, "min": 1, "max": 1000, "group": "Checkpoints"}, + "preview_steps": {"label": "Preview steps", "type": "int", "default": 50, "min": 1, "max": 500, "group": "Preview"}, + "training_intensity": {"label": "Training intensity", "type": "slider", "default": 100, "min": 10, "max": 100, "group": "Advanced", "advanced": True}, + "completed_epochs": {"label": "Completed epochs", "type": "int", "default": 0, "min": 0, "max": 100000, "group": "Internal", "advanced": True}, + "preview_enabled": {"label": "Generate previews while training", "type": "bool", "default": True, "group": "Preview"}, + "preview_every": {"label": "Preview interval", "type": "int", "default": 5, "min": 1, "max": 100000, "group": "Preview"}, + "preview_prompt": {"label": "Preview prompt", "type": "text", "default": "", "group": "Preview"}, + "preview_seed": {"label": "Preview seed", "type": "int", "default": 123456789, "min": 0, "max": 2147483647, "group": "Preview"}, +} + +GENERATION_SETTINGS = { + "prompt": {"label": "Creative note", "type": "multiline_text", "default": "", "group": "Generation"}, + "image_count": {"label": "Images", "type": "int", "default": 1, "min": 1, "max": 48, "group": "Generation"}, + "steps": {"label": "Sampling steps", "type": "int", "default": 50, "min": 5, "max": 500, "group": "Generation"}, + "sampler": {"label": "Sampler", "type": "choice", "options": ["DDIM", "DDPM"], "default": "DDIM", "group": "Generation"}, + "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"}, + "seed": {"label": "Seed", "type": "int", "default": 0, "min": 0, "max": 2147483647, "group": "Generation"}, + "reference_image": {"label": "Reference image", "type": "path", "default": "", "group": "Reference"}, + "reference_strength": {"label": "Reference strength", "type": "slider", "default": 65, "min": 0, "max": 100, "group": "Reference"}, + "width": {"label": "Custom width", "type": "int", "default": 0, "min": 0, "max": 2048, "group": "Advanced", "advanced": True}, + "height": {"label": "Custom height", "type": "int", "default": 0, "min": 0, "max": 2048, "group": "Advanced", "advanced": True}, + "preview_interval": {"label": "Steps per preview", "type": "int", "default": 0, "min": 0, "max": 500, "group": "Preview"}, + "smart_generation": {"label": "Smart Generation", "type": "bool", "default": False, "group": "Smart Generation"}, + "smart_wanted_results": {"label": "Wanted results", "type": "int", "default": 8, "min": 1, "max": 48, "group": "Smart Generation"}, + "smart_max_candidates": {"label": "Maximum candidates", "type": "int", "default": 32, "min": 1, "max": 256, "group": "Smart Generation"}, + "smart_min_score": {"label": "Minimum score", "type": "float", "default": 0.7, "min": 0, "max": 1, "group": "Smart Generation"}, + "smart_mode": {"label": "Selection mode", "type": "choice", "options": ["threshold", "top_n"], "default": "threshold", "group": "Smart Generation"}, + "smart_keep_rejected": {"label": "Keep rejected candidates", "type": "bool", "default": True, "group": "Smart Generation"}, +} + +TRAINING_TOOL = {"id": "ddpm_trainer"} +GENERATION_TOOL = {"id": "ddpm_generator"} diff --git a/adam/model_plugins_builtin/flow_matching/__init__.py b/adam/model_plugins_builtin/flow_matching/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..d245954847c1858b537b64b04aca83430c700957 --- /dev/null +++ b/adam/model_plugins_builtin/flow_matching/__init__.py @@ -0,0 +1 @@ +"""Flow Matching model plugin.""" diff --git a/adam/model_plugins_builtin/flow_matching/manifest.py b/adam/model_plugins_builtin/flow_matching/manifest.py new file mode 100644 index 0000000000000000000000000000000000000000..9df1fe4f165c44fc9b07da8bbaefcbd358c685c4 --- /dev/null +++ b/adam/model_plugins_builtin/flow_matching/manifest.py @@ -0,0 +1,51 @@ +PLUGIN_ID = "flow" + +MODEL_INFO = { + "name": "Flow Matching", + "version": "1.0", + "category": "Image Generation", + "description": "Rectified Flow image trainer and generator.", + "architecture": "rectified_flow", + "status": "stable", + "output_type": "image", + "capabilities": ["fresh_training", "resume_training", "image_generation", "smart_generation", "live_preview"], + "input_formats": ["image folder"], + "output_formats": ["diffusers unet folder", "flow metadata", "png preview"], + "hardware": {"recommended_vram_gb": 8, "recommended_system_ram_gb": 16}, + "vram_behavior": {"scales_with": ["resolution", "batch_size"], "estimate": "Moderate to high; flow runs usually want smaller batches at 512px."}, +} + +TRAINING_SETTINGS = { + "resolution": {"label": "Resolution", "type": "choice", "options": [64, 128, 256, 384, 512], "default": 256, "group": "Basic"}, + "batch_size": {"label": "Batch size", "type": "int", "default": 8, "min": 1, "max": 64, "group": "Basic"}, + "learning_rate": {"label": "Learning rate", "type": "float", "default": 0.0002, "min": 0.0000001, "max": 0.1, "decimals": 7, "step": 0.00005, "group": "Optimization"}, + "gradient_accumulation": {"label": "Gradient accumulation", "type": "int", "default": 1, "min": 1, "max": 64, "group": "Optimization"}, + "workers": {"label": "Loader workers", "type": "int", "default": 4, "min": 0, "max": 16, "group": "Dataset"}, + "mixed_precision": {"label": "Precision", "type": "choice", "options": ["fp16", "no"], "default": "fp16", "group": "Optimization"}, + "save_every": {"label": "Save every", "type": "int", "default": 10, "min": 1, "max": 1000, "group": "Checkpoints"}, + "preview_steps": {"label": "Preview steps", "type": "int", "default": 10, "min": 1, "max": 500, "group": "Preview"}, + "gradient_checkpointing": {"label": "Gradient checkpointing", "type": "bool", "default": False, "group": "Advanced", "advanced": True}, + "preview_enabled": {"label": "Generate previews while training", "type": "bool", "default": True, "group": "Preview"}, + "preview_every": {"label": "Preview interval", "type": "int", "default": 5, "min": 1, "max": 100000, "group": "Preview"}, + "preview_prompt": {"label": "Preview prompt", "type": "text", "default": "", "group": "Preview"}, + "preview_seed": {"label": "Preview seed", "type": "int", "default": 123456789, "min": 0, "max": 2147483647, "group": "Preview"}, +} + +GENERATION_SETTINGS = { + "prompt": {"label": "Creative note", "type": "multiline_text", "default": "", "group": "Generation"}, + "image_count": {"label": "Images", "type": "int", "default": 1, "min": 1, "max": 48, "group": "Generation"}, + "steps": {"label": "ODE steps", "type": "int", "default": 20, "min": 1, "max": 200, "group": "Generation"}, + "sampler": {"label": "Method", "type": "choice", "options": ["Heun", "Euler"], "default": "Heun", "group": "Generation"}, + "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"}, + "seed": {"label": "Seed", "type": "int", "default": 0, "min": 0, "max": 2147483647, "group": "Generation"}, + "preview_interval": {"label": "Steps per preview", "type": "int", "default": 0, "min": 0, "max": 500, "group": "Preview"}, + "smart_generation": {"label": "Smart Generation", "type": "bool", "default": False, "group": "Smart Generation"}, + "smart_wanted_results": {"label": "Wanted results", "type": "int", "default": 8, "min": 1, "max": 48, "group": "Smart Generation"}, + "smart_max_candidates": {"label": "Maximum candidates", "type": "int", "default": 32, "min": 1, "max": 256, "group": "Smart Generation"}, + "smart_min_score": {"label": "Minimum score", "type": "float", "default": 0.7, "min": 0, "max": 1, "group": "Smart Generation"}, + "smart_mode": {"label": "Selection mode", "type": "choice", "options": ["threshold", "top_n"], "default": "threshold", "group": "Smart Generation"}, + "smart_keep_rejected": {"label": "Keep rejected candidates", "type": "bool", "default": True, "group": "Smart Generation"}, +} + +TRAINING_TOOL = {"id": "flow_trainer"} +GENERATION_TOOL = {"id": "flow_generator"} diff --git a/adam/model_plugins_builtin/model_template/__init__.py b/adam/model_plugins_builtin/model_template/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..0f6ede77e7e58822c9b4bd92bbaaa63b3f181355 --- /dev/null +++ b/adam/model_plugins_builtin/model_template/__init__.py @@ -0,0 +1 @@ +"""Copyable model plugin template.""" diff --git a/adam/model_plugins_builtin/model_template/manifest.py b/adam/model_plugins_builtin/model_template/manifest.py new file mode 100644 index 0000000000000000000000000000000000000000..583e462446ed91be5143599f055cf3d52272fb0b --- /dev/null +++ b/adam/model_plugins_builtin/model_template/manifest.py @@ -0,0 +1,30 @@ +PLUGIN_ID = "model_template" + +MODEL_INFO = { + "name": "Model Template", + "version": "0.1", + "category": "Template", + "description": "Example manifest for adding a new ADAM model plugin.", + "status": "example", +} + +TRAINING_SETTINGS = { + "dataset_dir": {"label": "Dataset folder", "type": "folder", "required": True, "group": "Dataset"}, + "epochs": {"label": "Epochs", "type": "int", "default": 10, "min": 1, "max": 100000, "group": "Basic"}, + "learning_rate": {"label": "Learning rate", "type": "float", "default": 0.0001, "min": 0.0000001, "max": 0.1, "group": "Optimization"}, +} + +GENERATION_SETTINGS = { + "model_path": {"label": "Model file or folder", "type": "path", "required": True, "group": "Model Loading"}, + "seed": {"label": "Seed", "type": "int", "default": 0, "min": 0, "max": 2147483647, "group": "Generation"}, +} + +# A real plugin can point to its own backend: +# TRAINING_TOOL = { +# "backend": {"type": "python", "module": "models.my_model.trainer", "function": "train"}, +# } +# GENERATION_TOOL = { +# "backend": {"type": "python", "module": "models.my_model.generator", "function": "generate"}, +# } +TRAINING_TOOL = {} +GENERATION_TOOL = {} diff --git a/adam/model_plugins_builtin/oasis/__init__.py b/adam/model_plugins_builtin/oasis/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..3ca27a18777ab0342f86c3d7c84080c0d3f32d36 --- /dev/null +++ b/adam/model_plugins_builtin/oasis/__init__.py @@ -0,0 +1 @@ +"""Oasis action-conditioned world model plugin.""" diff --git a/adam/model_plugins_builtin/oasis/manifest.py b/adam/model_plugins_builtin/oasis/manifest.py new file mode 100644 index 0000000000000000000000000000000000000000..63ae2dc7a1672ca7bcde69a907974a25165144c6 --- /dev/null +++ b/adam/model_plugins_builtin/oasis/manifest.py @@ -0,0 +1,70 @@ +PLUGIN_ID = "oasis" + +MODEL_INFO = { + "name": "Oasis Action World Model", + "version": "1.0", + "category": "Playable World Models", + "description": "Action-conditioned playable world model trainer for gameplay frame sequences.", + "architecture": "action_conditioned_rectified_flow_video", + "status": "experimental", + "output_type": "playable_world", + "capabilities": ["fresh_training", "resume_training", "playable_inference", "live_preview"], + "input_formats": ["Oasis action dataset folder", "semicolon-separated Oasis dataset folders"], + "output_formats": ["action_flow_model_info.json", "diffusers unet folder", "png preview"], + "hardware": {"recommended_vram_gb": 12, "recommended_system_ram_gb": 32}, + "vram_behavior": { + "scales_with": ["resolution", "batch_size", "sequence_context"], + "estimate": "High; 256x144 with batch 2 is the conservative RTX 3060 starting point.", + }, +} + +TRAINING_SETTINGS = { + "resolution": {"label": "Resolution", "type": "choice", "options": ["128x72", "256x144", "384x216", "512x288"], "default": "256x144", "group": "Basic"}, + "batch_size": {"label": "Batch size", "type": "int", "default": 2, "min": 1, "max": 16, "group": "Basic"}, + "learning_rate": {"label": "Learning rate", "type": "float", "default": 0.00002, "min": 0.0000001, "max": 0.01, "decimals": 7, "step": 0.00001, "group": "Optimization"}, + "workers": {"label": "Loader workers", "type": "int", "default": 2, "min": 0, "max": 8, "group": "Dataset"}, + "mixed_precision": {"label": "Precision", "type": "choice", "options": ["fp32", "fp16", "no"], "default": "fp32", "group": "Optimization"}, + "gradient_accumulation": {"label": "Gradient accumulation", "type": "int", "default": 1, "min": 1, "max": 16, "group": "Optimization"}, + "frame_gap": {"label": "Prediction horizon", "type": "int", "default": 3, "min": 1, "max": 60, "group": "Sequence"}, + "sequence_context": {"label": "Context length", "type": "int", "default": 1, "min": 1, "max": 32, "group": "Sequence"}, + "action_aggregation": {"label": "Action aggregation", "type": "choice", "options": ["window", "mean", "last"], "default": "window", "group": "Sequence"}, + "validation_split": {"label": "Validation split", "type": "float", "default": 0.1, "min": 0.01, "max": 0.5, "decimals": 3, "step": 0.01, "group": "Dataset"}, + "validation_batches": {"label": "Validation batches", "type": "int", "default": 8, "min": 0, "max": 128, "group": "Dataset"}, + "save_every": {"label": "Save every", "type": "int", "default": 5, "min": 1, "max": 1000, "group": "Checkpoints"}, + "preview_enabled": {"label": "Generate previews while training", "type": "bool", "default": True, "group": "Preview"}, + "preview_every": {"label": "Preview interval", "type": "int", "default": 5, "min": 1, "max": 100000, "group": "Preview"}, + "preview_steps": {"label": "Preview steps", "type": "int", "default": 1, "min": 1, "max": 50, "group": "Preview"}, + "seed": {"label": "Random seed", "type": "int", "default": 1234, "min": 0, "max": 2147483647, "group": "Reproducibility"}, + "base_model": {"label": "Base video model", "type": "folder", "default": "", "group": "Checkpoints", "advanced": True}, + "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}, + "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}, + "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}, + "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}, + "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}, + "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}, + "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}, + "gradient_checkpointing": {"label": "Gradient checkpointing", "type": "bool", "default": False, "group": "Advanced", "advanced": True}, + "balance_actions": {"label": "Balance rare actions", "type": "bool", "default": False, "group": "Advanced", "advanced": True}, +} + +GENERATION_SETTINGS = { + "starting_frame": {"label": "Starting frame", "type": "path", "default": "", "group": "Player"}, + "seed": {"label": "Seed", "type": "int", "default": 0, "min": 0, "max": 2147483647, "group": "Player"}, +} + +TRAINING_TOOL = { + "id": "oasis_trainer", + "name": "Oasis Action World Model Trainer", + "backend": {"type": "python", "module": "adam.tools.oasis_adapter", "function": "train_oasis"}, + "capabilities": ["fresh_training", "resume_training", "progress", "pause", "cancel", "live_preview"], +} + +GENERATION_TOOL = { + "id": "oasis_player", + "name": "Oasis Playable Inference", + "model_trainers": ["oasis"], + "arguments": ["model_name", "model_path", "starting_frame", "seed"], + "required_arguments": ["model_path"], + "capabilities": ["playable_inference", "keyboard_actions", "progress", "cancel"], + "backend": {"type": "python", "module": "adam.tools.oasis_adapter", "function": "launch_oasis_player"}, +} diff --git a/adam/model_plugins_builtin/sdxl_lora/__init__.py b/adam/model_plugins_builtin/sdxl_lora/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..7b8c554554ccc2ccee0bd86f4d629460c27a65b2 --- /dev/null +++ b/adam/model_plugins_builtin/sdxl_lora/__init__.py @@ -0,0 +1 @@ +"""SDXL LoRA model plugin.""" diff --git a/adam/model_plugins_builtin/sdxl_lora/manifest.py b/adam/model_plugins_builtin/sdxl_lora/manifest.py new file mode 100644 index 0000000000000000000000000000000000000000..69cb5c15073c434705fbead9b6a0c86dfe70a0c4 --- /dev/null +++ b/adam/model_plugins_builtin/sdxl_lora/manifest.py @@ -0,0 +1,60 @@ +PLUGIN_ID = "lora" + +MODEL_INFO = { + "name": "SDXL LoRA", + "version": "1.0", + "category": "Image Generation", + "description": "Stable Diffusion XL LoRA adapter training and base-model plus adapter generation.", + "architecture": "sdxl_lora", + "status": "stable", + "output_type": "image", + "dependencies": ["diffusers", "safetensors"], + "capabilities": ["fresh_training", "resume_training", "lora_adapter", "image_generation", "reference_image"], + "input_formats": ["captioned image folder", "SDXL checkpoint"], + "output_formats": ["safetensors", "png generation"], + "hardware": {"recommended_vram_gb": 8, "recommended_system_ram_gb": 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."}, +} + +TRAINING_SETTINGS = { + "base_model": {"label": "Base model", "type": "path", "default": "", "required": True, "must_exist": True, "group": "Basic"}, + "trigger_word": {"label": "Trigger word", "type": "text", "default": "", "group": "LoRA"}, + "resolution": {"label": "Resolution", "type": "choice", "options": [512, 768, 1024], "default": 1024, "group": "Basic"}, + "rank": {"label": "Rank", "type": "int", "default": 16, "min": 1, "max": 256, "group": "LoRA"}, + "alpha": {"label": "Alpha", "type": "int", "default": 16, "min": 1, "max": 256, "group": "LoRA"}, + "learning_rate": {"label": "Learning rate", "type": "float", "default": 0.0001, "min": 0.0000001, "max": 0.01, "decimals": 7, "step": 0.00005, "group": "Optimization"}, + "batch_size": {"label": "Batch size", "type": "int", "default": 1, "min": 1, "max": 16, "group": "Basic"}, + "gradient_accumulation_steps": {"label": "Gradient accumulation", "type": "int", "default": 1, "min": 1, "max": 64, "group": "Optimization"}, + "mixed_precision": {"label": "Precision", "type": "choice", "options": ["fp16", "bf16", "no"], "default": "fp16", "group": "Optimization"}, + "caption_extension": {"label": "Caption extension", "type": "choice", "options": [".txt", ".caption"], "default": ".txt", "group": "Dataset"}, + "save_every": {"label": "Save every", "type": "int", "default": 10, "min": 1, "max": 1000, "group": "Checkpoints"}, + "optimizer": {"label": "Optimizer", "type": "choice", "options": ["AdamW", "AdamW8bit"], "default": "AdamW", "group": "Optimization", "advanced": True}, + "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}, + "preview_enabled": {"label": "Generate previews while training", "type": "bool", "default": True, "group": "Preview"}, + "preview_every": {"label": "Preview interval", "type": "int", "default": 5, "min": 1, "max": 100000, "group": "Preview"}, + "preview_prompt": {"label": "Preview prompt", "type": "text", "default": "", "group": "Preview"}, + "preview_seed": {"label": "Preview seed", "type": "int", "default": 123456789, "min": 0, "max": 2147483647, "group": "Preview"}, +} + +GENERATION_SETTINGS = { + "base_model_path": {"label": "Base model", "type": "path", "default": "", "required": True, "must_exist": True, "group": "Model Loading"}, + "model_path": {"label": "LoRA adapter", "type": "path", "default": "", "required": True, "must_exist": True, "group": "Model Loading"}, + "prompt": {"label": "Prompt", "type": "multiline_text", "default": "", "required": True, "group": "Prompt"}, + "negative_prompt": {"label": "Negative prompt", "type": "multiline_text", "default": "", "group": "Prompt"}, + "lora_strength": {"label": "LoRA strength", "type": "float", "default": 1.0, "min": 0.0, "max": 2.0, "decimals": 2, "step": 0.05, "group": "Generation"}, + "image_count": {"label": "Images", "type": "int", "default": 1, "min": 1, "max": 48, "group": "Generation"}, + "steps": {"label": "Steps", "type": "int", "default": 30, "min": 1, "max": 150, "group": "Generation"}, + "cfg_scale": {"label": "CFG scale", "type": "float", "default": 7.0, "min": 0.1, "max": 30.0, "decimals": 2, "step": 0.5, "group": "Generation"}, + "sampler": {"label": "Sampler", "type": "choice", "options": ["DPM++ 2M", "DPM++ SDE", "Euler", "Euler a", "DDIM"], "default": "DPM++ 2M", "group": "Generation"}, + "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"}, + "width": {"label": "Width", "type": "int", "default": 1024, "min": 256, "max": 2048, "group": "Generation"}, + "height": {"label": "Height", "type": "int", "default": 1024, "min": 256, "max": 2048, "group": "Generation"}, + "seed": {"label": "Seed", "type": "int", "default": 0, "min": 0, "max": 2147483647, "group": "Generation"}, + "reference_image": {"label": "Reference image", "type": "path", "default": "", "group": "Reference"}, + "denoise_strength": {"label": "Denoise strength", "type": "float", "default": 0.45, "min": 0.0, "max": 1.0, "decimals": 2, "step": 0.05, "group": "Reference"}, + "prompt_weighting": {"label": "Use prompt weights", "type": "bool", "default": True, "group": "Advanced", "advanced": True}, + "preview_interval": {"label": "Steps per preview", "type": "int", "default": 0, "min": 0, "max": 500, "group": "Preview"}, +} + +TRAINING_TOOL = {"id": "lora_trainer"} +GENERATION_TOOL = {"id": "lora_generator"} diff --git a/adam/model_profiles.py b/adam/model_profiles.py new file mode 100644 index 0000000000000000000000000000000000000000..5ed609ef9acd48f15fd5ee8d4a5ef8dd182934e5 --- /dev/null +++ b/adam/model_profiles.py @@ -0,0 +1,93 @@ +from __future__ import annotations + +from dataclasses import asdict, dataclass, field +from typing import Any + +from adam.model_plugins import ModelPlugin, ModelPluginRegistry + + +@dataclass(frozen=True, slots=True) +class ModelProfile: + """Normalized model profile built from ADAM's plugin manifests.""" + + id: str + name: str + category: str + architecture: str + version: str + description: str + status: str = "experimental" + output_type: str = "image" + training: dict[str, dict[str, Any]] = field(default_factory=dict) + generation: dict[str, dict[str, Any]] = field(default_factory=dict) + trainer_module: str = "" + generator_module: str = "" + trainer_tool: str = "" + generator_tool: str = "" + capabilities: list[str] = field(default_factory=list) + hardware: dict[str, Any] = field(default_factory=dict) + vram_behavior: dict[str, Any] = field(default_factory=dict) + input_formats: list[str] = field(default_factory=list) + output_formats: list[str] = field(default_factory=list) + + def to_dict(self) -> dict[str, Any]: + return asdict(self) + + +def profile_from_plugin(plugin: ModelPlugin) -> ModelProfile: + info = dict(plugin.info) + training_tool = dict(plugin.training_tool) + generation_tool = dict(plugin.generation_tool) + capabilities = list( + dict.fromkeys( + [ + *info.get("capabilities", []), + *training_tool.get("capabilities", []), + *generation_tool.get("capabilities", []), + ] + ) + ) + trainer_backend = dict(training_tool.get("backend", {})) + generator_backend = dict(generation_tool.get("backend", {})) + return ModelProfile( + id=plugin.id, + name=str(info.get("name", plugin.name)), + category=str(info.get("category", "")), + architecture=str(info.get("architecture", plugin.id)), + version=str(info.get("version", "")), + description=str(info.get("description", "")), + status=str(info.get("status", "experimental")), + output_type=str(info.get("output_type", "image")), + training=plugin.training_settings, + generation=plugin.generation_settings, + trainer_module=str(trainer_backend.get("module", "")), + generator_module=str(generator_backend.get("module", "")), + trainer_tool=plugin.trainer_id if plugin.training_settings else "", + generator_tool=plugin.generator_id if plugin.generation_settings else "", + capabilities=capabilities, + hardware=dict(info.get("hardware", {})), + vram_behavior=dict(info.get("vram_behavior", {})), + input_formats=list(info.get("input_formats", [])), + output_formats=list(info.get("output_formats", [])), + ) + + +class ModelProfileRegistry: + """Read-only view over plugin manifests for UI and automation features.""" + + def __init__(self, plugins: ModelPluginRegistry) -> None: + self.plugins = plugins + + def all(self) -> list[ModelProfile]: + return [ + profile_from_plugin(plugin) + for plugin in self.plugins.all() + if plugin.info.get("category") != "Template" + ] + + def get(self, profile_id: str) -> ModelProfile | None: + plugin = self.plugins.by_trainer(profile_id) + return profile_from_plugin(plugin) if plugin else None + + def as_catalog(self) -> list[dict[str, Any]]: + return [profile.to_dict() for profile in self.all()] diff --git a/adam/models.py b/adam/models.py index d7d5c79134729c83c6223a2ddb8010c80b9cf21a..a1c9e4694436f797423c1340f6e6ecede45a46c5 100644 --- a/adam/models.py +++ b/adam/models.py @@ -14,6 +14,7 @@ def utc_now() -> str: class JobStatus(str, Enum): DRAFT = "Draft" AWAITING_CONFIRMATION = "Awaiting confirmation" + SCHEDULED = "Scheduled" QUEUED = "Queued" RUNNING = "Running" PAUSED = "Paused" @@ -73,6 +74,7 @@ class Job: progress: int = 0 current_step: int = -1 created_at: str = field(default_factory=utc_now) + scheduled_for: str | None = None started_at: str | None = None ended_at: str | None = None output_folder: str | None = None @@ -89,8 +91,15 @@ class Job: preview_total: int = 0 preview_image_index: int = 0 preview_image_count: int = 0 + eta_seconds: int | None = None + estimated_completion_at: str | None = None + progress_current: int = 0 + progress_total: int = 0 + progress_rate: float = 0.0 + progress_unit: str = "step" atlas_report: dict[str, Any] = field(default_factory=dict) nova_report: dict[str, Any] = field(default_factory=dict) + metadata: dict[str, Any] = field(default_factory=dict) def to_dict(self) -> dict[str, Any]: payload = asdict(self) diff --git a/adam/oasis_dataset.py b/adam/oasis_dataset.py new file mode 100644 index 0000000000000000000000000000000000000000..f7e8910e76d20516a489c8ccc6f5b52aa262e155 --- /dev/null +++ b/adam/oasis_dataset.py @@ -0,0 +1,215 @@ +from __future__ import annotations + +import json +import re +from dataclasses import dataclass, field +from pathlib import Path +from typing import Any + +from PIL import Image + + +BINARY_ACTIONS = { + "w", "a", "s", "d", "jump", "arrow_up", "arrow_left", "arrow_down", + "arrow_right", "enter", "shift", "ctrl", "alt", "tab", "escape", + "q", "e", "r", "f", "z", "x", "c", "v", "key_1", "key_2", "key_3", + "key_4", "mouse_left", "mouse_middle", "mouse_right", +} +CONTINUOUS_ACTIONS = {"mouse_dx", "mouse_dy", "zoom"} +DERIVED_ACTIONS = { + "move_x", "move_y", "right_mouse", "camera_active", + "camera_yaw_delta_degrees", "camera_pitch_delta_degrees", + "mouse_raw_dx", "mouse_raw_dy", +} +SUPPORTED_ACTIONS = BINARY_ACTIONS | CONTINUOUS_ACTIONS | DERIVED_ACTIONS +REQUIRED_CANONICAL_ACTIONS = {"w", "a", "s", "d", "jump"} +IMAGE_EXTENSIONS = {".jpg", ".jpeg", ".png", ".webp", ".bmp"} +METADATA_FIELDS = { + "session_id", "session_started_at", "frame_index", "filename", + "timestamp_seconds", "camera_encoding", +} + + +@dataclass(slots=True) +class OasisDatasetReport: + dataset_folders: list[str] = field(default_factory=list) + frames: int = 0 + metadata_rows: int = 0 + valid_rows: int = 0 + valid_transitions: int = 0 + sessions: int = 0 + resolution: str = "" + action_counts: dict[str, int] = field(default_factory=dict) + errors: list[str] = field(default_factory=list) + warnings: list[str] = field(default_factory=list) + + @property + def ok(self) -> bool: + return not self.errors + + +def dataset_directories(value: str | list[str] | tuple[str, ...]) -> list[Path]: + entries = value if isinstance(value, (list, tuple)) else str(value or "").split(";") + directories: list[Path] = [] + for entry in entries: + text = str(entry).strip().strip('"') + if not text: + continue + path = Path(text).expanduser() + if path not in directories: + directories.append(path) + return directories + + +def _numeric_frame_index(path: Path) -> int | None: + match = re.search(r"frame_(\d+)", path.stem, re.I) or re.search(r"(\d+)", path.stem) + return int(match.group(1)) if match else None + + +def validate_oasis_dataset(value: str | list[str] | tuple[str, ...], *, frame_gap: int = 1) -> OasisDatasetReport: + report = OasisDatasetReport() + frame_gap = max(1, int(frame_gap)) + directories = dataset_directories(value) + if not directories: + report.errors.append("Select at least one Oasis action dataset folder.") + return report + seen_resolution: tuple[int, int] | None = None + action_counts = {name: 0 for name in sorted(REQUIRED_CANONICAL_ACTIONS | CONTINUOUS_ACTIONS)} + transition_total = 0 + session_ids: set[str] = set() + + for directory in directories: + resolved = directory.resolve() + report.dataset_folders.append(str(resolved)) + if not directory.is_dir(): + report.errors.append(f"Dataset folder does not exist: {directory}") + continue + frames_dir = directory / "frames" + actions_path = directory / "actions.jsonl" + if not frames_dir.is_dir(): + report.errors.append(f"{directory} is missing a frames folder.") + continue + if not actions_path.is_file(): + report.errors.append(f"{directory} is missing actions.jsonl.") + continue + frame_files = sorted(path for path in frames_dir.iterdir() if path.is_file() and path.suffix.casefold() in IMAGE_EXTENSIONS) + report.frames += len(frame_files) + if not frame_files: + report.errors.append(f"{directory} has no image frames.") + indexed_frames = [index for index in (_numeric_frame_index(path) for path in frame_files) if index is not None] + frame_paths_by_index: dict[int, list[Path]] = {} + for frame_file in frame_files: + index = _numeric_frame_index(frame_file) + if index is not None: + frame_paths_by_index.setdefault(index, []).append(frame_file) + if indexed_frames: + gaps = [ + (left, right) for left, right in zip(indexed_frames, indexed_frames[1:]) + if right != left + 1 + ] + if gaps: + report.warnings.append(f"{directory} has frame ordering gaps such as {gaps[0][0]} to {gaps[0][1]}; invalid transitions will be skipped.") + rows_by_session: dict[str, list[dict[str, Any]]] = {} + seen_keys: set[tuple[str, int]] = set() + for line_number, line in enumerate(actions_path.read_text(encoding="utf-8").splitlines(), 1): + line = line.strip() + if not line: + continue + report.metadata_rows += 1 + try: + row = json.loads(line) + except json.JSONDecodeError: + report.errors.append(f"{actions_path.name} line {line_number} is not valid JSON.") + continue + filename = str(row.get("filename", "")).strip() + if not filename: + report.errors.append(f"{actions_path.name} line {line_number} has no frame filename.") + continue + frame_path = frames_dir / filename + if not frame_path.is_file(): + try: + frame_index = int(row.get("frame_index")) + except (TypeError, ValueError): + report.warnings.append( + f"{actions_path.name} line {line_number} points to missing frame {filename}; skipping row." + ) + continue + candidates = frame_paths_by_index.get(frame_index, []) + if len(candidates) == 1: + frame_path = candidates[0] + report.warnings.append( + f"{actions_path.name} line {line_number} uses {frame_path.name} for missing legacy filename {filename}." + ) + else: + report.warnings.append( + f"{actions_path.name} line {line_number} points to missing frame {filename}; skipping row." + ) + continue + try: + with Image.open(frame_path) as image: + image.verify() + with Image.open(frame_path) as image: + size = image.size + except Exception as exc: + report.errors.append(f"Broken image file {frame_path.name}: {exc}") + continue + if seen_resolution is None: + seen_resolution = size + report.resolution = f"{size[0]}x{size[1]}" + elif size != seen_resolution: + report.errors.append(f"Inconsistent frame resolution: {frame_path.name} is {size[0]}x{size[1]}, expected {seen_resolution[0]}x{seen_resolution[1]}.") + unexpected = sorted(set(row) - SUPPORTED_ACTIONS - METADATA_FIELDS) + if unexpected: + report.errors.append(f"{actions_path.name} line {line_number} contains unsupported action field(s): {', '.join(unexpected[:6])}.") + missing = sorted(name for name in REQUIRED_CANONICAL_ACTIONS if name not in row) + if missing: + report.errors.append(f"{actions_path.name} line {line_number} is missing action label(s): {', '.join(missing)}.") + continue + try: + frame_index = int(row.get("frame_index")) + except (TypeError, ValueError): + report.errors.append(f"{actions_path.name} line {line_number} has an invalid frame_index.") + continue + session_id = str(row.get("session_id") or f"legacy-{directory.name}").strip() + key = (session_id, frame_index) + if key in seen_keys: + report.errors.append(f"{actions_path.name} repeats frame_index {frame_index} in session {session_id}.") + continue + seen_keys.add(key) + for name in action_counts: + try: + value = float(row.get(name, 0)) + except (TypeError, ValueError): + report.errors.append(f"{actions_path.name} line {line_number} has invalid {name} action value.") + value = 0.0 + if abs(value) > (0.5 if name in BINARY_ACTIONS else 0.02): + action_counts[name] += 1 + row["_session_id"] = session_id + rows_by_session.setdefault(session_id, []).append(row) + session_ids.add(f"{resolved}:{session_id}") + report.valid_rows += 1 + for session_id, rows in rows_by_session.items(): + if not rows: + report.errors.append(f"{directory} has an empty sequence {session_id}.") + continue + rows.sort(key=lambda item: int(item["frame_index"])) + for left, right in zip(rows, rows[1:]): + if int(right["frame_index"]) != int(left["frame_index"]) + 1: + report.warnings.append(f"{directory} session {session_id} has an ordering gap at frame {left['frame_index']}; invalid transitions will be skipped.") + transition_total += sum( + 1 for left, right in zip(rows, rows[frame_gap:]) + if int(right["frame_index"]) == int(left["frame_index"]) + frame_gap + ) + + report.sessions = len(session_ids) + report.action_counts = action_counts + report.valid_transitions = transition_total + if report.metadata_rows != report.frames: + report.warnings.append(f"Frame and label counts differ: {report.frames} frame files, {report.metadata_rows} action rows.") + if report.valid_rows < 2: + report.errors.append("The dataset needs at least two valid labelled frames.") + if report.valid_transitions < 1: + report.errors.append(f"No valid frame transitions were found for prediction horizon {frame_gap}.") + if report.valid_rows and not any(action_counts.values()): + report.errors.append("No non-idle action labels were found. Record idle plus at least one active control.") + return report diff --git a/adam/orion.py b/adam/orion.py index 17774343d2e28453d19bf9d7f883488973ca982c..6a372356260164dc00f79d38f3864c231c38970f 100644 --- a/adam/orion.py +++ b/adam/orion.py @@ -10,6 +10,8 @@ IMAGE_EXTENSIONS = {".jpg", ".jpeg", ".png", ".webp", ".bmp"} def dataset_image_count(raw_path: object) -> int: + if not str(raw_path or "").strip(): + return 0 path = Path(str(raw_path or "")).expanduser() if not path.is_dir(): return 0 @@ -36,6 +38,19 @@ def _available_vram_gb() -> float | None: return None +def _resolution_extent(value: object, fallback: int = 256) -> int: + raw = str(value or fallback).strip().lower() + if "x" in raw: + try: + return max(int(part.strip()) for part in raw.split("x", 1)) + except ValueError: + return fallback + try: + return int(raw) + except ValueError: + return fallback + + def recommend_training_settings( trainer: str, image_count: int, @@ -46,7 +61,7 @@ def recommend_training_settings( """Return an explainable, conservative starting recipe for manual review.""" trainer = str(trainer).casefold() images = max(10, int(image_count)) - resolution = max(64, min(512, int(resolution))) + resolution = max(64, min(512, _resolution_extent(resolution))) vram = _available_vram_gb() if vram_gb is None else vram_gb cpu_workers = max(2, min(8, (os.cpu_count() or 4) // 2)) @@ -131,7 +146,7 @@ def review_training_plan(plan: Any) -> dict[str, Any]: 1, int(args.get("gradient_accumulation_steps", args.get("gradient_accumulation", 1)) or 1), ) - resolution = max(64, int(args.get("resolution", 256) or 256)) + resolution = max(64, _resolution_extent(args.get("resolution", 256))) exposures = images * epochs if images else 0 optimizer_steps = math.ceil(images / batch / accumulation) * epochs if images else 0 total_steps += optimizer_steps diff --git a/adam/planner.py b/adam/planner.py index 5fbeed14457d6db6fe9a4d61b911277732eda51e..cb6989b724a9b5924c6cc7acb688af28058b9260 100644 --- a/adam/planner.py +++ b/adam/planner.py @@ -43,6 +43,22 @@ def _project_name(subject: str, suffix: str) -> str: return f"{safe.title()} {suffix}".strip()[:64] +def _trainer_label(trainer: str) -> str: + return { + "ddpm": "DDPM", + "flow": "Flow Matching", + "lora": "LoRA", + "oasis": "Oasis Action World Model", + }.get(trainer, trainer.replace("_", " ").title()) + + +def _friendly_model_name(asset: Asset) -> str: + name = str(asset.name or "").strip() + if re.search(r"^[A-Za-z]:[\\/]", name) or "/" in name or "\\" in name: + return Path(asset.path).name + return name or Path(asset.path).name + + def _collection_mode(request: str) -> str: """Return the user's requested stopping rule for internet collection.""" return ( @@ -318,6 +334,10 @@ class Planner: def _deterministic_plan(self, request: str) -> ExecutionPlan | None: lowered = request.lower() + oasis_player = self._oasis_player_plan(request) + if oasis_player: + return oasis_player + youtube_plan = self._youtube_dataset_plan(request) if youtube_plan: return youtube_plan @@ -556,10 +576,17 @@ class Planner: if not re.search(r"\b(train|fine[- ]?tune|retrain|continue|resume)\b", lowered): return None fine_tune_payload = self._fine_tune_payload(request) + marker = re.search(r"\[ADAM_TRAINER:([A-Za-z0-9_ -]+)\]", request, re.I) trainer = str(fine_tune_payload.get("trainer", "")) or ( + marker.group(1).strip().casefold().replace(" ", "_") if marker else "" + ) or ( "lora" if re.search(r"\blora\b", lowered) else "ddpm" if re.search(r"\bddpm\b", lowered) else "flow" if re.search(r"\bflow(?:\s+matching)?\b", lowered) + else "oasis" if re.search( + r"\b(oasis|action[- ]conditioned|playable\s+ai\s+games?|world\s+models?|gameplay[- ]frame|wasd|w/a/s/d)\b", + lowered, + ) else "" ) action = ( @@ -574,7 +601,7 @@ class Planner: model_query = "" resume_match = re.search( r"\b(?:fine[- ]?tune|retrain|continue|resume)\s+(?:the\s+)?(.+?)" - r"(?:\s+model)?\s+(?:from|on|with)\s+(?:the\s+)?(?:ddpm|lora)\b", + r"(?:\s+model)?\s+(?:from|on|with)\s+(?:the\s+)?(?:ddpm|lora|oasis)\b", request, re.I, ) @@ -619,7 +646,7 @@ class Planner: if action == "resume_training": candidates: list[Asset] = [] if model_query: - candidates = self.assets.find("model", model_query, trainer=trainer) + candidates = self._model_candidates(model_query, trainer=trainer) if not candidates: return ExecutionPlan( request=request, @@ -642,7 +669,13 @@ class Planner: trainer = trainer or model.trainer ddpm_pipeline = trainer == "ddpm" and (Path(model.path) / "model_index.json").is_file() flow_model = trainer == "flow" and self._valid_flow_model(Path(model.path)) - if (not model.checkpoint or not Path(model.checkpoint).exists()) and not ddpm_pipeline and not flow_model: + oasis_model = trainer == "oasis" and self._valid_oasis_model(Path(model.path)) + if ( + (not model.checkpoint or not Path(model.checkpoint).exists()) + and not ddpm_pipeline + and not flow_model + and not oasis_model + ): return ExecutionPlan( request=request, summary=( @@ -653,10 +686,10 @@ class Planner: ), steps=[], project_name="Resume training", - ) + ) dataset_mode = str(fine_tune_payload.get("dataset_mode", "original")) if dataset_mode == "existing": - dataset = self._asset_dataset(str(fine_tune_payload.get("dataset_name", ""))) + dataset = self._resolve_dataset_asset(str(fine_tune_payload.get("dataset_name", ""))) else: dataset = self._dataset_for_model(model) if dataset_mode == "new": @@ -677,21 +710,25 @@ class Planner: steps=[], project_name="Resume training", ) + resumed_model_name = _friendly_model_name(model) command = TrainingCommand.from_dict( { "action": "resume_training", "trainer": trainer, "dataset": dataset.path, - "model_name": model.name, + "model_name": resumed_model_name, "epochs": epochs, "output": ( - str(self._training_output(trainer, f"{model.name} Fine Tune") or model.path) + str(self._training_output(trainer, f"{resumed_model_name} Fine Tune") or model.path) if trainer == "flow" else model.path ), # The DDPM adapter can safely branch from a complete pipeline when # its exact Accelerate checkpoint has been cleaned up. "resume_from": model.checkpoint or model.path, - "base_model": self._lora_base_model() if trainer == "lora" else "", + "base_model": ( + str(training_options.get("base_model") or self._lora_base_model()) + if trainer == "lora" else "" + ), "training_options": training_options, } ) @@ -731,7 +768,10 @@ class Planner: "model_name": model_name, "epochs": epochs, "output": str(output), - "base_model": self._lora_base_model() if trainer == "lora" else "", + "base_model": ( + str(training_options.get("base_model") or self._lora_base_model()) + if trainer == "lora" else "" + ), "training_options": training_options, } ) @@ -787,23 +827,24 @@ class Planner: if dataset_dir.exists(): dataset_dir = dataset_dir.with_name(f"{dataset_dir.name} {datetime.now().strftime('%Y%m%d_%H%M%S')}") image_count = max(10, min(int(payload.get("image_count", 60)), 5000)) + model_name = _friendly_model_name(model) arguments: dict[str, Any] = { - "dataset_dir": str(dataset_dir), "model_name": model.name, + "dataset_dir": str(dataset_dir), "model_name": model_name, "epochs": epochs, "output_dir": ( - str(self._training_output(trainer, f"{model.name} Fine Tune") or model.path) + str(self._training_output(trainer, f"{model_name} Fine Tune") or model.path) if trainer == "flow" else model.path ), "resume_from": model.checkpoint or model.path, **training_options, } if trainer == "lora": - base_model = self._lora_base_model() + base_model = str(training_options.get("base_model") or self._lora_base_model()) if not base_model or not Path(base_model).is_file(): return ExecutionPlan(request=request, summary="Choose a valid SDXL base model in the LoRA app before fine-tuning.", steps=[], project_name="LoRA training") arguments["base_model"] = base_model return ExecutionPlan( request=request, - summary=f"Collect {image_count} new images for {subject}, then continue {model.name} for {epochs} additional epochs.", + summary=f"Collect {image_count} new images for {subject}, then continue {model_name} for {epochs} additional epochs.", steps=[ 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)}), 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: def _dataset_for_phrase(self, phrase: str) -> Asset | None: """Resolve a friendly dataset phrase, preferring the shortest clear folder match.""" - direct = self._asset_dataset(phrase) + direct = self._resolve_dataset_asset(phrase) if direct: return direct wanted = re.sub(r"[^a-z0-9]+", " ", phrase.casefold()).strip() @@ -900,6 +941,36 @@ class Planner: candidates.sort(key=lambda item: (len(item.name), item.name.casefold())) return candidates[0] + def _model_candidates(self, query: str, *, trainer: str = "") -> list[Asset]: + raw_query = str(query).strip().strip('"').replace("\\_", "_") + path = Path(raw_query).expanduser() + if path.exists(): + resolved = path.resolve() + matches = [ + asset for asset in self.assets.assets + if asset.kind == "model" + and (not trainer or asset.trainer == trainer) + and Path(asset.path).expanduser().resolve() == resolved + ] + if matches: + return matches + path_like = re.search(r"^[A-Za-z]:[\\/]", raw_query) or "/" in raw_query or "\\" in raw_query + if path_like and path.name: + matches = self.assets.find("model", path.name, trainer=trainer) + existing = [asset for asset in matches if Path(asset.path).expanduser().exists()] + if trainer == "oasis": + valid = [ + asset for asset in existing + if self._valid_oasis_model(Path(asset.path).expanduser()) + ] + if valid: + return valid + if existing: + return existing + if matches: + return matches + return self.assets.find("model", raw_query, trainer=trainer) + def _plan_training_command( self, request: str, @@ -917,25 +988,48 @@ class Planner: steps=[], project_name="Unsupported training request", ) - dataset_path = Path(command.dataset).expanduser() - if not dataset_path.is_dir(): - raise PlanningError("The validated training dataset does not exist.") + if command.trainer == "oasis": + dataset_paths = self._oasis_dataset_paths(command.dataset) + if not dataset_paths: + raise PlanningError( + "The validated Oasis dataset does not exist or contains no action dataset folders." + ) + dataset_path = dataset_paths[0] + resolved_dataset_paths = [str(path.resolve()) for path in dataset_paths] + dataset_argument = ( + resolved_dataset_paths[0] + if len(resolved_dataset_paths) == 1 + else resolved_dataset_paths + ) + dataset_label = ";".join(resolved_dataset_paths) + else: + dataset_path = Path(command.dataset).expanduser() + if not dataset_path.is_dir(): + raise PlanningError("The validated training dataset does not exist.") + dataset_argument = str(dataset_path.resolve()) + dataset_label = dataset_argument trainer_folder = self._configured_tool_folder(tool_id) - if not trainer_folder: + plugin = self.registry.model_plugins.by_trainer(command.trainer) + custom_plugin = plugin is not None and command.trainer not in {"ddpm", "flow", "lora", "oasis"} + if not trainer_folder and not custom_plugin: raise PlanningError(f"The {spec.name} folder is not connected.") - output_folder = "output_flow_models" if command.trainer == "flow" else "output" - output_root = (Path(trainer_folder) / output_folder).resolve() output_path = Path(command.output).expanduser().resolve() - try: - output_path.relative_to(output_root) - except ValueError as exc: - raise PlanningError( - f"{spec.name} outputs must stay inside {output_root}." - ) from exc + if custom_plugin: + output_root = (self.root / "data" / "model_plugin_outputs" / command.trainer).resolve() + output_root.mkdir(parents=True, exist_ok=True) + else: + output_folder = self._trainer_output_folder(command.trainer) + output_root = (Path(trainer_folder) / output_folder).resolve() + try: + output_path.relative_to(output_root) + except ValueError as exc: + raise PlanningError( + f"{spec.name} outputs must stay inside {output_root}." + ) from exc if command.resume_from and not Path(command.resume_from).exists(): raise PlanningError("The validated resume checkpoint does not exist.") arguments: dict[str, Any] = { - "dataset_dir": str(dataset_path.resolve()), + "dataset_dir": dataset_argument, "model_name": command.model_name, "epochs": command.epochs, "output_dir": str(output_path), @@ -944,31 +1038,37 @@ class Planner: if command.resume_from: arguments["resume_from"] = command.resume_from if command.trainer == "lora": - if not command.base_model or not Path(command.base_model).is_file(): + base_model = command.base_model or str((command.training_options or {}).get("base_model", "")) + if not base_model or not Path(base_model).is_file(): return ExecutionPlan( request=request, summary=( - "I found the LoRA dataset, but the connected LoRA trainer has no " - "valid SDXL base model selected. Choose one in the LoRA app first." + "I found the LoRA dataset, but no valid SDXL base model is selected. " + "Choose one in the generated LoRA settings before training." ), steps=[], project_name="LoRA training", ) - arguments["base_model"] = command.base_model + arguments["base_model"] = base_model + arguments["trigger_word"] = ( + command.trigger_word + or str((command.training_options or {}).get("trigger_word") or "") + or command.model_name + ) verb = "Continue" if command.action == "resume_training" else "Train" epoch_kind = "additional epochs" if command.action == "resume_training" else "epochs" return ExecutionPlan( request=request, summary=( - f"{verb} {command.model_name} with the registered {command.trainer.upper()} " - f"trainer for {command.epochs} {epoch_kind}. Dataset: {command.dataset}. " + f"{verb} {command.model_name} with the registered {_trainer_label(command.trainer)} " + f"trainer for {command.epochs} {epoch_kind}. Dataset: {dataset_label}. " f"Output: {command.output}." + (f" Training options: {command.training_options}." if command.training_options else "") ), steps=[ PlanStep( tool_id, - f"{verb} {command.trainer.upper()} model", + f"{verb} {_trainer_label(command.trainer)} model", "Launch the connected trainer with validated paths and stream progress.", arguments, ) @@ -987,7 +1087,7 @@ class Planner: # causing a valid Flow Matching request to fall back to its old # clarification screen instead of creating a training plan. explicit_path = re.search( - r"\bfrom\s+(?:the\s+)?([A-Za-z]:[\\/].+?)\s+dataset\s*" + r"\b(?:from|use)\s+(?:the\s+)?([A-Za-z]:[\\/].+?)\s+dataset\s*" r"(?=[,.;]?\s*(?:train|continue|resume|name|call|save|output|put)\b)", request, re.I, @@ -996,6 +1096,7 @@ class Planner: return explicit_path.group(1).strip() patterns = ( r"\btrain\s+(?:the\s+)?(.+?)\s+dataset\s+(?:on|with|for)\b", + r"\buse\s+(?:the\s+)?(.+?)\s+dataset\b", r"\bfrom\s+(?:the\s+)?(.+?)\s+dataset\b", r"\bwith\s+(?:the\s+)?(.+?)\s+dataset\b", r"\b(?:the\s+)?(.+?)\s+dataset\s*,?\s+(?:train|use)\b", @@ -1015,6 +1116,9 @@ class Planner: # metadata before looking for the user-facing name. request = re.sub(r"\s*\[ADAM_TRAINING_OPTIONS:\{.*?\}\]", "", request, flags=re.I | re.S) match = re.search(r"\b(?:name|call)\s+(?:the\s+)?model\s+(.+?)(?:[,\[\{]|$)", request, re.I) + if match: + return _clean_subject(match.group(1)) + match = re.search(r"\b(?:model|checkpoint)\s+called\s+(.+?)(?:[,\[\{]|$)", request, re.I) return _clean_subject(match.group(1)) if match else "" def _asset_dataset(self, name: str) -> Asset | None: @@ -1047,19 +1151,68 @@ class Planner: return ranked[0] return None + def _resolve_dataset_asset(self, value: str) -> Asset | None: + path = self._resolve_dataset(value) + if path: + return self.assets.register( + kind="dataset", + name=path.name, + path=str(path), + persist=False, + ) + return None + + @staticmethod + def _is_oasis_dataset_folder(path: Path) -> bool: + return path.is_dir() and (path / "frames").is_dir() and (path / "actions.jsonl").is_file() + + def _oasis_dataset_paths(self, value: str) -> list[Path]: + paths: list[Path] = [] + for raw in str(value or "").split(";"): + text = raw.strip().strip('"') + if not text: + continue + path = Path(text).expanduser() + if self._is_oasis_dataset_folder(path): + paths.append(path.resolve()) + continue + if path.is_dir(): + children = [ + child.resolve() + for child in sorted(path.rglob("*"), key=lambda item: str(item).casefold()) + if self._is_oasis_dataset_folder(child) + ] + paths.extend(children) + unique: list[Path] = [] + for path in paths: + if path not in unique: + unique.append(path) + return unique + def _training_output(self, trainer: str, model_name: str) -> Path | None: folder = self._configured_tool_folder(f"{trainer}_trainer") - if not folder: + plugin = self.registry.model_plugins.by_trainer(trainer) + if not folder and plugin is None: return None safe = re.sub(r"[^A-Za-z0-9._-]+", "_", model_name).strip("._") or "model" - output_root = "output_flow_models" if trainer == "flow" else "output" - candidate = (Path(folder) / output_root / safe).resolve() + if folder: + output_root = self._trainer_output_folder(trainer) + candidate = (Path(folder) / output_root / safe).resolve() + else: + candidate = (self.root / "data" / "model_plugin_outputs" / trainer / safe).resolve() if candidate.exists(): candidate = candidate.with_name( f"{candidate.name}_{datetime.now().strftime('%Y%m%d_%H%M%S')}" ) return candidate + @staticmethod + def _trainer_output_folder(trainer: str) -> str: + return { + "flow": "output_flow_models", + "oasis": "output_action_flow_models", + }.get(trainer, "output") + @staticmethod def _valid_flow_model(folder: Path) -> bool: try: @@ -1145,12 +1298,22 @@ class Planner: @staticmethod def _training_options_from_request(request: str) -> dict[str, Any]: match = re.search(r"\[ADAM_TRAINING_OPTIONS:(\{.*?\})\]", request, re.S) - if not match: - return {} - try: - options = json.loads(match.group(1)) - except json.JSONDecodeError as exc: - raise PlanningError("Training options could not be read safely.") from exc + options: dict[str, Any] = {} + if match: + try: + parsed = json.loads(match.group(1)) + except json.JSONDecodeError as exc: + raise PlanningError("Training options could not be read safely.") from exc + if not isinstance(parsed, dict): + raise PlanningError("Training options must be a settings object.") + options = parsed + trigger_match = re.search( + r"\b(?:trigger\s+word|concept\s+token)(?:\s+of|\s*=|\s*:)?\s*['\"\u201c\u201d]?([A-Za-z0-9_.-]{1,128})", + request, + re.I, + ) + if trigger_match and "trigger_word" not in options: + options["trigger_word"] = trigger_match.group(1).strip() if not isinstance(options, dict): raise PlanningError("Training options must be a settings object.") return options @@ -1475,6 +1638,82 @@ class Planner: project_name="DDPM training", ) + def _oasis_player_plan(self, request: str) -> ExecutionPlan | None: + lowered = request.casefold() + if not re.search(r"\b(launch|start|play|open|run)\b", lowered): + return None + if not re.search(r"\b(oasis|playable\s+ai\s+game|action\s+player|world\s+model)\b", lowered): + return None + self.assets.discover(self.config) + path_match = re.search(r"([A-Za-z]:[\\/][^,\n]+)", request) + model_path = Path(path_match.group(1).strip().strip("\"'")) if path_match else None + model_name = "" + if model_path is not None and not self._valid_oasis_model(model_path): + return ExecutionPlan( + request=request, + summary="I need a valid Oasis action model folder to launch the playable window.", + steps=[], + project_name="Oasis player", + ) + if model_path is None: + query_text = re.sub(r"\b(?:with\s+)?seed\s+\d+\b", " ", request, flags=re.I) + query_text = re.sub(r"\b(?:starting|reference)\s+frame\s*(?:is|:|=)?\s*[A-Za-z]:[\\/][^,\n]+", " ", query_text, flags=re.I) + query = _clean_subject( + re.sub(r"\b(launch|start|play|open|run|oasis|action player|world model)\b", " ", query_text, flags=re.I) + ) + candidates = self.assets.find("model", query, trainer="oasis") if query else [ + asset for asset in self.assets.assets if asset.kind == "model" and asset.trainer == "oasis" + ] + candidates = [asset for asset in candidates if Path(asset.path).is_dir()] + if len(candidates) != 1: + examples = ", ".join(asset.name for asset in candidates[:4]) + return ExecutionPlan( + request=request, + summary=( + "I need one Oasis checkpoint folder to launch." + + (f" Matching models: {examples}." if examples else "") + ), + steps=[], + project_name="Oasis player", + ) + model_name = candidates[0].name + model_path = Path(candidates[0].path) + starting_frame = "" + start_match = re.search(r"\b(?:starting|reference)\s+frame\s*(?:is|:|=)?\s*([A-Za-z]:[\\/][^,\n]+)", request, re.I) + if start_match: + starting_frame = start_match.group(1).strip().strip("\"'") + seed_match = re.search(r"\bseed\s+(\d+)", request, re.I) + return ExecutionPlan( + request=request, + summary=f"Launch the Oasis playable window for {model_name or model_path.name}.", + steps=[ + PlanStep( + "oasis_player", + "Launch Oasis player", + "Open the existing Oasis playable inference window in its own process.", + { + "model_name": model_name or model_path.name, + "model_path": str(model_path.resolve()), + "starting_frame": starting_frame, + "seed": int(seed_match.group(1)) if seed_match else 0, + }, + ) + ], + requires_confirmation=False, + project_name="Oasis player", + ) + + @staticmethod + def _valid_oasis_model(folder: Path) -> bool: + try: + metadata = json.loads((folder / "action_flow_model_info.json").read_text(encoding="utf-8")) + return ( + metadata.get("model_type") == "action_conditioned_rectified_flow_video" + and (folder / "unet" / "config.json").is_file() + ) + except (OSError, ValueError, TypeError, json.JSONDecodeError): + return False + @staticmethod def _missing_ddpm_message(fields: dict[str, Any]) -> str: missing = [ @@ -1539,17 +1778,31 @@ class Planner: def _configured_tool_folder(self, tool_id: str) -> str: folders = self.config.get("tool_folders", {}) if not isinstance(folders, dict): - return "" + folders = {} raw_path = str(folders.get(tool_id, "")).strip() - return raw_path if raw_path and Path(raw_path).is_dir() else "" + if raw_path and Path(raw_path).is_dir(): + return raw_path + if tool_id == "oasis_trainer": + external = self.root / "config" / "external_tools.json" + try: + payload = json.loads(external.read_text(encoding="utf-8")) + except (OSError, ValueError, TypeError, json.JSONDecodeError): + return "" + for entry in payload.get("tools", []): + if isinstance(entry, dict) and entry.get("id") == "external_oasis_game_trainer": + candidate = Path(str(entry.get("backend", {}).get("root", ""))).expanduser() + if candidate.is_dir(): + return str(candidate) + return "" def _lora_plan(self, request: str, subject: str) -> ExecutionPlan: requested_name = self._model_name_from_request(request) model_name = requested_name or subject + training_options = self._training_options_from_request(request) project = _project_name(model_name, "LoRA") collector_root = self._configured_tool_folder("dataset_collector") trainer_root = self._configured_tool_folder("lora_trainer") - base_model = self._lora_base_model() + base_model = str(training_options.get("base_model") or self._lora_base_model()) if not collector_root or not trainer_root: return ExecutionPlan( request=request, @@ -1564,8 +1817,8 @@ class Planner: return ExecutionPlan( request=request, summary=( - "Select a valid SDXL base model in the connected LoRA app first. " - "ADAM will reuse that reviewed setting." + "Choose a valid SDXL base model in the generated LoRA settings " + "before training." ), steps=[], project_name="LoRA training", @@ -1610,6 +1863,7 @@ class Planner: "epochs": epochs, "output_dir": str(output_dir), "base_model": base_model, + **training_options, }, ), ] diff --git a/adam/process_control.py b/adam/process_control.py index 0514cf41c01e0858e8b4e87a40af15a106ddca79..771410f0a2b0d27eda3b4b566859e2b0ad33e835 100644 --- a/adam/process_control.py +++ b/adam/process_control.py @@ -1,6 +1,7 @@ from __future__ import annotations import subprocess +from typing import Any 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: return True except Exception: return False + + +def terminate_process_tree(process: Any, *, timeout: float = 3.0) -> None: + """Terminate a process and any children it launched.""" + try: + import psutil + + parent = psutil.Process(process.pid) + children = parent.children(recursive=True) + targets = [*children, parent] + for target in targets: + try: + target.terminate() + except (psutil.NoSuchProcess, psutil.AccessDenied): + continue + _gone, alive = psutil.wait_procs(targets, timeout=timeout) + for target in alive: + try: + target.kill() + except (psutil.NoSuchProcess, psutil.AccessDenied): + continue + return + except Exception: + pass + try: + process.terminate() + except Exception: + return + try: + process.wait(timeout=timeout) + except Exception: + try: + process.kill() + except Exception: + pass diff --git a/adam/recommendations.py b/adam/recommendations.py new file mode 100644 index 0000000000000000000000000000000000000000..dca9bf954b5892a34495bb30443d607b8a067914 --- /dev/null +++ b/adam/recommendations.py @@ -0,0 +1,195 @@ +from __future__ import annotations + +import math +import os +from dataclasses import asdict, dataclass, field +from typing import Any + +from adam.model_profiles import ModelProfile +from adam.models import SystemSnapshot + + +@dataclass(slots=True) +class SettingsRecommendation: + profile_id: str + epochs: int + settings: dict[str, Any] = field(default_factory=dict) + reasons: list[str] = field(default_factory=list) + warnings: list[str] = field(default_factory=list) + summary: str = "" + estimated_vram_gb: float | None = None + risk_level: str = "normal" + confidence: str = "conservative" + + def to_dict(self) -> dict[str, Any]: + return asdict(self) + + +def _field_default(profile: ModelProfile, key: str, fallback: Any) -> Any: + return profile.training.get(key, {}).get("default", fallback) + + +def _clamp_to_schema(profile: ModelProfile, key: str, value: Any) -> Any: + spec = profile.training.get(key, {}) + kind = str(spec.get("type", "text")) + try: + if kind in {"int", "slider"}: + numeric = int(value) + return max(int(spec.get("min", numeric)), min(numeric, int(spec.get("max", numeric)))) + if kind == "float": + numeric = float(value) + return max(float(spec.get("min", numeric)), min(numeric, float(spec.get("max", numeric)))) + except (TypeError, ValueError): + return spec.get("default", value) + if kind == "choice": + options = list(spec.get("options", [])) + return value if value in options else (options[0] if options else value) + return value + + +def estimate_vram_gb(profile: ModelProfile, resolution: int | str, batch_size: int, base_model_gb: float = 0.0) -> float: + """Broad VRAM estimate used only for warnings and conservative defaults.""" + architecture = profile.architecture.casefold() + pixels = (max(64, resolution) / 512) ** 2 + if profile.id == "lora" or "lora" in architecture: + base = max(6.0, base_model_gb * 1.8) + return base + pixels * max(1, batch_size) * 1.2 + if profile.id == "oasis" or "action_conditioned" in architecture: + width, height = (resolution, resolution) + if isinstance(resolution, str) and "x" in resolution: + try: + width, height = (int(part) for part in resolution.lower().split("x", 1)) + except ValueError: + width, height = (256, 144) + pixels = (max(width, height) / 512) ** 2 + return 4.0 + pixels * max(1, batch_size) * 2.4 + if "flow" in architecture: + return 2.8 + pixels * max(1, batch_size) * 1.0 + if "diffusion" in architecture: + return 2.2 + pixels * max(1, batch_size) * 0.9 + return 3.0 + pixels * max(1, batch_size) * 0.8 + + +def recommend_for_profile( + profile: ModelProfile, + *, + dataset_items: int, + resolution: int | str | None = None, + snapshot: SystemSnapshot | None = None, + base_model_gb: float = 0.0, +) -> SettingsRecommendation: + images = max(10, int(dataset_items or 10)) + raw_resolution = resolution or _field_default(profile, "resolution", 256) or 256 + if isinstance(raw_resolution, str) and "x" in raw_resolution: + resolution = int(raw_resolution.lower().split("x", 1)[0]) + else: + resolution = int(raw_resolution) + reasons: list[str] = [] + vram_total = snapshot.vram_total_gb if snapshot and snapshot.vram_total_gb else None + available_vram = ( + max(0.0, snapshot.vram_total_gb - snapshot.vram_used_gb) + if snapshot and snapshot.vram_total_gb + else vram_total + ) + architecture = profile.architecture.casefold() + target_exposures = 80_000 if profile.id == "lora" else 180_000 if "diffusion" in architecture else 120_000 + max_epochs = 220 if profile.id == "lora" else 600 if "diffusion" in architecture else 300 + epochs = max(10 if profile.id == "lora" else 25, min(max_epochs, round(target_exposures / images))) + reasons.append( + f"Epochs target roughly {target_exposures:,} image exposures, then clamp to the profile's safe range." + ) + + batch_defaults = { + 64: 16, + 128: 12, + 256: 4, + 384: 2, + 512: 1, + 768: 1, + 1024: 1, + } + if profile.id == "flow": + batch_defaults.update({64: 12, 128: 8, 256: 4}) + if profile.id == "oasis": + batch_defaults.update({128: 4, 256: 2, 384: 1, 512: 1}) + if profile.id == "lora": + batch_defaults.update({512: 2, 768: 1, 1024: 1}) + nearest = min(batch_defaults, key=lambda size: abs(size - resolution)) + batch_size = batch_defaults[nearest] + reasons.append(f"Batch starts from the closest resolution preset ({nearest}px).") + if available_vram is not None and available_vram < 8: + batch_size = max(1, batch_size // 2) + reasons.append("Available VRAM is below 8 GB, so batch size is reduced conservatively.") + + settings: dict[str, Any] = {} + for key in ("resolution", "batch_size"): + if key in profile.training: + value = batch_size + if key == "resolution": + value = raw_resolution if profile.id == "oasis" else resolution + settings[key] = _clamp_to_schema( + profile, + key, + value, + ) + if "learning_rate" in profile.training: + settings["learning_rate"] = _clamp_to_schema( + profile, + "learning_rate", + 0.00002 if profile.id == "oasis" else 0.0001 if profile.id in {"ddpm", "lora"} else 0.0002, + ) + workers = max(1, min(8, (os.cpu_count() or 4) // 2)) + for key in ("dataloader_num_workers", "workers"): + if key in profile.training: + settings[key] = _clamp_to_schema(profile, key, workers) + for key, value in { + "gradient_accumulation_steps": 1, + "gradient_accumulation": 1, + "mixed_precision": "fp32" if profile.id == "oasis" else "fp16", + "save_every": max(5, min(25, max(1, epochs // 10))), + "preview_every": max(5, min(50, max(1, epochs // 10))), + "training_intensity": 100, + "gradient_checkpointing": resolution >= 384 or (available_vram is not None and available_vram < 8), + "rank": 16, + "alpha": 16, + "frame_gap": 3, + "sequence_context": 1, + "preview_steps": 1 if profile.id == "oasis" else 50 if profile.id == "ddpm" else 10, + }.items(): + if key in profile.training: + settings[key] = _clamp_to_schema(profile, key, value) + + estimated = estimate_vram_gb(profile, resolution, int(settings.get("batch_size", batch_size)), base_model_gb) + warnings: list[str] = [] + if available_vram is not None and estimated > available_vram * 0.9: + warnings.append( + f"Estimated VRAM need is about {estimated:.1f} GB, above the conservative {available_vram * 0.9:.1f} GB working limit." + ) + if "batch_size" in settings and int(settings["batch_size"]) > 1: + settings["batch_size"] = max(1, int(settings["batch_size"]) // 2) + estimated = estimate_vram_gb(profile, resolution, int(settings["batch_size"]), base_model_gb) + warnings.append(f"Batch size was reduced to {settings['batch_size']} for a safer first run.") + reasons.append("The initial VRAM estimate was high, so ADAM reduced the batch before applying the recipe.") + if images < 20: + warnings.append("Dataset is very small; expect overfitting unless this is just a smoke test.") + reasons.append("Very small datasets get a warning because quality usually depends more on data cleanup than long training.") + risk_level = "risky" if warnings else "normal" + + memory_note = ( + f" using about {available_vram:.1f} GB available VRAM" if available_vram is not None else " without detected VRAM" + ) + summary = ( + f"Recommended {epochs:,} epochs for {images:,} item(s), " + f"batch {settings.get('batch_size', batch_size)} at {resolution}px{memory_note}. " + "Treat this as a starting recipe, not a guarantee." + ) + return SettingsRecommendation( + profile_id=profile.id, + epochs=epochs, + settings=settings, + reasons=reasons, + warnings=warnings, + summary=summary, + estimated_vram_gb=estimated, + risk_level=risk_level, + ) diff --git a/adam/registry.py b/adam/registry.py index ca58a9505107c5c2bd8d6bf7ecedaad3c18816b7..6815a3f8cab338e5a0f249209da78d7890bbc574 100644 --- a/adam/registry.py +++ b/adam/registry.py @@ -1,10 +1,13 @@ from __future__ import annotations import json +import logging from dataclasses import dataclass, field from pathlib import Path from typing import Any +from adam.model_plugins import ModelPluginRegistry + class RegistryError(RuntimeError): pass @@ -60,6 +63,7 @@ class ToolRegistry: self.root = root.resolve() self.path = self.root / "config" / "tools.json" self._tools: dict[str, ToolSpec] = {} + self.model_plugins = ModelPluginRegistry(self.root, logging.getLogger(__name__)) self.load() def load(self) -> None: @@ -112,6 +116,26 @@ class ToolRegistry: safe_entry["demo"] = False spec = ToolSpec.from_dict(safe_entry) loaded[spec.id] = spec + for entry in self.model_plugins.training_tool_specs() + self.model_plugins.generation_tool_specs(): + if not isinstance(entry, dict): + continue + safe_entry = dict(entry) + backend = dict(safe_entry.get("backend", {})) + if not backend: + safe_entry["backend"] = { + "type": "python", + "module": "adam.model_plugin_backend", + "function": "train" if safe_entry.get("category") == "Training" else "generate", + } + try: + spec = ToolSpec.from_dict(safe_entry) + except RegistryError as exc: + self.model_plugins.errors.append(f"{safe_entry.get('id', 'unknown')}: {exc}") + continue + if spec.id in loaded: + loaded[spec.id] = _merge_tool_specs(loaded[spec.id], spec) + else: + loaded[spec.id] = spec self._tools = loaded def get(self, tool_id: str, *, require_enabled: bool = True) -> ToolSpec: @@ -144,3 +168,31 @@ class ToolRegistry: } for tool in self.enabled() ] + + +def _merge_tool_specs(existing: ToolSpec, plugin: ToolSpec) -> ToolSpec: + """Keep the existing backend while accepting plugin-declared schema arguments.""" + arguments = tuple(dict.fromkeys([*existing.arguments, *plugin.arguments])) + required_arguments = existing.required_arguments or plugin.required_arguments + capabilities = tuple(dict.fromkeys([*existing.capabilities, *plugin.capabilities])) + model_trainers = tuple( + dict.fromkeys([*existing.model_trainers, *plugin.model_trainers]) + ) + generation_options = dict(existing.generation_options) + generation_options.update(plugin.generation_options) + return ToolSpec( + id=existing.id, + name=existing.name, + description=existing.description, + category=existing.category, + entry_function=existing.entry_function, + arguments=arguments, + required_arguments=required_arguments, + capabilities=capabilities, + model_trainers=model_trainers, + generation_options=generation_options, + requires_confirmation=existing.requires_confirmation, + enabled=existing.enabled, + demo=existing.demo, + backend=existing.backend, + ) diff --git a/adam/remote_access.py b/adam/remote_access.py new file mode 100644 index 0000000000000000000000000000000000000000..e69230a29c5452147a6340307cdc9fc4b280d5b0 --- /dev/null +++ b/adam/remote_access.py @@ -0,0 +1,2751 @@ +from __future__ import annotations + +import json +import mimetypes +import secrets +import shutil +import socket +import subprocess +import threading +from dataclasses import dataclass +from datetime import datetime, timedelta, timezone +from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer +from hmac import compare_digest +from ipaddress import ip_address +from pathlib import Path +from typing import Any +from urllib.parse import parse_qs, urlencode, urlparse + +from adam.generations import ( + ChatGenerationRequest, + build_generation_plan, + generation_model_match_score, + generation_tools, + load_generation_history, + parse_chat_generation_request, +) +from adam.remote_dispatcher import RemoteCommandDispatcher +from adam.remote_media import OpaqueIdCodec, RemoteMediaStore +from adam.remote_v1 import RemoteV1Service + + +REMOTE_MODE_DISABLED = "disabled" +REMOTE_MODE_LOCAL = "local_wifi" +REMOTE_MODE_TAILSCALE = "tailscale" + + +class RemoteHTTPServer(ThreadingHTTPServer): + """Bound concurrent connections so slow clients cannot spawn unlimited threads.""" + + allow_reuse_address = True + daemon_threads = True + + def __init__(self, *args, **kwargs): + self._slots = threading.BoundedSemaphore(16) + super().__init__(*args, **kwargs) + + def process_request(self, request, client_address): + if not self._slots.acquire(blocking=False): + self.shutdown_request(request) + return + try: + super().process_request(request, client_address) + except BaseException: + self._slots.release() + raise + + def process_request_thread(self, request, client_address): + try: + super().process_request_thread(request, client_address) + finally: + self._slots.release() + + +@dataclass(frozen=True, slots=True) +class TailscaleStatus: + installed: bool = False + connected: bool = False + device_name: str = "" + dns_name: str = "" + tailscale_ip: str = "" + backend_state: str = "" + serve_available: bool = False + serve_running: bool = False + message: str = "Tailscale is not installed." + + +def default_remote_settings() -> dict[str, Any]: + return { + "enabled": False, + "remote_mode": REMOTE_MODE_LOCAL, + "bind_address": "127.0.0.1", + "port": 8765, + "token": secrets.token_urlsafe(24), + "allow_job_control": False, + "auto_approve_training": False, + } + + +def remote_scope(bind_address: str) -> str: + bind = bind_address.strip().casefold() + if bind == "localhost": + return "local-device only" + if bind in {"0.0.0.0", "::"}: + return "all network interfaces" + try: + address = ip_address(bind.strip("[]")) + except ValueError: + return "custom bind address" + if address.is_loopback: + return "local-device only" + if address.is_private or address.is_link_local: + return "local network" + return "custom bind address" + + +def local_network_host() -> str: + """Best-effort address other devices on the same network can use.""" + try: + with socket.socket(socket.AF_INET, socket.SOCK_DGRAM) as sock: + sock.connect(("8.8.8.8", 80)) + host = str(sock.getsockname()[0]) + except OSError: + try: + host = socket.gethostbyname(socket.gethostname()) + except OSError: + return "" + try: + address = ip_address(host) + except ValueError: + return "" + if address.is_loopback or address.is_unspecified: + return "" + return host if address.is_private or address.is_link_local else "" + + +def inspect_tailscale( + runner: Any | None = None, + which: Any | None = None, +) -> TailscaleStatus: + which = which or shutil.which + executable = which("tailscale") + if not executable: + return TailscaleStatus() + runner = runner or _run_tailscale + try: + status = runner([executable, "status", "--json"]) + except OSError as exc: + return TailscaleStatus(installed=True, message=f"Tailscale could not be checked: {exc}") + if getattr(status, "returncode", 1) != 0: + error = _command_text(getattr(status, "stderr", "")) or "Tailscale is installed but not connected." + return TailscaleStatus(installed=True, message=error) + try: + payload = json.loads(_command_text(getattr(status, "stdout", "")) or "{}") + except json.JSONDecodeError: + return TailscaleStatus(installed=True, message="Tailscale returned an unreadable status response.") + self_node = payload.get("Self") if isinstance(payload, dict) else {} + self_node = self_node if isinstance(self_node, dict) else {} + ips = [str(item) for item in self_node.get("TailscaleIPs", []) if str(item)] + backend = str(payload.get("BackendState", "") or "") + connected = backend.casefold() == "running" or bool(ips) + serve_status = _tailscale_serve_running(runner, executable) + return TailscaleStatus( + installed=True, + connected=connected, + device_name=str(self_node.get("HostName", "") or ""), + dns_name=str(self_node.get("DNSName", "") or "").rstrip("."), + tailscale_ip=next((ip for ip in ips if "." in ip), ips[0] if ips else ""), + backend_state=backend, + serve_available=serve_status is not None, + serve_running=bool(serve_status), + message="Tailscale is connected." if connected else "Tailscale is installed but disconnected.", + ) + + +def _tailscale_serve_running(runner: Any, executable: str) -> bool | None: + try: + result = runner([executable, "serve", "status", "--json"]) + except OSError: + return None + if getattr(result, "returncode", 1) != 0: + return None + text = _command_text(getattr(result, "stdout", "")).strip() + return bool(text and text not in {"{}", "null"}) + + +def _run_tailscale(command: list[str]) -> subprocess.CompletedProcess[str]: + return subprocess.run(command, capture_output=True, text=True, timeout=8, check=False) + + +def _command_text(value: Any) -> str: + if isinstance(value, bytes): + return value.decode("utf-8", errors="replace") + return str(value or "") + + +def _remote_prompt_from_payload(payload: dict[str, Any]) -> str: + for key in ("prompt", "message", "text", "request", "input"): + value = payload.get(key) + if isinstance(value, str) and value.strip(): + return value.strip() + return "" + + +def _remote_dashboard_html() -> str: + return _remote_dashboard_app_html() + return """ + + + + +ADAM Remote + + + +
+
+
+ ADAM + AI Development and Automation Manager +
+
Connecting
+
+ +
+
+

Prompt ADAM

+
+ +
+ + + + +
+ +

Plans that need desktop approval will wait safely inside the main ADAM app.

+
+
+
+

Generate

+
+ +
+ + + + + + + +
+
+ + + +
+ +

Generation uses ADAM's existing desktop backend and queue.

+
+
+
+

Active Job

+
+
+
Checking ADAM...
+
Waiting for status.
+
+
...
+
+
+
+
--
+
--
+
--
+
--
+
+
+
+
+

Live Preview

+
+ Latest ADAM preview +
The latest training or generation preview will appear here.
+
+

Waiting for preview output.

+
+
+

Latest Generation

+
+ Latest generated ADAM image +
Finished generated images will appear here.
+
+ +

Waiting for a completed generation.

+
+
+
+

System

+
+
--
+
--
+
--
+
--
+
+

+
+
+

Remote Control

+

--

+

Status-only remote access is loading.

+ + +

Remote settings are synced from ADAM.

+ +
+
+
+

Recent Queue

+
+
+
+

Completed Jobs

+
+
+
+

Failed Jobs

+
+
+
+
+ + +""" + + +def _remote_dashboard_app_html() -> str: + from adam.remote_dashboard import remote_dashboard_app_html + + return remote_dashboard_app_html() + return """ + + + + +ADAM Remote + + + +
+
ADAM Remote
Connecting...
+
+
+
Active Job
No active job.
+
Live Preview
Live Preview
Waiting for a preview.
+
Latest Generation
+
Prompt ADAM
+
+
+
+
Structured Training
+ + + + + +
+
+
+ +
+

+    
+
+
+
Generate Image
+ + + + + +
+
+
+
+
+
+
+
Datasets
+ + +
Models
+
+
+
System
+
Remote Control
+
Jobs
+
+
+ + + +""" + + +class RemoteAccessService: + """Small authenticated local API foundation for browser/device clients.""" + + def __init__(self, config: Any, jobs: Any, monitor: Any, planner: Any = None) -> None: + self.config = config + self.jobs = jobs + self.monitor = monitor + self.planner = planner + self.root = Path(getattr(config, "root", None) or getattr(planner, "root", None) or Path.cwd()).resolve() + self.dispatcher = RemoteCommandDispatcher() + token = "" + try: + token = str(self.settings().get("token", "")) + except Exception: + token = "" + self.codec = OpaqueIdCodec(f"{self.root}|{token}") + self.media = RemoteMediaStore(self.root, self.codec) + self.api_v1 = RemoteV1Service( + root=self.root, + config=config, + jobs=jobs, + planner=planner, + dispatcher=self.dispatcher, + codec=self.codec, + media=self.media, + auto_approve_training=self._should_auto_approve_training, + ) + self._server: ThreadingHTTPServer | None = None + self._thread: threading.Thread | None = None + + @property + def running(self) -> bool: + return self._server is not None + + def settings(self) -> dict[str, Any]: + values = default_remote_settings() + stored = self.config.get("remote_access", {}) + if isinstance(stored, dict): + values.update(stored) + if not isinstance(stored, dict) or not stored.get("token"): + values["token"] = secrets.token_urlsafe(24) + self.config.update({"remote_access": values}) + return values + + def save_settings(self, values: dict[str, Any]) -> None: + clean = self.settings() + port = clean["port"] + try: + port = int(values.get("port", clean["port"])) + except (TypeError, ValueError): + port = clean["port"] + clean.update( + { + "enabled": bool(values.get("enabled", clean["enabled"])), + "remote_mode": self._clean_mode(str(values.get("remote_mode", clean["remote_mode"]))), + "bind_address": str(values.get("bind_address", clean["bind_address"])).strip() or "127.0.0.1", + "port": max(1024, min(port, 65535)), + "token": str(values.get("token", clean["token"])).strip() or secrets.token_urlsafe(24), + "allow_job_control": bool(values.get("allow_job_control", clean["allow_job_control"])), + "auto_approve_training": bool(values.get("auto_approve_training", clean["auto_approve_training"])), + } + ) + self.config.update({"remote_access": clean}) + + def start(self) -> str: + if self.running: + return self.url() + settings = self.settings() + mode = self._clean_mode(str(settings.get("remote_mode", REMOTE_MODE_LOCAL))) + if mode == REMOTE_MODE_DISABLED: + raise RuntimeError("Remote access is disabled by Remote Mode.") + if not settings.get("enabled"): + raise RuntimeError("Remote access is disabled.") + bind = "127.0.0.1" if mode == REMOTE_MODE_TAILSCALE else str(settings["bind_address"]) + port = int(settings["port"]) + jobs = self.jobs + monitor = self.monitor + service = self + api_v1 = self.api_v1 + + class Handler(BaseHTTPRequestHandler): + def setup(self) -> None: + self.request.settimeout(10) + super().setup() + + def _authorized(self) -> bool: + current = service.settings() + if not current.get("enabled") or current.get("remote_mode") == REMOTE_MODE_DISABLED: + return False + token = str(current["token"]).encode("utf-8") + header = self.headers.get("Authorization", "") + query_token = "" + parsed = urlparse(self.path) + values = parse_qs(parsed.query) + if values.get("token"): + query_token = values["token"][0] + bearer = header.removeprefix("Bearer ").strip() + return compare_digest(bearer.encode("utf-8"), token) or ( + bool(query_token) and self._browser_token_allowed() and compare_digest(query_token.encode("utf-8"), token) + ) + + def _security_headers(self) -> None: + self.send_header("X-Content-Type-Options", "nosniff") + self.send_header("Referrer-Policy", "no-referrer") + self.send_header("X-Frame-Options", "DENY") + self.send_header("Content-Security-Policy", "default-src 'self'; script-src 'self' 'unsafe-inline'; style-src 'self' 'unsafe-inline'; img-src 'self' data: blob:; connect-src 'self'; object-src 'none'; base-uri 'none'; frame-ancestors 'none'; form-action 'self'") + + def _valid_post(self) -> bool: + origin = self.headers.get("Origin") + if self.headers.get("Sec-Fetch-Site") == "cross-site" or (origin and ( + urlparse(origin).scheme not in {"http", "https"} + or urlparse(origin).netloc.casefold() != self.headers.get("Host", "").casefold() + )): + self._send(403, {"error": "Cross-site requests are not allowed."}) + return False + if self.headers.get("Content-Type", "").split(";", 1)[0].strip().lower() != "application/json": + self._send(415, {"error": "Use application/json for remote commands."}) + return False + lengths = self.headers.get_all("Content-Length", []) + try: + length = int(lengths[0]) if len(lengths) == 1 else -1 + except ValueError: + length = -1 + if self.headers.get("Transfer-Encoding") or length < 0 or length > 20_000: + self._send(413 if length > 20_000 else 400, {"error": "Invalid request size (maximum 20000 bytes)."}) + return False + return True + + def _browser_token_allowed(self) -> bool: + try: + client = ip_address(str(self.client_address[0]).strip("[]")) + except ValueError: + return False + if client.is_loopback: + return True + return remote_scope(bind) != "local-device only" and (client.is_private or client.is_link_local) + + def _send(self, status: int, payload: dict[str, Any]) -> None: + body = json.dumps(payload).encode("utf-8") + self._send_bytes(status, body, "application/json") + + def _send_html(self, status: int, html: str) -> None: + self._send_bytes(status, html.encode("utf-8"), "text/html; charset=utf-8") + + def _send_bytes(self, status: int, body: bytes, content_type: str) -> None: + if status >= 400 and self.command == "POST" and not getattr(self, "_body_read", False): + # Drain a small, already-sent body before closing; Windows can + # otherwise reset the connection before the error is delivered. + try: + length = int(self.headers.get("Content-Length", "0")) + if 0 < length <= 20_000 and not self.headers.get("Transfer-Encoding"): + self.connection.settimeout(0.25) + self.rfile.read(length) + except (ValueError, OSError): + pass + finally: + self.connection.settimeout(10) + self._body_read = True + try: + self.send_response(status) + self.send_header("Content-Type", content_type) + self.send_header("Cache-Control", "no-store") + self._security_headers() + self.send_header("Content-Length", str(len(body))) + self.end_headers() + self.wfile.write(body) + except (BrokenPipeError, ConnectionAbortedError, ConnectionResetError): + return + + def do_GET(self) -> None: + if not self._authorized(): + self._send(401, {"error": "Missing or invalid remote access token."}) + return + path = urlparse(self.path).path + response = api_v1.route("GET", path, urlparse(self.path).query) + if response is not None: + self._send_response(response) + return + if path == "/": + self._send_html(200, _remote_dashboard_html()) + return + if path == "/api/preview": + self._send_preview() + return + if path == "/api/generation-image": + self._send_generation_image() + return + if path != "/api/status": + self._send(404, {"error": "Unknown endpoint."}) + return + self._send(200, service._status_payload(bind, service.settings())) + + def do_POST(self) -> None: + if not self._authorized(): + self._send(401, {"error": "Missing or invalid remote access token."}) + return + if not self._valid_post(): + return + path = urlparse(self.path).path + payload = None + if path.startswith("/api/v1/"): + payload = self._read_json_body() + if payload is None: + self._send(400, {"error": "Send a valid JSON object."}) + return + response = api_v1.route("POST", path, urlparse(self.path).query, payload) + if response is not None: + self._send_response(response) + return + if path == "/api/job": + self._handle_job_action() + return + if path == "/api/remote-settings": + self._handle_remote_settings() + return + if path != "/api/prompt": + self._send(404, {"error": "Unknown endpoint."}) + return + payload = self._read_json_body() + if payload is None: + self._send(400, {"error": "Send a valid prompt."}) + return + prompt = _remote_prompt_from_payload(payload) + if not prompt: + self._send(400, {"error": "Type a prompt for ADAM first."}) + return + if len(prompt) > 2_000: + self._send(400, {"error": "Keep remote prompts under 2000 characters."}) + return + result = service.submit_prompt(prompt) + self._send(200 if result.get("ok") else 400, result) + + def _send_response(self, response: Any) -> None: + try: + self.send_response(int(response.status)) + self.send_header("Content-Type", str(response.content_type)) + headers = response.headers or {} + if "Cache-Control" in headers: + self.send_header("Cache-Control", headers["Cache-Control"]) + else: + self.send_header("Cache-Control", "no-store") + self._security_headers() + self.send_header("Content-Length", str(len(response.body))) + self.end_headers() + self.wfile.write(response.body) + except (BrokenPipeError, ConnectionAbortedError, ConnectionResetError): + return + + def _read_json_body(self) -> dict[str, Any] | None: + self._body_read = True + try: + length = int(self.headers.get("Content-Length", "0") or "0") + except ValueError: + length = 0 + try: + payload = json.loads(self.rfile.read(length).decode("utf-8")) if length else {} + except (UnicodeDecodeError, json.JSONDecodeError, OSError): + return None + return payload if isinstance(payload, dict) else None + + def _handle_job_action(self) -> None: + payload = self._read_json_body() + if payload is None: + self._send(400, {"error": "Send a valid job action."}) + return + action = str(payload.get("action", "")).strip().casefold() + job_id = str(payload.get("job_id", "")).strip() + result = service.job_action(job_id, action, bool(service.settings().get("allow_job_control"))) + self._send(200 if result.get("ok") else 400, result) + + def _handle_remote_settings(self) -> None: + payload = self._read_json_body() + if payload is None: + self._send(400, {"error": "Send valid remote settings."}) + return + current = service.settings() + enabled = payload.get("auto_approve_training", False) + if not isinstance(enabled, bool): + self._send(400, {"error": "Auto-approval must be true or false."}) + return + if enabled and not current.get("allow_job_control"): + self._send(403, {"error": "Enable remote job controls in the desktop app before changing approval permissions."}) + return + def update_approval(): + # Recheck on the owning thread and only change this permission. + # A queued request must not restore an older token or settings. + if enabled and not service.settings().get("allow_job_control"): + return False + service.save_settings({"auto_approve_training": enabled}) + return True + + if not service.dispatcher.call_ui(update_approval): + self._send(403, {"error": "Remote job controls have been disabled."}) + return + state = "on" if enabled else "off" + self._send(200, {"ok": True, "message": f"Auto-approval is {state}."}) + + def _send_preview(self) -> None: + path = service.preview_path() + if path is None or not path.is_file(): + self._send(404, {"error": "No preview image is available yet."}) + return + content_type = mimetypes.guess_type(str(path))[0] or "image/png" + try: + body = path.read_bytes() + except OSError: + self._send(404, {"error": "Preview image is no longer available."}) + return + self._send_bytes(200, body, content_type) + + def _send_generation_image(self) -> None: + parsed = urlparse(self.path) + values = parse_qs(parsed.query) + try: + record_index = int(values.get("record", ["0"])[0]) + image_index = int(values.get("image", ["0"])[0]) + except (TypeError, ValueError): + self._send(400, {"error": "Choose a valid generation image."}) + return + path = service.generation_image_path(record_index, image_index) + if path is None or not path.is_file(): + self._send(404, {"error": "Generated image is no longer available."}) + return + content_type = mimetypes.guess_type(str(path))[0] or "image/png" + try: + body = path.read_bytes() + except OSError: + self._send(404, {"error": "Generated image is no longer available."}) + return + self._send_bytes(200, body, content_type) + + def log_message(self, _format: str, *_args: Any) -> None: + return + + self._server = RemoteHTTPServer((bind, port), Handler) + self._thread = threading.Thread(target=self._server.serve_forever, daemon=True) + self._thread.start() + return self.url() + + def _status_payload(self, bind: str, settings: dict[str, Any]) -> dict[str, Any]: + snapshot = self.monitor.snapshot() if self.monitor is not None else None + active = self.jobs.active_job if self.jobs is not None else None + return { + "app": "ADAM", + "scope": remote_scope(bind), + "permissions": { + "status": True, + "system": True, + "queue_view": True, + "prompt": self.planner is not None and self.jobs is not None, + "job_control": bool(settings.get("allow_job_control")), + "auto_approve_training": bool(settings.get("auto_approve_training")), + "dangerous_actions": False, + }, + "active_job": None + if active is None + else { + "id": active.id, + "project": active.plan.project_name, + "status": active.status.value, + "progress": active.progress, + "timing": self._job_timing(active), + "preview": self._job_preview(active), + }, + "preview": self.preview_payload(), + "latest_generation": self.latest_generation_payload(), + "queue": [ + self._job_summary(job) + for job in (self.jobs.jobs[:20] if self.jobs is not None else []) + if job.status.value not in {"Finished", "Failed", "Cancelled"} + ], + "completed_jobs": [ + self._job_summary(job) + for job in (self.jobs.jobs[:30] if self.jobs is not None else []) + if job.status.value == "Finished" + ], + "failed_jobs": [ + self._job_summary(job) + for job in (self.jobs.jobs[:30] if self.jobs is not None else []) + if job.status.value in {"Failed", "Cancelled", "Interrupted"} + ], + "system": { + "cpu_percent": snapshot.cpu_percent if snapshot else None, + "memory_percent": snapshot.memory_percent if snapshot else None, + "gpu_name": snapshot.gpu_name if snapshot else "", + "gpu_percent": snapshot.gpu_percent if snapshot else None, + "vram_percent": snapshot.vram_percent if snapshot else None, + "gpu_temperature": snapshot.gpu_temperature if snapshot else None, + }, + } + + @staticmethod + def _job_summary(job: Any) -> dict[str, Any]: + try: + current_step = int(getattr(job, "current_step", -1)) + except (TypeError, ValueError): + current_step = -1 + steps = list(getattr(getattr(job, "plan", None), "steps", []) or []) + current_step_title = "" + if 0 <= current_step < len(steps): + current_step_title = str(getattr(steps[current_step], "title", "") or "") + return { + "id": getattr(job, "id", ""), + "project": getattr(getattr(job, "plan", None), "project_name", ""), + "status": getattr(getattr(job, "status", None), "value", str(getattr(job, "status", ""))), + "progress": getattr(job, "progress", 0), + "timing": RemoteAccessService._job_timing(job), + "error": getattr(job, "error", "") or "", + "started_at": getattr(job, "started_at", "") or "", + "ended_at": getattr(job, "ended_at", "") or "", + "scheduled_for": getattr(job, "scheduled_for", "") or "", + "current_step": current_step, + "step_count": len(steps), + "current_step_title": current_step_title, + "requires_confirmation": bool(getattr(getattr(job, "plan", None), "requires_confirmation", False)), + "output_folder": Path(str(getattr(job, "output_folder", "") or "")).name, + "progress_current": getattr(job, "progress_current", 0) or 0, + "progress_total": getattr(job, "progress_total", 0) or 0, + "progress_unit": getattr(job, "progress_unit", "") or "", + "logs": list(getattr(job, "logs", []) or [])[-3:], + } + + @staticmethod + def _job_timing(job: Any) -> dict[str, Any]: + started = RemoteAccessService._parse_time(getattr(job, "started_at", "") or "") + ended = RemoteAccessService._parse_time(getattr(job, "ended_at", "") or "") + progress = max(0, min(100, int(getattr(job, "progress", 0) or 0))) + status = getattr(getattr(job, "status", None), "value", str(getattr(job, "status", ""))) + review = getattr(getattr(job, "plan", None), "orion_review", {}) or {} + try: + estimate_seconds = int(float(review.get("estimated_high_minutes", 0) or 0) * 60) + except (TypeError, ValueError): + estimate_seconds = 0 + + now = datetime.now(timezone.utc) + elapsed_seconds = 0 + if started is not None: + finish = ended or now + elapsed_seconds = max(0, int((finish - started).total_seconds())) + + remaining_seconds: int | None = None + basis = "" + if status in {"Finished", "Failed", "Cancelled", "Interrupted"}: + remaining_seconds = 0 + basis = "complete" + elif started is not None and progress > 0: + remaining_seconds = max(0, int(elapsed_seconds * (100 - progress) / progress)) + basis = "progress" + elif estimate_seconds: + remaining_seconds = max(0, estimate_seconds - elapsed_seconds) + basis = "planning estimate" + + finish_label = "" + if remaining_seconds is not None and remaining_seconds > 0: + finish_label = (now + timedelta(seconds=remaining_seconds)).astimezone().strftime("%I:%M %p").lstrip("0") + + return { + "elapsed_seconds": elapsed_seconds if started is not None else None, + "remaining_seconds": remaining_seconds, + "estimated_total_seconds": estimate_seconds or None, + "elapsed_label": RemoteAccessService._format_duration(elapsed_seconds) if started is not None else "", + "remaining_label": ( + "done" if remaining_seconds == 0 and basis == "complete" + else f"about {RemoteAccessService._format_duration(remaining_seconds)}" if remaining_seconds is not None else "" + ), + "finish_label": finish_label, + "estimate_label": f"up to {RemoteAccessService._format_duration(estimate_seconds)}" if estimate_seconds else "", + "basis": basis, + } + + @staticmethod + def _parse_time(value: str) -> datetime | None: + if not value: + return None + try: + parsed = datetime.fromisoformat(str(value).replace("Z", "+00:00")) + except ValueError: + return None + if parsed.tzinfo is None: + parsed = parsed.replace(tzinfo=timezone.utc) + return parsed.astimezone(timezone.utc) + + @staticmethod + def _format_duration(seconds: int | float | None) -> str: + if seconds is None: + return "" + total = max(0, int(seconds)) + if total < 60: + return f"{total}s" + minutes, sec = divmod(total, 60) + if minutes < 60: + return f"{minutes}m {sec}s" if sec else f"{minutes}m" + hours, minute = divmod(minutes, 60) + if hours < 24: + return f"{hours}h {minute}m" if minute else f"{hours}h" + days, hour = divmod(hours, 24) + return f"{days}d {hour}h" if hour else f"{days}d" + + def preview_payload(self) -> dict[str, Any]: + if self.jobs is None: + return {"available": False} + candidates = [] + if self.jobs.active_job is not None: + candidates.append(self.jobs.active_job) + candidates.extend(self.jobs.jobs[:20]) + for job in candidates: + preview = self._job_preview(job) + if preview["available"]: + return preview + return {"available": False} + + def preview_path(self) -> Path | None: + if self.jobs is None: + return None + candidates = [] + if self.jobs.active_job is not None: + candidates.append(self.jobs.active_job) + candidates.extend(self.jobs.jobs[:20]) + for job in candidates: + path = self._job_preview_path(job) + if path is not None: + return path + return None + + @staticmethod + def _job_preview(job: Any) -> dict[str, Any]: + if RemoteAccessService._job_preview_path(job) is None: + return {"available": False} + return { + "available": True, + "url": "/api/preview", + "kind": getattr(job, "preview_kind", ""), + "epoch": getattr(job, "preview_epoch", 0), + "current": getattr(job, "preview_current", 0), + "total": getattr(job, "preview_total", 0), + "prompt": getattr(job, "preview_prompt", ""), + } + + @staticmethod + def _job_preview_path(job: Any) -> Path | None: + path = str(getattr(job, "preview_path", "") or "") + if not path: + return None + target = Path(path) + return target if target.is_file() else None + + def latest_generation_payload(self) -> dict[str, Any]: + records = self._generation_records() + if not records: + return {"available": False} + record = records[0] + return { + "available": True, + "url": "/api/generation-image?record=0&image=0", + "images": [ + { + "index": index, + "url": f"/api/generation-image?record=0&image={index}", + } + for index, _path in enumerate(record.images) + ], + "model_name": record.model_name, + "provider_id": record.provider_id, + "provider_name": record.provider_name, + "prompt": record.prompt, + "seed": record.seed, + "steps": record.steps, + "sampler": record.sampler, + "aspect_ratio": record.aspect_ratio, + "created_at": record.created_at, + "image_count": len(record.images), + } + + def generation_image_path(self, record_index: int, image_index: int) -> Path | None: + if record_index < 0 or image_index < 0: + return None + records = self._generation_records(limit=max(1, record_index + 1)) + if record_index >= len(records): + return None + record = records[record_index] + if image_index >= len(record.images): + return None + path = record.images[image_index] + return path if path.is_file() else None + + def _generation_records(self, *, limit: int = 30): + root = self._generation_root() + if root is None: + return [] + return load_generation_history(root, limit=limit) + + def _generation_root(self) -> Path | None: + for candidate in ( + getattr(self.planner, "root", None), + getattr(getattr(self.planner, "registry", None), "root", None), + ): + if candidate: + return Path(candidate) + return None + + def submit_prompt(self, prompt: str) -> dict[str, Any]: + generation_result = self._submit_generation_prompt(prompt) + if generation_result is not None: + return generation_result + if self.planner is None or self.jobs is None: + return {"ok": False, "error": "Remote prompting is not available in this ADAM session."} + try: + def prepare_plan(): + from adam.training_assistant import append_preflight_summary + + plan = self.planner.plan(prompt) + append_preflight_summary(plan, self.config) + return plan + + plan = self.dispatcher.call_background(prepare_plan) + except Exception as exc: + return {"ok": False, "error": f"ADAM could not plan that request: {exc}"} + if not plan.steps: + return {"ok": True, "message": plan.summary or "ADAM received your message.", "requires_approval": False} + job = self.dispatcher.submit_job(self.jobs, plan) + if self._should_auto_approve_training(plan): + self.dispatcher.confirm_job(self.jobs, job.id) + return { + "ok": True, + "message": f"Queued {job.plan.project_name}. Remote training auto-approval is on.", + "job_id": job.id, + "requires_approval": False, + "auto_approved": True, + } + if plan.requires_confirmation: + return { + "ok": True, + "message": f"Plan created for {job.plan.project_name}. It needs approval in the desktop app before it runs.", + "job_id": job.id, + "requires_approval": True, + } + return { + "ok": True, + "message": f"Queued {job.plan.project_name}.", + "job_id": job.id, + "requires_approval": False, + } + + def job_action(self, job_id: str, action: str, allowed: bool) -> dict[str, Any]: + if not allowed: + return {"ok": False, "error": "Remote job controls are disabled in ADAM."} + if self.jobs is None: + return {"ok": False, "error": "Job controls are not available in this ADAM session."} + if not job_id: + return {"ok": False, "error": "Choose a job first."} + if action not in {"cancel", "retry", "confirm", "pause", "resume", "end"}: + return {"ok": False, "error": "Unsupported remote job action."} + try: + result = self.dispatcher.job_action(self.jobs, job_id, action) + job = result.get("job") + if action == "cancel": + return {"ok": True, "message": f"Cancellation requested for {job.plan.project_name}."} + if action == "confirm": + return {"ok": True, "message": f"Approved {job.plan.project_name}."} + if action == "pause": + return {"ok": True, "message": f"Paused {job.plan.project_name}."} + if action == "resume": + return {"ok": True, "message": f"Resumed {job.plan.project_name}."} + if action == "end": + return {"ok": True, "message": f"Ended {job.plan.project_name}."} + retried = result.get("retried") + return {"ok": True, "message": f"Retry queued for {retried.plan.project_name}.", "job_id": retried.id} + except Exception as exc: + return {"ok": False, "error": f"ADAM could not update that job: {exc}"} + + def _submit_generation_prompt(self, prompt: str) -> dict[str, Any] | None: + parsed = parse_chat_generation_request(prompt) + if parsed is None: + return None + if self.planner is None or self.jobs is None: + return {"ok": False, "error": "Remote image generation is not available in this ADAM session."} + if not all(hasattr(self.planner, name) for name in ("assets", "registry")): + return {"ok": False, "error": "Remote image generation needs the full ADAM planner session."} + try: + plan = self._generation_plan(parsed) + except ValueError as exc: + return {"ok": False, "error": str(exc)} + try: + job = self.dispatcher.submit_job(self.jobs, plan) + except Exception as exc: + return {"ok": False, "error": f"ADAM could not queue that generation: {exc}"} + return { + "ok": True, + "message": f"Queued {job.plan.project_name}.", + "job_id": job.id, + "requires_approval": bool(job.plan.requires_confirmation), + } + + def _should_auto_approve_training(self, plan: Any) -> bool: + if getattr(plan, "orion_review", {}).get("level") == "warning": + return False + if not getattr(plan, "requires_confirmation", False): + return False + if not bool(self.settings().get("auto_approve_training")): + return False + return any( + str(getattr(step, "tool_id", "")).endswith("_trainer") + for step in getattr(plan, "steps", []) + ) + + def _generation_plan(self, parsed: ChatGenerationRequest): + assets = self.planner.assets + registry = self.planner.registry + if hasattr(assets, "discover"): + assets.discover(self.config) + tools = generation_tools(registry) + if not tools: + raise ValueError("No image generators are currently available in ADAM.") + + stable_diffusion_request = ( + parsed.has_positive_prompt + or bool(parsed.base_model_query) + or bool(parsed.negative_prompt) + or parsed.cfg_scale is not None + or parsed.lora_strength is not None + or parsed.denoise_strength is not None + ) + plain_model_search = ( + not parsed.provider_hint + and not stable_diffusion_request + and not parsed.model_query + ) + base_only = ( + stable_diffusion_request + and parsed.provider_hint != "lora" + and not parsed.model_query + ) + preferred_id = { + "ddpm": "ddpm_generator", + "flow": "flow_generator", + "lora": "lora_generator", + }.get(parsed.provider_hint, "") + if stable_diffusion_request and parsed.provider_hint not in {"ddpm", "flow"}: + preferred_id = "lora_generator" + preferred_tool = next((item for item in tools if item.id == preferred_id), None) + if parsed.provider_hint and preferred_tool is None: + raise ValueError(f"The requested {parsed.provider_hint.upper()} image generator is not currently available.") + + candidate_tools = [preferred_tool] if preferred_tool else tools + candidates = [ + asset + for asset in getattr(assets, "assets", []) + if asset.kind == "model" + and (not plain_model_search or asset.trainer in {"ddpm", "flow"}) + and not (plain_model_search and parsed.reference_image and asset.trainer == "flow") + and any( + item is not None and asset.trainer in item.model_trainers + for item in candidate_tools + ) + and self._generation_model_is_ready(asset) + ] + model_query = parsed.model_query or (parsed.subject if not base_only else "") + scored = sorted( + ( + (generation_model_match_score(model_query, asset.name), asset) + for asset in candidates + ), + key=lambda item: item[0], + reverse=True, + ) + model = scored[0][1] if scored and scored[0][0] > 0 else None + if model is None and not model_query and len(candidates) == 1: + model = candidates[0] + if model is None and plain_model_search: + base_only = True + preferred_tool = next((item for item in tools if item.id == "lora_generator"), None) + candidate_tools = [preferred_tool] if preferred_tool else tools + model_query = "" + if model is None and not base_only: + detail = f' matching "{model_query}"' if model_query else "" + examples: list[str] = [] + for asset in candidates: + if asset.name not in examples: + examples.append(asset.name) + if len(examples) == 4: + break + available = f" Available examples: {', '.join(examples)}." if examples else "" + raise ValueError( + f"I could not find a completed image model{detail}.{available} " + 'Try: Generate an image using model "Model Name".' + ) + + tool = next( + ( + item + for item in candidate_tools + if item is not None and (base_only or model.trainer in item.model_trainers) + ), + None, + ) + if tool is None: + raise ValueError("The matching model does not have an available image generator.") + if parsed.reference_image and "reference_image" not in tool.capabilities: + raise ValueError( + f"{tool.name} does not support reference-image conditioning. " + "Remove the attachment or choose LoRA/Stable Diffusion or DDPM." + ) + + options = tool.generation_options + saved_generation = self.config.get("generation_settings", {}) + saved_generation = saved_generation if isinstance(saved_generation, dict) else {} + sampler_options = [str(value) for value in options.get("samplers", [])] + sampler = parsed.sampler or ( + str(saved_generation.get("sampler", "")) if tool.id == "lora_generator" else "" + ) + if sampler not in sampler_options: + sampler = sampler_options[0] if sampler_options else sampler or "DDIM" + aspect_options = [str(value) for value in options.get("aspect_ratios", [])] + aspect = parsed.aspect_ratio or ( + str(saved_generation.get("aspect", "")) if tool.id == "lora_generator" else "" + ) + if aspect and aspect not in aspect_options: + aspect = next( + (value for value in aspect_options if value.startswith(f"{aspect} ") or value == aspect), + "", + ) + if not aspect: + aspect = aspect_options[0] if aspect_options else "1:1 (Square)" + step_min = int(options.get("step_min", 1) or 1) + step_max = int(options.get("step_max", 500) or 500) + default_steps = ( + saved_generation.get("steps", options.get("step_default", 50)) + if tool.id == "lora_generator" + else options.get("step_default", 50) + ) + steps = parsed.steps if parsed.steps is not None else int(default_steps or 50) + steps = max(step_min, min(steps, step_max)) + count_limit = 8 if tool.id == "lora_generator" else 32 + default_count = ( + int(saved_generation.get("images", 1) or 1) + if tool.id == "lora_generator" + else 1 + ) + count = max(1, min(parsed.image_count or default_count, count_limit)) + seed = parsed.seed if parsed.seed is not None else 0 + extra_arguments: dict[str, Any] = {} + if tool.id == "ddpm_generator": + extra_arguments = { + "reference_image": parsed.reference_image, + "reference_strength": max( + 0, + min(parsed.reference_strength if parsed.reference_strength is not None else 65, 100), + ), + "width": 0, + "height": 0, + } + elif tool.id == "lora_generator": + base_model_path = self._stable_diffusion_base_model_path(parsed, saved_generation) + extra_arguments = { + "negative_prompt": parsed.negative_prompt or str(saved_generation.get("negative_prompt", "")), + "base_model_path": base_model_path, + "width": 0, + "height": 0, + "cfg_scale": parsed.cfg_scale if parsed.cfg_scale is not None else float(saved_generation.get("cfg_scale", 0) or 0), + "lora_strength": 0.0 if base_only else ( + parsed.lora_strength if parsed.lora_strength is not None else float(saved_generation.get("lora_strength", 0) or 0) + ), + "reference_image": parsed.reference_image, + "denoise_strength": parsed.denoise_strength if parsed.denoise_strength is not None else float(saved_generation.get("denoise_strength", 0) or 0), + "prompt_weighting": bool(saved_generation.get("prompt_weighting", True)), + } + + return build_generation_plan( + tool, + model_name=(Path(extra_arguments.get("base_model_path", "")).stem if base_only else model.name), + model_path="" if base_only else model.path, + prompt=parsed.prompt, + image_count=count, + steps=steps, + seed=seed, + sampler=sampler, + aspect_ratio=aspect, + extra_arguments=extra_arguments, + ) + + def _stable_diffusion_base_model_path( + self, + parsed: ChatGenerationRequest, + saved_generation: dict[str, Any], + ) -> str: + base_assets = [ + asset + for asset in getattr(self.planner.assets, "assets", []) + if asset.kind == "base_model" and Path(asset.path).exists() + ] + if parsed.base_model_query: + scored_bases = sorted( + ( + (generation_model_match_score(parsed.base_model_query, asset.name), asset) + for asset in base_assets + ), + key=lambda item: item[0], + reverse=True, + ) + if scored_bases and scored_bases[0][0] > 0: + return scored_bases[0][1].path + raise ValueError( + f'I could not find a Stable Diffusion base model matching "{parsed.base_model_query}".' + ) + preferred_base = next( + ( + asset + for asset in base_assets + if "waiillustrious" in "".join( + character for character in asset.name.casefold() if character.isalnum() + ) + or "wallilustrious" in "".join( + character for character in asset.name.casefold() if character.isalnum() + ) + ), + None, + ) + if preferred_base is not None: + return preferred_base.path + selected_base = str(saved_generation.get("base_model_path", "")) + if selected_base and Path(selected_base).expanduser().exists(): + return selected_base + trainer_root = Path(str(self.config.get("tool_folders", {}).get("lora_trainer", ""))) + try: + trainer_settings = json.loads( + (trainer_root / "config" / "app_settings.json").read_text(encoding="utf-8") + ) + configured_base = str( + trainer_settings.get("generate_model") + or trainer_settings.get("last_model") + or "" + ) + configured_path = Path(configured_base).expanduser() + if configured_base and not configured_path.is_absolute(): + configured_path = trainer_root / configured_path + if configured_base and configured_path.exists(): + return str(configured_path.resolve()) + except (OSError, ValueError, TypeError, json.JSONDecodeError): + pass + if len(base_assets) == 1: + return base_assets[0].path + names = ", ".join(asset.name for asset in base_assets[:4]) + available = f" Available base models: {names}." if names else "" + raise ValueError( + "LoRA generation also needs a Stable Diffusion base model. Put one in " + '"LoRA StableDiffusionModels Here", or select one in the Generations tab.' + + available + ) + + @staticmethod + def _generation_model_is_ready(asset: Any) -> bool: + path = Path(asset.path) + if asset.trainer == "ddpm": + return path.is_dir() and (path / "model_index.json").is_file() + if asset.trainer == "flow": + return ( + path.is_dir() + and (path / "flow_model_info.json").is_file() + and (path / "unet" / "config.json").is_file() + ) + if asset.trainer == "lora": + return ( + path.is_file() + and path.suffix.casefold() == ".safetensors" + and "_comfy" not in path.stem.casefold() + ) or ( + path.is_dir() + and any( + item.is_file() + and item.suffix.casefold() == ".safetensors" + and "_comfy" not in item.stem.casefold() + for item in path.glob("*.safetensors") + ) + ) + return path.exists() + + def stop(self) -> None: + if self._server is None: + return + self._server.shutdown() + self._server.server_close() + self._server = None + self._thread = None + + def shutdown(self) -> None: + self.stop() + self.dispatcher.shutdown() + + def url(self) -> str: + settings = self.settings() + host = "127.0.0.1" if self._clean_mode(str(settings.get("remote_mode"))) == REMOTE_MODE_TAILSCALE else str(settings["bind_address"]) + return f"http://{self._url_host(host)}:{int(settings['port'])}/api/status" + + def local_test_url(self) -> str: + settings = self.settings() + host = "127.0.0.1" if self._clean_mode(str(settings.get("remote_mode"))) == REMOTE_MODE_TAILSCALE else str(settings["bind_address"]) + if host in {"0.0.0.0", "::"}: + host = "127.0.0.1" + query = urlencode({"token": str(settings["token"])}) + return f"http://{self._url_host(host)}:{int(settings['port'])}/?{query}" + + def phone_test_url(self) -> str: + settings = self.settings() + if self._clean_mode(str(settings.get("remote_mode"))) == REMOTE_MODE_TAILSCALE: + return self.tailscale_url() + host = str(settings["bind_address"]) + scope = remote_scope(host) + if scope == "local-device only": + return "" + if host in {"0.0.0.0", "::"}: + host = local_network_host() + if not host: + return "" + query = urlencode({"token": str(settings["token"])}) + return f"http://{self._url_host(host)}:{int(settings['port'])}/?{query}" + + def tailscale_url(self) -> str: + status = inspect_tailscale() + if not status.installed or not status.connected: + return "" + host = status.dns_name or status.tailscale_ip + if not host: + return "" + settings = self.settings() + query = urlencode({"token": str(settings["token"])}) + return f"https://{self._url_host(host)}/?{query}" + + def tailscale_status(self) -> TailscaleStatus: + return inspect_tailscale() + + def start_tailscale_serve(self) -> tuple[bool, str]: + status = inspect_tailscale() + if not status.installed: + return False, "Tailscale is not installed." + if not status.connected: + return False, "Tailscale is installed but not connected." + executable = shutil.which("tailscale") + if not executable: + return False, "Tailscale is not installed." + port = int(self.settings()["port"]) + try: + result = _run_tailscale([executable, "serve", "--bg", str(port)]) + except (OSError, subprocess.TimeoutExpired) as exc: + return False, f"Tailscale Serve could not start: {exc}" + if result.returncode != 0: + return False, _command_text(result.stderr) or "Tailscale Serve could not start." + return True, "Tailscale Serve is forwarding private tailnet traffic to ADAM." + + def stop_tailscale_serve(self) -> tuple[bool, str]: + executable = shutil.which("tailscale") + if not executable: + return False, "Tailscale is not installed." + try: + result = _run_tailscale([executable, "serve", "reset"]) + except (OSError, subprocess.TimeoutExpired) as exc: + return False, f"Tailscale Serve could not stop: {exc}" + if result.returncode != 0: + return False, _command_text(result.stderr) or "Tailscale Serve could not stop." + return True, "Tailscale Serve forwarding was reset." + + @staticmethod + def _url_host(host: str) -> str: + value = host.strip() or "127.0.0.1" + if ":" in value and not value.startswith("["): + return f"[{value}]" + return value + + @staticmethod + def _clean_mode(value: str) -> str: + mode = value.strip().casefold() + return mode if mode in {REMOTE_MODE_DISABLED, REMOTE_MODE_LOCAL, REMOTE_MODE_TAILSCALE} else REMOTE_MODE_LOCAL diff --git a/adam/remote_api.py b/adam/remote_api.py new file mode 100644 index 0000000000000000000000000000000000000000..aeb493ab5e2f79fa4fed37cf3a0491bb3a4976ea --- /dev/null +++ b/adam/remote_api.py @@ -0,0 +1,147 @@ +from __future__ import annotations + +import json +import math +from dataclasses import dataclass +from typing import Any + + +class RemoteApiError(ValueError): + def __init__(self, message: str, *, status: int = 400) -> None: + super().__init__(message) + self.status = status + + +@dataclass(frozen=True, slots=True) +class RemoteResponse: + status: int + body: bytes + content_type: str = "application/json" + headers: dict[str, str] | None = None + + +def json_response(payload: dict[str, Any], *, status: int = 200) -> RemoteResponse: + return RemoteResponse( + status=status, + body=json.dumps(payload, separators=(",", ":")).encode("utf-8"), + content_type="application/json", + ) + + +def error_response(message: str, *, status: int = 400) -> RemoteResponse: + return json_response({"ok": False, "error": message}, status=status) + + +def media_response(body: bytes, content_type: str, *, cache_seconds: int = 86400) -> RemoteResponse: + return RemoteResponse( + status=200, + body=body, + content_type=content_type, + headers={"Cache-Control": f"private, max-age={max(0, int(cache_seconds))}"}, + ) + + +def bounded_text(value: Any, *, max_length: int, label: str, required: bool = False) -> str: + if value is None: + value = "" + if not isinstance(value, str): + value = str(value) + text = value.strip() + if required and not text: + raise RemoteApiError(f"{label} is required.") + if len(text) > max_length: + raise RemoteApiError(f"{label} must be {max_length} characters or shorter.") + return text + + +def bounded_int( + value: Any, + *, + minimum: int, + maximum: int, + default: int, + label: str, +) -> int: + if value in (None, ""): + return default + try: + if isinstance(value, bool): + raise ValueError + number = int(value) + except (TypeError, ValueError) as exc: + raise RemoteApiError(f"{label} must be a whole number.") from exc + if number < minimum or number > maximum: + raise RemoteApiError(f"{label} must be between {minimum} and {maximum}.") + return number + + +def bounded_float( + value: Any, + *, + minimum: float, + maximum: float, + default: float, + label: str, +) -> float: + if value in (None, ""): + return default + try: + if isinstance(value, bool): + raise ValueError + number = float(value) + except (TypeError, ValueError) as exc: + raise RemoteApiError(f"{label} must be a number.") from exc + if not math.isfinite(number) or number < minimum or number > maximum: + raise RemoteApiError(f"{label} must be between {minimum:g} and {maximum:g}.") + return number + + +def parse_pagination(query: dict[str, list[str]], *, default_size: int = 24, max_size: int = 80) -> dict[str, int]: + page = bounded_int( + (query.get("page") or ["1"])[0], + minimum=1, + maximum=1_000_000, + default=1, + label="Page", + ) + page_size = bounded_int( + (query.get("page_size") or [str(default_size)])[0], + minimum=1, + maximum=max_size, + default=default_size, + label="Page size", + ) + return { + "page": page, + "page_size": page_size, + "offset": (page - 1) * page_size, + "limit": page_size, + } + + +def coerce_json_object(payload: Any) -> dict[str, Any]: + if not isinstance(payload, dict): + raise RemoteApiError("Send a JSON object.") + return payload + + +def sanitized_arguments(arguments: dict[str, Any]) -> dict[str, Any]: + """Return client-safe arguments without absolute filesystem paths.""" + hidden = { + "dataset_dir", + "output_dir", + "model_path", + "base_model", + "base_model_path", + "resume_from", + "reference_image", + } + clean: dict[str, Any] = {} + for key, value in arguments.items(): + if key in hidden: + text = str(value or "") + clean[f"{key}_name"] = text.replace("\\", "/").rstrip("/").rsplit("/", 1)[-1] if text else "" + continue + if isinstance(value, (str, int, float, bool)) or value is None: + clean[key] = value + return clean diff --git a/adam/remote_dashboard.py b/adam/remote_dashboard.py new file mode 100644 index 0000000000000000000000000000000000000000..1ea04dbfa5efd64c1da1cc3a50fc88d093e7cf28 --- /dev/null +++ b/adam/remote_dashboard.py @@ -0,0 +1,141 @@ +from __future__ import annotations + + +def remote_dashboard_app_html() -> str: + return """ + + + + +ADAM Remote + + + +
+
ADAM RemoteConnecting
+ +
+
+
Active Job
Checking ADAM...
Time Left -Finish -
+
Live Preview
Waiting for a preview.
+
Prompt ADAM
+
Latest Generation
+
+
+ +
+

Datasets

+
Q
+
+
Favorites
View All
+
Recently Used
View All
+
Locations
+
All Remembered
+ + +
+ +
+

Create Model

+
+
+
Save time with a preset configuration.
+
+
+
No dataset selected
+
+
+
Advanced Settings
+ + + +

+      
+
+
+ +
Jobs
+
Remote Control
Remembered Locations
System
+
+ + + +""" diff --git a/adam/remote_dispatcher.py b/adam/remote_dispatcher.py new file mode 100644 index 0000000000000000000000000000000000000000..7f5ec7d1ac6ca91108f498fc7989d6cd3843783f --- /dev/null +++ b/adam/remote_dispatcher.py @@ -0,0 +1,120 @@ +from __future__ import annotations + +import threading +from concurrent.futures import ThreadPoolExecutor +from typing import Any, Callable, TypeVar + + +T = TypeVar("T") + + +class _NoQtInvoker: + def invoke(self, payload: dict[str, Any]) -> None: + fn = payload["fn"] + try: + payload["result"] = fn() + except Exception as exc: + payload["error"] = exc + finally: + payload["event"].set() + + +def _qt_invoker_class(): + try: + from PySide6.QtCore import QObject, Signal, Slot + except Exception: + return None + + class _QtInvoker(QObject): + invoke_requested = Signal(object) + + def __init__(self) -> None: + super().__init__() + self.invoke_requested.connect(self._invoke) + + @Slot(object) + def _invoke(self, payload: dict[str, Any]) -> None: + fn = payload["fn"] + try: + payload["result"] = fn() + except Exception as exc: + payload["error"] = exc + finally: + payload["event"].set() + + def invoke(self, payload: dict[str, Any]) -> None: + self.invoke_requested.emit(payload) + + return _QtInvoker + + +class RemoteCommandDispatcher: + """Boundary between HTTP worker threads and ADAM's Qt-owned services.""" + + def __init__(self, *, timeout_seconds: float = 30.0) -> None: + self.timeout_seconds = timeout_seconds + self.main_thread_id = threading.get_ident() + self._executor = ThreadPoolExecutor(max_workers=2, thread_name_prefix="adam-remote-plan") + self._invoker: Any = _NoQtInvoker() + invoker_class = _qt_invoker_class() + if invoker_class is not None: + try: + from PySide6.QtCore import QCoreApplication + + if QCoreApplication.instance() is not None: + self._invoker = invoker_class() + except Exception: + self._invoker = _NoQtInvoker() + + def shutdown(self) -> None: + self._executor.shutdown(wait=False, cancel_futures=True) + + def call_background(self, fn: Callable[[], T], *, timeout: float | None = None) -> T: + future = self._executor.submit(fn) + return future.result(timeout=timeout or self.timeout_seconds) + + def call_ui(self, fn: Callable[[], T], *, timeout: float | None = None) -> T: + if threading.get_ident() == self.main_thread_id: + return fn() + event = threading.Event() + payload: dict[str, Any] = {"fn": fn, "event": event} + self._invoker.invoke(payload) + if not event.wait(timeout or self.timeout_seconds): + raise TimeoutError("ADAM did not finish the remote command in time.") + if "error" in payload: + raise payload["error"] + return payload.get("result") + + def submit_job(self, jobs: Any, plan: Any) -> Any: + return self.call_ui(lambda: jobs.submit(plan)) + + def confirm_job(self, jobs: Any, job_id: str) -> Any: + return self.call_ui(lambda: jobs.confirm(job_id)) + + def job_action(self, jobs: Any, job_id: str, action: str) -> Any: + def run() -> Any: + job = jobs.get(job_id) if hasattr(jobs, "get") else None + if job is None: + raise ValueError("ADAM could not find that job.") + if action == "cancel": + jobs.cancel(job_id) + return {"job": job} + if action == "confirm": + jobs.confirm(job_id) + return {"job": job} + if action == "pause": + jobs.pause(job_id) + return {"job": job} + if action == "resume": + jobs.resume(job_id) + return {"job": job} + if action == "end": + ended = jobs.end_task(job_id) if hasattr(jobs, "end_task") else False + if not ended: + raise ValueError("Only interrupted jobs can be ended remotely.") + return {"job": job} + if action == "retry": + return {"job": job, "retried": jobs.retry(job_id)} + raise ValueError("Unsupported remote job action.") + + return self.call_ui(run) diff --git a/adam/remote_media.py b/adam/remote_media.py new file mode 100644 index 0000000000000000000000000000000000000000..67e4225c5baee5aab975a3c85bc8fe92bc517731 --- /dev/null +++ b/adam/remote_media.py @@ -0,0 +1,104 @@ +from __future__ import annotations + +import base64 +import hashlib +import hmac +import json +import mimetypes +from dataclasses import dataclass +from pathlib import Path +from typing import Any + +from adam.remote_api import RemoteApiError + + +MAX_SOURCE_BYTES = 60 * 1024 * 1024 +MAX_THUMBNAIL_SIZE = 1200 +DEFAULT_THUMBNAIL_SIZE = 320 + + +class OpaqueIdCodec: + def __init__(self, secret: str) -> None: + self.secret = (secret or "adam-remote").encode("utf-8") + + def encode(self, payload: dict[str, Any]) -> str: + body = json.dumps(payload, sort_keys=True, separators=(",", ":")).encode("utf-8") + token = base64.urlsafe_b64encode(body).decode("ascii").rstrip("=") + signature = hmac.new(self.secret, token.encode("ascii"), hashlib.sha256).hexdigest()[:24] + return f"{token}.{signature}" + + def decode(self, value: str) -> dict[str, Any]: + try: + token, signature = value.rsplit(".", 1) + except ValueError as exc: + raise RemoteApiError("Unknown resource id.", status=404) from exc + expected = hmac.new(self.secret, token.encode("ascii"), hashlib.sha256).hexdigest()[:24] + if not hmac.compare_digest(signature, expected): + raise RemoteApiError("Unknown resource id.", status=404) + padding = "=" * (-len(token) % 4) + try: + payload = json.loads(base64.urlsafe_b64decode((token + padding).encode("ascii")).decode("utf-8")) + except (ValueError, TypeError, json.JSONDecodeError) as exc: + raise RemoteApiError("Unknown resource id.", status=404) from exc + if not isinstance(payload, dict): + raise RemoteApiError("Unknown resource id.", status=404) + return payload + + +@dataclass(frozen=True, slots=True) +class RemoteMediaFile: + path: Path + content_type: str + cache_hit: bool + + +class RemoteMediaStore: + def __init__(self, root: Path, codec: OpaqueIdCodec) -> None: + self.root = root.resolve() + self.codec = codec + self.cache_root = self.root / "data" / "remote_thumbnails" + self.cache_root.mkdir(parents=True, exist_ok=True) + + def media_id(self, *, kind: str, asset_id: str, index: int) -> str: + return self.codec.encode({"kind": kind, "asset_id": asset_id, "index": int(index)}) + + def thumbnail(self, source: Path, *, size: int = DEFAULT_THUMBNAIL_SIZE) -> RemoteMediaFile: + source = source.expanduser().resolve() + if not source.is_file(): + raise RemoteApiError("Media file was not found.", status=404) + try: + stat = source.stat() + except OSError as exc: + raise RemoteApiError("Media file could not be read.", status=404) from exc + if stat.st_size > MAX_SOURCE_BYTES: + raise RemoteApiError("Media file is too large for remote preview.", status=413) + bounded_size = max(64, min(int(size or DEFAULT_THUMBNAIL_SIZE), MAX_THUMBNAIL_SIZE)) + key = hashlib.sha256( + f"{source}|{stat.st_mtime_ns}|{stat.st_size}|{bounded_size}".encode("utf-8", errors="ignore") + ).hexdigest() + target = self.cache_root / f"{key}.jpg" + if target.is_file(): + return RemoteMediaFile(target, "image/jpeg", True) + try: + from PIL import Image, ImageOps + + with Image.open(source) as image: + image = ImageOps.exif_transpose(image) + image.thumbnail((bounded_size, bounded_size), Image.Resampling.LANCZOS) + if image.mode not in {"RGB", "L"}: + image = image.convert("RGB") + image.save(target, format="JPEG", quality=82, optimize=True) + except Exception as exc: + raise RemoteApiError("Thumbnail could not be generated.", status=415) from exc + return RemoteMediaFile(target, "image/jpeg", False) + + def original(self, source: Path) -> RemoteMediaFile: + source = source.expanduser().resolve() + if not source.is_file(): + raise RemoteApiError("Media file was not found.", status=404) + try: + if source.stat().st_size > MAX_SOURCE_BYTES: + raise RemoteApiError("Media file is too large for remote viewing.", status=413) + except OSError as exc: + raise RemoteApiError("Media file could not be read.", status=404) from exc + return RemoteMediaFile(source, mimetypes.guess_type(str(source))[0] or "application/octet-stream", True) diff --git a/adam/remote_v1.py b/adam/remote_v1.py new file mode 100644 index 0000000000000000000000000000000000000000..d552bba563b2adc852dc02451420ce7d8695cb74 --- /dev/null +++ b/adam/remote_v1.py @@ -0,0 +1,758 @@ +from __future__ import annotations + +from pathlib import Path +import os +import tempfile +from typing import Any +from urllib.parse import parse_qs + +from adam.assets import Asset, AssetRegistry +from adam.commands import CommandValidationError, TrainingCommand +from adam.dataset_registry import DatasetRecord, DatasetRegistry +from adam.dataset_lab import scan_dataset +from adam.generations import build_generation_plan, generation_tools +from adam.remote_api import ( + RemoteApiError, + RemoteResponse, + bounded_float, + bounded_int, + bounded_text, + coerce_json_object, + error_response, + json_response, + media_response, + parse_pagination, + sanitized_arguments, +) +from adam.remote_dispatcher import RemoteCommandDispatcher +from adam.remote_media import OpaqueIdCodec, RemoteMediaStore +from adam.studio import caption_path, image_files, StudioStore +from adam.training_assistant import append_preflight_summary + + +class RemoteV1Service: + """Versioned ADAM Remote API built around existing ADAM services.""" + + def __init__( + self, + *, + root: Path, + config: Any, + jobs: Any, + planner: Any, + dispatcher: RemoteCommandDispatcher, + codec: OpaqueIdCodec, + media: RemoteMediaStore, + auto_approve_training: Any, + ) -> None: + self.root = root.resolve() + self.config = config + self.jobs = jobs + self.planner = planner + self.dispatcher = dispatcher + self.codec = codec + self.media = media + self.auto_approve_training = auto_approve_training + self._asset_fallback: AssetRegistry | None = None + self._studio: StudioStore | None = None + + def route( + self, + method: str, + path: str, + query: str = "", + payload: dict[str, Any] | None = None, + ) -> RemoteResponse | None: + if not path.startswith("/api/v1/"): + return None + parts = [part for part in path.removeprefix("/api/v1/").split("/") if part] + parsed_query = parse_qs(query) + try: + if method == "GET": + return self._get(parts, parsed_query) + if method == "POST": + return self._post(parts, coerce_json_object(payload or {})) + except RemoteApiError as exc: + return error_response(str(exc), status=exc.status) + except (CommandValidationError, ValueError) as exc: + return error_response(str(exc), status=400) + except Exception as exc: + return error_response(f"ADAM could not finish that remote request: {exc}", status=500) + return error_response("Unsupported Remote API method.", status=405) + + def _get(self, parts: list[str], query: dict[str, list[str]]) -> RemoteResponse: + if parts == ["datasets"]: + return json_response({"ok": True, "datasets": self.datasets()}) + if parts == ["datasets", "locations"]: + return json_response({"ok": True, "locations": self.dataset_locations()}) + if parts == ["training", "presets"]: + return json_response({"ok": True, "presets": self.training_presets()}) + if len(parts) == 2 and parts[0] == "datasets": + return json_response({"ok": True, "dataset": self.dataset_detail(parts[1])}) + if len(parts) == 3 and parts[0] == "datasets" and parts[2] == "thumbnail": + return self.dataset_thumbnail(parts[1], query) + if len(parts) == 3 and parts[0] == "datasets" and parts[2] == "items": + return json_response({"ok": True, **self.dataset_items(parts[1], query)}) + if len(parts) == 2 and parts[0] == "media": + return self.media_file(parts[1], query) + if parts == ["models"]: + return json_response({"ok": True, "models": self.models()}) + if parts == ["training", "schema"]: + return json_response({"ok": True, **self.training_schema()}) + if parts == ["generation", "schema"]: + return json_response({"ok": True, **self.generation_schema()}) + if len(parts) == 2 and parts[0] == "jobs": + return json_response({"ok": True, "job": self.job_detail(parts[1])}) + raise RemoteApiError("Unknown Remote API endpoint.", status=404) + + def _post(self, parts: list[str], payload: dict[str, Any]) -> RemoteResponse: + if len(parts) == 5 and parts[0] == "datasets" and parts[2] == "items" and parts[4] == "caption": + return json_response({"ok": True, **self.update_caption(parts[1], parts[3], payload)}) + if len(parts) == 5 and parts[0] == "datasets" and parts[2] == "items" and parts[4] == "decision": + return json_response({"ok": True, **self.update_decision(parts[1], parts[3], payload)}) + if len(parts) == 3 and parts[0] == "datasets" and parts[2] == "favorite": + return json_response({"ok": True, **self.favorite_dataset(parts[1], payload)}) + if len(parts) == 3 and parts[0] == "datasets" and parts[2] == "use": + return json_response({"ok": True, "dataset": self.use_dataset(parts[1])}) + if parts == ["training", "plan"]: + return json_response({"ok": True, "plan": self.training_plan(payload)}) + if parts == ["training", "start"]: + return json_response({"ok": True, **self.start_training(payload)}) + if parts == ["generation", "start"]: + return json_response({"ok": True, **self.start_generation(payload)}) + raise RemoteApiError("Unknown Remote API endpoint.", status=404) + + def _assets(self) -> AssetRegistry: + assets = getattr(self.planner, "assets", None) + if assets is None: + if self._asset_fallback is None: + self._asset_fallback = AssetRegistry(self.root) + assets = self._asset_fallback + if hasattr(assets, "discover"): + assets.discover(self.config) + return assets + + def _registry(self) -> Any: + registry = getattr(self.planner, "registry", None) + if registry is None: + raise RemoteApiError("The ADAM tool registry is not available.", status=503) + return registry + + def _studio_store(self) -> StudioStore: + if self._studio is None: + self._studio = StudioStore(self.root) + return self._studio + + def _dataset_registry(self) -> DatasetRegistry: + return DatasetRegistry(self.root, self.config) + + def _asset_public_id(self, asset: Asset) -> str: + return self.codec.encode({"kind": asset.kind, "asset_id": asset.id}) + + def _asset_from_id(self, public_id: str, *, kind: str = "") -> Asset: + payload = self.codec.decode(public_id) + asset_id = str(payload.get("asset_id", "")) + expected_kind = kind or str(payload.get("kind", "")) + for asset in self._assets().assets: + if asset.id == asset_id and (not expected_kind or asset.kind == expected_kind): + return asset + raise RemoteApiError("Unknown resource id.", status=404) + + def _item_path(self, item_id: str, *, dataset: Asset | None = None) -> tuple[Asset, Path, int]: + payload = self.codec.decode(item_id) + if payload.get("kind") != "dataset_image": + raise RemoteApiError("Unknown media id.", status=404) + asset_id = str(payload.get("asset_id", "")) + index = bounded_int(payload.get("index"), minimum=0, maximum=100_000, default=0, label="Image index") + if dataset is not None and dataset.id != asset_id: + raise RemoteApiError("Dataset image does not belong to that dataset.", status=403) + dataset_asset = dataset or next((item for item in self._assets().assets if item.id == asset_id and item.kind == "dataset"), None) + if dataset_asset is None: + raise RemoteApiError("Unknown dataset image.", status=404) + paths = image_files(dataset_asset.path, limit=100_000) + if index >= len(paths): + raise RemoteApiError("Dataset image is no longer available.", status=404) + path = paths[index].resolve() + try: + path.relative_to(Path(dataset_asset.path).expanduser().resolve()) + except ValueError as exc: + raise RemoteApiError("Dataset image is outside its dataset.", status=403) from exc + return dataset_asset, path, index + + @staticmethod + def _dataset_path(asset: Asset, path: Path) -> Path: + resolved = path.expanduser().resolve() + if not resolved.is_relative_to(Path(asset.path).expanduser().resolve()): + raise RemoteApiError("File is outside its dataset.", status=403) + return resolved + + def _dataset_counts(self, asset: Asset) -> dict[str, int]: + review = self._studio_store().review(asset.path) + images = image_files(asset.path, limit=100_000) + resolved = {str(path.resolve()) for path in images} + keep = sum(1 for path, decision in review.decisions.items() if decision == "keep" and path in resolved) + reject = sum(1 for path, decision in review.decisions.items() if decision == "reject" and path in resolved) + return { + "accepted": keep, + "rejected": reject, + "unreviewed": max(0, len(images) - keep - reject), + } + + def _dimensions(self, path: Path) -> dict[str, int]: + try: + from PIL import Image + + with Image.open(path) as image: + return {"width": int(image.width), "height": int(image.height)} + except Exception: + return {"width": 0, "height": 0} + + def datasets(self) -> list[dict[str, Any]]: + assets = self._assets() + records = self._dataset_registry().discover(asset_registry=assets) + by_path = {str(Path(asset.path).expanduser().resolve()): asset for asset in assets.assets if asset.kind == "dataset"} + rows = [] + for record in records[:300]: + asset = by_path.get(str(Path(record.path).expanduser().resolve())) + if asset is None and record.exists: + asset = assets.register( + kind="dataset", + name=record.name, + path=record.path, + metadata={ + "dataset_registry_source": record.source, + "dataset_location_id": record.location_id, + }, + persist=False, + ) + if asset is None: + continue + rows.append(self._dataset_summary(asset, record)) + if rows: + assets.save() + return rows + + def _dataset_summary(self, asset: Asset, record: DatasetRecord | None = None) -> dict[str, Any]: + if record is None: + record = self._dataset_registry().record_for_path(asset.path) + exists = Path(asset.path).is_dir() + public_id = self._asset_public_id(asset) + if exists: + review_state = self._studio_store().review(asset.path) + accepted = sum(1 for decision in review_state.decisions.values() if decision == "keep") + rejected = sum(1 for decision in review_state.decisions.values() if decision == "reject") + review = { + "accepted": accepted, + "rejected": rejected, + "unreviewed": max(0, record.image_count - accepted - rejected), + } + else: + review = {"accepted": 0, "rejected": 0, "unreviewed": 0} + warnings = list(record.warnings) + if not exists and "Dataset folder is unavailable." not in warnings: + warnings.append("Dataset folder is unavailable.") + return { + "id": public_id, + "name": asset.name or record.name, + "exists": exists, + "available": exists, + "source": record.source, + "location_id": record.location_id, + "favorite": record.favorite, + "last_used_at": record.last_used_at, + "item_count": record.item_count, + "image_count": record.image_count, + "video_count": record.video_count, + "caption_count": record.caption_count, + "missing_caption_count": record.missing_caption_count, + "dataset_format": record.dataset_format, + "warnings": warnings[:6], + "thumbnail_url": f"/api/v1/datasets/{public_id}/thumbnail" if exists else "", + "review": review, + } + + def dataset_locations(self) -> list[dict[str, Any]]: + registry = self._dataset_registry() + return [ + { + "id": item.id, + "name": item.name, + "source": item.source, + "exists": Path(item.path).is_dir(), + "available": Path(item.path).is_dir(), + "last_seen_at": item.last_seen_at, + } + for item in registry.known_locations() + ][:200] + + def _first_image(self, folder: str | Path) -> Path | None: + root = Path(folder).expanduser() + if not root.is_dir(): + return None + try: + for path in root.rglob("*"): + if path.is_file() and path.suffix.casefold() in {".png", ".jpg", ".jpeg", ".webp", ".bmp", ".gif"}: + return path + except OSError: + return None + return None + + def dataset_detail(self, dataset_id: str) -> dict[str, Any]: + asset = self._asset_from_id(dataset_id, kind="dataset") + registry = self._dataset_registry() + record = registry.record_for_path(asset.path) + if not Path(asset.path).is_dir(): + return self._dataset_summary(asset, record) + report = scan_dataset(asset.path, limit=300) + registry.refresh_async(asset.path, source=record.source, location_id=record.location_id) + return { + "id": self._asset_public_id(asset), + "name": asset.name, + "exists": True, + "available": True, + "source": record.source, + "location_id": record.location_id, + "favorite": record.favorite, + "last_used_at": record.last_used_at, + "image_count": report.image_count, + "video_count": report.video_count, + "item_count": report.image_count + report.video_count, + "caption_count": report.caption_count, + "missing_caption_count": report.missing_caption_count, + "duplicate_groups": report.duplicate_groups, + "dimensions": dict(list(report.dimensions.items())[:20]), + "extensions": report.extensions, + "warnings": report.warnings, + "dataset_format": record.dataset_format, + "thumbnail_url": f"/api/v1/datasets/{dataset_id}/thumbnail", + "review": self._dataset_counts(asset), + } + + def dataset_thumbnail(self, dataset_id: str, query: dict[str, list[str]]) -> RemoteResponse: + asset = self._asset_from_id(dataset_id, kind="dataset") + source = self._first_image(asset.path) + if source is None: + raise RemoteApiError("Dataset thumbnail is not available.", status=404) + source = self._dataset_path(asset, source) + size = bounded_int((query.get("size") or ["320"])[0], minimum=64, maximum=640, default=320, label="Media size") + media = self.media.thumbnail(source, size=size) + return media_response(media.path.read_bytes(), media.content_type, cache_seconds=86400) + + def dataset_items(self, dataset_id: str, query: dict[str, list[str]]) -> dict[str, Any]: + asset = self._asset_from_id(dataset_id, kind="dataset") + page = parse_pagination(query, default_size=24, max_size=60) + paths = image_files(asset.path, limit=100_000) + total = len(paths) + review = self._studio_store().review(asset.path) + rows = [] + for index, path in enumerate(paths[page["offset"]: page["offset"] + page["limit"]], page["offset"]): + path = self._dataset_path(asset, path) + item_id = self.media.media_id(kind="dataset_image", asset_id=asset.id, index=index) + caption_file = self._dataset_path(asset, caption_path(path)) + try: + caption = "" + if caption_file.is_file(): + with caption_file.open(encoding="utf-8") as handle: + caption = handle.read(4000) + except (OSError, UnicodeError): + caption = "" + rows.append({ + "id": item_id, + "display_name": path.name, + "dimensions": self._dimensions(path), + "caption": caption, + "has_caption": bool(caption.strip()), + "decision": review.decisions.get(str(path.resolve()), "unreviewed"), + "thumbnail_url": f"/api/v1/media/{item_id}?size=320", + "preview_url": f"/api/v1/media/{item_id}?size=960", + }) + return { + "dataset": self.dataset_detail(dataset_id), + "items": rows, + "pagination": { + "page": page["page"], + "page_size": page["page_size"], + "total": total, + "has_next": page["offset"] + page["limit"] < total, + }, + } + + def media_file(self, media_id: str, query: dict[str, list[str]]) -> RemoteResponse: + _dataset, source, _index = self._item_path(media_id) + size = bounded_int((query.get("size") or ["320"])[0], minimum=64, maximum=1200, default=320, label="Media size") + media = self.media.thumbnail(source, size=size) + return media_response(media.path.read_bytes(), media.content_type, cache_seconds=86400) + + def update_caption(self, dataset_id: str, item_id: str, payload: dict[str, Any]) -> dict[str, Any]: + asset = self._asset_from_id(dataset_id, kind="dataset") + _asset, image, _index = self._item_path(item_id, dataset=asset) + caption = bounded_text(payload.get("caption", ""), max_length=4000, label="Caption") + target = self._dataset_path(asset, caption_path(image)) + # Replace the file instead of writing through a possible hard link. + temporary = None + try: + with tempfile.NamedTemporaryFile(mode="w", encoding="utf-8", dir=target.parent, delete=False) as handle: + temporary = Path(handle.name) + handle.write(caption.rstrip() + ("\n" if caption else "")) + os.replace(temporary, target) + finally: + if temporary is not None: + temporary.unlink(missing_ok=True) + return {"message": "Caption saved.", "item": {"id": item_id, "caption": caption, "has_caption": bool(caption)}} + + def update_decision(self, dataset_id: str, item_id: str, payload: dict[str, Any]) -> dict[str, Any]: + asset = self._asset_from_id(dataset_id, kind="dataset") + _asset, image, _index = self._item_path(item_id, dataset=asset) + decision = bounded_text(payload.get("decision", "unreviewed"), max_length=32, label="Decision") + if decision not in {"keep", "reject", "unreviewed"}: + raise RemoteApiError("Decision must be keep, reject, or unreviewed.") + self._studio_store().set_decision(asset.path, str(image), decision) + return {"message": "Review decision saved.", "item": {"id": item_id, "decision": decision}} + + def favorite_dataset(self, dataset_id: str, payload: dict[str, Any]) -> dict[str, Any]: + asset = self._asset_from_id(dataset_id, kind="dataset") + enabled = bool(payload.get("favorite", True)) + record = self._dataset_registry().favorite(asset.path, enabled) + return { + "message": "Dataset favorite updated.", + "dataset": self._dataset_summary(asset, record), + } + + def use_dataset(self, dataset_id: str) -> dict[str, Any]: + asset = self._asset_from_id(dataset_id, kind="dataset") + record = self._dataset_registry().touch(asset.path) + return self._dataset_summary(asset, record) + + def models(self) -> list[dict[str, Any]]: + assets = self._assets() + experiment_by_model: dict[str, Any] = {} + try: + for run in getattr(getattr(self.jobs, "experiments", None), "list_runs", lambda limit=100: [])(limit=100): + experiment_by_model.setdefault(run.model_name, run) + except Exception: + experiment_by_model = {} + rows = [] + for asset in assets.assets: + if asset.kind not in {"model", "base_model"} or not Path(asset.path).exists(): + continue + dataset = next((item for item in assets.assets if item.id == asset.dataset_id), None) + metadata = dict(asset.metadata or {}) + trigger_word = str(metadata.get("trigger_word") or (asset.name if asset.trainer == "lora" else "")) + latest = experiment_by_model.get(asset.name) + rows.append({ + "id": self._asset_public_id(asset), + "name": asset.name, + "kind": asset.kind, + "architecture": asset.trainer or ("stable_diffusion" if asset.kind == "base_model" else ""), + "trainer": asset.trainer, + "checkpoint_name": Path(asset.checkpoint or asset.path).name, + "dataset": None if dataset is None else {"id": self._asset_public_id(dataset), "name": dataset.name}, + "epochs": asset.epochs, + "trigger_word": trigger_word, + "latest_experiment": None if latest is None else { + "id": latest.id, + "status": latest.status, + "resolution": latest.resolution, + "dataset_name": latest.dataset_name, + "trigger_word": getattr(latest, "trigger_word", ""), + }, + }) + return rows[:300] + + def training_schema(self) -> dict[str, Any]: + registry = self._registry() + trainers = [] + for tool in registry.enabled(): + if not tool.id.endswith("_trainer"): + continue + trainer = tool.id.removesuffix("_trainer") + schema = registry.model_plugins.training_schema(trainer) + trainers.append({ + "id": trainer, + "tool_id": tool.id, + "name": tool.name, + "requires_confirmation": tool.requires_confirmation, + "settings": schema, + "arguments": [arg for arg in tool.arguments if arg not in {"dataset_dir", "output_dir"}], + "required_arguments": [arg for arg in tool.required_arguments if arg not in {"dataset_dir", "output_dir"}], + }) + return { + "trainers": trainers, + "datasets": self.datasets(), + "base_models": [item for item in self.models() if item["kind"] == "base_model"], + "presets": self.training_presets(), + } + + def training_presets(self) -> list[dict[str, Any]]: + presets = [ + { + "id": "standard_ddpm", + "name": "Standard DDPM", + "trainer": "ddpm", + "epochs": 100, + "settings": { + "resolution": 128, + "batch_size": 1, + "learning_rate": 0.0001, + "save_every": 10, + "preview_enabled": True, + "preview_every": 5, + }, + }, + { + "id": "quick_ddpm_test", + "name": "Quick DDPM Test", + "trainer": "ddpm", + "epochs": 3, + "settings": { + "resolution": 64, + "batch_size": 1, + "learning_rate": 0.0001, + "save_every": 1, + "preview_enabled": True, + "preview_every": 1, + "training_intensity": 25, + }, + }, + { + "id": "oasis_training", + "name": "Oasis Training", + "trainer": "oasis", + "epochs": 25, + "settings": { + "resolution": "256x144", + "batch_size": 2, + "learning_rate": 0.00002, + "workers": 2, + "preview_enabled": True, + "preview_every": 5, + }, + }, + ] + for recipe in self._studio_store().recipes: + presets.append({ + "id": f"recipe_{recipe.id}", + "name": recipe.name, + "trainer": recipe.trainer, + "epochs": recipe.epochs, + "settings": { + "preview_prompt": recipe.preview_prompt, + **({"base_model": recipe.base_model} if recipe.base_model else {}), + }, + "user": True, + }) + available = {trainer["id"] for trainer in self.training_schema_no_presets()} + return [preset for preset in presets if preset["trainer"] in available] + + def training_schema_no_presets(self) -> list[dict[str, Any]]: + registry = self._registry() + trainers = [] + for tool in registry.enabled(): + if not tool.id.endswith("_trainer"): + continue + trainer = tool.id.removesuffix("_trainer") + trainers.append({"id": trainer}) + return trainers + + def _build_training_plan(self, payload: dict[str, Any]) -> Any: + if self.planner is None: + raise RemoteApiError("Planning is not available in this ADAM session.", status=503) + trainer = bounded_text(payload.get("trainer", "lora"), max_length=64, label="Trainer", required=True).casefold() + dataset = self._asset_from_id(bounded_text(payload.get("dataset_id"), max_length=4000, label="Dataset", required=True), kind="dataset") + if not Path(dataset.path).is_dir(): + raise RemoteApiError("That dataset is unavailable on the PC.", status=404) + model_name = bounded_text(payload.get("model_name") or dataset.name, max_length=96, label="Model name", required=True) + epochs = bounded_int(payload.get("epochs"), minimum=1, maximum=100_000, default=10, label="Epochs") + options = dict(payload.get("settings") or payload.get("training_options") or {}) + if not isinstance(options, dict): + raise RemoteApiError("Training settings must be an object.") + trigger_word = bounded_text(payload.get("trigger_word") or options.get("trigger_word") or "", max_length=128, label="Trigger word") + if trigger_word: + options["trigger_word"] = trigger_word + base_model = "" + if trainer == "lora": + base_id = bounded_text(payload.get("base_model_id", ""), max_length=4000, label="Base model") + if base_id: + base_model = self._asset_from_id(base_id, kind="base_model").path + else: + base_model = str(options.get("base_model") or "") + if not base_model and hasattr(self.planner, "_lora_base_model"): + base_model = str(self.planner._lora_base_model()) + if base_model: + options["base_model"] = base_model + output = "" + if hasattr(self.planner, "_training_output"): + output_path = self.planner._training_output(trainer, model_name) + output = str(output_path) if output_path else "" + command = TrainingCommand.from_dict({ + "action": "train", + "trainer": trainer, + "dataset": dataset.path, + "model_name": model_name, + "epochs": epochs, + "output": output, + "base_model": base_model, + "trigger_word": trigger_word, + "training_options": options, + }) + if not hasattr(self.planner, "_plan_training_command"): + raise RemoteApiError("Structured training validation is not available.", status=503) + plan = self.planner._plan_training_command("Remote structured training", command) + append_preflight_summary(plan, self.config) + return plan + + def training_plan(self, payload: dict[str, Any]) -> dict[str, Any]: + plan = self._build_training_plan(payload) + return self._plan_payload(plan) + + def start_training(self, payload: dict[str, Any]) -> dict[str, Any]: + if self.jobs is None: + raise RemoteApiError("Jobs are not available in this ADAM session.", status=503) + plan = self._build_training_plan(payload) + if not plan.steps: + raise RemoteApiError(plan.summary or "ADAM could not build a runnable training plan.") + job = self.dispatcher.submit_job(self.jobs, plan) + auto_approved = False + if self.auto_approve_training(plan): + self.dispatcher.confirm_job(self.jobs, job.id) + auto_approved = True + try: + dataset = self._asset_from_id(bounded_text(payload.get("dataset_id"), max_length=4000, label="Dataset"), kind="dataset") + self._dataset_registry().touch(dataset.path) + except Exception: + pass + return { + "message": f"Queued {job.plan.project_name}.", + "job_id": job.id, + "requires_approval": bool(plan.requires_confirmation and not auto_approved), + "auto_approved": auto_approved, + "plan": self._plan_payload(plan), + } + + def generation_schema(self) -> dict[str, Any]: + registry = self._registry() + providers = [] + for tool in generation_tools(registry): + providers.append({ + "id": tool.id, + "name": tool.name, + "model_trainers": tool.model_trainers, + "requires_confirmation": tool.requires_confirmation, + "options": tool.generation_options, + "settings": registry.model_plugins.generation_schema_for_tool(tool.id), + }) + return {"providers": providers, "models": [item for item in self.models() if item["kind"] == "model"], "base_models": [item for item in self.models() if item["kind"] == "base_model"]} + + def start_generation(self, payload: dict[str, Any]) -> dict[str, Any]: + if self.jobs is None: + raise RemoteApiError("Jobs are not available in this ADAM session.", status=503) + registry = self._registry() + provider_id = bounded_text(payload.get("provider_id", ""), max_length=96, label="Provider", required=True) + tool = next((item for item in generation_tools(registry) if item.id == provider_id), None) + if tool is None: + raise RemoteApiError("Unknown generation provider.", status=404) + model = self._asset_from_id(bounded_text(payload.get("model_id", ""), max_length=4000, label="Model", required=True), kind="model") + if model.trainer not in tool.model_trainers: + raise RemoteApiError("That model is not compatible with the selected provider.") + prompt = bounded_text(payload.get("prompt", ""), max_length=2000, label="Prompt") + count = bounded_int(payload.get("image_count"), minimum=1, maximum=32, default=1, label="Image count") + if provider_id == "lora_generator": + count = min(count, 8) + options = tool.generation_options + step_max = int(options.get("step_max", 500) or 500) + steps = bounded_int(payload.get("steps"), minimum=1, maximum=step_max, default=int(options.get("step_default", 30) or 30), label="Steps") + seed = bounded_int(payload.get("seed"), minimum=0, maximum=2_147_483_647, default=0, label="Seed") + sampler_options = [str(item) for item in options.get("samplers", [])] + sampler = bounded_text(payload.get("sampler") or (sampler_options[0] if sampler_options else "DDIM"), max_length=64, label="Sampler") + if sampler_options and sampler not in sampler_options: + raise RemoteApiError("Sampler is not supported by that provider.") + aspect_options = [str(item) for item in options.get("aspect_ratios", [])] + aspect = bounded_text(payload.get("aspect_ratio") or (aspect_options[0] if aspect_options else "1:1 (Square)"), max_length=64, label="Aspect ratio") + if aspect_options and aspect not in aspect_options: + raise RemoteApiError("Aspect ratio is not supported by that provider.") + extra: dict[str, Any] = {} + if provider_id == "lora_generator": + base_id = bounded_text(payload.get("base_model_id", ""), max_length=4000, label="Base model") + base_model_path = self._asset_from_id(base_id, kind="base_model").path if base_id else "" + extra.update({ + "negative_prompt": bounded_text(payload.get("negative_prompt", ""), max_length=3000, label="Negative prompt"), + "base_model_path": base_model_path, + "width": bounded_int(payload.get("width"), minimum=0, maximum=2048, default=0, label="Width"), + "height": bounded_int(payload.get("height"), minimum=0, maximum=2048, default=0, label="Height"), + "cfg_scale": bounded_float(payload.get("cfg_scale"), minimum=0.0, maximum=30.0, default=0.0, label="CFG scale"), + "lora_strength": bounded_float(payload.get("lora_strength"), minimum=0.0, maximum=3.0, default=1.0, label="LoRA strength"), + "denoise_strength": bounded_float(payload.get("denoise_strength"), minimum=0.0, maximum=1.0, default=0.0, label="Denoise strength"), + "prompt_weighting": bool(payload.get("prompt_weighting", True)), + }) + elif provider_id == "ddpm_generator": + extra.update({ + "reference_strength": bounded_int(payload.get("reference_strength"), minimum=0, maximum=100, default=65, label="Reference strength"), + "width": bounded_int(payload.get("width"), minimum=0, maximum=2048, default=0, label="Width"), + "height": bounded_int(payload.get("height"), minimum=0, maximum=2048, default=0, label="Height"), + }) + plugin_settings = payload.get("settings") or {} + if isinstance(plugin_settings, dict): + extra.update(plugin_settings) + plan = build_generation_plan( + tool, + model_name=model.name, + model_path=model.path, + prompt=prompt, + image_count=count, + steps=steps, + seed=seed, + sampler=sampler, + aspect_ratio=aspect, + extra_arguments=extra, + ) + job = self.dispatcher.submit_job(self.jobs, plan) + return { + "message": f"Queued {job.plan.project_name}.", + "job_id": job.id, + "requires_approval": bool(plan.requires_confirmation), + "plan": self._plan_payload(plan), + } + + def job_detail(self, job_id: str) -> dict[str, Any]: + if self.jobs is None or not hasattr(self.jobs, "get"): + raise RemoteApiError("Jobs are not available in this ADAM session.", status=503) + job = self.jobs.get(job_id) + if job is None: + raise RemoteApiError("ADAM could not find that job.", status=404) + return { + "id": job.id, + "project": job.plan.project_name, + "status": job.status.value, + "progress": job.progress, + "current_step": job.current_step, + "logs": list(job.logs)[-40:], + "steps": [ + { + "tool_id": step.tool_id, + "title": step.title, + "description": step.description, + "status": step.status.value, + "arguments": sanitized_arguments(dict(step.arguments)), + } + for step in job.plan.steps + ], + } + + @staticmethod + def _plan_payload(plan: Any) -> dict[str, Any]: + return { + "id": getattr(plan, "id", ""), + "summary": getattr(plan, "summary", ""), + "project_name": getattr(plan, "project_name", ""), + "requires_confirmation": bool(getattr(plan, "requires_confirmation", False)), + "confirmation_reason": getattr(plan, "confirmation_reason", ""), + "steps": [ + { + "tool_id": step.tool_id, + "title": step.title, + "description": step.description, + "arguments": sanitized_arguments(dict(step.arguments)), + } + for step in getattr(plan, "steps", []) + ], + } diff --git a/adam/tool_folders.py b/adam/tool_folders.py index 47117a5ba2edf8ca89ec1b14f78956c521f865fb..a7e9c43cb034dcb65104502e2dd354664eadba3b 100644 --- a/adam/tool_folders.py +++ b/adam/tool_folders.py @@ -51,6 +51,11 @@ TOOL_FOLDER_DEFINITIONS = ( "Flow Matching Trainer", ("flow_matching_app.py", "roblox_action_flow_app.py"), ), + ToolFolderDefinition( + "oasis_trainer", + "Oasis Game Trainer", + ("roblox_action_flow_app.py", "roblox_action_dataset_recorder.py", "roblox_oasis.py"), + ), ToolFolderDefinition( "preview_generator", "Preview Generator", @@ -154,6 +159,9 @@ class ToolFolderManager: "flow_trainer": ( r"(?im)^\s*Flow(?:\s+Matching)?(?:\s+Trainer)?\s*:\s*(.+?)\s*$" ), + "oasis_trainer": ( + r"(?im)^\s*Oasis(?:\s+(?:Game\s+)?Trainer)?\s*:\s*(.+?)\s*$" + ), "lora_trainer": r"(?im)^\s*LoRA(?:\s+Trainer)?\s*:\s*(.+?)\s*$", "dataset_collector": ( r"(?im)^\s*Dataset(?:\s+Collector)?\s*:\s*(.+?)\s*$" diff --git a/adam/tools/ddpm_adapter.py b/adam/tools/ddpm_adapter.py index 896c0da18532cfb80c3ea3980d50391c0044524e..2707c38a94c0439e6c456bdb0add110f142cd311 100644 --- a/adam/tools/ddpm_adapter.py +++ b/adam/tools/ddpm_adapter.py @@ -14,11 +14,29 @@ import time from pathlib import Path from adam.config import ConfigManager -from adam.executor import ToolCancelled, ToolContext, ToolExecutionError -from adam.process_control import set_process_tree_paused +from adam.executor import ToolAdjustmentRequested, ToolCancelled, ToolContext, ToolExecutionError +from adam.process_control import set_process_tree_paused, terminate_process_tree IMAGE_EXTENSIONS = {".jpg", ".jpeg", ".png", ".webp", ".bmp"} +FORCE_STOP_TIMEOUT_SECONDS = 30 + + +def _saved_unet_resolution(path: Path) -> int | None: + config_path = path / "unet" / "config.json" + if not config_path.is_file(): + return None + try: + data = json.loads(config_path.read_text(encoding="utf-8")) + except (OSError, json.JSONDecodeError): + return None + sample_size = data.get("sample_size") + if isinstance(sample_size, list): + sample_size = sample_size[0] if sample_size else None + try: + return int(sample_size) if sample_size is not None else None + except (TypeError, ValueError): + return None def _latest_preview(folder: Path) -> Path | None: @@ -52,11 +70,25 @@ def _parse_progress( current = int(payload.get("global_step", 0) or 0) total = int(payload.get("total_steps", 0) or 0) if total: - context.progress(max(1, min(99, round(current * 100 / total))), f"Training step {current:,} of {total:,}") + context.progress( + max(1, min(99, round(current * 100 / total))), + f"Training step {current:,} of {total:,}", + current_step=current, + total_steps=total, + epoch=int(payload.get("epoch", 0) or 0), + total_epochs=epochs, + unit="step", + ) return if event == "epoch_end": epoch = int(payload.get("epoch", 0) or 0) - context.progress(max(1, min(99, round(epoch * 100 / max(epochs, 1)))), f"Finished epoch {epoch} of {epochs}") + context.progress( + max(1, min(99, round(epoch * 100 / max(epochs, 1)))), + f"Finished epoch {epoch} of {epochs}", + epoch=epoch, + total_epochs=epochs, + unit="epoch", + ) if preview_enabled and epoch and epoch % preview_every == 0: candidate = Path(str(payload.get("preview_path", ""))) if payload.get("preview_path") else _latest_preview(output) if candidate: @@ -79,6 +111,7 @@ def train_ddpm( mixed_precision: str = "fp16", save_every: int = 10, preview_steps: int = 50, training_intensity: int = 100, preview_enabled: bool = True, preview_every: int = 5, preview_prompt: str = "", preview_seed: int = 123456789, + completed_epochs: int = 0, ) -> dict[str, object]: """Run the registered DDPM project without shell interpolation or overwrites.""" trainer_root = Path(str(ConfigManager(context.root).get("tool_folders", {}).get("ddpm_trainer", ""))).expanduser() @@ -157,14 +190,25 @@ def train_ddpm( raise ToolExecutionError("The DDPM checkpoint must be inside its model folder.") from exc if not resume.is_dir() or not resume.name.startswith("checkpoint-"): raise ToolExecutionError("A valid DDPM checkpoint-* folder is required to resume.") - steps_per_epoch = max(1, math.ceil(image_count / int(batch_size))) - checkpoint_step = int(resume.name.rsplit("-", 1)[-1]) - completed_epochs = checkpoint_step // steps_per_epoch - epochs = completed_epochs + int(epochs) - context.log( - f"Continuing after approximately {completed_epochs} completed epochs " - f"for {int(epochs) - completed_epochs} additional epochs." - ) + checkpoint_resolution = _saved_unet_resolution(resume) + if checkpoint_resolution and checkpoint_resolution != int(resolution) and (output / "model_index.json").is_file(): + pretrained_model = output + timestamp = time.strftime("%Y%m%d_%H%M%S") + output = output.with_name(f"{output.name}_finetuned_{timestamp}") + resume = None + context.log( + f"Changing DDPM resolution from {checkpoint_resolution}px to {int(resolution)}px. " + "Starting a fresh fine-tune from the saved model weights instead of resuming the old optimizer schedule." + ) + if resume is not None: + steps_per_epoch = max(1, math.ceil(image_count / int(batch_size))) + checkpoint_step = int(resume.name.rsplit("-", 1)[-1]) + prior_epochs = int(completed_epochs) if int(completed_epochs) > 0 else checkpoint_step // steps_per_epoch + epochs = prior_epochs + int(epochs) + context.log( + f"Continuing after approximately {prior_epochs} completed epochs " + f"for {int(epochs) - prior_epochs} additional epochs." + ) if not resume: if output.exists(): raise ToolExecutionError("The chosen DDPM output folder already exists; ADAM will not overwrite it.") @@ -182,11 +226,13 @@ def train_ddpm( "--save_model_epochs", str(int(save_every)), "--training_intensity", str(int(training_intensity)), "--dataloader_num_workers", str(int(dataloader_num_workers)), "--gradient_accumulation_steps", str(int(gradient_accumulation_steps)), "--preview_num_inference_steps", str(int(preview_steps)), "--preview_sampler", "DDIM", "--pin_memory", "true", "--stop_signal_file", str(stop_file), - "--checkpointing_steps", str(max(1, image_count)), + "--checkpointing_steps", str(max(1, math.ceil(image_count / int(batch_size)))), "--checkpoints_total_limit", "1", "--keep_latest_resume_checkpoint", "--gui_progress", ] if resume: command.extend(["--resume_from_checkpoint", resume.name]) + if int(completed_epochs) > 0: + command.extend(["--resume_completed_epochs", str(int(completed_epochs))]) if pretrained_model: command.extend(["--pretrained_model_path", str(pretrained_model)]) context.log(f"Starting real DDPM training with {image_count} images at {resolution}px, batch {batch_size}, lr {learning_rate}.") @@ -205,13 +251,23 @@ def train_ddpm( context.progress(1, "Starting DDPM trainer") stop_requested_at: float | None = None cancelled = False + force_stop_sent = False suspended = False + adjustment_requested = False + stopped_details: dict[str, object] = {} while True: should_pause = not context.run_event.is_set() if should_pause != suspended: if set_process_tree_paused(process, should_pause): suspended = should_pause context.log("DDPM trainer paused safely." if suspended else "DDPM trainer resumed.") + if context.adjustment_event and context.adjustment_event.is_set() and not adjustment_requested: + if suspended: + set_process_tree_paused(process, False) + suspended = False + stop_file.write_text("after_epoch", encoding="utf-8") + adjustment_requested = True + context.log("Settings change accepted; finishing this epoch and saving a resume checkpoint.") if context.cancel_event.is_set() and stop_requested_at is None: if suspended: set_process_tree_paused(process, False) @@ -220,16 +276,20 @@ def train_ddpm( stop_requested_at = time.monotonic() cancelled = True context.log("Safe stop requested; waiting for DDPM to finish its current step.") - if stop_requested_at and time.monotonic() - stop_requested_at > 75: - process.terminate() + if stop_requested_at and not force_stop_sent and time.monotonic() - stop_requested_at > FORCE_STOP_TIMEOUT_SECONDS: + terminate_process_tree(process, timeout=3) + force_stop_sent = True context.log("DDPM did not stop in time; terminating the trainer process.") try: line = lines.get(timeout=0.12) if line: if line.startswith("PROGRESS_JSON:"): try: + payload = json.loads(line.split(":", 1)[1]) + if str(payload.get("event", "")) == "stopped": + stopped_details = payload if not cancelled: - _parse_progress(context, json.loads(line.split(":", 1)[1]), int(epochs), output, + _parse_progress(context, payload, int(epochs), output, bool(preview_enabled), int(preview_every), preview_prompt, int(preview_seed), int(preview_steps)) except json.JSONDecodeError: @@ -242,6 +302,18 @@ def train_ddpm( break if cancelled: raise ToolCancelled("DDPM training stopped by user.") + if adjustment_requested: + checkpoints = sorted(output.glob("checkpoint-*"), key=lambda path: int(path.name.rsplit("-", 1)[-1])) + if not checkpoints: + raise ToolExecutionError("Training stopped for adjustment, but no complete resume checkpoint was found.") + raise ToolAdjustmentRequested( + "DDPM training reached a safe epoch boundary.", + { + "checkpoint": str(checkpoints[-1]), + "completed_epochs": int(stopped_details.get("completed_epochs", 0) or 0), + "updates": dict(context.adjustment_request or {}), + }, + ) if process.returncode != 0: raise ToolExecutionError(f"DDPM trainer exited with code {process.returncode}. See the job log for details.") context.progress(100, "DDPM training completed") diff --git a/adam/tools/ddpm_generator.py b/adam/tools/ddpm_generator.py index 9a2a5fd0e128abc4e7d2c0439ded66bb5adc38b7..c0678cb86d45fb1dfe8b7ca0a654e133a204bada 100644 --- a/adam/tools/ddpm_generator.py +++ b/adam/tools/ddpm_generator.py @@ -15,6 +15,7 @@ from adam.config import ConfigManager from adam.executor import ToolContext, ToolExecutionError from adam.generations import generation_metadata_path, generation_output_folder from adam.generation_previews import accepts_preview_callback, publish_generation_preview +from adam.image_preferences import GenerationPreferenceEvaluator, PreferenceProfile _backend_module: ModuleType | None = None @@ -67,6 +68,12 @@ def generate_ddpm_images( width: int = 0, height: int = 0, preview_interval: int = 0, + smart_generation: bool = False, + smart_wanted_results: int = 0, + smart_max_candidates: int = 0, + smart_min_score: float = 0.70, + smart_mode: str = "threshold", + smart_keep_rejected: bool = True, ) -> dict[str, object]: config = ConfigManager(context.root) trainer_root = Path( @@ -104,6 +111,19 @@ def generate_ddpm_images( count = int(image_count) step_count = int(steps) + smart_enabled = bool(smart_generation) + wanted_results = int(smart_wanted_results or count) + max_candidates = int(smart_max_candidates or count) + if smart_enabled: + if not 1 <= wanted_results <= 48: + raise ToolExecutionError("Wanted Smart Generation results must be between 1 and 48.") + if not wanted_results <= max_candidates <= 256: + raise ToolExecutionError("Maximum Smart Generation candidates must be between wanted results and 256.") + if not 0.0 <= float(smart_min_score) <= 1.0: + raise ToolExecutionError("Minimum Smart Generation score must be between 0.00 and 1.00.") + if str(smart_mode).casefold() not in {"threshold", "top_n"}: + raise ToolExecutionError("Smart Generation mode must be threshold or top_n.") + count = wanted_results if not 1 <= count <= 48: raise ToolExecutionError("Image count must be between 1 and 48.") if not 5 <= step_count <= 500: @@ -147,7 +167,8 @@ def generate_ddpm_images( base_seed = int(seed) if base_seed <= 0: base_seed = random.randint(1, 2_147_483_647 - count) - if base_seed + count - 1 > 2_147_483_647: + seed_count = max_candidates if smart_enabled else count + if base_seed + seed_count - 1 > 2_147_483_647: raise ToolExecutionError("The seed is too large for this image count.") timestamp = datetime.now(timezone.utc).strftime("%Y%m%d_%H%M%S") @@ -169,12 +190,19 @@ def generate_ddpm_images( if preview_enabled and not preview_supported: context.log("This connected DDPM generator does not yet expose denoising previews; generation will continue normally.") image_paths: list[str] = [] - for index in range(count): + image_evaluations: dict[str, dict[str, object]] = {} + selected_paths: list[str] = [] + generated_total = max_candidates if smart_enabled else count + profile = PreferenceProfile(context.root, context.tool.id, _safe_label(model_name), str(model)) if smart_enabled else None + evaluator = GenerationPreferenceEvaluator(context.root) if smart_enabled else None + threshold = float(smart_min_score) + top_n_mode = str(smart_mode).casefold() == "top_n" + for index in range(generated_total): context.checkpoint() current_seed = base_seed + index context.progress( - max(1, round(index * 100 / count)), - f"Generating image {index + 1} of {count}", + max(1, round(index * 100 / generated_total)), + f"Generating image {index + 1} of {generated_total}" if smart_enabled else f"Generating image {index + 1} of {count}", ) try: settings = { @@ -204,8 +232,31 @@ def generate_ddpm_images( except Exception as exc: raise ToolExecutionError(f"DDPM generation failed: {exc}") from exc image_paths.append(str(destination)) + if smart_enabled and profile is not None and evaluator is not None: + score = evaluator.score(profile, [destination], keep_threshold=threshold, reject_threshold=profile.reject_threshold)[0] + image_evaluations[str(destination.resolve())] = { + "score": score.score, + "confidence": score.confidence, + "category": score.category, + "reason": score.reason, + } + if not top_n_mode and score.score is not None and score.score >= threshold: + selected_paths.append(str(destination)) + if len(selected_paths) >= wanted_results: + break created_at = datetime.now(timezone.utc).isoformat() + if smart_enabled and top_n_mode: + ranked = sorted( + image_paths, + key=lambda path: float(image_evaluations.get(str(Path(path).resolve()), {}).get("score") or -1.0), + reverse=True, + ) + selected_paths = ranked[:wanted_results] + if smart_enabled: + ordered_images = [*selected_paths, *[path for path in image_paths if path not in set(selected_paths)]] + else: + ordered_images = image_paths metadata = { "version": 1, "provider_id": context.tool.id, @@ -216,7 +267,7 @@ def generate_ddpm_images( "prompt_behavior": "label_only", "seed": base_seed, "image_seeds": [base_seed + index for index in range(count)], - "image_count": count, + "image_count": len(ordered_images) if smart_enabled and smart_keep_rejected else count, "steps": step_count, "sampler": sampler, "aspect_ratio": aspect_ratio, @@ -226,13 +277,30 @@ def generate_ddpm_images( "reference_strength": strength if reference_path else None, "preview_interval": int(preview_interval), "preview_supported": preview_supported, - "images": image_paths, + "images": ordered_images if smart_keep_rejected or not smart_enabled else selected_paths, + "image_evaluations": image_evaluations, + "smart_generation": { + "enabled": smart_enabled, + "mode": str(smart_mode), + "wanted_results": wanted_results if smart_enabled else count, + "maximum_candidates": max_candidates if smart_enabled else count, + "minimum_score": threshold, + "selected_count": len(selected_paths) if smart_enabled else count, + "candidate_count": len(image_paths), + "profile_id": profile.id if profile else "", + "keep_rejected_candidates": bool(smart_keep_rejected), + }, "created_at": created_at, } generation_metadata_path(output, timestamp, context.job_id).write_text( json.dumps(metadata, indent=2), encoding="utf-8" ) - context.progress(100, f"Generated {count} image(s)") + if evaluator is not None: + evaluator.vision.unload() + if smart_enabled: + context.progress(100, f"Smart Generation selected {len(selected_paths)} of {wanted_results} requested image(s) from {len(image_paths)} candidate(s)") + else: + context.progress(100, f"Generated {count} image(s)") return { "output_folder": str(output), "assets": [ diff --git a/adam/tools/flow_adapter.py b/adam/tools/flow_adapter.py index ac33b44c0287e576fd957fa7ce4bc51688271a19..f31a1b710072d28ef821b6c0a756d889ef0efd45 100644 --- a/adam/tools/flow_adapter.py +++ b/adam/tools/flow_adapter.py @@ -7,14 +7,16 @@ import queue import subprocess import sys import threading +import time from pathlib import Path from adam.config import ConfigManager from adam.executor import ToolCancelled, ToolContext, ToolExecutionError -from adam.process_control import set_process_tree_paused +from adam.process_control import set_process_tree_paused, terminate_process_tree IMAGE_EXTENSIONS = {".jpg", ".jpeg", ".png", ".webp", ".bmp"} +FORCE_STOP_TIMEOUT_SECONDS = 30 def _latest_preview(folder: Path) -> Path | None: @@ -104,6 +106,8 @@ def train_flow( threading.Thread(target=read_output, daemon=True).start() context.progress(1, "Starting Flow Matching trainer") stopped = False + stop_requested_at: float | None = None + force_stop_sent = False suspended = False stop_file = output / "stop_flow_training.flag" while True: @@ -118,7 +122,12 @@ def train_flow( suspended = False stop_file.touch(exist_ok=True) stopped = True + stop_requested_at = time.monotonic() context.log("Safe stop requested; waiting for Flow Matching to finish its current batch.") + if stop_requested_at and not force_stop_sent and time.monotonic() - stop_requested_at > FORCE_STOP_TIMEOUT_SECONDS: + terminate_process_tree(process, timeout=3) + force_stop_sent = True + context.log("Flow Matching did not stop in time; terminating the trainer process.") try: line = lines.get(timeout=0.15) if line and line.startswith("FLOW_EVENT:"): diff --git a/adam/tools/flow_generator.py b/adam/tools/flow_generator.py index c8d68bc598caceb5cb51d5c60e3625549ad9a552..833c44bc5a235f1684349d9b1f0a411dbc058944 100644 --- a/adam/tools/flow_generator.py +++ b/adam/tools/flow_generator.py @@ -15,6 +15,7 @@ from adam.config import ConfigManager from adam.executor import ToolCancelled, ToolContext, ToolExecutionError from adam.generations import generation_metadata_path, generation_output_folder from adam.generation_previews import accepts_preview_callback, publish_generation_preview +from adam.image_preferences import GenerationPreferenceEvaluator, PreferenceProfile _backend_module: ModuleType | None = None @@ -67,6 +68,12 @@ def generate_flow_images( sampler: str, aspect_ratio: str, preview_interval: int = 0, + smart_generation: bool = False, + smart_wanted_results: int = 0, + smart_max_candidates: int = 0, + smart_min_score: float = 0.70, + smart_mode: str = "threshold", + smart_keep_rejected: bool = True, ) -> dict[str, object]: global _loaded_model, _loaded_model_path config = ConfigManager(context.root) @@ -108,6 +115,19 @@ def generate_flow_images( ) count = int(image_count) step_count = int(steps) + smart_enabled = bool(smart_generation) + wanted_results = int(smart_wanted_results or count) + max_candidates = int(smart_max_candidates or count) + if smart_enabled: + if not 1 <= wanted_results <= 48: + raise ToolExecutionError("Wanted Smart Generation results must be between 1 and 48.") + if not wanted_results <= max_candidates <= 256: + raise ToolExecutionError("Maximum Smart Generation candidates must be between wanted results and 256.") + if not 0.0 <= float(smart_min_score) <= 1.0: + raise ToolExecutionError("Minimum Smart Generation score must be between 0.00 and 1.00.") + if str(smart_mode).casefold() not in {"threshold", "top_n"}: + raise ToolExecutionError("Smart Generation mode must be threshold or top_n.") + count = wanted_results if not 1 <= count <= 48: raise ToolExecutionError("Image count must be between 1 and 48.") if not 1 <= step_count <= 200: @@ -129,7 +149,8 @@ def generate_flow_images( base_seed = int(seed) if base_seed <= 0: base_seed = random.randint(1, 2_147_483_647 - count) - if base_seed + count - 1 > 2_147_483_647: + seed_count = max_candidates if smart_enabled else count + if base_seed + seed_count - 1 > 2_147_483_647: raise ToolExecutionError("The seed is too large for this image count.") backend = _load_backend(script) @@ -158,15 +179,22 @@ def generate_flow_images( output = generation_output_folder(context.root, context.tool.id, _safe_label(model_name)) context.log("Flow Matching models generate learned visual samples; the label is metadata, not a text prompt.") image_paths: list[str] = [] - for index in range(count): + image_evaluations: dict[str, dict[str, object]] = {} + selected_paths: list[str] = [] + generated_total = max_candidates if smart_enabled else count + profile = PreferenceProfile(context.root, context.tool.id, _safe_label(model_name), str(model)) if smart_enabled else None + evaluator = GenerationPreferenceEvaluator(context.root) if smart_enabled else None + threshold = float(smart_min_score) + top_n_mode = str(smart_mode).casefold() == "top_n" + for index in range(generated_total): context.checkpoint() current_seed = base_seed + index def on_progress(done: int, total: int, image_index: int = index) -> None: completed = image_index + (done / max(1, total)) context.progress( - max(1, min(99, round(completed * 100 / count))), - f"Generating image {image_index + 1} of {count} · flow step {done} of {total}", + max(1, min(99, round(completed * 100 / generated_total))), + f"Generating image {image_index + 1} of {generated_total} · flow step {done} of {total}" if smart_enabled else f"Generating image {image_index + 1} of {count} · flow step {done} of {total}", ) try: @@ -195,8 +223,31 @@ def generate_flow_images( except Exception as exc: raise ToolExecutionError(f"Flow Matching generation failed: {exc}") from exc image_paths.append(str(destination)) + if smart_enabled and profile is not None and evaluator is not None: + score = evaluator.score(profile, [destination], keep_threshold=threshold, reject_threshold=profile.reject_threshold)[0] + image_evaluations[str(destination.resolve())] = { + "score": score.score, + "confidence": score.confidence, + "category": score.category, + "reason": score.reason, + } + if not top_n_mode and score.score is not None and score.score >= threshold: + selected_paths.append(str(destination)) + if len(selected_paths) >= wanted_results: + break created_at = datetime.now(timezone.utc).isoformat() + if smart_enabled and top_n_mode: + ranked = sorted( + image_paths, + key=lambda path: float(image_evaluations.get(str(Path(path).resolve()), {}).get("score") or -1.0), + reverse=True, + ) + selected_paths = ranked[:wanted_results] + if smart_enabled: + ordered_images = [*selected_paths, *[path for path in image_paths if path not in set(selected_paths)]] + else: + ordered_images = image_paths metadata = { "version": 1, "provider_id": context.tool.id, @@ -207,19 +258,36 @@ def generate_flow_images( "prompt_behavior": "label_only", "seed": base_seed, "image_seeds": [base_seed + index for index in range(count)], - "image_count": count, + "image_count": len(ordered_images) if smart_enabled and smart_keep_rejected else count, "steps": step_count, "sampler": method, "aspect_ratio": aspect_ratio, "preview_interval": int(preview_interval), "preview_supported": preview_supported, - "images": image_paths, + "images": ordered_images if smart_keep_rejected or not smart_enabled else selected_paths, + "image_evaluations": image_evaluations, + "smart_generation": { + "enabled": smart_enabled, + "mode": str(smart_mode), + "wanted_results": wanted_results if smart_enabled else count, + "maximum_candidates": max_candidates if smart_enabled else count, + "minimum_score": threshold, + "selected_count": len(selected_paths) if smart_enabled else count, + "candidate_count": len(image_paths), + "profile_id": profile.id if profile else "", + "keep_rejected_candidates": bool(smart_keep_rejected), + }, "created_at": created_at, } generation_metadata_path(output, timestamp, context.job_id).write_text( json.dumps(metadata, indent=2), encoding="utf-8" ) - context.progress(100, f"Generated {count} image(s)") + if evaluator is not None: + evaluator.vision.unload() + if smart_enabled: + context.progress(100, f"Smart Generation selected {len(selected_paths)} of {wanted_results} requested image(s) from {len(image_paths)} candidate(s)") + else: + context.progress(100, f"Generated {count} image(s)") return { "output_folder": str(output), "assets": [ diff --git a/adam/tools/lora_adapter.py b/adam/tools/lora_adapter.py index 2d1bbb6bd5e0c09443708bbd4d443a3bc3ff99a7..dc2578c621178eeaf03524e54f5af599d8438fca 100644 --- a/adam/tools/lora_adapter.py +++ b/adam/tools/lora_adapter.py @@ -3,6 +3,8 @@ from __future__ import annotations import json +import multiprocessing +import queue import sys import time from pathlib import Path @@ -10,9 +12,11 @@ from typing import Any from adam.config import ConfigManager from adam.executor import ToolCancelled, ToolContext, ToolExecutionError +from adam.process_control import terminate_process_tree IMAGE_EXTENSIONS = {".jpg", ".jpeg", ".png", ".webp", ".bmp"} +FORCE_STOP_TIMEOUT_SECONDS = 30 class _Control: @@ -36,6 +40,25 @@ class _Control: return +class _ProcessControl: + def __init__(self, cancel_event: Any, run_event: Any) -> None: + self._cancel_event = cancel_event + self._run_event = run_event + + @property + def cancel_requested(self) -> bool: + return self._cancel_event.is_set() + + @property + def pause_requested(self) -> bool: + return not self._run_event.is_set() + + def wait_if_paused(self) -> None: + while not self._run_event.wait(timeout=0.15): + if self._cancel_event.is_set(): + return + + def _safe_name(value: str) -> str: name = value.strip() if not name or len(name) > 96 or any(char in name for char in '<>:"/\\|?*\x00'): @@ -43,6 +66,77 @@ def _safe_name(value: str) -> str: return name +def _run_lora_training_worker( + payload: dict[str, Any], + events: Any, + cancel_event: Any, + run_event: Any, +) -> None: + try: + source_root = Path(str(payload["source_root"])) + sys.path.insert(0, str(source_root)) + from loratrainer.models.training_config import TrainingConfig + from loratrainer.trainer.diffusers_sdxl_lora_backend import ( + DiffusersSDXLLoRABackend, + ) + + valid_fields = set(TrainingConfig.__dataclass_fields__) + blocked = {"dataset_dir", "base_model_path", "output_dir", "resume_checkpoint"} + overrides = { + key: value + for key, value in dict(payload.get("settings", {})).items() + if key in valid_fields and key not in blocked + } + overrides.update( + { + key: value + for key, value in dict(payload.get("training_overrides", {})).items() + if key in valid_fields and key not in {*blocked, "epochs"} + } + ) + overrides["trigger_word"] = str(payload.get("trigger_word") or payload["model_name"]) + overrides["epochs"] = int(payload["epochs"]) + config = TrainingConfig( + dataset_dir=Path(str(payload["dataset"])), + base_model_path=Path(str(payload["base"])), + output_dir=Path(str(payload["output"])), + resume_checkpoint=Path(str(payload["resume"])) if payload.get("resume") else None, + **overrides, + ) + control = _ProcessControl(cancel_event, run_event) + + def progress(update: Any) -> None: + if cancel_event.is_set(): + return + events.put( + { + "type": "progress", + "total_steps": int(getattr(update, "total_steps", 0) or 0), + "step": int(getattr(update, "step", 0) or 0), + "epoch": int(getattr(update, "epoch", 0) or 0), + "total_epochs": int(getattr(update, "total_epochs", payload["epochs"]) or payload["epochs"]), + "message": str(getattr(update, "message", "") or ""), + "preview_path": str(getattr(update, "preview_path", "") or ""), + } + ) + + final_path = Path( + DiffusersSDXLLoRABackend().train(config, control, progress) + ).resolve() + events.put({"type": "result", "final_path": str(final_path)}) + except Exception as exc: + if cancel_event.is_set(): + events.put({"type": "cancelled", "message": "LoRA training stopped by user."}) + return + events.put( + { + "type": "error", + "message": str(exc), + "exception": type(exc).__name__, + } + ) + + def train_lora( context: ToolContext, dataset_dir: str, @@ -50,9 +144,11 @@ def train_lora( epochs: int, output_dir: str, base_model: str, + trigger_word: str = "", resume_from: str = "", preview_enabled: bool = True, preview_every: int = 5, preview_prompt: str = "", preview_seed: int = 123456789, + **training_overrides: Any, ) -> dict[str, Any]: folders = ConfigManager(context.root).get("tool_folders", {}) trainer_root = Path(str(folders.get("lora_trainer", ""))).expanduser().resolve() @@ -68,6 +164,9 @@ def train_lora( output = Path(output_dir).expanduser().resolve() resume = Path(resume_from).expanduser().resolve() if resume_from else None name = _safe_name(model_name) + trigger = str(trigger_word or name).strip() + if not trigger or len(trigger) > 128 or any(char in trigger for char in '<>:"/\\|?*\x00'): + raise ToolExecutionError("Choose a short LoRA trigger word without reserved characters.") if not dataset.is_dir(): raise ToolExecutionError("The selected LoRA dataset folder no longer exists.") images = [ @@ -110,48 +209,35 @@ def train_lora( except (OSError, ValueError, TypeError, json.JSONDecodeError): pass - sys.path.insert(0, str(source_root)) - try: - from loratrainer.models.training_config import TrainingConfig - from loratrainer.trainer.diffusers_sdxl_lora_backend import ( - DiffusersSDXLLoRABackend, - ) - except Exception as exc: - raise ToolExecutionError(f"Could not load the connected LoRA backend: {exc}") from exc - - valid_fields = set(TrainingConfig.__dataclass_fields__) - overrides = { - key: value - for key, value in settings.items() - if key in valid_fields - and key not in {"dataset_dir", "base_model_path", "output_dir", "resume_checkpoint"} - } - overrides["trigger_word"] = name - overrides["epochs"] = int(epochs) - config = TrainingConfig( - dataset_dir=dataset, - base_model_path=base, - output_dir=output, - resume_checkpoint=resume, - **overrides, - ) - control = _Control(context) - - def progress(update: Any) -> None: - if context.cancel_event.is_set(): - return + def relay_progress(update: dict[str, Any]) -> None: total_steps = int(getattr(update, "total_steps", 0) or 0) - step = int(getattr(update, "step", 0) or 0) - epoch = int(getattr(update, "epoch", 0) or 0) - total_epochs = int(getattr(update, "total_epochs", epochs) or epochs) + if isinstance(update, dict): + total_steps = int(update.get("total_steps", 0) or 0) + step = int(update.get("step", 0) or 0) + epoch = int(update.get("epoch", 0) or 0) + total_epochs = int(update.get("total_epochs", epochs) or epochs) + message = str(update.get("message", "") or f"LoRA epoch {epoch}/{total_epochs}") + preview_path = str(update.get("preview_path", "") or "") + else: + step = int(getattr(update, "step", 0) or 0) + epoch = int(getattr(update, "epoch", 0) or 0) + total_epochs = int(getattr(update, "total_epochs", epochs) or epochs) + message = str(getattr(update, "message", "") or f"LoRA epoch {epoch}/{total_epochs}") + preview_path = str(getattr(update, "preview_path", "") or "") percent = ( round(step * 100 / total_steps) if total_steps else round(epoch * 100 / max(total_epochs, 1)) ) - message = str(getattr(update, "message", "") or f"LoRA epoch {epoch}/{total_epochs}") - context.progress(max(1, min(percent, 99)), message) - preview_path = str(getattr(update, "preview_path", "") or "") + context.progress( + max(1, min(percent, 99)), + message, + current_step=step, + total_steps=total_steps, + epoch=epoch, + total_epochs=total_epochs, + unit="step" if total_steps else "epoch", + ) if preview_enabled and preview_path and epoch and epoch % max(1, int(preview_every)) == 0: context.preview(preview_path, epoch=epoch, next_epoch=min(total_epochs, epoch + max(1, int(preview_every))), @@ -161,19 +247,97 @@ def train_lora( context.log(f"Base model: {base}") if resume: context.log(f"Continuing from LoRA: {resume}") - try: - final_path = Path( - DiffusersSDXLLoRABackend().train(config, control, progress) - ).resolve() - except Exception as exc: - if context.cancel_event.is_set(): - raise ToolCancelled("LoRA training stopped by user.") from exc - raise ToolExecutionError(f"LoRA trainer failed: {exc}") from exc + payload = { + "source_root": str(source_root), + "dataset": str(dataset), + "base": str(base), + "output": str(output), + "resume": str(resume) if resume else "", + "model_name": name, + "trigger_word": trigger, + "epochs": int(epochs), + "settings": settings, + "training_overrides": training_overrides, + } + mp_context = multiprocessing.get_context("spawn") + events = mp_context.Queue() + process_cancel = mp_context.Event() + process_run = mp_context.Event() + process_run.set() + process = mp_context.Process( + target=_run_lora_training_worker, + args=(payload, events, process_cancel, process_run), + daemon=True, + ) + process.start() + final_path: Path | None = None + cancelled = False + stop_requested_at: float | None = None + force_stop_sent = False + process_paused = False + + def handle_child_event(event: dict[str, Any]) -> None: + nonlocal final_path, cancelled + event_type = str(event.get("type", "")) + if event_type == "progress" and not cancelled: + relay_progress(event) + elif event_type == "result": + final_path = Path(str(event.get("final_path", ""))).resolve() + elif event_type == "cancelled": + cancelled = True + elif event_type == "error": + message = str(event.get("message", "LoRA trainer failed.")) + terminate_process_tree(process, timeout=3) + process.join(timeout=1) + raise ToolExecutionError(f"LoRA trainer failed: {message}") + + while True: + if context.cancel_event.is_set() and stop_requested_at is None: + process_cancel.set() + process_run.set() + stop_requested_at = time.monotonic() + cancelled = True + context.log("Safe stop requested; waiting for LoRA training to finish its current step.") + should_pause = not context.run_event.is_set() and not cancelled + if should_pause != process_paused: + if should_pause: + process_run.clear() + context.log("LoRA trainer paused safely.") + else: + process_run.set() + context.log("LoRA trainer resumed.") + process_paused = should_pause + if stop_requested_at and not force_stop_sent and time.monotonic() - stop_requested_at > FORCE_STOP_TIMEOUT_SECONDS: + terminate_process_tree(process, timeout=3) + force_stop_sent = True + context.log("LoRA trainer did not stop in time; terminating the trainer process.") + try: + event = events.get(timeout=0.15) + except queue.Empty: + event = None + if event: + handle_child_event(event) + if not process.is_alive() and event is None: + break + + process.join(timeout=1) + while True: + try: + handle_child_event(events.get_nowait()) + except queue.Empty: + break + if cancelled: + raise ToolCancelled("LoRA training stopped by user.") + if process.exitcode not in {0, None}: + raise ToolExecutionError(f"LoRA trainer exited with code {process.exitcode}.") + if final_path is None: + raise ToolExecutionError("LoRA trainer finished without reporting a checkpoint.") context.progress(100, "LoRA training completed") return { "output_folder": str(output), "model_name": name, + "trigger_word": trigger, "assets": [ { "kind": "model", @@ -183,6 +347,8 @@ def train_lora( "dataset_path": str(dataset), "checkpoint": str(final_path), "epochs": int(epochs), + "metadata": {"trigger_word": trigger}, + "trigger_word": trigger, } ], } diff --git a/adam/tools/oasis_adapter.py b/adam/tools/oasis_adapter.py new file mode 100644 index 0000000000000000000000000000000000000000..6ab6a5d2a545e9dfba7c6046839271b82ffcb772 --- /dev/null +++ b/adam/tools/oasis_adapter.py @@ -0,0 +1,346 @@ +"""Execution bridge for the connected Oasis action world model trainer.""" + +from __future__ import annotations + +import importlib.util +import json +import os +import queue +import subprocess +import sys +import threading +import time +from pathlib import Path +from typing import Any + +from adam.config import ConfigManager +from adam.executor import ToolCancelled, ToolContext, ToolExecutionError +from adam.oasis_dataset import validate_oasis_dataset +from adam.process_control import set_process_tree_paused, terminate_process_tree + + +FORCE_STOP_TIMEOUT_SECONDS = 30 +IMAGE_EXTENSIONS = {".jpg", ".jpeg", ".png", ".webp", ".bmp"} + + +def _external_oasis_root(context: ToolContext) -> Path | None: + path = context.root / "config" / "external_tools.json" + try: + payload = json.loads(path.read_text(encoding="utf-8")) + except (OSError, json.JSONDecodeError): + return None + for entry in payload.get("tools", []): + if not isinstance(entry, dict): + continue + if str(entry.get("id", "")) != "external_oasis_game_trainer": + continue + root = Path(str(entry.get("backend", {}).get("root", ""))).expanduser() + return root if root.is_dir() else None + return None + + +def _oasis_root(context: ToolContext) -> Path: + folders = ConfigManager(context.root).get("tool_folders", {}) + raw = str(folders.get("oasis_trainer", "")) if isinstance(folders, dict) else "" + root = Path(raw).expanduser() if raw else (_external_oasis_root(context) or Path()) + script = root / "roblox_action_flow_app.py" + if not script.is_file(): + raise ToolExecutionError( + "Oasis roblox_action_flow_app.py was not found. Connect the Oasis-Game-Trainer folder in Settings." + ) + return root.resolve() + + +def _safe_model_name(value: str) -> str: + name = value.strip() + if not name or len(name) > 96 or any(char in name for char in "<>:\"/\\|?*\x00"): + raise ToolExecutionError("Choose a short Oasis model name without filesystem-reserved characters.") + return name + + +def _valid_action_model(path: Path) -> bool: + try: + info = json.loads((path / "action_flow_model_info.json").read_text(encoding="utf-8")) + return ( + info.get("model_type") == "action_conditioned_rectified_flow_video" + and (path / "unet" / "config.json").is_file() + ) + except (OSError, ValueError, TypeError, json.JSONDecodeError): + return False + + +def _latest_preview(folder: Path) -> Path | None: + try: + images = [ + path for path in folder.rglob("*") + if path.is_file() + and path.suffix.casefold() in IMAGE_EXTENSIONS + and any(token in path.name.casefold() for token in ("preview", "sample", "epoch")) + ] + return max(images, key=lambda path: path.stat().st_mtime) if images else None + except OSError: + return None + + +def _parse_event(context: ToolContext, event: dict[str, Any], epochs: int, output: Path, preview_enabled: bool, preview_every: int, preview_steps: int) -> None: + kind = str(event.get("type", "")) + if kind == "start": + context.log( + f"Oasis dataset ready: {event.get('transitions', 0):,} transitions, " + f"{event.get('training', 0):,} training, {event.get('validation', 0):,} validation on {event.get('device', 'device')}." + ) + context.progress(1, "Starting Oasis trainer") + elif kind == "progress": + update = int(event.get("update", 0) or 0) + total = int(event.get("total_updates", 0) or 0) + epoch = int(event.get("epoch", 0) or 0) + loss = event.get("loss") + eta = event.get("eta") + message = f"Oasis epoch {epoch} of {epochs}" + if total: + message += f", step {update:,} of {total:,}" + if isinstance(loss, (int, float)): + message += f", loss {loss:.4f}" + context.progress( + max(1, min(99, round(update * 100 / total))) if total else max(1, min(99, round(epoch * 100 / max(epochs, 1)))), + message, + current_step=update, + total_steps=total, + epoch=epoch, + total_epochs=epochs, + eta_seconds=eta if isinstance(eta, (int, float)) else None, + unit="step" if total else "epoch", + ) + elif kind == "validation": + validation = event.get("validation", {}) + loss = validation.get("total_loss") if isinstance(validation, dict) else None + context.log(f"Oasis validation epoch {event.get('epoch')}: loss {loss:.4f}" if isinstance(loss, (int, float)) else f"Oasis validation epoch {event.get('epoch')}.") + elif kind == "preview" and preview_enabled: + path = Path(str(event.get("path", ""))) if event.get("path") else _latest_preview(output) + if path: + epoch = int(event.get("epoch", 0) or 0) + context.preview(path, epoch=epoch, next_epoch=epoch + preview_every, steps=preview_steps, kind="training") + elif kind == "saved": + context.log(f"Oasis checkpoint saved at epoch {event.get('epoch')}.") + elif kind == "recovery_saved": + context.log(f"Oasis recovery checkpoint saved at batch {event.get('batch')}.") + elif kind == "warning": + context.log(str(event.get("message", "Oasis trainer warning."))) + elif kind == "complete": + context.progress(100, "Oasis training completed") + + +def train_oasis( + context: ToolContext, + dataset_dir: str | list[str], + model_name: str, + epochs: int, + output_dir: str, + resume_from: str = "", + resolution: str = "256x144", + batch_size: int = 2, + learning_rate: float = 0.00002, + workers: int = 2, + mixed_precision: str = "fp32", + gradient_accumulation: int = 1, + frame_gap: int = 3, + sequence_context: int = 1, + action_aggregation: str = "window", + validation_split: float = 0.1, + validation_batches: int = 8, + save_every: int = 5, + preview_enabled: bool = True, + preview_every: int = 5, + preview_steps: int = 1, + seed: int = 1234, + base_model: str = "", + condition_noise: float = 0.03, + temporal_loss_weight: float = 0.1, + motion_loss_weight: float = 2.0, + action_input_scale: float = 8.0, + neutral_action_dropout: float = 0.15, + action_contrast_weight: float = 0.35, + action_contrast_margin: float = 0.02, + gradient_checkpointing: bool = False, + balance_actions: bool = False, +) -> dict[str, object]: + root = _oasis_root(context) + script = root / "roblox_action_flow_app.py" + model_name = _safe_model_name(model_name) + dataset_value = ( + ";".join(str(item) for item in dataset_dir) + if isinstance(dataset_dir, list) + else str(dataset_dir) + ) + output = Path(output_dir).expanduser().resolve() + output_root = (root / "output_action_flow_models").resolve() + try: + output.relative_to(output_root) + except ValueError as exc: + raise ToolExecutionError("Oasis outputs must stay inside output_action_flow_models.") from exc + if output.exists() and not resume_from: + raise ToolExecutionError("The chosen Oasis output folder already exists; ADAM will not overwrite it.") + if mixed_precision not in {"fp32", "fp16", "no"}: + raise ToolExecutionError("Oasis precision must be fp32, fp16, or no.") + if int(sequence_context) != 1: + context.log("Oasis currently trains previous-frame context length 1; the setting is stored for future multi-frame support.") + report = validate_oasis_dataset(dataset_value, frame_gap=int(frame_gap)) + if not report.ok: + raise ToolExecutionError("Oasis dataset validation failed: " + " ".join(report.errors[:12])) + context.log( + f"Oasis dataset validated: {report.valid_rows:,} labelled frames, " + f"{report.valid_transitions:,} transitions, resolution {report.resolution}." + ) + missing_packages = [ + package for package in ("torch", "torchvision", "diffusers", "PIL") + if importlib.util.find_spec(package) is None + ] + if missing_packages: + raise ToolExecutionError( + "ADAM's Python environment is missing Oasis packages: " + + ", ".join(missing_packages) + + ". Install the Oasis requirements, then restart ADAM." + ) + output.parent.mkdir(parents=True, exist_ok=True) + resume = Path(resume_from).expanduser().resolve() if resume_from else None + if resume and not _valid_action_model(resume): + raise ToolExecutionError("Choose a valid Oasis action model folder to resume.") + base = Path(base_model).expanduser().resolve() if base_model else None + command = [ + sys.executable, str(script), "--train-worker", + "--dataset-dir", dataset_value, + "--output-dir", str(output), + "--model-name", model_name, + "--epochs", str(int(epochs)), + "--resolution", str(resolution), + "--batch-size", str(int(batch_size)), + "--learning-rate", str(float(learning_rate)), + "--workers", str(int(workers)), + "--gradient-accumulation", str(int(gradient_accumulation)), + "--condition-noise", str(float(condition_noise)), + "--temporal-loss-weight", str(float(temporal_loss_weight)), + "--motion-loss-weight", str(float(motion_loss_weight)), + "--validation-split", str(float(validation_split)), + "--validation-batches", str(int(validation_batches)), + "--metrics-every", "10", + "--frame-gap", str(int(frame_gap)), + "--action-aggregation", str(action_aggregation), + "--action-input-scale", str(float(action_input_scale)), + "--neutral-action-dropout", str(float(neutral_action_dropout)), + "--action-contrast-weight", str(float(action_contrast_weight)), + "--action-contrast-margin", str(float(action_contrast_margin)), + "--chunk-seed", str(int(seed)), + "--save-every", str(int(save_every)), + "--preview-every", str(int(preview_every) if preview_enabled else int(epochs) + 1), + "--preview-steps", str(int(preview_steps)), + "--mixed-precision", str(mixed_precision), + "--tf32", + ] + if resume: + command.extend(["--continue-action-model", str(resume)]) + if base: + command.extend(["--base-model", str(base)]) + if gradient_checkpointing: + command.append("--gradient-checkpointing") + if balance_actions: + command.append("--balance-actions") + process = subprocess.Popen( + command, cwd=str(root), stdout=subprocess.PIPE, stderr=subprocess.STDOUT, + text=True, encoding="utf-8", errors="replace", shell=False, + creationflags=subprocess.CREATE_NO_WINDOW if os.name == "nt" else 0, + ) + lines: queue.Queue[str | None] = queue.Queue() + + def read_output() -> None: + assert process.stdout is not None + for line in process.stdout: + lines.put(line.rstrip()) + lines.put(None) + + threading.Thread(target=read_output, daemon=True).start() + context.progress(1, "Starting Oasis trainer") + stopped = False + stop_requested_at: float | None = None + force_stop_sent = False + suspended = False + stop_file = output / "stop_action_training.flag" + while True: + should_pause = not context.run_event.is_set() + if should_pause != suspended and set_process_tree_paused(process, should_pause): + suspended = should_pause + context.log("Oasis trainer paused safely." if suspended else "Oasis trainer resumed.") + if context.cancel_event.is_set() and not stopped: + if suspended: + set_process_tree_paused(process, False) + suspended = False + output.mkdir(parents=True, exist_ok=True) + stop_file.touch(exist_ok=True) + stopped = True + stop_requested_at = time.monotonic() + context.log("Safe stop requested; waiting for Oasis to save a usable checkpoint.") + if stop_requested_at and not force_stop_sent and time.monotonic() - stop_requested_at > FORCE_STOP_TIMEOUT_SECONDS: + terminate_process_tree(process, timeout=3) + force_stop_sent = True + context.log("Oasis did not stop in time; terminating the trainer process.") + try: + line = lines.get(timeout=0.15) + if line and line.startswith("ACTION_FLOW_EVENT:"): + try: + _parse_event(context, json.loads(line.split(":", 1)[1]), int(epochs), output, bool(preview_enabled), int(preview_every), int(preview_steps)) + except json.JSONDecodeError: + context.log(line) + elif line: + context.log(line) + except queue.Empty: + pass + if process.poll() is not None and lines.empty(): + break + if stopped: + raise ToolCancelled("Oasis training stopped by user.") + if process.returncode != 0: + raise ToolExecutionError(f"Oasis trainer exited with code {process.returncode}. See the job log for details.") + if not _valid_action_model(output): + raise ToolExecutionError("Oasis training ended, but no valid action model metadata was written.") + context.progress(100, "Oasis training completed") + return { + "output_folder": str(output), + "model_name": model_name, + "assets": [{ + "kind": "model", + "name": model_name, + "path": str(output), + "trainer": "oasis", + "dataset_path": dataset_value, + "checkpoint": str(output), + "epochs": int(epochs), + }], + } + + +def launch_oasis_player( + context: ToolContext, + model_path: str, + model_name: str = "", + starting_frame: str = "", + seed: int = 0, +) -> dict[str, object]: + root = _oasis_root(context) + model = Path(model_path).expanduser().resolve() + if not _valid_action_model(model): + raise ToolExecutionError("Choose a valid Oasis action model folder.") + if starting_frame: + start = Path(starting_frame).expanduser() + if not start.is_file(): + raise ToolExecutionError("The selected Oasis starting frame does not exist.") + command = [sys.executable, str(root / "roblox_action_flow_app.py")] + env = os.environ.copy() + env["ADAM_OASIS_MODEL_PATH"] = str(model) + env["ADAM_OASIS_STARTING_FRAME"] = starting_frame + env["ADAM_OASIS_SEED"] = str(int(seed)) + process = subprocess.Popen( + command, cwd=str(root), env=env, stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL, + shell=False, creationflags=subprocess.CREATE_NO_WINDOW if os.name == "nt" else 0, + ) + context.log(f"Launched Oasis playable inference for {model_name or model.name}.") + context.progress(100, "Oasis player launched") + return {"player_pid": process.pid, "model_path": str(model)} diff --git a/adam/training_assistant.py b/adam/training_assistant.py index cf6b60656bc00e16c27277f92093331011c27038..cc67f81472cb7de7296bd7ead701d124dcdd8b1d 100644 --- a/adam/training_assistant.py +++ b/adam/training_assistant.py @@ -32,6 +32,11 @@ DEFAULT_PRESETS: dict[str, dict[str, Any]] = { "trainer": "flow", "epochs": 25, "image_count": 40, "description": "A short Flow Matching setup check.", }, + "Oasis Pipeline Test": { + "trainer": "oasis", "epochs": 5, "image_count": 500, + "description": "A short action-world-model run for validating gameplay frames and controls.", + "training_options": {"resolution": "256x144", "batch_size": 2, "workers": 2, "mixed_precision": "fp32"}, + }, } @@ -170,9 +175,10 @@ def estimate_plan(plan: Any) -> list[PreflightItem]: continue trainer = step.tool_id.removesuffix("_trainer") epochs = max(1, int(step.arguments.get("epochs", 1) or 1)) - dataset = Path(str(step.arguments.get("dataset_dir", ""))).expanduser() + raw_dataset = str(step.arguments.get("dataset_dir", "") or "").strip() + dataset = Path(raw_dataset).expanduser() image_count = 0 - if dataset.is_dir(): + if raw_dataset and dataset.is_dir(): try: image_count = sum( 1 @@ -189,6 +195,7 @@ def estimate_plan(plan: Any) -> list[PreflightItem]: "lora": 0.12, "ddpm": 0.07, "flow": 0.10, + "oasis": 0.18, }.get(trainer, 0.10) center_minutes = max(1, int(workload * seconds_per_image_epoch / 60)) low = max(1, center_minutes // 2) @@ -197,10 +204,11 @@ def estimate_plan(plan: Any) -> list[PreflightItem]: "lora": 0.25, "ddpm": 1.0, "flow": 1.0, + "oasis": 1.0, }.get(trainer, 0.75) checkpoint_count = max(1, min(20, epochs // 25 + 1)) disk_gb = checkpoint_gb * checkpoint_count - typical_vram = {"lora": 8, "ddpm": 6, "flow": 8}.get(trainer, 8) + typical_vram = {"lora": 8, "ddpm": 6, "flow": 8, "oasis": 12}.get(trainer, 8) estimates.extend( [ PreflightItem( @@ -248,7 +256,10 @@ def build_training_request( subject = subject.strip() dataset_name = dataset_name.strip() model_name = model_name.strip() or subject or dataset_name - trainer_label = {"lora": "LoRA", "ddpm": "DDPM", "flow": "Flow Matching"}[trainer] + trainer_label = {"lora": "LoRA", "ddpm": "DDPM", "flow": "Flow Matching", "oasis": "Oasis Action World Model"}.get( + trainer, + trainer.replace("_", " ").title(), + ) if create_dataset: collection_phrase = ( "as many available images as Bing returns (up to 5,000)" @@ -278,6 +289,7 @@ def build_training_request( ) if training_options: request += " [ADAM_TRAINING_OPTIONS:" + json.dumps(training_options, sort_keys=True) + "]" + request += " [ADAM_TRAINER:" + trainer + "]" return request @@ -293,7 +305,7 @@ def build_fine_tune_request( training_options: dict[str, Any] | None = None, ) -> str: """Build the explicit continuation request used by the Fine-Tune assistant.""" - labels = {"lora": "LoRA", "ddpm": "DDPM", "flow": "Flow Matching"} + labels = {"lora": "LoRA", "ddpm": "DDPM", "flow": "Flow Matching", "oasis": "Oasis Action World Model"} if trainer not in labels: raise ValueError("Fine-tuning requires a supported trainer.") if not model_name.strip(): @@ -328,8 +340,9 @@ def inspect_plan(plan: Any, config: Any) -> list[PreflightItem]: for step in plan.steps: if step.tool_id.endswith("_trainer") or step.tool_id == "dataset_collector": if step.tool_id not in checked_tools: - configured = Path(str(folders.get(step.tool_id, ""))).expanduser() - if configured.is_dir(): + raw_folder = str(folders.get(step.tool_id, "") or "").strip() + configured = Path(raw_folder).expanduser() + if raw_folder and configured.is_dir(): items.append(PreflightItem("ready", f"{step.title}: connected")) else: items.append(PreflightItem("warning", f"{step.title}: program folder is not connected")) @@ -395,12 +408,16 @@ def append_preflight_summary(plan: Any, config: Any) -> None: for item in items ] plan.summary += "\n\nPre-flight:\n" + "\n".join(f"• {line}" for line in lines) - if not getattr(plan, "orion_review", None) and "ORION —" not in plan.summary: + if not getattr(plan, "orion_review", None): apply_orion_review(plan) def completion_recommendation(plan: Any) -> str: tools = {step.tool_id for step in plan.steps} + if "oasis_trainer" in tools: + return ( + "Recommended next step: launch the Oasis player with the saved checkpoint and a starting frame from the same game, then record more control-balanced gameplay before long training." + ) if "lora_trainer" in tools or "ddpm_trainer" in tools or "flow_trainer" in tools: return ( "Recommended next step: generate a few preview images and compare them with " diff --git a/adam/transcript_dataset.py b/adam/transcript_dataset.py new file mode 100644 index 0000000000000000000000000000000000000000..a1fc11b5259c3ebe05fbce6cf1049506fdc6c9e1 --- /dev/null +++ b/adam/transcript_dataset.py @@ -0,0 +1,132 @@ +from __future__ import annotations + +import json +import re +import shutil +from dataclasses import dataclass +from importlib.util import find_spec +from pathlib import Path +from typing import Any + + +@dataclass(slots=True) +class TranscriptResult: + available: bool + message: str + samples: list[dict[str, Any]] + output_folder: str = "" + backend: str = "" + + +@dataclass(frozen=True, slots=True) +class TranscriptionBackend: + id: str + name: str + available: bool + requirement: str + + +def clean_transcript(text: str) -> str: + text = re.sub(r"\s+", " ", text.replace("\r", " ").replace("\n", " ")).strip() + text = re.sub(r"\s+([,.!?;:])", r"\1", text) + return text + + +def split_transcript(text: str, *, max_chars: int = 900) -> list[str]: + cleaned = clean_transcript(text) + if not cleaned: + return [] + sentences = re.split(r"(?<=[.!?])\s+", cleaned) + samples: list[str] = [] + current = "" + for sentence in sentences: + if not sentence: + continue + if current and len(current) + 1 + len(sentence) > max_chars: + samples.append(current.strip()) + current = sentence + else: + current = f"{current} {sentence}".strip() + if current: + samples.append(current.strip()) + return samples + + +def available_transcription_backends() -> list[TranscriptionBackend]: + return [ + TranscriptionBackend("whisper", "OpenAI Whisper", find_spec("whisper") is not None, "pip install openai-whisper"), + TranscriptionBackend("faster_whisper", "faster-whisper", find_spec("faster_whisper") is not None, "pip install faster-whisper"), + ] + + +def _select_backend(preferred: str = "auto") -> TranscriptionBackend | None: + backends = available_transcription_backends() + if preferred != "auto": + return next((backend for backend in backends if backend.id == preferred and backend.available), None) + return next((backend for backend in backends if backend.available), None) + + +def _transcribe_with_backend(video: Path, backend: TranscriptionBackend) -> str: + if backend.id == "whisper": + import whisper + + model = whisper.load_model("base") + result = model.transcribe(str(video)) + return str(result.get("text", "")) + if backend.id == "faster_whisper": + from faster_whisper import WhisperModel + + model = WhisperModel("base", device="auto", compute_type="auto") + segments, _info = model.transcribe(str(video)) + return " ".join(segment.text for segment in segments) + raise RuntimeError(f"Unsupported transcription backend: {backend.id}") + + +def transcript_videos_to_dataset( + videos: list[str | Path], + output_folder: str | Path, + *, + backend: str = "auto", + max_chars: int = 900, + preserve_metadata: bool = True, +) -> TranscriptResult: + output = Path(output_folder).expanduser().resolve() + if not videos: + return TranscriptResult(False, "Choose at least one local video.", []) + if not shutil.which("ffmpeg"): + return TranscriptResult( + False, + "FFmpeg was not found. ADAM can keep the workflow ready, but audio extraction/transcription needs FFmpeg.", + [], + ) + selected_backend = _select_backend(backend) + if selected_backend is None: + requirements = ", ".join(item.requirement for item in available_transcription_backends()) + return TranscriptResult( + False, + f"No local transcription backend is installed. Install one of: {requirements}.", + [], + ) + samples: list[dict[str, Any]] = [] + try: + for video in videos: + path = Path(video).expanduser().resolve() + if not path.is_file(): + return TranscriptResult(False, f"Video not found: {path}", []) + text = _transcribe_with_backend(path, selected_backend) + for index, sample in enumerate(split_transcript(text, max_chars=max_chars), 1): + item = {"text": sample} + if preserve_metadata: + item.update({"source_video": str(path), "sample_index": index}) + samples.append(item) + except Exception as exc: + return TranscriptResult(False, f"{selected_backend.name} could not transcribe the selected video(s): {exc}", [], backend=selected_backend.id) + output.mkdir(parents=True, exist_ok=True) + txt_path = output / "transcript_samples.txt" + jsonl_path = output / "transcript_samples.jsonl" + txt_path.write_text("\n\n".join(item["text"] for item in samples), encoding="utf-8") + jsonl_path.write_text( + "\n".join(json.dumps(item, ensure_ascii=False) for item in samples) + ("\n" if samples else ""), + encoding="utf-8", + ) + return TranscriptResult(True, f"Exported {len(samples):,} transcript sample(s) with {selected_backend.name}.", samples, str(output), selected_backend.id) diff --git a/adam/ui/generations.py b/adam/ui/generations.py index fa18e9aae06e72a85b43ab67b3e0d3457b34e138..863cdf1a2bcc9f182ae1751bf750e94175f7e971 100644 --- a/adam/ui/generations.py +++ b/adam/ui/generations.py @@ -3,8 +3,8 @@ from __future__ import annotations import json from pathlib import Path -from PySide6.QtCore import QSize, Qt, QTimer, QUrl -from PySide6.QtGui import QDesktopServices, QIcon, QImageReader, QPixmap +from PySide6.QtCore import QRectF, QSize, Qt, QTimer, QUrl +from PySide6.QtGui import QColor, QDesktopServices, QIcon, QImageReader, QPainter, QPainterPath, QPixmap from PySide6.QtWidgets import ( QAbstractItemView, QComboBox, @@ -15,6 +15,7 @@ from PySide6.QtWidgets import ( QFileDialog, QFrame, QGridLayout, + QGroupBox, QHBoxLayout, QLabel, QLineEdit, @@ -36,12 +37,16 @@ from adam.generations import ( GenerationRecord, build_generation_plan, generation_tools, + generation_model_key, + group_generation_records, load_generation_history, combine_generation_plans, ) +from adam.image_preferences import PreferenceProfile, score_generated_images from adam.job_manager import JobManager from adam.models import Job, JobStatus from adam.registry import ToolRegistry, ToolSpec +from adam.ui.settings_ui import SettingsForm def _card() -> QFrame: @@ -81,6 +86,38 @@ def _thumbnail(path: Path, width: int, height: int) -> QPixmap: return QPixmap.fromImage(image) if not image.isNull() else QPixmap() +def _folder_cover(path: Path | None, width: int = 176, height: int = 118) -> QPixmap: + """Draw a compact folder tile with a generated image set into its face.""" + canvas = QPixmap(width, height) + canvas.fill(Qt.transparent) + painter = QPainter(canvas) + painter.setRenderHint(QPainter.Antialiasing) + painter.setPen(Qt.NoPen) + painter.setBrush(QColor("#236a9d")) + painter.drawRoundedRect(QRectF(8, 7, 72, 28), 7, 7) + painter.setBrush(QColor("#174f78")) + painter.drawRoundedRect(QRectF(5, 23, width - 10, height - 28), 9, 9) + inset = QPainterPath() + inset.addRoundedRect(QRectF(12, 30, width - 24, height - 42), 6, 6) + painter.setClipPath(inset) + if path is not None: + cover = _thumbnail(path, width - 24, height - 42) + if not cover.isNull(): + scaled = cover.scaled( + width - 24, height - 42, Qt.KeepAspectRatioByExpanding, + Qt.SmoothTransformation, + ) + x = 12 + (width - 24 - scaled.width()) // 2 + y = 30 + (height - 42 - scaled.height()) // 2 + painter.drawPixmap(x, y, scaled) + painter.setClipping(False) + painter.setPen(QColor("#55b8ff")) + painter.setBrush(Qt.NoBrush) + painter.drawRoundedRect(QRectF(5.5, 23.5, width - 11, height - 29), 9, 9) + painter.end() + return canvas + + class GenerationCycleDialog(QDialog): """Choose several completed models and shared playback settings.""" @@ -240,6 +277,11 @@ class GenerationsPage(QWidget): self.records: list[GenerationRecord] = [] self.hidden_history_images: set[str] = set() self._cycle_jobs: dict[str, dict] = {} + self._history_load_token = 0 + self._history_load_index = 0 + self._history_entries: list[tuple[int, int, Path]] = [] + self._history_selected_image = "" + self._active_model_folder = "" layout = QVBoxLayout(self) layout.setContentsMargins(22, 18, 22, 18) @@ -273,10 +315,13 @@ class GenerationsPage(QWidget): self.refresh_button.clicked.connect(self.refresh) self.gallery.currentItemChanged.connect(self._selection_changed) self.gallery.itemDoubleClicked.connect(lambda _item: self._open_image()) + self.model_folders.itemDoubleClicked.connect(self._open_model_folder) + self.folder_back_button.clicked.connect(self._show_model_folders) self.open_image_button.clicked.connect(self._open_image) self.open_folder_button.clicked.connect(self._open_folder) self.reuse_button.clicked.connect(self._reuse_settings) self.clear_history_button.clicked.connect(self._clear_displayed_history) + self.auto_sort_button.clicked.connect(self._auto_sort_history) self.random_seed_button.clicked.connect(lambda: self.seed.clear()) self.jobs.job_updated.connect(self._job_updated) self.cycle_button.clicked.connect(self._open_generation_cycle) @@ -433,6 +478,13 @@ class GenerationsPage(QWidget): self.ddpm_reference_options.setVisible(False) layout.addWidget(self.ddpm_reference_options) + self.plugin_generation_group = QGroupBox("Model-specific settings") + plugin_generation_layout = QVBoxLayout(self.plugin_generation_group) + self.plugin_generation_form = SettingsForm() + plugin_generation_layout.addWidget(self.plugin_generation_form) + self.plugin_generation_group.setVisible(False) + layout.addWidget(self.plugin_generation_group) + grid = QGridLayout() grid.setHorizontalSpacing(8) grid.setVerticalSpacing(6) @@ -488,6 +540,38 @@ class GenerationsPage(QWidget): seed_row.addWidget(self.random_seed_button) layout.addLayout(seed_row) + self.smart_group = QGroupBox("Smart Generation") + smart_layout = QGridLayout(self.smart_group) + self.smart_enabled = QCheckBox("Smart Generation") + self.smart_wanted = QSpinBox() + self.smart_wanted.setRange(1, 48) + self.smart_wanted.setValue(8) + self.smart_max_candidates = QSpinBox() + self.smart_max_candidates.setRange(1, 256) + self.smart_max_candidates.setValue(32) + self.smart_min_score = QDoubleSpinBox() + self.smart_min_score.setRange(0.0, 1.0) + self.smart_min_score.setSingleStep(0.05) + self.smart_min_score.setDecimals(2) + self.smart_min_score.setValue(0.70) + self.smart_mode = QComboBox() + self.smart_mode.addItem("Stop at threshold", "threshold") + self.smart_mode.addItem("Rank fixed pool", "top_n") + self.smart_keep_rejected = QCheckBox("Keep rejected candidates for review") + self.smart_keep_rejected.setChecked(True) + smart_layout.addWidget(self.smart_enabled, 0, 0, 1, 2) + for row, (label, widget) in enumerate(( + ("Wanted results", self.smart_wanted), + ("Maximum candidates", self.smart_max_candidates), + ("Minimum score", self.smart_min_score), + ("Mode", self.smart_mode), + ), 1): + smart_layout.addWidget(QLabel(label), row, 0) + smart_layout.addWidget(widget, row, 1) + smart_layout.addWidget(self.smart_keep_rejected, 5, 0, 1, 2) + self.smart_group.setVisible(False) + layout.addWidget(self.smart_group) + self.generate_button = QPushButton("Generate images →") self.generate_button.setProperty("primary", True) layout.addWidget(self.generate_button) @@ -504,7 +588,11 @@ class GenerationsPage(QWidget): layout.setContentsMargins(16, 15, 16, 16) layout.setSpacing(9) top = QHBoxLayout() - top.addWidget(_title("GENERATION HISTORY")) + self.history_title = _title("MODEL FOLDERS") + top.addWidget(self.history_title) + self.folder_back_button = QPushButton("← All models") + self.folder_back_button.setVisible(False) + top.addWidget(self.folder_back_button) top.addStretch() self.history_summary = QLabel("No generations yet") self.history_summary.setProperty("muted", True) @@ -515,8 +603,28 @@ class GenerationsPage(QWidget): ) self.clear_history_button.setEnabled(False) top.addWidget(self.clear_history_button) + self.auto_sort_button = QPushButton("Auto Sort") + self.auto_sort_button.setToolTip( + "Score visible generations against their model preference profiles and show stronger matches first." + ) + top.addWidget(self.auto_sort_button) layout.addLayout(top) + self.generation_target = QLabel( + "Double-click a model folder to browse its images and generate with that model." + ) + self.generation_target.setProperty("muted", True) + self.generation_target.setWordWrap(True) + layout.addWidget(self.generation_target) + + self.model_folders = QListWidget() + self.model_folders.setViewMode(QListWidget.IconMode) + self.model_folders.setIconSize(QSize(176, 118)) + self.model_folders.setGridSize(QSize(220, 184)) + self.model_folders.setResizeMode(QListWidget.Adjust) + self.model_folders.setSelectionMode(QAbstractItemView.SingleSelection) + layout.addWidget(self.model_folders, 1) + self.gallery = QListWidget() self.gallery.setViewMode(QListWidget.IconMode) self.gallery.setIconSize(QSize(170, 128)) @@ -524,11 +632,32 @@ class GenerationsPage(QWidget): self.gallery.setResizeMode(QListWidget.Adjust) self.gallery.setSelectionMode(QAbstractItemView.SingleSelection) layout.addWidget(self.gallery, 1) + self.gallery.setVisible(False) + self.image_detail_panel = QWidget() + image_detail_layout = QVBoxLayout(self.image_detail_panel) + image_detail_layout.setContentsMargins(0, 0, 0, 0) + image_detail_layout.setSpacing(7) self.detail = QLabel("Select an image to see its reproducibility settings.") self.detail.setProperty("muted", True) self.detail.setWordWrap(True) - layout.addWidget(self.detail) + image_detail_layout.addWidget(self.detail) + rating_row = QHBoxLayout() + self.favorite_button = QPushButton("Favorite") + self.keep_button = QPushButton("Keep") + self.unsure_button = QPushButton("Unsure") + self.reject_button = QPushButton("Reject") + for button, rating in ( + (self.favorite_button, "favorite"), + (self.keep_button, "keep"), + (self.unsure_button, "unsure"), + (self.reject_button, "reject"), + ): + button.setEnabled(False) + button.clicked.connect(lambda _checked=False, value=rating: self._rate_selection(value)) + rating_row.addWidget(button) + rating_row.addStretch() + image_detail_layout.addLayout(rating_row) actions = QHBoxLayout() self.open_image_button = QPushButton("Open image") self.open_folder_button = QPushButton("Open folder") @@ -541,7 +670,9 @@ class GenerationsPage(QWidget): button.setEnabled(False) actions.addWidget(button) actions.addStretch() - layout.addLayout(actions) + image_detail_layout.addLayout(actions) + layout.addWidget(self.image_detail_panel) + self.image_detail_panel.setVisible(False) return card def refresh(self) -> None: @@ -617,12 +748,16 @@ class GenerationsPage(QWidget): ) self.lora_options.setVisible(tool.id == "lora_generator") self.ddpm_reference_options.setVisible(tool.id == "ddpm_generator") + self.smart_group.setVisible("smart_generation" in tool.capabilities) self.aspect.setEnabled( tool.id != "ddpm_generator" or not self.ddpm_custom_size.isChecked() ) if tool.id == "lora_generator": self._load_lora_options() + self._load_plugin_generation_settings(tool) self.model.blockSignals(False) + if not tool: + self.smart_group.setVisible(False) self.generate_button.setEnabled(bool(tool and self.model.count())) self._apply_preset() self._model_changed() @@ -728,6 +863,29 @@ class GenerationsPage(QWidget): if self._current_tool() and self._current_tool().id == "ddpm_generator": self.aspect.setEnabled(not enabled) + def _load_plugin_generation_settings(self, tool: ToolSpec) -> None: + schema = self.registry.model_plugins.generation_schema_for_tool(tool.id) + built_in = { + "model_name", "model_path", "prompt", "image_count", "steps", "seed", + "sampler", "aspect_ratio", "preview_interval", + "smart_generation", "smart_wanted_results", "smart_max_candidates", + "smart_min_score", "smart_mode", "smart_keep_rejected", + } + if tool.id == "lora_generator": + built_in.update({ + "base_model_path", "negative_prompt", "width", "height", + "cfg_scale", "lora_strength", "reference_image", + "denoise_strength", "prompt_weighting", + }) + elif tool.id == "ddpm_generator": + built_in.update({"reference_image", "reference_strength", "width", "height"}) + extra_schema = { + key: spec for key, spec in schema.items() + if key not in built_in and key in tool.arguments + } + self.plugin_generation_form.set_schema(extra_schema) + self.plugin_generation_group.setVisible(bool(extra_schema)) + @staticmethod def _model_is_ready(asset: Asset) -> bool: path = Path(asset.path) @@ -837,6 +995,19 @@ class GenerationsPage(QWidget): } elif tool.id == "flow_generator": extra_arguments = {"preview_interval": self.preview_interval.value()} + if self.smart_group.isVisible() and self.smart_enabled.isChecked(): + if self.smart_max_candidates.value() < self.smart_wanted.value(): + QMessageBox.warning(self, "Check Smart Generation", "Maximum candidates must be at least the wanted result count.") + return + extra_arguments.update({ + "smart_generation": True, + "smart_wanted_results": self.smart_wanted.value(), + "smart_max_candidates": self.smart_max_candidates.value(), + "smart_min_score": self.smart_min_score.value(), + "smart_mode": str(self.smart_mode.currentData() or "threshold"), + "smart_keep_rejected": self.smart_keep_rejected.isChecked(), + }) + extra_arguments.update(self.plugin_generation_form.values()) plan = build_generation_plan( tool, model_name=self.model.currentText(), @@ -980,6 +1151,12 @@ class GenerationsPage(QWidget): "ddpm_custom_size": self.ddpm_custom_size.isChecked(), "ddpm_width": self.ddpm_width.value(), "ddpm_height": self.ddpm_height.value(), + "smart_generation": self.smart_enabled.isChecked(), + "smart_wanted_results": self.smart_wanted.value(), + "smart_max_candidates": self.smart_max_candidates.value(), + "smart_min_score": self.smart_min_score.value(), + "smart_mode": self.smart_mode.currentData(), + "smart_keep_rejected": self.smart_keep_rejected.isChecked(), } }) @@ -1006,6 +1183,12 @@ class GenerationsPage(QWidget): self.ddpm_custom_size.toggled.connect(self._save_generation_settings) self.ddpm_width.valueChanged.connect(self._save_generation_settings) self.ddpm_height.valueChanged.connect(self._save_generation_settings) + self.smart_enabled.toggled.connect(self._save_generation_settings) + self.smart_wanted.valueChanged.connect(self._save_generation_settings) + self.smart_max_candidates.valueChanged.connect(self._save_generation_settings) + self.smart_min_score.valueChanged.connect(self._save_generation_settings) + self.smart_mode.currentIndexChanged.connect(self._save_generation_settings) + self.smart_keep_rejected.toggled.connect(self._save_generation_settings) def _restore_generation_settings(self) -> None: saved = self.config.get("generation_settings", {}) @@ -1044,6 +1227,20 @@ class GenerationsPage(QWidget): self.prompt_weighting.setChecked(bool(saved.get("prompt_weighting", self.prompt_weighting.isChecked()))) self.ddpm_reference_image.setText(str(saved.get("ddpm_reference_image", ""))) self.ddpm_custom_size.setChecked(bool(saved.get("ddpm_custom_size", False))) + self.smart_enabled.setChecked(bool(saved.get("smart_generation", False))) + for spin, key in ((self.smart_wanted, "smart_wanted_results"), (self.smart_max_candidates, "smart_max_candidates")): + try: + spin.setValue(max(spin.minimum(), min(int(saved.get(key, spin.value())), spin.maximum()))) + except (TypeError, ValueError): + pass + try: + self.smart_min_score.setValue(float(saved.get("smart_min_score", self.smart_min_score.value()))) + except (TypeError, ValueError): + pass + mode_index = self.smart_mode.findData(saved.get("smart_mode")) + if mode_index >= 0: + self.smart_mode.setCurrentIndex(mode_index) + self.smart_keep_rejected.setChecked(bool(saved.get("smart_keep_rejected", True))) def _load_history(self) -> None: selected_image = None @@ -1051,40 +1248,144 @@ class GenerationsPage(QWidget): if current: selected_image = current.data(Qt.UserRole) self.records = load_generation_history(self.root) + folders = group_generation_records(self.records) + if self._active_model_folder and not any( + folder.key == self._active_model_folder for folder in folders + ): + self._active_model_folder = "" + self._history_load_token += 1 + token = self._history_load_token + self._history_load_index = 0 + self._history_selected_image = ( + str(selected_image.get("image", "")) if isinstance(selected_image, dict) else "" + ) self.gallery.clear() - image_total = 0 - selected_row = -1 + self.model_folders.clear() + if not self._active_model_folder: + for folder in folders: + item = QListWidgetItem( + QIcon(_folder_cover(folder.cover_image)), + f"{folder.model_name}\n{folder.provider_name} · " + f"{folder.image_count} image{'s' if folder.image_count != 1 else ''}", + ) + item.setData(Qt.UserRole, folder.key) + item.setToolTip( + f"Open {folder.model_name}\n" + f"{len(folder.records)} batch{'es' if len(folder.records) != 1 else ''} · " + f"{folder.image_count} image{'s' if folder.image_count != 1 else ''}" + ) + self.model_folders.addItem(item) + self.model_folders.setVisible(True) + self.gallery.setVisible(False) + self.image_detail_panel.setVisible(False) + self.history_title.setText("MODEL FOLDERS") + self.folder_back_button.setVisible(False) + self.generation_target.setText( + "Double-click a model folder to browse its images and generate with that model." + ) + total_images = sum(folder.image_count for folder in folders) + self.history_summary.setText( + f"{len(folders)} models · {len(self.records)} batches · {total_images} images" + if folders else "No generations yet" + ) + self.clear_history_button.setEnabled(False) + self.auto_sort_button.setEnabled(False) + self._selection_changed(None, None) + return + active_folder = next(folder for folder in folders if folder.key == self._active_model_folder) + self.model_folders.setVisible(False) + self.gallery.setVisible(True) + self.image_detail_panel.setVisible(True) + self.history_title.setText(active_folder.model_name.upper()) + self.folder_back_button.setVisible(True) + self.generation_target.setText( + f"Generating now will add the new batch to {active_folder.model_name}." + ) + self.auto_sort_button.setEnabled(True) + self._history_entries = [] + active_batch_count = 0 for record_index, record in enumerate(self.records): + if generation_model_key(record) != self._active_model_folder: + continue + active_batch_count += 1 for image_index, path in enumerate(record.images): if str(path) in self.hidden_history_images: continue - image_total += 1 - seed = record.seed + image_index - item = QListWidgetItem( - QIcon(_thumbnail(path, 170, 128)), - f"{record.model_name}\nSeed {seed}", - ) - payload = { - "record": record_index, - "image": str(path), - "seed": seed, - } - item.setData(Qt.UserRole, payload) - self.gallery.addItem(item) - if selected_image and selected_image.get("image") == str(path): - selected_row = self.gallery.count() - 1 + self._history_entries.append((record_index, image_index, path)) + image_total = len(self._history_entries) self.history_summary.setText( - f"{len(self.records)} batches · {image_total} images" + f"{active_batch_count} batches · {image_total} images · Loading thumbnails…" if image_total else "No generations yet" ) self.clear_history_button.setEnabled(image_total > 0) - if selected_row >= 0: - self.gallery.setCurrentRow(selected_row) - elif self.gallery.count(): - self.gallery.setCurrentRow(0) - else: + if not image_total: self._selection_changed(None, None) + return + QTimer.singleShot(0, lambda: self._load_next_history_thumbnail(token)) + + def _load_next_history_thumbnail(self, token: int) -> None: + """Decode one history image per event-loop turn so the page opens immediately.""" + if token != self._history_load_token: + return + if self._history_load_index >= len(self._history_entries): + batch_count = len({entry[0] for entry in self._history_entries}) + self.history_summary.setText( + f"{batch_count} batches · {len(self._history_entries)} images" + ) + if self.gallery.currentRow() < 0 and self.gallery.count(): + self.gallery.setCurrentRow(0) + return + record_index, image_index, path = self._history_entries[self._history_load_index] + record = self.records[record_index] + seed = record.seed + image_index + rating = self._stored_rating(record, path) + score = record.image_evaluations.get(str(path.resolve()), {}).get("score") + badges = [] + if rating: + badges.append(rating.title()) + if isinstance(score, (int, float)): + badges.append(f"{float(score):.2f}") + item = QListWidgetItem( + QIcon(_thumbnail(path, 170, 128)), + f"{record.model_name}\nSeed {seed}" + (f" · {' · '.join(badges)}" if badges else ""), + ) + item.setData(Qt.UserRole, { + "record": record_index, + "image": str(path), + "seed": seed, + }) + self.gallery.addItem(item) + if self._history_selected_image == str(path): + self.gallery.setCurrentItem(item) + self._history_load_index += 1 + QTimer.singleShot(0, lambda: self._load_next_history_thumbnail(token)) + + def _open_model_folder(self, item: QListWidgetItem) -> None: + key = str(item.data(Qt.UserRole) or "") + folder = next( + (candidate for candidate in group_generation_records(self.records) if candidate.key == key), + None, + ) + if folder is None: + return + self._active_model_folder = folder.key + provider_index = self.provider.findData(folder.provider_id) + if provider_index >= 0: + self.provider.setCurrentIndex(provider_index) + model_index = self.model.findData(folder.model_path) + if model_index >= 0: + self.model.setCurrentIndex(model_index) + self.status.setText(f"Opened {folder.model_name}. New images will be generated here.") + else: + self.status.setText( + f"Opened {folder.model_name}. Its history is available, but the model is not currently connected." + ) + self._load_history() + + def _show_model_folders(self) -> None: + self._active_model_folder = "" + self._load_history() def _clear_displayed_history(self) -> None: for index in range(self.gallery.count()): @@ -1097,6 +1398,64 @@ class GenerationsPage(QWidget): self._selection_changed(None, None) self.status.setText("Displayed generation history cleared. Image files were not deleted.") + def _auto_sort_history(self) -> None: + changed = 0 + skipped = 0 + for record in self.records: + if self._active_model_folder and generation_model_key(record) != self._active_model_folder: + continue + profile = PreferenceProfile(self.root, record.provider_id, record.model_name, record.model_path) + if not profile.has_signal(): + skipped += 1 + continue + try: + scores = score_generated_images( + self.root, + provider_id=record.provider_id, + model_name=record.model_name, + model_path=record.model_path, + image_paths=record.images, + keep_threshold=profile.keep_threshold, + reject_threshold=profile.reject_threshold, + ) + evaluations = { + score.image_path: { + "score": score.score, + "confidence": score.confidence, + "category": score.category, + "reason": score.reason, + } + for score in scores + } + ordered = sorted( + record.images, + key=lambda path: float(evaluations.get(str(path.resolve()), {}).get("score") or -1.0), + reverse=True, + ) + payload = json.loads(record.metadata_path.read_text(encoding="utf-8")) + payload["images"] = [str(path) for path in ordered] + payload["image_evaluations"] = evaluations + smart = dict(payload.get("smart_generation") or {}) + smart.update({ + "auto_sorted": True, + "profile_id": profile.id, + "selected_count": sum( + 1 for score in scores + if score.score is not None and score.score >= profile.keep_threshold + ), + "candidate_count": len(scores), + }) + payload["smart_generation"] = smart + temporary = record.metadata_path.with_suffix(".tmp") + temporary.write_text(json.dumps(payload, indent=2), encoding="utf-8") + temporary.replace(record.metadata_path) + changed += 1 + except Exception as exc: + skipped += 1 + self.status.setText(f"Auto Sort skipped a batch: {exc}") + self._load_history() + self.status.setText(f"Auto Sort updated {changed} batch(es). {skipped} had no profile signal or could not be scored.") + def _selection(self) -> tuple[GenerationRecord, Path, int] | None: item = self.gallery.currentItem() if not item: @@ -1117,16 +1476,64 @@ class GenerationsPage(QWidget): self.open_image_button.setEnabled(enabled) self.open_folder_button.setEnabled(enabled) self.reuse_button.setEnabled(enabled) + for button in (self.favorite_button, self.keep_button, self.unsure_button, self.reject_button): + button.setEnabled(enabled) if not selection: self.detail.setText("Select an image to see its reproducibility settings.") return - record, _path, seed = selection + record, path, seed = selection note = f" · Note: {record.prompt}" if record.prompt else "" + rating = self._stored_rating(record, path) + evaluation = record.image_evaluations.get(str(path.resolve()), {}) + score = evaluation.get("score") + score_text = f" · Smart score {float(score):.2f}" if isinstance(score, (int, float)) else "" + category = str(evaluation.get("category", "")) + category_text = f" · {category}" if category else "" + rating_text = f" · Rating: {rating.title()}" if rating else " · Not rated" self.detail.setText( f"{record.provider_name} · {record.model_name} · Seed {seed} · " - f"{record.steps} steps · {record.sampler} · {record.aspect_ratio}{note}" + f"{record.steps} steps · {record.sampler} · {record.aspect_ratio}{rating_text}{score_text}{category_text}{note}" ) + def _stored_rating(self, record: GenerationRecord, path: Path) -> str: + profile = PreferenceProfile(self.root, record.provider_id, record.model_name, record.model_path) + rating = profile.rating_for(path) + return rating.rating if rating else "" + + def _generation_settings_for_record(self, record: GenerationRecord, seed: int) -> dict: + settings = { + "provider_id": record.provider_id, + "provider_name": record.provider_name, + "model_name": record.model_name, + "model_path": record.model_path, + "prompt": record.prompt, + "seed": seed, + "steps": record.steps, + "sampler": record.sampler, + "aspect_ratio": record.aspect_ratio, + } + settings.update(record.smart_generation) + return settings + + def _rate_selection(self, rating: str) -> None: + selection = self._selection() + if not selection: + return + record, path, seed = selection + profile = PreferenceProfile(self.root, record.provider_id, record.model_name, record.model_path) + profile.set_rating( + path, + rating, + seed=seed, + sampler=record.sampler, + steps=record.steps, + resolution=record.aspect_ratio, + generation_settings=self._generation_settings_for_record(record, seed), + generation_created_at=record.created_at, + ) + self.status.setText(f"Saved {rating.title()} for {record.model_name}.") + self._load_history() + def _open_image(self) -> None: selection = self._selection() if selection: diff --git a/adam/ui/main_window.py b/adam/ui/main_window.py index 9a953d58c9eb10084c0b483d86a5920c49fcb67a..7025678cff3ee49f2689d7fa4c14eb858286d530 100644 --- a/adam/ui/main_window.py +++ b/adam/ui/main_window.py @@ -1,5 +1,6 @@ from __future__ import annotations +from io import BytesIO from datetime import datetime from dataclasses import replace import json @@ -8,7 +9,7 @@ import re import shutil import sys -from PySide6.QtCore import QThread, QTimer, Qt, QUrl, Signal +from PySide6.QtCore import QDateTime, QThread, QTimer, Qt, QUrl, Signal from PySide6.QtGui import QCloseEvent, QDesktopServices, QIcon, QPixmap from PySide6.QtWidgets import ( QAbstractItemView, @@ -17,6 +18,7 @@ from PySide6.QtWidgets import ( QComboBox, QDialog, QDialogButtonBox, + QDateTimeEdit, QDoubleSpinBox, QFileDialog, QFrame, @@ -39,6 +41,7 @@ from PySide6.QtWidgets import ( QStackedWidget, QSystemTrayIcon, QTabBar, + QTabWidget, QTableWidget, QTableWidgetItem, QVBoxLayout, @@ -46,6 +49,8 @@ from PySide6.QtWidgets import ( ) from adam.config import ConfigManager +from adam.dataset_lab import scan_dataset +from adam.experiment_tracker import ExperimentStore, ExperimentRun from adam.generations import ( ChatGenerationRequest, build_generation_plan, @@ -53,6 +58,8 @@ from adam.generations import ( generation_tools, parse_chat_generation_request, ) +from adam.model_inspector import ModelComparison, ModelInspection, compare_models, inspect_model +from adam.model_inspector.statistics import bytes_label, shape_label from adam.external_tools import ( ExternalToolStore, ToolAnalysis, @@ -61,12 +68,24 @@ from adam.external_tools import ( ) from adam.job_manager import JobManager from adam.models import Job, JobStatus, SystemSnapshot +from adam.model_plugins import ModelPluginError, safe_plugin_id, scaffold_model_plugin +from adam.model_profiles import ModelProfileRegistry from adam.monitoring import SystemMonitor from adam.orion import dataset_image_count, recommend_training_settings from adam.ollama import OllamaClient from adam.planner import Planner, PlanningError +from adam.recommendations import recommend_for_profile +from adam.remote_access import ( + REMOTE_MODE_DISABLED, + REMOTE_MODE_LOCAL, + REMOTE_MODE_TAILSCALE, + RemoteAccessService, + remote_scope, +) from adam.registry import ToolRegistry +from adam.studio import StudioStore from adam.tool_folders import ToolFolderManager, ToolFolderStatus +from adam.transcript_dataset import available_transcription_backends, transcript_videos_to_dataset from adam.training_assistant import ( append_preflight_summary, build_fine_tune_request, @@ -82,6 +101,7 @@ from adam.ui.theme import APP_STYLESHEET, COLORS from adam.ui.studio import StudioPage from adam.ui.generations import GenerationsPage from adam.ui.showcase import ShowcasePage +from adam.ui.settings_ui import SettingsForm from adam.ui.widgets import ( ActiveJobPanel, ChatBubble, @@ -340,6 +360,33 @@ class ToolScanWorker(QThread): self.failed.emit(str(exc)) +class DatasetScanWorker(QThread): + scanned = Signal(object) + failed = Signal(str) + + def __init__(self, folder: str) -> None: + super().__init__() + self.folder = folder + + def run(self) -> None: + try: + self.scanned.emit(scan_dataset(self.folder)) + except Exception as exc: + self.failed.emit(str(exc)) + + +class TranscriptExportWorker(QThread): + completed = Signal(object) + + def __init__(self, videos: list[str], output: str) -> None: + super().__init__() + self.videos = videos + self.output = output + + def run(self) -> None: + self.completed.emit(transcript_videos_to_dataset(self.videos, self.output)) + + class ChatWorker(QThread): chunk = Signal(str) answered = Signal(str) @@ -366,9 +413,100 @@ class ChatWorker(QThread): self.failed.emit(str(exc)) +class ModelPluginWizardDialog(QDialog): + """Scaffold a copyable model plugin folder from a few fields.""" + + def __init__(self, root: Path, parent: QWidget | None = None) -> None: + super().__init__(parent) + self.root = root.resolve() + self.created_folder = "" + self.setWindowTitle("Create Model Plugin") + self.setMinimumWidth(520) + layout = QVBoxLayout(self) + layout.setSpacing(10) + layout.addWidget( + _page_header( + "Create a model plugin", + "Scaffold the manifest and Python files for a new architecture.", + ) + ) + form = QGridLayout() + form.setHorizontalSpacing(12) + form.setVerticalSpacing(9) + self.name = QLineEdit() + self.name.setPlaceholderText("Example: Neural Cellular Automata") + self.plugin_id = QLineEdit() + self.plugin_id.setPlaceholderText("neural_cellular_automata") + self.architecture = QLineEdit() + self.architecture.setPlaceholderText("nca, maskgit, vae, autoregressive") + self.output_type = QComboBox() + for label in ("image", "video", "audio", "text", "other"): + self.output_type.addItem(label.title(), label) + self.training = QCheckBox("Training") + self.training.setChecked(True) + self.generation = QCheckBox("Generation") + self.generation.setChecked(True) + capability_row = QHBoxLayout() + capability_row.addWidget(self.training) + capability_row.addWidget(self.generation) + capability_row.addStretch() + form.addWidget(QLabel("Model name"), 0, 0) + form.addWidget(self.name, 0, 1) + form.addWidget(QLabel("Plugin folder"), 1, 0) + form.addWidget(self.plugin_id, 1, 1) + form.addWidget(QLabel("Architecture"), 2, 0) + form.addWidget(self.architecture, 2, 1) + form.addWidget(QLabel("Output type"), 3, 0) + form.addWidget(self.output_type, 3, 1) + form.addWidget(QLabel("Capabilities"), 4, 0) + form.addLayout(capability_row, 4, 1) + layout.addLayout(form) + self.status = QLabel() + self.status.setProperty("muted", True) + self.status.setWordWrap(True) + layout.addWidget(self.status) + buttons = QDialogButtonBox(QDialogButtonBox.Cancel | QDialogButtonBox.Ok) + buttons.button(QDialogButtonBox.Ok).setText("Create plugin folder") + buttons.accepted.connect(self._create) + buttons.rejected.connect(self.reject) + layout.addWidget(buttons) + self.name.textChanged.connect(self._suggest_id) + + def _suggest_id(self, value: str) -> None: + if not self.plugin_id.text().strip(): + self.plugin_id.setPlaceholderText(safe_plugin_id(value)) + + def _create(self) -> None: + name = self.name.text().strip() + if not name: + self.status.setText("Enter a model name.") + return + plugin_id = self.plugin_id.text().strip() or self.plugin_id.placeholderText() + if not self.training.isChecked() and not self.generation.isChecked(): + self.status.setText("Choose training, generation, or both.") + return + try: + folder = scaffold_model_plugin( + self.root, + plugin_id=plugin_id, + name=name, + architecture=self.architecture.text().strip() or "custom", + output_type=str(self.output_type.currentData() or "image"), + include_training=self.training.isChecked(), + include_generation=self.generation.isChecked(), + ) + except ModelPluginError as exc: + self.status.setText(str(exc)) + return + self.created_folder = str(folder) + self.accept() + + class ModelCreationDialog(QDialog): """Collects training choices in plain language and produces a planner request.""" + BUILTIN_PREVIEW_TRAINERS = {"ddpm", "flow", "lora", "oasis"} + def __init__(self, planner: Planner, config: ConfigManager, parent: QWidget | None = None) -> None: super().__init__(parent) self.planner = planner @@ -376,17 +514,21 @@ class ModelCreationDialog(QDialog): self.request = "" self.requests: list[str] = [] self.collection_only = False + self.scheduled_for: str | None = None self._model_states: list[dict[str, object]] = [] self._current_model_index = 0 self.setWindowTitle("Model Creation Assistant") self.setMinimumWidth(560) + available = self.screen().availableGeometry() + self.resize(min(940, available.width() - 40), min(850, available.height() - 60)) outer = QVBoxLayout(self) outer.setContentsMargins(0, 0, 0, 0) scroll = QScrollArea() scroll.setWidgetResizable(True) - scroll.setHorizontalScrollBarPolicy(Qt.ScrollBarAlwaysOff) + scroll.setHorizontalScrollBarPolicy(Qt.ScrollBarAsNeeded) content = QWidget() root = QVBoxLayout(content) + root.setSizeConstraint(QVBoxLayout.SetMinimumSize) scroll.setWidget(content) outer.addWidget(scroll) root.setSpacing(10) @@ -396,9 +538,11 @@ class ModelCreationDialog(QDialog): "Choose what you know. ADAM will turn it into a complete, reviewable training request.", ) ) + planner.assets.discover(config) journey = QLabel( "1 GOAL → 2 DATASET → 3 TRAINING RECIPE → 4 REVIEW & APPROVE" ) + journey.setWordWrap(True) journey.setStyleSheet( f"color: {COLORS['blue_2']}; background: #081a27; " f"border: 1px solid {COLORS['border_bright']}; border-radius: 8px; " @@ -420,20 +564,21 @@ class ModelCreationDialog(QDialog): model_tabs_row.addWidget(self.add_model_button) root.addLayout(model_tabs_row) - batch_tools = QHBoxLayout() + batch_tools = QGridLayout() self.bulk_add_button = QPushButton("Paste model list…") + self.create_plugin_button = QPushButton("Create plugin…") self.apply_many_button = QPushButton("Apply current settings…") self.save_draft_button = QPushButton("Save draft") self.load_draft_button = QPushButton("Load draft") self.match_existing_button = QPushButton("Match existing datasets") self.refresh_datasets_button = QPushButton("Find collected datasets") - for button in ( + for index, button in enumerate(( self.bulk_add_button, self.apply_many_button, self.save_draft_button, self.load_draft_button, self.match_existing_button, self.refresh_datasets_button, - ): + self.create_plugin_button, + )): button.setProperty("chip", True) - batch_tools.addWidget(button) - batch_tools.addStretch() + batch_tools.addWidget(button, index // 3, index % 3) root.addLayout(batch_tools) form = QGridLayout() @@ -441,14 +586,27 @@ class ModelCreationDialog(QDialog): form.setVerticalSpacing(9) self.preset = QComboBox() self.presets = presets_from_config(config) - self.preset.addItems(self.presets) self.trainer = QComboBox() - self.trainer.addItem("LoRA", "lora") - self.trainer.addItem("DDPM", "ddpm") - self.trainer.addItem("Flow Matching", "flow") + preferred_trainers = ["lora", "ddpm", "flow", "oasis"] + plugins = { + plugin.id: plugin + for plugin in planner.registry.model_plugins.all() + if plugin.training_settings and plugin.info.get("category") != "Template" + } + for trainer_id in preferred_trainers: + plugin = plugins.pop(trainer_id, None) + if plugin: + self.trainer.addItem(plugin.name, plugin.id) + for plugin in sorted(plugins.values(), key=lambda item: item.name.casefold()): + self.trainer.addItem(plugin.name, plugin.id) + self.training_mode = QComboBox() + self.training_mode.addItem("Train a new model", "new") + self.training_mode.addItem("Continue one of my models", "continue") + self.continue_model = QComboBox() self.source = QComboBox() self.source.addItem("Create a new dataset", "new") self.source.addItem("Use an existing dataset", "existing") + self.source.addItem("Use the continued model's original dataset", "original") self.subject = QLineEdit() self.subject.setPlaceholderText("Example: Hatsune Miku") self.dataset = QComboBox() @@ -459,6 +617,8 @@ class ModelCreationDialog(QDialog): self.dataset.addItem(asset.name) self.model_name = QLineEdit() self.model_name.setPlaceholderText("Defaults to the subject or dataset name") + self.trigger_word = QLineEdit() + self.trigger_word.setPlaceholderText("Defaults to the model name") self.epochs = QSpinBox() self.epochs.setRange(1, 100_000) self.images = QSpinBox() @@ -474,39 +634,37 @@ class ModelCreationDialog(QDialog): rows = [ ("Preset", self.preset), ("Trainer", self.trainer), + ("Starting point", self.training_mode), + ("Model to continue", self.continue_model), ("Dataset choice", self.source), ("What should it learn?", self.subject), ("Existing dataset", self.dataset), ("Model name", self.model_name), + ("LoRA trigger word", self.trigger_word), ("Training length", self.epochs), ("Internet image collection", self.collection_mode), ("New dataset size", self.images), ] + self.form_labels: dict[str, QLabel] = {} for row, (label, widget) in enumerate(rows): - form.addWidget(QLabel(label), row, 0) + label_widget = QLabel(label) + self.form_labels[label] = label_widget + form.addWidget(label_widget, row, 0) form.addWidget(widget, row, 1) root.addLayout(form) - self.options_group = QGroupBox("Training options") - options = QGridLayout(self.options_group) - self.resolution = QComboBox(); self.resolution.addItems(["64", "128", "256", "384", "512"]) - self.batch_size = QSpinBox(); self.batch_size.setRange(1, 64) - self.learning_rate = QDoubleSpinBox(); self.learning_rate.setDecimals(7); self.learning_rate.setRange(0.0000001, 0.1); self.learning_rate.setSingleStep(0.00005) - self.gradient_accumulation = QSpinBox(); self.gradient_accumulation.setRange(1, 64) - self.workers = QSpinBox(); self.workers.setRange(0, 16) - self.precision = QComboBox(); self.precision.addItem("FP16 (faster / less VRAM)", "fp16"); self.precision.addItem("Full precision (more stable / slower)", "no") - self.save_every = QSpinBox(); self.save_every.setRange(1, 1000) - self.preview_steps = QSpinBox(); self.preview_steps.setRange(1, 500) - self.intensity = QSpinBox(); self.intensity.setRange(10, 100); self.intensity.setSuffix("%") - self.gradient_checkpointing = QCheckBox("Gradient checkpointing (uses less VRAM)") + self.options_group = QGroupBox() + options = QVBoxLayout(self.options_group) + options_title = QLabel("Training options") + options_title.setProperty("sectionTitle", True) + self.training_form = SettingsForm() self.options_hint = QLabel(); self.options_hint.setProperty("muted", True); self.options_hint.setWordWrap(True) - fields = [("Resolution", self.resolution), ("Batch size", self.batch_size), ("Learning rate", self.learning_rate), ("Gradient accumulation", self.gradient_accumulation), ("Loader workers", self.workers), ("Precision", self.precision), ("Save every", self.save_every), ("Preview steps", self.preview_steps), ("DDPM training intensity", self.intensity)] - for row, (label, widget) in enumerate(fields): options.addWidget(QLabel(label), row, 0); options.addWidget(widget, row, 1) self.orion_settings_button = QPushButton("ORION: apply a starting recipe") self.orion_settings_button.setToolTip("Fill in a conservative draft from the image count and resolution. You can change every value afterward.") self.orion_settings_button.setProperty("chip", True) - options.addWidget(self.orion_settings_button, len(fields), 0, 1, 2) - options.addWidget(self.gradient_checkpointing, len(fields) + 1, 0, 1, 2) - options.addWidget(self.options_hint, len(fields) + 2, 0, 1, 2) + options.addWidget(options_title) + options.addWidget(self.training_form) + options.addWidget(self.orion_settings_button) + options.addWidget(self.options_hint) root.addWidget(self.options_group) self.preview_group = QGroupBox("Live training preview") preview_form = QGridLayout(self.preview_group) @@ -550,6 +708,25 @@ class ModelCreationDialog(QDialog): self.review_summary.setProperty("muted", True) root.addWidget(self.review_summary) + schedule_group = QGroupBox("Model Scheduler") + schedule_layout = QGridLayout(schedule_group) + self.schedule_enabled = QCheckBox("Start this training batch at a specific time") + self.schedule_time = QDateTimeEdit(QDateTime.currentDateTime().addSecs(3600)) + self.schedule_time.setCalendarPopup(True) + self.schedule_time.setDisplayFormat("MMM d, yyyy h:mm AP") + self.schedule_time.setMinimumDateTime(QDateTime.currentDateTime()) + self.schedule_time.setEnabled(False) + schedule_note = QLabel( + "This is the earliest start time. If another job is still running, ADAM starts this batch when that job finishes." + ) + schedule_note.setProperty("muted", True) + schedule_note.setWordWrap(True) + schedule_layout.addWidget(self.schedule_enabled, 0, 0, 1, 2) + schedule_layout.addWidget(QLabel("Schedule training"), 1, 0) + schedule_layout.addWidget(self.schedule_time, 1, 1) + schedule_layout.addWidget(schedule_note, 2, 0, 1, 2) + root.addWidget(schedule_group) + save_row = QHBoxLayout() self.preset_name = QLineEdit() self.preset_name.setPlaceholderText("Optional custom preset name") @@ -574,31 +751,38 @@ class ModelCreationDialog(QDialog): buttons.button(QDialogButtonBox.Ok).setText("Build training plan") buttons.accepted.connect(self._accept_request) buttons.rejected.connect(self.reject) - root.addWidget(buttons) + outer.addWidget(buttons) self.preset.currentTextChanged.connect(self._apply_preset) self.trainer.currentIndexChanged.connect(self._trainer_changed) + self.trainer.currentIndexChanged.connect(self._refresh_continue_models) + self.training_mode.currentIndexChanged.connect(self._training_mode_changed) + self.training_mode.currentIndexChanged.connect(self._update_review) + self.continue_model.currentIndexChanged.connect(self._training_mode_changed) + self.continue_model.currentIndexChanged.connect(self._update_review) self.source.currentIndexChanged.connect(self._update_source) self.source.currentIndexChanged.connect(self._update_review) self.subject.textChanged.connect(self._suggest_name) self.subject.textChanged.connect(self._update_review) self.dataset.currentTextChanged.connect(self._update_review) self.model_name.textChanged.connect(self._update_review) + self.trigger_word.textChanged.connect(self._update_review) self.epochs.valueChanged.connect(self._update_review) self.images.valueChanged.connect(self._update_review) self.collection_mode.currentIndexChanged.connect(self._update_collection_mode) self.collection_mode.currentIndexChanged.connect(self._update_review) self.trainer.currentIndexChanged.connect(self._update_review) - for widget in (self.resolution, self.batch_size, self.learning_rate, self.gradient_accumulation, self.workers, self.precision, self.save_every, self.preview_steps, self.intensity, self.gradient_checkpointing): - signal = getattr(widget, "valueChanged", None) or getattr(widget, "currentIndexChanged", None) or getattr(widget, "stateChanged", None) - if signal: signal.connect(self._update_review) + self.training_form.changed.connect(self._update_review) self.preview_enabled.toggled.connect(self._update_preview_controls) self.preview_enabled.toggled.connect(self._update_review) self.preview_every.valueChanged.connect(self._update_review) self.preview_prompt.textChanged.connect(self._update_review) self.preview_seed.valueChanged.connect(self._update_review) + self._refresh_presets_for_trainer() + self._refresh_continue_models() self._apply_preset(self.preset.currentText()) self._set_training_defaults() + self._training_mode_changed() self._update_source() self._update_review() self._model_states = [self._capture_state()] @@ -606,6 +790,7 @@ class ModelCreationDialog(QDialog): self.model_tabs.tabMoved.connect(self._move_model) self.add_model_button.clicked.connect(self._add_model) self.bulk_add_button.clicked.connect(self._bulk_add_models) + self.create_plugin_button.clicked.connect(self._create_model_plugin) self.apply_many_button.clicked.connect(self._apply_settings_to_models) self.save_draft_button.clicked.connect(self._save_batch_draft) self.load_draft_button.clicked.connect(self._load_batch_draft) @@ -615,12 +800,18 @@ class ModelCreationDialog(QDialog): self.dataset_reviewed.toggled.connect(self._update_review) self.approve_all_datasets_button.clicked.connect(self._approve_all_datasets) self.orion_settings_button.clicked.connect(self._apply_orion_settings) + self.schedule_enabled.toggled.connect(self.schedule_time.setEnabled) + self.schedule_enabled.toggled.connect(self._update_review) + self.schedule_time.dateTimeChanged.connect(self._update_review) def _capture_state(self) -> dict[str, object]: return { "preset": self.preset.currentText(), "trainer": self.trainer.currentData(), + "training_mode": self.training_mode.currentData(), + "continue_model_id": getattr(self.continue_model.currentData(), "id", ""), "source": self.source.currentData(), "subject": self.subject.text(), "dataset": self.dataset.currentText(), "model_name": self.model_name.text(), + "trigger_word": self.trigger_word.text(), "epochs": self.epochs.value(), "images": self.images.value(), "collection_mode": self.collection_mode.currentData(), "dataset_reviewed": self.dataset_reviewed.isChecked(), @@ -628,16 +819,18 @@ class ModelCreationDialog(QDialog): } def _load_state(self, state: dict[str, object]) -> None: - preset_index = self.preset.findText(str(state.get("preset", ""))) - if preset_index >= 0: - self.preset.setCurrentIndex(preset_index) trainer_index = self.trainer.findData(state.get("trainer", "lora")) self.trainer.setCurrentIndex(max(0, trainer_index)) + self._refresh_presets_for_trainer(str(state.get("preset", ""))) + self._refresh_continue_models(str(state.get("continue_model_id", ""))) + mode_index = self.training_mode.findData(state.get("training_mode", "new")) + self.training_mode.setCurrentIndex(max(0, mode_index)) source_index = self.source.findData(state.get("source", "new")) self.source.setCurrentIndex(max(0, source_index)) self.subject.setText(str(state.get("subject", ""))) self.dataset.setCurrentText(str(state.get("dataset", ""))) self.model_name.setText(str(state.get("model_name", ""))) + self.trigger_word.setText(str(state.get("trigger_word", ""))) self.epochs.setValue(int(state.get("epochs", 100))) self.images.setValue(int(state.get("images", 60))) mode_index = self.collection_mode.findData(state.get("collection_mode", "target")) @@ -645,22 +838,13 @@ class ModelCreationDialog(QDialog): self.dataset_reviewed.setChecked(bool(state.get("dataset_reviewed", False))) options = state.get("training_options", {}) if isinstance(options, dict): - self.resolution.setCurrentText(str(options.get("resolution", self.resolution.currentText()))) - self.batch_size.setValue(int(options.get("batch_size", self.batch_size.value()))) - self.learning_rate.setValue(float(options.get("learning_rate", self.learning_rate.value()))) - self.gradient_accumulation.setValue(int(options.get("gradient_accumulation_steps", options.get("gradient_accumulation", self.gradient_accumulation.value())))) - self.workers.setValue(int(options.get("dataloader_num_workers", options.get("workers", self.workers.value())))) - precision_index = self.precision.findData(options.get("mixed_precision", self.precision.currentData())) - self.precision.setCurrentIndex(max(0, precision_index)) - self.save_every.setValue(int(options.get("save_every", self.save_every.value()))) - self.preview_steps.setValue(int(options.get("preview_steps", self.preview_steps.value()))) - self.intensity.setValue(int(options.get("training_intensity", self.intensity.value()))) - self.gradient_checkpointing.setChecked(bool(options.get("gradient_checkpointing", False))) + self.training_form.set_values(options) self.preview_enabled.setChecked(bool(options.get("preview_enabled", True))) self.preview_every.setValue(int(options.get("preview_every", 5))) self.preview_prompt.setText(str(options.get("preview_prompt", ""))) self.preview_seed.setValue(int(options.get("preview_seed", 123456789))) self._update_source() + self._training_mode_changed() self._update_review() def _switch_model(self, index: int) -> None: @@ -674,7 +858,7 @@ class ModelCreationDialog(QDialog): def _add_model(self) -> None: self._model_states[self._current_model_index] = self._capture_state() blank = dict(self._model_states[0]) - blank.update({"subject": "", "dataset": "", "model_name": ""}) + blank.update({"subject": "", "dataset": "", "model_name": "", "continue_model_id": ""}) self._model_states.append(blank) index = self.model_tabs.addTab(f"Model {len(self._model_states)}") self._install_remove_button(index) @@ -723,6 +907,14 @@ class ModelCreationDialog(QDialog): self.model_tabs.setCurrentIndex(0 if current_blank else len(self._model_states) - len(states)) self.validation.setText(f"Added {len(names)} models. Their shared settings came from the current model.") + def _create_model_plugin(self) -> None: + dialog = ModelPluginWizardDialog(self.planner.root, self) + if dialog.exec() != QDialog.Accepted: + return + self.validation.setText( + f"Created {dialog.created_folder}. Restart ADAM after editing the plugin code." + ) + def _rebuild_model_tabs(self) -> None: self.model_tabs.blockSignals(True) while self.model_tabs.count(): @@ -755,7 +947,10 @@ class ModelCreationDialog(QDialog): if dialog.exec() != QDialog.Accepted or not choices.selectedItems(): return source = self._capture_state() - shared_keys = {"preset", "trainer", "epochs", "images", "collection_mode", "training_options"} + shared_keys = { + "preset", "trainer", "training_mode", "continue_model_id", + "epochs", "images", "collection_mode", "training_options", + } for item in choices.selectedItems(): target = self._model_states[int(item.data(Qt.UserRole))] for key in shared_keys: @@ -929,20 +1124,53 @@ class ModelCreationDialog(QDialog): name = str(state.get("model_name", "")).strip() self.model_tabs.setTabText(index, name or f"Model {index + 1}") + def _refresh_presets_for_trainer(self, preferred: str = "") -> None: + trainer = str(self.trainer.currentData() or "") + names = [ + name for name, values in self.presets.items() + if str(values.get("trainer", "")) == trainer + ] + self.preset.blockSignals(True) + self.preset.clear() + if names: + self.preset.addItems(names) + target = preferred if preferred in names else names[0] + self.preset.setCurrentText(target) + else: + self.preset.addItem("Plugin defaults") + self.preset.blockSignals(False) + def _apply_preset(self, name: str) -> None: values = self.presets.get(name, {}) - index = self.trainer.findData(values.get("trainer", "lora")) - self.trainer.setCurrentIndex(max(0, index)) + trainer = str(self.trainer.currentData() or "") + if not values: + plugin = self.planner.registry.model_plugins.by_trainer(trainer) + label = plugin.name if plugin else self.trainer.currentText() + self.preset_hint.setText(f"{label}: using the plugin's default settings.") + self._set_training_defaults() + return + if str(values.get("trainer", "")) != trainer: + return + self._set_training_defaults() self.epochs.setValue(int(values.get("epochs", 100))) self.images.setValue(int(values.get("image_count", 60))) + options = values.get("training_options", {}) + if isinstance(options, dict): + self.training_form.set_values(options) + if self.preview_group.isVisible(): + self.preview_enabled.setChecked(bool(options.get("preview_enabled", True))) + self.preview_every.setValue(int(options.get("preview_every", 5))) + self.preview_prompt.setText(str(options.get("preview_prompt", ""))) + self.preview_seed.setValue(int(options.get("preview_seed", 123456789))) self.preset_hint.setText(str(values.get("description", ""))) def _update_source(self) -> None: + continuing = self.training_mode.currentData() == "continue" creating = self.source.currentData() == "new" self.subject.setEnabled(creating) self.collection_mode.setEnabled(creating) self.images.setEnabled(creating and self.collection_mode.currentData() == "target") - self.dataset.setEnabled(not creating) + self.dataset.setEnabled(not creating and self.source.currentData() != "original") if not creating: self._suggest_name(self.dataset.currentText()) @@ -952,34 +1180,125 @@ class ModelCreationDialog(QDialog): and self.collection_mode.currentData() == "target" ) + def _model_can_continue(self, asset: object) -> bool: + trainer = str(getattr(asset, "trainer", "")) + path = Path(str(getattr(asset, "path", ""))) + checkpoint = str(getattr(asset, "checkpoint", "")) + checkpoint_ready = bool(checkpoint and Path(checkpoint).exists()) + ddpm_pipeline = trainer == "ddpm" and (path / "model_index.json").is_file() + flow_model = trainer == "flow" and ( + (path / "flow_model_info.json").is_file() + and (path / "unet" / "config.json").is_file() + ) + oasis_model = trainer == "oasis" and ( + (path / "action_flow_model_info.json").is_file() + and (path / "unet" / "config.json").is_file() + ) + try: + supports_resume = "resume_training" in self.planner.registry.get( + f"{trainer}_trainer" + ).capabilities + except Exception: + supports_resume = False + return supports_resume and (checkpoint_ready or ddpm_pipeline or flow_model or oasis_model) + + def _refresh_continue_models(self, preferred_id: str = "") -> None: + self.planner.assets.discover(self.config) + trainer = str(self.trainer.currentData() or "") + current = preferred_id or str(getattr(self.continue_model.currentData(), "id", "")) + self.continue_model.blockSignals(True) + self.continue_model.clear() + models = [ + asset for asset in self.planner.assets.assets + if asset.kind == "model" and asset.trainer == trainer and self._model_can_continue(asset) + ] + for asset in sorted(models, key=lambda item: item.name.casefold()): + self.continue_model.addItem(asset.name, asset) + if not models: + self.continue_model.addItem("No resumable models found") + self.continue_model.model().item(0).setEnabled(False) + elif current: + index = next( + ( + i for i in range(self.continue_model.count()) + if getattr(self.continue_model.itemData(i), "id", "") == current + ), + -1, + ) + if index >= 0: + self.continue_model.setCurrentIndex(index) + self.continue_model.blockSignals(False) + + def _continued_model_has_original_dataset(self) -> bool: + asset = self.continue_model.currentData() + if not asset: + return False + return any( + item.kind == "dataset" + and item.id == getattr(asset, "dataset_id", "") + and Path(item.path).is_dir() + for item in self.planner.assets.assets + ) + + def _training_mode_changed(self) -> None: + continuing = self.training_mode.currentData() == "continue" + has_original_dataset = self._continued_model_has_original_dataset() + self.continue_model.setVisible(continuing) + if hasattr(self, "form_labels"): + self.form_labels["Model to continue"].setVisible(continuing) + self.source.model().item(self.source.findData("original")).setEnabled( + continuing and has_original_dataset + ) + if not continuing and self.source.currentData() == "original": + self.source.setCurrentIndex(self.source.findData("existing")) + if continuing and self.source.currentData() == "new" and self.trainer.currentData() in {"flow", "oasis"}: + target = "original" if has_original_dataset else "existing" + self.source.setCurrentIndex(self.source.findData(target)) + if continuing and self.source.currentData() == "original" and not has_original_dataset: + self.source.setCurrentIndex(self.source.findData("existing")) + self._update_source() + def _trainer_changed(self) -> None: + self._refresh_presets_for_trainer() flow = self.trainer.currentData() == "flow" - if flow: + oasis = self.trainer.currentData() == "oasis" + self.trigger_word.setVisible(self.trainer.currentData() == "lora") + if self.training_mode.currentData() != "continue" and (flow or oasis): self.source.setCurrentIndex(self.source.findData("existing")) - self.source.model().item(self.source.findData("new")).setEnabled(not flow) + self.source.model().item(self.source.findData("new")).setEnabled( + self.training_mode.currentData() == "continue" or not (flow or oasis) + ) if flow: self.preset_hint.setText( "Flow Matching currently uses an existing reviewed dataset. " "Create a dataset first if you do not have one yet." ) - self._set_training_defaults() + if oasis: + self.preset_hint.setText( + "Oasis uses existing gameplay folders with frames and synchronized action labels. " + "Record or convert an action dataset before training." + ) + self._apply_preset(self.preset.currentText()) def _set_training_defaults(self) -> None: - trainer = self.trainer.currentData() - enabled = trainer in {"ddpm", "flow"} - self.options_group.setEnabled(enabled) - if trainer == "flow": - values = ("256", 8, 0.0002, 1, 4, 10, 10, 30) - self.intensity.hide(); self.gradient_checkpointing.show() - self.options_hint.setText("Flow: higher resolution and batch size need substantially more VRAM. Heun/preview settings remain in the Flow app.") - elif trainer == "ddpm": - values = ("128", 1, 0.0001, 1, 4, 10, 50, 100) - self.intensity.show(); self.gradient_checkpointing.hide() - self.options_hint.setText("DDPM: resolution has the biggest speed and VRAM impact. Keep batch size at 1 if you are unsure.") + trainer = str(self.trainer.currentData()) + full_schema = self.planner.registry.model_plugins.training_schema(trainer) + builtin_preview = trainer in self.BUILTIN_PREVIEW_TRAINERS + schema = { + key: spec for key, spec in full_schema.items() + if not (builtin_preview and key.startswith("preview_")) + } + self.training_form.set_schema(schema) + self.preview_group.setVisible(builtin_preview) + self.orion_settings_button.setEnabled(trainer in {"ddpm", "flow", "lora", "oasis"}) + plugin = self.planner.registry.model_plugins.by_trainer(trainer) + self.options_group.setEnabled(bool(schema)) + if plugin: + self.options_hint.setText( + f"{plugin.name}: settings are generated from the model plugin manifest." + ) else: - self.options_hint.setText("LoRA uses its connected trainer's saved settings for now.") - return - self.resolution.setCurrentText(values[0]); self.batch_size.setValue(values[1]); self.learning_rate.setValue(values[2]); self.gradient_accumulation.setValue(values[3]); self.workers.setValue(values[4]); self.save_every.setValue(values[5]); self.preview_steps.setValue(values[6]); self.intensity.setValue(values[7]); self.gradient_checkpointing.setChecked(False) + self.options_hint.setText("This trainer has no model plugin manifest yet.") def _orion_image_count(self) -> int: if self.source.currentData() == "new": @@ -994,40 +1313,57 @@ class ModelCreationDialog(QDialog): def _apply_orion_settings(self) -> None: trainer = str(self.trainer.currentData()) images = self._orion_image_count() - recommendation = recommend_training_settings( - trainer, images, int(self.resolution.currentText()) - ) - self.epochs.setValue(int(recommendation["epochs"])) - settings = recommendation["settings"] - if trainer in {"ddpm", "flow"}: - self.batch_size.setValue(int(settings["batch_size"])) - self.learning_rate.setValue(float(settings["learning_rate"])) - self.gradient_accumulation.setValue(int(settings["gradient_accumulation_steps"])) - self.workers.setValue(int(settings["dataloader_num_workers"])) - precision = self.precision.findData(settings["mixed_precision"]) - self.precision.setCurrentIndex(max(0, precision)) - self.save_every.setValue(int(settings["save_every"])) - self.preview_steps.setValue(int(settings["preview_steps"])) + current = self.training_form.values() + raw_resolution = current.get("resolution", 128) or 128 + if isinstance(raw_resolution, str) and "x" in raw_resolution: + resolution = int(raw_resolution.lower().split("x", 1)[0]) + else: + resolution = int(raw_resolution) + profile = ModelProfileRegistry(self.planner.registry.model_plugins).get(trainer) + if profile: + result = recommend_for_profile( + profile, + dataset_items=images, + resolution=resolution, + snapshot=getattr(self.parent(), "latest_snapshot", None), + ).to_dict() + else: + result = recommend_training_settings(trainer, images, resolution) + self.epochs.setValue(int(result["epochs"])) + settings = result["settings"] + self.training_form.set_values(settings) + if "preview_every" in settings: self.preview_every.setValue(int(settings["preview_every"])) - self.intensity.setValue(int(settings["training_intensity"])) - self.gradient_checkpointing.setChecked(bool(settings["gradient_checkpointing"])) - self.preset_hint.setText(str(recommendation["summary"])) - self.validation.setText("ORION applied a reviewable starting recipe. Nothing has been queued or started.") + warnings = result.get("warnings", []) + self.preset_hint.setText(str(result["summary"])) + reasons = result.get("reasons", []) + explanation = "" + if reasons: + explanation = "\nWhy: " + " ".join(str(reason) for reason in reasons[:3]) + self.validation.setText( + "ORION applied a reviewable starting recipe. Nothing has been queued or started." + + explanation + + ("\n" + "\n".join(str(item) for item in warnings) if warnings else "") + ) self._update_review() def _training_options(self) -> dict[str, object]: - trainer = self.trainer.currentData() - common = { - "preview_enabled": self.preview_enabled.isChecked(), - "preview_every": self.preview_every.value(), - "preview_prompt": self.preview_prompt.text().strip(), - "preview_seed": self.preview_seed.value(), - } - if trainer == "ddpm": - return {"resolution": int(self.resolution.currentText()), "batch_size": self.batch_size.value(), "learning_rate": self.learning_rate.value(), "gradient_accumulation_steps": self.gradient_accumulation.value(), "dataloader_num_workers": self.workers.value(), "mixed_precision": self.precision.currentData(), "save_every": self.save_every.value(), "preview_steps": self.preview_steps.value(), "training_intensity": self.intensity.value(), **common} - if trainer == "flow": - return {"resolution": int(self.resolution.currentText()), "batch_size": self.batch_size.value(), "learning_rate": self.learning_rate.value(), "gradient_accumulation": self.gradient_accumulation.value(), "workers": self.workers.value(), "mixed_precision": self.precision.currentData(), "save_every": self.save_every.value(), "preview_steps": self.preview_steps.value(), "gradient_checkpointing": self.gradient_checkpointing.isChecked(), **common} - return common + plugin_options = self.training_form.values() + if self.preview_group.isVisible(): + plugin_options.update({ + "preview_enabled": self.preview_enabled.isChecked(), + "preview_every": self.preview_every.value(), + }) + if self.trainer.currentData() != "oasis": + plugin_options.update({ + "preview_prompt": self.preview_prompt.text().strip(), + "preview_seed": self.preview_seed.value(), + }) + if self.trainer.currentData() == "lora": + trigger_word = self.trigger_word.text().strip() + if trigger_word: + plugin_options["trigger_word"] = trigger_word + return plugin_options def _update_preview_controls(self) -> None: enabled = self.preview_enabled.isChecked() @@ -1040,10 +1376,14 @@ class ModelCreationDialog(QDialog): self.model_name.setPlaceholderText(value.strip() or "Model name") def _update_review(self) -> None: + continuing = self.training_mode.currentData() == "continue" creating = self.source.currentData() == "new" subject = self.subject.text().strip() if creating else self.dataset.currentText().strip() model = self.model_name.text().strip() or subject or "Unnamed model" - if creating: + if self.source.currentData() == "original": + base = self.continue_model.currentData() + dataset = f"reuse the dataset linked to {getattr(base, 'name', 'the selected model')}" + elif creating: dataset = ( f"collect every result Bing makes available (up to 5,000) for {subject or 'the subject'}" if self.collection_mode.currentData() == "all_available" @@ -1051,16 +1391,51 @@ class ModelCreationDialog(QDialog): ) else: dataset = f"use the registered {subject or 'selected'} dataset" + if continuing: + base = self.continue_model.currentData() + action = f"continue {getattr(base, 'name', 'a selected model')}" + else: + action = f"train {model}" self.review_summary.setText( - f"Review: {dataset}; train {model} with " + f"Review: {dataset}; {action} with " f"{self.trainer.currentText()} for {self.epochs.value():,} epochs. " - + (f"{self.resolution.currentText()}px · batch {self.batch_size.value()} · lr {self.learning_rate.value():.7f}. " if self.trainer.currentData() in {"ddpm", "flow"} else "") - + (f"Live preview every {self.preview_every.value()} epochs. " if self.preview_enabled.isChecked() else "Live previews off. ") + + ( + f"Trigger: {self.trigger_word.text().strip() or model}. " + if self.trainer.currentData() == "lora" and not continuing else "" + ) + + self._training_option_summary() + + self._preview_summary() + + ( + f"Scheduled for {self.schedule_time.dateTime().toString('MMM d, yyyy h:mm AP')}. " + if hasattr(self, "schedule_enabled") and self.schedule_enabled.isChecked() + else "" + ) + "ADAM will run preflight checks and still ask for approval." ) if hasattr(self, "model_tabs") and self.model_tabs.count(): self.model_tabs.setTabText(self.model_tabs.currentIndex(), model) + def _preview_summary(self) -> str: + if self.preview_group.isVisible(): + if self.preview_enabled.isChecked(): + return f"Live preview every {self.preview_every.value()} epochs. " + return "Live previews off. " + values = self.training_form.values() + if values.get("preview_every"): + return f"Plugin preview every {values['preview_every']} steps. " + return "" + + def _training_option_summary(self) -> str: + values = self.training_form.values() + details = [] + if values.get("resolution"): + details.append(f"{values['resolution']}px") + if values.get("batch_size"): + details.append(f"batch {values['batch_size']}") + if values.get("learning_rate"): + details.append(f"lr {float(values['learning_rate']):.7f}") + return " · ".join(details) + ". " if details else "" + def _save_preset(self) -> None: name = self.preset_name.text().strip() if not name: @@ -1077,8 +1452,7 @@ class ModelCreationDialog(QDialog): } self.config.update({"training_presets": stored}) self.presets[name] = stored[name] - if self.preset.findText(name) < 0: - self.preset.addItem(name) + self._refresh_presets_for_trainer(name) self.preset.setCurrentText(name) self.validation.setText(f"Saved preset: {name}") @@ -1086,9 +1460,44 @@ class ModelCreationDialog(QDialog): self._model_states[self._current_model_index] = self._capture_state() requests: list[str] = [] for index, state in enumerate(self._model_states, 1): + continuing = state.get("training_mode") == "continue" creating = state.get("source") == "new" subject = str(state.get("subject", "")).strip() dataset = str(state.get("dataset", "")).strip() + if continuing: + model = next( + ( + asset for asset in self.planner.assets.assets + if asset.kind == "model" + and asset.id == str(state.get("continue_model_id", "")) + and self._model_can_continue(asset) + ), + None, + ) + if not model: + self.validation.setText(f"Model {index}: choose a completed model to continue.") + self.model_tabs.setCurrentIndex(index - 1) + return + if state.get("source") == "existing" and not dataset: + self.validation.setText(f"Model {index}: choose the gameplay dataset for continuation.") + self.model_tabs.setCurrentIndex(index - 1) + return + if state.get("source") == "new" and not subject: + self.validation.setText(f"Model {index}: tell ADAM what the new dataset should contain.") + self.model_tabs.setCurrentIndex(index - 1) + return + options = state.get("training_options", {}) + requests.append(build_fine_tune_request( + model_name=model.name, + trainer=model.trainer, + epochs=int(state.get("epochs", 100)), + dataset_mode=str(state.get("source", "original")), + dataset_name=dataset, + new_subject=subject, + image_count=int(state.get("images", 60)), + training_options=options if isinstance(options, dict) else {}, + )) + continue if creating and not subject: self.validation.setText(f"Model {index}: tell ADAM what it should learn.") self.model_tabs.setCurrentIndex(index - 1) @@ -1114,6 +1523,14 @@ class ModelCreationDialog(QDialog): )) self.requests = requests self.request = requests[0] + if self.schedule_enabled.isChecked(): + scheduled = self.schedule_time.dateTime() + if scheduled <= QDateTime.currentDateTime(): + self.validation.setText("Choose a scheduled training time in the future.") + return + self.scheduled_for = scheduled.toPython().astimezone().isoformat() + else: + self.scheduled_for = None self.config.update({"model_batch_draft": { "saved_at": datetime.now().isoformat(timespec="seconds"), "models": self._model_states, @@ -1245,9 +1662,9 @@ class FineTuneDialog(QDialog): def _add_model_groups(self) -> None: """Show every model family, while allowing only safe continuation choices.""" - labels = {"ddpm": "DDPM models", "flow": "Flow Matching models", "lora": "LoRA models"} + labels = {"ddpm": "DDPM models", "flow": "Flow Matching models", "lora": "LoRA models", "oasis": "Oasis models"} models = [asset for asset in self.planner.assets.assets if asset.kind == "model"] - for trainer in ("ddpm", "flow", "lora"): + for trainer in ("ddpm", "flow", "lora", "oasis"): header_index = self.model.count() self.model.addItem(f"— {labels[trainer]} —") self.model.model().item(header_index).setEnabled(False) @@ -1264,13 +1681,17 @@ class FineTuneDialog(QDialog): (Path(asset.path) / "flow_model_info.json").is_file() and (Path(asset.path) / "unet" / "config.json").is_file() ) + oasis_model = trainer == "oasis" and ( + (Path(asset.path) / "action_flow_model_info.json").is_file() + and (Path(asset.path) / "unet" / "config.json").is_file() + ) try: supports_resume = "resume_training" in self.planner.registry.get( f"{trainer}_trainer" ).capabilities except Exception: supports_resume = False - ready = supports_resume and (checkpoint_ready or ddpm_pipeline or flow_model) + ready = supports_resume and (checkpoint_ready or ddpm_pipeline or flow_model or oasis_model) if ready: detail = ( "saved Flow model" if flow_model else @@ -1327,7 +1748,7 @@ class FineTuneDialog(QDialog): self.dataset_mode.model().item(original_index).setEnabled(has_original_dataset or not asset) if asset and not has_original_dataset and self.dataset_mode.currentData() == "original": self.dataset_mode.setCurrentIndex(self.dataset_mode.findData("existing")) - self.options_group.setEnabled(trainer in {"ddpm", "flow"}) + self.options_group.setEnabled(trainer in {"ddpm", "flow", "oasis"}) self.resolution.setEnabled(trainer != "flow") if trainer == "flow": try: @@ -1345,6 +1766,8 @@ class FineTuneDialog(QDialog): ) elif trainer == "ddpm": self.options_hint.setText("These settings are passed to the DDPM trainer for this continuation run.") + elif trainer == "oasis": + self.options_hint.setText("Oasis continuation uses an existing action model folder and a reviewed gameplay action dataset.") else: self.options_hint.setText("The connected LoRA trainer currently reuses its saved training settings; choose the additional epochs above.") self._update_summary() @@ -1370,6 +1793,14 @@ class FineTuneDialog(QDialog): "save_every": self.save_every.value(), "preview_every": self.save_every.value(), "preview_steps": self.preview_steps.value(), "gradient_checkpointing": False, } + if trainer == "oasis": + return { + "resolution": "256x144", "batch_size": min(2, self.batch_size.value()), + "learning_rate": self.learning_rate.value(), "gradient_accumulation": self.gradient_accumulation.value(), + "workers": min(2, self.workers.value()), "mixed_precision": "fp32", + "save_every": self.save_every.value(), "preview_every": self.save_every.value(), + "preview_steps": 1, + } return {} def _accept_request(self) -> None: @@ -1596,7 +2027,7 @@ class RecentPlansPanel(QFrame): if item.widget(): item.widget().deleteLater() queued = sum( - job.status in {JobStatus.QUEUED, JobStatus.AWAITING_CONFIRMATION} + job.status in {JobStatus.SCHEDULED, JobStatus.QUEUED, JobStatus.AWAITING_CONFIRMATION} for job in self.jobs.jobs ) self.queue_label.setText(f"{queued} queued" if queued else "Queue clear") @@ -1615,6 +2046,7 @@ class RecentPlansPanel(QFrame): JobStatus.CANCELLED: "×", JobStatus.INTERRUPTED: "!", JobStatus.AWAITING_CONFIRMATION: "?", + JobStatus.SCHEDULED: "◷", JobStatus.QUEUED: "…", } for job in recent: @@ -1637,7 +2069,7 @@ class SystemSummaryPanel(QFrame): def __init__(self) -> None: super().__init__() self.setProperty("card", True) - self.setMaximumHeight(124) + self.setMinimumHeight(124) root = QVBoxLayout(self) root.setContentsMargins(16, 11, 16, 11) root.setSpacing(6) @@ -1752,6 +2184,9 @@ class CommandCenterPage(QWidget): provider_changed = Signal(str) tool_folders_changed = Signal() open_jobs_requested = Signal() + open_dataset_lab_requested = Signal() + open_experiments_requested = Signal() + open_remote_requested = Signal() history_changed = Signal() def __init__( @@ -1781,6 +2216,8 @@ class CommandCenterPage(QWidget): self.history_store = ChatHistoryStore(root_path) self._generation_cards: dict[str, GenerationChatCard] = {} self._prompt_reference_image = "" + self._pending_schedule_for: str | None = None + self.latest_snapshot = SystemSnapshot() root = QVBoxLayout(self) root.setContentsMargins(24, 20, 24, 17) @@ -1861,6 +2298,7 @@ class CommandCenterPage(QWidget): self.active_panel.resume_requested.connect(self.jobs.resume) self.active_panel.cancel_requested.connect(self.jobs.cancel) self.active_panel.open_requested.connect(self.open_output) + self.jobs.job_created.connect(self._job_created) self.jobs.job_updated.connect(self._job_updated) self.jobs.active_changed.connect(self.active_panel.set_job) self.jobs.active_changed.connect( @@ -1939,23 +2377,21 @@ class CommandCenterPage(QWidget): actions_content_layout.addLayout(self.actions_grid) self.action_cards: list[QPushButton] = [] action_specs = ( - ("Collect a dataset", "Gather captioned images\nwith filters.", "Adam, collect a dataset of liminal spaces"), - ("Train a LoRA", "Fine-tune Stable Diffusion\nwith your dataset.", "Adam, train a LoRA of Hatsune Miku"), - ("Train DDPM", "Train a diffusion model\nfrom scratch.", "Adam, train a DDPM model"), - ("Train Flow Matching", "Build an image or video\nflow model.", "Adam, train a Flow Matching model"), - ("Generate previews", "Review model outputs\nbefore export.", "Adam, generate 4 previews"), - ("Inspect GPU", "Check VRAM, utilization,\ndrivers and heat.", "Adam, check GPU status"), - ("Generate an image", "Create it here from a\ncompleted model.", 'Generate a DDPM image of "A new sample" for 100 steps on DDIM sampler with aspect ratio 16:9'), - ) - for title, description, command in action_specs: + ("Auto-detect settings", "Recommend a safe\ntraining recipe.", self._open_model_assistant), + ("Dataset Lab / EVE", "Inspect datasets,\ncaptions, and files.", self.open_dataset_lab_requested.emit), + ("Compare runs", "Review experiment\ndifferences side by side.", self.open_experiments_requested.emit), + ("Remote control", "Configure browser\nstatus access.", self.open_remote_requested.emit), + ("Transcript to dataset", "Convert local video\nspeech into samples.", self.open_dataset_lab_requested.emit), + ("Overnight queue", "Plan sequential jobs\nfor unattended runs.", self._open_model_batch_assistant), + ("Generate an image", "Create it here from a\ncompleted model.", lambda: self.submit('Generate a DDPM image of "A new sample" for 100 steps on DDIM sampler with aspect ratio 16:9')), + ) + for title, description, callback in action_specs: button = QPushButton(f"{title}\n{description}") button.setProperty("workflowCard", True) button.setMinimumWidth(0) button.setSizePolicy(QSizePolicy.Ignored, QSizePolicy.Preferred) button.setToolTip(f"Start: {title}") - button.clicked.connect( - lambda _checked=False, text=command: self.submit(text) - ) + button.clicked.connect(lambda _checked=False, action=callback: action()) self.action_cards.append(button) self._reflow_actions(2) actions_root.addWidget(self.actions_content) @@ -2181,6 +2617,7 @@ class CommandCenterPage(QWidget): self.mode_selector.setCurrentIndex(self.mode_selector.findData("trainer")) dialog = ModelCreationDialog(self.planner, self.config, self) if dialog.exec() == QDialog.Accepted and dialog.requests: + self._pending_schedule_for = dialog.scheduled_for if len(dialog.requests) == 1: self.submit(dialog.requests[0]) else: @@ -2193,6 +2630,7 @@ class CommandCenterPage(QWidget): dialog.setWindowTitle("Model Batch Builder") QTimer.singleShot(0, dialog._bulk_add_models) if dialog.exec() == QDialog.Accepted and dialog.requests: + self._pending_schedule_for = dialog.scheduled_for self.submit_training_batch(dialog.requests) def submit_training_batch(self, requests: list[str]) -> None: @@ -2610,6 +3048,7 @@ class CommandCenterPage(QWidget): self._planning_bubble.set_text(self._streamed_text) def _planning_failed(self, message: str) -> None: + self._pending_schedule_for = None if self._planning_bubble: self._planning_bubble.set_label("ADAM · NEEDS INPUT") self._planning_bubble.set_text(message) @@ -2703,6 +3142,7 @@ class CommandCenterPage(QWidget): {"text": final_response, "user": False, "label": "ADAM"} ) if not plan.steps: + self._pending_schedule_for = None label = { "Conversation": "ADAM", "DDPM training": "ADAM · NEEDS DETAILS", @@ -2715,7 +3155,9 @@ class CommandCenterPage(QWidget): self._type_into(self._planning_bubble, plan.summary) return append_preflight_summary(plan, self.config) - job = self.jobs.submit(plan) + scheduled_for = self._pending_schedule_for + self._pending_schedule_for = None + job = self.jobs.submit(plan, scheduled_for=scheduled_for) self.selected_job = job self.plan_panel.set_job(job) trusted_start = self._can_trusted_start(plan) @@ -2726,6 +3168,8 @@ class CommandCenterPage(QWidget): if trusted_start else "Review the plan at right. I’m waiting for your approval." if plan.requires_confirmation + else "The plan is scheduled and will start automatically when its time and the training slot are available." + if job.status == JobStatus.SCHEDULED else "The plan uses safe, read-only or output-only tools, so it has been queued." ) text = f"{plan.summary}\n\n{len(plan.steps)} registered steps · {state}" @@ -2775,7 +3219,14 @@ class CommandCenterPage(QWidget): self.selected_job = job self.plan_panel.set_job(job) + def _job_created(self, job: Job) -> None: + self.selected_job = job + self.plan_panel.set_job(job) + if job.status == JobStatus.AWAITING_CONFIRMATION: + self.plan_shell.set_collapsed(False, persist=False) + def update_snapshot(self, snapshot: SystemSnapshot) -> None: + self.latest_snapshot = snapshot self.system_summary.update_snapshot(snapshot) def _job_updated(self, job: Job) -> None: @@ -2914,6 +3365,7 @@ class JobsPage(QWidget): actions.setHorizontalSpacing(7) actions.setVerticalSpacing(7) self.pause_button = QPushButton("Pause") + self.adjust_button = QPushButton("Adjust after epoch") self.stop_button = QPushButton("Stop") self.stop_button.setProperty("danger", True) self.end_task_button = QPushButton("End task") @@ -2926,9 +3378,10 @@ class JobsPage(QWidget): self.clear_terminal_button = QPushButton("Remove completed / failed") self.clear_terminal_button.setProperty("danger", True) actions.addWidget(self.pause_button, 0, 0) - actions.addWidget(self.stop_button, 0, 1) - actions.addWidget(self.end_task_button, 0, 2) - actions.addWidget(self.retry_button, 0, 3) + actions.addWidget(self.adjust_button, 0, 1) + actions.addWidget(self.stop_button, 0, 2) + actions.addWidget(self.end_task_button, 0, 3) + actions.addWidget(self.retry_button, 0, 4) actions.addWidget(self.export_button, 1, 0) actions.addWidget(self.full_log_button, 1, 1) actions.addWidget(self.export_all_button, 1, 2) @@ -2939,6 +3392,7 @@ class JobsPage(QWidget): root.addLayout(body, 1) self.pause_button.clicked.connect(self._pause_or_resume) + self.adjust_button.clicked.connect(self._adjust_after_epoch) self.stop_button.clicked.connect(self._stop) self.end_task_button.clicked.connect(self._end_task) self.output_button.clicked.connect(self._open_output) @@ -3022,6 +3476,11 @@ class JobsPage(QWidget): step = f" · {job.plan.steps[job.current_step].title}" self.detail_status.setText( f"{job.status.value} · {job.progress}% · {len(job.plan.steps)} steps{step}" + + ( + f" · starts {self._format_time(job.scheduled_for)}" + if job.status == JobStatus.SCHEDULED and job.scheduled_for + else "" + ) ) reports = [] if job.plan.orion_review: @@ -3047,10 +3506,17 @@ class JobsPage(QWidget): ) running = job.status in {JobStatus.RUNNING, JobStatus.PAUSED} self.pause_button.setEnabled(running) + active_ddpm = ( + running + and 0 <= job.current_step < len(job.plan.steps) + and job.plan.steps[job.current_step].tool_id == "ddpm_trainer" + ) + self.adjust_button.setEnabled(active_ddpm) self.pause_button.setText("Resume" if job.status == JobStatus.PAUSED else "Pause") self.stop_button.setEnabled( running or job.status in { + JobStatus.SCHEDULED, JobStatus.QUEUED, JobStatus.AWAITING_CONFIRMATION, } @@ -3058,7 +3524,12 @@ class JobsPage(QWidget): self.end_task_button.setEnabled(job.status == JobStatus.INTERRUPTED) self.output_button.setEnabled(bool(job.output_folder)) awaiting_confirmation = job.status == JobStatus.AWAITING_CONFIRMATION - self.retry_button.setText("Approve and run" if awaiting_confirmation else "Retry plan") + vram_failure = job.status == JobStatus.FAILED and self.jobs._looks_like_vram_failure(job) + self.retry_button.setText( + "Approve and run" if awaiting_confirmation + else "Retry with safer batch" if vram_failure + else "Retry plan" + ) self.retry_button.setEnabled( awaiting_confirmation or job.status in { JobStatus.FINISHED, @@ -3112,6 +3583,49 @@ class JobsPage(QWidget): else: self.jobs.pause(job.id) + def _adjust_after_epoch(self) -> None: + job = self._selected() + if not job or not (0 <= job.current_step < len(job.plan.steps)): + return + arguments = job.plan.steps[job.current_step].arguments + dialog = QDialog(self) + dialog.setWindowTitle("Adjust training after this epoch") + dialog.setMinimumWidth(440) + layout = QVBoxLayout(dialog) + explanation = QLabel( + "Training will finish the current epoch, save a complete checkpoint, " + "release VRAM, and resume with these settings." + ) + explanation.setWordWrap(True) + layout.addWidget(explanation) + grid = QGridLayout() + batch = QSpinBox(); batch.setRange(1, 64); batch.setValue(int(arguments.get("batch_size", 1))) + accumulation = QSpinBox(); accumulation.setRange(1, 64); accumulation.setValue(int(arguments.get("gradient_accumulation_steps", 1))) + intensity = QSpinBox(); intensity.setRange(10, 100); intensity.setValue(int(arguments.get("training_intensity", 100))); intensity.setSuffix("%") + grid.addWidget(QLabel("Batch size"), 0, 0); grid.addWidget(batch, 0, 1) + grid.addWidget(QLabel("Gradient accumulation"), 1, 0); grid.addWidget(accumulation, 1, 1) + grid.addWidget(QLabel("Training intensity"), 2, 0); grid.addWidget(intensity, 2, 1) + layout.addLayout(grid) + note = QLabel("Tip: when lowering batch size, increase gradient accumulation to preserve a similar effective batch.") + note.setWordWrap(True) + note.setProperty("muted", True) + layout.addWidget(note) + buttons = QDialogButtonBox(QDialogButtonBox.Cancel | QDialogButtonBox.Ok) + buttons.button(QDialogButtonBox.Ok).setText("Apply after epoch") + buttons.accepted.connect(dialog.accept) + buttons.rejected.connect(dialog.reject) + layout.addWidget(buttons) + if dialog.exec() != QDialog.Accepted: + return + try: + self.jobs.request_training_adjustment(job.id, { + "batch_size": batch.value(), + "gradient_accumulation_steps": accumulation.value(), + "training_intensity": intensity.value(), + }) + except ValueError as exc: + QMessageBox.information(self, "Settings not queued", str(exc)) + def _stop(self) -> None: job = self._selected() if job: @@ -3135,7 +3649,15 @@ class JobsPage(QWidget): if job.status == JobStatus.AWAITING_CONFIRMATION: self.jobs.confirm(job.id) return - retried = self.jobs.retry(job.id) + try: + retried = ( + self.jobs.safer_vram_retry(job.id) + if job.status == JobStatus.FAILED and self.jobs._looks_like_vram_failure(job) + else self.jobs.retry(job.id) + ) + except ValueError as exc: + QMessageBox.information(self, "Retry not created", str(exc)) + return self.selected_job_id = retried.id self.refresh() @@ -3217,6 +3739,1269 @@ class JobsPage(QWidget): return False +class ModelInspectionWorker(QThread): + progress_changed = Signal(int, str) + inspection_ready = Signal(object) + failed = Signal(str) + + def __init__(self, path: str, architecture: str = "", settings: dict | None = None) -> None: + super().__init__() + self.path = path + self.architecture = architecture + self.settings = settings or {} + self._cancelled = False + + def cancel(self) -> None: + self._cancelled = True + + def run(self) -> None: + try: + summary = inspect_model( + self.path, + recorded_architecture=self.architecture, + run_settings=self.settings, + progress=self.progress_changed.emit, + cancelled=lambda: self._cancelled, + ) + except Exception as exc: + self.failed.emit(str(exc)) + return + self.inspection_ready.emit(summary) + + +class ModelComparisonWorker(QThread): + progress_changed = Signal(int, str) + comparison_ready = Signal(object) + failed = Signal(str) + + def __init__(self, path_a: str, path_b: str, run_a: ExperimentRun | None = None, run_b: ExperimentRun | None = None) -> None: + super().__init__() + self.path_a = path_a + self.path_b = path_b + self.run_a = run_a + self.run_b = run_b + self._cancelled = False + + def cancel(self) -> None: + self._cancelled = True + + def run(self) -> None: + try: + comparison = compare_models( + self.path_a, + self.path_b, + arch_a=self.run_a.model_architecture if self.run_a else "", + arch_b=self.run_b.model_architecture if self.run_b else "", + settings_a=self.run_a.settings if self.run_a else {}, + settings_b=self.run_b.settings if self.run_b else {}, + progress=self.progress_changed.emit, + cancelled=lambda: self._cancelled, + ) + except Exception as exc: + self.failed.emit(str(exc)) + return + self.comparison_ready.emit(comparison) + + +class ModelInspectorPlot(QWidget): + def __init__(self) -> None: + super().__init__() + self.layout = QVBoxLayout(self) + self.layout.setContentsMargins(8, 8, 8, 8) + self.placeholder = QLabel("Inspect a model to see parameter and weight-distribution charts.") + self.placeholder.setProperty("muted", True) + self.placeholder.setAlignment(Qt.AlignCenter) + self.layout.addWidget(self.placeholder, 1) + self.canvas = None + self.figure = None + try: + from matplotlib.backends.backend_qtagg import FigureCanvasQTAgg + from matplotlib.figure import Figure + + self.figure = Figure(figsize=(9, 4.5), facecolor=COLORS["surface"]) + self.canvas = FigureCanvasQTAgg(self.figure) + self.layout.addWidget(self.canvas, 1) + self.canvas.hide() + except Exception: + self.placeholder.setText("Charts are unavailable because matplotlib could not initialize.") + + def plot(self, summary: ModelInspection | None) -> None: + if self.canvas is None or self.figure is None or summary is None: + return + self.placeholder.hide() + self.canvas.show() + self.figure.clear() + axes = self.figure.subplots(1, 3) + for axis in axes: + axis.set_facecolor(COLORS["surface"]) + axis.tick_params(colors=COLORS["muted"], labelsize=8) + for spine in axis.spines.values(): + spine.set_color(COLORS["border"]) + + components = sorted(summary.components.items(), key=lambda item: item[1], reverse=True)[:10] + axes[0].set_title("Component Parameters", color=COLORS["text"], fontsize=10) + if components: + labels = [name[:18] for name, _ in components] + values = [count / 1_000_000 for _, count in components] + axes[0].barh(labels[::-1], values[::-1], color=COLORS["blue"]) + axes[0].set_xlabel("Millions", color=COLORS["muted"], fontsize=8) + + sizes = summary.tensor_size_distribution[:20] + axes[1].set_title("Largest Tensors", color=COLORS["text"], fontsize=10) + if sizes: + labels = [Path(name).name[:16] for name, _ in sizes[:10]] + values = [count / 1_000_000 for _, count in sizes[:10]] + axes[1].bar(range(len(values)), values, color=COLORS["purple"]) + axes[1].set_xticks(range(len(labels))) + axes[1].set_xticklabels(labels, rotation=70, ha="right") + axes[1].set_ylabel("Millions", color=COLORS["muted"], fontsize=8) + + axes[2].set_title("Abs Mean Distribution", color=COLORS["text"], fontsize=10) + bins = summary.histogram.get("abs_mean_bins", []) + counts = summary.histogram.get("counts", []) + if len(bins) > 1 and counts: + axes[2].bar(range(len(counts)), counts, color=COLORS["green"]) + axes[2].set_xticks([0, len(counts) - 1]) + axes[2].set_xticklabels([f"{bins[0]:.2g}", f"{bins[-1]:.2g}"]) + axes[2].set_ylabel("Tensors", color=COLORS["muted"], fontsize=8) + self.figure.tight_layout() + self.canvas.draw_idle() + + +class ExperimentTrackerPage(QWidget): + clone_requested = Signal(str) + + def __init__(self, store: ExperimentStore) -> None: + super().__init__() + self.store = store + self.selected_run_id = "" + self.current_inspection: ModelInspection | None = None + self.current_model_path = "" + self.inspection_cache: dict[tuple[str, float, int], ModelInspection] = {} + self.inspection_worker: ModelInspectionWorker | None = None + self.comparison_worker: ModelComparisonWorker | None = None + root = QVBoxLayout(self) + root.setContentsMargins(24, 20, 24, 17) + root.setSpacing(12) + header = QHBoxLayout() + header.addWidget( + _page_header( + "Experiment tracker", + "Search, score, clone, and compare local training runs recorded from ADAM jobs.", + ), + 1, + ) + refresh = QPushButton("Refresh") + refresh.clicked.connect(self.refresh) + header.addWidget(refresh, 0, Qt.AlignTop) + root.addLayout(header) + + filters = QHBoxLayout() + self.search = QLineEdit() + self.search.setPlaceholderText("Search runs, datasets, or notes") + self.architecture = QComboBox() + self.architecture.addItem("All architectures", "") + self.dataset = QLineEdit() + self.dataset.setPlaceholderText("Dataset filter") + filters.addWidget(self.search, 2) + filters.addWidget(self.architecture) + filters.addWidget(self.dataset, 1) + root.addLayout(filters) + + body = QHBoxLayout() + self.table = QTableWidget(0, 8) + self.table.setHorizontalHeaderLabels( + ["RUN", "MODEL", "ARCH", "DATASET", "EPOCHS", "LOSS", "TIME", "QUALITY"] + ) + self.table.setSelectionBehavior(QAbstractItemView.SelectRows) + self.table.setSelectionMode(QAbstractItemView.MultiSelection) + self.table.setEditTriggers(QAbstractItemView.NoEditTriggers) + self.table.setSortingEnabled(True) + self.table.verticalHeader().hide() + header_view = self.table.horizontalHeader() + header_view.setSectionResizeMode(0, QHeaderView.ResizeToContents) + header_view.setSectionResizeMode(1, QHeaderView.Stretch) + header_view.setSectionResizeMode(2, QHeaderView.ResizeToContents) + header_view.setSectionResizeMode(3, QHeaderView.Stretch) + for column in range(4, 8): + header_view.setSectionResizeMode(column, QHeaderView.ResizeToContents) + body.addWidget(self.table, 3) + + details = _card() + details.setMinimumWidth(390) + details_layout = QVBoxLayout(details) + details_layout.setContentsMargins(17, 16, 17, 16) + details_layout.setSpacing(8) + details_layout.addWidget(_card_title("RUN DETAILS")) + self.detail_title = QLabel("Select a run") + self.detail_title.setStyleSheet("font-size: 17px; font-weight: 650;") + self.detail = QLabel("Recorded training settings will appear here.") + self.detail.setWordWrap(True) + self.detail.setProperty("muted", True) + self.notes = QPlainTextEdit() + self.notes.setPlaceholderText("Notes") + self.notes.setMaximumHeight(110) + rating_row = QHBoxLayout() + rating_row.addWidget(QLabel("Quality")) + self.quality = QSpinBox() + self.quality.setRange(0, 100) + self.quality.setSuffix(" / 100") + rating_row.addWidget(self.quality) + rating_row.addStretch() + actions = QGridLayout() + self.save_button = QPushButton("Save notes") + self.clone_button = QPushButton("Clone settings") + self.compare_button = QPushButton("Compare selected") + self.output_button = QPushButton("Open output") + self.inspect_button = QPushButton("Inspect model") + self.manual_inspect_button = QPushButton("Choose checkpoint") + self.cancel_inspect_button = QPushButton("Cancel scan") + self.cancel_inspect_button.setEnabled(False) + actions.addWidget(self.save_button, 0, 0) + actions.addWidget(self.clone_button, 0, 1) + actions.addWidget(self.compare_button, 1, 0) + actions.addWidget(self.output_button, 1, 1) + actions.addWidget(self.inspect_button, 2, 0) + actions.addWidget(self.manual_inspect_button, 2, 1) + actions.addWidget(self.cancel_inspect_button, 3, 0, 1, 2) + self.inspect_status = QLabel("No model inspected yet.") + self.inspect_status.setProperty("muted", True) + self.inspect_status.setWordWrap(True) + self.inspect_progress = QProgressBar() + self.inspect_progress.setRange(0, 100) + self.inspect_progress.setValue(0) + details_layout.addWidget(self.detail_title) + details_layout.addWidget(self.detail) + details_layout.addLayout(rating_row) + details_layout.addWidget(self.notes) + details_layout.addLayout(actions) + details_layout.addWidget(self.inspect_status) + details_layout.addWidget(self.inspect_progress) + details_layout.addStretch(1) + body.addWidget(details, 2) + root.addLayout(body, 3) + + self.analysis_tabs = QTabWidget() + self.overview_view = QPlainTextEdit() + self.overview_view.setReadOnly(True) + self.overview_view.setPlaceholderText("Inspect a model to see architecture, parameters, configs, and major components.") + weights_page = QWidget() + weights_layout = QVBoxLayout(weights_page) + weights_layout.setContentsMargins(10, 10, 10, 10) + self.weight_search = QLineEdit() + self.weight_search.setPlaceholderText("Search tensor names") + self.weights_table = QTableWidget(0, 11) + self.weights_table.setHorizontalHeaderLabels( + ["TENSOR", "SHAPE", "DTYPE", "PARAMS", "MIN", "MAX", "MEAN", "STD", "ABS MEAN", "L2", "ZEROS"] + ) + self.weights_table.setSelectionBehavior(QAbstractItemView.SelectRows) + self.weights_table.setEditTriggers(QAbstractItemView.NoEditTriggers) + weights_header = self.weights_table.horizontalHeader() + weights_header.setSectionResizeMode(0, QHeaderView.Stretch) + for column in range(1, 11): + weights_header.setSectionResizeMode(column, QHeaderView.ResizeToContents) + weights_layout.addWidget(self.weight_search) + weights_layout.addWidget(self.weights_table, 1) + self.health_view = QPlainTextEdit() + self.health_view.setReadOnly(True) + self.plot_widget = ModelInspectorPlot() + self.compare_view = QPlainTextEdit() + self.compare_view.setReadOnly(True) + self.compare_view.setPlaceholderText("Select 2 runs and choose Compare selected.") + self.timeline_view = QPlainTextEdit() + self.timeline_view.setReadOnly(True) + self.timeline_view.setPlaceholderText("Checkpoint timeline appears when an inspected output has multiple checkpoints.") + self.analysis_tabs.addTab(self.overview_view, "Overview") + self.analysis_tabs.addTab(weights_page, "Weights") + self.analysis_tabs.addTab(self.health_view, "Health") + self.analysis_tabs.addTab(self.plot_widget, "Plots") + self.analysis_tabs.addTab(self.compare_view, "Compare") + self.analysis_tabs.addTab(self.timeline_view, "Timeline") + root.addWidget(self.analysis_tabs, 2) + + self.search.textChanged.connect(self.refresh) + self.architecture.currentIndexChanged.connect(self.refresh) + self.dataset.textChanged.connect(self.refresh) + self.table.itemSelectionChanged.connect(self._selection_changed) + self.save_button.clicked.connect(self._save_notes) + self.clone_button.clicked.connect(self._clone) + self.compare_button.clicked.connect(self._compare) + self.output_button.clicked.connect(self._open_output) + self.inspect_button.clicked.connect(self._inspect_selected) + self.manual_inspect_button.clicked.connect(self._inspect_manual) + self.cancel_inspect_button.clicked.connect(self._cancel_inspection) + self.weight_search.textChanged.connect(self._populate_weights) + self.refresh() + + def refresh(self) -> None: + current = self.selected_run_id + runs = self.store.list_runs( + self.search.text().strip(), + str(self.architecture.currentData() or ""), + self.dataset.text().strip(), + ) + known = sorted({run.model_architecture for run in self.store.list_runs(limit=1000) if run.model_architecture}) + self.architecture.blockSignals(True) + selected_arch = str(self.architecture.currentData() or "") + self.architecture.clear() + self.architecture.addItem("All architectures", "") + for architecture in known: + self.architecture.addItem(architecture.upper(), architecture) + index = self.architecture.findData(selected_arch) + self.architecture.setCurrentIndex(max(0, index)) + self.architecture.blockSignals(False) + self.table.setRowCount(len(runs)) + self.table.setSortingEnabled(False) + for row, run in enumerate(runs): + values = [ + run.id, + run.model_name, + run.model_architecture.upper(), + run.dataset_name or Path(run.dataset_path).name, + str(run.epochs), + "—" if run.final_loss is None else f"{run.final_loss:.5f}", + self._duration(run.training_time_seconds), + "—" if run.quality_score is None else str(run.quality_score), + ] + for column, value in enumerate(values): + item = QTableWidgetItem(value) + item.setData(Qt.UserRole, run.id) + if column in {0, 2, 4, 5, 6, 7}: + item.setTextAlignment(Qt.AlignCenter) + self.table.setItem(row, column, item) + self.table.setSortingEnabled(True) + if not runs: + self.detail_title.setText("No experiments recorded") + self.detail.setText("Training runs are recorded here when a registered trainer finishes, fails, or is cancelled.") + self.notes.clear() + self.compare_view.clear() + self.inspect_button.setEnabled(False) + self.output_button.setEnabled(False) + if current: + for row in range(self.table.rowCount()): + if self.table.item(row, 0).text() == current: + self.table.selectRow(row) + break + + def _selection_changed(self) -> None: + run = self._selected_run() + if not run: + return + self.selected_run_id = run.id + self.detail_title.setText(run.model_name) + details = [ + f"{run.model_architecture.upper()} · {run.status}", + f"Epochs {run.epochs:,} · Batch {run.batch_size or '—'} · LR {run.learning_rate or '—'}", + f"Resolution {run.resolution or '—'} · Loss {'—' if run.final_loss is None else f'{run.final_loss:.5f}'}", + f"Dataset: {run.dataset_path or '—'}", + f"Output: {run.output_folder or '—'}", + ] + if run.peak_vram_gb: + details.append(f"VRAM at record time: {run.peak_vram_gb:.1f} GB") + if run.preview_images: + details.append(f"Previews recorded: {len(run.preview_images)}") + self.detail.setText("\n".join(details)) + self.notes.setPlainText(run.notes) + self.quality.setValue(int(run.quality_score or 0)) + model_path = self._default_model_path(run) + self.output_button.setEnabled(bool(run.output_folder and Path(run.output_folder).exists())) + self.inspect_button.setEnabled(bool(model_path)) + if model_path: + self.inspect_status.setText(f"Ready to inspect: {model_path}") + + def _selected_ids(self) -> list[str]: + ids = [] + for index in self.table.selectionModel().selectedRows(): + item = self.table.item(index.row(), 0) + if item: + ids.append(item.text()) + return ids + + def _selected_run(self) -> ExperimentRun | None: + ids = self._selected_ids() + return self.store.get(ids[0]) if ids else None + + def _save_notes(self) -> None: + if not self.selected_run_id: + return + quality = self.quality.value() + self.store.update_notes(self.selected_run_id, self.notes.toPlainText().strip(), quality if quality else None) + self.refresh() + + def _clone(self) -> None: + run = self._selected_run() + if run: + self.clone_requested.emit(self.store.clone_request(run.id)) + + def _compare(self) -> None: + ids = self._selected_ids() + if len(ids) < 2: + self.compare_view.setPlainText("Select at least two experiment runs.") + return + runs = [self.store.get(run_id) for run_id in ids[:2]] + paths = [self._default_model_path(run) if run else "" for run in runs] + if len(paths) == 2 and paths[0] and paths[1]: + self._start_comparison(paths[0], paths[1], runs[0], runs[1]) + return + rows = self.store.compare(ids) + lines = ["Compare selected runs", "Run IDs: " + ", ".join(ids), ""] + for row in rows: + field = str(row.get("field", "")).replace("_", " ").title() + marker = "Changed" if row.get("changed") else "Same" + values = " | ".join(self._format_compare_value(field, row.get(run_id, "—")) for run_id in ids) + lines.append(f"{marker:7} {field}: {values}") + self.compare_view.setPlainText("\n".join(lines)) + + def _open_output(self) -> None: + run = self._selected_run() + if run and run.output_folder and Path(run.output_folder).exists(): + QDesktopServices.openUrl(QUrl.fromLocalFile(run.output_folder)) + + def _default_model_path(self, run: ExperimentRun) -> str: + candidates = [path for path in run.checkpoint_paths if path] + candidates.append(run.output_folder) + if run.model_architecture == "flow" and run.output_folder: + candidates.insert(0, run.output_folder) + for value in reversed(candidates): + path = Path(value).expanduser() + if path.exists(): + return str(path) + return "" + + def _cache_key(self, path: str) -> tuple[str, float, int] | None: + target = Path(path).expanduser() + if not target.exists(): + return None + try: + if target.is_file(): + stat = target.stat() + return (str(target.resolve()), stat.st_mtime, stat.st_size) + latest = 0.0 + size = 0 + for item in target.rglob("*"): + if item.is_file(): + stat = item.stat() + latest = max(latest, stat.st_mtime) + size += stat.st_size + return (str(target.resolve()), latest, size) + except OSError: + return None + + def _inspect_selected(self) -> None: + run = self._selected_run() + if not run: + self.inspect_status.setText("Select a run first.") + return + path = self._default_model_path(run) + if not path: + self.inspect_status.setText("No checkpoint or output folder was found for this run.") + return + self._start_inspection(path, run) + + def _inspect_manual(self) -> None: + file_path, _ = QFileDialog.getOpenFileName( + self, + "Choose checkpoint", + str(Path.cwd()), + "Model checkpoints (*.safetensors *.pt *.pth *.bin *.ckpt);;All files (*)", + ) + selected = file_path + if not selected: + selected = QFileDialog.getExistingDirectory(self, "Choose model folder", str(Path.cwd())) + if selected: + self._start_inspection(selected, self._selected_run()) + + def _start_inspection(self, path: str, run: ExperimentRun | None = None) -> None: + if self.inspection_worker and self.inspection_worker.isRunning(): + self.inspect_status.setText("A model scan is already running.") + return + key = self._cache_key(path) + if key and key in self.inspection_cache: + self._show_inspection(self.inspection_cache[key]) + self.inspect_status.setText(f"Loaded cached inspection: {Path(path).name}") + return + self.current_model_path = path + self.inspect_progress.setValue(0) + self.inspect_status.setText(f"Inspecting model: {path}") + self.cancel_inspect_button.setEnabled(True) + self.inspect_button.setEnabled(False) + self.manual_inspect_button.setEnabled(False) + self.inspection_worker = ModelInspectionWorker( + path, + run.model_architecture if run else "", + run.settings if run else {}, + ) + self.inspection_worker.progress_changed.connect(self._inspection_progress) + self.inspection_worker.inspection_ready.connect(lambda summary, cache_key=key: self._inspection_finished(summary, cache_key)) + self.inspection_worker.failed.connect(self._inspection_failed) + self.inspection_worker.finished.connect(self._inspection_worker_done) + self.inspection_worker.start() + + def _cancel_inspection(self) -> None: + if self.inspection_worker and self.inspection_worker.isRunning(): + self.inspection_worker.cancel() + self.inspect_status.setText("Cancelling model scan...") + if self.comparison_worker and self.comparison_worker.isRunning(): + self.comparison_worker.cancel() + self.inspect_status.setText("Cancelling comparison...") + + def _inspection_progress(self, value: int, message: str) -> None: + self.inspect_progress.setValue(value) + self.inspect_status.setText(message) + + def _inspection_finished(self, summary: ModelInspection, cache_key: tuple[str, float, int] | None) -> None: + if cache_key: + self.inspection_cache[cache_key] = summary + self._show_inspection(summary) + self.inspect_status.setText(summary.messages[0] if summary.messages else "Inspection complete.") + self.inspect_progress.setValue(100) + + def _inspection_failed(self, message: str) -> None: + self.inspect_status.setText(message or "Model inspection failed.") + self.health_view.setPlainText(message or "Model inspection failed.") + + def _inspection_worker_done(self) -> None: + self.cancel_inspect_button.setEnabled(False) + self.inspect_button.setEnabled(True) + self.manual_inspect_button.setEnabled(True) + + def _show_inspection(self, summary: ModelInspection) -> None: + self.current_inspection = summary + self.overview_view.setPlainText(self._format_overview(summary)) + self.health_view.setPlainText("\n".join(summary.health)) + self._populate_weights() + self.plot_widget.plot(summary) + self.timeline_view.setPlainText(self._format_timeline(summary)) + self.analysis_tabs.setCurrentWidget(self.overview_view) + + def _populate_weights(self) -> None: + summary = self.current_inspection + query = self.weight_search.text().strip().casefold() + tensors = summary.tensors if summary else [] + if query: + tensors = [tensor for tensor in tensors if query in tensor.name.casefold()] + tensors = sorted(tensors, key=lambda item: item.parameter_count, reverse=True)[:1000] + self.weights_table.setRowCount(len(tensors)) + for row, tensor in enumerate(tensors): + values = [ + tensor.name, + shape_label(tensor.shape), + tensor.dtype, + f"{tensor.parameter_count:,}", + self._metric(tensor.minimum), + self._metric(tensor.maximum), + self._metric(tensor.mean), + self._metric(tensor.std), + self._metric(tensor.abs_mean), + self._metric(tensor.l2_norm), + "-" if tensor.zero_percent is None else f"{tensor.zero_percent:.2f}%", + ] + for column, value in enumerate(values): + item = QTableWidgetItem(value) + if column > 0: + item.setTextAlignment(Qt.AlignRight | Qt.AlignVCenter) + self.weights_table.setItem(row, column, item) + + def _format_overview(self, summary: ModelInspection) -> str: + lines = [ + f"Detected architecture: {summary.architecture} ({summary.confidence:.0%} confidence)", + f"Model/checkpoint path: {summary.resolved_path}", + f"Checkpoint size: {bytes_label(summary.size_bytes)}", + f"Configuration files found: {len(summary.config_files)}", + f"Resolution: {summary.resolution or '-'}", + f"Epoch: {summary.epoch or '-'}", + f"Step: {summary.step or '-'}", + f"Number of tensors: {summary.tensor_count:,}", + f"Total parameter count: {summary.total_parameters:,}", + f"Trainable parameter count: {'-' if summary.trainable_parameters is None else f'{summary.trainable_parameters:,}'}", + f"Parameter memory size: {bytes_label(summary.parameter_memory_bytes)}", + "", + "Tensor data types:", + ] + lines.extend(f" {dtype}: {count:,} parameters" for dtype, count in sorted(summary.dtypes.items())) + lines.extend(["", "Major model components:"]) + lines.extend(f" {name}: {count:,} parameters" for name, count in sorted(summary.components.items(), key=lambda item: item[1], reverse=True)[:20]) + lines.extend(["", "Largest tensors/layers:"]) + lines.extend(f" {tensor.name} {shape_label(tensor.shape)} {tensor.parameter_count:,}" for tensor in summary.largest_tensors[:15]) + if summary.config_files: + lines.extend(["", "Configuration files:"]) + lines.extend(f" {path}" for path in summary.config_files[:20]) + if summary.lora: + lines.extend(["", "LoRA adapter details:"]) + lines.extend(f" {key.replace('_', ' ').title()}: {value}" for key, value in summary.lora.items()) + lines.append(" Tensor norms are adapter statistics; they do not directly tell visual strength.") + if summary.messages: + lines.extend(["", "Messages:"]) + lines.extend(f" {message}" for message in summary.messages) + return "\n".join(lines) + + def _format_timeline(self, summary: ModelInspection) -> str: + if not summary.checkpoints: + return "No timeline checkpoints were found near this model output." + lines = [ + "Checkpoint timeline", + "Select several checkpoints with Compare selected runs, or inspect individual checkpoints from this list.", + "", + ] + for path in summary.checkpoints[:80]: + item = Path(path) + try: + size = bytes_label(sum(file.stat().st_size for file in item.rglob("*") if file.is_file()) if item.is_dir() else item.stat().st_size) + except OSError: + size = "-" + lines.append(f"{item.name} | {size} | {path}") + if len(summary.checkpoints) > 80: + lines.append(f"... {len(summary.checkpoints) - 80} more checkpoints omitted.") + return "\n".join(lines) + + def _start_comparison(self, path_a: str, path_b: str, run_a: ExperimentRun | None, run_b: ExperimentRun | None) -> None: + if self.comparison_worker and self.comparison_worker.isRunning(): + self.compare_view.setPlainText("A model comparison is already running.") + return + self.compare_view.setPlainText("Comparing checkpoints...") + self.analysis_tabs.setCurrentWidget(self.compare_view) + self.cancel_inspect_button.setEnabled(True) + self.comparison_worker = ModelComparisonWorker(path_a, path_b, run_a, run_b) + self.comparison_worker.progress_changed.connect(self._inspection_progress) + self.comparison_worker.comparison_ready.connect(self._comparison_finished) + self.comparison_worker.failed.connect(self._comparison_failed) + self.comparison_worker.finished.connect(self._comparison_worker_done) + self.comparison_worker.start() + + def _comparison_finished(self, comparison: ModelComparison) -> None: + self.compare_view.setPlainText(self._format_model_comparison(comparison)) + self.inspect_progress.setValue(100) + self.inspect_status.setText("Model comparison complete.") + + def _comparison_failed(self, message: str) -> None: + self.compare_view.setPlainText(message or "Model comparison failed.") + self.inspect_status.setText(message or "Model comparison failed.") + + def _comparison_worker_done(self) -> None: + self.cancel_inspect_button.setEnabled(False) + + def _format_model_comparison(self, comparison: ModelComparison) -> str: + lines = [ + "Compare checkpoints", + f"A: {comparison.path_a}", + f"B: {comparison.path_b}", + f"Architecture: {comparison.architecture_a} vs {comparison.architecture_b} ({'match' if comparison.architecture_match else 'mismatch'})", + f"Parameter-count difference: {comparison.parameter_count_difference:+,}", + ] + if comparison.resolution_difference: + lines.append(f"Resolution difference: {comparison.resolution_difference[0]} vs {comparison.resolution_difference[1]}") + lines.extend(["", "Messages:"]) + lines.extend(f" {message}" for message in comparison.messages) + if comparison.config_differences: + lines.extend(["", "Config differences:"]) + lines.extend(f" {item}" for item in comparison.config_differences[:30]) + if comparison.only_a or comparison.only_b or comparison.shape_mismatches: + lines.extend(["", "Tensor availability:"]) + lines.append(f" Tensors only present in A: {len(comparison.only_a)}") + lines.append(f" Tensors only present in B: {len(comparison.only_b)}") + lines.append(f" Tensors with different shapes: {len(comparison.shape_mismatches)}") + if comparison.group_comparisons: + lines.extend(["", "Group summary:"]) + for name, values in sorted(comparison.group_comparisons.items(), key=lambda item: item[1].get("mean_change_score", 0), reverse=True)[:20]: + lines.append( + f" {name}: {int(values['tensors'])} tensors, " + f"change score {values['mean_change_score']:.6f}, " + f"mean abs diff {values['mean_abs_difference']:.6f}" + ) + if comparison.tensor_comparisons: + lines.extend(["", "Most changed tensors:"]) + for item in comparison.tensor_comparisons[:40]: + lines.append( + f" {item.name} | score {self._metric(item.change_score)} | " + f"mean abs {self._metric(item.mean_abs_difference)} | " + f"relative {self._metric(item.relative_difference)} | " + f"cosine {self._metric(item.cosine_similarity)} | L2 {self._metric(item.l2_distance)}" + ) + return "\n".join(lines) + + @staticmethod + def _metric(value: float | None) -> str: + if value is None: + return "-" + if abs(value) >= 1000 or (abs(value) < 0.001 and value != 0): + return f"{value:.4e}" + return f"{value:.6f}" + + @staticmethod + def _duration(seconds: int) -> str: + if seconds >= 3600: + return f"{seconds // 3600}h {(seconds % 3600) // 60}m" + return f"{seconds // 60}m {seconds % 60}s" + + @classmethod + def _format_compare_value(cls, field: str, value) -> str: + if value in (None, ""): + return "—" + if field == "Training Time Seconds": + return cls._duration(int(value)) + if field == "Learning Rate": + try: + return f"{float(value):.7f}" + except (TypeError, ValueError): + return str(value) + if field in {"Final Loss", "Peak Vram Gb"}: + try: + return f"{float(value):.5f}" + except (TypeError, ValueError): + return str(value) + return str(value) + + +class DatasetLabPage(QWidget): + def __init__(self, root_path: Path) -> None: + super().__init__() + self.root_path = root_path + self.store = StudioStore(root_path) + self.current_folder = "" + self.selected_videos: list[str] = [] + self.selected_item_path = "" + self._scan_worker: DatasetScanWorker | None = None + self._transcript_worker: TranscriptExportWorker | None = None + root = QVBoxLayout(self) + root.setContentsMargins(24, 20, 24, 17) + root.setSpacing(12) + root.addWidget( + _page_header( + "Dataset Lab", + "Inspect local datasets, caption coverage, duplicates, dimensions, and transcript exports.", + ) + ) + chooser = QHBoxLayout() + self.folder = QLineEdit() + self.folder.setPlaceholderText("Choose a dataset folder") + browse = QPushButton("Browse") + scan = QPushButton("Scan") + browse.clicked.connect(self._browse_dataset) + scan.clicked.connect(self._scan) + chooser.addWidget(self.folder, 1) + chooser.addWidget(browse) + chooser.addWidget(scan) + root.addLayout(chooser) + + self.summary = QLabel("No dataset scanned yet.") + self.summary.setWordWrap(True) + self.summary.setProperty("muted", True) + root.addWidget(self.summary) + + body = QHBoxLayout() + self.files = QTableWidget(0, 6) + self.files.setHorizontalHeaderLabels(["FILE", "TYPE", "SIZE", "DIMENSIONS", "CAPTION", "DUPLICATE"]) + self.files.setEditTriggers(QAbstractItemView.NoEditTriggers) + self.files.setSelectionBehavior(QAbstractItemView.SelectRows) + self.files.setSelectionMode(QAbstractItemView.SingleSelection) + self.files.verticalHeader().hide() + self.files.horizontalHeader().setSectionResizeMode(0, QHeaderView.Stretch) + for column in range(1, 6): + self.files.horizontalHeader().setSectionResizeMode(column, QHeaderView.ResizeToContents) + self.files.itemSelectionChanged.connect(self._show_selected_file) + body.addWidget(self.files, 3) + + side = QVBoxLayout() + side.setSpacing(12) + inspector = _card() + inspector_layout = QVBoxLayout(inspector) + inspector_layout.setContentsMargins(17, 16, 17, 16) + inspector_layout.setSpacing(8) + inspector_layout.addWidget(_card_title("DATASET ITEM")) + self.preview = QLabel("Select an image or caption") + self.preview.setAlignment(Qt.AlignCenter) + self.preview.setFixedSize(260, 190) + self.preview.setProperty("muted", True) + self.preview.setStyleSheet( + f"background: #050d14; border: 1px solid {COLORS['border_bright']}; border-radius: 8px;" + ) + self.caption = QPlainTextEdit() + self.caption.setPlaceholderText("Caption text appears here.") + self.caption.setMaximumHeight(105) + decision_row = QHBoxLayout() + decision_row.addWidget(QLabel("Decision")) + self.decision = QComboBox() + self.decision.addItem("Uncertain", "unreviewed") + self.decision.addItem("Accept", "keep") + self.decision.addItem("Reject", "reject") + decision_row.addWidget(self.decision) + save_caption = QPushButton("Save caption") + open_file = QPushButton("Open file") + inspector_layout.addWidget(self.preview, 0, Qt.AlignHCenter) + inspector_layout.addWidget(self.caption) + inspector_layout.addLayout(decision_row) + inspector_layout.addWidget(save_caption) + inspector_layout.addWidget(open_file) + side.addWidget(inspector) + + transcript = _card() + transcript.setMinimumWidth(360) + transcript_layout = QVBoxLayout(transcript) + transcript_layout.setContentsMargins(17, 16, 17, 16) + transcript_layout.setSpacing(9) + transcript_layout.addWidget(_card_title("TRANSCRIPT TO DATASET")) + self.video_summary = QLabel("No videos selected.") + self.video_summary.setWordWrap(True) + self.video_summary.setProperty("muted", True) + backend_status = ", ".join( + f"{backend.name}: {'ready' if backend.available else 'missing'}" + for backend in available_transcription_backends() + ) + self.transcript_backend = QLabel(backend_status) + self.transcript_backend.setWordWrap(True) + self.transcript_backend.setProperty("muted", True) + choose_videos = QPushButton("Choose videos") + export = QPushButton("Export TXT / JSONL") + choose_videos.clicked.connect(self._choose_videos) + export.clicked.connect(self._export_transcripts) + self.transcript_status = QPlainTextEdit() + self.transcript_status.setReadOnly(True) + self.transcript_status.setMaximumHeight(230) + transcript_layout.addWidget(self.video_summary) + transcript_layout.addWidget(self.transcript_backend) + transcript_layout.addWidget(choose_videos) + transcript_layout.addWidget(export) + transcript_layout.addWidget(self.transcript_status, 1) + side.addWidget(transcript, 1) + body.addLayout(side, 1) + root.addLayout(body, 1) + save_caption.clicked.connect(self._save_caption) + open_file.clicked.connect(self._open_selected_file) + self.decision.currentIndexChanged.connect(self._decision_changed) + + def _browse_dataset(self) -> None: + selected = QFileDialog.getExistingDirectory(self, "Choose dataset", self.folder.text() or str(self.root_path / "ADAM_Datasets")) + if selected: + self.folder.setText(selected) + self._scan() + + def _scan(self) -> None: + if self._scan_worker and self._scan_worker.isRunning(): + self.summary.setText("Dataset scan is already running.") + return + folder = self.folder.text().strip() + if not folder: + self.summary.setText("Choose a dataset folder first.") + return + self.summary.setText("Scanning dataset...") + self.files.setRowCount(0) + self._scan_worker = DatasetScanWorker(folder) + self._scan_worker.scanned.connect(self._scan_finished) + self._scan_worker.failed.connect(self._scan_failed) + self._scan_worker.finished.connect(self._scan_worker.deleteLater) + self._scan_worker.finished.connect(lambda: setattr(self, "_scan_worker", None)) + self._scan_worker.start() + + def _scan_finished(self, report) -> None: + self.current_folder = report.path + self.store.load() + self.summary.setText( + f"{report.image_count:,} images · {report.video_count:,} videos · {report.caption_count:,} captions · " + f"{report.missing_caption_count:,} missing captions · {report.duplicate_groups:,} duplicate groups\n" + + (" ".join(report.warnings) if report.warnings else "Dataset sample looks ready for review.") + ) + self.files.setRowCount(len(report.items)) + for row, item in enumerate(report.items): + values = [ + Path(item.path).name, + item.kind, + f"{item.size_bytes / 1024:.1f} KB", + f"{item.width}x{item.height}" if item.width and item.height else "—", + "Yes" if item.caption_path else "—", + item.duplicate_key[:8] if item.duplicate_key else "—", + ] + for column, value in enumerate(values): + table_item = QTableWidgetItem(value) + table_item.setToolTip(item.path) + if column: + table_item.setTextAlignment(Qt.AlignCenter) + table_item.setData(Qt.UserRole, item.path) + self.files.setItem(row, column, table_item) + if report.items: + self.files.selectRow(0) + else: + self.preview.setPixmap(QPixmap()) + self.preview.setText("No files to preview") + self.caption.clear() + + def _scan_failed(self, message: str) -> None: + self.summary.setText(f"Dataset scan failed: {message}") + + def _show_selected_file(self) -> None: + rows = self.files.selectionModel().selectedRows() + if not rows: + return + item = self.files.item(rows[0].row(), 0) + path = Path(str(item.data(Qt.UserRole) if item else "")).expanduser() + self.selected_item_path = str(path) + pixmap = QPixmap(str(path)) + if not pixmap.isNull(): + self.preview.setPixmap(pixmap.scaled(self.preview.size(), Qt.KeepAspectRatio, Qt.SmoothTransformation)) + self.preview.setText("") + caption_path = next((path.with_suffix(ext) for ext in (".txt", ".caption") if path.with_suffix(ext).is_file()), path.with_suffix(".txt")) + else: + self.preview.setPixmap(QPixmap()) + self.preview.setText(path.name or "No preview") + caption_path = path if path.suffix.casefold() in {".txt", ".caption"} else path.with_suffix(".txt") + try: + text = caption_path.read_text(encoding="utf-8") if caption_path.is_file() else "" + except OSError as exc: + text = f"Caption could not be read: {exc}" + self.caption.setPlainText(text) + self.caption.setProperty("caption_path", str(caption_path)) + decision = "unreviewed" + if self.current_folder and pixmap.isNull() is False: + decision = self.store.review(self.current_folder).decisions.get(str(path.resolve()), "unreviewed") + index = self.decision.findData(decision) + self.decision.blockSignals(True) + self.decision.setCurrentIndex(max(0, index)) + self.decision.blockSignals(False) + + def _save_caption(self) -> None: + raw_path = str(self.caption.property("caption_path") or "") + if not raw_path: + return + caption_path = Path(raw_path).expanduser() + try: + caption_path.write_text(self.caption.toPlainText().strip() + "\n", encoding="utf-8") + self.summary.setText(f"Saved caption: {caption_path.name}") + except OSError as exc: + QMessageBox.warning(self, "Caption not saved", str(exc)) + + def _decision_changed(self) -> None: + if not self.current_folder or not self.selected_item_path: + return + path = Path(self.selected_item_path) + if path.suffix.casefold() not in {".jpg", ".jpeg", ".png", ".webp", ".bmp"}: + return + try: + self.store.set_decision(self.current_folder, self.selected_item_path, str(self.decision.currentData())) + self.summary.setText(f"Saved review decision for {path.name}.") + except (OSError, ValueError) as exc: + QMessageBox.warning(self, "Decision not saved", str(exc)) + + def _open_selected_file(self) -> None: + if self.selected_item_path and Path(self.selected_item_path).exists(): + QDesktopServices.openUrl(QUrl.fromLocalFile(self.selected_item_path)) + + def _choose_videos(self) -> None: + files, _ = QFileDialog.getOpenFileNames( + self, + "Choose local videos", + str(self.root_path), + "Videos (*.mp4 *.mov *.mkv *.webm *.avi)", + ) + self.selected_videos = files + self.video_summary.setText(f"{len(files)} video(s) selected." if files else "No videos selected.") + + def _export_transcripts(self) -> None: + if self._transcript_worker and self._transcript_worker.isRunning(): + self.transcript_status.setPlainText("Transcript export is already running.") + return + output = QFileDialog.getExistingDirectory( + self, + "Choose transcript dataset output", + self.current_folder or str(self.root_path / "data" / "transcript_datasets"), + ) + if not output: + return + self.transcript_status.setPlainText("Preparing transcript dataset...") + self._transcript_worker = TranscriptExportWorker(self.selected_videos, output) + self._transcript_worker.completed.connect(self._transcript_finished) + self._transcript_worker.finished.connect(self._transcript_worker.deleteLater) + self._transcript_worker.finished.connect(lambda: setattr(self, "_transcript_worker", None)) + self._transcript_worker.start() + + def _transcript_finished(self, result) -> None: + self.transcript_status.setPlainText( + result.message + + (f"\nOutput: {result.output_folder}" if result.output_folder else "") + ) + + +class RemoteAccessPage(QWidget): + def __init__(self, service: RemoteAccessService) -> None: + super().__init__() + self.service = service + root = QVBoxLayout(self) + root.setContentsMargins(24, 20, 24, 17) + root.setSpacing(12) + root.addWidget( + _page_header( + "Remote Access", + "Optional authenticated local-network status API for a browser or phone. Disabled by default.", + ) + ) + card = _card() + layout = QGridLayout(card) + layout.setContentsMargins(18, 17, 18, 17) + layout.setHorizontalSpacing(12) + layout.setVerticalSpacing(10) + self.enabled = QCheckBox("Enable remote access service") + self.remote_mode = QComboBox() + self.remote_mode.addItem("Local / Wi-Fi Only", REMOTE_MODE_LOCAL) + self.remote_mode.addItem("Private Tailscale", REMOTE_MODE_TAILSCALE) + self.remote_mode.addItem("Disabled", REMOTE_MODE_DISABLED) + self.bind = QLineEdit() + self.port = QSpinBox() + self.port.setRange(1024, 65535) + self.token = QLineEdit() + self.token.setEchoMode(QLineEdit.PasswordEchoOnEdit) + self.allow_job_control = QCheckBox("Allow remote job controls and approval changes") + self.allow_job_control.setToolTip("Allows approval, pause, resume, stop and retry from devices with your access token.") + self.auto_approve_training = QCheckBox("Auto-approve remote training prompts") + self.status = QLabel() + self.status.setWordWrap(True) + self.status.setProperty("muted", True) + self.qr_code = QLabel("Use phone access to generate a QR code.") + self.qr_code.setAlignment(Qt.AlignCenter) + self.qr_code.setMinimumSize(220, 220) + self.qr_code.setStyleSheet( + f"background: #ffffff; color: #00101b; border: 1px solid {COLORS['border_bright']}; border-radius: 8px;" + ) + self.qr_caption = QLabel() + self.qr_caption.setWordWrap(True) + self.qr_caption.setProperty("muted", True) + self.tailscale_status = QLabel() + self.tailscale_status.setWordWrap(True) + self.tailscale_status.setProperty("muted", True) + save = QPushButton("Save") + toggle = QPushButton("Start / Stop") + open_url = QPushButton("Open local test URL") + phone_access = QPushButton("Use phone access") + copy_phone_url = QPushButton("Copy phone URL") + start_tailscale = QPushButton("Start private Tailscale") + stop_tailscale = QPushButton("Stop private Tailscale") + regenerate = QPushButton("New token") + layout.addWidget(self.enabled, 0, 0, 1, 2) + layout.addWidget(QLabel("Remote mode"), 1, 0) + layout.addWidget(self.remote_mode, 1, 1) + layout.addWidget(QLabel("Bind address"), 2, 0) + layout.addWidget(self.bind, 2, 1) + layout.addWidget(QLabel("Port"), 3, 0) + layout.addWidget(self.port, 3, 1) + layout.addWidget(QLabel("Token"), 4, 0) + layout.addWidget(self.token, 4, 1) + layout.addWidget(self.allow_job_control, 5, 0, 1, 2) + layout.addWidget(self.auto_approve_training, 6, 0, 1, 2) + layout.addWidget(save, 7, 0) + layout.addWidget(toggle, 7, 1) + layout.addWidget(regenerate, 8, 0) + layout.addWidget(open_url, 8, 1) + layout.addWidget(phone_access, 9, 0) + layout.addWidget(copy_phone_url, 9, 1) + layout.addWidget(start_tailscale, 10, 0) + layout.addWidget(stop_tailscale, 10, 1) + layout.addWidget(self.tailscale_status, 11, 0, 1, 2) + layout.addWidget(self.qr_code, 12, 0, 1, 2) + layout.addWidget(self.qr_caption, 13, 0, 1, 2) + layout.addWidget(self.status, 14, 0, 1, 2) + root.addWidget(card) + root.addStretch() + save.clicked.connect(self._save) + toggle.clicked.connect(self._toggle) + open_url.clicked.connect(self._open_status_url) + phone_access.clicked.connect(self._enable_phone_access) + copy_phone_url.clicked.connect(self._copy_phone_url) + start_tailscale.clicked.connect(self._start_tailscale) + stop_tailscale.clicked.connect(self._stop_tailscale) + regenerate.clicked.connect(self._regenerate) + self.refresh() + + def refresh(self) -> None: + settings = self.service.settings() + self.enabled.setChecked(bool(settings["enabled"])) + mode_index = self.remote_mode.findData(settings.get("remote_mode", REMOTE_MODE_LOCAL)) + self.remote_mode.setCurrentIndex(max(0, mode_index)) + self.bind.setText(str(settings["bind_address"])) + self.port.setValue(int(settings["port"])) + self.token.setText(str(settings["token"])) + self.allow_job_control.setChecked(bool(settings["allow_job_control"])) + self.auto_approve_training.setChecked(bool(settings.get("auto_approve_training", False))) + scope = remote_scope(str(settings["bind_address"])) + phone_url = self.service.phone_test_url() + tailscale = self.service.tailscale_status() + tailscale_url = self.service.tailscale_url() + mode = str(settings.get("remote_mode", REMOTE_MODE_LOCAL)) + self.bind.setEnabled(mode != REMOTE_MODE_TAILSCALE) + tailscale_lines = [ + f"Tailscale: {'installed' if tailscale.installed else 'not installed'}", + f"Status: {'connected' if tailscale.connected else 'disconnected'}", + ] + if tailscale.device_name: + tailscale_lines.append(f"Device: {tailscale.device_name}") + if tailscale.tailscale_ip: + tailscale_lines.append(f"Private IP: {tailscale.tailscale_ip}") + if tailscale_url: + tailscale_lines.append(f"Private URL: {tailscale_url}") + elif mode == REMOTE_MODE_TAILSCALE: + tailscale_lines.append("Private URL appears after Tailscale is installed, connected, and Serve is started.") + tailscale_lines.append(tailscale.message) + self.tailscale_status.setText("\n".join(tailscale_lines)) + phone_hint = ( + f" Phone URL: {phone_url}" + if phone_url + else " For the phone dashboard and QR code, use the button below while your phone is on the same Wi-Fi." + ) + self.status.setText( + ( + f"Running at {self.service.url()} · Scope: {scope}. " + "Use Open local test URL on this computer, scan the QR code on your phone, or send the token as a Bearer token from another device. " + "Dangerous actions are unavailable remotely." + + phone_hint + ) + if self.service.running + else ( + f"Stopped · Scope when started: {scope}. Keep 127.0.0.1 for this device only. " + "Remote clients can view status, system usage, and queue state." + + phone_hint + ) + ) + self._refresh_qr(phone_url) + + def _save(self) -> bool: + bind = self.bind.text().strip() or "127.0.0.1" + was_running = self.service.running + mode = str(self.remote_mode.currentData() or REMOTE_MODE_LOCAL) + enabled = self.enabled.isChecked() and mode != REMOTE_MODE_DISABLED + if mode == REMOTE_MODE_TAILSCALE: + bind = "127.0.0.1" + if enabled and mode == REMOTE_MODE_LOCAL and bind not in {"127.0.0.1", "localhost"}: + answer = QMessageBox.question( + self, + "Enable local-network access", + "This can expose ADAM status to other devices on your network. Continue only on a trusted network.", + QMessageBox.Yes | QMessageBox.No, + QMessageBox.No, + ) + if answer != QMessageBox.Yes: + return False + self.service.save_settings( + { + "enabled": enabled, + "remote_mode": mode, + "bind_address": bind, + "port": self.port.value(), + "token": self.token.text().strip(), + "allow_job_control": self.allow_job_control.isChecked(), + "auto_approve_training": self.auto_approve_training.isChecked(), + } + ) + if was_running: + self.service.stop() + if enabled: + try: + self.service.start() + except (OSError, RuntimeError) as exc: + self.status.setText(f"Remote access settings were saved, but restart failed: {exc}") + return False + self.refresh() + return True + + def _toggle(self) -> None: + if self.service.running: + self.service.stop() + self.refresh() + return + self.enabled.setChecked(True) + if not self._save(): + return + try: + self.service.start() + self.refresh() + except (OSError, RuntimeError) as exc: + self.status.setText(f"Remote access could not start: {exc}") + + def _open_status_url(self) -> None: + QDesktopServices.openUrl(QUrl(self.service.local_test_url())) + + def _enable_phone_access(self) -> None: + self.enabled.setChecked(True) + self.remote_mode.setCurrentIndex(max(0, self.remote_mode.findData(REMOTE_MODE_LOCAL))) + self.bind.setText("0.0.0.0") + if self._save() and not self.service.running: + try: + self.service.start() + except (OSError, RuntimeError) as exc: + self.status.setText(f"Remote access could not start: {exc}") + return + self.refresh() + + def _copy_phone_url(self) -> None: + url = self.service.phone_test_url() + if not url: + self.status.setText("No phone URL yet. Choose Use phone access, then keep your phone on the same Wi-Fi.") + return + QApplication.clipboard().setText(url) + self.status.setText(f"Copied phone URL: {url}") + + def _start_tailscale(self) -> None: + self.enabled.setChecked(True) + self.remote_mode.setCurrentIndex(max(0, self.remote_mode.findData(REMOTE_MODE_TAILSCALE))) + self.bind.setText("127.0.0.1") + if not self._save(): + return + if not self.service.running: + try: + self.service.start() + except (OSError, RuntimeError) as exc: + self.status.setText(f"ADAM Remote could not start for Tailscale: {exc}") + return + ok, message = self.service.start_tailscale_serve() + self.status.setText(message) + self.refresh() + + def _stop_tailscale(self) -> None: + answer = QMessageBox.question( + self, + "Stop Tailscale Serve", + "This resets Tailscale Serve forwarding on this PC. ADAM will keep running locally. Continue?", + QMessageBox.Yes | QMessageBox.No, + QMessageBox.No, + ) + if answer != QMessageBox.Yes: + return + ok, message = self.service.stop_tailscale_serve() + self.status.setText(message) + self.refresh() + + def _refresh_qr(self, url: str) -> None: + if not url: + self.qr_code.setPixmap(QPixmap()) + self.qr_code.setText("Use phone access to generate a QR code.") + self.qr_caption.setText("The QR code appears here after ADAM is available to devices on your Wi-Fi.") + return + pixmap = self._qr_pixmap(url) + if pixmap is None: + self.qr_code.setPixmap(QPixmap()) + self.qr_code.setText("QR package not installed") + self.qr_caption.setText("Install ADAM requirements, then restart the app to generate phone QR codes.") + return + self.qr_code.setText("") + self.qr_code.setPixmap(pixmap) + self.qr_caption.setText("Scan this code to open the ADAM mobile dashboard on your phone.") + + @staticmethod + def _qr_pixmap(url: str) -> QPixmap | None: + try: + import qrcode + except ImportError: + return None + image = qrcode.make(url).convert("RGB") + buffer = BytesIO() + image.save(buffer, format="PNG") + pixmap = QPixmap() + if not pixmap.loadFromData(buffer.getvalue(), "PNG"): + return None + return pixmap.scaled(220, 220, Qt.KeepAspectRatio, Qt.FastTransformation) + + def _regenerate(self) -> None: + from adam.remote_access import default_remote_settings + + self.token.setText(str(default_remote_settings()["token"])) + self._save() + + class ToolsPage(QWidget): setup_requested = Signal() @@ -3998,6 +5783,19 @@ class SettingsPage(QWidget): class MainWindow(QMainWindow): + PAGE_COMMAND = 0 + PAGE_CHAT_HISTORY = 1 + PAGE_STUDIO = 2 + PAGE_DATASET_LAB = 3 + PAGE_EXPERIMENTS = 4 + PAGE_GENERATIONS = 5 + PAGE_SHOWCASE = 6 + PAGE_JOBS = 7 + PAGE_TOOLS = 8 + PAGE_SYSTEM = 9 + PAGE_REMOTE = 10 + PAGE_SETTINGS = 11 + def __init__( self, root_path: Path, @@ -4013,6 +5811,7 @@ class MainWindow(QMainWindow): self.jobs = jobs self.config = config self.monitor = monitor + self.remote_service = RemoteAccessService(config, jobs, monitor, planner) self.setWindowTitle("ADAM — AI Development and Automation Manager") self.resize(1480, 900) self.setMinimumSize(1120, 760) @@ -4052,6 +5851,8 @@ class MainWindow(QMainWindow): self.command_scroll.setWidget(self.command_page) self.jobs_page = JobsPage(jobs) self.studio_page = StudioPage(root_path, jobs, jobs.assets, config) + self.dataset_lab_page = DatasetLabPage(root_path) + self.experiment_page = ExperimentTrackerPage(jobs.experiments) self.generations_page = GenerationsPage( root_path, registry, jobs, jobs.assets, config ) @@ -4060,18 +5861,22 @@ class MainWindow(QMainWindow): ) self.tools_page = ToolsPage(registry, tool_folders) self.system_page = SystemPage() + self.remote_page = RemoteAccessPage(self.remote_service) self.settings_page = SettingsPage(config, tool_folders) self.chat_history_page = ChatHistoryPage(self.command_page.history_store) for page in ( self.command_scroll, + self.chat_history_page, self.studio_page, + self.dataset_lab_page, + self.experiment_page, self.generations_page, self.showcase_page, self.jobs_page, self.tools_page, self.system_page, + self.remote_page, self.settings_page, - self.chat_history_page, ): self.stack.addWidget(page) content_layout.addWidget(self.stack, 1) @@ -4081,24 +5886,35 @@ class MainWindow(QMainWindow): self.settings_page.saved.connect(self.command_page.refresh_provider_badge) self.settings_page.saved.connect(self.tools_page.reload) - self.tools_page.setup_requested.connect(lambda: self._switch_page(7)) + self.tools_page.setup_requested.connect(lambda: self._switch_page(self.PAGE_SETTINGS)) self.command_page.tool_folders_changed.connect( self.settings_page.refresh_tool_folders ) self.command_page.tool_folders_changed.connect(self.tools_page.reload) - self.command_page.open_jobs_requested.connect(lambda: self._switch_page(4)) + self.command_page.open_jobs_requested.connect(lambda: self._switch_page(self.PAGE_JOBS)) + self.command_page.open_dataset_lab_requested.connect(lambda: self._switch_page(self.PAGE_DATASET_LAB)) + self.command_page.open_experiments_requested.connect(lambda: self._switch_page(self.PAGE_EXPERIMENTS)) + self.command_page.open_remote_requested.connect(lambda: self._switch_page(self.PAGE_REMOTE)) self.command_page.history_changed.connect(self.chat_history_page.refresh) self.chat_history_page.open_requested.connect(self._open_saved_conversation) + self.experiment_page.clone_requested.connect(self._clone_experiment_request) self.jobs.active_changed.connect(self.system_page.set_active_job) self.jobs.job_updated.connect(self._update_system_job) + self.experiment_refresh_timer = QTimer(self) + self.experiment_refresh_timer.setInterval(1000) + self.experiment_refresh_timer.setSingleShot(True) + self.experiment_refresh_timer.timeout.connect(self._refresh_visible_experiments) + self.jobs.job_updated.connect(self._schedule_experiment_refresh) self.jobs.notification.connect(self._show_notification) self.studio_page.plan_requested.connect(self._plan_from_studio) - self._switch_page(0) + self._switch_page(self.PAGE_COMMAND) self.monitor_timer = QTimer(self) self.monitor_timer.timeout.connect(self._refresh_monitor) self.monitor_timer.start(1500) self._refresh_monitor() + if self.remote_service.settings().get("enabled"): + QTimer.singleShot(500, self._start_saved_remote_access) if any(job.status == JobStatus.INTERRUPTED for job in self.jobs.jobs): QTimer.singleShot(350, self._offer_recovery) @@ -4106,13 +5922,23 @@ class MainWindow(QMainWindow): sidebar = QFrame() sidebar.setObjectName("Sidebar") sidebar.setFixedWidth(230) - layout = QVBoxLayout(sidebar) + outer = QVBoxLayout(sidebar) + outer.setContentsMargins(0, 0, 0, 0) + scroll = QScrollArea() + scroll.setObjectName("SidebarScroll") + scroll.setWidgetResizable(True) + scroll.setHorizontalScrollBarPolicy(Qt.ScrollBarAlwaysOff) + content = QWidget() + layout = QVBoxLayout(content) + layout.setSizeConstraint(QVBoxLayout.SetMinimumSize) + scroll.setWidget(content) + outer.addWidget(scroll) layout.setContentsMargins(0, 20, 0, 16) layout.setSpacing(3) brand = QWidget() brand_layout = QHBoxLayout(brand) - brand_layout.setContentsMargins(16, 0, 12, 18) + brand_layout.setContentsMargins(0, 0, 0, 18) brand_layout.setSpacing(7) logo = QLabel() pixmap = QPixmap(str(self.root_path / "assets" / "adam_atom.png")) @@ -4143,15 +5969,18 @@ class MainWindow(QMainWindow): ) layout.addWidget(section) nav_items = [ - ("COMMAND CENTER", 0), - ("CHAT HISTORY", 8), - ("TRAINING STUDIO", 1), - ("GENERATIONS", 2), - ("SHOWCASE VIDEO", 3), - ("JOBS / HISTORY", 4), - ("TOOL REGISTRY", 5), - ("SYSTEM MONITOR", 6), - ("SETTINGS", 7), + ("COMMAND CENTER", self.PAGE_COMMAND), + ("CHAT HISTORY", self.PAGE_CHAT_HISTORY), + ("TRAINING STUDIO", self.PAGE_STUDIO), + ("DATASET LAB", self.PAGE_DATASET_LAB), + ("EXPERIMENT TRACKER", self.PAGE_EXPERIMENTS), + ("GENERATIONS", self.PAGE_GENERATIONS), + ("SHOWCASE VIDEO", self.PAGE_SHOWCASE), + ("JOBS / HISTORY", self.PAGE_JOBS), + ("TOOL REGISTRY", self.PAGE_TOOLS), + ("SYSTEM MONITOR", self.PAGE_SYSTEM), + ("REMOTE ACCESS", self.PAGE_REMOTE), + ("SETTINGS", self.PAGE_SETTINGS), ] self.nav_buttons: list[QPushButton] = [] for text, index in nav_items: @@ -4173,6 +6002,8 @@ class MainWindow(QMainWindow): ("◉", "LoRA Trainer", "lora"), ("◎", "DDPM Trainer", "ddpm"), ("⌁", "Flow Matching", "flow"), + ("▤", "Dataset Lab", "dataset_lab"), + ("▥", "Compare Runs", "experiments"), ("▧", "Image Generator", "generations"), ("▣", "Video Generator", "video"), ): @@ -4189,11 +6020,13 @@ class MainWindow(QMainWindow): safety_layout = QVBoxLayout(safety) safety_layout.setContentsMargins(12, 11, 12, 11) safety_layout.setSpacing(4) - safe_title = QLabel("● SAFE MODE") + safe_title = QLabel("● APPROVAL SETTINGS") safe_title.setStyleSheet( f"color: {COLORS['green']}; font-size: 10px; font-weight: 700;" ) - safe_body = QLabel("Approval gates on\nFull action logging") + safe_body = QLabel() + self.safety_summary = safe_body + safe_body.setWordWrap(True) safe_body.setProperty("muted", True) safe_body.setStyleSheet("font-size: 11px;") safety_layout.addWidget(safe_title) @@ -4207,6 +6040,14 @@ class MainWindow(QMainWindow): layout.setContentsMargins(12, 20, 12, 16) return sidebar + def _schedule_experiment_refresh(self, _job: Job) -> None: + if self.experiment_page.isVisible() and not self.experiment_refresh_timer.isActive(): + self.experiment_refresh_timer.start() + + def _refresh_visible_experiments(self) -> None: + if self.experiment_page.isVisible(): + self.experiment_page.refresh() + def _build_status_bar(self) -> QFrame: bar = QFrame() bar.setObjectName("TopBar") @@ -4242,14 +6083,20 @@ class MainWindow(QMainWindow): if not hasattr(self, "stack"): return self.stack.setCurrentIndex(index) - if index == 1: - self.studio_page.refresh() - elif index == 2: - self.generations_page.refresh() - elif index == 3: - self.showcase_page.refresh() - if index == 8: - self.chat_history_page.refresh() + if index == self.PAGE_STUDIO: + QTimer.singleShot(0, self.studio_page.refresh) + elif index == self.PAGE_DATASET_LAB: + self.dataset_lab_page.folder.setFocus() + elif index == self.PAGE_EXPERIMENTS: + QTimer.singleShot(0, self.experiment_page.refresh) + elif index == self.PAGE_GENERATIONS: + QTimer.singleShot(0, self.generations_page.refresh) + elif index == self.PAGE_SHOWCASE: + QTimer.singleShot(0, self.showcase_page.refresh) + elif index == self.PAGE_REMOTE: + QTimer.singleShot(0, self.remote_page.refresh) + if index == self.PAGE_CHAT_HISTORY: + QTimer.singleShot(0, self.chat_history_page.refresh) for button in self.nav_buttons: button.setProperty("navActive", button.property("pageIndex") == index) button.style().unpolish(button) @@ -4257,16 +6104,28 @@ class MainWindow(QMainWindow): def _open_saved_conversation(self, conversation: dict) -> None: self.command_page.open_conversation(conversation) - self._switch_page(0) + self._switch_page(self.PAGE_COMMAND) + + def _clone_experiment_request(self, request: str) -> None: + if not request: + return + self._switch_page(self.PAGE_COMMAND) + self.command_page.submit(request) def _quick_access(self, action: str) -> None: if action == "generations": - self._switch_page(2) + self._switch_page(self.PAGE_GENERATIONS) return if action == "video": - self._switch_page(3) + self._switch_page(self.PAGE_SHOWCASE) + return + if action == "dataset_lab": + self._switch_page(self.PAGE_DATASET_LAB) + return + if action == "experiments": + self._switch_page(self.PAGE_EXPERIMENTS) return - self._switch_page(0) + self._switch_page(self.PAGE_COMMAND) if action == "dataset": self.command_page.submit("Adam, collect a dataset") elif action == "lora": @@ -4277,10 +6136,14 @@ class MainWindow(QMainWindow): self.command_page.submit("Adam, train a Flow Matching model") def _plan_from_studio(self, request: str) -> None: - self._switch_page(0) + self._switch_page(self.PAGE_COMMAND) self.command_page.submit(request) def _refresh_monitor(self) -> None: + remote = self.remote_service.settings() + approval = "Remote auto-approval on" if remote.get("auto_approve_training") else "Remote approval required" + local = "Long-task prompts on" if self.config.get("ask_before_long_tasks", True) else "Long-task prompts off" + self.safety_summary.setText(f"{local}\n{approval}") snapshot = self.monitor.snapshot() self.jobs.supervise(snapshot) self.system_page.update_snapshot(snapshot) @@ -4319,6 +6182,13 @@ class MainWindow(QMainWindow): if self.config.get("sound_notifications"): QApplication.beep() + def _start_saved_remote_access(self) -> None: + try: + self.remote_service.start() + self.remote_page.refresh() + except OSError as exc: + self.remote_page.status.setText(f"Remote access is enabled, but it could not start: {exc}") + def _offer_recovery(self) -> None: interrupted = [ job for job in self.jobs.jobs if job.status == JobStatus.INTERRUPTED @@ -4339,5 +6209,6 @@ class MainWindow(QMainWindow): def closeEvent(self, event: QCloseEvent) -> None: if self.tray_icon: self.tray_icon.hide() + self.remote_service.stop() self.studio_page.shutdown() super().closeEvent(event) diff --git a/adam/ui/settings_ui.py b/adam/ui/settings_ui.py new file mode 100644 index 0000000000000000000000000000000000000000..799d987d861ef64c6bcc6642ae5c4912bf7ed65d --- /dev/null +++ b/adam/ui/settings_ui.py @@ -0,0 +1,281 @@ +from __future__ import annotations + +from pathlib import Path +from typing import Any + +from PySide6.QtWidgets import ( + QCheckBox, + QComboBox, + QDoubleSpinBox, + QFileDialog, + QFrame, + QGridLayout, + QGroupBox, + QHBoxLayout, + QLabel, + QLineEdit, + QPlainTextEdit, + QPushButton, + QScrollArea, + QSlider, + QSpinBox, + QVBoxLayout, + QWidget, +) +from PySide6.QtCore import Qt, Signal + + +class SettingsForm(QWidget): + """Builds editable PySide controls from a model plugin setting schema.""" + + changed = Signal() + + def __init__( + self, + schema: dict[str, dict[str, Any]] | None = None, + parent: QWidget | None = None, + ) -> None: + super().__init__(parent) + self.schema: dict[str, dict[str, Any]] = {} + self.widgets: dict[str, QWidget] = {} + self.rows: dict[str, tuple[QLabel, QWidget]] = {} + self.show_advanced = QCheckBox("Show advanced settings") + self.show_advanced.toggled.connect(self._advanced_changed) + + layout = QVBoxLayout(self) + layout.setContentsMargins(0, 4, 0, 0) + layout.setSpacing(10) + layout.addWidget(self.show_advanced) + self.scroll = QScrollArea() + self.scroll.setWidgetResizable(True) + self.scroll.setFrameShape(QFrame.NoFrame) + self.scroll.setMinimumHeight(180) + self.scroll.setMaximumHeight(390) + self.content = QWidget() + self.content_layout = QVBoxLayout(self.content) + self.content_layout.setContentsMargins(4, 4, 4, 4) + self.content_layout.setSpacing(10) + self.scroll.setWidget(self.content) + layout.addWidget(self.scroll) + if schema: + self.set_schema(schema) + + def set_schema(self, schema: dict[str, dict[str, Any]]) -> None: + self.schema = {str(key): dict(value) for key, value in schema.items()} + self.widgets.clear() + self.rows.clear() + while self.content_layout.count(): + item = self.content_layout.takeAt(0) + if item.widget(): + item.widget().deleteLater() + + grouped: dict[str, list[tuple[str, dict[str, Any]]]] = {} + for key, spec in self.schema.items(): + grouped.setdefault(str(spec.get("group", "Basic")), []).append((key, spec)) + has_advanced = any(bool(spec.get("advanced")) for spec in self.schema.values()) + self.show_advanced.setVisible(has_advanced) + + for group, items in grouped.items(): + box = QGroupBox() + grid = QGridLayout(box) + grid.setContentsMargins(12, 12, 12, 12) + grid.setHorizontalSpacing(12) + grid.setVerticalSpacing(8) + grid.setColumnStretch(0, 0) + grid.setColumnStretch(1, 1) + group_title = QLabel(group) + group_title.setProperty("sectionTitle", True) + grid.addWidget(group_title, 0, 0, 1, 2) + for row, (key, spec) in enumerate(items): + required = " *" if spec.get("required") else "" + label = QLabel(str(spec.get("label", key)) + required) + label.setMinimumWidth(128) + label.setMaximumWidth(210) + label.setWordWrap(True) + widget = self._build_widget(key, spec) + tooltip = self._tooltip_for(key, spec) + if tooltip: + label.setToolTip(tooltip) + widget.setToolTip(tooltip) + if str(spec.get("type", "text")) not in {"path", "folder", "multiline_text"}: + widget.setMaximumWidth(420) + grid.addWidget(label, row + 1, 0, Qt.AlignLeft | Qt.AlignVCenter) + grid.addWidget(widget, row + 1, 1, Qt.AlignLeft | Qt.AlignVCenter) + self.widgets[key] = widget + self.rows[key] = (label, widget) + self.content_layout.addWidget(box) + self.content_layout.addStretch() + self._advanced_changed(self.show_advanced.isChecked()) + + def _build_widget(self, key: str, spec: dict[str, Any]) -> QWidget: + kind = str(spec.get("type", "text")) + default = spec.get("default", "") + if kind == "bool": + widget = QCheckBox() + widget.setChecked(bool(default)) + widget.setMaximumWidth(420) + widget.toggled.connect(self.changed) + return widget + if kind == "choice": + widget = QComboBox() + for option in spec.get("options", []): + widget.addItem(str(option), option) + index = widget.findData(default) + if index < 0: + index = widget.findText(str(default)) + widget.setCurrentIndex(max(0, index)) + widget.currentIndexChanged.connect(self.changed) + widget.setMinimumWidth(170) + return widget + if kind in {"int", "slider"}: + if kind == "slider": + container = QWidget() + row = QHBoxLayout(container) + row.setContentsMargins(0, 0, 0, 0) + slider = QSlider(Qt.Horizontal) + slider.setRange(int(spec.get("min", 0)), int(spec.get("max", 100))) + value = QSpinBox() + value.setRange(slider.minimum(), slider.maximum()) + slider.setValue(int(default or slider.minimum())) + value.setValue(slider.value()) + slider.valueChanged.connect(value.setValue) + value.valueChanged.connect(slider.setValue) + slider.valueChanged.connect(self.changed) + row.addWidget(slider, 1) + row.addWidget(value) + container.setProperty("value_widget", value) + container.setMaximumWidth(420) + return container + widget = QSpinBox() + widget.setRange(int(spec.get("min", -2_147_000_000)), int(spec.get("max", 2_147_000_000))) + widget.setValue(int(default or 0)) + widget.valueChanged.connect(self.changed) + widget.setMinimumWidth(120) + return widget + if kind == "float": + widget = QDoubleSpinBox() + widget.setDecimals(int(spec.get("decimals", 7))) + widget.setRange(float(spec.get("min", -1_000_000.0)), float(spec.get("max", 1_000_000.0))) + widget.setSingleStep(float(spec.get("step", 0.0001))) + widget.setValue(float(default or 0.0)) + widget.valueChanged.connect(self.changed) + widget.setMinimumWidth(140) + return widget + if kind == "multiline_text": + widget = QPlainTextEdit() + widget.setMaximumHeight(int(spec.get("height", 76))) + widget.setPlainText(str(default or "")) + widget.textChanged.connect(self.changed) + widget.setMinimumWidth(260) + return widget + if kind in {"path", "folder"}: + container = QWidget() + row = QHBoxLayout(container) + row.setContentsMargins(0, 0, 0, 0) + edit = QLineEdit(str(default or "")) + button = QPushButton("Browse") + button.setFixedWidth(82) + button.clicked.connect(lambda _checked=False, e=edit, k=kind: self._browse(e, k)) + edit.editingFinished.connect(self.changed) + row.addWidget(edit, 1) + row.addWidget(button) + container.setProperty("value_widget", edit) + container.setMinimumWidth(260) + return container + widget = QLineEdit(str(default or "")) + widget.editingFinished.connect(self.changed) + widget.setMinimumWidth(180) + return widget + + @staticmethod + def _tooltip_for(key: str, spec: dict[str, Any]) -> str: + parts = [] + description = str(spec.get("description", "")).strip() + if description: + parts.append(description) + if spec.get("required"): + parts.append("Required.") + if "min" in spec or "max" in spec: + low = spec.get("min", "—") + high = spec.get("max", "—") + parts.append(f"Range: {low} to {high}.") + if str(spec.get("type", "")) == "choice" and spec.get("options"): + parts.append("Choices: " + ", ".join(str(option) for option in spec.get("options", [])) + ".") + if spec.get("advanced"): + parts.append("Advanced setting.") + if not parts: + parts.append(key.replace("_", " ").title()) + return "\n".join(parts) + + def _browse(self, edit: QLineEdit, kind: str) -> None: + current = Path(edit.text()).expanduser() if edit.text().strip() else Path.home() + if kind == "folder": + selected = QFileDialog.getExistingDirectory(self, "Select folder", str(current)) + else: + selected, _filter = QFileDialog.getOpenFileName(self, "Select file", str(current)) + if selected: + edit.setText(selected) + self.changed.emit() + + def _advanced_changed(self, enabled: bool) -> None: + for key, spec in self.schema.items(): + label, widget = self.rows[key] + visible = enabled or not bool(spec.get("advanced")) + label.setVisible(visible) + widget.setVisible(visible) + + def values(self) -> dict[str, Any]: + return {key: self.value(key) for key in self.schema} + + def value(self, key: str) -> Any: + widget = self.widgets[key] + spec = self.schema[key] + kind = str(spec.get("type", "text")) + value_widget = widget.property("value_widget") + if value_widget: + widget = value_widget + if kind == "bool": + return bool(widget.isChecked()) + if kind == "choice": + data = widget.currentData() + return data if data is not None else widget.currentText() + if kind in {"int", "slider"}: + return int(widget.value()) + if kind == "float": + return float(widget.value()) + if kind == "multiline_text": + return widget.toPlainText() + return widget.text() + + def set_values(self, values: dict[str, Any]) -> None: + for key, value in values.items(): + if key not in self.widgets: + continue + self._set_value(key, value) + + def _set_value(self, key: str, value: Any) -> None: + widget = self.widgets[key] + spec = self.schema[key] + kind = str(spec.get("type", "text")) + value_widget = widget.property("value_widget") + if value_widget: + widget = value_widget + try: + if kind == "bool": + widget.setChecked(bool(value)) + elif kind == "choice": + index = widget.findData(value) + if index < 0: + index = widget.findText(str(value)) + if index >= 0: + widget.setCurrentIndex(index) + elif kind in {"int", "slider"}: + widget.setValue(int(value)) + elif kind == "float": + widget.setValue(float(value)) + elif kind == "multiline_text": + widget.setPlainText(str(value)) + else: + widget.setText(str(value)) + except (TypeError, ValueError): + return diff --git a/adam/ui/studio.py b/adam/ui/studio.py index adf339ee38b0333d87552bc5f8bd96ae0a6d2cc1..6fc525040a6aafa2ae2c9148afad48243ee32ed3 100644 --- a/adam/ui/studio.py +++ b/adam/ui/studio.py @@ -36,6 +36,7 @@ from PySide6.QtWidgets import ( from adam.assets import Asset, AssetRegistry from adam.config import ConfigManager +from adam.dataset_registry import DatasetRegistry from adam.job_manager import JobManager from adam.models import Job from adam.eve import EveResult, EveVisionModel, save_eve_results @@ -399,11 +400,84 @@ class EveReviewDialog(QDialog): super().closeEvent(event) +class DatasetLocationDialog(QDialog): + def __init__( + self, + root_path: Path, + config: ConfigManager, + assets: AssetRegistry, + parent: QWidget | None = None, + ) -> None: + super().__init__(parent) + self.registry = DatasetRegistry(root_path, config) + self.assets = assets + self.setWindowTitle("Remembered Dataset Locations") + self.setMinimumSize(740, 420) + root = QVBoxLayout(self) + root.addWidget(_header( + "Remembered dataset locations", + "Add dataset parent folders here so ADAM Remote can find them without exposing arbitrary PC browsing.", + )) + self.list = QListWidget() + root.addWidget(self.list, 1) + actions = QHBoxLayout() + add = QPushButton("Add location...") + remove = QPushButton("Remove selected") + refresh = QPushButton("Refresh") + add.setProperty("primary", True) + add.clicked.connect(self._add) + remove.clicked.connect(self._remove) + refresh.clicked.connect(self.refresh) + actions.addWidget(add) + actions.addWidget(remove) + actions.addWidget(refresh) + actions.addStretch() + root.addLayout(actions) + buttons = QDialogButtonBox(QDialogButtonBox.Close) + buttons.rejected.connect(self.reject) + root.addWidget(buttons) + self.refresh() + + def refresh(self) -> None: + self.registry.load() + self.registry.discover_into_assets(self.assets, persist=True) + self.list.clear() + for location in self.registry.known_locations(): + status = "available" if Path(location.path).is_dir() else "unavailable" + label = f"{location.name}\n{status} - {location.source} - {location.path}" + item = QListWidgetItem(label) + item.setData(Qt.UserRole, location.id) + self.list.addItem(item) + + def _add(self) -> None: + selected = QFileDialog.getExistingDirectory( + self, "Remember dataset location", str(Path.home()) + ) + if not selected: + return + try: + self.registry.register_location(selected, name=Path(selected).name, source="user") + except ValueError as exc: + QMessageBox.warning(self, "Location not saved", str(exc)) + return + self.refresh() + + def _remove(self) -> None: + item = self.list.currentItem() + if not item: + return + location_id = str(item.data(Qt.UserRole) or "") + self.registry.remove_location(location_id) + self.refresh() + + class DatasetReviewTab(QWidget): - def __init__(self, assets: AssetRegistry, store: StudioStore) -> None: + def __init__(self, assets: AssetRegistry, store: StudioStore, config: ConfigManager) -> None: super().__init__() self.assets = assets self.store = store + self.config = config + self.registry = DatasetRegistry(assets.path.parent.parent, config) self.paths: list[Path] = [] self.dataset_path = "" self._load_index = 0 @@ -420,6 +494,8 @@ class DatasetReviewTab(QWidget): self.dataset.setMinimumWidth(300) browse = QPushButton("Open another dataset…") browse.clicked.connect(self._browse) + locations = QPushButton("Remember locations…") + locations.clicked.connect(self._manage_locations) refresh = QPushButton("Refresh") refresh.clicked.connect(self.refresh) duplicates = QPushButton("Check duplicates") @@ -445,6 +521,7 @@ class DatasetReviewTab(QWidget): dataset_row.addWidget(QLabel("Dataset")) dataset_row.addWidget(self.dataset, 1) dataset_row.addWidget(browse) + dataset_row.addWidget(locations) dataset_row.addWidget(refresh) review_actions.addWidget(duplicates) review_actions.addWidget(keep_all) @@ -514,6 +591,7 @@ class DatasetReviewTab(QWidget): self.reload_assets() def reload_assets(self) -> None: + self.registry.discover_into_assets(self.assets, persist=True) current = self.dataset.currentData() self.dataset.blockSignals(True) self.dataset.clear() @@ -534,10 +612,27 @@ class DatasetReviewTab(QWidget): return index = self.dataset.findData(selected) if index < 0: - self.dataset.addItem(Path(selected).name, selected) + asset = self.assets.register( + kind="dataset", + name=Path(selected).name, + path=selected, + metadata={"dataset_registry_source": "studio"}, + ) + self.registry.record_for_path(asset.path) + self.dataset.addItem(asset.name, asset.path) index = self.dataset.count() - 1 self.dataset.setCurrentIndex(index) + def _manage_locations(self) -> None: + dialog = DatasetLocationDialog( + self.assets.path.parent.parent, + self.config, + self.assets, + self, + ) + dialog.exec() + self.reload_assets() + def refresh(self) -> None: previous_path = self.dataset_path self.dataset_path = str(self.dataset.currentData() or "") @@ -857,8 +952,24 @@ class ExperimentsTab(QWidget): actions.addWidget(self.open) detail_layout.addLayout(actions) root.addWidget(detail, 1) - jobs.job_created.connect(lambda _job: self.refresh()) - jobs.job_updated.connect(lambda _job: self.refresh()) + self.refresh_timer = QTimer(self) + self.refresh_timer.setSingleShot(True) + self.refresh_timer.setInterval(1000) + self.refresh_timer.timeout.connect(self._refresh_if_visible) + jobs.job_created.connect(self._schedule_refresh) + jobs.job_updated.connect(self._schedule_refresh) + self.refresh() + + def _schedule_refresh(self, _job: Job) -> None: + if self.isVisible() and not self.refresh_timer.isActive(): + self.refresh_timer.start() + + def _refresh_if_visible(self) -> None: + if self.isVisible(): + self.refresh() + + def showEvent(self, event) -> None: + super().showEvent(event) self.refresh() @staticmethod @@ -1354,7 +1465,6 @@ class StudioPage(QWidget): config: ConfigManager, ) -> None: super().__init__() - del config self.store = StudioStore(root_path) root = QVBoxLayout(self) root.setContentsMargins(24, 20, 24, 17) @@ -1366,7 +1476,7 @@ class StudioPage(QWidget): ) ) self.tabs = QTabWidget() - self.datasets = DatasetReviewTab(assets, self.store) + self.datasets = DatasetReviewTab(assets, self.store, config) self.experiments = ExperimentsTab(jobs, assets, self.store) self.previews = PreviewLabTab(assets, self.store) self.recipes = RecipesTab(self.store) @@ -1377,15 +1487,22 @@ class StudioPage(QWidget): root.addWidget(self.tabs, 1) self.previews.plan_requested.connect(self.plan_requested) self.recipes.plan_requested.connect(self.plan_requested) + self.tabs.currentChanged.connect(self._refresh_current_tab) def refresh(self) -> None: self.datasets.assets.load() - self.datasets.reload_assets() - self.experiments.assets.load() - self.experiments.refresh() - self.previews.assets.load() - self.previews.reload_assets() - self.recipes.refresh() + self._refresh_current_tab() + + def _refresh_current_tab(self, _index: int = 0) -> None: + current = self.tabs.currentWidget() + if current is self.datasets: + self.datasets.reload_assets() + elif current is self.experiments: + self.experiments.refresh() + elif current is self.previews: + self.previews.reload_assets() + elif current is self.recipes: + self.recipes.refresh() def shutdown(self) -> None: workers = [ diff --git a/adam/ui/theme.py b/adam/ui/theme.py index 6056eb0dfa3cd9bde3b7f76ed6411254eab5fb0d..7caafbfd2945f62215a42baed3ebbdff46b55d08 100644 --- a/adam/ui/theme.py +++ b/adam/ui/theme.py @@ -26,7 +26,7 @@ APP_STYLESHEET = f""" font-size: 13px; color: {COLORS["text"]}; }} -QMainWindow, QWidget#Root {{ +QMainWindow, QDialog, QWidget#Root {{ background: {COLORS["bg"]}; }} QFrame#Sidebar {{ @@ -47,6 +47,17 @@ QFrame[innerCard="true"] {{ border: 1px solid {COLORS["border"]}; border-radius: 9px; }} +QGroupBox {{ + margin-top: 14px; + padding-top: 12px; +}} +QGroupBox::title {{ + subcontrol-origin: margin; + subcontrol-position: top left; + left: 12px; + padding: 0 5px; + color: {COLORS["blue_2"]}; +}} QLabel {{ background: transparent; }} @@ -80,6 +91,12 @@ QLabel#CardTitle {{ font-weight: 700; letter-spacing: 1.5px; }} +QLabel[sectionTitle="true"] {{ + color: {COLORS["blue_2"]}; + font-size: 12px; + font-weight: 700; + padding: 1px 0 3px 0; +}} QPushButton {{ min-height: 34px; padding: 0 16px; @@ -177,7 +194,7 @@ QPushButton[quick="true"]:hover {{ color: {COLORS["blue_2"]}; background: #0a1924; }} -QLineEdit, QTextEdit, QPlainTextEdit, QComboBox, QSpinBox, QDoubleSpinBox {{ +QLineEdit, QTextEdit, QPlainTextEdit, QComboBox, QSpinBox, QDoubleSpinBox, QDateTimeEdit {{ background: #050d14; border: 1px solid {COLORS["border_bright"]}; border-radius: 8px; @@ -203,6 +220,9 @@ QScrollArea {{ QScrollArea > QWidget > QWidget {{ background: {COLORS["surface"]}; }} +QScrollArea#SidebarScroll, QScrollArea#SidebarScroll > QWidget > QWidget {{ + background: {COLORS["sidebar"]}; +}} QScrollBar:vertical {{ width: 8px; background: transparent; @@ -217,8 +237,8 @@ QScrollBar::add-line:vertical, QScrollBar::sub-line:vertical {{ height: 0; }} QProgressBar {{ - min-height: 7px; - max-height: 7px; + min-height: 18px; + max-height: 18px; border: 0; border-radius: 3px; background: #13232e; diff --git a/adam/ui/widgets.py b/adam/ui/widgets.py index fcdc2fbe3b3ae7fd3d359710f9f38994e44215ac..e5f69b46d0386082d1c6a2cc5b3c26be9c2649c2 100644 --- a/adam/ui/widgets.py +++ b/adam/ui/widgets.py @@ -326,6 +326,8 @@ class GenerationChatCard(QFrame): self.cancel_button.hide() elif job.status == JobStatus.AWAITING_CONFIRMATION: self.status.setText("Waiting for approval in Current Plan") + elif job.status == JobStatus.SCHEDULED: + self.status.setText("Scheduled · waiting for its start time") elif job.status == JobStatus.QUEUED: self.status.setText("Queued · waiting for the generator") elif job.status == JobStatus.PAUSED: @@ -646,6 +648,7 @@ class ActiveJobPanel(QFrame): JobStatus.CANCELLED: COLORS["red"], JobStatus.INTERRUPTED: COLORS["orange"], JobStatus.AWAITING_CONFIRMATION: COLORS["purple"], + JobStatus.SCHEDULED: COLORS["purple"], }.get(job.status, COLORS["muted"]) self.status.setStyleSheet( f"font-size: 11px; font-weight: 700; color: {color};" @@ -663,7 +666,15 @@ class ActiveJobPanel(QFrame): f"{atlas.get('message', 'Monitoring begins with the training step.')}" ) self.progress.setValue(job.progress) - self.progress_text.setText(f"{job.progress}% · Job {job.id}") + progress_bits = [f"{job.progress}%", f"Job {job.id}"] + if job.progress_current and job.progress_total: + progress_bits.insert( + 1, + f"{job.progress_unit.title()} {job.progress_current:,} of {job.progress_total:,}", + ) + if job.progress_rate > 0: + progress_bits.append(self._format_rate(job.progress_rate, job.progress_unit)) + self.progress_text.setText(" · ".join(progress_bits)) losses: list[float] = [] for line in job.logs: match = re.search( @@ -763,16 +774,39 @@ class ActiveJobPanel(QFrame): else f"{elapsed // 60}m {elapsed % 60}s" ) eta = "—" - if 0 < job.progress < 100: - remaining = int(elapsed * (100 - job.progress) / job.progress) - eta = ( - f"{remaining // 3600}h {(remaining % 3600) // 60}m" - if remaining >= 3600 - else f"{remaining // 60}m {remaining % 60}s" - ) - elif job.progress == 100: + completion = "" + if job.status == JobStatus.FINISHED and job.ended_at: eta = "complete" - self.timing.setText(f"Elapsed {elapsed_text} · ETA {eta}") + completion = f" · Completed at {self._format_clock_time(job.ended_at)}" + elif 0 < job.progress < 100: + remaining = job.eta_seconds + if remaining is None: + remaining = int(elapsed * (100 - job.progress) / job.progress) + eta = self._format_duration(max(0, int(remaining))) + if job.estimated_completion_at: + completion = f" · Completion at {self._format_clock_time(job.estimated_completion_at)}" + self.timing.setText(f"Elapsed {elapsed_text} · ETA {eta}{completion}") + + @staticmethod + def _format_duration(seconds: int) -> str: + if seconds >= 3600: + return f"{seconds // 3600}h {(seconds % 3600) // 60}m" + return f"{seconds // 60}m {seconds % 60}s" + + @staticmethod + def _format_clock_time(value: str) -> str: + try: + text = datetime.fromisoformat(value).astimezone().strftime("%I:%M %p") + except ValueError: + return "—" + return text.lstrip("0") + + @staticmethod + def _format_rate(rate: float, unit: str) -> str: + if rate >= 1: + return f"{rate:.2f} {unit}s/s" + seconds_per_unit = 1 / rate if rate > 0 else 0 + return f"{seconds_per_unit:.2f}s/{unit}" def _update_buttons(self) -> None: running = bool( diff --git a/config/tools.json b/config/tools.json index 2bb16b2645991b4e66c8a68fc48cd9066b914b4f..74b80340f7554247730816eda830e02b6bc994b4 100644 --- a/config/tools.json +++ b/config/tools.json @@ -74,7 +74,7 @@ "description": "Launches and monitors a registered SD/SDXL LoRA training backend.", "category": "Training", "entry_function": "train_lora", - "arguments": ["dataset_dir", "model_name", "epochs", "output_dir", "base_model", "resume_from", "preview_enabled", "preview_every", "preview_prompt", "preview_seed"], + "arguments": ["dataset_dir", "model_name", "epochs", "output_dir", "base_model", "trigger_word", "resume_from", "preview_enabled", "preview_every", "preview_prompt", "preview_seed"], "required_arguments": ["dataset_dir", "model_name", "epochs", "output_dir", "base_model"], "capabilities": ["fresh_training", "resume_training", "progress", "pause", "cancel"], "requires_confirmation": true, @@ -116,9 +116,9 @@ "description": "Generates reproducible image batches from completed models in the connected DDPM project.", "category": "Output", "entry_function": "generate_ddpm_images", - "arguments": ["model_name", "model_path", "prompt", "image_count", "steps", "seed", "sampler", "aspect_ratio", "reference_image", "reference_strength", "width", "height", "preview_interval"], + "arguments": ["model_name", "model_path", "prompt", "image_count", "steps", "seed", "sampler", "aspect_ratio", "reference_image", "reference_strength", "width", "height", "preview_interval", "smart_generation", "smart_wanted_results", "smart_max_candidates", "smart_min_score", "smart_mode", "smart_keep_rejected"], "required_arguments": ["model_name", "model_path", "prompt", "image_count", "steps", "seed", "sampler", "aspect_ratio"], - "capabilities": ["image_generation", "seed", "sampler", "batch", "aspect_ratio", "reference_image", "live_preview", "progress", "cancel"], + "capabilities": ["image_generation", "smart_generation", "seed", "sampler", "batch", "aspect_ratio", "reference_image", "live_preview", "progress", "cancel"], "model_trainers": ["ddpm"], "generation_options": { "samplers": ["DDIM", "DDPM"], @@ -142,9 +142,9 @@ "description": "Generates reproducible image batches from completed models in the connected Flow Matching project.", "category": "Output", "entry_function": "generate_flow_images", - "arguments": ["model_name", "model_path", "prompt", "image_count", "steps", "seed", "sampler", "aspect_ratio", "preview_interval"], + "arguments": ["model_name", "model_path", "prompt", "image_count", "steps", "seed", "sampler", "aspect_ratio", "preview_interval", "smart_generation", "smart_wanted_results", "smart_max_candidates", "smart_min_score", "smart_mode", "smart_keep_rejected"], "required_arguments": ["model_name", "model_path", "prompt", "image_count", "steps", "seed", "sampler", "aspect_ratio"], - "capabilities": ["image_generation", "seed", "ode_method", "batch", "aspect_ratio", "live_preview", "progress", "cancel"], + "capabilities": ["image_generation", "smart_generation", "seed", "ode_method", "batch", "aspect_ratio", "live_preview", "progress", "cancel"], "model_trainers": ["flow"], "generation_options": { "samplers": ["Heun", "Euler"], diff --git a/docs/oasis_integration.md b/docs/oasis_integration.md new file mode 100644 index 0000000000000000000000000000000000000000..fe5d9a565a2a398ad100cfc842eb3b04f3aa945d --- /dev/null +++ b/docs/oasis_integration.md @@ -0,0 +1,52 @@ +# Oasis Action World Model Integration + +ADAM treats Oasis as a first-class trainer named `Oasis Action World Model`. +It uses the existing `Oasis-Game-Trainer/roblox_action_flow_app.py` worker and +keeps training and playable inference in separate processes. + +## Dataset Layout + +Use one or more recorder folders. Multiple folders can be supplied with +semicolon separators. + +```text +My_Oasis_Dataset/ + dataset_info.json + actions.jsonl + frames/ + frame_00000000.png + frame_00000001.png +``` + +Each `actions.jsonl` row must identify the frame and its labels: + +```json +{"session_id":"run1","frame_index":1,"filename":"frame_00000001.png","w":1,"a":0,"s":0,"d":0,"jump":0,"mouse_dx":0.0,"mouse_dy":0.0,"zoom":0.0} +``` + +Rows describe the action for the transition into that frame. ADAM validates +missing frames, invalid labels, duplicate or gapped frame indexes, broken +images, inconsistent resolutions, empty sequences, and missing action columns. +It does not guess missing labels. + +## Minecraft Beta-Style Recording Notes + +Record your own gameplay. Start with `256x144` at 8-12 FPS for a first pipeline +test, then increase data before long runs. Keep every clip as its own session or +recording folder so training never learns fake transitions between worlds. + +Capture a mix of walking forward, strafing, turning, jumping, climbing, terrain +changes, interiors, caves, water, sky, and short idle moments. Avoid long +standing-still stretches; idle is useful, but a dataset dominated by idle frames +will make the model ignore controls. + +A good smoke dataset is 300-1,000 labelled frames. A more meaningful first +Minecraft experiment is 5,000-20,000 labelled frames with clean W/A/S/D/Space +coverage and synchronized frame/action timestamps. + +## RTX 3060 12 GB Starting Point + +Use `256x144`, batch size `2`, mixed precision `fp32`, gradient accumulation +`1`, prediction horizon `3`, loader workers `2`, save every `5` epochs, preview +every `5` epochs, and 5-10 epochs for a pipeline test. Try `fp16` only after +FP32 is stable. diff --git a/docs/screenshots/command-center.png b/docs/screenshots/command-center.png new file mode 100644 index 0000000000000000000000000000000000000000..695999b92fa99c5de2cefc23ef298daaf9423654 --- /dev/null +++ b/docs/screenshots/command-center.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:a3e92b6b7189532e72b99ec16702721ed239177530836ecd5ce9b3f1389ffc21 +size 151984 diff --git a/docs/screenshots/create-model.png b/docs/screenshots/create-model.png new file mode 100644 index 0000000000000000000000000000000000000000..575858555fddf82ccd37cd6f17e88534e306770c Binary files /dev/null and b/docs/screenshots/create-model.png differ diff --git a/docs/youtube_video_collector_plan.md b/docs/youtube_video_collector_plan.md new file mode 100644 index 0000000000000000000000000000000000000000..c83c177632602f7b9caf60c0977fd5fe0aa9f553 --- /dev/null +++ b/docs/youtube_video_collector_plan.md @@ -0,0 +1,60 @@ +# ADAM Video Dataset Collector — implementation plan + +## Existing architecture and integration point + +ADAM plans only allow-listed tools from `config/tools.json`. Python backends are loaded by `adam/executor.py`, receive a `ToolContext`, run through the existing background job manager, and use that context for logs, progress, pause, and cancellation. `adam/planner.py` handles deterministic workflows before asking Ollama to propose a registry-validated plan. + +The current image collector (`adam/tools/real_dataset_collector.py`) is a separate Bing/Selenium workflow and remains unchanged. The external VTI installation is user-configured. Its three scripts are Tk/CustomTkinter GUI applications. Extraction is embedded in window classes, imports initialize UI dependencies, and none exposes a stable headless function or command contract. Importing or driving that GUI from ADAM would be fragile and would bypass ADAM's job controls. + +The clean first integration is one registered Python backend, `adam.tools.youtube_video_collector.collect_youtube_dataset`. It owns collection, normalization, provenance, and headless OpenCV extraction. A follow-up VTI cleanup should move extraction into a shared module and have both the VTI GUI and collector call it. The external VTI files are deliberately not rewritten in version one. + +## Files implemented + +- `adam/tools/youtube_video_collector.py`: settings validation, metadata preview, filters, bounded playlist expansion, download/archive handling, MP4 output, extraction, quality checks, provenance, credits, CLI, and partial-failure logic. +- `config/tools.json`: registered ADAM tool and supported argument contract. +- `adam/planner.py`: deterministic recognition of supplied YouTube URLs and common natural-language options. +- `requirements.txt`: `yt-dlp` and OpenCV dependencies. +- `tests/test_youtube_video_collector.py`: offline metadata, filtering, synthetic MP4/frame manifest, and incomplete-file cleanup tests. + +## Data flow + +1. Validate settings and create a new, non-overwriting dataset directory. +2. Expand supplied video/playlist URLs with `playlistend=max_videos`. +3. Store title, IDs, channel, date, duration, resolution, thumbnail, URL, and reported license. Apply video, total-duration, and estimated-size limits. +4. Persist preview provenance before download. A dry run stops here. +5. Download one source at a time with strict retries and a shared archive; select video-only formats unless audio is enabled; ask FFmpeg for MP4. +6. Extract timestamped frames. Sequential mode retains ordering and disables near-duplicate rejection. Image mode applies blur, black-frame, and adjacent-similarity checks. +7. Atomically refresh per-source JSON, `sources.json`, the frame manifest, and credits after every source so one failure does not discard completed work. +8. On cancellation, remove downloader temporary files and preserve completed work. + +## Dependencies and compatibility + +- Python 3.10 or newer. +- Current `yt-dlp`; YouTube changes periodically, so keep it current. +- FFmpeg and ffprobe on `PATH`. +- OpenCV and NumPy for headless extraction and quality scores. +- Windows-safe names and no shell invocation are used. Long path support can still depend on Windows configuration. + +## Remaining staged work + +### Version 1.1 + +- Move sampling into a shared package and update the external VTI GUI to consume it after separately backing up and testing that project. +- Implement true mixed-output copies and a non-separated layout; version one always preserves safer per-source directories even though the options are accepted for forward compatibility. +- Add subtitle downloads and generated visual captions. Current captions use the source title. +- Enforce total bytes during transfer; version one filters on reported estimated size, which can be unavailable. +- Add an ADAM settings panel and thumbnail approval table. Current integration uses chat, plan approval, and job logs. +- Upgrade adjacent pixel similarity to perceptual hashes or embeddings. + +### Version 2 + +- Add a `CandidateProvider` interface. A YouTube Search API provider can return candidate IDs without changing download or extraction. +- Rank candidates by reviewable features, show thumbnails and metadata, and require approval before passing URLs to the collector. +- Prefer an official API or supplied candidate list; do not require autonomous browser interaction. + +## Risks + +- Reported license fields can be absent or inaccurate. ADAM records them but never treats attribution as permission. +- Private, deleted, restricted, or login-required videos may need cookies or remain unavailable. Version one records failures and continues; it does not bypass controls. +- Cancellation is responsive between download fragments, but a blocked network operation may last until its timeout. +- FFmpeg transcoding can be CPU and disk intensive, so actual collection retains ADAM's approval gate. diff --git a/main.py b/main.py index a9e66531f0395c74b58f4b49ba473df9dec5b305..4eb652734e9a29d7f6adf1b4e2bc0ff1170e9081 100644 --- a/main.py +++ b/main.py @@ -2,6 +2,7 @@ from __future__ import annotations import argparse import ctypes +import multiprocessing import os import shutil import sys @@ -99,4 +100,5 @@ def main() -> int: if __name__ == "__main__": + multiprocessing.freeze_support() raise SystemExit(main()) diff --git a/models/__init__.py b/models/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..4eacb5237daeff2178531b34352d083f08f24478 --- /dev/null +++ b/models/__init__.py @@ -0,0 +1 @@ +"""User-installable ADAM model plugins.""" diff --git a/models/model_template/__init__.py b/models/model_template/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..929070fe870b9bdd3428bbf59f3bb1c4740915d1 --- /dev/null +++ b/models/model_template/__init__.py @@ -0,0 +1 @@ +"""Copy this folder to create a new ADAM model plugin.""" diff --git a/models/model_template/generator.py b/models/model_template/generator.py new file mode 100644 index 0000000000000000000000000000000000000000..a9acbe168d0c0a00d3a94cbb185afa36c4d38080 --- /dev/null +++ b/models/model_template/generator.py @@ -0,0 +1,13 @@ +from __future__ import annotations + +from typing import Any + + +def generate(context, **settings: Any) -> dict[str, Any]: + """Put your generation implementation here. + + Return a dictionary with output paths or assets when generation is done. + """ + context.log("Template generator loaded. Replace this with real generation code.") + context.progress(100, "Template generation complete") + return {} diff --git a/models/model_template/manifest.py b/models/model_template/manifest.py new file mode 100644 index 0000000000000000000000000000000000000000..c72b0b5ca8eb9dc70c10be6df951dcb9f62b4a15 --- /dev/null +++ b/models/model_template/manifest.py @@ -0,0 +1,39 @@ +PLUGIN_ID = "my_model" + +MODEL_INFO = { + "name": "My Model", + "version": "0.1", + "category": "Image Generation", + "description": "Describe what this model architecture does.", + "architecture": "custom", + "status": "experimental", +} + +TRAINING_SETTINGS = { + "dataset_dir": {"label": "Dataset folder", "type": "folder", "required": True, "must_exist": True, "group": "Dataset"}, + "output_dir": {"label": "Output folder", "type": "folder", "required": True, "group": "Checkpoints"}, + "epochs": {"label": "Epochs", "type": "int", "default": 10, "min": 1, "max": 100000, "group": "Basic"}, + "learning_rate": {"label": "Learning rate", "type": "float", "default": 0.0001, "min": 0.0000001, "max": 0.1, "group": "Optimization"}, +} + +GENERATION_SETTINGS = { + "model_path": {"label": "Model file or folder", "type": "path", "required": True, "must_exist": True, "group": "Model Loading"}, + "output_dir": {"label": "Output folder", "type": "folder", "required": True, "group": "Generation"}, + "seed": {"label": "Seed", "type": "int", "default": 0, "min": 0, "max": 2147483647, "group": "Generation"}, +} + +# Uncomment these after copying this template and replacing the module path. +# TRAINING_TOOL = { +# "id": "my_model_trainer", +# "name": "My Model Trainer", +# "backend": {"type": "python", "module": "models.my_model.trainer", "function": "train"}, +# } +# +# GENERATION_TOOL = { +# "id": "my_model_generator", +# "name": "My Model Generator", +# "model_trainers": ["my_model"], +# "backend": {"type": "python", "module": "models.my_model.generator", "function": "generate"}, +# } +TRAINING_TOOL = {} +GENERATION_TOOL = {} diff --git a/models/model_template/model.py b/models/model_template/model.py new file mode 100644 index 0000000000000000000000000000000000000000..62bedace57284ebf90337b1d84f504ca844f48f1 --- /dev/null +++ b/models/model_template/model.py @@ -0,0 +1,8 @@ +from __future__ import annotations + +from typing import Any + + +def load_model(model_path: str, settings: dict[str, Any] | None = None) -> Any: + """Load and return your model or inference pipeline here.""" + raise NotImplementedError("Replace load_model with your architecture code.") diff --git a/models/model_template/trainer.py b/models/model_template/trainer.py new file mode 100644 index 0000000000000000000000000000000000000000..234d3cbc4c9e18c235005a51bc0dd8c786781233 --- /dev/null +++ b/models/model_template/trainer.py @@ -0,0 +1,14 @@ +from __future__ import annotations + +from typing import Any + + +def train(context, **settings: Any) -> dict[str, Any]: + """Put your training implementation here. + + Use context.log("message"), context.progress(percent, "message"), and + context.preview(path) to report back to ADAM without touching the GUI. + """ + context.log("Template trainer loaded. Replace this with real training code.") + context.progress(100, "Template training complete") + return {} diff --git a/models/neural_cellular_automata/__init__.py b/models/neural_cellular_automata/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..a9cfdbd9d9d567338263e40b7a557564116018b5 --- /dev/null +++ b/models/neural_cellular_automata/__init__.py @@ -0,0 +1 @@ +"""Neural Cellular Automata model plugin for ADAM.""" diff --git a/models/neural_cellular_automata/config.py b/models/neural_cellular_automata/config.py new file mode 100644 index 0000000000000000000000000000000000000000..99908e57e9d0924828c7f137d7649092d075b794 --- /dev/null +++ b/models/neural_cellular_automata/config.py @@ -0,0 +1,49 @@ +MODEL_INFO = { + "name": "Neural Cellular Automata", + "short_name": "NCA", + "category": "Cellular Automata", + "description": "Learns local cellular update rules that grow images from a single living seed.", + "supports_training": True, + "supports_generation": True, + "order": 10, +} + +TRAINING_SETTINGS = { + "training_mode": { + "label": "Training mode", + "type": "choice", + "options": ["Single target", "Dataset average (experimental)"], + "default": "Single target", + }, + "target_image": {"label": "Target image", "type": "file", "default": ""}, + "epochs": {"label": "Training steps", "type": "int", "default": 500, "min": 1, "max": 100000}, + "batch_size": {"type": "int", "default": 8, "min": 1, "max": 128}, + "learning_rate": {"type": "float", "default": 0.001, "min": 0.0000001}, + "resolution": {"label": "Target resolution", "type": "choice", "options": [32, 64, 128], "default": 64}, + "cell_channels": {"type": "int", "default": 16, "min": 4, "max": 64}, + "hidden_size": {"type": "int", "default": 128, "min": 16, "max": 512}, + "min_growth_steps": {"type": "int", "default": 64, "min": 1, "max": 512}, + "max_growth_steps": {"type": "int", "default": 96, "min": 1, "max": 1024}, + "stable_growth": {"label": "Stable growth training", "type": "bool", "default": False}, + "stable_max_growth_steps": {"label": "Stable maximum growth steps", "type": "int", "default": 160, "min": 1, "max": 2048}, + "fire_rate": {"label": "Update probability", "type": "slider", "default": 0.5, "min": 0.05, "max": 1.0, "resolution": 0.05}, + "random_seed": {"type": "int", "default": 42}, + "mixed_precision": {"type": "bool", "default": True}, + "device": {"type": "choice", "options": ["auto", "cuda", "cpu"], "default": "auto"}, + "preview_every": {"label": "Preview every N steps", "type": "int", "default": 10, "min": 1}, + "checkpoint_every": {"label": "Checkpoint every N steps", "type": "int", "default": 50, "min": 1}, + "loss_smoothing_window": {"label": "Loss smoothing window", "type": "int", "default": 25, "min": 1, "max": 1000}, + "fit_mode": {"label": "Image fit", "type": "choice", "options": ["contain", "crop", "stretch"], "default": "contain"}, + "background": {"type": "choice", "options": ["transparent", "white", "black"], "default": "transparent"}, + "horizontal_flip": {"label": "Random horizontal flip", "type": "bool", "default": False}, + "resume_checkpoint": {"label": "Resume checkpoint (optional)", "type": "file", "default": ""}, +} + +GENERATION_SETTINGS = { + "growth_steps": {"type": "int", "default": 100, "min": 1, "max": 5000}, + "seed": {"type": "int", "default": -1}, + "fire_rate": {"label": "Update probability", "type": "slider", "default": 0.5, "min": 0.05, "max": 1.0, "resolution": 0.05}, + "animate_growth": {"type": "bool", "default": True}, + "frame_every": {"label": "Save every Nth step", "type": "int", "default": 5, "min": 1}, + "export_gif": {"type": "bool", "default": True}, +} diff --git a/models/neural_cellular_automata/generator.py b/models/neural_cellular_automata/generator.py new file mode 100644 index 0000000000000000000000000000000000000000..b1a86d04b1db97b67b1c6b5f0ea4a7acf0628721 --- /dev/null +++ b/models/neural_cellular_automata/generator.py @@ -0,0 +1,204 @@ +from __future__ import annotations + +import json +import random +from datetime import datetime +from pathlib import Path + +import numpy as np +import torch +from PIL import Image + +from .model import create_model, create_seed +from .image_display import over_checkerboard, side_by_side +from .image_preprocessing import prepare_image + + +GENERATION_KEYS = ("growth_steps", "seed", "fire_rate", "animate_growth", "frame_every", "export_gif") + + +def _to_image(state: torch.Tensor) -> Image.Image: + rgba = state[0, :4].detach().float().clamp(0, 1).permute(1, 2, 0).cpu().numpy() + return Image.fromarray((rgba * 255).astype(np.uint8), "RGBA") + + +def _meaningfully_visible(image: Image.Image) -> bool: + alpha = image.getchannel("A") + visible_pixels = sum(alpha.histogram()[16:]) + return visible_pixels >= max(4, int(image.width * image.height * 0.01)) + + +def load_model(checkpoint_path): + checkpoint_path = Path(checkpoint_path) + if not checkpoint_path.is_file(): + raise FileNotFoundError(f"Checkpoint not found: {checkpoint_path}") + device = torch.device("cuda" if torch.cuda.is_available() else "cpu") + data = torch.load(checkpoint_path, map_location=device, weights_only=False) + model = create_model(data["config"]).to(device) + model.load_state_dict(data["model_state"]) + model.eval() + return model, data["config"] + + +def generation_warning(config: dict) -> str | None: + steps = int(config.get("growth_steps", 100)) + minimum = int(config.get("min_growth_steps", 1)) + normal_maximum = int(config.get("max_growth_steps", steps)) + safe_maximum = ( + int(config.get("stable_max_growth_steps", normal_maximum)) + if config.get("stable_growth", False) + else normal_maximum + ) + if steps > safe_maximum: + return ( + f"This checkpoint trained for {minimum}-{safe_maximum} growth steps. " + f"Generating {steps} steps may distort, explode, or erase the image." + ) + if steps < minimum: + return f"This checkpoint trained for at least {minimum} growth steps; {steps} steps may look incomplete." + return None + + +def _save_json(path: Path, data: dict) -> None: + temporary = path.with_suffix(".json.tmp") + with temporary.open("w", encoding="utf-8") as handle: + json.dump(data, handle, indent=2) + temporary.replace(path) + + +def generate_nca(model, config, output_path, progress_callback=None, preview_callback=None, stop_event=None): + output = Path(output_path) + frames_dir = output / "frames" + output.mkdir(parents=True, exist_ok=True) + if config.get("animate_growth", True): + frames_dir.mkdir(exist_ok=True) + + seed = int(config.get("seed", -1)) + if seed < 0: + seed = random.SystemRandom().randint(0, 2**31 - 1) + torch.manual_seed(seed) + device = next(model.parameters()).device + resolution = int(config.get("resolution", 64)) + state = create_seed(1, model.channels, resolution, device) + steps = int(config.get("growth_steps", 100)) + frame_every = max(1, int(config.get("frame_every", 5))) + frames: list[Image.Image] = [] + frame_steps: list[int] = [] + warning = generation_warning(config) + metadata = { + "created_at": datetime.now().isoformat(timespec="seconds"), + "status": "running", + "settings": {key: config.get(key) for key in GENERATION_KEYS}, + "resolved_seed": seed, + "trained_growth_range": { + "minimum": int(config.get("min_growth_steps", 1)), + "maximum": int(config.get("stable_max_growth_steps", config.get("max_growth_steps", steps))) + if config.get("stable_growth", False) + else int(config.get("max_growth_steps", steps)), + }, + "warning": warning, + } + _save_json(output / "generation_config.json", metadata) + + with torch.inference_mode(): + for step in range(steps + 1): + if stop_event is not None and stop_event.is_set(): + break + if step == 0 or step % frame_every == 0 or step == steps: + image = _to_image(state) + if config.get("animate_growth", True): + image.save(frames_dir / f"step_{step:05d}.png") + frames.append(image.copy()) + frame_steps.append(step) + if preview_callback: + preview_callback(image.copy()) + if progress_callback: + progress_callback({"epoch": step, "total_epochs": steps, "loss": None, "seed": seed}) + if step < steps: + state = model(state, fire_rate=float(config.get("fire_rate", model.fire_rate))) + + final_path = output / "final.png" + final_image = _to_image(state) + final_image.save(final_path) + gif_path = None + if config.get("export_gif", True) and frames: + gif_path = output / "growth.gif" + visible_index = next( + (index for index, frame in enumerate(frames) if _meaningfully_visible(frame)), + len(frames) - 1, + ) + # The raw frame export remains complete. GIFs begin at a meaningful frame + # and use a checkerboard so transparent growth is visible in all viewers. + gif_frames = [over_checkerboard(frame) for frame in frames[visible_index:]] + gif_frames[0].save(gif_path, save_all=True, append_images=gif_frames[1:], duration=100, loop=0, disposal=2) + + comparison_path = None + target_path = str(config.get("target_image", "")).strip() + if target_path and Path(target_path).is_file(): + try: + with Image.open(target_path) as target: + target = prepare_image( + target, + resolution, + str(config.get("fit_mode", "contain")), + str(config.get("background", "transparent")), + ) + comparison_path = output / "comparison.png" + side_by_side(target, final_image).save(comparison_path) + except OSError: + comparison_path = None + + metadata.update({ + "status": "stopped" if stop_event is not None and stop_event.is_set() else "complete", + "saved_frame_steps": frame_steps, + "gif_first_frame_step": frame_steps[visible_index] if frames and gif_path else None, + "final_path": str(final_path), + "gif_path": str(gif_path) if gif_path else None, + "comparison_path": str(comparison_path) if comparison_path else None, + }) + _save_json(output / "generation_config.json", metadata) + return { + "final_path": str(final_path), + "gif_path": str(gif_path) if gif_path else None, + "comparison_path": str(comparison_path) if comparison_path else None, + "generation_config_path": str(output / "generation_config.json"), + "seed": seed, + "warning": warning, + } + + +def generate(context, model_name: str, model_path: str, output_dir: str = "", **settings): + output = Path(output_dir) if output_dir else context.root / "data" / "generations" / "neural_cellular_automata" / model_name + checkpoint = Path(model_path) + if checkpoint.is_dir(): + checkpoint = checkpoint / "checkpoint.pt" + model, config = load_model(checkpoint) + config.update(settings) + + def progress(update): + step = int(update.get("epoch", 0) or 0) + total = int(update.get("total_epochs", config.get("growth_steps", 1)) or 1) + context.progress(max(1, min(99, round(step * 100 / max(total, 1)))), f"NCA growth step {step:,} of {total:,}") + + result = generate_nca( + model, + config, + output, + progress_callback=progress, + stop_event=context.cancel_event, + ) + final_path = result.get("final_path") + if final_path: + context.preview(final_path, kind="generation") + context.progress(100, "NCA generation completed") + return { + "output_folder": str(output), + "assets": [ + { + "kind": "generation", + "name": model_name, + "path": str(output), + "trainer": "neural_cellular_automata", + } + ], + } diff --git a/models/neural_cellular_automata/image_display.py b/models/neural_cellular_automata/image_display.py new file mode 100644 index 0000000000000000000000000000000000000000..fbb87f86329f6c7f1066e2c6cfafcc3671a1a5b4 --- /dev/null +++ b/models/neural_cellular_automata/image_display.py @@ -0,0 +1,41 @@ +from __future__ import annotations + +from PIL import Image, ImageDraw + + +def checkerboard(size: tuple[int, int], tile: int = 12) -> Image.Image: + image = Image.new("RGB", size, (210, 210, 210)) + draw = ImageDraw.Draw(image) + for y in range(0, size[1], tile): + for x in range(0, size[0], tile): + if (x // tile + y // tile) % 2: + draw.rectangle( + (x, y, min(x + tile - 1, size[0] - 1), min(y + tile - 1, size[1] - 1)), + fill=(160, 160, 160), + ) + return image + + +def over_checkerboard(image: Image.Image, tile: int = 12) -> Image.Image: + rgba = image.convert("RGBA") + background = checkerboard(rgba.size, tile).convert("RGBA") + background.alpha_composite(rgba) + return background.convert("RGB") + + +def side_by_side( + left: Image.Image, + right: Image.Image, + labels: tuple[str, str] = ("Target", "Generated"), +) -> Image.Image: + left = over_checkerboard(left) + right = over_checkerboard(right) + width = left.width + right.width + header = 26 + result = Image.new("RGB", (width, max(left.height, right.height) + header), "white") + result.paste(left, (0, header)) + result.paste(right, (left.width, header)) + draw = ImageDraw.Draw(result) + draw.text((6, 6), labels[0], fill="black") + draw.text((left.width + 6, 6), labels[1], fill="black") + return result diff --git a/models/neural_cellular_automata/image_preprocessing.py b/models/neural_cellular_automata/image_preprocessing.py new file mode 100644 index 0000000000000000000000000000000000000000..335b5da1708e368703da931fee0de71b4cbf188c --- /dev/null +++ b/models/neural_cellular_automata/image_preprocessing.py @@ -0,0 +1,52 @@ +from __future__ import annotations + +import random + +import numpy as np +import torch +from PIL import Image, ImageOps + + +BACKGROUND_COLORS = { + "transparent": (0, 0, 0, 0), + "white": (255, 255, 255, 255), + "black": (0, 0, 0, 255), +} + + +def prepare_image( + image: Image.Image, + resolution: int, + fit_mode: str = "contain", + background: str = "transparent", + horizontal_flip: bool = False, +) -> Image.Image: + if resolution < 1: + raise ValueError("Target resolution must be positive.") + if fit_mode not in {"contain", "crop", "stretch"}: + raise ValueError(f"Unknown fit mode: {fit_mode}") + if background not in BACKGROUND_COLORS: + raise ValueError(f"Unknown background: {background}") + image = image.convert("RGBA") + size = (resolution, resolution) + if fit_mode == "crop": + result = ImageOps.fit(image, size, method=Image.Resampling.LANCZOS) + elif fit_mode == "stretch": + result = image.resize(size, Image.Resampling.LANCZOS) + else: + contained = ImageOps.contain(image, size, method=Image.Resampling.LANCZOS) + result = Image.new("RGBA", size, BACKGROUND_COLORS[background]) + result.alpha_composite( + contained, + ((resolution - contained.width) // 2, (resolution - contained.height) // 2), + ) + return ImageOps.mirror(result) if horizontal_flip else result + + +def image_to_tensor(image: Image.Image) -> torch.Tensor: + array = np.asarray(image, dtype=np.float32) / 255.0 + return torch.from_numpy(array.copy()).permute(2, 0, 1) + + +def randomly_flip(enabled: bool) -> bool: + return bool(enabled and random.random() < 0.5) diff --git a/models/neural_cellular_automata/manifest.py b/models/neural_cellular_automata/manifest.py new file mode 100644 index 0000000000000000000000000000000000000000..905446ec2cdfec94afdfd14fdf6d7c952e63cd6b --- /dev/null +++ b/models/neural_cellular_automata/manifest.py @@ -0,0 +1,58 @@ +PLUGIN_ID = "neural_cellular_automata" + +MODEL_INFO = { + "name": "Neural Cellular Automata", + "version": "0.1", + "category": "Cellular Automata", + "description": "Learns local cellular update rules that grow images from a single living seed.", + "architecture": "nca", + "status": "experimental", + "output_type": "image", +} + +TRAINING_SETTINGS = { + "training_mode": {"label": "Training mode", "type": "choice", "options": ["Single target", "Dataset average (experimental)"], "default": "Dataset average (experimental)", "group": "Dataset"}, + "target_image": {"label": "Target image", "type": "path", "default": "", "group": "Dataset"}, + "resolution": {"label": "Target resolution", "type": "choice", "options": [32, 64, 128], "default": 64, "group": "Basic"}, + "batch_size": {"label": "Batch size", "type": "int", "default": 4, "min": 1, "max": 32, "group": "Basic"}, + "learning_rate": {"label": "Learning rate", "type": "float", "default": 0.001, "min": 0.0000001, "max": 0.1, "decimals": 7, "group": "Optimization"}, + "cell_channels": {"label": "Cell channels", "type": "int", "default": 16, "min": 4, "max": 64, "group": "Cellular Automata"}, + "hidden_size": {"label": "Hidden size", "type": "int", "default": 128, "min": 16, "max": 512, "group": "Cellular Automata"}, + "min_growth_steps": {"label": "Min growth steps", "type": "int", "default": 64, "min": 1, "max": 512, "group": "Growth"}, + "max_growth_steps": {"label": "Max growth steps", "type": "int", "default": 96, "min": 1, "max": 1024, "group": "Growth"}, + "fire_rate": {"label": "Update probability", "type": "float", "default": 0.5, "min": 0.05, "max": 1.0, "decimals": 2, "step": 0.05, "group": "Growth"}, + "mixed_precision": {"label": "Mixed precision", "type": "bool", "default": True, "group": "Optimization"}, + "device": {"label": "Device", "type": "choice", "options": ["auto", "cuda", "cpu"], "default": "auto", "group": "Optimization"}, + "fit_mode": {"label": "Image fit", "type": "choice", "options": ["contain", "crop", "stretch"], "default": "contain", "group": "Dataset"}, + "background": {"label": "Background", "type": "choice", "options": ["transparent", "white", "black"], "default": "transparent", "group": "Dataset"}, + "horizontal_flip": {"label": "Random horizontal flip", "type": "bool", "default": False, "group": "Dataset"}, + "stable_growth": {"label": "Stable growth training", "type": "bool", "default": False, "group": "Advanced", "advanced": True}, + "stable_max_growth_steps": {"label": "Stable maximum growth steps", "type": "int", "default": 160, "min": 1, "max": 2048, "group": "Advanced", "advanced": True}, + "checkpoint_every": {"label": "Checkpoint every N steps", "type": "int", "default": 50, "min": 1, "max": 5000, "group": "Checkpoints"}, + "loss_smoothing_window": {"label": "Loss smoothing window", "type": "int", "default": 25, "min": 1, "max": 1000, "group": "Advanced", "advanced": True}, + "preview_every": {"label": "Preview every N steps", "type": "int", "default": 10, "min": 1, "max": 100000, "group": "Preview"}, + "preview_prompt": {"label": "Preview prompt", "type": "text", "default": "", "group": "Preview"}, + "preview_seed": {"label": "Preview seed", "type": "int", "default": 123456789, "min": 0, "max": 2147483647, "group": "Preview"}, +} + +GENERATION_SETTINGS = { + "growth_steps": {"label": "Growth steps", "type": "int", "default": 100, "min": 1, "max": 5000, "group": "Generation"}, + "seed": {"label": "Seed", "type": "int", "default": -1, "min": -1, "max": 2147483647, "group": "Generation"}, + "fire_rate": {"label": "Update probability", "type": "float", "default": 0.5, "min": 0.05, "max": 1.0, "decimals": 2, "step": 0.05, "group": "Generation"}, + "animate_growth": {"label": "Save growth frames", "type": "bool", "default": True, "group": "Animation"}, + "frame_every": {"label": "Save every Nth step", "type": "int", "default": 5, "min": 1, "max": 500, "group": "Animation"}, + "export_gif": {"label": "Export GIF", "type": "bool", "default": True, "group": "Animation"}, +} + +TRAINING_TOOL = { + "id": "neural_cellular_automata_trainer", + "name": "Neural Cellular Automata Trainer", + "backend": {"type": "python", "module": "models.neural_cellular_automata.trainer", "function": "train"}, +} + +GENERATION_TOOL = { + "id": "neural_cellular_automata_generator", + "name": "Neural Cellular Automata Generator", + "model_trainers": ["neural_cellular_automata"], + "backend": {"type": "python", "module": "models.neural_cellular_automata.generator", "function": "generate"}, +} diff --git a/models/neural_cellular_automata/model.py b/models/neural_cellular_automata/model.py new file mode 100644 index 0000000000000000000000000000000000000000..0cd35ae16af24fe309484e86c9039ea4f8c36b7b --- /dev/null +++ b/models/neural_cellular_automata/model.py @@ -0,0 +1,60 @@ +from __future__ import annotations + +import torch +from torch import nn +from torch.nn import functional as F + + +class NeuralCellularAutomata(nn.Module): + def __init__(self, channels: int = 16, hidden_size: int = 128, fire_rate: float = 0.5): + super().__init__() + if channels < 4: + raise ValueError("NCA requires at least 4 cell channels (RGBA).") + self.channels = channels + self.hidden_size = hidden_size + self.fire_rate = fire_rate + self.update_net = nn.Sequential( + nn.Conv2d(channels * 3, hidden_size, kernel_size=1), + nn.ReLU(), + nn.Conv2d(hidden_size, channels, kernel_size=1, bias=False), + ) + nn.init.zeros_(self.update_net[-1].weight) + + identity = torch.tensor([[0.0, 0.0, 0.0], [0.0, 1.0, 0.0], [0.0, 0.0, 0.0]]) + sobel_x = torch.tensor([[-1.0, 0.0, 1.0], [-2.0, 0.0, 2.0], [-1.0, 0.0, 1.0]]) / 8.0 + sobel_y = sobel_x.t() + kernels = torch.stack((identity, sobel_x, sobel_y))[:, None] + self.register_buffer("perception_kernels", kernels) + + def perceive(self, state: torch.Tensor) -> torch.Tensor: + kernels = self.perception_kernels.repeat(self.channels, 1, 1, 1) + return F.conv2d(state, kernels, padding=1, groups=self.channels) + + @staticmethod + def living_mask(state: torch.Tensor) -> torch.Tensor: + return F.max_pool2d(state[:, 3:4], kernel_size=3, stride=1, padding=1) > 0.1 + + def forward(self, state: torch.Tensor, fire_rate: float | None = None) -> torch.Tensor: + pre_life = self.living_mask(state) + delta = self.update_net(self.perceive(state)) + rate = self.fire_rate if fire_rate is None else fire_rate + stochastic = (torch.rand_like(state[:, :1]) <= rate).to(state.dtype) + state = state + delta * stochastic + post_life = self.living_mask(state) + return state * (pre_life & post_life).to(state.dtype) + + +def create_seed(batch_size: int, channels: int, resolution: int, device: torch.device) -> torch.Tensor: + state = torch.zeros(batch_size, channels, resolution, resolution, device=device) + center = resolution // 2 + state[:, 3, center, center] = 1.0 + return state + + +def create_model(config: dict) -> NeuralCellularAutomata: + return NeuralCellularAutomata( + channels=int(config.get("cell_channels", 16)), + hidden_size=int(config.get("hidden_size", 128)), + fire_rate=float(config.get("fire_rate", 0.5)), + ) + diff --git a/models/neural_cellular_automata/trainer.py b/models/neural_cellular_automata/trainer.py new file mode 100644 index 0000000000000000000000000000000000000000..78db98259fe3421a01e5562b025a9dffacc3e10a --- /dev/null +++ b/models/neural_cellular_automata/trainer.py @@ -0,0 +1,286 @@ +from __future__ import annotations + +import json +import random +from pathlib import Path + +import numpy as np +import torch +from PIL import Image +from torch.nn import functional as F + +from .model import create_model, create_seed +from .image_preprocessing import image_to_tensor, prepare_image, randomly_flip + + +EXTENSIONS = {".png", ".jpg", ".jpeg", ".webp"} + + +def _device(config: dict) -> torch.device: + requested = config.get("device", "auto") + if requested == "cuda" and not torch.cuda.is_available(): + raise RuntimeError("CUDA was selected, but PyTorch cannot access an NVIDIA CUDA device.") + return torch.device("cuda" if requested == "cuda" or (requested == "auto" and torch.cuda.is_available()) else "cpu") + + +def _load_targets( + folder: str, + resolution: int, + fit_mode: str, + background: str, + training_mode: str = "Dataset average (experimental)", + target_image: str = "", +) -> tuple[list[torch.Tensor], list[Path]]: + if training_mode == "Single target": + selected = Path(target_image).expanduser() if target_image else None + if selected is None or not selected.is_file(): + raise ValueError("Single Target Growth requires an existing target image. Choose one on the Dataset tab or in Target Image.") + if selected.suffix.lower() not in EXTENSIONS: + raise ValueError("The selected target must be a PNG, JPG, JPEG, or WEBP image.") + paths = [selected.resolve()] + elif training_mode == "Dataset average (experimental)": + paths = sorted(file for file in Path(folder).rglob("*") if file.suffix.lower() in EXTENSIONS) + else: + raise ValueError(f"Unsupported NCA training mode: {training_mode}") + targets = [] + loaded_paths = [] + for path in paths: + try: + with Image.open(path) as image: + prepared = prepare_image(image, resolution, fit_mode, background) + targets.append(image_to_tensor(prepared)) + loaded_paths.append(path) + except OSError: + continue + if not targets: + raise ValueError("No readable PNG, JPG, JPEG, or WEBP images were found in the dataset folder.") + return targets, loaded_paths + + +def _state_image(state: torch.Tensor) -> Image.Image: + rgba = state[0, :4].detach().float().clamp(0, 1).permute(1, 2, 0).cpu().numpy() + return Image.fromarray((rgba * 255).astype(np.uint8), "RGBA") + + +def _save_checkpoint(path, model, optimizer, scaler, config, losses, completed_step, best_smoothed_loss): + temporary = path.with_suffix(".pt.tmp") + torch.save( + { + "format_version": 2, + "model_state": model.state_dict(), + "optimizer_state": optimizer.state_dict(), + "scaler_state": scaler.state_dict(), + "config": dict(config), + "loss_history": list(losses), + "completed_step": completed_step, + "best_smoothed_loss": best_smoothed_loss, + "python_random_state": random.getstate(), + "numpy_random_state": np.random.get_state(), + "torch_random_state": torch.get_rng_state(), + }, + temporary, + ) + temporary.replace(path) + + +def train_model(config, dataset_path, output_path, progress_callback=None, preview_callback=None, stop_event=None): + output = Path(output_path) + output.mkdir(parents=True, exist_ok=True) + (output / "previews").mkdir(exist_ok=True) + + seed = int(config.get("random_seed", 42)) + random.seed(seed) + np.random.seed(seed) + torch.manual_seed(seed) + device = _device(config) + if device.type == "cuda": + torch.cuda.manual_seed_all(seed) + + resolution = int(config["resolution"]) + channels = int(config["cell_channels"]) + batch_size = int(config["batch_size"]) + total = int(config["epochs"]) + min_steps = int(config["min_growth_steps"]) + max_steps = int(config["max_growth_steps"]) + if min_steps > max_steps: + raise ValueError("Minimum growth steps cannot exceed maximum growth steps.") + stable_growth = bool(config.get("stable_growth", False)) + stable_max_steps = int(config.get("stable_max_growth_steps", max_steps)) + if stable_growth and stable_max_steps < max_steps: + raise ValueError("Stable maximum growth steps cannot be less than Maximum Growth Steps.") + training_max_steps = stable_max_steps if stable_growth else max_steps + + training_mode = str(config.get("training_mode", "Dataset average (experimental)")) + targets, target_paths = _load_targets( + dataset_path, + resolution, + str(config.get("fit_mode", "contain")), + str(config.get("background", "transparent")), + training_mode, + str(config.get("target_image", "")), + ) + model = create_model(config).to(device) + optimizer = torch.optim.Adam(model.parameters(), lr=float(config["learning_rate"])) + use_amp = bool(config.get("mixed_precision", True)) and device.type == "cuda" + scaler = torch.amp.GradScaler("cuda", enabled=use_amp) + losses: list[float] = [] + start_step = 1 + best_smoothed_loss = float("inf") + resume_path = str(config.get("resume_checkpoint", "")).strip() + if resume_path: + checkpoint = torch.load(resume_path, map_location=device, weights_only=False) + old_config = checkpoint.get("config", {}) + for key in ("resolution", "cell_channels", "hidden_size"): + if int(old_config.get(key, config[key])) != int(config[key]): + raise ValueError(f"Cannot resume because {key} differs from the checkpoint.") + model.load_state_dict(checkpoint["model_state"]) + if "optimizer_state" in checkpoint: + optimizer.load_state_dict(checkpoint["optimizer_state"]) + if "scaler_state" in checkpoint: + scaler.load_state_dict(checkpoint["scaler_state"]) + losses = list(checkpoint.get("loss_history", [])) + best_smoothed_loss = float(checkpoint.get("best_smoothed_loss", float("inf"))) + start_step = int(checkpoint.get("completed_step", len(losses))) + 1 + if "python_random_state" in checkpoint: + random.setstate(checkpoint["python_random_state"]) + if "numpy_random_state" in checkpoint: + np.random.set_state(checkpoint["numpy_random_state"]) + if "torch_random_state" in checkpoint: + torch.set_rng_state(checkpoint["torch_random_state"].cpu()) + log_path = output / "training_log.txt" + checkpoint_path = output / "checkpoint.pt" + best_checkpoint_path = output / "best_checkpoint.pt" + completed_step = start_step - 1 + + with log_path.open("a", encoding="utf-8", buffering=1) as log: + log.write(f"Training on {device}; mode={training_mode}; {len(targets)} target(s); mixed precision={use_amp}\n") + log.write(f"Growth rollout range: {min_steps}-{training_max_steps}; stable growth={stable_growth}\n") + if training_mode == "Single target": + log.write(f"Target image: {target_paths[0]}\n") + if resume_path: + log.write(f"Resumed from {resume_path} at step {start_step}.\n") + for step in range(start_step, total + 1): + if stop_event is not None and stop_event.is_set(): + log.write("Training stopped by user.\n") + break + selected = random.choices(targets, k=batch_size) + target = torch.stack(selected).to(device) + if randomly_flip(bool(config.get("horizontal_flip", False))): + target = torch.flip(target, dims=(3,)) + state = create_seed(batch_size, channels, resolution, device) + growth_steps = random.randint(min_steps, training_max_steps) + + optimizer.zero_grad(set_to_none=True) + with torch.amp.autocast(device_type=device.type, enabled=use_amp): + for _ in range(growth_steps): + state = model(state) + loss = F.mse_loss(state[:, :4], target) + scaler.scale(loss).backward() + scaler.unscale_(optimizer) + torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) + scaler.step(optimizer) + scaler.update() + + value = float(loss.detach().cpu()) + losses.append(value) + smoothing_window = max(1, int(config.get("loss_smoothing_window", 25))) + smoothed_loss = sum(losses[-smoothing_window:]) / min(len(losses), smoothing_window) + completed_step = step + update = {"step": step, "total_steps": total, "epoch": step, "total_epochs": total, "loss": value, "smoothed_loss": smoothed_loss, "device": str(device), "growth_steps": growth_steps} + if progress_callback: + progress_callback(update) + log.write(json.dumps(update) + "\n") + + preview_every = max(1, int(config.get("preview_every", 10))) + if step == 1 or step % preview_every == 0 or step == total: + preview = _state_image(state) + preview.save(output / "previews" / f"step_{step:06d}.png") + if preview_callback: + preview_callback(preview.copy()) + + checkpoint_every = max(1, int(config.get("checkpoint_every", 50))) + if step % checkpoint_every == 0: + _save_checkpoint(checkpoint_path, model, optimizer, scaler, config, losses, completed_step, best_smoothed_loss) + log.write(f"Saved recovery checkpoint at step {step}.\n") + if smoothed_loss < best_smoothed_loss: + best_smoothed_loss = smoothed_loss + _save_checkpoint(best_checkpoint_path, model, optimizer, scaler, config, losses, completed_step, best_smoothed_loss) + log.write(f"Saved best checkpoint at step {step}; smoothed loss={smoothed_loss:.6f}.\n") + + final_smoothed = sum(losses[-max(1, int(config.get("loss_smoothing_window", 25))):]) / min( + len(losses), max(1, int(config.get("loss_smoothing_window", 25))) + ) if losses else float("inf") + _save_checkpoint(checkpoint_path, model, optimizer, scaler, config, losses, completed_step, best_smoothed_loss) + if not best_checkpoint_path.exists() or final_smoothed < best_smoothed_loss: + best_smoothed_loss = final_smoothed + _save_checkpoint(best_checkpoint_path, model, optimizer, scaler, config, losses, completed_step, best_smoothed_loss) + return { + "checkpoint_path": str(checkpoint_path), + "best_checkpoint_path": str(best_checkpoint_path), + "best_smoothed_loss": best_smoothed_loss, + "loss_history": losses, + "completed_step": completed_step, + "stopped": bool(stop_event and stop_event.is_set()), + } + + +def train( + context, + dataset_dir: str, + model_name: str, + epochs: int, + output_dir: str, + resume_from: str = "", + **settings, +): + from .manifest import TRAINING_SETTINGS + + config = { + key: spec.get("default") + for key, spec in TRAINING_SETTINGS.items() + if isinstance(spec, dict) and "default" in spec + } + config.update(settings) + config["epochs"] = int(epochs) + if resume_from: + config["resume_checkpoint"] = resume_from + + output = Path(output_dir) + + def progress(update): + step = int(update.get("step", update.get("epoch", 0)) or 0) + total = int(update.get("total_steps", update.get("total_epochs", epochs)) or epochs) + loss = update.get("loss") + message = f"NCA step {step:,} of {total:,}" + if loss is not None: + message += f" · loss {float(loss):.5f}" + context.progress(max(1, min(99, round(step * 100 / max(total, 1)))), message) + + result = train_model( + config, + dataset_dir, + output, + progress_callback=progress, + stop_event=context.cancel_event, + ) + checkpoint = str(result.get("checkpoint_path") or output / "checkpoint.pt") + preview = output / "previews" + if preview.is_dir(): + latest = max(preview.glob("*.png"), key=lambda path: path.stat().st_mtime, default=None) + if latest: + context.preview(latest, epoch=int(result.get("completed_step", epochs) or epochs)) + context.progress(100, "NCA training completed") + return { + "output_folder": str(output), + "assets": [ + { + "kind": "model", + "name": model_name, + "path": str(output), + "trainer": "neural_cellular_automata", + "dataset_path": str(Path(dataset_dir).resolve()), + "checkpoint": checkpoint, + "epochs": int(epochs), + } + ], + } diff --git a/pytest.ini b/pytest.ini new file mode 100644 index 0000000000000000000000000000000000000000..69c9c4067ad54a4cc94770f463867c3ae40c199e --- /dev/null +++ b/pytest.ini @@ -0,0 +1,3 @@ +[pytest] +testpaths = tests +norecursedirs = build dist hf-release-staging hf-release-verify artifacts diff --git a/requirements.txt b/requirements.txt index dba51922acaee23346b48e22e852a54446ac748d..c8425a2ed0a0f4821c1efc212ddfc67b10d9e105 100644 --- a/requirements.txt +++ b/requirements.txt @@ -5,6 +5,7 @@ pytest>=8.0 yt-dlp>=2025.6.30 opencv-python>=4.10 Pillow>=10.0 +qrcode[pil]>=7.4 torch>=2.2 torchvision>=0.17 transformers>=4.45 diff --git a/tests/test_adam_manager_foundation.py b/tests/test_adam_manager_foundation.py new file mode 100644 index 0000000000000000000000000000000000000000..2f9f73a86cb3229a2cf9269d2cb2c896def218d9 --- /dev/null +++ b/tests/test_adam_manager_foundation.py @@ -0,0 +1,699 @@ +from __future__ import annotations + +from datetime import datetime, timedelta, timezone +from pathlib import Path +from urllib.error import HTTPError +from urllib.request import Request, urlopen +import json +import socket + +from PIL import Image + +from adam.assets import AssetRegistry +from adam.dataset_lab import scan_dataset +from adam.model_profiles import ModelProfileRegistry +from adam.recommendations import recommend_for_profile +from adam.registry import ToolRegistry +from adam.remote_access import ( + REMOTE_MODE_DISABLED, + REMOTE_MODE_TAILSCALE, + RemoteAccessService, + inspect_tailscale, + remote_scope, +) +from adam.models import ExecutionPlan, Job, JobStatus, PlanStep, SystemSnapshot +from adam.transcript_dataset import clean_transcript, split_transcript, transcript_videos_to_dataset + + +def test_model_profiles_expose_manifest_metadata() -> None: + profiles = {profile.id: profile for profile in ModelProfileRegistry(ToolRegistry(Path.cwd()).model_plugins).all()} + + assert profiles["ddpm"].hardware["recommended_vram_gb"] >= 6 + assert "resolution" in profiles["flow"].training + assert "safetensors" in profiles["lora"].output_formats + + +def test_profile_recommendation_uses_schema_and_warns_on_small_vram() -> None: + profile = ModelProfileRegistry(ToolRegistry(Path.cwd()).model_plugins).get("ddpm") + + result = recommend_for_profile(profile, dataset_items=1_000, resolution=128) # type: ignore[arg-type] + + assert result.epochs == 180 + assert result.settings["batch_size"] == 12 + assert result.estimated_vram_gb is not None + + +def test_dataset_lab_scans_images_captions_and_duplicates(tmp_path: Path) -> None: + image = Image.new("RGB", (16, 24), (10, 20, 30)) + first = tmp_path / "first.png" + second = tmp_path / "second.png" + image.save(first) + image.save(second) + first.with_suffix(".txt").write_text("caption", encoding="utf-8") + + report = scan_dataset(tmp_path) + + assert report.image_count == 2 + assert report.caption_count == 1 + assert report.missing_caption_count == 1 + assert report.duplicate_groups == 1 + assert report.dimensions["16x24"] == 2 + + +def test_transcript_cleaning_and_splitting() -> None: + text = "Hello world .\nThis is ADAM! " * 20 + + samples = split_transcript(clean_transcript(text), max_chars=80) + + assert len(samples) > 1 + assert all(len(sample) <= 90 for sample in samples) + + +def test_remote_access_defaults_to_disabled(tmp_path: Path) -> None: + class Config: + values = {} + + def get(self, key, default=None): + return self.values.get(key, default) + + def update(self, values): + self.values.update(values) + + service = RemoteAccessService(Config(), jobs=None, monitor=None) + + assert service.settings()["enabled"] is False + assert service.settings()["bind_address"] == "127.0.0.1" + assert service.settings()["remote_mode"] == "local_wifi" + assert service.phone_test_url() == "" + + +def test_remote_access_requires_token_and_reports_safe_permissions() -> None: + class Config: + def __init__(self) -> None: + self.values = {} + + def get(self, key, default=None): + return self.values.get(key, default) + + def update(self, values): + self.values.update(values) + + class Jobs: + active_job = None + jobs = [] + + class Monitor: + @staticmethod + def snapshot(): + return SystemSnapshot(cpu_percent=12, memory_percent=34) + + with socket.socket() as sock: + sock.bind(("127.0.0.1", 0)) + port = sock.getsockname()[1] + config = Config() + service = RemoteAccessService(config, Jobs(), Monitor()) + settings = service.settings() + service.save_settings({"enabled": True, "port": port, "token": settings["token"]}) + try: + url = service.start() + try: + urlopen(url, timeout=3) + except HTTPError as exc: + assert exc.code == 401 + else: + raise AssertionError("Remote status URL should require a token") + local_payload = urlopen(service.local_test_url(), timeout=3).read().decode("utf-8") + assert "ADAM Remote" in local_payload + request = Request(url, headers={"Authorization": f"Bearer {settings['token']}"}) + payload = urlopen(request, timeout=3).read().decode("utf-8") + assert '"dangerous_actions": false' in payload + assert '"cpu_percent": 12' in payload + finally: + service.stop() + + +def test_remote_access_scope_labels() -> None: + assert remote_scope("127.0.0.1") == "local-device only" + assert remote_scope("172.16.0.5") == "local network" + assert remote_scope("172.15.0.5") == "custom bind address" + assert remote_scope("0.0.0.0") == "all network interfaces" + + +def test_remote_access_phone_url_uses_network_bind_address() -> None: + class Config: + def __init__(self) -> None: + self.values = {} + + def get(self, key, default=None): + return self.values.get(key, default) + + def update(self, values): + self.values.update(values) + + service = RemoteAccessService(Config(), jobs=None, monitor=None) + token = service.settings()["token"] + service.save_settings({"enabled": True, "bind_address": "192.168.1.25", "port": 8765, "token": token}) + + assert service.phone_test_url() == f"http://192.168.1.25:8765/?token={token}" + + +def test_remote_mode_disabled_refuses_to_start() -> None: + class Config: + def __init__(self) -> None: + self.values = {} + + def get(self, key, default=None): + return self.values.get(key, default) + + def update(self, values): + self.values.update(values) + + service = RemoteAccessService(Config(), jobs=None, monitor=None) + service.save_settings({"enabled": True, "remote_mode": REMOTE_MODE_DISABLED}) + + try: + service.start() + except RuntimeError as exc: + assert "Remote Mode" in str(exc) + else: + raise AssertionError("Disabled remote mode should not start") + + +def test_tailscale_unavailable_and_available_states() -> None: + assert inspect_tailscale(which=lambda _name: None).installed is False + + class Result: + def __init__(self, returncode=0, stdout="", stderr=""): + self.returncode = returncode + self.stdout = stdout + self.stderr = stderr + + disconnected = inspect_tailscale( + which=lambda _name: "tailscale", + runner=lambda _command: Result(1, stderr="not logged in"), + ) + assert disconnected.installed is True + assert disconnected.connected is False + + def runner(command): + if command[1:3] == ["status", "--json"]: + return Result( + stdout=json.dumps( + { + "BackendState": "Running", + "Self": { + "HostName": "adam-pc", + "DNSName": "adam-pc.tailnet.ts.net.", + "TailscaleIPs": ["100.64.0.12", "fd7a:115c:a1e0::12"], + }, + } + ) + ) + return Result(stdout='{"TCP":{}}') + + connected = inspect_tailscale(which=lambda _name: "tailscale", runner=runner) + assert connected.installed is True + assert connected.connected is True + assert connected.device_name == "adam-pc" + assert connected.dns_name == "adam-pc.tailnet.ts.net" + assert connected.tailscale_ip == "100.64.0.12" + + +def test_remote_access_mobile_dashboard_prompt_and_preview(tmp_path: Path) -> None: + class Config: + def __init__(self) -> None: + self.values = {} + + def get(self, key, default=None): + return self.values.get(key, default) + + def update(self, values): + self.values.update(values) + + class Planner: + def plan(self, prompt: str) -> ExecutionPlan: + return ExecutionPlan( + request=prompt, + summary="Generate one image.", + steps=[PlanStep("preview_generator", "Generate preview", "Generate one preview.", {})], + project_name="Phone Prompt", + ) + + class Jobs: + def __init__(self) -> None: + preview = tmp_path / "preview.png" + Image.new("RGB", (8, 8), (10, 120, 240)).save(preview) + plan = ExecutionPlan( + "train", + "Training.", + [], + project_name="Phone Preview", + orion_review={"estimated_high_minutes": 20}, + ) + self.active_job = Job( + plan=plan, + status=JobStatus.RUNNING, + progress=55, + started_at=(datetime.now(timezone.utc) - timedelta(minutes=5)).isoformat(), + preview_path=str(preview), + preview_epoch=3, + preview_kind="training", + ) + self.jobs = [self.active_job] + self.submitted: list[ExecutionPlan] = [] + + def submit(self, plan: ExecutionPlan) -> Job: + self.submitted.append(plan) + job = Job(plan=plan, status=JobStatus.QUEUED) + self.jobs.insert(0, job) + return job + + class Monitor: + @staticmethod + def snapshot(): + return SystemSnapshot(cpu_percent=12, memory_percent=34, gpu_percent=56, vram_percent=78) + + with socket.socket() as sock: + sock.bind(("127.0.0.1", 0)) + port = sock.getsockname()[1] + jobs = Jobs() + config = Config() + service = RemoteAccessService(config, jobs, Monitor(), Planner()) + token = service.settings()["token"] + service.save_settings({"enabled": True, "port": port, "token": token}) + try: + service.start() + html = urlopen(service.local_test_url(), timeout=3).read().decode("utf-8") + assert "Prompt ADAM" in html + assert "Create Model" in html + assert "Collect Dataset" in html + assert "Quick" in html + assert "Auto-approve remote training" in html + assert "Keep screen updated" in html + assert "Time Left" in html + assert "Live Preview" in html + active_script = html.split("", 1)[0] + assert "XMLHttpRequest" in active_script + assert "window.addEventListener(\"error\"" in active_script + assert "fetch(" not in active_script + assert "async " not in active_script + assert "=>" not in active_script + assert "?." not in active_script + assert "??" not in active_script + + status = json.loads(urlopen(service.url() + f"?token={token}", timeout=3).read().decode("utf-8")) + assert status["preview"]["available"] is True + assert status["permissions"]["prompt"] is True + assert status["permissions"]["auto_approve_training"] is False + assert status["active_job"]["timing"]["remaining_seconds"] is not None + assert status["active_job"]["timing"]["estimate_label"] + + settings_request = Request( + f"http://127.0.0.1:{port}/api/remote-settings?token={token}", + data=json.dumps({"auto_approve_training": True}).encode("utf-8"), + headers={"Content-Type": "application/json"}, + method="POST", + ) + # The desktop owner must grant control before a phone may enable approval. + service.save_settings({"allow_job_control": True}) + settings_response = json.loads(urlopen(settings_request, timeout=3).read().decode("utf-8")) + assert settings_response["ok"] is True + status = json.loads(urlopen(service.url() + f"?token={token}", timeout=3).read().decode("utf-8")) + assert status["permissions"]["auto_approve_training"] is True + + preview = urlopen(f"http://127.0.0.1:{port}/api/preview?token={token}", timeout=3).read() + assert preview.startswith(b"\x89PNG") + + request = Request( + f"http://127.0.0.1:{port}/api/prompt?token={token}", + data=json.dumps({"prompt": "generate one preview"}).encode("utf-8"), + headers={"Content-Type": "application/json"}, + method="POST", + ) + response = json.loads(urlopen(request, timeout=3).read().decode("utf-8")) + assert response["ok"] is True + assert response["job_id"] + assert jobs.submitted[0].project_name == "Phone Prompt" + + message_request = Request( + f"http://127.0.0.1:{port}/api/prompt?token={token}", + data=json.dumps({"message": "generate one preview"}).encode("utf-8"), + headers={"Content-Type": "application/json"}, + method="POST", + ) + message_response = json.loads(urlopen(message_request, timeout=3).read().decode("utf-8")) + assert message_response["ok"] is True + assert jobs.submitted[0].project_name == "Phone Prompt" + finally: + service.stop() + + +def test_remote_job_actions_include_existing_controls() -> None: + class Config: + def __init__(self) -> None: + self.values = {} + + def get(self, key, default=None): + return self.values.get(key, default) + + def update(self, values): + self.values.update(values) + + plan = ExecutionPlan( + request="train", + summary="Training.", + steps=[PlanStep("ddpm_trainer", "Train model", "Run training")], + requires_confirmation=True, + project_name="Phone Controls", + ) + job = Job(plan=plan, status=JobStatus.AWAITING_CONFIRMATION, current_step=0) + + class Jobs: + active_job = job + jobs = [job] + + def __init__(self) -> None: + self.actions: list[tuple[str, str]] = [] + + def get(self, job_id: str) -> Job: + assert job_id == job.id + return job + + def confirm(self, job_id: str) -> None: + self.actions.append(("confirm", job_id)) + job.status = JobStatus.QUEUED + + def pause(self, job_id: str) -> None: + self.actions.append(("pause", job_id)) + job.status = JobStatus.PAUSED + + def resume(self, job_id: str) -> None: + self.actions.append(("resume", job_id)) + job.status = JobStatus.RUNNING + + def cancel(self, job_id: str) -> None: + self.actions.append(("cancel", job_id)) + job.status = JobStatus.CANCELLED + + def retry(self, job_id: str) -> Job: + self.actions.append(("retry", job_id)) + return Job(plan=plan, status=JobStatus.AWAITING_CONFIRMATION) + + def end_task(self, job_id: str) -> bool: + self.actions.append(("end", job_id)) + return True + + jobs = Jobs() + service = RemoteAccessService(Config(), jobs, monitor=None) + + summary = service._job_summary(job) + assert summary["current_step_title"] == "Train model" + + assert service.job_action(job.id, "confirm", True)["ok"] is True + assert service.job_action(job.id, "pause", True)["ok"] is True + assert service.job_action(job.id, "resume", True)["ok"] is True + assert service.job_action(job.id, "retry", True)["ok"] is True + assert service.job_action(job.id, "end", True)["ok"] is True + assert [action for action, _job_id in jobs.actions] == [ + "confirm", + "pause", + "resume", + "retry", + "end", + ] + assert service.job_action(job.id, "pause", False)["ok"] is False + + +def test_remote_prompt_queues_image_generation_before_conversation(tmp_path: Path) -> None: + class Config: + def __init__(self) -> None: + self.values = {"tool_folders": {}, "generation_settings": {}} + + def get(self, key, default=None): + return self.values.get(key, default) + + def update(self, values): + self.values.update(values) + + class Jobs: + def __init__(self) -> None: + self.active_job = None + self.jobs = [] + self.submitted: list[ExecutionPlan] = [] + + def submit(self, plan: ExecutionPlan) -> Job: + self.submitted.append(plan) + job = Job(plan=plan, status=JobStatus.QUEUED) + self.jobs.insert(0, job) + return job + + class Planner: + def __init__(self, assets: AssetRegistry, registry: ToolRegistry) -> None: + self.assets = assets + self.registry = registry + + def plan(self, _prompt: str) -> ExecutionPlan: + raise AssertionError("Image generation should not fall back to conversation planning.") + + model = tmp_path / "ddpm" / "output" / "Minecraft" + model.mkdir(parents=True) + (model / "model_index.json").write_text("{}", encoding="utf-8") + assets = AssetRegistry(tmp_path) + assets.register(kind="model", name="Minecraft", path=str(model), trainer="ddpm") + jobs = Jobs() + service = RemoteAccessService( + Config(), + jobs=jobs, + monitor=None, + planner=Planner(assets, ToolRegistry(Path.cwd())), + ) + + response = service.submit_prompt("Generate an image of Minecraft") + + assert response["ok"] is True + assert response["job_id"] + assert jobs.submitted[0].steps[0].tool_id == "ddpm_generator" + assert jobs.submitted[0].steps[0].arguments["model_name"] == "Minecraft" + assert jobs.submitted[0].steps[0].arguments["prompt"] == "Minecraft" + + +def test_remote_prompt_can_auto_approve_training_plan(tmp_path: Path) -> None: + class Config: + def __init__(self) -> None: + self.values = {"remote_access": {"auto_approve_training": True}} + + def get(self, key, default=None): + return self.values.get(key, default) + + def update(self, values): + self.values.update(values) + + class Planner: + def plan(self, prompt: str) -> ExecutionPlan: + return ExecutionPlan( + request=prompt, + summary="Train DDPM.", + steps=[PlanStep("ddpm_trainer", "Train", "Train a model.", {})], + requires_confirmation=True, + confirmation_reason="This starts training.", + project_name="Remote Training", + ) + + class Jobs: + def __init__(self) -> None: + self.active_job = None + self.jobs = [] + self.confirmed: list[str] = [] + + def submit(self, plan: ExecutionPlan) -> Job: + job = Job(plan=plan, status=JobStatus.AWAITING_CONFIRMATION) + self.jobs.insert(0, job) + return job + + def confirm(self, job_id: str) -> None: + self.confirmed.append(job_id) + self.jobs[0].status = JobStatus.QUEUED + + jobs = Jobs() + service = RemoteAccessService(Config(), jobs=jobs, monitor=None, planner=Planner()) + + response = service.submit_prompt("train a DDPM model") + + assert response["ok"] is True + assert response["auto_approved"] is True + assert response["requires_approval"] is False + assert jobs.confirmed == [response["job_id"]] + + +def test_remote_prompt_selects_flow_match_generation_model(tmp_path: Path) -> None: + class Config: + def __init__(self) -> None: + self.values = {"tool_folders": {}, "generation_settings": {}} + + def get(self, key, default=None): + return self.values.get(key, default) + + def update(self, values): + self.values.update(values) + + class Jobs: + def __init__(self) -> None: + self.active_job = None + self.jobs = [] + self.submitted: list[ExecutionPlan] = [] + + def submit(self, plan: ExecutionPlan) -> Job: + self.submitted.append(plan) + job = Job(plan=plan, status=JobStatus.QUEUED) + self.jobs.insert(0, job) + return job + + class Planner: + def __init__(self, assets: AssetRegistry, registry: ToolRegistry) -> None: + self.assets = assets + self.registry = registry + + def plan(self, _prompt: str) -> ExecutionPlan: + raise AssertionError("Flow image generation should not fall back to conversation planning.") + + model = tmp_path / "flow" / "output_flow_models" / "Minecraft Flow Match" + (model / "unet").mkdir(parents=True) + (model / "flow_model_info.json").write_text("{}", encoding="utf-8") + (model / "unet" / "config.json").write_text("{}", encoding="utf-8") + assets = AssetRegistry(tmp_path) + assets.register(kind="model", name="Minecraft Flow Match", path=str(model), trainer="flow") + jobs = Jobs() + service = RemoteAccessService( + Config(), + jobs=jobs, + monitor=None, + planner=Planner(assets, ToolRegistry(Path.cwd())), + ) + + response = service.submit_prompt("Generate an image of Minecraft Flow Match") + + assert response["ok"] is True + assert jobs.submitted[0].steps[0].tool_id == "flow_generator" + assert jobs.submitted[0].steps[0].arguments["model_name"] == "Minecraft Flow Match" + assert jobs.submitted[0].steps[0].arguments["sampler"] == "Heun" + + +def test_remote_status_serves_latest_generation_image(tmp_path: Path) -> None: + class Config: + def __init__(self) -> None: + self.values = {} + + def get(self, key, default=None): + return self.values.get(key, default) + + def update(self, values): + self.values.update(values) + + class Planner: + root = tmp_path + + def __init__(self) -> None: + self.registry = ToolRegistry(Path.cwd()) + + class Jobs: + active_job = None + jobs = [] + + folder = tmp_path / "data" / "generations" / "flow_generator" / "Minecraft Flow Match" + folder.mkdir(parents=True) + image = folder / "image_001.png" + second_image = folder / "image_002.png" + Image.new("RGB", (8, 8), (20, 160, 80)).save(image) + Image.new("RGB", (8, 8), (140, 60, 220)).save(second_image) + (folder / "generation_20260830_TEST.json").write_text( + json.dumps( + { + "provider_id": "flow_generator", + "provider_name": "Flow Matching Generator", + "model_name": "Minecraft Flow Match", + "model_path": "D:/Flow/Minecraft Flow Match", + "prompt": "Minecraft Flow Match", + "seed": 0, + "steps": 20, + "sampler": "Heun", + "aspect_ratio": "1:1 (Square)", + "images": [str(image), str(second_image)], + "created_at": "2026-08-30T00:00:00+00:00", + } + ), + encoding="utf-8", + ) + + with socket.socket() as sock: + sock.bind(("127.0.0.1", 0)) + port = sock.getsockname()[1] + config = Config() + service = RemoteAccessService(config, Jobs(), monitor=None, planner=Planner()) + token = service.settings()["token"] + service.save_settings({"enabled": True, "port": port, "token": token}) + try: + service.start() + status = json.loads(urlopen(service.url() + f"?token={token}", timeout=3).read().decode("utf-8")) + assert status["latest_generation"]["available"] is True + assert status["latest_generation"]["model_name"] == "Minecraft Flow Match" + assert status["latest_generation"]["image_count"] == 2 + assert status["latest_generation"]["images"][1]["url"].endswith("image=1") + + generated = urlopen( + f"http://127.0.0.1:{port}/api/generation-image?record=0&image=0&token={token}", + timeout=3, + ).read() + assert generated.startswith(b"\x89PNG") + generated_second = urlopen( + f"http://127.0.0.1:{port}/api/generation-image?record=0&image=1&token={token}", + timeout=3, + ).read() + assert generated_second.startswith(b"\x89PNG") + finally: + service.stop() + + +def test_remote_job_controls_are_gated(tmp_path: Path) -> None: + class Config: + def __init__(self) -> None: + self.values = {} + + def get(self, key, default=None): + return self.values.get(key, default) + + def update(self, values): + self.values.update(values) + + class Jobs: + def __init__(self) -> None: + self.job = Job( + plan=ExecutionPlan("run", "Run.", [], project_name="Remote Job"), + status=JobStatus.QUEUED, + ) + self.jobs = [self.job] + self.active_job = None + self.cancelled: list[str] = [] + + def get(self, job_id: str): + return self.job if job_id == self.job.id else None + + def cancel(self, job_id: str) -> None: + self.cancelled.append(job_id) + + def retry(self, job_id: str) -> Job: + return Job(plan=self.job.plan, status=JobStatus.QUEUED) + + jobs = Jobs() + service = RemoteAccessService(Config(), jobs=jobs, monitor=None) + assert service.job_action(jobs.job.id, "cancel", False)["ok"] is False + assert service.job_action(jobs.job.id, "cancel", True)["ok"] is True + assert jobs.cancelled == [jobs.job.id] + + +def test_transcript_export_reports_missing_ffmpeg(monkeypatch, tmp_path: Path) -> None: + monkeypatch.setattr("adam.transcript_dataset.shutil.which", lambda _name: None) + + result = transcript_videos_to_dataset([tmp_path / "video.mp4"], tmp_path / "out") + + assert result.available is False + assert "FFmpeg" in result.message diff --git a/tests/test_agents.py b/tests/test_agents.py index 311a6df641fa71539adf8def3d5a8d45f9463f40..c5c6879d11646621b7d0e550965baaa9a2693536 100644 --- a/tests/test_agents.py +++ b/tests/test_agents.py @@ -30,6 +30,39 @@ def _training_job(tmp_path: Path) -> Job: ) +def test_atlas_checks_training_output_drive(tmp_path: Path, monkeypatch) -> None: + job = _training_job(tmp_path) + job.plan.steps[0].arguments['output_dir'] = str(tmp_path / 'future-output') + checked = [] + + def usage(path): + checked.append(path) + return SimpleNamespace(free=512 * 1024**2) + + monkeypatch.setattr('adam.atlas.shutil.disk_usage', usage) + atlas = AtlasSupervisor() + snapshot = SystemSnapshot(disk_total_gb=1000, disk_used_gb=100) + atlas.observe(job, snapshot, now=0) + result = atlas.observe(job, snapshot, now=1) + assert result.action == 'pause' + assert 'output drive' in result.message + assert checked == [tmp_path, tmp_path] + assert snapshot.disk_used_gb == 100 + + +def test_atlas_reports_unavailable_output_space(tmp_path: Path, monkeypatch) -> None: + job = _training_job(tmp_path) + job.plan.steps[0].arguments['output_dir'] = str(tmp_path) + + def unavailable(_path): + raise OSError('disconnected') + + monkeypatch.setattr('adam.atlas.shutil.disk_usage', unavailable) + result = AtlasSupervisor().observe(job, SystemSnapshot(disk_total_gb=100, disk_used_gb=1)) + assert result.severity == 'warning' + assert 'could not check' in result.message + + def test_orion_flags_small_dataset_preset_applied_to_large_dataset(tmp_path: Path) -> None: dataset = tmp_path / "dataset" dataset.mkdir() @@ -107,6 +140,28 @@ def test_orion_recipe_becomes_more_conservative_at_high_resolution() -> None: assert recipe["settings"]["workers"] == recipe["settings"]["dataloader_num_workers"] +def test_orion_accepts_oasis_wide_resolution(tmp_path: Path) -> None: + dataset = tmp_path / "oasis-dataset" + dataset.mkdir() + for index in range(20): + (dataset / f"{index}.png").touch() + plan = ExecutionPlan( + request="train oasis", + summary="Train Oasis.", + steps=[PlanStep("oasis_trainer", "Train Oasis", "Train", { + "dataset_dir": str(dataset), + "epochs": 5, + "batch_size": 2, + "resolution": "256x144", + })], + ) + + report = apply_orion_review(plan) + + assert report["estimated_optimizer_steps"] == 50 + assert "ORION" in plan.summary + + def test_atlas_pauses_for_a_new_non_finite_loss_message(tmp_path: Path) -> None: job = _training_job(tmp_path) job.logs.append("loss = NaN") diff --git a/tests/test_ddpm_adapter.py b/tests/test_ddpm_adapter.py new file mode 100644 index 0000000000000000000000000000000000000000..917660ea34903e6ddd99eb13db62b8c1c97af734 --- /dev/null +++ b/tests/test_ddpm_adapter.py @@ -0,0 +1,79 @@ +from __future__ import annotations + +import io +import threading +from pathlib import Path + +from adam.config import ConfigManager +from adam.executor import ToolContext +from adam.registry import ToolSpec +from adam.tools import ddpm_adapter + + +def _context(root: Path, logs: list[str]) -> ToolContext: + running = threading.Event() + running.set() + return ToolContext( + root=root, + job_id="DDPMTEST", + tool=ToolSpec("ddpm_trainer", "DDPM Trainer", "test", "Training", "train_ddpm"), + cancel_event=threading.Event(), + run_event=running, + progress_callback=lambda *_args: None, + log_callback=logs.append, + ) + + +def test_resolution_change_branches_from_pipeline_instead_of_resuming_checkpoint(tmp_path: Path, monkeypatch) -> None: + trainer_root = tmp_path / "DDPM" + model = trainer_root / "output" / "Anime" + checkpoint = model / "checkpoint-40" + dataset = tmp_path / "dataset" + for folder in (checkpoint / "unet", model / "unet", model / "scheduler", dataset): + folder.mkdir(parents=True) + (trainer_root / "train.py").write_text("# fake trainer", encoding="utf-8") + (model / "model_index.json").write_text("{}", encoding="utf-8") + (checkpoint / "optimizer.bin").write_bytes(b"optimizer") + (checkpoint / "scheduler.bin").write_bytes(b"scheduler") + (checkpoint / "unet" / "diffusion_pytorch_model.safetensors").write_bytes(b"weights") + (checkpoint / "unet" / "config.json").write_text('{"sample_size": 64}', encoding="utf-8") + for index in range(2): + (dataset / f"{index}.png").write_bytes(b"image") + + ConfigManager(tmp_path).update({"tool_folders": {"ddpm_trainer": str(trainer_root)}}) + monkeypatch.setattr(ddpm_adapter.importlib.util, "find_spec", lambda _name: object()) + commands: list[list[str]] = [] + + class FakeProcess: + stdout = io.StringIO("PROGRESS_JSON:{\"event\":\"done\"}\n") + returncode = 0 + + def poll(self): + return 0 + + def terminate(self): + self.returncode = -15 + + def fake_popen(command, **_kwargs): + commands.append(command) + return FakeProcess() + + monkeypatch.setattr(ddpm_adapter.subprocess, "Popen", fake_popen) + + logs: list[str] = [] + result = ddpm_adapter.train_ddpm( + _context(tmp_path, logs), + dataset_dir=str(dataset), + model_name="Anime", + epochs=5, + output_dir=str(model), + resume_from=str(checkpoint), + resolution=256, + ) + + command = commands[0] + assert "--pretrained_model_path" in command + assert str(model.resolve()) in command + assert "--resume_from_checkpoint" not in command + assert result["output_folder"] != str(model.resolve()) + assert any("Changing DDPM resolution from 64px to 256px" in line for line in logs) diff --git a/tests/test_experiment_tracker.py b/tests/test_experiment_tracker.py new file mode 100644 index 0000000000000000000000000000000000000000..cc7f802f2a1aeb56f55dc5208cb301a8f4c86bc4 --- /dev/null +++ b/tests/test_experiment_tracker.py @@ -0,0 +1,108 @@ +from __future__ import annotations + +from datetime import datetime, timedelta, timezone +from pathlib import Path +import sqlite3 + +from adam.experiment_tracker import ExperimentStore +from adam.models import ExecutionPlan, Job, JobStatus, PlanStep, SystemSnapshot + + +def test_experiment_store_records_training_job_and_clone_request(tmp_path: Path) -> None: + dataset = tmp_path / "dataset" + dataset.mkdir() + for index in range(3): + (dataset / f"{index}.png").write_bytes(b"image") + output = tmp_path / "output" + output.mkdir() + checkpoint = output / "model.safetensors" + checkpoint.write_bytes(b"weights") + now = datetime.now(timezone.utc) + job = Job( + id="ABC123", + plan=ExecutionPlan( + request="train", + summary="Train", + steps=[ + PlanStep( + "ddpm_trainer", + "Train DDPM", + "Train", + { + "dataset_dir": str(dataset), + "model_name": "Demo Model", + "epochs": 12, + "output_dir": str(output), + "resolution": 128, + "batch_size": 2, + "learning_rate": 0.0001, + "preview_seed": 44, + }, + ) + ], + ), + status=JobStatus.FINISHED, + started_at=(now - timedelta(minutes=5)).isoformat(), + ended_at=now.isoformat(), + output_folder=str(output), + logs=["loss: 0.25"], + ) + + store = ExperimentStore(tmp_path) + run = store.record_job(job, SystemSnapshot(gpu_name="Test GPU", vram_used_gb=4, vram_total_gb=8)) + + assert run is not None + assert store.list_runs()[0].dataset_item_count == 3 + assert store.list_runs()[0].final_loss == 0.25 + request = store.clone_request("EXP-ABC123") + assert "Demo Model Clone" in request + assert "dataset_dir" not in request + + +def test_experiment_store_updates_notes_and_compare(tmp_path: Path) -> None: + store = ExperimentStore(tmp_path) + for suffix in ("A", "B"): + job = Job( + id=f"JOB{suffix}", + plan=ExecutionPlan( + request="train", + summary="Train", + steps=[ + PlanStep( + "flow_trainer", + "Train", + "Train", + {"dataset_dir": str(tmp_path), "model_name": suffix, "epochs": 5, "output_dir": str(tmp_path)}, + ) + ], + ), + status=JobStatus.FINISHED, + ended_at=datetime.now(timezone.utc).isoformat(), + ) + store.record_job(job) + + store.update_notes("EXP-JOBA", "best so far", 89) + + assert store.get("EXP-JOBA").notes == "best so far" # type: ignore[union-attr] + comparison = store.compare(["EXP-JOBA", "EXP-JOBB"]) + assert any(row["field"] == "quality_score" and row["EXP-JOBA"] == 89 for row in comparison) + + +def test_experiment_store_migrates_older_sqlite_schema(tmp_path: Path) -> None: + path = tmp_path / "data" / "experiments.sqlite3" + path.parent.mkdir() + with sqlite3.connect(path) as db: + db.execute( + "CREATE TABLE experiments (id TEXT PRIMARY KEY, job_id TEXT UNIQUE NOT NULL, timestamp TEXT NOT NULL, model_architecture TEXT NOT NULL, model_name TEXT NOT NULL)" + ) + db.execute( + "INSERT INTO experiments (id, job_id, timestamp, model_architecture, model_name) VALUES ('EXP-OLD', 'OLD', '2026-01-01T00:00:00+00:00', 'ddpm', 'Old Run')" + ) + + store = ExperimentStore(tmp_path) + run = store.get("EXP-OLD") + + assert run is not None + assert run.dataset_path == "" + assert run.checkpoint_paths == [] + assert run.settings == {} diff --git a/tests/test_generations.py b/tests/test_generations.py index ea4376d91bdf79fb706494cb030940c0e6b28568..a9ca379d6d794fcae56ea82ae0b87b948d73fac0 100644 --- a/tests/test_generations.py +++ b/tests/test_generations.py @@ -11,7 +11,9 @@ from adam.generations import ( build_generation_plan, generation_output_folder, generation_model_match_score, + generation_model_key, generation_tools, + group_generation_records, load_generation_history, parse_chat_generation_request, combine_generation_plans, @@ -20,6 +22,7 @@ from adam.registry import ToolRegistry from adam.config import ConfigManager from adam.executor import ToolContext from adam.generation_previews import accepts_preview_callback, publish_generation_preview +from adam.image_preferences import PreferenceScore from adam.registry import ToolSpec from adam.tools import ddpm_generator from adam.tools import flow_generator @@ -391,6 +394,47 @@ def test_generation_history_reads_new_model_folder_layout(tmp_path: Path) -> Non assert records[0].images == (image,) +def test_generation_history_groups_model_folders_with_latest_image_cover(tmp_path: Path) -> None: + mario_model = tmp_path / "models" / "Mario V2" + luigi_model = tmp_path / "models" / "Luigi" + mario_model.mkdir(parents=True) + luigi_model.mkdir(parents=True) + mario_folder = generation_output_folder(tmp_path, "ddpm_generator", "Mario V2") + luigi_folder = generation_output_folder(tmp_path, "flow_generator", "Luigi") + + def add_batch(folder: Path, stamp: str, model_name: str, model_path: Path, images: int) -> list[Path]: + paths = [] + for index in range(images): + path = folder / f"{stamp}_{index}.png" + path.write_bytes(b"image") + paths.append(path) + (folder / f"generation_{stamp}.json").write_text(json.dumps({ + "provider_id": "ddpm_generator" if "Mario" in model_name else "flow_generator", + "provider_name": "DDPM Generator" if "Mario" in model_name else "Flow Generator", + "model_name": model_name, + "model_path": str(model_path), + "images": [str(path) for path in paths], + "created_at": stamp, + }), encoding="utf-8") + return paths + + old_mario = add_batch(mario_folder, "2026-08-01", "Mario V2", mario_model, 2) + new_mario = add_batch(mario_folder, "2026-08-03", "Mario V2", mario_model, 1) + add_batch(luigi_folder, "2026-08-02", "Luigi", luigi_model, 1) + + records = load_generation_history(tmp_path) + folders = group_generation_records(records) + + assert [folder.model_name for folder in folders] == ["Mario V2", "Luigi"] + assert folders[0].image_count == 3 + assert folders[0].cover_image == new_mario[0] + assert tuple(record.created_at for record in folders[0].records) == ( + "2026-08-03", "2026-08-01", + ) + assert generation_model_key(folders[0].records[0]).startswith("path:") + assert old_mario[0] in folders[0].records[1].images + + def test_ddpm_adapter_writes_images_and_reproducibility_metadata( tmp_path: Path, monkeypatch ) -> None: @@ -541,6 +585,103 @@ def test_ddpm_adapter_passes_enabled_custom_dimensions(tmp_path: Path, monkeypat assert metadata["height"] == 192 +def test_ddpm_smart_generation_stops_after_enough_passing_candidates(tmp_path: Path, monkeypatch) -> None: + trainer_root = tmp_path / "connected-ddpm" + model = trainer_root / "output" / "Mario" + model.mkdir(parents=True) + (model / "model_index.json").write_text("{}", encoding="utf-8") + script = trainer_root / "appStableDiffusion.py" + script.write_text("# test backend", encoding="utf-8") + ConfigManager(tmp_path).update({"tool_folders": {"ddpm_trainer": str(trainer_root)}}) + + class FakeImage: + def save(self, path: Path, *, format: str) -> None: + path.write_bytes(b"fake png") + + backend = ModuleType("fake_ddpm_smart") + backend.generate_images = lambda *_args, **_settings: [FakeImage()] # type: ignore[attr-defined] + monkeypatch.setattr(ddpm_generator, "_backend_module", backend) + monkeypatch.setattr(ddpm_generator, "_backend_script", script.resolve()) + monkeypatch.setattr(ddpm_generator.importlib.util, "find_spec", lambda _name: object()) + + class FakeVision: + def unload(self) -> None: + pass + + class FakeEvaluator: + def __init__(self, _root): + self.vision = FakeVision() + + def score(self, _profile, paths, *, keep_threshold=None, reject_threshold=None): + seed = int(str(paths[0]).rsplit("_seed_", 1)[1].split(".", 1)[0]) + score = {10: 0.2, 11: 0.8, 12: 0.9}.get(seed, 0.1) + category = "Strong Keep" if score >= float(keep_threshold or 0.7) else "Likely Reject" + return [PreferenceScore(str(Path(paths[0]).resolve()), score, score, category)] + + monkeypatch.setattr(ddpm_generator, "GenerationPreferenceEvaluator", FakeEvaluator) + event = threading.Event(); event.set() + context = ToolContext(tmp_path, "SMARTDDPM", ToolSpec("ddpm_generator", "DDPM Generator", "test", "Output", "generate_ddpm_images"), threading.Event(), event, lambda *_args: None, lambda *_args: None, 0) + + result = ddpm_generator.generate_ddpm_images( + context, "Mario", str(model), "Smart study", 2, 20, 10, "DDIM", + "1:1 (Square)", smart_generation=True, smart_wanted_results=2, + smart_max_candidates=5, smart_min_score=0.7, + ) + + output = Path(str(result["output_folder"])) + metadata = json.loads(next(output.glob("generation_*.json")).read_text(encoding="utf-8")) + assert metadata["smart_generation"]["selected_count"] == 2 + assert metadata["smart_generation"]["candidate_count"] == 3 + assert len(metadata["images"]) == 3 + assert "_seed_11" in metadata["images"][0] + assert "_seed_12" in metadata["images"][1] + + +def test_ddpm_smart_generation_reports_insufficient_passing_candidates(tmp_path: Path, monkeypatch) -> None: + trainer_root = tmp_path / "connected-ddpm" + model = trainer_root / "output" / "Mario" + model.mkdir(parents=True) + (model / "model_index.json").write_text("{}", encoding="utf-8") + script = trainer_root / "appStableDiffusion.py" + script.write_text("# test backend", encoding="utf-8") + ConfigManager(tmp_path).update({"tool_folders": {"ddpm_trainer": str(trainer_root)}}) + + class FakeImage: + def save(self, path: Path, *, format: str) -> None: + path.write_bytes(b"fake png") + + backend = ModuleType("fake_ddpm_smart_insufficient") + backend.generate_images = lambda *_args, **_settings: [FakeImage()] # type: ignore[attr-defined] + monkeypatch.setattr(ddpm_generator, "_backend_module", backend) + monkeypatch.setattr(ddpm_generator, "_backend_script", script.resolve()) + monkeypatch.setattr(ddpm_generator.importlib.util, "find_spec", lambda _name: object()) + + class FakeVision: + def unload(self) -> None: + pass + + class FakeEvaluator: + def __init__(self, _root): + self.vision = FakeVision() + + def score(self, _profile, paths, *, keep_threshold=None, reject_threshold=None): + return [PreferenceScore(str(Path(paths[0]).resolve()), 0.4, 0.6, "Needs Review")] + + monkeypatch.setattr(ddpm_generator, "GenerationPreferenceEvaluator", FakeEvaluator) + event = threading.Event(); event.set() + context = ToolContext(tmp_path, "SMARTLOW", ToolSpec("ddpm_generator", "DDPM Generator", "test", "Output", "generate_ddpm_images"), threading.Event(), event, lambda *_args: None, lambda *_args: None, 0) + + result = ddpm_generator.generate_ddpm_images( + context, "Mario", str(model), "Smart study", 3, 20, 20, "DDIM", + "1:1 (Square)", smart_generation=True, smart_wanted_results=3, + smart_max_candidates=4, smart_min_score=0.7, + ) + + metadata = json.loads(next(Path(str(result["output_folder"])).glob("generation_*.json")).read_text(encoding="utf-8")) + assert metadata["smart_generation"]["selected_count"] == 0 + assert metadata["smart_generation"]["candidate_count"] == 4 + + def test_flow_adapter_uses_registered_model_and_tracks_each_seed( tmp_path: Path, monkeypatch ) -> None: @@ -618,3 +759,59 @@ def test_flow_adapter_uses_registered_model_and_tracks_each_seed( assert [call["seed"] for call in calls] == [700, 701] assert all(call["method"] == "Heun" for call in calls) assert all(call["aspect_ratio"] == "4:3 (Landscape)" for call in calls) + + +def test_flow_smart_generation_can_return_top_ranked_pool(tmp_path: Path, monkeypatch) -> None: + flow_root = tmp_path / "connected-flow" + model = flow_root / "output_flow_models" / "Rooms" + (model / "unet").mkdir(parents=True) + (model / "unet" / "config.json").write_text("{}", encoding="utf-8") + (model / "flow_model_info.json").write_text( + json.dumps({"model_type": "rectified_flow", "model_name": "Rooms"}), + encoding="utf-8", + ) + script = flow_root / "flow_matching_app.py" + script.write_text("# test backend", encoding="utf-8") + ConfigManager(tmp_path).update({"tool_folders": {"flow_trainer": str(flow_root)}}) + + class FakeImage: + def save(self, path: Path, *, format: str) -> None: + path.write_bytes(b"fake flow png") + + backend = ModuleType("fake_flow_smart") + backend.load_unet = lambda *_args, **_kwargs: object() # type: ignore[attr-defined] + backend.sample_flow = lambda _model, _count, steps, _device, _dtype, seed, method, progress, **_settings: (progress(steps, steps), [FakeImage()])[1] # type: ignore[attr-defined] + monkeypatch.setattr(flow_generator, "_backend_module", backend) + monkeypatch.setattr(flow_generator, "_backend_script", script.resolve()) + monkeypatch.setattr(flow_generator, "_loaded_model", None) + monkeypatch.setattr(flow_generator, "_loaded_model_path", None) + monkeypatch.setattr(flow_generator.importlib.util, "find_spec", lambda _name: object()) + + class FakeVision: + def unload(self) -> None: + pass + + class FakeEvaluator: + def __init__(self, _root): + self.vision = FakeVision() + + def score(self, _profile, paths, *, keep_threshold=None, reject_threshold=None): + seed = int(str(paths[0]).rsplit("_seed_", 1)[1].split(".", 1)[0]) + score = {50: 0.1, 51: 0.9, 52: 0.3, 53: 0.8}[seed] + return [PreferenceScore(str(Path(paths[0]).resolve()), score, score, "Strong Keep" if score >= 0.7 else "Needs Review")] + + monkeypatch.setattr(flow_generator, "GenerationPreferenceEvaluator", FakeEvaluator) + event = threading.Event(); event.set() + context = ToolContext(tmp_path, "SMARTFLOW", ToolSpec("flow_generator", "Flow Matching Generator", "test", "Output", "generate_flow_images"), threading.Event(), event, lambda *_args: None, lambda *_args: None, 0) + + result = flow_generator.generate_flow_images( + context, "Rooms", str(model), "Room study", 2, 8, 50, "Heun", + "4:3 (Landscape)", smart_generation=True, smart_wanted_results=2, + smart_max_candidates=4, smart_min_score=0.7, smart_mode="top_n", + ) + + metadata = json.loads(next(Path(str(result["output_folder"])).glob("generation_*.json")).read_text(encoding="utf-8")) + assert metadata["smart_generation"]["candidate_count"] == 4 + assert metadata["smart_generation"]["selected_count"] == 2 + assert "_seed_51" in metadata["images"][0] + assert "_seed_53" in metadata["images"][1] diff --git a/tests/test_image_preferences.py b/tests/test_image_preferences.py new file mode 100644 index 0000000000000000000000000000000000000000..d73161a32445a5690b9dfc94ac7b5f2c042ff2e5 --- /dev/null +++ b/tests/test_image_preferences.py @@ -0,0 +1,72 @@ +from __future__ import annotations + +from pathlib import Path + +from adam.image_preferences import GenerationPreferenceEvaluator, PreferenceProfile + + +def _image(path: Path) -> Path: + path.write_bytes(b"image") + return path + + +def test_generation_ratings_save_load_and_manual_correction(tmp_path: Path) -> None: + image = _image(tmp_path / "good.png") + profile = PreferenceProfile(tmp_path, "ddpm_generator", "SML", str(tmp_path / "model")) + + profile.set_rating(image, "keep", seed=12, sampler="DDIM", steps=50) + loaded = PreferenceProfile(tmp_path, "ddpm_generator", "SML", str(tmp_path / "model")) + + assert loaded.rating_for(image).rating == "keep" # type: ignore[union-attr] + loaded.set_rating(image, "reject", seed=12, sampler="DDIM", steps=50) + + corrected = PreferenceProfile(tmp_path, "ddpm_generator", "SML", str(tmp_path / "model")) + assert corrected.rating_for(image).rating == "reject" # type: ignore[union-attr] + + +def test_preference_profiles_are_model_specific(tmp_path: Path) -> None: + image = _image(tmp_path / "sample.png") + one = PreferenceProfile(tmp_path, "ddpm_generator", "SML", str(tmp_path / "sml")) + two = PreferenceProfile(tmp_path, "ddpm_generator", "Minecraft", str(tmp_path / "minecraft")) + + one.set_rating(image, "favorite") + + assert one.id != two.id + assert PreferenceProfile(tmp_path, "ddpm_generator", "SML", str(tmp_path / "sml")).rating_for(image) + assert PreferenceProfile(tmp_path, "ddpm_generator", "Minecraft", str(tmp_path / "minecraft")).rating_for(image) is None + + +def test_preference_scoring_ranks_candidates_with_positive_and_negative_examples(tmp_path: Path) -> None: + favorite = _image(tmp_path / "favorite.png") + rejected = _image(tmp_path / "rejected.png") + close = _image(tmp_path / "close.png") + far = _image(tmp_path / "far.png") + vectors = { + str(favorite): [[1.0, 0.0]], + str(rejected): [[0.0, 1.0]], + str(close): [[0.95, 0.05]], + str(far): [[0.05, 0.95]], + } + profile = PreferenceProfile(tmp_path, "flow_generator", "Rooms", str(tmp_path / "flow-model")) + profile.set_rating(favorite, "favorite") + profile.set_rating(rejected, "reject") + evaluator = GenerationPreferenceEvaluator(tmp_path, embedder=lambda paths: [vectors[str(Path(path))][0] for path in paths]) + + scores = evaluator.score(profile, [close, far], keep_threshold=0.70, reject_threshold=0.30) + + assert scores[0].score is not None + assert scores[1].score is not None + assert scores[0].score > scores[1].score # type: ignore[operator] + assert scores[0].category == "Strong Keep" + assert scores[1].category == "Likely Reject" + + +def test_scoring_without_profile_signal_needs_review(tmp_path: Path) -> None: + image = _image(tmp_path / "candidate.png") + profile = PreferenceProfile(tmp_path, "ddpm_generator", "Empty", str(tmp_path / "empty")) + evaluator = GenerationPreferenceEvaluator(tmp_path, embedder=lambda _paths: [[1.0, 0.0]]) + + scores = evaluator.score(profile, [image]) + + assert scores[0].score is None + assert scores[0].category == "Needs Review" diff --git a/tests/test_job_manager.py b/tests/test_job_manager.py index 30f08fb2a7b9c355a433c0d7df58db4054531a7c..22d23ca3b250d4aae2ad443f04f7a95461a6ea2b 100644 --- a/tests/test_job_manager.py +++ b/tests/test_job_manager.py @@ -1,10 +1,59 @@ from __future__ import annotations import logging +from datetime import datetime, timedelta, timezone from pathlib import Path -from adam.job_manager import JobManager -from adam.models import ExecutionPlan, Job, JobStatus +from adam.job_manager import JobManager, JobWorker +from adam.models import ExecutionPlan, Job, JobStatus, PlanStep +from adam.training_assistant import append_preflight_summary + + +def test_submission_reviews_training_before_queueing(tmp_path: Path, monkeypatch) -> None: + manager = JobManager(tmp_path, None, logging.getLogger("test.jobs")) + starts = [] + monkeypatch.setattr(manager, "_start_next", lambda: starts.append(True)) + dataset = tmp_path / "planned_dataset" + plan = ExecutionPlan( + request="collect and train", summary="Collect and train.", + steps=[ + PlanStep("dataset_collector", "Collect", "Collect", { + "output_dir": str(dataset), "image_count": 2000, + }), + PlanStep("ddpm_trainer", "Train", "Train", { + "dataset_dir": str(dataset), "epochs": 600, "batch_size": 1, + }), + ], + ) + + job = manager.submit(plan) + + assert job.status == JobStatus.AWAITING_CONFIRMATION + assert plan.orion_review["level"] == "warning" + assert "Pre-flight:" in plan.summary + assert starts == [] + assert manager._queue == [] + restored = JobManager(tmp_path, None, logging.getLogger("test.jobs")) + assert restored.jobs[0].plan.orion_review == plan.orion_review + assert restored.jobs[0].status == JobStatus.AWAITING_CONFIRMATION + + +def test_submission_preserves_an_already_reviewed_plan(tmp_path: Path) -> None: + manager = JobManager(tmp_path, None, logging.getLogger("test.jobs")) + plan = ExecutionPlan( + request="train", summary="Train.", requires_confirmation=True, + steps=[PlanStep("ddpm_trainer", "Train", "Train", {"epochs": 10})], + ) + append_preflight_summary(plan, {}) + summary = plan.summary + arguments = dict(plan.steps[0].arguments) + + job = manager.submit(plan) + + assert job.plan.summary == summary + assert job.plan.steps[0].arguments == arguments + assert summary.count("ORION —") == 1 + assert summary.count("Pre-flight:") == 1 def _job(index: int, status: JobStatus = JobStatus.FINISHED) -> Job: @@ -58,3 +107,210 @@ def test_end_task_acknowledges_an_interrupted_job(tmp_path: Path) -> None: restored = JobManager(tmp_path, None, logging.getLogger("test.jobs")) # type: ignore[arg-type] assert restored.jobs[0].status == JobStatus.CANCELLED + + +def test_approved_future_job_stays_scheduled_and_survives_restart(tmp_path: Path) -> None: + manager = JobManager(tmp_path, None, logging.getLogger("test.jobs")) # type: ignore[arg-type] + plan = ExecutionPlan( + request="train later", summary="Scheduled training", steps=[], + requires_confirmation=True, + ) + start = (datetime.now(timezone.utc) + timedelta(hours=2)).isoformat() + + job = manager.submit(plan, scheduled_for=start) + assert job.status == JobStatus.AWAITING_CONFIRMATION + manager.confirm(job.id) + assert job.status == JobStatus.SCHEDULED + + restored = JobManager(tmp_path, None, logging.getLogger("test.jobs")) # type: ignore[arg-type] + assert restored.jobs[0].status == JobStatus.SCHEDULED + assert restored.jobs[0].scheduled_for == start + + +def test_due_schedule_queues_behind_an_active_job(tmp_path: Path) -> None: + manager = JobManager(tmp_path, None, logging.getLogger("test.jobs")) # type: ignore[arg-type] + job = _job(2, JobStatus.SCHEDULED) + job.scheduled_for = (datetime.now(timezone.utc) - timedelta(minutes=1)).isoformat() + manager.jobs = [job] + + class BusyWorker: + @staticmethod + def isRunning() -> bool: + return True + + manager._worker = BusyWorker() # type: ignore[assignment] + manager._release_due_scheduled() + + assert job.status == JobStatus.QUEUED + assert manager._queue == [job.id] + + +def test_restart_preserves_queued_jobs_and_interrupts_only_active_work(tmp_path: Path) -> None: + manager = JobManager(tmp_path, None, logging.getLogger("test.jobs")) # type: ignore[arg-type] + queued = _job(1, JobStatus.QUEUED) + running = _job(2, JobStatus.RUNNING) + manager.jobs = [queued, running] + manager._save() + + restored = JobManager(tmp_path, None, logging.getLogger("test.jobs")) # type: ignore[arg-type] + + assert restored.jobs[0].status == JobStatus.QUEUED + assert restored.jobs[1].status == JobStatus.INTERRUPTED + assert restored._queue == [queued.id] + + +def test_worker_coalesces_rapid_progress_events() -> None: + plan = ExecutionPlan( + request="train", + summary="Training", + steps=[PlanStep("ddpm_trainer", "Train", "Run training")], + ) + job = Job(id="FAST0001", plan=plan, status=JobStatus.RUNNING) + + class NoisyExecutor: + def execute(self, _tool_id, _arguments, **kwargs): + for percent in range(1, 101): + kwargs["progress_callback"](percent, "same training burst") + return {} + + worker = JobWorker(job, NoisyExecutor()) # type: ignore[arg-type] + events = [] + worker.event.connect(events.append) + + worker.run() + + progress_events = [event for event in events if event.get("type") == "progress"] + assert 1 <= len(progress_events) <= 2 + assert progress_events[-1]["overall"] == 100 + + +def test_step_eta_uses_measured_progress_cadence() -> None: + samples: list[dict[str, float]] = [] + + first = JobWorker._estimate_step_eta( + {"current_step": 10, "total_steps": 110, "unit": "step"}, + samples, + 100.0, + ) + second = JobWorker._estimate_step_eta( + {"current_step": 20, "total_steps": 110, "unit": "step"}, + samples, + 120.0, + ) + + assert "eta_seconds" not in first + assert second["eta_seconds"] == 180 + assert second["progress_current"] == 20 + assert second["progress_total"] == 110 + assert second["progress_unit"] == "step" + assert second["progress_rate"] == 0.5 + assert second["estimated_completion_at"] + + +def test_active_ddpm_adjustment_is_queued_on_worker(tmp_path: Path, monkeypatch) -> None: + manager = JobManager(tmp_path, None, logging.getLogger("test.jobs")) # type: ignore[arg-type] + job = Job( + plan=ExecutionPlan( + request="train", summary="Train", + steps=[PlanStep("ddpm_trainer", "Train", "Train", { + "batch_size": 8, "gradient_accumulation_steps": 1, + "training_intensity": 100, "epochs": 20, + })], + ), + status=JobStatus.RUNNING, + current_step=0, + ) + + class Worker: + updates = None + + def request_adjustment(self, updates): + self.updates = updates + + worker = Worker() + manager.jobs = [job] + manager._active_job = job + manager._worker = worker # type: ignore[assignment] + monkeypatch.setattr(manager, "_save", lambda: None) + + manager.request_training_adjustment(job.id, { + "batch_size": 4, + "gradient_accumulation_steps": 2, + "training_intensity": 75, + }) + + assert worker.updates == { + "batch_size": 4, + "gradient_accumulation_steps": 2, + "training_intensity": 75, + } + assert job.plan.steps[0].arguments["batch_size"] == 8 + + +def test_vram_retry_uses_old_batch_to_calculate_completed_epochs(tmp_path: Path, monkeypatch) -> None: + dataset = tmp_path / "dataset" + output = tmp_path / "output" + checkpoint = output / "checkpoint-20" + (checkpoint / "unet").mkdir(parents=True) + dataset.mkdir() + for index in range(40): + (dataset / f"{index}.png").write_bytes(b"image") + (checkpoint / "unet" / "diffusion_pytorch_model.safetensors").write_bytes(b"weights") + (checkpoint / "optimizer.bin").write_bytes(b"optimizer") + (checkpoint / "scheduler.bin").write_bytes(b"scheduler") + manager = JobManager(tmp_path, None, logging.getLogger("test.jobs")) # type: ignore[arg-type] + failed = Job( + plan=ExecutionPlan( + request="train", summary="Train", + steps=[PlanStep("ddpm_trainer", "Train", "Train", { + "dataset_dir": str(dataset), "output_dir": str(output), + "epochs": 10, "batch_size": 4, "gradient_accumulation_steps": 1, + })], + ), + status=JobStatus.FAILED, + current_step=0, + error="CUDA out of memory", + ) + manager.jobs = [failed] + monkeypatch.setattr(manager, "_start_next", lambda: None) + + retry = manager.safer_vram_retry(failed.id) + arguments = retry.plan.steps[0].arguments + + assert arguments["batch_size"] == 2 + assert arguments["gradient_accumulation_steps"] == 2 + assert arguments["completed_epochs"] == 2 + assert arguments["epochs"] == 8 + assert arguments["resume_from"] == str(checkpoint) + + +def test_adjustment_ready_requeues_same_job_from_checkpoint(tmp_path: Path, monkeypatch) -> None: + manager = JobManager(tmp_path, None, logging.getLogger("test.jobs")) # type: ignore[arg-type] + job = Job( + plan=ExecutionPlan( + request="train", summary="Train", + steps=[PlanStep("ddpm_trainer", "Train", "Train", { + "epochs": 20, "batch_size": 8, "training_intensity": 100, + })], + ), + status=JobStatus.RUNNING, + current_step=0, + ) + manager.jobs = [job] + manager._active_job = job + monkeypatch.setattr(manager, "_save", lambda: None) + + manager._handle_event({ + "type": "adjustment_ready", + "checkpoint": str(tmp_path / "checkpoint-40"), + "completed_epochs": 4, + "updates": {"batch_size": 4, "training_intensity": 75}, + }) + + arguments = job.plan.steps[0].arguments + assert job.status == JobStatus.QUEUED + assert manager._queue == [job.id] + assert arguments["epochs"] == 16 + assert arguments["completed_epochs"] == 4 + assert arguments["batch_size"] == 4 + assert arguments["training_intensity"] == 75 diff --git a/tests/test_model_inspector.py b/tests/test_model_inspector.py new file mode 100644 index 0000000000000000000000000000000000000000..2be2c30466ac32ec65b373c50f4957598f4b8d73 --- /dev/null +++ b/tests/test_model_inspector.py @@ -0,0 +1,74 @@ +from __future__ import annotations + +from pathlib import Path + +import torch + +from adam.model_inspector import compare_models, inspect_model + + +def test_inspector_reads_generic_torch_checkpoint(tmp_path: Path) -> None: + checkpoint = tmp_path / "checkpoint.pt" + torch.save( + { + "state_dict": { + "encoder.weight": torch.tensor([[1.0, 2.0], [3.0, 4.0]]), + "attention.query.bias": torch.zeros(2), + }, + "epoch": 3, + }, + checkpoint, + ) + + summary = inspect_model(checkpoint) + + assert summary.architecture == "Generic / Unknown" + assert summary.tensor_count == 2 + assert summary.total_parameters == 6 + assert summary.epoch == 3 + assert summary.components["attention"] == 2 + + +def test_inspector_detects_lora_and_adapter_details(tmp_path: Path) -> None: + checkpoint = tmp_path / "adapter_model.safetensors" + try: + from safetensors.torch import save_file + except Exception: + return + save_file( + { + "unet.block.lora_down.weight": torch.ones(4, 8), + "unet.block.lora_up.weight": torch.ones(8, 4) * 0.5, + }, + str(checkpoint), + ) + + summary = inspect_model(checkpoint, recorded_architecture="lora") + + assert summary.architecture == "LoRA" + assert summary.lora["rank"] == "4" + assert summary.trainable_parameters == 64 + + +def test_compare_models_ranks_changed_tensors(tmp_path: Path) -> None: + first = tmp_path / "a.pt" + second = tmp_path / "b.pt" + torch.save({"state_dict": {"layer.weight": torch.ones(4), "same.weight": torch.zeros(2)}}, first) + torch.save({"state_dict": {"layer.weight": torch.ones(4) * 3, "same.weight": torch.zeros(2)}}, second) + + comparison = compare_models(first, second) + + assert comparison.architecture_match + assert comparison.parameter_count_difference == 0 + assert comparison.tensor_comparisons[0].name == "layer.weight" + assert comparison.tensor_comparisons[0].change_score is not None + + +def test_malformed_checkpoint_returns_warning(tmp_path: Path) -> None: + checkpoint = tmp_path / "broken.pt" + checkpoint.write_bytes(b"not a checkpoint") + + summary = inspect_model(checkpoint) + + assert summary.status == "warning" + assert "no readable tensor checkpoint" in "\n".join(summary.health).casefold() diff --git a/tests/test_model_plugins.py b/tests/test_model_plugins.py new file mode 100644 index 0000000000000000000000000000000000000000..ecff29111b3ec60333e5b0f11886051ba9d47e92 --- /dev/null +++ b/tests/test_model_plugins.py @@ -0,0 +1,91 @@ +from __future__ import annotations + +from pathlib import Path +import json + +from adam.commands import TrainingCommand +from adam.model_plugins import ( + ModelPluginRegistry, + scaffold_model_plugin, + validate_settings, +) +from adam.registry import ToolRegistry + + +def test_builtin_model_plugins_are_discovered() -> None: + registry = ModelPluginRegistry(Path.cwd()) + + assert {"ddpm", "flow", "lora"}.issubset(registry.plugins) + assert registry.errors == [] + assert registry.training_schema("ddpm")["resolution"]["type"] == "choice" + assert registry.generation_schema_for_tool("lora_generator")["base_model_path"]["required"] + + +def test_plugin_schema_validation_reports_clear_errors(tmp_path: Path) -> None: + schema = { + "batch_size": {"label": "Batch size", "type": "int", "min": 1, "max": 8}, + "base_model": {"label": "Base model", "type": "path", "required": True, "must_exist": True}, + } + + errors = validate_settings(schema, {"batch_size": 0, "base_model": str(tmp_path / "missing.safetensors")}) + + assert "Batch size must be at least 1." in errors + assert "Base model must point to an existing file." in errors + + +def test_plugin_schema_extends_existing_tool_arguments_without_required_breakage() -> None: + registry = ToolRegistry(Path.cwd()) + lora = registry.get("lora_trainer") + + assert "rank" in lora.arguments + assert "alpha" in lora.arguments + assert set(lora.required_arguments) == { + "dataset_dir", + "model_name", + "epochs", + "output_dir", + "base_model", + } + + +def test_scaffolded_plugin_is_discovered_and_gets_standard_training_arguments(tmp_path: Path) -> None: + folder = scaffold_model_plugin( + tmp_path, + plugin_id="Neural Cellular Automata", + name="Neural Cellular Automata", + architecture="nca", + ) + + config = tmp_path / "config" + config.mkdir() + (config / "tools.json").write_text(json.dumps({"tools": []}), encoding="utf-8") + registry = ToolRegistry(tmp_path) + tool = registry.get("neural_cellular_automata_trainer") + + assert folder.name == "neural_cellular_automata" + assert "dataset_dir" in tool.arguments + assert "output_dir" in tool.required_arguments + assert registry.model_plugins.training_schema("neural_cellular_automata")["resolution"]["default"] == 256 + + +def test_training_command_accepts_discovered_custom_plugin(tmp_path: Path, monkeypatch) -> None: + scaffold_model_plugin( + tmp_path, + plugin_id="maskgit", + name="MaskGIT", + architecture="maskgit", + ) + monkeypatch.chdir(tmp_path) + + command = TrainingCommand.from_dict( + { + "action": "train", + "trainer": "maskgit", + "dataset": "D:/data", + "model_name": "Mask Test", + "epochs": 5, + "training_options": {"resolution": 256, "batch_size": 1}, + } + ) + + assert command.trainer == "maskgit" diff --git a/tests/test_models.py b/tests/test_models.py index 7709a9a0a4f6d2f7a1203432c1df6cfe2e027e47..2255e52a10589ff3c57d672268918d6c1e09eecc 100644 --- a/tests/test_models.py +++ b/tests/test_models.py @@ -26,6 +26,12 @@ def test_job_round_trip_preserves_enums_and_plan() -> None: preview_prompt="test prompt", preview_seed=42, preview_steps=30, + eta_seconds=125, + estimated_completion_at="2026-08-30T23:32:00+00:00", + progress_current=75, + progress_total=100, + progress_rate=0.75, + progress_unit="step", ) restored = Job.from_dict(job.to_dict()) @@ -38,6 +44,12 @@ def test_job_round_trip_preserves_enums_and_plan() -> None: assert restored.preview_prompt == "test prompt" assert restored.preview_seed == 42 assert restored.preview_steps == 30 + assert restored.eta_seconds == 125 + assert restored.estimated_completion_at == "2026-08-30T23:32:00+00:00" + assert restored.progress_current == 75 + assert restored.progress_total == 100 + assert restored.progress_rate == 0.75 + assert restored.progress_unit == "step" def test_training_command_accepts_live_preview_options() -> None: diff --git a/tests/test_oasis_integration.py b/tests/test_oasis_integration.py new file mode 100644 index 0000000000000000000000000000000000000000..519dcefd1e56f10d04fd1aee488f9d837c6dd4bc --- /dev/null +++ b/tests/test_oasis_integration.py @@ -0,0 +1,441 @@ +from __future__ import annotations + +import json +import threading +from pathlib import Path + +from PIL import Image + +from adam.config import ConfigManager +from adam.executor import ToolContext +from adam.oasis_dataset import validate_oasis_dataset +from adam.planner import Planner +from adam.registry import ToolRegistry, ToolSpec +from adam.tools import oasis_adapter + + +def _write_oasis_dataset(root: Path, *, frames: int = 4) -> Path: + dataset = root / "Minecraft Action Dataset" + frames_dir = dataset / "frames" + frames_dir.mkdir(parents=True) + rows = [] + for index in range(frames): + filename = f"frame_{index:08d}.png" + Image.new("RGB", (256, 144), (index * 20, 50, 80)).save(frames_dir / filename) + rows.append({ + "session_id": "session-a", + "frame_index": index, + "filename": filename, + "timestamp_seconds": index / 10, + "w": 1 if index == 1 else 0, + "a": 0, + "s": 0, + "d": 1 if index == 2 else 0, + "jump": 1 if index == 3 else 0, + "mouse_dx": 0.0, + "mouse_dy": 0.0, + "zoom": 0.0, + }) + (dataset / "actions.jsonl").write_text( + "\n".join(json.dumps(row) for row in rows), + encoding="utf-8", + ) + (dataset / "dataset_info.json").write_text( + json.dumps({"capture_fps": 10, "output_resolution": "256x144"}), + encoding="utf-8", + ) + return dataset + + +def _write_registry(root: Path) -> None: + config = root / "config" + config.mkdir() + (config / "tools.json").write_text(json.dumps({"tools": []}), encoding="utf-8") + + +def _write_oasis_model(root: Path, name: str = "Smoke") -> Path: + model = root / "Oasis-Game-Trainer" / "output_action_flow_models" / name + (model / "unet").mkdir(parents=True) + (model / "unet" / "config.json").write_text("{}", encoding="utf-8") + (model / "action_flow_model_info.json").write_text( + json.dumps({"model_type": "action_conditioned_rectified_flow_video", "model_name": name}), + encoding="utf-8", + ) + return model + + +def test_oasis_plugin_is_registered() -> None: + registry = ToolRegistry(Path.cwd()) + + tool = registry.get("oasis_trainer") + player = registry.get("oasis_player") + + assert tool.model_trainers == () + assert "resume_training" in tool.capabilities + assert registry.model_plugins.training_schema("oasis")["resolution"]["default"] == "256x144" + assert player.model_trainers == ("oasis",) + + +def test_oasis_dataset_validation_accepts_legacy_actions_jsonl(tmp_path: Path) -> None: + dataset = _write_oasis_dataset(tmp_path) + + report = validate_oasis_dataset(str(dataset), frame_gap=1) + + assert report.ok + assert report.valid_rows == 4 + assert report.valid_transitions == 3 + assert report.action_counts["w"] == 1 + + +def test_oasis_dataset_validation_rejects_missing_labels(tmp_path: Path) -> None: + dataset = _write_oasis_dataset(tmp_path) + rows = (dataset / "actions.jsonl").read_text(encoding="utf-8").splitlines() + first = json.loads(rows[0]) + del first["jump"] + rows[0] = json.dumps(first) + (dataset / "actions.jsonl").write_text("\n".join(rows), encoding="utf-8") + + report = validate_oasis_dataset(str(dataset), frame_gap=1) + + assert not report.ok + assert any("missing action label" in error for error in report.errors) + + +def test_oasis_dataset_validation_accepts_legacy_action_suffix_mismatch(tmp_path: Path) -> None: + dataset = _write_oasis_dataset(tmp_path, frames=3) + frames_dir = dataset / "frames" + (frames_dir / "frame_00000001.png").rename(frames_dir / "frame_00000001_CAM.png") + + report = validate_oasis_dataset(str(dataset), frame_gap=1) + + assert report.ok + assert report.valid_rows == 3 + assert any("missing legacy filename" in warning for warning in report.warnings) + + +def test_oasis_dataset_validation_skips_deleted_frames(tmp_path: Path) -> None: + dataset = _write_oasis_dataset(tmp_path, frames=5) + (dataset / "frames" / "frame_00000004.png").unlink() + + report = validate_oasis_dataset(str(dataset), frame_gap=1) + + assert report.ok + assert report.valid_rows == 4 + assert report.valid_transitions == 3 + assert any("skipping row" in warning for warning in report.warnings) + + +def test_oasis_training_request_creates_standard_plan(tmp_path: Path, monkeypatch) -> None: + _write_registry(tmp_path) + dataset = _write_oasis_dataset(tmp_path) + oasis_root = tmp_path / "Oasis-Game-Trainer" + oasis_root.mkdir() + (oasis_root / "roblox_action_flow_app.py").write_text("# oasis", encoding="utf-8") + config = ConfigManager(tmp_path) + config.settings["provider"] = "manual" + config.settings["tool_folders"] = {"oasis_trainer": str(oasis_root)} + monkeypatch.chdir(tmp_path) + planner = Planner(tmp_path, ToolRegistry(tmp_path), config) + + plan = planner.plan( + f"Create an Oasis model called Old Minecraft Beta, use the {dataset} dataset, " + "train it for 5 epochs at 256x144. " + '[ADAM_TRAINING_OPTIONS:{"resolution":"256x144","batch_size":2,"workers":0}]' + ) + + assert plan.requires_confirmation is True + assert [step.tool_id for step in plan.steps] == ["oasis_trainer"] + assert plan.steps[0].arguments["model_name"] == "Old Minecraft Beta" + assert Path(plan.steps[0].arguments["output_dir"]).parent.name == "output_action_flow_models" + assert plan.steps[0].arguments["resolution"] == "256x144" + + +def test_oasis_player_request_finds_registered_model(tmp_path: Path) -> None: + _write_registry(tmp_path) + oasis_root = tmp_path / "Oasis-Game-Trainer" + _write_oasis_model(tmp_path, "Beta World") + config = ConfigManager(tmp_path) + config.settings["tool_folders"] = {"oasis_trainer": str(oasis_root)} + planner = Planner(tmp_path, ToolRegistry(tmp_path), config) + + plan = planner.plan("Open Oasis Beta World with seed 42") + + assert plan.requires_confirmation is False + assert [step.tool_id for step in plan.steps] == ["oasis_player"] + assert plan.steps[0].arguments["model_name"] == "Beta World" + assert plan.steps[0].arguments["seed"] == 42 + + +def test_oasis_player_request_rejects_invalid_explicit_folder(tmp_path: Path) -> None: + _write_registry(tmp_path) + invalid = tmp_path / "not-a-model" + invalid.mkdir() + planner = Planner(tmp_path, ToolRegistry(tmp_path), ConfigManager(tmp_path)) + + plan = planner.plan(f"Launch Oasis from {invalid}") + + assert plan.steps == [] + assert "valid Oasis action model folder" in plan.summary + + +def test_oasis_fine_tune_accepts_existing_dataset_path(tmp_path: Path) -> None: + _write_registry(tmp_path) + bundle = tmp_path / "Roblox Dataset" + dataset = _write_oasis_dataset(bundle) + oasis_root = tmp_path / "Oasis-Game-Trainer" + _write_oasis_model(tmp_path, "Roblox Oasis V 2.3.6") + config = ConfigManager(tmp_path) + config.settings["provider"] = "manual" + config.settings["tool_folders"] = {"oasis_trainer": str(oasis_root)} + planner = Planner(tmp_path, ToolRegistry(tmp_path), config) + + plan = planner.plan( + "Fine-tune Roblox Oasis V 2.3.6 for 30 epochs with Oasis Action World Model. " + "[ADAM_FINE_TUNE:" + + json.dumps({ + "dataset_mode": "existing", + "dataset_name": str(bundle), + "epochs": 30, + "image_count": 500, + "model_name": "Roblox Oasis V 2.3.6", + "new_subject": "", + "trainer": "oasis", + "training_options": {"resolution": "256x144", "workers": 0}, + }) + + "]" + ) + + assert plan.requires_confirmation is True + step = plan.steps[0] + assert step.tool_id == "oasis_trainer" + assert step.arguments["dataset_dir"] == str(dataset.resolve()) + assert step.arguments["resume_from"] == str( + (oasis_root / "output_action_flow_models" / "Roblox Oasis V 2.3.6").resolve() + ) + + +def test_oasis_fine_tune_normalizes_path_shaped_model_name(tmp_path: Path) -> None: + _write_registry(tmp_path) + dataset = _write_oasis_dataset(tmp_path) + oasis_root = tmp_path / "Oasis-Game-Trainer" + model = _write_oasis_model(tmp_path, "Roblox Oasis V 2.3.6") + (model / "action_flow_model_info.json").write_text( + json.dumps({ + "model_type": "action_conditioned_rectified_flow_video", + "model_name": str(model), + }), + encoding="utf-8", + ) + config = ConfigManager(tmp_path) + config.settings["provider"] = "manual" + config.settings["tool_folders"] = {"oasis_trainer": str(oasis_root)} + planner = Planner(tmp_path, ToolRegistry(tmp_path), config) + + plan = planner.plan( + f"Fine-tune {model} for 30 epochs with Oasis Action World Model. " + "[ADAM_FINE_TUNE:" + + json.dumps({ + "dataset_mode": "existing", + "dataset_name": str(dataset), + "epochs": 30, + "image_count": 500, + "model_name": str(model), + "new_subject": "", + "trainer": "oasis", + "training_options": {"resolution": "256x144", "workers": 0}, + }) + + "]" + ) + + assert plan.requires_confirmation is True + assert plan.steps[0].arguments["model_name"] == "Roblox Oasis V 2.3.6" + + +def test_oasis_fine_tune_finds_model_when_payload_path_has_old_parent(tmp_path: Path) -> None: + _write_registry(tmp_path) + dataset = _write_oasis_dataset(tmp_path) + oasis_root = tmp_path / "FlowMatchImageGenerator" / "Oasis-Game-Trainer" + model = _write_oasis_model(tmp_path / "FlowMatchImageGenerator", "Roblox Oasis V 2.3.6") + old_parent_path = ( + tmp_path + / "FlowMatchImageGenerator" + / "output_action_flow_models" + / "Roblox Oasis V 2.3.6" + ) + config = ConfigManager(tmp_path) + config.settings["provider"] = "manual" + config.settings["tool_folders"] = {"oasis_trainer": str(oasis_root)} + planner = Planner(tmp_path, ToolRegistry(tmp_path), config) + + plan = planner.plan( + f"Fine-tune {old_parent_path} for 30 epochs with Oasis Action World Model. " + "[ADAM_FINE_TUNE:" + + json.dumps({ + "dataset_mode": "existing", + "dataset_name": str(dataset), + "epochs": 30, + "image_count": 500, + "model_name": str(old_parent_path), + "new_subject": "", + "trainer": "oasis", + "training_options": {"resolution": "256x144", "workers": 0}, + }) + + "]" + ) + + assert plan.requires_confirmation is True + assert plan.steps[0].arguments["resume_from"] == str(model.resolve()) + assert plan.steps[0].arguments["model_name"] == "Roblox Oasis V 2.3.6" + + +def test_oasis_fine_tune_expands_dataset_bundle_to_safe_list(tmp_path: Path) -> None: + _write_registry(tmp_path) + bundle = tmp_path / "Roblox Dataset" + first = _write_oasis_dataset(bundle / "one") + second = _write_oasis_dataset(bundle / "two") + oasis_root = tmp_path / "Oasis-Game-Trainer" + _write_oasis_model(tmp_path, "Roblox Oasis V 2.3.6") + config = ConfigManager(tmp_path) + config.settings["provider"] = "manual" + config.settings["tool_folders"] = {"oasis_trainer": str(oasis_root)} + planner = Planner(tmp_path, ToolRegistry(tmp_path), config) + + plan = planner.plan( + "Fine-tune Roblox Oasis V 2.3.6 for 30 epochs with Oasis Action World Model. " + "[ADAM_FINE_TUNE:" + + json.dumps({ + "dataset_mode": "existing", + "dataset_name": str(bundle), + "epochs": 30, + "image_count": 500, + "model_name": "Roblox Oasis V 2.3.6", + "new_subject": "", + "trainer": "oasis", + "training_options": {"resolution": "256x144", "workers": 0}, + }) + + "]" + ) + + dataset_dir = plan.steps[0].arguments["dataset_dir"] + assert dataset_dir == [str(first.resolve()), str(second.resolve())] + + +def test_oasis_adapter_builds_worker_command_and_registers_model(tmp_path: Path, monkeypatch) -> None: + dataset = _write_oasis_dataset(tmp_path) + oasis_root = tmp_path / "Oasis-Game-Trainer" + output = oasis_root / "output_action_flow_models" / "Smoke" + oasis_root.mkdir() + (oasis_root / "roblox_action_flow_app.py").write_text("# oasis", encoding="utf-8") + config = ConfigManager(tmp_path) + config.update({"tool_folders": {"oasis_trainer": str(oasis_root)}}) + captured: dict[str, object] = {} + + class FakeProcess: + pid = 123 + returncode = 0 + stdout = iter([ + 'ACTION_FLOW_EVENT:{"type":"start","transitions":3,"training":2,"validation":1,"device":"cpu"}\n', + 'ACTION_FLOW_EVENT:{"type":"progress","epoch":1,"epochs":1,"update":1,"total_updates":1,"loss":0.5}\n', + 'ACTION_FLOW_EVENT:{"type":"complete","output_dir":"x"}\n', + ]) + + def poll(self): + return 0 + + def fake_popen(command, **kwargs): + captured["command"] = command + captured["cwd"] = kwargs.get("cwd") + (output / "unet").mkdir(parents=True) + (output / "unet" / "config.json").write_text("{}", encoding="utf-8") + (output / "action_flow_model_info.json").write_text( + json.dumps({"model_type": "action_conditioned_rectified_flow_video"}), + encoding="utf-8", + ) + return FakeProcess() + + monkeypatch.setattr(oasis_adapter.subprocess, "Popen", fake_popen) + context = ToolContext( + tmp_path, + "OASIS", + ToolSpec("oasis_trainer", "Oasis", "test", "Training", "train_oasis"), + threading.Event(), + threading.Event(), + lambda *_args, **_kwargs: None, + lambda *_args: None, + ) + context.run_event.set() + + result = oasis_adapter.train_oasis( + context, + dataset_dir=str(dataset), + model_name="Smoke", + epochs=1, + output_dir=str(output), + workers=0, + preview_enabled=False, + ) + + command = captured["command"] + assert "--train-worker" in command + assert "--dataset-dir" in command + assert str(dataset) in command + assert "--mixed-precision" in command + assert result["assets"][0]["trainer"] == "oasis" + + +def test_oasis_adapter_accepts_dataset_folder_list(tmp_path: Path, monkeypatch) -> None: + first = _write_oasis_dataset(tmp_path / "one") + second = _write_oasis_dataset(tmp_path / "two") + oasis_root = tmp_path / "Oasis-Game-Trainer" + output = oasis_root / "output_action_flow_models" / "Smoke" + oasis_root.mkdir() + (oasis_root / "roblox_action_flow_app.py").write_text("# oasis", encoding="utf-8") + ConfigManager(tmp_path).update({"tool_folders": {"oasis_trainer": str(oasis_root)}}) + captured: dict[str, object] = {} + + class FakeProcess: + pid = 123 + returncode = 0 + stdout = iter([ + 'ACTION_FLOW_EVENT:{"type":"complete","output_dir":"x"}\n', + ]) + + def poll(self): + return 0 + + def fake_popen(command, **kwargs): + captured["command"] = command + (output / "unet").mkdir(parents=True) + (output / "unet" / "config.json").write_text("{}", encoding="utf-8") + (output / "action_flow_model_info.json").write_text( + json.dumps({"model_type": "action_conditioned_rectified_flow_video"}), + encoding="utf-8", + ) + return FakeProcess() + + monkeypatch.setattr(oasis_adapter.subprocess, "Popen", fake_popen) + context = ToolContext( + tmp_path, + "OASIS", + ToolSpec("oasis_trainer", "Oasis", "test", "Training", "train_oasis"), + threading.Event(), + threading.Event(), + lambda *_args, **_kwargs: None, + lambda *_args: None, + ) + context.run_event.set() + + result = oasis_adapter.train_oasis( + context, + dataset_dir=[str(first), str(second)], + model_name="Smoke", + epochs=1, + output_dir=str(output), + workers=0, + preview_enabled=False, + ) + + command = captured["command"] + dataset_arg = command[command.index("--dataset-dir") + 1] + assert dataset_arg == f"{first};{second}" + assert result["assets"][0]["dataset_path"] == f"{first};{second}" diff --git a/tests/test_planner.py b/tests/test_planner.py index 62c56e2a67b1f09cbf4b1903d32082bdc1ec1b19..dd5def56665978defa451a4a025ff7feabe7a5af 100644 --- a/tests/test_planner.py +++ b/tests/test_planner.py @@ -1,6 +1,8 @@ from __future__ import annotations from pathlib import Path +import json +import tempfile import shutil from adam.config import ConfigManager @@ -9,12 +11,39 @@ from adam.registry import ToolRegistry ROOT = Path(__file__).resolve().parents[1] +_TEST_DIRECTORIES: list[tempfile.TemporaryDirectory] = [] def make_planner() -> Planner: - config = ConfigManager(ROOT) + temporary = tempfile.TemporaryDirectory(prefix="adam-planner-test-") + _TEST_DIRECTORIES.append(temporary) + root = Path(temporary.name) + (root / "config").mkdir() + shutil.copy2(ROOT / "config" / "tools.json", root / "config" / "tools.json") + collector = root / "collector" + for name in ("Hatsune Miku", "Liminal Spaces Dataset", "Mario"): + (collector / "Datasets" / name).mkdir(parents=True) + tool_folders = { + "dataset_collector": str(collector), + "lora_trainer": str(root / "lora"), + "ddpm_trainer": str(root / "ddpm"), + "flow_trainer": str(root / "flow"), + } + for folder in tool_folders.values(): + Path(folder).mkdir(parents=True, exist_ok=True) + base_models = root / "LoRA StableDiffusionModels Here" + base_models.mkdir() + base_model = base_models / "test-sdxl.safetensors" + base_model.touch() + lora_config = Path(tool_folders["lora_trainer"]) / "config" + lora_config.mkdir() + (lora_config / "app_settings.json").write_text( + json.dumps({"last_model": str(base_model)}), encoding="utf-8" + ) + config = ConfigManager(root) config.settings["provider"] = "manual" - return Planner(ROOT, ToolRegistry(ROOT), config) + config.settings["tool_folders"] = tool_folders + return Planner(root, ToolRegistry(root), config) def test_lora_request_builds_real_confirmed_pipeline() -> None: @@ -139,16 +168,20 @@ def test_natural_ddpm_followup_resolves_registered_dataset_and_defaults() -> Non assert [step.tool_id for step in followup.steps] == ["ddpm_trainer"] assert followup.steps[0].arguments["model_name"] == "Hatsune Miku" assert followup.steps[0].arguments["epochs"] == 300 - assert Path(followup.steps[0].arguments["dataset_dir"]).name == "On Hatsune Miku LoRA" + dataset_dir = Path(followup.steps[0].arguments["dataset_dir"]) + assert dataset_dir.is_dir() + assert "hatsune miku" in dataset_dir.name.casefold() -def test_ddpm_followup_accepts_an_absolute_windows_dataset_path() -> None: +def test_ddpm_followup_accepts_an_absolute_windows_dataset_path(tmp_path: Path) -> None: planner = make_planner() - dataset = Path(r"D:\Users\PlayRobloxAllDay\Desktop\Programs\GoogleImageDatasetCollector\Datasets\Dantdm Dataset") + dataset = tmp_path / "DatasetCollector" / "Datasets" / "Dantdm Dataset" + dataset.mkdir(parents=True) + output = tmp_path / "DDPM" / "output" request = ( f"From the {dataset} dataset, train a DDPM model for 100 epochs. " "Name the model DanTDM. Output Folder " - r"D:\Users\PlayRobloxAllDay\Desktop\Programs\DDPM\output" + f"{output}" ) fields = planner._parse_ddpm_fields(request) diff --git a/tests/test_remote_hardening.py b/tests/test_remote_hardening.py new file mode 100644 index 0000000000000000000000000000000000000000..fd9123fc424a9a006eceed6fd06aa5732bf2a028 --- /dev/null +++ b/tests/test_remote_hardening.py @@ -0,0 +1,145 @@ +from __future__ import annotations + +import http.client +import json +from pathlib import Path + +import pytest + +from adam.config import ConfigManager +from adam.remote_access import RemoteAccessService +from adam.remote_api import RemoteApiError, bounded_float + + +@pytest.fixture +def remote(tmp_path): + config = ConfigManager(tmp_path) + service = RemoteAccessService(config, None, None) + service.save_settings({'enabled': True}) + # Let the OS allocate a port, avoiding a bind/close/rebind race in the test. + config.settings['remote_access']['port'] = 0 + service.start() + try: + yield service + finally: + service.shutdown() + + +def request(service, path='/api/status', *, method='GET', body=None, headers=None, token=None): + connection = http.client.HTTPConnection('127.0.0.1', service._server.server_port, timeout=3) + auth = token if token is not None else service.settings()['token'] + fields = {'Authorization': f'Bearer {auth}'} + fields.update(headers or {}) + try: + connection.request(method, path, body=body, headers=fields) + response = connection.getresponse() + return response.status, dict(response.getheaders()), response.read() + finally: + connection.close() + + +def test_token_rotation_and_disable_take_effect_without_restart(remote): + old = remote.settings()['token'] + assert request(remote, token=old)[0] == 200 + remote.save_settings({'token': 'replacement-token'}) + assert request(remote, token=old)[0] == 401 + assert request(remote)[0] == 200 + remote.save_settings({'enabled': False}) + assert request(remote)[0] == 401 + + +def test_missing_token_is_persisted_once(tmp_path): + config = ConfigManager(tmp_path) + config.update({'remote_access': {'enabled': False}}) + service = RemoteAccessService(config, None, None) + try: + first = service.settings()['token'] + assert service.settings()['token'] == first + assert ConfigManager(tmp_path).get('remote_access')['token'] == first + finally: + service.shutdown() + + +def test_remote_cannot_grant_its_own_control(remote): + payload = json.dumps({'auto_approve_training': True}) + args = dict(method='POST', body=payload, headers={'Content-Type': 'application/json'}) + assert request(remote, '/api/remote-settings', **args)[0] == 403 + assert remote.settings()['auto_approve_training'] is False + remote.save_settings({'allow_job_control': True}) + assert request(remote, '/api/remote-settings', **args)[0] == 200 + remote.save_settings({'allow_job_control': False}) + args['body'] = json.dumps({'auto_approve_training': False}) + assert request(remote, '/api/remote-settings', **args)[0] == 200 + + +@pytest.mark.parametrize('headers,body,status', [ + ({'Content-Type': 'application/json', 'Origin': 'https://attacker.invalid'}, '{}', 403), + ({'Content-Type': 'application/json', 'Sec-Fetch-Site': 'cross-site'}, '{}', 403), + ({'Content-Type': 'text/plain'}, '{}', 415), + ({'Content-Type': 'application/json', 'Content-Length': '-1'}, '', 400), + ({'Content-Type': 'application/json', 'Content-Length': '20001'}, '', 413), + ({'Content-Type': 'application/json'}, '[]', 400), + ({'Content-Type': 'application/json'}, '{', 400), + ({'Content-Type': 'application/json'}, '{"auto_approve_training":"false"}', 400), +]) +def test_rejects_unsafe_requests_before_mutation(remote, headers, body, status): + assert request(remote, '/api/remote-settings', method='POST', body=body, headers=headers)[0] == status + assert remote.settings()['auto_approve_training'] is False + + +def test_browser_security_headers(remote): + status, headers, _ = request(remote, '/') + assert status == 200 + assert headers['Referrer-Policy'] == 'no-referrer' + assert headers['X-Frame-Options'] == 'DENY' + assert "frame-ancestors 'none'" in headers['Content-Security-Policy'] + + +@pytest.mark.parametrize('value', ['nan', 'inf', '-inf', float('nan')]) +def test_nonfinite_remote_settings_are_rejected(value): + with pytest.raises(RemoteApiError): + bounded_float(value, minimum=0, maximum=10, default=1, label='Guidance') + + +def test_dashboard_treats_remote_labels_as_text(tmp_path): + import shutil + import subprocess + from adam.remote_dashboard import remote_dashboard_app_html + + node = shutil.which('node') + if not node: + pytest.skip('Node is needed for the dashboard JavaScript regression check') + html = remote_dashboard_app_html() + script = html.split('', 1)[0] + names = ['$', 'list', 'clear', 'text', 'appendText', 'renderQueues', 'renderSystem', 'renderLocations', 'makeSetting', 'fieldId'] + functions = '\n'.join(line for line in script.splitlines() if any(line.startswith('function ' + name + '(') for name in names)) + harness = r''' +const assert = require('assert'); +class Element { + constructor(tag) { this.tagName=tag; this.children=[]; this.textContent=''; } + set innerHTML(value) { throw new Error('Untrusted text reached HTML parsing'); } + appendChild(child) { this.children.push(child); } + get firstChild() { return this.children[0]; } + removeChild(child) { this.children.splice(this.children.indexOf(child),1); } + cloneNode() { return this; } + setAttribute() {} +} +const elements={}; +const document={createElement:tag=>new Element(tag),getElementById:id=>elements[id]||(elements[id]=new Element('div'))}; +const attack=''; +var state={locationFilter:'',locations:[{id:'1',name:attack,source:attack,available:true}]}; +''' + checks = r''' +renderQueues({queue:[{project:attack,status:attack,progress:0}]}); +assert.equal(elements.queues.children[0].children[0].textContent,attack); +renderLocations(); +assert.equal(elements.locationsList.children[0].children[0].textContent,attack); +makeSetting('preview',{type:'bool',label:attack},false); +renderSystem({cpu_percent:attack}); +assert.equal(globalThis.compromised,undefined); +''' + # Parse the complete shipped script too, not just the rendering functions. + path = tmp_path / 'dashboard-test.js' + path.write_text('new Function(' + json.dumps(script) + ');\n' + harness + functions + checks, encoding='utf-8') + result = subprocess.run([node, str(path)], capture_output=True, text=True, timeout=10) + assert result.returncode == 0, result.stderr diff --git a/tests/test_remote_phase3a.py b/tests/test_remote_phase3a.py new file mode 100644 index 0000000000000000000000000000000000000000..4d55c7a91389ed18c6e3cdbe9772dd1438c777c5 --- /dev/null +++ b/tests/test_remote_phase3a.py @@ -0,0 +1,508 @@ +from __future__ import annotations + +import json +import logging +import threading +import time +from pathlib import Path +from urllib.error import HTTPError +from urllib.request import Request, urlopen + +from PIL import Image +import pytest + +from adam.assets import AssetRegistry +from adam.config import ConfigManager +from adam.dataset_registry import DatasetRegistry +from adam.executor import ToolContext +from adam.generations import build_generation_plan +from adam.job_manager import JobManager +from adam.models import ExecutionPlan, Job, JobStatus, PlanStep +from adam.planner import Planner +from adam.registry import ToolRegistry +from adam.remote_access import RemoteAccessService +from adam.remote_dispatcher import RemoteCommandDispatcher +from adam.remote_media import OpaqueIdCodec, RemoteMediaStore +from adam.remote_v1 import RemoteV1Service +from adam.studio import caption_path +from adam.tools.lora_adapter import train_lora + + +def _image(path: Path, color: tuple[int, int, int] = (40, 120, 210)) -> None: + path.parent.mkdir(parents=True, exist_ok=True) + Image.new("RGB", (16, 12), color).save(path) + + +@pytest.mark.parametrize("structured", [False, True]) +@pytest.mark.parametrize("epochs,needs_review", [(10, False), (1000, True)]) +def test_remote_training_reviews_before_auto_approval( + tmp_path: Path, monkeypatch, structured: bool, epochs: int, needs_review: bool, +) -> None: + dataset = tmp_path / "dataset" + _image(dataset / "image.png") + trainer = tmp_path / "trainer" + trainer.mkdir() + config = _config(tmp_path, { + "tool_folders": {"ddpm_trainer": str(trainer)}, + "remote_access": {"auto_approve_training": True, "token": "test-token"}, + }) + planner = Planner(tmp_path, ToolRegistry(Path.cwd()), config) + asset = planner.assets.register(kind="dataset", name="Test Dataset", path=str(dataset)) + jobs = JobManager(tmp_path, None, logging.getLogger("test.remote.review"), config) + monkeypatch.setattr(jobs, "_start_next", lambda: None) + service = RemoteAccessService(config, jobs, None, planner) + try: + if structured: + payload = { + "trainer": "ddpm", "model_name": "Test Model", "epochs": epochs, + "dataset_id": service.codec.encode({"kind": "dataset", "asset_id": asset.id}), + } + preview = service.api_v1.training_plan(payload) + assert "Pre-flight:" in preview["summary"] + assert "ORION —" in preview["summary"] + response = service.api_v1.start_training(payload) + else: + response = service.submit_prompt( + f"From the Test Dataset dataset, train a DDPM model for {epochs} epochs. " + "Name the model Test Model." + ) + assert response["ok"] is True + + job = jobs.get(response["job_id"]) + assert response["requires_approval"] is needs_review + assert job.status == (JobStatus.AWAITING_CONFIRMATION if needs_review else JobStatus.QUEUED) + assert job.plan.orion_review["level"] == ("warning" if needs_review else "ready") + assert job.plan.summary.count("Pre-flight:") == 1 + assert job.plan.summary.count("ORION —") == 1 + assert job.plan.steps[0].arguments["epochs"] == epochs + assert config.get("remote_access")["auto_approve_training"] is True + finally: + service.shutdown() + jobs.shutdown() + + +def _config(root: Path, values: dict | None = None) -> ConfigManager: + config = ConfigManager(root) + if values: + config.update(values) + return config + + +def _remote_v1(root: Path, *, jobs=None, planner=None) -> RemoteV1Service: + config = _config(root) + planner = planner or Planner(root, ToolRegistry(Path.cwd()), config) + return RemoteV1Service( + root=root, + config=config, + jobs=jobs, + planner=planner, + dispatcher=RemoteCommandDispatcher(), + codec=OpaqueIdCodec("test-secret"), + media=RemoteMediaStore(root, OpaqueIdCodec("test-secret")), + auto_approve_training=lambda _plan: False, + ) + + +def test_remote_dispatcher_uses_invoker_from_worker_thread() -> None: + dispatcher = RemoteCommandDispatcher() + calls: list[str] = [] + + class Invoker: + def invoke(self, payload): + calls.append("invoked") + payload["result"] = payload["fn"]() + payload["event"].set() + + dispatcher._invoker = Invoker() + result: list[str] = [] + thread = threading.Thread(target=lambda: result.append(dispatcher.call_ui(lambda: "done"))) + thread.start() + thread.join(timeout=3) + + dispatcher.shutdown() + assert calls == ["invoked"] + assert result == ["done"] + + +def test_remote_v1_datasets_are_paginated_redacted_and_editable(tmp_path: Path) -> None: + dataset = tmp_path / "datasets" / "Minecraft Steve" + for index in range(3): + image = dataset / f"image_{index}.png" + _image(image, (index * 40, 100, 200)) + caption_path(image).write_text(f"caption {index}\n", encoding="utf-8") + assets = AssetRegistry(tmp_path) + dataset_asset = assets.register(kind="dataset", name="Minecraft Steve", path=str(dataset)) + + api = _remote_v1(tmp_path) + public_id = api.codec.encode({"kind": "dataset", "asset_id": dataset_asset.id}) + listed = json.loads(api.route("GET", "/api/v1/datasets").body.decode("utf-8")) + page = json.loads(api.route("GET", f"/api/v1/datasets/{public_id}/items", "page=1&page_size=2").body.decode("utf-8")) + + assert listed["datasets"][0]["name"] == "Minecraft Steve" + assert "path" not in listed["datasets"][0] + assert page["pagination"]["total"] == 3 + assert len(page["items"]) == 2 + item = page["items"][0] + assert "path" not in item + assert item["caption"] == "caption 0\n" + + caption = json.loads(api.route( + "POST", + f"/api/v1/datasets/{public_id}/items/{item['id']}/caption", + payload={"caption": "new caption"}, + ).body.decode("utf-8")) + decision = json.loads(api.route( + "POST", + f"/api/v1/datasets/{public_id}/items/{item['id']}/decision", + payload={"decision": "reject"}, + ).body.decode("utf-8")) + + assert caption["item"]["caption"] == "new caption" + assert (dataset / "image_0.txt").read_text(encoding="utf-8") == "new caption\n" + assert decision["item"]["decision"] == "reject" + + +def test_dataset_registry_discovers_registered_locations_into_remote(tmp_path: Path) -> None: + location = tmp_path / "Remembered" + dataset = location / "Minecraft Oasis V3" + _image(dataset / "frame_0001.png") + registry = DatasetRegistry(tmp_path, _config(tmp_path)) + registry.register_location(location, name="Oasis datasets") + + api = _remote_v1(tmp_path) + listed = json.loads(api.route("GET", "/api/v1/datasets").body.decode("utf-8")) + locations = json.loads(api.route("GET", "/api/v1/datasets/locations").body.decode("utf-8")) + + assert listed["datasets"][0]["name"] == "Minecraft Oasis V3" + assert listed["datasets"][0]["available"] is True + assert listed["datasets"][0]["thumbnail_url"] + assert "path" not in listed["datasets"][0] + assert locations["locations"][0]["name"] == "Oasis datasets" + assert "path" not in locations["locations"][0] + + +def test_remote_caption_cannot_escape_dataset(tmp_path: Path, monkeypatch) -> None: + dataset = tmp_path / 'dataset' + _image(dataset / 'image.png') + assets = AssetRegistry(tmp_path) + asset = assets.register(kind='dataset', name='Test', path=str(dataset)) + api = _remote_v1(tmp_path) + dataset_id = api.codec.encode({'kind': 'dataset', 'asset_id': asset.id}) + item_id = api.media.media_id(kind='dataset_image', asset_id=asset.id, index=0) + outside = tmp_path / 'private.txt' + outside.write_text('private', encoding='utf-8') + monkeypatch.setattr('adam.remote_v1.caption_path', lambda _path: outside) + response = api.route('POST', f'/api/v1/datasets/{dataset_id}/items/{item_id}/caption', payload={'caption': 'overwritten'}) + assert response.status == 403 + assert outside.read_text(encoding='utf-8') == 'private' + assert api.route('GET', f'/api/v1/datasets/{dataset_id}/items').status == 403 + api.dispatcher.shutdown() + + +def test_remote_caption_replaces_hard_link_without_overwriting_target(tmp_path: Path) -> None: + import os + dataset = tmp_path / 'dataset' + _image(dataset / 'image.png') + outside = tmp_path / 'private.txt' + outside.write_text('private', encoding='utf-8') + os.link(outside, dataset / 'image.txt') + assets = AssetRegistry(tmp_path) + asset = assets.register(kind='dataset', name='Test', path=str(dataset)) + api = _remote_v1(tmp_path) + dataset_id = api.codec.encode({'kind': 'dataset', 'asset_id': asset.id}) + item_id = api.media.media_id(kind='dataset_image', asset_id=asset.id, index=0) + response = api.route('POST', f'/api/v1/datasets/{dataset_id}/items/{item_id}/caption', payload={'caption': 'new caption'}) + assert response.status == 200 + assert outside.read_text(encoding='utf-8') == 'private' + assert (dataset / 'image.txt').read_text(encoding='utf-8') == 'new caption\n' + api.dispatcher.shutdown() + + +def test_remote_dataset_favorite_and_use_are_persistent_without_paths(tmp_path: Path) -> None: + dataset = tmp_path / "datasets" / "Minecraft Oasis V3" + _image(dataset / "frame_0001.png") + assets = AssetRegistry(tmp_path) + asset = assets.register(kind="dataset", name="Minecraft Oasis V3", path=str(dataset)) + api = _remote_v1(tmp_path) + dataset_id = api.codec.encode({"kind": "dataset", "asset_id": asset.id}) + + favorite = json.loads(api.route( + "POST", + f"/api/v1/datasets/{dataset_id}/favorite", + payload={"favorite": True}, + ).body.decode("utf-8")) + used = json.loads(api.route( + "POST", + f"/api/v1/datasets/{dataset_id}/use", + payload={}, + ).body.decode("utf-8")) + + assert favorite["dataset"]["favorite"] is True + assert used["dataset"]["last_used_at"] + registry = DatasetRegistry(tmp_path, _config(tmp_path)) + record = registry.record_for_path(dataset) + assert record.favorite is True + assert record.last_used_at + + +def test_remote_v1_opaque_item_id_cannot_cross_datasets(tmp_path: Path) -> None: + first = tmp_path / "first" + second = tmp_path / "second" + _image(first / "a.png") + _image(second / "b.png") + assets = AssetRegistry(tmp_path) + one = assets.register(kind="dataset", name="One", path=str(first)) + two = assets.register(kind="dataset", name="Two", path=str(second)) + api = _remote_v1(tmp_path) + first_id = api.codec.encode({"kind": "dataset", "asset_id": one.id}) + wrong_dataset = api.codec.encode({"kind": "dataset", "asset_id": two.id}) + item_id = api.media.media_id(kind="dataset_image", asset_id=one.id, index=0) + + response = api.route( + "POST", + f"/api/v1/datasets/{wrong_dataset}/items/{item_id}/decision", + payload={"decision": "keep"}, + ) + + assert response.status == 403 + assert first_id + + +def test_remote_thumbnail_cache_reuses_and_invalidates_changed_source(tmp_path: Path) -> None: + source = tmp_path / "image.png" + _image(source, (10, 20, 30)) + media = RemoteMediaStore(tmp_path, OpaqueIdCodec("cache-test")) + + first = media.thumbnail(source, size=180) + second = media.thumbnail(source, size=180) + time.sleep(0.02) + _image(source, (200, 40, 30)) + third = media.thumbnail(source, size=180) + + assert first.path == second.path + assert second.cache_hit is True + assert third.path != first.path + assert third.cache_hit is False + + +def test_remote_v1_models_include_lora_trigger_word_without_paths(tmp_path: Path) -> None: + model = tmp_path / "models" / "Adam_OC_LoRA_v2" + model.mkdir(parents=True) + checkpoint = model / "adam.safetensors" + checkpoint.write_bytes(b"weights") + assets = AssetRegistry(tmp_path) + assets.register( + kind="model", + name="Adam_OC_LoRA_v2", + path=str(model), + trainer="lora", + checkpoint=str(checkpoint), + metadata={"trigger_word": "adam_oc"}, + ) + api = _remote_v1(tmp_path) + + payload = json.loads(api.route("GET", "/api/v1/models").body.decode("utf-8")) + + assert payload["models"][0]["trigger_word"] == "adam_oc" + assert "path" not in payload["models"][0] + assert payload["models"][0]["checkpoint_name"] == "adam.safetensors" + + +def test_structured_generation_queues_existing_generation_plan(tmp_path: Path) -> None: + model = tmp_path / "ddpm" / "Model" + model.mkdir(parents=True) + (model / "model_index.json").write_text("{}", encoding="utf-8") + assets = AssetRegistry(tmp_path) + model_asset = assets.register(kind="model", name="Minecraft", path=str(model), trainer="ddpm") + + class Jobs: + def __init__(self) -> None: + self.jobs = [] + self.active_job = None + + def submit(self, plan: ExecutionPlan) -> Job: + job = Job(plan=plan, status=JobStatus.QUEUED) + self.jobs.insert(0, job) + return job + + api = _remote_v1(tmp_path, jobs=Jobs()) + model_id = api.codec.encode({"kind": "model", "asset_id": model_asset.id}) + response = json.loads(api.route( + "POST", + "/api/v1/generation/start", + payload={ + "provider_id": "ddpm_generator", + "model_id": model_id, + "prompt": "Minecraft", + "image_count": 1, + "steps": 20, + "seed": 5, + "sampler": "DDIM", + "aspect_ratio": "1:1 (Square)", + }, + ).body.decode("utf-8")) + + assert response["job_id"] + assert response["plan"]["steps"][0]["tool_id"] == "ddpm_generator" + + +def test_lora_trigger_word_survives_planner_command_job_and_experiment(tmp_path: Path) -> None: + dataset = tmp_path / "dataset" + dataset.mkdir() + for index in range(2): + _image(dataset / f"{index}.png") + caption_path(dataset / f"{index}.png").write_text("adam_oc\n", encoding="utf-8") + lora_root = tmp_path / "lora" + (lora_root / "output").mkdir(parents=True) + base = tmp_path / "base.safetensors" + base.write_bytes(b"base") + config = _config(tmp_path, {"tool_folders": {"lora_trainer": str(lora_root)}}) + registry = ToolRegistry(Path.cwd()) + planner = Planner(tmp_path, registry, config) + planner.assets.register(kind="dataset", name="Adam Dataset", path=str(dataset)) + request = ( + "From the Adam Dataset dataset, train a LoRA model for 3 epochs. " + "Name the model Adam_OC_LoRA_v2. " + f"[ADAM_TRAINING_OPTIONS:{{\"base_model\":{json.dumps(str(base))},\"trigger_word\":\"adam_oc\"}}] " + "[ADAM_TRAINER:lora]" + ) + + plan = planner.plan(request) + args = plan.steps[0].arguments + job = Job(plan=plan, status=JobStatus.FINISHED, output_folder=args["output_dir"]) + run = planner.assets + experiment = __import__("adam.experiment_tracker", fromlist=["ExperimentStore"]).ExperimentStore(tmp_path).record_job(job) + + assert args["model_name"] == "Adam_OC_LoRA_v2" + assert args["trigger_word"] == "adam_oc" + assert experiment is not None + assert experiment.trigger_word == "adam_oc" + assert run + + +def test_lora_adapter_passes_explicit_trigger_word_to_native_payload(tmp_path: Path) -> None: + trainer = tmp_path / "trainer" + backend = trainer / "src" / "loratrainer" / "trainer" + model_pkg = trainer / "src" / "loratrainer" / "models" + backend.mkdir(parents=True) + model_pkg.mkdir(parents=True) + for package in (trainer / "src" / "loratrainer", backend, model_pkg): + (package / "__init__.py").write_text("", encoding="utf-8") + (model_pkg / "training_config.py").write_text( + "from dataclasses import dataclass\n" + "from pathlib import Path\n" + "@dataclass\n" + "class TrainingConfig:\n" + " dataset_dir: Path\n" + " base_model_path: Path\n" + " output_dir: Path\n" + " resume_checkpoint: Path | None = None\n" + " trigger_word: str = ''\n" + " epochs: int = 1\n", + encoding="utf-8", + ) + (backend / "diffusers_sdxl_lora_backend.py").write_text( + "import json\n" + "class DiffusersSDXLLoRABackend:\n" + " def train(self, config, control, progress):\n" + " config.output_dir.mkdir(parents=True, exist_ok=True)\n" + " (config.output_dir / 'payload.json').write_text(json.dumps({'trigger_word': config.trigger_word, 'epochs': config.epochs}), encoding='utf-8')\n" + " final = config.output_dir / 'final.safetensors'\n" + " final.write_bytes(b'weights')\n" + " return final\n", + encoding="utf-8", + ) + dataset = tmp_path / "dataset" + for index in range(2): + _image(dataset / f"{index}.png") + caption_path(dataset / f"{index}.png").write_text("adam_oc\n", encoding="utf-8") + base = tmp_path / "base.safetensors" + base.write_bytes(b"base") + _config(tmp_path, {"tool_folders": {"lora_trainer": str(trainer)}}) + context = ToolContext( + root=tmp_path, + job_id="LORA1", + tool="lora_trainer", + cancel_event=threading.Event(), + run_event=threading.Event(), + progress_callback=lambda *_args, **_kwargs: None, + log_callback=lambda _message: None, + preview_callback=lambda _payload: None, + ) + context.run_event.set() + + result = train_lora( + context, + dataset_dir=str(dataset), + model_name="Adam_OC_LoRA_v2", + trigger_word="adam_oc", + epochs=1, + output_dir=str(trainer / "output" / "Adam_OC_LoRA_v2"), + base_model=str(base), + ) + + payload = json.loads((Path(result["output_folder"]) / "payload.json").read_text(encoding="utf-8")) + assert payload["trigger_word"] == "adam_oc" + assert result["trigger_word"] == "adam_oc" + assert result["assets"][0]["metadata"]["trigger_word"] == "adam_oc" + + +def test_remote_v1_routes_are_authenticated_and_legacy_status_is_redacted(tmp_path: Path) -> None: + assets = AssetRegistry(tmp_path) + dataset = tmp_path / "dataset" + _image(dataset / "a.png") + assets.register(kind="dataset", name="Dataset", path=str(dataset)) + + class Config: + root = tmp_path + values = {} + + def get(self, key, default=None): + return self.values.get(key, default) + + def update(self, values): + self.values.update(values) + + asset_registry = assets + + class PlannerStub: + root = tmp_path + registry = ToolRegistry(Path.cwd()) + assets = asset_registry + + class Jobs: + def __init__(self) -> None: + self.active_job = None + self.jobs = [ + Job( + plan=ExecutionPlan("run", "Run", [PlanStep("preview_generator", "Preview", "Preview")]), + status=JobStatus.QUEUED, + output_folder=str(tmp_path / "secret" / "output"), + ) + ] + + service = RemoteAccessService(Config(), Jobs(), monitor=None, planner=PlannerStub()) + token = service.settings()["token"] + service.save_settings({"enabled": True, "port": 0, "token": token}) + import socket + + with socket.socket() as sock: + sock.bind(("127.0.0.1", 0)) + port = sock.getsockname()[1] + service.save_settings({"enabled": True, "port": port, "token": token}) + try: + service.start() + try: + urlopen(f"http://127.0.0.1:{port}/api/v1/datasets", timeout=3) + except HTTPError as exc: + assert exc.code == 401 + else: + raise AssertionError("v1 route should require authentication") + payload = json.loads(urlopen(f"http://127.0.0.1:{port}/api/status?token={token}", timeout=3).read().decode("utf-8")) + datasets = json.loads(urlopen(f"http://127.0.0.1:{port}/api/v1/datasets?token={token}", timeout=3).read().decode("utf-8")) + finally: + service.stop() + + assert payload["queue"][0]["output_folder"] == "output" + assert str(tmp_path) not in json.dumps(payload) + assert datasets["datasets"][0]["name"] == "Dataset" diff --git a/tests/test_training_assistant.py b/tests/test_training_assistant.py index 618852bbf402d79e120280263b686b731f83b4f9..108c67625aec7aeedd17fb4740a1b7af0b20cf14 100644 --- a/tests/test_training_assistant.py +++ b/tests/test_training_assistant.py @@ -25,6 +25,41 @@ class FakeConfig: return self.values.get(key, default) +def test_preflight_does_not_treat_empty_tool_path_as_connected(monkeypatch) -> None: + def unexpected_scan(*args, **kwargs): + raise AssertionError("An empty dataset path must not scan the working directory") + + monkeypatch.setattr(Path, "rglob", unexpected_scan) + plan = ExecutionPlan( + request="train", summary="Train.", + steps=[PlanStep("ddpm_trainer", "Train DDPM", "Train", {"epochs": 10})], + ) + + append_preflight_summary(plan, FakeConfig()) + + assert "program folder is not connected" in plan.summary + assert "Train DDPM: connected" not in plan.summary + + +def test_orion_text_without_a_report_does_not_skip_review(tmp_path: Path) -> None: + plan = ExecutionPlan( + request="train", summary="ORION — mentioned in a user summary.", + steps=[ + PlanStep("dataset_collector", "Collect", "Collect", { + "output_dir": str(tmp_path / "future"), "image_count": 2000, + }), + PlanStep("ddpm_trainer", "Train", "Train", { + "dataset_dir": str(tmp_path / "future"), "epochs": 600, + }), + ], + ) + + append_preflight_summary(plan, FakeConfig()) + + assert plan.orion_review["level"] == "warning" + assert plan.requires_confirmation is True + + def test_wizard_builds_specific_new_ddpm_request() -> None: request = build_training_request( trainer="ddpm",