diff --git a/studio/backend/core/inference/video.py b/studio/backend/core/inference/video.py index ee6cd6f7c2..977d284eca 100644 --- a/studio/backend/core/inference/video.py +++ b/studio/backend/core/inference/video.py @@ -520,7 +520,13 @@ class VideoBackend: raise RuntimeError("A video load is already in progress.") self._load_token += 1 token = self._load_token - self._cancel_event.clear() + # A NEW event per load, never a clear() of the shared one. unload() sets the event the + # running worker holds, but it also drops _loading, so the next begin_load could start + # before that worker had exited and clear the very object it was watching -- the + # cancelled multi-gigabyte pull then resumed and ran alongside the replacement load + # until the token check at the end. A fresh object leaves the old worker's event set. + cancel_event = threading.Event() + self._cancel_event = cancel_event self._loading = _VideoLoadingState(repo_id = repo_id, base_repo = fam.base_repo) threading.Thread( @@ -540,6 +546,7 @@ class VideoBackend: text_encoder_quant = text_encoder_quant, model_kind = model_kind, _load_token = token, + _cancel_event = cancel_event, ), daemon = True, ).start() @@ -547,6 +554,9 @@ class VideoBackend: def _run_load(self, **kwargs: Any) -> None: token = kwargs.get("_load_token") + # This load's own event: a later load replaces self._cancel_event rather than clearing it, + # so a cancelled worker stays cancelled. + cancel_event = kwargs.pop("_cancel_event", None) or self._cancel_event try: fam = _detect_load_family( kwargs["repo_id"], kwargs.get("gguf_filename"), kwargs.get("family_override") @@ -583,7 +593,7 @@ class VideoBackend: kwargs["repo_id"], kwargs["gguf_filename"], kwargs.get("hf_token"), - cancel_event = self._cancel_event, + cancel_event = cancel_event, ) ) # An LTX-2.3 checkpoint supplies the VAEs/vocoder/connectors, so the base pull shrinks to @@ -624,7 +634,9 @@ class VideoBackend: if self._load_token == token and self._loading is not None: self._loading.expected_bytes = expected # Only a pre-cast checkpoint actually on disk earns the dense skip below. - te_skipped = self._fetch_te_prequant(te_sources, kwargs.get("hf_token")) + te_skipped = self._fetch_te_prequant( + te_sources, kwargs.get("hf_token"), cancel_event = cancel_event + ) kwargs["_te_prequant_skipped"] = te_skipped base_local = self._predownload_base( base, @@ -632,6 +644,7 @@ class VideoBackend: kind, ltx23 = ltx23, skip_te_components = te_skipped, + cancel_event = cancel_event, ) # The 2.3 assembly pulls per component from the hub id (its snapshot lacks the base VAEs), so it # only gets the warmed cache; generic paths get the full local snapshot. @@ -936,7 +949,11 @@ class VideoBackend: return "2.3" in text or "2_3" in text or "23b" in text def _fetch_te_prequant( - self, sources: dict[str, Any], hf_token: Optional[str] + self, + sources: dict[str, Any], + hf_token: Optional[str], + *, + cancel_event: Optional[threading.Event] = None, ) -> tuple[str, ...]: """Pre-fetch the hosted pre-cast encoder checkpoints; return the components whose dense weights the base pull can therefore skip. @@ -946,6 +963,7 @@ class VideoBackend: drop the dense shards, and the pull becomes cancellable and resumable like every other load download instead of an untracked stall. A component whose fetch fails keeps its dense weights, so the load still has an encoder to fall back to.""" + cancel = cancel_event if cancel_event is not None else self._cancel_event from utils.hf_xet_fallback import hf_hub_download_with_xet_fallback fetched: list[str] = [] @@ -959,10 +977,10 @@ class VideoBackend: source.location, source.filename, hf_token, - cancel_event = self._cancel_event, + cancel_event = cancel, ) except Exception as exc: # noqa: BLE001 -- no pre-cast file just means the dense encoder - if self._cancel_event.is_set(): + if cancel.is_set(): raise logger.warning( "video.te_prequant_fetch_failed: %s/%s: %s", @@ -982,6 +1000,7 @@ class VideoBackend: *, ltx23: bool = False, skip_te_components: tuple[str, ...] = (), + cancel_event: Optional[threading.Event] = None, ) -> Optional[str]: """Pull exactly the base-repo files the load needs; return the local snapshot dir. @@ -992,6 +1011,7 @@ class VideoBackend: diffusers' own expected-files sweep. None -> caller keeps the hub id (local path, non-diffusers layout, or any metadata failure: from_pretrained then resolves the repo exactly as before).""" + cancel = cancel_event if cancel_event is not None else self._cancel_event try: if not base or Path(base).expanduser().exists(): return None @@ -1009,18 +1029,18 @@ class VideoBackend: for name, _ in files: # Explicit check: a cached file returns without consulting the event, so a warm-cache sweep would # otherwise run to completion after an unload cancelled. - if self._cancel_event.is_set(): + if cancel.is_set(): raise RuntimeError(VIDEO_CANCELLED_MSG) local = Path( hf_hub_download_with_xet_fallback( - base, name, hf_token, cancel_event = self._cancel_event + base, name, hf_token, cancel_event = cancel ) ) if name == "model_index.json": snapshot_root = local.parent return str(snapshot_root) if snapshot_root is not None else None except Exception as exc: # noqa: BLE001 -- fall back to from_pretrained's own pull - if self._cancel_event.is_set(): + if cancel.is_set(): raise logger.warning("video.predownload_fallback: %s", exc) return None diff --git a/studio/backend/routes/models.py b/studio/backend/routes/models.py index e7d6b188e6..71af09f8be 100644 --- a/studio/backend/routes/models.py +++ b/studio/backend/routes/models.py @@ -3651,12 +3651,18 @@ def _cached_repo_task(repo_info) -> Optional[str]: pass if not _repo_is_diffusers(repo_info): return None - # Same trust rule the image load path applies, so every advertised row can actually load: a - # cached community pipeline has a model_index.json like any other, but validate_load_request - # refuses its repo id, so tagging it text-to-image put a row in the Images picker that 400s. + # BOTH gates, mirroring the video branch above: the load path's trust rule AND a detected image + # family. A model_index.json only proves the repo is a diffusers pipeline, not that it is one of + # the image families this backend can assemble, so an unsloth-hosted pipeline of an unsupported + # class cleared the trust gate, was advertised as text-to-image, and then deterministically + # failed validate_load_request, which detects the family the same way. try: from core.inference.diffusion import _is_trusted_diffusion_repo - return "text-to-image" if _is_trusted_diffusion_repo(repo_id) else None + from core.inference.diffusion_families import detect_family + + if not _is_trusted_diffusion_repo(repo_id) or detect_family(repo_id) is None: + return None + return "text-to-image" except Exception: # noqa: BLE001 -- an import failure must not hide a usable repo return "text-to-image" diff --git a/studio/backend/routes/training.py b/studio/backend/routes/training.py index 64f0c71213..ad10021c0a 100644 --- a/studio/backend/routes/training.py +++ b/studio/backend/routes/training.py @@ -1394,6 +1394,22 @@ async def start_diffusion_training( except ValueError as e: raise HTTPException(status_code = 400, detail = str(e)) + # Only the DiT trainer reads cond_cache_dir. The SDXL trainer builds a per-process in-memory + # latent cache and never touches the persistent store, so accepting the option there promised + # cross-run reuse that never happened and silently re-encoded the dataset every run. Refuse it + # instead of ignoring it. Checked against the RESOLVED family, not the request field, so a + # request that omits model_family and lets an SDXL base be detected is caught too. + if cond_cache and normalized_cfg.resolved_family == "sdxl": + raise HTTPException( + status_code = 400, + detail = ( + "cond_cache_dir is not supported for the sdxl family: its trainer uses a " + "per-run in-memory latent cache and would ignore the persistent one. Omit it, " + "or train a DiT family (flux.1, flux.2-klein, flux.2-dev, qwen-image, " + "z-image, krea-2), which reuses conditioning across runs." + ), + ) + # Preflight the requested DiT precision BEFORE freeing GPU residents: the trainer's own checks # (bf16-capable GPU required; explicit int8 needs a functional torchao) fire only in the child, # AFTER _free_gpu_for_diffusion_training() evicted the user's model. Fail fast (400) so a diff --git a/studio/backend/tests/test_cached_gguf_routes.py b/studio/backend/tests/test_cached_gguf_routes.py index 57f8246a42..90af537a66 100644 --- a/studio/backend/tests/test_cached_gguf_routes.py +++ b/studio/backend/tests/test_cached_gguf_routes.py @@ -1811,3 +1811,46 @@ def test_pipeline_class_guard_fires_before_any_download(): assert "0.39" in msg and "0.37.0" in msg assert "3.10" in msg # names the Python floor that carries a new enough diffusers assert diffusers is not None + + +def test_cached_pipeline_needs_a_detectable_image_family(monkeypatch): + # A top-level model_index.json only proves the repo is a diffusers pipeline. An unsloth-hosted + # pipeline of a class this backend cannot assemble cleared the trust gate, was advertised to the + # Images picker as text-to-image, and then deterministically failed validate_load_request, which + # resolves the family the same way. Both gates now, mirroring the video branch above. + monkeypatch.setattr(models_route, "_repo_has_pipeline_index", lambda info: True) + + def _task(repo_id): + return models_route._cached_repo_task(SimpleNamespace(repo_id = repo_id, repo_path = "/x")) + + # Trusted AND a detected family -> claimed by Images. + assert _task("unsloth/Z-Image-Turbo") == "text-to-image" + assert _task("unsloth/FLUX.1-dev") == "text-to-image" + # Trusted but no image family the loader can detect -> not advertised. + assert _task("unsloth/some-unsupported-pipeline") is None + # Untrusted keeps its existing refusal. + assert _task("someone/random-diffusers-pipeline") is None + + +def test_cached_repo_task_agrees_with_the_image_loader(monkeypatch): + # Same invariant as the GGUF arch test: whatever the picker advertises as a loadable image + # model, validate_load_request must accept. + from core.inference.diffusion import DiffusionBackend + + monkeypatch.setattr(models_route, "_repo_has_pipeline_index", lambda info: True) + backend = DiffusionBackend.__new__(DiffusionBackend) + for repo_id in ( + "unsloth/Z-Image-Turbo", + "unsloth/FLUX.1-dev", + "unsloth/some-unsupported-pipeline", + "unsloth/stable-audio-open-1.0", + ): + task = models_route._cached_repo_task(SimpleNamespace(repo_id = repo_id, repo_path = "/x")) + try: + backend.validate_load_request(repo_id) + loader_accepts = True + except (ValueError, FileNotFoundError, RuntimeError): + loader_accepts = False + assert (task == "text-to-image") == loader_accepts, ( + f"{repo_id}: picker task={task} but loader accepts={loader_accepts}" + ) diff --git a/studio/backend/tests/test_diffusion_training.py b/studio/backend/tests/test_diffusion_training.py index 8f0bd01f03..7716a35ad0 100644 --- a/studio/backend/tests/test_diffusion_training.py +++ b/studio/backend/tests/test_diffusion_training.py @@ -1918,10 +1918,14 @@ def test_route_start_carries_and_contains_the_conditioning_cache_dir(client): # multi-GB text encoders on a rerun, but the start schema omitted the field, so Pydantic # dropped it silently and every API-driven run fell back to the in-memory cache. It also has to # be contained like output_dir: the trainer subprocess would otherwise resolve it against its - # own cwd. + # own cwd. Sent against a DiT family, the only one whose trainer reads the option (see + # test_route_rejects_cond_cache_dir_for_sdxl). from pathlib import Path - r = client.post("/api/train/diffusion/start", json = {**_BODY, "cond_cache_dir": "cond-cache"}) + r = client.post( + "/api/train/diffusion/start", + json = {**_BODY, "model_family": "z-image", "cond_cache_dir": "cond-cache"}, + ) assert r.status_code == 200, r.text resolved = client._fake.started_with["cond_cache_dir"] assert Path(resolved).is_absolute() @@ -1934,3 +1938,43 @@ def test_route_start_carries_and_contains_the_conditioning_cache_dir(client): r = client.post("/api/train/diffusion/start", json = {**_BODY, "cond_cache_dir": " "}) assert r.status_code == 200, r.text assert client._fake.started_with["cond_cache_dir"] is None + + +def test_route_rejects_cond_cache_dir_for_sdxl(client): + # Only the DiT trainer reads cond_cache_dir. The SDXL trainer builds a per-process in-memory + # latent cache and never touches the persistent store, so accepting the option there promised + # cross-run reuse that never happened. Refuse it rather than ignore it. + r = client.post( + "/api/train/diffusion/start", + json = {**_BODY, "model_family": "sdxl", "cond_cache_dir": "cond-cache"}, + ) + assert r.status_code == 400, r.text + detail = r.json()["detail"] + assert "cond_cache_dir" in detail and "sdxl" in detail + # It names the families that DO support it, so the message is actionable. + assert "z-image" in detail + + # The check is on the RESOLVED family, so omitting model_family and letting an SDXL base be + # detected is refused too -- otherwise the common request shape kept the silent no-op. + r = client.post( + "/api/train/diffusion/start", + json = {**_BODY, "cond_cache_dir": "cond-cache"}, + ) + assert r.status_code == 400, r.text + assert "cond_cache_dir" in r.json()["detail"] + + # A DiT family still accepts it, resolved and contained like output_dir. + r = client.post( + "/api/train/diffusion/start", + json = {**_BODY, "model_family": "z-image", "cond_cache_dir": "cond-cache"}, + ) + assert r.status_code == 200, r.text + from pathlib import Path + + assert Path(client._fake.started_with["cond_cache_dir"]).is_absolute() + # And omitting it stays off (the trainer's in-memory default), not resolved to the outputs root. + r = client.post( + "/api/train/diffusion/start", json = {**_BODY, "model_family": "sdxl"} + ) + assert r.status_code == 200, r.text + assert client._fake.started_with["cond_cache_dir"] is None diff --git a/studio/backend/tests/test_video_backend.py b/studio/backend/tests/test_video_backend.py index 73266ef882..042e978acc 100644 --- a/studio/backend/tests/test_video_backend.py +++ b/studio/backend/tests/test_video_backend.py @@ -2196,3 +2196,49 @@ def test_download_plan_keeps_the_wide_base_for_a_plain_ltx2_pick(monkeypatch): assert any(f.startswith("connectors/") for f in base["files"]) # The GGUF still replaces the DiT, so the transformer shards stay out. assert not any(f.startswith("transformer/") for f in base["files"]) + + +def test_each_video_load_gets_its_own_cancel_event(monkeypatch): + """A cancelled load must STAY cancelled once the next one starts. + + unload() sets the event the running worker holds and drops _loading, so the next begin_load + could arrive before that worker had exited. Clearing a shared event there un-cancelled the old + worker, and its multi-gigabyte checkpoint pull resumed alongside the replacement load until the + token check at the very end. + """ + import threading + from types import SimpleNamespace + + from core.inference.video import VideoBackend + + backend = VideoBackend.__new__(VideoBackend) + backend._lock = threading.RLock() + backend._generate_lock = threading.RLock() + backend._cancel_event = threading.Event() + backend._load_token = 0 + backend._loading = None + backend._active_generate_cancel = None + backend._state = None + + started: list[threading.Event] = [] + monkeypatch.setattr( + threading, "Thread", lambda *a, **k: SimpleNamespace(start = lambda: None, daemon = True) + ) + fam = SimpleNamespace(base_repo = "org/base", name = "wan2.2-ti2v-5b") + monkeypatch.setattr(VideoBackend, "validate_load_request", lambda self, *a, **k: fam) + monkeypatch.setattr(VideoBackend, "status", lambda self: {}) + + backend.begin_load("org/base") + first = backend._cancel_event + started.append(first) + + # unload() signals the in-flight worker and clears _loading, letting a new load through. + backend._teardown_state = lambda: None + backend.unload() + assert first.is_set(), "unload must cancel the in-flight load" + + backend.begin_load("org/other") + second = backend._cancel_event + assert second is not first, "each load needs its own event" + assert first.is_set(), "the replaced load must stay cancelled" + assert not second.is_set(), "a fresh load starts uncancelled"