From eee5a53658c35db485cc3bba54040356b0986cf5 Mon Sep 17 00:00:00 2001 From: Daniel Han Date: Tue, 7 Jul 2026 08:55:23 +0000 Subject: [PATCH] Offload the gated-base HEAD preflight to a worker thread start_diffusion_training is async but called _preflight_gated_base inline; it does a blocking urllib urlopen HEAD to Hugging Face (up to a 5s timeout) to detect a gated/unauthorized base repo. On a slow or unreachable network that stalled the FastAPI event loop, freezing every concurrent status/progress/cancel request until it returned or timed out. Wrap it in asyncio.to_thread, matching the dataset preflight and GPU cleanup just below it. Regression test asserts it runs off the coroutine thread. --- studio/backend/routes/training.py | 7 ++++- .../backend/tests/test_diffusion_training.py | 30 +++++++++++++++++++ 2 files changed, 36 insertions(+), 1 deletion(-) diff --git a/studio/backend/routes/training.py b/studio/backend/routes/training.py index d663c31034..b9fdc1ee04 100644 --- a/studio/backend/routes/training.py +++ b/studio/backend/routes/training.py @@ -1301,7 +1301,12 @@ async def start_diffusion_training( # Preflight access to a gated base repo with the user's token BEFORE freeing GPU # residents, so a missing/insufficient token fails fast (400) without tearing down the # user's loaded chat/Images model, and never surfaces as a confusing mid-load 401. - _preflight_gated_base(config.get("base_model", ""), config.get("hf_token")) + # Offloaded to a worker thread: it does a blocking urlopen HEAD (up to a 5s timeout) to + # Hugging Face, which would otherwise stall the event loop and every concurrent + # status/progress/cancel request, as the filesystem preflight just below already does. + await asyncio.to_thread( + _preflight_gated_base, config.get("base_model", ""), config.get("hf_token") + ) # Preflight the dataset too: a missing/empty/uncaptionable data_dir otherwise # fails inside the spawned trainer AFTER the user's chat/Images model was diff --git a/studio/backend/tests/test_diffusion_training.py b/studio/backend/tests/test_diffusion_training.py index 8c49e9a5ed..a5277ccda6 100644 --- a/studio/backend/tests/test_diffusion_training.py +++ b/studio/backend/tests/test_diffusion_training.py @@ -391,6 +391,36 @@ def test_route_start_frees_gpu_off_the_coroutine_thread(client, monkeypatch): assert threads["cleanup"] is not threads["inline"] # offloaded to a worker, not run inline +def test_route_start_preflights_gated_base_off_the_coroutine_thread(client, monkeypatch): + # _preflight_gated_base does a blocking urlopen HEAD (up to a 5s timeout) to Hugging Face, so + # the async start route must offload it via asyncio.to_thread rather than run it inline and + # freeze the event loop for concurrent status/progress/cancel requests. Assert it runs on a + # DIFFERENT thread than the inline coroutine body (service.start), which an inline call could + # not. + import threading + + import routes.training as tr + + threads: dict = {} + + def _record_preflight(base_model, hf_token): + threads["preflight"] = threading.current_thread() + + monkeypatch.setattr(tr, "_preflight_gated_base", _record_preflight) + + orig_start = client._fake.start + + def _record_start(config): + threads["inline"] = threading.current_thread() + return orig_start(config) + + monkeypatch.setattr(client._fake, "start", _record_start) + + r = client.post("/api/train/diffusion/start", json = _BODY) + assert r.status_code == 200, r.text + assert threads["preflight"] is not threads["inline"] # offloaded to a worker, not run inline + + 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"]}