feat(lora-worker): stand up in-arbo LoRA training worker on irv-ml1 (arbo Phase 1 §4.1)
Host service (runs as llmuser, owns /opt/fluxgym + GPU access) that runs
sd-scripts SDXL LoRA training on demand for arbo — the infra-ops half of the
in-arbo LoRA training Phase 1 ownership split (vh/arbo
docs/contracts/in-arbo-lora-training-phase1.contract.md §4.1/§2).
- Fixed-invocation only (INV-T7): bounded params -> one sd-scripts command
shape; every param range/allowlist/path-containment checked before spawn;
bad request = 422, never a silent downgrade. 14 unit tests green.
- Thin supervisor: never imports torch; subprocesses the fluxgym venv's
accelerate. 1-job-at-a-time (arbo lease is the serializer, 409 is backstop).
Durable job records + boot reconciliation (§4.3).
- API: POST /train, GET /train/{id}[/log], POST /train/{id}/cancel,
GET /gpu-status (per-device VRAM + tts_on_3090 co-OOM signal), GET /healthz.
- Wire-shape (§7 resolved with comfy-dev): shared /worktank/arbo/train handoff
(group arbotrain, setgid 2770); worker binds 0.0.0.0:8203, arbo reaches via
host.docker.internal:host-gateway (reachability proven on 172.20.0.1:8203);
device-aware TTS steering via /gpu-status.
Deployed to irv-ml1 via playbooks/deploy-lora-training-worker.yaml (elway,
idempotent); systemd unit active; /healthz + /gpu-status verified live.
This commit is contained in:
@@ -0,0 +1,5 @@
|
||||
"""LoRA training worker — a host service (runs as llmuser on irv-ml1) that runs sd-scripts
|
||||
on demand for arbo's in-arbo LoRA training (contract: vh/arbo
|
||||
docs/contracts/in-arbo-lora-training-phase1.contract.md §4.1/§2). Fixed-invocation only."""
|
||||
|
||||
__version__ = "0.1.0"
|
||||
@@ -0,0 +1,85 @@
|
||||
"""FastAPI surface for the LoRA training worker (§4.1 API — arbo is the client).
|
||||
|
||||
Endpoints:
|
||||
POST /train dispatch a train (409 if busy, 422 on invalid params) → {worker_job_id}
|
||||
GET /train/{id} status {status, step, total_steps, loss, eta_s, lora_path?, error?}
|
||||
GET /train/{id}/log tail of the run log
|
||||
POST /train/{id}/cancel best-effort kill → cancelled
|
||||
GET /gpu-status per-device VRAM + tts_on_3090 (arbo's device-aware steering, §7)
|
||||
GET /healthz liveness + whether a train is active
|
||||
|
||||
The worker holds ONE job at a time; arbo's lease is the real serializer (INV-T2), the 409
|
||||
here is the backstop. Nothing here interprets free-form arguments — see invocation.py (INV-T7).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from fastapi import FastAPI, HTTPException
|
||||
from fastapi.responses import PlainTextResponse
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from . import config
|
||||
from .gpu import gpu_status
|
||||
from .invocation import InvalidTrainRequest
|
||||
from .jobs import Busy, manager
|
||||
|
||||
app = FastAPI(title="LoRA Training Worker", version="0.1.0")
|
||||
|
||||
|
||||
class TrainRequest(BaseModel):
|
||||
# Bounded params only (INV-T7). Full validation (ranges, path containment, tier/device
|
||||
# fit) happens in invocation.validate_request; pydantic just pins the shape + types.
|
||||
dataset_dir: str
|
||||
base_model_path: str
|
||||
output_dir: str
|
||||
output_name: str
|
||||
trigger: str
|
||||
subject_class: str
|
||||
repeats: int = Field(ge=1, le=100)
|
||||
tier: str
|
||||
device_index: int = Field(ge=0, le=1)
|
||||
seed: int = Field(default=42, ge=0)
|
||||
|
||||
|
||||
@app.post("/train")
|
||||
def post_train(req: TrainRequest):
|
||||
try:
|
||||
return manager.submit(req.model_dump())
|
||||
except InvalidTrainRequest as exc:
|
||||
raise HTTPException(status_code=422, detail=str(exc))
|
||||
except Busy as exc:
|
||||
raise HTTPException(status_code=409, detail=str(exc))
|
||||
|
||||
|
||||
@app.get("/train/{job_id}")
|
||||
def get_train(job_id: str):
|
||||
status = manager.get(job_id)
|
||||
if status is None:
|
||||
raise HTTPException(status_code=404, detail=f"unknown worker_job_id: {job_id}")
|
||||
return status
|
||||
|
||||
|
||||
@app.get("/train/{job_id}/log", response_class=PlainTextResponse)
|
||||
def get_train_log(job_id: str):
|
||||
tail = manager.log_tail(job_id)
|
||||
if tail is None:
|
||||
raise HTTPException(status_code=404, detail=f"no log for worker_job_id: {job_id}")
|
||||
return tail
|
||||
|
||||
|
||||
@app.post("/train/{job_id}/cancel")
|
||||
def post_cancel(job_id: str):
|
||||
status = manager.cancel(job_id)
|
||||
if status is None:
|
||||
raise HTTPException(status_code=404, detail=f"unknown worker_job_id: {job_id}")
|
||||
return status
|
||||
|
||||
|
||||
@app.get("/gpu-status")
|
||||
def get_gpu_status():
|
||||
return gpu_status()
|
||||
|
||||
|
||||
@app.get("/healthz")
|
||||
def healthz():
|
||||
return {"ok": True, "active_job": manager.active(), "port": config.PORT}
|
||||
@@ -0,0 +1,92 @@
|
||||
"""Static configuration for the LoRA training worker.
|
||||
|
||||
Everything load-bearing is an explicit module constant here (explicit-over-implicit): the
|
||||
paths the worker is allowed to touch, the fixed sd-scripts invocation surface, the tier
|
||||
table, and the SDXL-LoRA hyperparameters. Nothing about a training run is free-form — the
|
||||
worker only ever builds ONE command shape (INV-T7), and every knob that shapes it is
|
||||
visible in this file.
|
||||
|
||||
The worker runs as `llmuser` on irv-ml1. It never imports torch / sd-scripts; it only
|
||||
*subprocesses* the fluxgym venv's `accelerate launch`. So this process stays tiny and the
|
||||
fluxgym venv stays pristine.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from pathlib import Path
|
||||
|
||||
# ---- Network ---------------------------------------------------------------------------
|
||||
# 0.0.0.0 (not 127.0.0.1): arbo runs CONTAINERIZED on the traefik-net bridge and reaches
|
||||
# the host worker via host.docker.internal:host-gateway — host loopback is unreachable from
|
||||
# that bridge. irv-ml1 is WireGuard-only + ACL'd, so 0.0.0.0 exposure is bounded to the tunnel.
|
||||
HOST = os.environ.get("LORA_WORKER_HOST", "0.0.0.0")
|
||||
PORT = int(os.environ.get("LORA_WORKER_PORT", "8203"))
|
||||
|
||||
# ---- Fluxgym / sd-scripts (the ONLY thing the worker executes) -------------------------
|
||||
FLUXGYM_ROOT = Path("/opt/fluxgym")
|
||||
FLUXGYM_VENV = FLUXGYM_ROOT / ".venv"
|
||||
ACCELERATE_BIN = FLUXGYM_VENV / "bin" / "accelerate"
|
||||
SD_SCRIPTS_DIR = FLUXGYM_ROOT / "sd-scripts"
|
||||
SDXL_TRAIN_SCRIPT = SD_SCRIPTS_DIR / "sdxl_train_network.py"
|
||||
|
||||
# ---- Allowed path roots (INV-T7 path safety) -------------------------------------------
|
||||
# dataset_dir + output_dir MUST live under the shared handoff root (the group-shared,
|
||||
# setgid /worktank/arbo/train that both arbo-container-uid and llmuser can rw). base_model
|
||||
# must live under a known model root. Anything else → rejected before a process is spawned.
|
||||
HANDOFF_ROOT = Path(os.environ.get("LORA_WORKER_HANDOFF_ROOT", "/worktank/arbo/train"))
|
||||
ALLOWED_MODEL_ROOTS = tuple(
|
||||
Path(p)
|
||||
for p in os.environ.get(
|
||||
"LORA_WORKER_MODEL_ROOTS",
|
||||
"/opt/fluxgym/models:/worktank/comfyui:/worktank/models:/worktank/arbo",
|
||||
).split(":")
|
||||
if p
|
||||
)
|
||||
|
||||
# ---- Worker state + logs (survives a worker restart for boot reconciliation, §4.3) -----
|
||||
STATE_DIR = Path(os.environ.get("LORA_WORKER_STATE_DIR", "/opt/lora-training-worker/state"))
|
||||
LOG_DIR = Path(os.environ.get("LORA_WORKER_LOG_DIR", "/opt/lora-training-worker/logs"))
|
||||
JOBS_FILE = STATE_DIR / "jobs.json"
|
||||
MAX_RETAINED_JOBS = 50 # keep terminal jobs queryable for arbo's status proxy
|
||||
|
||||
# ---- Tier table (§4.6) -----------------------------------------------------------------
|
||||
# tier -> (max_train_steps, network_dim, resolution). Quality (1024) is A6000-only; the
|
||||
# device-fit guard lives in invocation.build_command (quality on the 3090 → rejected).
|
||||
TIERS = {
|
||||
"fast": {"steps": 400, "dim": 16, "resolution": 768},
|
||||
"balanced": {"steps": 1500, "dim": 32, "resolution": 768},
|
||||
"quality": {"steps": 3000, "dim": 32, "resolution": 1024},
|
||||
}
|
||||
|
||||
# ---- Device model (CUDA_DEVICE_ORDER=PCI_BUS_ID) ---------------------------------------
|
||||
# Under PCI_BUS_ID (which the recipe pins), index 0 = RTX 3090, index 1 = RTX A6000.
|
||||
# Quality (1024) will not fit the 3090's ~24GB LoRA envelope → allowed on the A6000 only.
|
||||
DEVICE_3090 = 0
|
||||
DEVICE_A6000 = 1
|
||||
QUALITY_ONLY_DEVICES = (DEVICE_A6000,)
|
||||
|
||||
# ---- SDXL-LoRA hyperparameters (agent-discretion defaults; confirm vs Sindra) ----------
|
||||
# The contract (§4.6) pins the flag SET + tier->(steps,dim,res) but NOT lr/scheduler/batch.
|
||||
# These are standard, conservative SDXL-LoRA values, surfaced here so they're one-line
|
||||
# auditable + tunable. Flagged to comfy-dev to cross-check against the proven Sindra runs.
|
||||
SDXL_HPARAMS = {
|
||||
"learning_rate": "1e-4",
|
||||
"lr_scheduler": "cosine",
|
||||
"lr_warmup_steps": "0",
|
||||
"train_batch_size_lean": "1", # 3090
|
||||
"train_batch_size_full": "2", # A6000
|
||||
"network_alpha_ratio": 1.0, # alpha = dim * ratio
|
||||
"min_snr_gamma": "5",
|
||||
"noise_offset": "0.1",
|
||||
"save_precision": "bf16",
|
||||
"mixed_precision": "bf16",
|
||||
"max_data_loader_n_workers": "2",
|
||||
}
|
||||
|
||||
# ---- TTS-liveness heuristic (GET /gpu-status → tts_on_3090) -----------------------------
|
||||
# arbo's scheduler steers a lean train off the 3090 when TTS is live there. We report the
|
||||
# raw per-device VRAM + a boolean derived from (a) a compute-app on the 3090 whose cmdline
|
||||
# matches a TTS marker, else (b) a used-MB floor fallback.
|
||||
TTS_CMDLINE_MARKERS = ("chatterbox", "omnivoice", "tts_server", "chatterbox_fast")
|
||||
TTS_3090_USED_MB_FLOOR = int(os.environ.get("LORA_WORKER_TTS_FLOOR_MB", "2000"))
|
||||
@@ -0,0 +1,122 @@
|
||||
"""GPU status for arbo's device-aware scheduler (GET /gpu-status).
|
||||
|
||||
arbo steers a lean train OFF the 3090 when TTS is live there (the 3090/TTS co-OOM concern,
|
||||
contract §7). The worker can't decide that — it just reports the raw per-device VRAM plus a
|
||||
`tts_on_3090` boolean derived from nvidia-smi compute-apps. Everything is nvidia-smi CSV
|
||||
parsing; no torch import.
|
||||
|
||||
Device indices are reported under CUDA_DEVICE_ORDER=PCI_BUS_ID (index 0 = RTX 3090,
|
||||
index 1 = RTX A6000) — the same ordering the training recipe pins, so `device_index` in a
|
||||
/train request and the indices here mean the same physical card.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import subprocess
|
||||
|
||||
from . import config
|
||||
|
||||
_SMI_ENV = {"CUDA_DEVICE_ORDER": "PCI_BUS_ID"}
|
||||
|
||||
|
||||
def _run(args: list[str]) -> str:
|
||||
return subprocess.run(
|
||||
["nvidia-smi", *args], capture_output=True, text=True, timeout=15, env={**_SMI_ENV}
|
||||
).stdout
|
||||
|
||||
|
||||
def _query_devices() -> list[dict]:
|
||||
# index,name,memory.total,memory.used,memory.free (MiB) under PCI_BUS_ID ordering
|
||||
out = _run(["--query-gpu=index,name,memory.total,memory.used,memory.free",
|
||||
"--format=csv,noheader,nounits"])
|
||||
devices = []
|
||||
for line in out.strip().splitlines():
|
||||
parts = [p.strip() for p in line.split(",")]
|
||||
if len(parts) != 5:
|
||||
continue
|
||||
idx, name, total, used, free = parts
|
||||
devices.append({
|
||||
"index": int(idx), "name": name,
|
||||
"total_mb": int(total), "used_mb": int(used), "free_mb": int(free),
|
||||
"procs": [],
|
||||
})
|
||||
return devices
|
||||
|
||||
|
||||
def _query_compute_apps() -> list[dict]:
|
||||
# pid,used_memory,gpu_bus_id — map each compute app to a device by bus id
|
||||
out = _run(["--query-compute-apps=pid,used_memory,gpu_bus_id", "--format=csv,noheader,nounits"])
|
||||
apps = []
|
||||
for line in out.strip().splitlines():
|
||||
parts = [p.strip() for p in line.split(",")]
|
||||
if len(parts) != 3:
|
||||
continue
|
||||
pid, mem, bus = parts
|
||||
if not pid.isdigit():
|
||||
continue
|
||||
apps.append({"pid": int(pid), "used_mb": int(mem) if mem.isdigit() else 0, "bus_id": bus})
|
||||
return apps
|
||||
|
||||
|
||||
def _bus_id_by_index() -> dict[int, str]:
|
||||
out = _run(["--query-gpu=index,gpu_bus_id", "--format=csv,noheader"])
|
||||
mapping = {}
|
||||
for line in out.strip().splitlines():
|
||||
parts = [p.strip() for p in line.split(",")]
|
||||
if len(parts) == 2 and parts[0].isdigit():
|
||||
mapping[int(parts[0])] = parts[1]
|
||||
return mapping
|
||||
|
||||
|
||||
def _cmdline(pid: int) -> str:
|
||||
try:
|
||||
with open(f"/proc/{pid}/cmdline", "rb") as fh:
|
||||
return fh.read().replace(b"\x00", b" ").decode("utf-8", "replace").lower()
|
||||
except (OSError, ValueError):
|
||||
return ""
|
||||
|
||||
|
||||
def gpu_status() -> dict:
|
||||
"""Return `{devices: [...], tts_on_3090: bool}`. Degrades to an error field on nvidia-smi failure."""
|
||||
try:
|
||||
devices = _query_devices()
|
||||
apps = _query_compute_apps()
|
||||
bus_by_idx = _bus_id_by_index()
|
||||
except (subprocess.SubprocessError, OSError, ValueError) as exc:
|
||||
return {"devices": [], "tts_on_3090": False, "error": f"nvidia-smi failed: {exc}"}
|
||||
|
||||
idx_by_bus = {bus: idx for idx, bus in bus_by_idx.items()}
|
||||
dev_by_idx = {d["index"]: d for d in devices}
|
||||
|
||||
# Attach compute apps (with a marker flag) to their device.
|
||||
for app in apps:
|
||||
idx = idx_by_bus.get(app["bus_id"])
|
||||
cmd = _cmdline(app["pid"])
|
||||
app["is_tts"] = any(m in cmd for m in config.TTS_CMDLINE_MARKERS)
|
||||
if idx is not None and idx in dev_by_idx:
|
||||
dev_by_idx[idx]["procs"].append(
|
||||
{"pid": app["pid"], "used_mb": app["used_mb"], "is_tts": app["is_tts"]}
|
||||
)
|
||||
|
||||
# tts_on_3090 — the co-OOM-risk signal for arbo's device-aware steering. Honest derivation
|
||||
# (transparent via `reason`), because the worker can't reliably fingerprint TTS: the TTS
|
||||
# containers run generic cmdlines (`uvicorn app:app`, `python app.py`) so only the container
|
||||
# NAME identifies them, and the worker's llmuser isn't in the docker group. So:
|
||||
# (a) proc_marker — a compute app on device 0 whose cmdline matches a TTS marker (best effort), else
|
||||
# (b) used_mb_floor — device 0 used >= floor (a busy 3090 = lean-train co-OOM risk, TTS/STT/audio alike).
|
||||
# The PRECISE lever arbo should key on is `devices[0].free_mb`; `tts_on_3090` is the convenience boolean.
|
||||
dev0 = dev_by_idx.get(config.DEVICE_3090)
|
||||
tts_on_3090, reason = False, "idle"
|
||||
free_3090 = None
|
||||
if dev0 is not None:
|
||||
free_3090 = dev0["free_mb"]
|
||||
if any(p["is_tts"] for p in dev0["procs"]):
|
||||
tts_on_3090, reason = True, "proc_marker"
|
||||
elif dev0["used_mb"] >= config.TTS_3090_USED_MB_FLOOR:
|
||||
tts_on_3090, reason = True, "used_mb_floor"
|
||||
return {
|
||||
"devices": devices,
|
||||
"tts_on_3090": tts_on_3090,
|
||||
"tts_on_3090_reason": reason,
|
||||
"gpu3090_free_mb": free_3090, # the precise co-OOM lever; steer lean off the 3090 when low
|
||||
}
|
||||
@@ -0,0 +1,170 @@
|
||||
"""Fixed-invocation command builder — the enforcement point for INV-T7.
|
||||
|
||||
arbo hands the worker a set of BOUNDED parameters; this module turns them into the ONE
|
||||
`accelerate launch sdxl_train_network.py …` command the worker is allowed to run. Every
|
||||
parameter is validated (type, range, allowlist, path-containment) before a single argv
|
||||
element is produced. There is NO path from an arbo request to a free-form argument — the
|
||||
flag set is a constant, only the *values* of vetted parameters vary.
|
||||
|
||||
`build_command` is a pure function (no I/O, no subprocess) so it is trivially unit-testable:
|
||||
feed it a request dict, assert the argv + env. `InvalidTrainRequest` maps to HTTP 422.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
from pathlib import Path
|
||||
|
||||
from . import config
|
||||
|
||||
# Safe token for names that land in a filename / kohya folder path. No shell metachars,
|
||||
# no path separators, no leading dot — even though we never use a shell (argv list), this
|
||||
# also protects the on-disk paths built from `output_name` / `trigger` / `subject_class`.
|
||||
_SAFE_TOKEN = re.compile(r"^[A-Za-z0-9][A-Za-z0-9 _.\-]{0,63}$")
|
||||
_SAFE_OUTPUT_NAME = re.compile(r"^[A-Za-z0-9][A-Za-z0-9_\-]{0,63}$") # filename stem, no spaces
|
||||
|
||||
|
||||
class InvalidTrainRequest(ValueError):
|
||||
"""A train request failed validation → HTTP 422 (never a silent downgrade)."""
|
||||
|
||||
|
||||
def _require(cond: bool, msg: str) -> None:
|
||||
if not cond:
|
||||
raise InvalidTrainRequest(msg)
|
||||
|
||||
|
||||
def _validate_under(path_str: str, roots: tuple[Path, ...], field: str) -> Path:
|
||||
"""Resolve `path_str` and require it to sit under one of `roots` (no traversal escape)."""
|
||||
_require(isinstance(path_str, str) and path_str.startswith("/"), f"{field} must be an absolute path")
|
||||
p = Path(path_str)
|
||||
# Reject traversal explicitly even before resolve() (belt): no '..' components.
|
||||
_require(".." not in p.parts, f"{field} must not contain '..'")
|
||||
resolved = p.resolve()
|
||||
for root in roots:
|
||||
try:
|
||||
resolved.relative_to(root.resolve())
|
||||
return resolved
|
||||
except ValueError:
|
||||
continue
|
||||
allowed = ", ".join(str(r) for r in roots)
|
||||
raise InvalidTrainRequest(f"{field} {resolved} is not under an allowed root ({allowed})")
|
||||
|
||||
|
||||
def validate_request(req: dict) -> dict:
|
||||
"""Validate + normalize an arbo /train request. Returns a clean param dict or raises."""
|
||||
# ---- tier -> (steps, dim, res) --------------------------------------------------
|
||||
tier_raw = req.get("tier")
|
||||
_require(tier_raw in config.TIERS, f"tier must be one of {sorted(config.TIERS)}; got {tier_raw!r}")
|
||||
tier = str(tier_raw)
|
||||
spec = config.TIERS[tier]
|
||||
|
||||
# ---- device_index (int, 0=3090 / 1=A6000 under PCI_BUS_ID) ----------------------
|
||||
device_index = req.get("device_index")
|
||||
_require(isinstance(device_index, int) and device_index in (config.DEVICE_3090, config.DEVICE_A6000),
|
||||
f"device_index must be {config.DEVICE_3090} (3090) or {config.DEVICE_A6000} (A6000)")
|
||||
|
||||
# ---- tier/device fit: quality (1024) is A6000-only (§4.6) -----------------------
|
||||
if tier == "quality":
|
||||
_require(device_index in config.QUALITY_ONLY_DEVICES,
|
||||
"tier 'quality' (1024) requires the A6000 (device_index=1); the 3090 cannot fit it")
|
||||
|
||||
# ---- names (safe tokens; used in argv AND on-disk paths) ------------------------
|
||||
trigger = req.get("trigger")
|
||||
subject_class = req.get("subject_class")
|
||||
output_name = req.get("output_name")
|
||||
_require(isinstance(trigger, str) and bool(_SAFE_TOKEN.match(trigger)), "trigger is not a safe token")
|
||||
_require(isinstance(subject_class, str) and bool(_SAFE_TOKEN.match(subject_class)),
|
||||
"subject_class is not a safe token")
|
||||
_require(isinstance(output_name, str) and bool(_SAFE_OUTPUT_NAME.match(output_name)),
|
||||
"output_name is not a safe filename stem")
|
||||
|
||||
# ---- repeats + seed (bounded ints) ----------------------------------------------
|
||||
repeats = req.get("repeats")
|
||||
_require(isinstance(repeats, int) and 1 <= repeats <= 100, "repeats must be an int in [1, 100]")
|
||||
seed = req.get("seed", 42)
|
||||
_require(isinstance(seed, int) and 0 <= seed <= 2**31 - 1, "seed must be a non-negative int32")
|
||||
|
||||
# ---- paths (containment-checked) ------------------------------------------------
|
||||
dataset_dir = _validate_under(req.get("dataset_dir", ""), (config.HANDOFF_ROOT,), "dataset_dir")
|
||||
output_dir = _validate_under(req.get("output_dir", ""), (config.HANDOFF_ROOT,), "output_dir")
|
||||
base_model_path = _validate_under(req.get("base_model_path", ""), config.ALLOWED_MODEL_ROOTS, "base_model_path")
|
||||
|
||||
return {
|
||||
"tier": tier,
|
||||
"steps": spec["steps"],
|
||||
"dim": spec["dim"],
|
||||
"resolution": spec["resolution"],
|
||||
"device_index": device_index,
|
||||
"trigger": trigger,
|
||||
"subject_class": subject_class,
|
||||
"output_name": output_name,
|
||||
"repeats": repeats,
|
||||
"seed": seed,
|
||||
"dataset_dir": str(dataset_dir),
|
||||
"output_dir": str(output_dir),
|
||||
"base_model_path": str(base_model_path),
|
||||
}
|
||||
|
||||
|
||||
def build_command(req: dict) -> tuple[list[str], dict[str, str], dict]:
|
||||
"""Validate `req` and return `(argv, env_overlay, params)` for the fixed sd-scripts invocation.
|
||||
|
||||
argv is a LIST (never a shell string) — no metachar interpretation is possible. env_overlay
|
||||
is merged onto os.environ by the caller (the CUDA ordering + allocator knobs from §4.6).
|
||||
`params` is the vetted, normalized parameter dict (so the caller need not re-validate).
|
||||
"""
|
||||
p = validate_request(req)
|
||||
h = config.SDXL_HPARAMS
|
||||
batch = h["train_batch_size_full"] if p["device_index"] == config.DEVICE_A6000 else h["train_batch_size_lean"]
|
||||
alpha = int(round(p["dim"] * float(h["network_alpha_ratio"])))
|
||||
res = f"{p['resolution']},{p['resolution']}"
|
||||
|
||||
argv = [
|
||||
str(config.ACCELERATE_BIN), "launch",
|
||||
"--num_processes", "1",
|
||||
"--num_cpu_threads_per_process", "1",
|
||||
"--mixed_precision", h["mixed_precision"],
|
||||
"--dynamo_backend", "no",
|
||||
str(config.SDXL_TRAIN_SCRIPT),
|
||||
"--pretrained_model_name_or_path", p["base_model_path"],
|
||||
"--train_data_dir", p["dataset_dir"],
|
||||
"--output_dir", p["output_dir"],
|
||||
"--output_name", p["output_name"],
|
||||
"--save_model_as", "safetensors",
|
||||
"--save_precision", h["save_precision"],
|
||||
"--mixed_precision", h["mixed_precision"],
|
||||
# LoRA network (§4.6 tier dim; unet-only per the lean recipe)
|
||||
"--network_module", "networks.lora",
|
||||
"--network_dim", str(p["dim"]),
|
||||
"--network_alpha", str(alpha),
|
||||
"--network_train_unet_only",
|
||||
# steps + resolution (§4.6 tier table)
|
||||
"--max_train_steps", str(p["steps"]),
|
||||
"--resolution", res,
|
||||
"--enable_bucket",
|
||||
"--train_batch_size", batch,
|
||||
# memory-lean base (§4.6 — the proven Sindra flag set)
|
||||
"--gradient_checkpointing",
|
||||
"--cache_latents", "--cache_latents_to_disk",
|
||||
"--cache_text_encoder_outputs", "--cache_text_encoder_outputs_to_disk",
|
||||
"--optimizer_type", "adamw8bit",
|
||||
"--no_half_vae",
|
||||
"--sdpa",
|
||||
"--caption_extension", ".txt", # CRITICAL: sd-scripts defaults to .caption → would skip our captions
|
||||
# optimizer schedule (agent-discretion defaults; confirm vs Sindra)
|
||||
"--learning_rate", h["learning_rate"],
|
||||
"--lr_scheduler", h["lr_scheduler"],
|
||||
"--lr_warmup_steps", h["lr_warmup_steps"],
|
||||
"--min_snr_gamma", h["min_snr_gamma"],
|
||||
"--noise_offset", h["noise_offset"],
|
||||
"--max_data_loader_n_workers", h["max_data_loader_n_workers"],
|
||||
"--persistent_data_loader_workers",
|
||||
"--seed", str(p["seed"]),
|
||||
]
|
||||
|
||||
env_overlay = {
|
||||
"CUDA_DEVICE_ORDER": "PCI_BUS_ID", # index 0=3090, 1=A6000 (recipe-pinned ordering)
|
||||
"CUDA_VISIBLE_DEVICES": str(p["device_index"]),
|
||||
"PYTORCH_CUDA_ALLOC_CONF": "expandable_segments:True",
|
||||
}
|
||||
return argv, env_overlay, p
|
||||
@@ -0,0 +1,325 @@
|
||||
"""Training-job lifecycle: one job at a time, durable, with boot reconciliation.
|
||||
|
||||
The worker runs at most ONE sd-scripts process (§4.1 — 1-at-a-time; arbo's lease is the
|
||||
real serializer, the worker's 409 is a backstop). A monitor thread owns the subprocess:
|
||||
it spawns `accelerate launch …`, tails the log for step/loss, and lands a terminal status.
|
||||
|
||||
Durability (§4.3 boot reconciliation): job records persist to JSON. On startup any job
|
||||
left non-terminal whose OS process is gone is marked `failed` ("worker restarted") — so
|
||||
arbo's `GET /train/{id}` never sees a phantom `training` after a worker blip; arbo then
|
||||
releases its lease per its boot rule.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import os
|
||||
import re
|
||||
import signal
|
||||
import subprocess
|
||||
import threading
|
||||
import time
|
||||
import uuid
|
||||
from pathlib import Path
|
||||
from typing import Optional
|
||||
|
||||
from . import config
|
||||
from .invocation import build_command # raises InvalidTrainRequest (→422), propagated by caller
|
||||
|
||||
# sd-scripts / tqdm progress: `steps: 12%|█▏ | 50/400 [00:30<03:30, 1.66it/s, avr_loss=0.123]`
|
||||
_STEP_RE = re.compile(r"(\d+)\s*/\s*(\d+)")
|
||||
_LOSS_RE = re.compile(r"avr_loss=([0-9]*\.?[0-9]+)")
|
||||
|
||||
TERMINAL = {"succeeded", "failed", "cancelled"}
|
||||
|
||||
|
||||
class Job:
|
||||
def __init__(self, worker_job_id: str, params: dict, log_path: Path):
|
||||
self.id = worker_job_id
|
||||
self.params = params
|
||||
self.log_path = log_path
|
||||
self.status = "queued" # queued|preparing|training|succeeded|failed|cancelled
|
||||
self.step = 0
|
||||
self.total_steps = params.get("steps", 0)
|
||||
self.loss: Optional[float] = None
|
||||
self.eta_s: Optional[int] = None
|
||||
self.lora_path: Optional[str] = None
|
||||
self.error: Optional[str] = None
|
||||
self.pid: Optional[int] = None
|
||||
self.started_at: Optional[float] = None
|
||||
self.finished_at: Optional[float] = None
|
||||
|
||||
def public(self) -> dict:
|
||||
return {
|
||||
"worker_job_id": self.id,
|
||||
"status": self.status,
|
||||
"step": self.step,
|
||||
"total_steps": self.total_steps,
|
||||
"loss": self.loss,
|
||||
"eta_s": self.eta_s,
|
||||
"lora_path": self.lora_path,
|
||||
"error": self.error,
|
||||
}
|
||||
|
||||
def to_record(self) -> dict:
|
||||
return {
|
||||
**self.public(),
|
||||
"params": self.params,
|
||||
"log_path": str(self.log_path),
|
||||
"pid": self.pid,
|
||||
"started_at": self.started_at,
|
||||
"finished_at": self.finished_at,
|
||||
}
|
||||
|
||||
@classmethod
|
||||
def from_record(cls, rec: dict) -> "Job":
|
||||
job = cls(rec["worker_job_id"], rec.get("params", {}), Path(rec.get("log_path", "/dev/null")))
|
||||
job.status = rec.get("status", "failed")
|
||||
job.step = rec.get("step", 0)
|
||||
job.total_steps = rec.get("total_steps", 0)
|
||||
job.loss = rec.get("loss")
|
||||
job.eta_s = rec.get("eta_s")
|
||||
job.lora_path = rec.get("lora_path")
|
||||
job.error = rec.get("error")
|
||||
job.pid = rec.get("pid")
|
||||
job.started_at = rec.get("started_at")
|
||||
job.finished_at = rec.get("finished_at")
|
||||
return job
|
||||
|
||||
|
||||
def _pid_alive(pid: Optional[int]) -> bool:
|
||||
if not pid:
|
||||
return False
|
||||
try:
|
||||
os.kill(pid, 0)
|
||||
return True
|
||||
except (OSError, ProcessLookupError):
|
||||
return False
|
||||
|
||||
|
||||
class JobManager:
|
||||
"""Singleton owner of the one active job + retained history. Thread-safe."""
|
||||
|
||||
def __init__(self):
|
||||
self._lock = threading.Lock()
|
||||
self._jobs: dict[str, Job] = {} # id -> Job (bounded to MAX_RETAINED_JOBS)
|
||||
self._active_id: Optional[str] = None
|
||||
self._proc: Optional[subprocess.Popen] = None
|
||||
config.STATE_DIR.mkdir(parents=True, exist_ok=True)
|
||||
config.LOG_DIR.mkdir(parents=True, exist_ok=True)
|
||||
self._load_and_reconcile()
|
||||
|
||||
# ---- persistence ----------------------------------------------------------------
|
||||
def _persist(self) -> None:
|
||||
recs = [self._jobs[i].to_record() for i in list(self._jobs)]
|
||||
tmp = config.JOBS_FILE.with_suffix(".tmp")
|
||||
tmp.write_text(json.dumps(recs, indent=2))
|
||||
tmp.replace(config.JOBS_FILE) # atomic
|
||||
|
||||
def _load_and_reconcile(self) -> None:
|
||||
if not config.JOBS_FILE.exists():
|
||||
return
|
||||
try:
|
||||
recs = json.loads(config.JOBS_FILE.read_text())
|
||||
except (json.JSONDecodeError, OSError):
|
||||
return
|
||||
for rec in recs:
|
||||
job = Job.from_record(rec)
|
||||
# A non-terminal job from a prior worker life whose process is gone → failed.
|
||||
if job.status not in TERMINAL and not _pid_alive(job.pid):
|
||||
job.status = "failed"
|
||||
job.error = "worker restarted; training process not re-attachable"
|
||||
job.finished_at = job.finished_at or time.time()
|
||||
self._jobs[job.id] = job
|
||||
# If a job somehow survived as alive, keep it active (best-effort re-attach of state).
|
||||
for job in self._jobs.values():
|
||||
if job.status not in TERMINAL and _pid_alive(job.pid):
|
||||
self._active_id = job.id
|
||||
self._proc = None # can't re-wrap the Popen; monitor via pid liveness
|
||||
threading.Thread(target=self._reattach_monitor, args=(job.id,), daemon=True).start()
|
||||
self._persist()
|
||||
|
||||
# ---- public API -----------------------------------------------------------------
|
||||
def submit(self, req: dict) -> dict:
|
||||
"""Start a train. Raises InvalidTrainRequest (→422) or Busy (→409)."""
|
||||
argv, env_overlay, params = build_command(req) # validates; raises InvalidTrainRequest (→422)
|
||||
with self._lock:
|
||||
if self._active_id and self._jobs[self._active_id].status not in TERMINAL:
|
||||
raise Busy(f"a training job is already running: {self._active_id}")
|
||||
job_id = "wjob_" + uuid.uuid4().hex[:12]
|
||||
log_path = config.LOG_DIR / f"{job_id}.log"
|
||||
job = Job(job_id, params, log_path)
|
||||
self._jobs[job_id] = job
|
||||
self._active_id = job_id
|
||||
self._trim()
|
||||
self._persist()
|
||||
threading.Thread(target=self._run, args=(job_id, argv, env_overlay), daemon=True).start()
|
||||
return {"worker_job_id": job_id}
|
||||
|
||||
def get(self, job_id: str) -> Optional[dict]:
|
||||
with self._lock:
|
||||
job = self._jobs.get(job_id)
|
||||
return job.public() if job else None
|
||||
|
||||
def log_tail(self, job_id: str, max_bytes: int = 16384) -> Optional[str]:
|
||||
job = self._jobs.get(job_id)
|
||||
if not job or not job.log_path.exists():
|
||||
return None
|
||||
with open(job.log_path, "rb") as fh:
|
||||
fh.seek(0, os.SEEK_END)
|
||||
size = fh.tell()
|
||||
fh.seek(max(0, size - max_bytes))
|
||||
return fh.read().decode("utf-8", "replace")
|
||||
|
||||
def cancel(self, job_id: str) -> Optional[dict]:
|
||||
with self._lock:
|
||||
job = self._jobs.get(job_id)
|
||||
if not job:
|
||||
return None
|
||||
if job.status in TERMINAL:
|
||||
return job.public()
|
||||
pid = job.pid
|
||||
# best-effort kill of the whole process group (accelerate spawns children)
|
||||
if pid:
|
||||
try:
|
||||
os.killpg(os.getpgid(pid), signal.SIGTERM)
|
||||
except (OSError, ProcessLookupError):
|
||||
pass
|
||||
with self._lock:
|
||||
job.status = "cancelled"
|
||||
job.error = "cancelled by request"
|
||||
job.finished_at = time.time()
|
||||
self._persist()
|
||||
return job.public()
|
||||
|
||||
def active(self) -> Optional[str]:
|
||||
with self._lock:
|
||||
if self._active_id and self._jobs[self._active_id].status not in TERMINAL:
|
||||
return self._active_id
|
||||
return None
|
||||
|
||||
# ---- internals ------------------------------------------------------------------
|
||||
def _trim(self) -> None:
|
||||
if len(self._jobs) <= config.MAX_RETAINED_JOBS:
|
||||
return
|
||||
# drop oldest terminal jobs first
|
||||
terminal = [i for i, j in self._jobs.items() if j.status in TERMINAL]
|
||||
for i in terminal[: len(self._jobs) - config.MAX_RETAINED_JOBS]:
|
||||
self._jobs.pop(i, None)
|
||||
|
||||
def _set(self, job_id: str, **kw) -> None:
|
||||
with self._lock:
|
||||
job = self._jobs.get(job_id)
|
||||
if not job:
|
||||
return
|
||||
for k, v in kw.items():
|
||||
setattr(job, k, v)
|
||||
self._persist()
|
||||
|
||||
def _run(self, job_id: str, argv: list[str], env_overlay: dict[str, str]) -> None:
|
||||
job = self._jobs[job_id]
|
||||
self._set(job_id, status="preparing", started_at=time.time())
|
||||
env = {**os.environ, **env_overlay}
|
||||
try:
|
||||
with open(job.log_path, "wb") as logf:
|
||||
proc = subprocess.Popen(
|
||||
argv, cwd=str(config.SD_SCRIPTS_DIR), env=env,
|
||||
stdout=logf, stderr=subprocess.STDOUT,
|
||||
start_new_session=True, # own process group → cancel can killpg
|
||||
)
|
||||
with self._lock:
|
||||
self._proc = proc
|
||||
job.pid = proc.pid
|
||||
job.status = "training"
|
||||
self._persist()
|
||||
self._poll_until_exit(job_id, proc)
|
||||
except (OSError, ValueError) as exc:
|
||||
self._set(job_id, status="failed", error=f"spawn failed: {exc}", finished_at=time.time())
|
||||
return
|
||||
|
||||
def _poll_until_exit(self, job_id: str, proc: subprocess.Popen) -> None:
|
||||
job = self._jobs[job_id]
|
||||
while proc.poll() is None:
|
||||
self._update_progress(job_id)
|
||||
time.sleep(3)
|
||||
rc = proc.returncode
|
||||
self._update_progress(job_id)
|
||||
with self._lock:
|
||||
if job.status == "cancelled":
|
||||
return # cancel already set the terminal state
|
||||
if rc == 0:
|
||||
lora = self._find_lora(job)
|
||||
if lora:
|
||||
job.status, job.lora_path = "succeeded", str(lora)
|
||||
else:
|
||||
job.status, job.error = "failed", "process exited 0 but no .safetensors found"
|
||||
else:
|
||||
job.status = "failed"
|
||||
job.error = job.error or f"training process exited {rc} (see log)"
|
||||
job.finished_at = time.time()
|
||||
self._persist()
|
||||
|
||||
def _reattach_monitor(self, job_id: str) -> None:
|
||||
# A job whose pid survived a worker restart: watch pid liveness (no Popen handle).
|
||||
job = self._jobs[job_id]
|
||||
while _pid_alive(job.pid):
|
||||
self._update_progress(job_id)
|
||||
time.sleep(3)
|
||||
self._update_progress(job_id)
|
||||
with self._lock:
|
||||
if job.status not in TERMINAL:
|
||||
lora = self._find_lora(job)
|
||||
job.status = "succeeded" if lora else "failed"
|
||||
job.lora_path = str(lora) if lora else None
|
||||
if not lora:
|
||||
job.error = "reattached process ended without a .safetensors"
|
||||
job.finished_at = time.time()
|
||||
self._persist()
|
||||
|
||||
def _update_progress(self, job_id: str) -> None:
|
||||
job = self._jobs.get(job_id)
|
||||
if not job or not job.log_path.exists():
|
||||
return
|
||||
try:
|
||||
with open(job.log_path, "rb") as fh:
|
||||
fh.seek(0, os.SEEK_END)
|
||||
fh.seek(max(0, fh.tell() - 8192))
|
||||
tail = fh.read().decode("utf-8", "replace")
|
||||
except OSError:
|
||||
return
|
||||
# tqdm uses \r; split on both so we see the latest bar frame
|
||||
frames = re.split(r"[\r\n]", tail)
|
||||
step, loss = job.step, job.loss
|
||||
for frame in frames:
|
||||
m = _STEP_RE.search(frame)
|
||||
if m and int(m.group(2)) == job.total_steps: # match the step bar, not a random N/M
|
||||
step = int(m.group(1))
|
||||
lm = _LOSS_RE.search(frame)
|
||||
if lm:
|
||||
loss = float(lm.group(1))
|
||||
eta = None
|
||||
if step > 0 and job.started_at and job.total_steps:
|
||||
elapsed = time.time() - job.started_at
|
||||
eta = int((job.total_steps - step) * (elapsed / step))
|
||||
self._set(job_id, step=step, loss=loss, eta_s=eta)
|
||||
|
||||
def _find_lora(self, job: Job) -> Optional[Path]:
|
||||
out_dir = Path(job.params.get("output_dir", ""))
|
||||
name = job.params.get("output_name", "")
|
||||
candidate = out_dir / f"{name}.safetensors"
|
||||
if candidate.exists():
|
||||
return candidate
|
||||
# fall back to newest .safetensors in the output dir
|
||||
try:
|
||||
sfts = sorted(out_dir.glob("*.safetensors"), key=lambda p: p.stat().st_mtime)
|
||||
return sfts[-1] if sfts else None
|
||||
except OSError:
|
||||
return None
|
||||
|
||||
|
||||
class Busy(RuntimeError):
|
||||
"""A second train was dispatched while one is active → HTTP 409."""
|
||||
|
||||
|
||||
# module-level singleton
|
||||
manager = JobManager()
|
||||
Reference in New Issue
Block a user