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:
parent
aa6d4ff5ef
commit
74e7854e15
3 changed files with 157 additions and 3 deletions
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue