diffusion: address review round (FBCache context guard, aiter/ROCm, video cleanup, prequant + ControlNet gating)
- diffusion_cache: do not engage FBCache when the selected pipeline opens no cache_context. A CacheMixin transformer is necessary but not sufficient -- Flux Kontext / img2img / inpaint / controlnet reuse the CacheMixin FluxTransformer2DModel yet their __call__ never opens a cache_context, so the First-Block-Cache hook raised 'No context is set' on the first forward, crashing every default FLUX.1-Kontext edit (28 steps, above the FBCache threshold). Detect it from the pipeline __call__ source, resolved off the instance so the per-expert proxy view delegates to the real pipe. - diffusion_attention: honor an explicit aiter backend on ROCm/AMD targets instead of dropping it via the NVIDIA-only guard (aiter is the AMD ROCm kernel; it only works there). - video: clear the CUDA cache on a failed load so a partially built pipeline's reserved VRAM does not OOM the next load (mirrors the image backend), and re-check cancellation after the export/mux so a clip cancelled during the blocking encode is discarded, not persisted. - diffusion_auto_policy / diffusion_prequant: validate a request-supplied prequant path override (present AND allowlisted) before budgeting the small prequant plan, so the loader does not skip the dense shards and then rebuild dense after evicting the resident pipeline. - diffusion_controlnet: family-gate a curated ControlNet addressed by its full repo id, not only its short catalog id, so a cross-family repo id 400s up front instead of downloading and loading through the wrong ControlNet class.
This commit is contained in:
parent
3b659ba075
commit
c00eb20958
12 changed files with 193 additions and 13 deletions
|
|
@ -129,6 +129,14 @@ def select_attention_backend(
|
|||
backend = _ALIASES[alias]
|
||||
if backend == "native":
|
||||
return None
|
||||
# AITER is the AMD ROCm kernel, not an NVIDIA one: honor it on a ROCm (AMD) CUDA
|
||||
# target and drop it everywhere else (diffusers' own set-time check rejects it off
|
||||
# ROCm anyway). Without this special-case the NVIDIA-only guard below would silently
|
||||
# drop the one explicit backend that only ever works on ROCm.
|
||||
if backend == "aiter":
|
||||
if getattr(target, "device", None) == "cuda" and not _is_cuda_nvidia(target):
|
||||
return backend
|
||||
return None
|
||||
# Every explicit kernel here (cuDNN / flash* / sage) is CUDA+NVIDIA-only; on
|
||||
# ROCm / MPS / CPU diffusers accepts the name at set time and the first
|
||||
# generation crashes, so drop to the native default up front.
|
||||
|
|
|
|||
|
|
@ -188,10 +188,15 @@ def resolve_dense_quant_candidate(
|
|||
return None
|
||||
prequant_available = False
|
||||
try:
|
||||
from .diffusion_prequant import resolve_prequant_source
|
||||
prequant_available = (
|
||||
resolve_prequant_source(fam, scheme, path_override = prequant_path) is not None
|
||||
)
|
||||
from .diffusion_prequant import local_prequant_path_ready, resolve_prequant_source
|
||||
src = resolve_prequant_source(fam, scheme, path_override = prequant_path)
|
||||
# A request-supplied local path override is only usable if the loader will accept it
|
||||
# (allowlisted AND present); otherwise load_prequantized_transformer refuses it and
|
||||
# rebuilds dense after the resident pipe is unloaded -- the evict-then-OOM this
|
||||
# small-plan prefetch exists to avoid. Hosted-repo sources keep the existing signal.
|
||||
if src is not None and src.kind == "path" and not local_prequant_path_ready(src.location):
|
||||
src = None
|
||||
prequant_available = src is not None
|
||||
except Exception: # noqa: BLE001 -- prequant probing must never sink the candidate
|
||||
prequant_available = False
|
||||
estimate = estimate_dense_quant(
|
||||
|
|
|
|||
|
|
@ -64,6 +64,27 @@ def normalize_transformer_cache(value: Optional[str]) -> Optional[str]:
|
|||
return normalized
|
||||
|
||||
|
||||
def _pipeline_opens_cache_context(pipe: Any) -> bool:
|
||||
"""Whether the pipeline enters ``transformer.cache_context(...)`` in its denoise loop.
|
||||
The First-Block-Cache hook requires it at run time, and a CacheMixin transformer alone
|
||||
does NOT guarantee it: Flux Kontext / img2img / inpaint / controlnet reuse the CacheMixin
|
||||
FluxTransformer2DModel but never open a cache_context. Read from the pipeline ``__call__``
|
||||
source, resolved off the instance so a per-expert proxy view (``_SecondDiTView``)
|
||||
delegates to the real pipe; if it cannot be read, report False so the cache stays off."""
|
||||
import inspect
|
||||
|
||||
call = getattr(pipe, "__call__", None)
|
||||
if call is None:
|
||||
return False
|
||||
try:
|
||||
src = inspect.getsource(call)
|
||||
except (OSError, TypeError):
|
||||
return False
|
||||
# Match the actual call `cache_context(` -- a bare mention in a comment/docstring lacks
|
||||
# the paren, so this does not false-positive on prose.
|
||||
return "cache_context(" in src
|
||||
|
||||
|
||||
def apply_step_cache(
|
||||
pipe: Any,
|
||||
*,
|
||||
|
|
@ -89,17 +110,23 @@ def apply_step_cache(
|
|||
if threshold is not None
|
||||
else (QUANT_FBCACHE_THRESHOLD if quant_active else DEFAULT_FBCACHE_THRESHOLD)
|
||||
)
|
||||
# Only engage via the transformer's native enable_cache (the diffusers CacheMixin path).
|
||||
# That mixin is present exactly when the pipeline wraps the transformer call in a
|
||||
# cache_context, which the First-Block-Cache hook requires at run time. The lower-level
|
||||
# apply_first_block_cache hook would install on a non-CacheMixin transformer too (e.g.
|
||||
# Z-Image), but its pipeline opens no cache_context, so the first generation would crash
|
||||
# inside the hook -- so a model without enable_cache runs uncached per the best-effort
|
||||
# contract instead of being reported as cached and then failing.
|
||||
# Engage only via the transformer's native enable_cache (the diffusers CacheMixin path):
|
||||
# the lower-level apply_first_block_cache hook would install on a non-CacheMixin
|
||||
# transformer too (e.g. Z-Image), whose pipeline opens no cache_context and would crash
|
||||
# the first generation -- so a model without enable_cache runs uncached per the
|
||||
# best-effort contract instead of being reported as cached and then failing.
|
||||
enable_cache = getattr(transformer, "enable_cache", None)
|
||||
if not callable(enable_cache):
|
||||
_warn(logger, mode, RuntimeError("transformer has no cache_context (not a CacheMixin)"))
|
||||
return None
|
||||
# A CacheMixin transformer is necessary but NOT sufficient: the First-Block-Cache hook
|
||||
# raises "No context is set" on the first forward unless the PIPELINE wraps its denoise
|
||||
# loop in transformer.cache_context(...). Flux Kontext / img2img / inpaint / controlnet
|
||||
# reuse the CacheMixin FluxTransformer2DModel yet their __call__ opens no cache_context,
|
||||
# so engaging FBCache there would crash every default generation -- run uncached instead.
|
||||
if not _pipeline_opens_cache_context(pipe):
|
||||
_warn(logger, mode, RuntimeError("pipeline __call__ opens no cache_context; running uncached"))
|
||||
return None
|
||||
try:
|
||||
try:
|
||||
from diffusers import FirstBlockCacheConfig
|
||||
|
|
|
|||
|
|
@ -171,6 +171,11 @@ def resolve_controlnet(spec_id: str, *, family: Optional[str] = None) -> Resolve
|
|||
with a clear error rather than being loaded through the wrong pipeline class later.
|
||||
"""
|
||||
entry = _catalog_by_id().get(spec_id)
|
||||
if entry is None:
|
||||
# A curated entry addressed by its full repo id (owner/name) rather than its catalog
|
||||
# id must still hit the family gate below, not slip through to the bare-repo fallback
|
||||
# and get downloaded + loaded through the wrong family's ControlNet class.
|
||||
entry = next((e for e in _CURATED if e.repo_id and e.repo_id == spec_id), None)
|
||||
if entry is not None:
|
||||
# A curated/local entry may declare the families it is built for. A client that
|
||||
# bypasses the UI filter (direct API call) could send an entry for another family;
|
||||
|
|
|
|||
|
|
@ -86,6 +86,19 @@ def _local_prequant_path_allowed(path: str) -> bool:
|
|||
return any(real == r or real.startswith(r + os.sep) for r in roots)
|
||||
|
||||
|
||||
def local_prequant_path_ready(path: str) -> bool:
|
||||
"""True only when a request-supplied local pre-quant path would actually load: it is
|
||||
inside an allowlisted root AND the checkpoint file is present. The auto-policy planner
|
||||
uses this before budgeting the small prequant plan, so it never skips the dense shards
|
||||
for a path ``load_prequantized_transformer`` will refuse -- which would otherwise evict
|
||||
the resident pipeline and then rebuild dense under an undersized plan (OOM)."""
|
||||
import os
|
||||
|
||||
if not _local_prequant_path_allowed(path):
|
||||
return False
|
||||
return os.path.isfile(os.path.expanduser(path))
|
||||
|
||||
|
||||
@dataclass(frozen = True)
|
||||
class PrequantSource:
|
||||
"""Where a pre-quantized transformer checkpoint lives. ``kind`` is "path" (a local
|
||||
|
|
|
|||
|
|
@ -671,6 +671,15 @@ class VideoBackend:
|
|||
if self._load_token != token:
|
||||
return
|
||||
logger.error("video.load_failed: %s", exc)
|
||||
# Free the debris of a failed construction (mirrors diffusion.py's _run_load):
|
||||
# no _VideoLoadState was committed, so no later unload releases the VRAM a
|
||||
# partially built pipeline (OOM in from_pretrained / quant / placement) left
|
||||
# reserved in the caching allocator -- which would OOM the next load. Guarded so
|
||||
# a sticky CUDA error cannot skip stamping the real error below.
|
||||
try:
|
||||
clear_gpu_cache()
|
||||
except Exception: # noqa: BLE001 -- cleanup is best-effort
|
||||
pass
|
||||
from utils.native_path_leases import redact_native_paths
|
||||
|
||||
with self._lock:
|
||||
|
|
@ -1644,6 +1653,11 @@ class VideoBackend:
|
|||
mp4_bytes = self._encode_mp4(
|
||||
video_frames, out_fps, audio_track, pipe if fam.has_audio else None
|
||||
)
|
||||
# A cancel that landed during the (blocking, uncancellable) export/mux must
|
||||
# still discard the clip: cancel_generate() already reported success for it,
|
||||
# so re-check here before it is returned and persisted to the gallery.
|
||||
if cancel.is_set():
|
||||
raise RuntimeError(VIDEO_CANCELLED_MSG)
|
||||
duration_s = len(video_frames) / float(out_fps) if out_fps else 0.0
|
||||
self._gen = {"active": False}
|
||||
return {
|
||||
|
|
|
|||
|
|
@ -92,6 +92,21 @@ def test_explicit_backend_dropped_off_nvidia_cuda(monkeypatch):
|
|||
assert select_attention_backend(_target(device = "mps"), alias, speed_active = True) is None
|
||||
|
||||
|
||||
def test_aiter_honored_on_rocm(monkeypatch):
|
||||
# AITER is the AMD ROCm kernel; on a ROCm CUDA target it must be honored, not dropped by
|
||||
# the NVIDIA-only guard -- it is the one explicit backend that only ever works on ROCm.
|
||||
monkeypatch.setattr(att, "_is_cuda_nvidia", lambda target: False) # hip build
|
||||
assert select_attention_backend(_target(), "aiter", speed_active = False) == "aiter"
|
||||
|
||||
|
||||
def test_aiter_dropped_off_rocm(monkeypatch):
|
||||
# aiter on NVIDIA CUDA (or MPS / CPU) is not usable, so it drops to the native default.
|
||||
monkeypatch.setattr(att, "_is_cuda_nvidia", lambda target: True) # NVIDIA
|
||||
assert select_attention_backend(_target(), "aiter", speed_active = False) is None
|
||||
monkeypatch.setattr(att, "_is_cuda_nvidia", lambda target: False)
|
||||
assert select_attention_backend(_target(device = "mps"), "aiter", speed_active = False) is None
|
||||
|
||||
|
||||
def test_explicit_native_returns_none():
|
||||
# native is the default -> nothing to set.
|
||||
assert select_attention_backend(_target(), "native", speed_active = True) is None
|
||||
|
|
|
|||
|
|
@ -169,7 +169,8 @@ def test_candidate_none_for_an_unlisted_family(monkeypatch):
|
|||
|
||||
|
||||
def test_candidate_uses_prequant_transient_when_available(monkeypatch):
|
||||
_patch_selector(monkeypatch, prequant = object())
|
||||
# A hosted-repo prequant source (kind="repo") is available without a local-path check.
|
||||
_patch_selector(monkeypatch, prequant = SimpleNamespace(kind = "repo", location = "org/int8"))
|
||||
est = resolve_dense_quant_candidate(fam = _fam("z-image"), target = object(), requested = "int8")
|
||||
assert est is not None and est.prequant is True
|
||||
assert est.transient_transformer_mib == est.steady_transformer_mib
|
||||
|
|
|
|||
|
|
@ -68,8 +68,32 @@ class _NonCacheMixinTransformer:
|
|||
the load runs uncached instead (e.g. Z-Image)."""
|
||||
|
||||
|
||||
class _CtxPipe:
|
||||
"""A pipeline whose denoise loop opens ``transformer.cache_context(...)`` (like FluxPipeline)
|
||||
-- the First-Block-Cache hook needs it, so FBCache may engage here."""
|
||||
|
||||
def __init__(self, transformer):
|
||||
self.transformer = transformer
|
||||
|
||||
def __call__(self, *args, **kwargs):
|
||||
with self.transformer.cache_context("cond"):
|
||||
return None
|
||||
|
||||
|
||||
class _NoCtxPipe:
|
||||
"""A pipeline that never enters a caching context (like FluxKontextPipeline / img2img /
|
||||
inpaint / controlnet, which reuse the CacheMixin FluxTransformer2DModel): FBCache must NOT
|
||||
engage or the hook raises "No context is set" on the first forward."""
|
||||
|
||||
def __init__(self, transformer):
|
||||
self.transformer = transformer
|
||||
|
||||
def __call__(self, *args, **kwargs):
|
||||
return None
|
||||
|
||||
|
||||
def _pipe(transformer):
|
||||
return types.SimpleNamespace(transformer = transformer)
|
||||
return _CtxPipe(transformer)
|
||||
|
||||
|
||||
def _stub_diffusers(monkeypatch, *, hook_recorder = None):
|
||||
|
|
@ -129,6 +153,17 @@ def test_non_cachemixin_runs_uncached(monkeypatch):
|
|||
assert rec == {} # the standalone hook was never called
|
||||
|
||||
|
||||
def test_pipeline_without_cache_context_runs_uncached(monkeypatch):
|
||||
# A CacheMixin transformer whose PIPELINE never opens a cache_context (Flux Kontext /
|
||||
# img2img / inpaint / controlnet reuse the CacheMixin FluxTransformer2DModel) must run
|
||||
# uncached -- otherwise the First-Block-Cache hook raises "No context is set" on the
|
||||
# first forward, crashing every default generation.
|
||||
_stub_diffusers(monkeypatch)
|
||||
t = _MixinTransformer()
|
||||
assert apply_step_cache(_NoCtxPipe(t), mode = "fbcache") is None
|
||||
assert t.enabled_with is None # enable_cache was never called
|
||||
|
||||
|
||||
def test_incompatible_model_runs_uncached(monkeypatch):
|
||||
# enable_cache raising (e.g. unrecognised block signature) must not fail the load.
|
||||
_stub_diffusers(monkeypatch)
|
||||
|
|
|
|||
|
|
@ -58,6 +58,16 @@ def test_resolve_controlnet_enforces_family_match():
|
|||
assert dc.resolve_controlnet("qwen-union").path
|
||||
|
||||
|
||||
def test_resolve_controlnet_repo_id_still_family_gated():
|
||||
# A curated ControlNet addressed by its full repo id (not its short catalog id) must still
|
||||
# hit the family gate, not slip through the bare-repo fallback and load through the wrong
|
||||
# family's ControlNet class.
|
||||
with pytest.raises(ValueError, match = "is for"):
|
||||
dc.resolve_controlnet("InstantX/Qwen-Image-ControlNet-Union", family = "flux.1")
|
||||
r = dc.resolve_controlnet("InstantX/Qwen-Image-ControlNet-Union", family = "qwen-image")
|
||||
assert r.path == "InstantX/Qwen-Image-ControlNet-Union" and not r.is_local
|
||||
|
||||
|
||||
def test_union_control_mode_maps_only_union_entries():
|
||||
# Union entries map a known control type to its integer mode; a union model always
|
||||
# needs a concrete mode, so an unmapped type (passthrough) defaults to 0. A non-union
|
||||
|
|
|
|||
|
|
@ -59,6 +59,22 @@ def test_resolve_nothing_configured_is_none():
|
|||
assert resolve_prequant_source(_fam(), "fp8", path_override = "") is None
|
||||
|
||||
|
||||
def test_local_prequant_path_ready(tmp_path, monkeypatch):
|
||||
# The auto-policy planner budgets the small prequant plan only when a request-supplied
|
||||
# path would actually load: present AND inside an allowlisted root. Missing or not
|
||||
# allowlisted -> not ready, else the loader refuses it and rebuilds dense after evict.
|
||||
import os
|
||||
|
||||
ckpt = tmp_path / "model.pt"
|
||||
ckpt.write_bytes(b"x")
|
||||
root = os.path.realpath(str(tmp_path))
|
||||
monkeypatch.setattr(pq, "_allowed_prequant_roots", lambda: [root])
|
||||
assert pq.local_prequant_path_ready(str(ckpt)) is True
|
||||
assert pq.local_prequant_path_ready(str(tmp_path / "missing.pt")) is False
|
||||
monkeypatch.setattr(pq, "_allowed_prequant_roots", lambda: [])
|
||||
assert pq.local_prequant_path_ready(str(ckpt)) is False
|
||||
|
||||
|
||||
# ── load_prequantized_transformer ────────────────────────────────────────────────
|
||||
class _FakeTransformer:
|
||||
calls: dict = {}
|
||||
|
|
|
|||
|
|
@ -151,6 +151,12 @@ class _FakeWanDiT:
|
|||
def set_attention_backend(self, backend) -> None:
|
||||
self.attention = backend
|
||||
|
||||
@contextlib.contextmanager
|
||||
def cache_context(self, name):
|
||||
# Real Wan / HV15 / LTX pipelines open a cache_context around the denoise loop; the
|
||||
# First-Block-Cache hook needs it, so the fake transformer provides it too.
|
||||
yield
|
||||
|
||||
|
||||
class _FakeWanVae:
|
||||
def __init__(self) -> None:
|
||||
|
|
@ -229,6 +235,8 @@ class _FakeWanPipeSingle(_FakeWanPipeBase):
|
|||
"num_frames": num_frames,
|
||||
**kwargs,
|
||||
}
|
||||
with self.transformer.cache_context("cond"): # real Wan pipeline wraps the denoise loop
|
||||
pass
|
||||
return self._finish(num_inference_steps, num_frames, callback_on_step_end)
|
||||
|
||||
|
||||
|
|
@ -264,6 +272,8 @@ class _FakeWanPipeMoE(_FakeWanPipeBase):
|
|||
"num_frames": num_frames,
|
||||
**kwargs,
|
||||
}
|
||||
with self.transformer.cache_context("cond"): # real Wan pipeline wraps the denoise loop
|
||||
pass
|
||||
return self._finish(num_inference_steps, num_frames, callback_on_step_end)
|
||||
|
||||
|
||||
|
|
@ -344,6 +354,8 @@ class _FakeHV15Pipe:
|
|||
"num_frames": num_frames,
|
||||
**kwargs,
|
||||
}
|
||||
with self.transformer.cache_context("cond"): # real HV15 pipeline wraps the denoise loop
|
||||
pass
|
||||
for _ in range(int(num_inference_steps or 1)):
|
||||
self.scheduler.step()
|
||||
frames = [[object() for _ in range(int(num_frames or 1))]]
|
||||
|
|
@ -1062,6 +1074,25 @@ def test_hv15_cancel_unwinds_scheduler_loop(fake_runtime):
|
|||
assert pipe.hooks_freed == 1
|
||||
|
||||
|
||||
def test_cancel_during_export_discards_clip(fake_runtime, monkeypatch):
|
||||
# A cancel that lands during the (blocking, uncancellable) export/mux must still discard
|
||||
# the clip: cancel_generate() already reported success for it, so generate() must raise
|
||||
# the cancelled sentinel rather than return the clip to be persisted to the gallery.
|
||||
backend = VideoBackend()
|
||||
backend.load_pipeline(
|
||||
"hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-480p_t2v",
|
||||
model_kind = "pipeline",
|
||||
)
|
||||
|
||||
def _encode_and_cancel(frames, fps, audio, pipe):
|
||||
backend.cancel_generate() # cancel arrives mid-mux, after the last denoise-step check
|
||||
return b"MP4"
|
||||
|
||||
monkeypatch.setattr(VideoBackend, "_encode_mp4", staticmethod(_encode_and_cancel))
|
||||
with pytest.raises(RuntimeError, match = VIDEO_CANCELLED_MSG):
|
||||
backend.generate(prompt = "a fox", steps = 4)
|
||||
|
||||
|
||||
def test_singleton():
|
||||
assert get_video_backend() is get_video_backend()
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue