Diffusion training API: LLM interlock, pre-spawn VRAM free, path containment, no dropped knobs

Four review findings on the diffusion training start path:
- It spawned the SDXL trainer without checking the LLM TrainingBackend, so a
  start while an LLM run was active put two trainers on the same GPU. Add a
  symmetric interlock: diffusion start returns 409 when LLM training is active,
  and LLM start refuses while a diffusion job is active.
- It went straight to service.start() without freeing GPU residents. Add a
  pre-spawn free of the export subprocess, the resident Images pipeline (with an
  arbiter release), and chat models, mirroring the LLM start path.
- data_dir / output_dir were passed through unresolved, so Studio-relative names
  failed and absolute paths bypassed containment. Resolve them with
  resolve_dataset_path / resolve_output_dir before spawn (400 on an uncontained
  path).
- The request model dropped max_grad_norm and lora_target_modules, so runs that
  set them trained with defaults. Add both fields.

The gemini pump-join deadlock was already fixed earlier (join outside the lock +
proc-identity fence). Note: honoring a stop DURING model load is a trainer-loop
change owned by the diffusion training engine PR (should_stop polled before the
first optimizer step). Adds route + model regression tests.
This commit is contained in:
Daniel Han 2026-07-02 01:09:56 +00:00
commit 74e7854e15
3 changed files with 157 additions and 3 deletions

View file

@ -693,6 +693,14 @@ class DiffusionTrainingStartRequest(BaseModel):
lora_rank: int = Field(16, ge = 1, le = 320)
lora_alpha: Optional[int] = Field(None, ge = 1, le = 640, description = "Defaults to lora_rank")
lora_dropout: float = Field(0.0, ge = 0.0, le = 1.0)
# Mirror the remaining training-affecting knobs of DiffusionLoraConfig so a client that
# sets them is not silently trained with defaults. Default the target list to the SDXL
# attention projections (the trainer's DEFAULT_LORA_TARGETS) so it is never None.
lora_target_modules: List[str] = Field(
default_factory = lambda: ["to_k", "to_q", "to_v", "to_out.0"],
description = "U-Net modules to attach LoRA to",
)
max_grad_norm: float = Field(1.0, gt = 0, description = "Gradient clipping max-norm")
seed: int = Field(42)
mixed_precision: Literal["bf16", "fp16", "no"] = Field("bf16")
snr_gamma: Optional[float] = Field(5.0, description = "Min-SNR loss weighting; null disables")

View file

@ -159,6 +159,20 @@ async def start_training(
error = "Training already active",
)
# A diffusion (SDXL) LoRA job runs in its own subprocess on the same GPU, so an
# LLM start must also refuse while one is active -- otherwise the two trainers
# contend for VRAM and both fail. Symmetric with the check in start_diffusion_training.
if _diffusion_training_active():
return TrainingJobResponse(
job_id = "",
status = "error",
message = (
"A diffusion (Images) LoRA training job is already running. "
"Stop it before starting an LLM training run."
),
error = "Diffusion training already active",
)
# Job ID; start_training() sets it on the backend only after the old
# pump thread is dead.
job_id = f"job_{datetime.now().strftime('%Y%m%d_%H%M%S')}_{_uuid.uuid4().hex[:8]}"
@ -1031,6 +1045,62 @@ async def stream_training_progress(
# lifecycle (DB run rows, plots, transfer-to-chat-inference).
def _diffusion_training_active() -> bool:
"""Whether a diffusion (SDXL) LoRA job is currently running. Best-effort so the
interlock never blocks a start just because the service could not be imported."""
try:
from core.training.diffusion_training_service import get_diffusion_training_service
return get_diffusion_training_service().is_active()
except Exception: # noqa: BLE001
return False
def _free_gpu_for_diffusion_training() -> None:
"""Free GPU residents before the diffusion trainer spawns its own SDXL pipeline.
The trainer subprocess loads a full SDXL pipeline; an export worker, a resident
Images pipeline, or loaded chat models would otherwise keep their VRAM allocated and
OOM the run. Mirrors the LLM start path's pre-spawn cleanup (export + diffusion
pipeline + chat). Best-effort: a failure to free one resident never blocks the start."""
try:
from core.export import get_export_backend
exp_backend = get_export_backend()
if exp_backend.current_checkpoint or exp_backend.is_export_active():
logger.info("Shutting down export subprocess to free GPU memory for diffusion training")
exp_backend._shutdown_subprocess()
exp_backend.current_checkpoint = None
exp_backend.is_vision = False
exp_backend.is_peft = False
except Exception as e: # noqa: BLE001
logger.warning("Could not shut down export subprocess: %s", e)
try:
from core.inference import gpu_arbiter
from core.inference.diffusion import get_diffusion_backend
diffusion = get_diffusion_backend()
if diffusion.is_loaded:
logger.info("Unloading resident Images pipeline to free GPU memory for training")
diffusion.unload() # no-op when nothing is loaded; also preempts an in-flight load
gpu_arbiter.release(gpu_arbiter.DIFFUSION)
except Exception as e: # noqa: BLE001
logger.warning("Could not unload Images pipeline for diffusion training: %s", e)
try:
# The SDXL trainer's footprint can't be cheaply sized against a resident chat
# model, so free chat unconditionally (same conservative choice the LLM path
# makes for an in-flight chat load) rather than risk an OOM.
from routes.training_vram import free_chat_models_for_training, summarize_resident_chat
if summarize_resident_chat()["any"]:
freed = free_chat_models_for_training(reason = "diffusion training starting")
logger.info("Freed chat model(s) for diffusion training: %s", freed)
except Exception as e: # noqa: BLE001
logger.warning("Could not free chat models for diffusion training: %s", e)
@router.post("/diffusion/start", response_model = DiffusionTrainingStartResponse)
async def start_diffusion_training(
body: DiffusionTrainingStartRequest, current_subject: str = Depends(get_current_subject)
@ -1038,9 +1108,41 @@ async def start_diffusion_training(
"""Start an SDXL LoRA training job from an image + caption dataset."""
from core.training.diffusion_training_service import get_diffusion_training_service
# Interlock: refuse while an LLM training run holds the GPU (symmetric with the
# diffusion check in start_training), so the two trainers never contend for VRAM.
try:
if get_training_backend().is_training_active():
raise HTTPException(
status_code = 409,
detail = (
"An LLM training job is already running. "
"Stop it before starting diffusion (Images) training."
),
)
except HTTPException:
raise
except Exception: # noqa: BLE001 -- backend import/health issue must not block a start
pass
# Resolve + contain the dataset and output paths BEFORE spawning, so Studio-relative
# names ("uploads/my-images") work and absolute paths stay under a Studio root -- the
# trainer subprocess otherwise resolves them relative to its own cwd.
config = body.model_dump()
try:
from utils.paths import resolve_dataset_path, resolve_output_dir
config["data_dir"] = str(resolve_dataset_path(config["data_dir"]))
config["output_dir"] = str(resolve_output_dir(config["output_dir"]))
except ValueError as e:
raise HTTPException(status_code = 400, detail = str(e))
# Free resident GPU workloads (export / Images pipeline / chat) before the trainer
# loads its own SDXL pipeline.
_free_gpu_for_diffusion_training()
service = get_diffusion_training_service()
try:
job_id = service.start(body.model_dump())
job_id = service.start(config)
except ValueError as e:
raise HTTPException(status_code = 400, detail = str(e))
except RuntimeError as e:

View file

@ -225,12 +225,26 @@ class _FakeService:
}
class _FakeLLMBackend:
def __init__(self, active = False):
self._active = active
def is_training_active(self):
return self._active
@pytest.fixture
def client(monkeypatch):
fake = _FakeService()
monkeypatch.setattr(
"core.training.diffusion_training_service.get_diffusion_training_service", lambda: fake
)
# Neutralize the LLM interlock + GPU-free for the wiring tests (their own tests below
# exercise those behaviors). The route imports get_training_backend at module scope.
import routes.training as tr
monkeypatch.setattr(tr, "get_training_backend", lambda: _FakeLLMBackend(active = False))
monkeypatch.setattr(tr, "_free_gpu_for_diffusion_training", lambda: None)
app = FastAPI()
app.include_router(training_router, prefix = "/api/train")
app.dependency_overrides[get_current_subject] = lambda: "test-user"
@ -239,10 +253,11 @@ def client(monkeypatch):
return c
# Studio-relative paths: the route resolves/contains them before spawn.
_BODY = {
"base_model": "stabilityai/sdxl-turbo",
"data_dir": "/data",
"output_dir": "/out",
"data_dir": "uploads/my-images",
"output_dir": "my-lora-run",
"train_steps": 10,
}
@ -252,6 +267,35 @@ def test_route_start_ok(client):
assert r.status_code == 200, r.text
assert r.json() == {"job_id": "job-123", "status": "running"}
assert client._fake.started_with["base_model"] == "stabilityai/sdxl-turbo"
# Paths were resolved to absolute Studio-contained locations before spawn.
from pathlib import Path
assert Path(client._fake.started_with["data_dir"]).is_absolute()
assert Path(client._fake.started_with["output_dir"]).is_absolute()
def test_route_start_forwards_extra_training_knobs(client):
# max_grad_norm and lora_target_modules must reach the service, not be silently dropped.
body = {**_BODY, "max_grad_norm": 0.5, "lora_target_modules": ["to_q", "to_v"]}
r = client.post("/api/train/diffusion/start", json = body)
assert r.status_code == 200, r.text
assert client._fake.started_with["max_grad_norm"] == 0.5
assert client._fake.started_with["lora_target_modules"] == ["to_q", "to_v"]
def test_route_start_rejects_uncontained_paths(client):
# An absolute path outside the Studio dataset roots is a 400, not silently accepted.
r = client.post("/api/train/diffusion/start", json = {**_BODY, "data_dir": "/etc"})
assert r.status_code == 400
def test_route_start_blocked_by_active_llm_training(client, monkeypatch):
import routes.training as tr
monkeypatch.setattr(tr, "get_training_backend", lambda: _FakeLLMBackend(active = True))
r = client.post("/api/train/diffusion/start", json = _BODY)
assert r.status_code == 409
assert "LLM training" in r.json()["detail"]
def test_route_start_missing_required_is_422(client):