SyntheticMDProductions's picture
Some of Adams structure
e0265b9 verified
Raw
History Blame Contribute Delete
8.38 kB
"""Execution bridge for the connected Rectified Flow image trainer."""
from __future__ import annotations
import json
import queue
import subprocess
import sys
import threading
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
IMAGE_EXTENSIONS = {".jpg", ".jpeg", ".png", ".webp", ".bmp"}
def _latest_preview(folder: Path) -> Path | None:
try:
images = [path for path in folder.rglob("*") if path.is_file()
and path.suffix.lower() in IMAGE_EXTENSIONS
and any(token in path.name.lower() 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 train_flow(
context: ToolContext,
dataset_dir: str,
model_name: str,
epochs: int,
output_dir: str,
resume_from: str = "",
resolution: int = 256, batch_size: int = 8, learning_rate: float = 0.0002,
gradient_accumulation: int = 1, workers: int = 4, mixed_precision: str = "fp16",
save_every: int = 10, preview_every: int = 10, preview_steps: int = 10,
gradient_checkpointing: bool = False, preview_enabled: bool = True,
preview_prompt: str = "", preview_seed: int = 123456789,
) -> dict[str, object]:
"""Launch the user's Flow Matching worker and relay its structured progress."""
root = Path(str(ConfigManager(context.root).get("tool_folders", {}).get("flow_trainer", ""))).expanduser()
script = root / "flow_matching_app.py"
dataset = Path(dataset_dir).expanduser().resolve()
output = Path(output_dir).expanduser().resolve()
if not script.is_file():
raise ToolExecutionError("Flow Matching flow_matching_app.py was not found. Re-scan its folder in Settings.")
if not dataset.is_dir():
raise ToolExecutionError("The selected Flow Matching dataset folder no longer exists.")
if sum(1 for path in dataset.iterdir() if path.is_file() and path.suffix.lower() in IMAGE_EXTENSIONS) < 2:
raise ToolExecutionError("The Flow Matching dataset needs at least two image files before training can start.")
if not 1 <= int(epochs) <= 100_000:
raise ToolExecutionError("Epoch count must be between 1 and 100000.")
if not 64 <= int(resolution) <= 512 or int(resolution) % 16 or not 1 <= int(batch_size) <= 64 or not 1e-7 <= float(learning_rate) <= 0.1 or not 1 <= int(gradient_accumulation) <= 64 or not 0 <= int(workers) <= 16 or mixed_precision not in {"fp16", "no"} or min(int(save_every), int(preview_every), int(preview_steps)) < 1:
raise ToolExecutionError("Flow training options are outside ADAM's safe range.")
safe_name = model_name.strip()
if not safe_name or len(safe_name) > 96 or any(char in safe_name for char in "<>:\\|?*\x00"):
raise ToolExecutionError("Choose a short model name without filesystem-reserved characters.")
output_root = (root / "output_flow_models").resolve()
try:
output.relative_to(output_root)
except ValueError as exc:
raise ToolExecutionError("Flow Matching outputs must stay inside output_flow_models.") from exc
if output.exists():
raise ToolExecutionError("The chosen Flow Matching output folder already exists; ADAM will not overwrite it.")
resume = Path(resume_from).expanduser().resolve() if resume_from else None
if resume:
try:
metadata = json.loads((resume / "flow_model_info.json").read_text(encoding="utf-8"))
if metadata.get("model_type") != "rectified_flow" or not (resume / "unet" / "config.json").is_file():
raise ValueError
saved_resolution = int(metadata.get("resolution", 0) or 0)
except (OSError, ValueError, TypeError, json.JSONDecodeError) as exc:
raise ToolExecutionError("Choose a valid completed Flow Matching model to continue.") from exc
if saved_resolution != int(resolution):
raise ToolExecutionError(
f"The selected Flow model is {saved_resolution}px; continuation must use the same resolution."
)
output.parent.mkdir(parents=True, exist_ok=True)
command = [
sys.executable, str(script), "--train-worker", "--data-dir", str(dataset),
"--output-dir", str(output), "--model-name", safe_name, "--epochs", str(int(epochs)),
"--resolution", str(int(resolution)), "--batch-size", str(int(batch_size)), "--learning-rate", str(float(learning_rate)),
"--workers", str(int(workers)), "--gradient-accumulation", str(int(gradient_accumulation)), "--mixed-precision", mixed_precision,
"--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)), "--tf32",
]
if gradient_checkpointing:
command.append("--gradient-checkpointing")
if resume:
command.extend(["--continue-model", str(resume)])
context.log(f"Starting real Flow Matching training. Output folder: {output}")
process = subprocess.Popen(command, cwd=str(root), stdout=subprocess.PIPE, stderr=subprocess.STDOUT,
text=True, encoding="utf-8", errors="replace", shell=False)
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 Flow Matching trainer")
stopped = False
suspended = False
stop_file = output / "stop_flow_training.flag"
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("Flow Matching trainer paused safely." if suspended else "Flow Matching trainer resumed.")
if context.cancel_event.is_set() and not stopped:
if suspended:
set_process_tree_paused(process, False)
suspended = False
stop_file.touch(exist_ok=True)
stopped = True
context.log("Safe stop requested; waiting for Flow Matching to finish its current batch.")
try:
line = lines.get(timeout=0.15)
if line and line.startswith("FLOW_EVENT:"):
event = json.loads(line.split(":", 1)[1])
if event.get("type") == "progress" and not stopped:
current = int(event.get("epoch", 0) or 0)
context.progress(max(1, min(99, round(current * 100 / int(epochs)))),
f"Finished epoch {current} of {epochs}")
if preview_enabled and current and current % int(preview_every) == 0:
candidate = Path(str(event.get("preview_path", ""))) if event.get("preview_path") else _latest_preview(output)
if candidate:
context.preview(candidate, epoch=current,
next_epoch=min(int(epochs), current + int(preview_every)),
prompt=preview_prompt, seed=int(preview_seed), steps=int(preview_steps))
elif event.get("type") == "warning":
context.log(str(event.get("message", "Flow trainer warning.")))
elif line:
context.log(line)
except queue.Empty:
pass
if process.poll() is not None and lines.empty():
break
if stopped:
raise ToolCancelled("Flow Matching training stopped by user.")
if process.returncode != 0:
raise ToolExecutionError(f"Flow Matching trainer exited with code {process.returncode}. See the job log for details.")
context.progress(100, "Flow Matching training completed")
return {"output_folder": str(output), "model_name": safe_name, "assets": [{
"kind": "model", "name": safe_name, "path": str(output), "trainer": "flow",
"dataset_path": str(dataset), "checkpoint": str(output), "epochs": int(epochs),
}]}