Fix/adjust diffusion lifecycle for round 3 findings (PR #5754)
- detect_family adds _FAMILY_EXCLUDE so 'stable-diffusion-3.5' no longer matches the SD3 Medium family and 'qwen-image-edit' no longer matches Qwen-Image. Both were misleading silent loads. - from_single_file now forwards config=<effective_base>, subfolder='transformer', and the HF token. Diffusers-format GGUFs (FLUX.2 klein, Qwen-Image, SD3) need the matching base config or the transformer load picks the wrong shapes; gated GGUFs need the token both for download and config read. - Move _release_chat_backend_for_diffusion + new _release_other_gpu_owners_for_diffusion to AFTER the GGUF download and pipeline class lookup so a typo or transient Hub error does not kill the user's currently-loaded chat model. Peak VRAM still stays at one model's worth because the releases run right before from_pretrained. - _release_other_gpu_owners_for_diffusion: shut down the export subprocess and any active training subprocess before a diffusion load. Symmetric with the export load path. - routes/training.py: unload diffusion before starting training so the new subprocess does not race FLUX/Qwen for VRAM. - routes/export.py: also unload the GGUF llama-server before export load (the existing inference-backend unload only covered the safetensors path).
This commit is contained in:
parent
0f9b19bb56
commit
65f7a2680d
4 changed files with 160 additions and 20 deletions
|
|
@ -172,16 +172,30 @@ def _smart_base_repo(fam: DiffusionFamily, repo_id: str) -> str:
|
|||
return "black-forest-labs/FLUX.2-klein-4B"
|
||||
|
||||
|
||||
# Negative substrings that disqualify a candidate family even when its
|
||||
# name appears as a substring of the repo id. Prevents
|
||||
# "stable-diffusion-3" matching SD3.5 and "qwen-image" matching
|
||||
# Qwen-Image-Edit. Each entry maps a family name to substrings that
|
||||
# must NOT appear anywhere in the repo id.
|
||||
_FAMILY_EXCLUDE: dict[str, tuple[str, ...]] = {
|
||||
"stable-diffusion-3": ("3.5", "3-5", "stable-diffusion-3.5"),
|
||||
"qwen-image": ("qwen-image-edit", "qwenimage-edit"),
|
||||
}
|
||||
|
||||
|
||||
def detect_family(
|
||||
repo_id: str, *, override_family: Optional[str] = None
|
||||
) -> Optional[DiffusionFamily]:
|
||||
"""Return the diffusion family matching ``repo_id``.
|
||||
|
||||
Matching is substring-based and case-insensitive. ``override_family``
|
||||
bypasses substring matching and looks up by ``DiffusionFamily.name``
|
||||
or (when explicitly asked) by ``_FULL_REPO_FAMILIES.name``.
|
||||
Returns ``None`` when no family applies so callers can surface a
|
||||
clear "unsupported model" error rather than guessing wrong.
|
||||
Matching is substring-based and case-insensitive, with a small
|
||||
deny list (``_FAMILY_EXCLUDE``) for known false positives such as
|
||||
SD3.5 (would otherwise match SD3 Medium) and Qwen-Image-Edit
|
||||
(would otherwise match Qwen-Image). ``override_family`` bypasses
|
||||
substring matching and looks up by ``DiffusionFamily.name`` or
|
||||
(when explicitly asked) by ``_FULL_REPO_FAMILIES.name``. Returns
|
||||
``None`` when no family applies so callers can surface a clear
|
||||
"unsupported model" error rather than guessing wrong.
|
||||
"""
|
||||
if override_family:
|
||||
wanted = override_family.strip().lower()
|
||||
|
|
@ -193,6 +207,9 @@ def detect_family(
|
|||
if not needle:
|
||||
return None
|
||||
for fam in _FAMILIES:
|
||||
excludes = _FAMILY_EXCLUDE.get(fam.name, ())
|
||||
if any(e in needle for e in excludes):
|
||||
continue
|
||||
if fam.name in needle:
|
||||
return fam
|
||||
for alias in fam.aliases:
|
||||
|
|
@ -345,12 +362,6 @@ class DiffusionBackend:
|
|||
self._loading = True
|
||||
self._last_error = None
|
||||
try:
|
||||
# Unload any chat model that is holding GPU memory so the
|
||||
# diffusion load does not OOM on a < 24 GB GPU. Best
|
||||
# effort: if the llama-cpp backend module is absent (eg
|
||||
# tests, headless tooling) we just continue.
|
||||
_release_chat_backend_for_diffusion()
|
||||
|
||||
pipeline_cls = getattr(diffusers, fam.pipeline_class, None)
|
||||
if pipeline_cls is None:
|
||||
raise RuntimeError(
|
||||
|
|
@ -412,10 +423,23 @@ class DiffusionBackend:
|
|||
token = hf_token,
|
||||
)
|
||||
quant_config = diffusers.GGUFQuantizationConfig(compute_dtype = dtype)
|
||||
# Diffusers-format GGUFs (FLUX.2 klein / Qwen-Image /
|
||||
# SD3) need the matching base repo's component config
|
||||
# at config=<base_repo>, subfolder="transformer".
|
||||
# Older city96-style GGUFs ignore those kwargs. The
|
||||
# token is also passed because gated GGUF repos
|
||||
# require it both at download and at config read time.
|
||||
single_file_kwargs: dict[str, Any] = {
|
||||
"quantization_config": quant_config,
|
||||
"torch_dtype": dtype,
|
||||
"config": effective_base,
|
||||
"subfolder": "transformer",
|
||||
}
|
||||
if hf_token:
|
||||
single_file_kwargs["token"] = hf_token
|
||||
transformer = transformer_cls.from_single_file(
|
||||
local_gguf_path,
|
||||
quantization_config = quant_config,
|
||||
torch_dtype = dtype,
|
||||
**single_file_kwargs,
|
||||
)
|
||||
|
||||
pipe_kwargs: dict[str, Any] = {
|
||||
|
|
@ -433,10 +457,15 @@ class DiffusionBackend:
|
|||
if hf_token:
|
||||
pipe_kwargs["token"] = hf_token
|
||||
|
||||
# Release the previous pipeline BEFORE allocating the
|
||||
# new one so peak VRAM stays at one model's worth, not
|
||||
# two. This matters on 16-24 GB consumer GPUs where the
|
||||
# combined footprint would OOM the from_pretrained call.
|
||||
# Cheap failure modes (bad gguf_filename, gated token,
|
||||
# transient Hub error) have all happened by now. Only
|
||||
# release the current chat backend + previous diffusion
|
||||
# pipeline right before the expensive allocation so a
|
||||
# typo does not kill the user's loaded chat model. Peak
|
||||
# VRAM still stays at one model's worth because the
|
||||
# release happens before from_pretrained.
|
||||
_release_chat_backend_for_diffusion()
|
||||
_release_other_gpu_owners_for_diffusion()
|
||||
old = self._pipe
|
||||
if old is not None:
|
||||
with self._lock:
|
||||
|
|
@ -644,6 +673,40 @@ def _release_chat_backend_for_diffusion() -> None:
|
|||
logger.debug("safetensors unload skipped: %s", exc)
|
||||
|
||||
|
||||
def _release_other_gpu_owners_for_diffusion() -> None:
|
||||
"""Best-effort: shut down export subprocess + active training before
|
||||
a diffusion load. Both can hold multi-GB of VRAM and would OOM the
|
||||
diffusion allocation on consumer GPUs."""
|
||||
# Export subprocess
|
||||
try:
|
||||
from core.export import get_export_backend # type: ignore
|
||||
|
||||
exp = get_export_backend()
|
||||
if getattr(exp, "current_checkpoint", None):
|
||||
logger.info("Shutting down export subprocess before diffusion load")
|
||||
exp._shutdown_subprocess()
|
||||
exp.current_checkpoint = None
|
||||
exp.is_vision = False
|
||||
exp.is_peft = False
|
||||
except Exception as exc:
|
||||
logger.debug("export unload skipped: %s", exc)
|
||||
|
||||
# Active training subprocess
|
||||
try:
|
||||
from core.training import get_training_backend # type: ignore
|
||||
|
||||
trn = get_training_backend()
|
||||
if trn.is_training_active():
|
||||
logger.info("Stopping training subprocess before diffusion load")
|
||||
trn.stop_training()
|
||||
for _ in range(60):
|
||||
if not trn.is_training_active():
|
||||
break
|
||||
time.sleep(0.5)
|
||||
except Exception as exc:
|
||||
logger.debug("training unload skipped: %s", exc)
|
||||
|
||||
|
||||
def _release(obj: Any) -> None:
|
||||
"""Best-effort GPU-memory release for a pipeline being swapped out."""
|
||||
if obj is None:
|
||||
|
|
|
|||
|
|
@ -81,6 +81,21 @@ async def load_checkpoint(
|
|||
except Exception as e:
|
||||
logger.warning("Could not unload inference model: %s", e)
|
||||
|
||||
# Also unload any active GGUF llama-server (the inference unload
|
||||
# above only covers the safetensors / Unsloth backend; GGUF
|
||||
# chat runs as a separate subprocess).
|
||||
try:
|
||||
from routes.inference import get_llama_cpp_backend
|
||||
|
||||
llama = get_llama_cpp_backend()
|
||||
if getattr(llama, "is_loaded", False):
|
||||
logger.info(
|
||||
"Unloading GGUF chat model to free GPU memory for export"
|
||||
)
|
||||
llama.unload_model()
|
||||
except Exception as e:
|
||||
logger.debug("llama-server unload skipped for export: %s", e)
|
||||
|
||||
# Also unload any active diffusion pipeline (Images page); it
|
||||
# competes for the same GPU and would survive the inference
|
||||
# shutdown above. Best effort; silently skip if the module is
|
||||
|
|
|
|||
|
|
@ -297,6 +297,20 @@ async def start_training(
|
|||
except Exception as e:
|
||||
logger.warning("Could not shut down export subprocess: %s", e)
|
||||
|
||||
# Also unload any loaded diffusion pipeline (Images page); it
|
||||
# holds the same GPU and would survive the inference shutdown.
|
||||
try:
|
||||
from core.inference.diffusion import get_diffusion_backend
|
||||
|
||||
diff_backend = get_diffusion_backend()
|
||||
if diff_backend.is_loaded:
|
||||
logger.info(
|
||||
"Unloading diffusion model to free GPU memory for training"
|
||||
)
|
||||
diff_backend.unload_model()
|
||||
except Exception as e:
|
||||
logger.warning("Could not unload diffusion model: %s", e)
|
||||
|
||||
# start_training now spawns a subprocess (non-blocking)
|
||||
success = backend.start_training(job_id = job_id, **training_kwargs)
|
||||
|
||||
|
|
|
|||
|
|
@ -111,6 +111,24 @@ def test_detect_family_unknown_returns_none():
|
|||
assert detect_family("") is None
|
||||
|
||||
|
||||
def test_detect_family_sd35_is_not_sd3():
|
||||
"""SD3.5 must NOT be matched as SD3 Medium. Pairing SD3.5 GGUFs
|
||||
with the Medium base produces a misleading load."""
|
||||
from core.inference.diffusion import detect_family
|
||||
|
||||
assert detect_family("unsloth/SD3.5-large-GGUF") is None
|
||||
assert detect_family("unsloth/stable-diffusion-3.5-large-GGUF") is None
|
||||
|
||||
|
||||
def test_detect_family_qwen_image_edit_is_not_qwen_image():
|
||||
"""Qwen-Image-Edit must NOT be matched as Qwen-Image. The Edit
|
||||
variant uses a different pipeline (image-to-image)."""
|
||||
from core.inference.diffusion import detect_family
|
||||
|
||||
assert detect_family("unsloth/Qwen-Image-Edit-GGUF") is None
|
||||
assert detect_family("unsloth/Qwen-Image-Edit-2509-GGUF") is None
|
||||
|
||||
|
||||
def test_supported_families_payload_shape():
|
||||
from core.inference.diffusion import supported_families
|
||||
|
||||
|
|
@ -285,11 +303,14 @@ def _install_fake_diffusers(monkeypatch, *, raise_on_pipeline = False):
|
|||
|
||||
class _FakeTransformer:
|
||||
@classmethod
|
||||
def from_single_file(cls, path, quantization_config = None, torch_dtype = None):
|
||||
def from_single_file(cls, path, **kw):
|
||||
inst = cls()
|
||||
inst.path = path
|
||||
inst.qc = quantization_config
|
||||
inst.dtype = torch_dtype
|
||||
inst.qc = kw.get("quantization_config")
|
||||
inst.dtype = kw.get("torch_dtype")
|
||||
inst.config = kw.get("config")
|
||||
inst.subfolder = kw.get("subfolder")
|
||||
inst.token = kw.get("token")
|
||||
return inst
|
||||
|
||||
class _FakePipeline:
|
||||
|
|
@ -477,6 +498,33 @@ def test_smart_base_repo_picks_base_4b(monkeypatch):
|
|||
assert status["base_repo"] == "black-forest-labs/FLUX.2-klein-base-4B"
|
||||
|
||||
|
||||
def test_gguf_transformer_load_passes_config_subfolder_token(monkeypatch):
|
||||
"""Diffusers-format GGUFs require config=<base_repo>+subfolder=
|
||||
transformer at from_single_file time; gated GGUFs also need the
|
||||
token. Verify all three kwargs are forwarded."""
|
||||
fake = _install_fake_diffusers(monkeypatch)
|
||||
from core.inference.diffusion import get_diffusion_backend
|
||||
|
||||
captured: dict = {}
|
||||
original = fake.Flux2Transformer2DModel.from_single_file.__func__
|
||||
|
||||
def _capture(cls, path, **kw):
|
||||
captured.update(kw)
|
||||
return original(cls, path, **kw)
|
||||
|
||||
fake.Flux2Transformer2DModel.from_single_file = classmethod(_capture)
|
||||
|
||||
backend = get_diffusion_backend()
|
||||
backend.load_model(
|
||||
"unsloth/FLUX.2-klein-4B-GGUF",
|
||||
gguf_filename = "flux-2-klein-4b-Q4_K_S.gguf",
|
||||
hf_token = "hf_test_token",
|
||||
)
|
||||
assert captured.get("config") == "black-forest-labs/FLUX.2-klein-4B"
|
||||
assert captured.get("subfolder") == "transformer"
|
||||
assert captured.get("token") == "hf_test_token"
|
||||
|
||||
|
||||
def test_release_chat_backend_calls_unload_with_model_name(monkeypatch):
|
||||
"""The safetensors backend unload helper must call unload_model
|
||||
with the active model name (the orchestrator's signature requires
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue