diff --git a/scripts/video_speedmem_bench.py b/scripts/video_speedmem_bench.py index 477c179580..d2fcb1ae2d 100644 --- a/scripts/video_speedmem_bench.py +++ b/scripts/video_speedmem_bench.py @@ -402,6 +402,7 @@ def _apply_levers( normalize_cache_quality, FBCACHE_MIN_STEPS, ) + from core.inference.video import _step_cache_all_or_none tgt = _target() engaged = { @@ -488,11 +489,14 @@ def _apply_levers( cache_request = cfg["cache"] if cache_request is not None: # Quality preset like the loader: an unset request takes the family's auto default. - # Expert names zip with the views so a dual-expert MoE resolves per-expert curves. quality = normalize_cache_quality(cache_quality) or auto_cache_quality(fam_name) - experts = ("transformer", "transformer_2") - for v, expert in zip(views, experts): - engaged["cache"] = apply_step_cache( + + # All-or-none across MoE experts, exactly like the loader: overwriting engaged["cache"] + # per expert would leave one expert cached and one dense on a partial engage while the + # row reports the cache off -- a config that never runs in production. The shared helper + # rolls the engaged expert(s) back so the row measures a real configuration. + def _engage_cache(v: Any, expert: str) -> Optional[str]: + return apply_step_cache( v, mode = cache_request, threshold = cache_threshold, @@ -503,6 +507,12 @@ def _apply_levers( expert = expert, logger = logger, ) + + engaged["cache"], cache_partial_reason = _step_cache_all_or_none( + pipe, fam_obj, _engage_cache, logger = logger + ) + if cache_partial_reason and logger is not None: + logger.warning("benchmark cache disabled: %s", cache_partial_reason) cache_active = engaged["cache"] not in (None, "off") # HunyuanVideo-1.5 joint-attention trim (per expert), BEFORE the backend set like the loader. diff --git a/studio/backend/core/inference/diffusion_attention.py b/studio/backend/core/inference/diffusion_attention.py index 4b5ef8d357..90baaf328c 100644 --- a/studio/backend/core/inference/diffusion_attention.py +++ b/studio/backend/core/inference/diffusion_attention.py @@ -146,6 +146,32 @@ def _cudnn_attention_supported() -> bool: return have is None or have >= (8, 0) +def attention_backend_supported_on_device(backend: Optional[str], device_index: int) -> bool: + """Whether an already-resolved dispatcher backend can actually RUN on CUDA ``device_index``. + + ``select_attention_backend`` arch-gates a backend against the ACTIVE device, but a CFG-parallel + replica lives on a possibly HETEROGENEOUS second GPU (FA3 is Hopper-SM90 only, FA4 needs + Blackwell-SM100, cuDNN needs Ampere+). Installing the primary-resolved backend there without + re-checking would set fine then crash on the replica's first attention kernel. Re-applies the + same arch gate to a specific device index. None (native) is always fine; an unqueryable + capability returns True (best-effort, matching ``_backend_arch_supported``).""" + if backend is None: + return True + try: + import torch + have = tuple(torch.cuda.get_device_capability(device_index)) # type: ignore[assignment] + except Exception: # noqa: BLE001 -- unqueryable device: don't block on a guess + return True + bounds = _ARCH_CAPABILITY.get(backend) + if bounds is not None: + low, high = bounds + if not (have >= low and (high is None or have < high)): + return False + if backend == "_native_cudnn" and have < (8, 0): + return False + return True + + # Optional-kernel backends installable on demand: dispatcher name -> (probe module, pip # package). Wheels only (--only-binary=:all:): a source build needs a CUDA toolchain a Studio # host may lack; no wheel means a native fallback. cuDNN/native ship with torch. @@ -398,6 +424,16 @@ _NULL_ATTN_FLAG = "_unsloth_null_attn_mask" _NULL_PROCESSOR_CACHE: dict = {} +def _set_hunyuan_null_mask(module: Any, enabled: bool) -> None: + """Set the null-mask flag on every block's attention of ``module``. The flag is valid ONLY for + the forward whose pre-hook removed the padding, so a post-hook clears it back to False after + each call (see the module note and _hunyuan_trim_post_hook).""" + for blk in getattr(module, "transformer_blocks", []): + attn = getattr(blk, "attn", None) + if attn is not None: + setattr(attn, _NULL_ATTN_FLAG, enabled) + + def _null_mask_processor_cls(): """Build (once, lazily) a HunyuanVideo15AttnProcessor2_0 subclass whose ``__call__`` runs attn_mask=None when the DiT is flagged (padding already removed by the pre-hook); otherwise it @@ -572,24 +608,26 @@ def _hunyuan_trim_pre_hook(module, args, kwargs): kwargs.update(original) null_ok = False - for blk in getattr(module, "transformer_blocks", []): - attn = getattr(blk, "attn", None) - if attn is not None: - setattr(attn, _NULL_ATTN_FLAG, null_ok) - + _set_hunyuan_null_mask(module, null_ok) return args, kwargs except Exception: # noqa: BLE001 — optimisation only; never break the forward # We may have trimmed some kwargs before failing. Restore the caller's untrimmed inputs so # the stock dense-mask path (flag False) runs on exactly what it expects. kwargs.clear() kwargs.update(original) - for blk in getattr(module, "transformer_blocks", []): - attn = getattr(blk, "attn", None) - if attn is not None: - setattr(attn, _NULL_ATTN_FLAG, False) + _set_hunyuan_null_mask(module, False) return args, kwargs +def _hunyuan_trim_post_hook(module, _args, output): + """Clear the null-mask flag after each hooked forward, scoping the authorisation to exactly the + call whose pre-hook removed the padding. Registered with ``always_call=True`` so the flag is + also cleared when the forward raises -- otherwise a latched True would null the mask over + un-trimmed padding on any later direct ``module.forward(...)``. Returns the output unchanged.""" + _set_hunyuan_null_mask(module, False) + return output + + def _install_null_processors(dit: Any, logger: Any) -> bool: """Swap every stock block attention processor on ``dit`` for the null-mask subclass. Only touches blocks whose processor is exactly the stock class (so a diffusers change or an @@ -643,11 +681,25 @@ def install_hunyuan_attention_trim( continue if not _install_null_processors(dit, logger): continue + # Installation (and every idle period between generations) starts in the conservative + # state: the flag is only ever True inside the exact forward its pre-hook trimmed. + _set_hunyuan_null_mask(dit, False) if getattr(dit, "_unsloth_trim_hook", None) is None: + pre_handle = None try: - handle = dit.register_forward_pre_hook(_hunyuan_trim_pre_hook, with_kwargs = True) - dit._unsloth_trim_hook = handle + pre_handle = dit.register_forward_pre_hook( + _hunyuan_trim_pre_hook, with_kwargs = True + ) + # always_call: clear the flag even when the forward raises, so an exception can + # never leave the null-mask authorisation latched for a later direct forward. + post_handle = dit.register_forward_hook( + _hunyuan_trim_post_hook, always_call = True + ) + dit._unsloth_trim_hook = (pre_handle, post_handle) except Exception as exc: # noqa: BLE001 — optimisation only + if pre_handle is not None: + pre_handle.remove() + _set_hunyuan_null_mask(dit, False) _warn(logger, "hunyuan_attn_trim", exc) continue engaged = True diff --git a/studio/backend/core/inference/diffusion_cfg_parallel.py b/studio/backend/core/inference/diffusion_cfg_parallel.py index 7ba139e210..79729492d5 100644 --- a/studio/backend/core/inference/diffusion_cfg_parallel.py +++ b/studio/backend/core/inference/diffusion_cfg_parallel.py @@ -268,6 +268,12 @@ class CFGParallelProxy: ) -> dict: """Resolve routing + dispatch for the next generation (call AFTER the cache toggle so the engaged state is current). Returns the plan for logging.""" + # Prompt/conditioning constants are reusable only WITHIN one generation (the same tensor + # objects flow through every denoise step); a new generation brings new ids. Clearing here + # releases the previous generation's replica-side copies (text embeds ~200 MiB each, held on + # BOTH GPUs) up front instead of pinning them until the 16-entry churn cap or teardown, and + # also cleans up after a cancelled/failed run that never reached note_generation_done(). + self._const_cache.clear() # The engaged-cache marker may live on the proxy (post-install toggle) or the # primary (pre-install engage); the delegating getattr covers both. marker = getattr(self, "_unsloth_step_cache", None) @@ -445,18 +451,26 @@ def _device_identity(idx: int) -> Optional[tuple]: return None -def _pick_secondary_device(primary_index: int) -> tuple[Optional[int], int, bool]: +def _pick_secondary_device( + primary_index: int, *, min_free_bytes: int = 0 +) -> tuple[Optional[int], int, bool]: """(secondary CUDA device != primary, its free bytes, identity-match flag). Bit-identity needs the SAME kernels on both branches, and eager kernel selection is - arch-dependent, so the picker prefers the most-free device whose (name, capability) - MATCH the primary's; only if none matches does it fall back to the most-free mismatched - one (so explicit ``on`` can still engage, lossy). An unqueryable identity counts as a - match (best-effort).""" + arch-dependent, so the picker prefers a device whose (name, capability) MATCH the primary's; + only if none matches does it fall back to a mismatched one (so explicit ``on`` can still + engage, lossy). An unqueryable identity counts as a match (best-effort). + + But a MATCHING device that cannot actually hold the replica is useless: rank a device that + fits ``min_free_bytes`` above one that does not FIRST, so a viable heterogeneous GPU wins over + an identical GPU too small for the replica (explicit ``on`` would otherwise fail with an + insufficient-memory gate while a usable device sat idle). Among devices in the same viability + tier, prefer the identity match, then the most free.""" import torch primary_id = _device_identity(primary_index) best, best_free, best_match = None, -1, False + best_key: tuple = (False, False, -1) for idx in range(torch.cuda.device_count()): if idx == primary_index: continue @@ -466,8 +480,9 @@ def _pick_secondary_device(primary_index: int) -> tuple[Optional[int], int, bool continue candidate_id = _device_identity(idx) match = primary_id is None or candidate_id is None or candidate_id == primary_id - if (match, free) > (best_match, best_free): - best, best_free, best_match = idx, free, match + key = (free >= min_free_bytes, match, free) + if best is None or key > best_key: + best, best_free, best_match, best_key = idx, free, match, key return best, best_free, best_match @@ -485,6 +500,7 @@ def maybe_enable_cfg_parallel( compiled: bool, attention_backend: Optional[str], speed_active: bool, + speed_mode: Optional[str] = None, logger: Any = None, ) -> tuple[Optional[CFGParallelProxy], str]: """Gate, build and install the CFG-parallel proxy on ``pipe``. Returns @@ -539,8 +555,12 @@ def maybe_enable_cfg_parallel( return None, f"primary DiT is on {p_dev.type}, not cuda" weight_bytes = sum(p.numel() * p.element_size() for p in primary.parameters()) primary_index = p_dev.index or 0 - secondary, free, device_match = _pick_secondary_device(primary_index) need = weight_bytes + _REPLICA_HEADROOM_BYTES + # Filter by the replica's memory need FIRST so a viable heterogeneous GPU is preferred + # over an identical GPU too small to hold the replica. + secondary, free, device_match = _pick_secondary_device( + primary_index, min_free_bytes = need + ) if secondary is None: return None, "no queryable secondary CUDA device" if not device_match: @@ -591,19 +611,53 @@ def maybe_enable_cfg_parallel( _warn(logger, "cfg-parallel replica load", exc) return None, "replica load failed" try: - from .diffusion_attention import apply_attention_backend, install_hunyuan_attention_trim + from .diffusion_attention import ( + apply_attention_backend, + attention_backend_supported_on_device, + install_hunyuan_attention_trim, + ) view = _ReplicaView(pipe, replica) if speed_active: install_hunyuan_attention_trim(view, fam, logger = logger) if attention_backend is not None: - apply_attention_backend(view, attention_backend, logger = logger) + # The backend was arch-gated against the PRIMARY device; re-validate it on the + # replica's (possibly heterogeneous) GPU before installing it -- FA3 is SM90-only, + # FA4 needs SM100, cuDNN needs Ampere+, so a mismatched replica would set fine then + # crash on its first attention kernel. Unsupported -> pin native on the replica. + replica_backend = attention_backend + if not attention_backend_supported_on_device(replica_backend, secondary): + if logger is not None: + logger.warning( + "diffusion.cfg_parallel: attention backend %r resolved for the primary " + "is unsupported on the replica cuda:%d (different arch); pinning native " + "there", + replica_backend, + secondary, + ) + replica_backend = None + apply_attention_backend(view, replica_backend, logger = logger) if compiled: - from .diffusion_speed import _compile_repeated_blocks + from .diffusion_speed import SPEED_MAX, _compile_repeated_blocks - # Same tier the primary got; a cache may engage/toggle on this DiT, so - # fullgraph stays off like the loader. - _compile_repeated_blocks(view, logger, cache_active = True) + # Mirror the primary's tier: under speed=max the primary compiles max-autotune + + # fuses QKV, so the replica must too or it becomes the slower branch and throttles + # the whole parallel run. A cache may engage/toggle on this DiT, so fullgraph stays + # off like the loader. + max_speed = str(speed_mode) == SPEED_MAX + _compile_repeated_blocks( + view, logger, max_autotune = max_speed, cache_active = True + ) + if max_speed: + # Fuse the REPLICA's QKV projections directly: _fuse_qkv(view) would resolve the + # pipe-level fuse_qkv_projections through _ReplicaView delegation and re-fuse the + # PRIMARY instead, leaving the replica unfused. + fuse = getattr(replica, "fuse_qkv_projections", None) + if callable(fuse): + try: + fuse() + except Exception as exc: # noqa: BLE001 -- optimisation only + _warn(logger, "cfg-parallel replica fuse_qkv", exc) if not _install_threadsafe_cudnn_attention(logger): raise RuntimeError("thread-safe attention patch failed") # The proxy's class name hides the transformer's from the metadata probe, so diff --git a/studio/backend/core/inference/diffusion_precision.py b/studio/backend/core/inference/diffusion_precision.py index 752eff7af3..58158ee9fc 100644 --- a/studio/backend/core/inference/diffusion_precision.py +++ b/studio/backend/core/inference/diffusion_precision.py @@ -335,7 +335,17 @@ def quantize_text_encoders( except Exception as exc: # noqa: BLE001 — leave this encoder dense # A mid-pass caster failure may have left the encoder PARTIALLY quantized (can't # run as dense), so fail the load for that; a clean miss stays best-effort dense. + # raise_if_partially_quantized only recognises torchao parameter subclasses, so it + # cannot see a partial layerwise fp8 mutation (diffusers apply_layerwise_casting installs + # upcast hooks + fp8 storage in place, leaving no torchao params). Detect a leftover + # layerwise hook directly and fail closed there too; a clean failure stays dense. from .diffusion_transformer_quant import raise_if_partially_quantized + if mode == TE_QUANT_FP8 and _has_layerwise_casting(encoder): + raise RuntimeError( + f"text_encoder_quant fp8:{attr} failed after partially installing layerwise " + "casting (leftover fp8 hooks); reload the model instead of a dense fallback " + f"(original error: {exc})" + ) from exc raise_if_partially_quantized(encoder, what = f"text_encoder_quant {mode}:{attr}", exc = exc) _warn(logger, f"{mode}:{attr}", exc) return mode if cast else None @@ -490,6 +500,29 @@ def _cast_nvfp4(encoder: Any, target: Any) -> None: quantize_(encoder, NVFP4WeightOnlyConfig(), filter_fn = filter_fn) +def _has_layerwise_casting(module: Any) -> bool: + """True when any submodule still carries a diffusers layerwise-casting hook -- i.e. an + ``apply_layerwise_casting`` pass installed an fp8-storage upcast hook before failing. torchao's + partial-quant detector cannot see these, so a mid-pass layerwise failure would otherwise report + a dense fallback over a half-cast encoder. Best-effort: a module without ``.modules()`` or a + moved diffusers internal returns False (defer to the torchao check).""" + try: + hook_name = "layerwise_casting" + try: + from diffusers.hooks.layerwise_casting import _LAYERWISE_CASTING_HOOK + hook_name = _LAYERWISE_CASTING_HOOK + except Exception: # noqa: BLE001 -- const moved: fall back to the stable literal + pass + for sub in module.modules(): + registry = getattr(sub, "_diffusers_hook", None) + get_hook = getattr(registry, "get_hook", None) + if callable(get_hook) and get_hook(hook_name) is not None: + return True + except Exception: # noqa: BLE001 -- unqueryable module: defer to the torchao check + return False + return False + + def _warn(logger: Any, what: str, exc: Exception) -> None: if logger is not None: logger.warning("diffusion.precision: text-encoder quant (%s) failed: %s", what, exc) diff --git a/studio/backend/core/inference/diffusion_vae_quant.py b/studio/backend/core/inference/diffusion_vae_quant.py index 878c6c289f..951ba53dc7 100644 --- a/studio/backend/core/inference/diffusion_vae_quant.py +++ b/studio/backend/core/inference/diffusion_vae_quant.py @@ -310,10 +310,19 @@ def quantize_vae( return mode except Exception as exc: # noqa: BLE001 — leave the VAE dense # fp8_dynamic's quantize_ swaps weights module-by-module, so a mid-pass failure may - # leave the VAE PARTIALLY quantized -- fail the load for that instead of a dense - # fallback. A clean miss (e.g. layerwise fp8) stays best-effort dense. + # leave the VAE PARTIALLY quantized -- fail the load for that instead of a dense fallback. + # raise_if_partially_quantized only recognises torchao parameter subclasses, so it CANNOT + # see a partial layerwise fp8 mutation (diffusers apply_layerwise_casting installs upcast + # hooks + fp8 storage in place, leaving no torchao params). Detect a leftover layerwise + # hook directly and fail closed there too; a clean failure (no hook installed) still falls + # back to dense (best-effort), matching the fp8 storage-only contract. from .diffusion_transformer_quant import raise_if_partially_quantized + if mode == VAE_QUANT_FP8 and _has_layerwise_casting(vae): + raise RuntimeError( + "vae_quant fp8 failed after partially installing layerwise casting (leftover " + f"fp8 hooks); reload the model instead of a dense fallback (original error: {exc})" + ) from exc raise_if_partially_quantized(vae, what = f"vae_quant {mode}", exc = exc) _warn(logger, mode, exc) return None @@ -372,6 +381,30 @@ def _cast_vae_fp8(vae: Any, target: Any) -> None: ) +def _has_layerwise_casting(module: Any) -> bool: + """True when any submodule still carries a diffusers layerwise-casting hook -- i.e. an + ``apply_layerwise_casting`` pass mutated the module (installed an fp8-storage upcast hook) + before failing. torchao's partial-quant detector cannot see these, so a mid-pass layerwise + failure would otherwise report a dense fallback over a half-cast module. Best-effort: a module + without ``.modules()`` or a moved diffusers internal returns False (defer to the torchao check). + """ + try: + hook_name = "layerwise_casting" + try: + from diffusers.hooks.layerwise_casting import _LAYERWISE_CASTING_HOOK + hook_name = _LAYERWISE_CASTING_HOOK + except Exception: # noqa: BLE001 -- const moved: fall back to the stable literal + pass + for sub in module.modules(): + registry = getattr(sub, "_diffusers_hook", None) + get_hook = getattr(registry, "get_hook", None) + if callable(get_hook) and get_hook(hook_name) is not None: + return True + except Exception: # noqa: BLE001 -- unqueryable module: defer to the torchao check + return False + return False + + def _warn(logger: Any, what: str, exc: Exception) -> None: if logger is not None: logger.warning("diffusion.vae_quant: (%s) failed: %s", what, exc) diff --git a/studio/backend/core/inference/video.py b/studio/backend/core/inference/video.py index e84c26cf88..b521eabc17 100644 --- a/studio/backend/core/inference/video.py +++ b/studio/backend/core/inference/video.py @@ -457,11 +457,22 @@ def _step_cache_all_or_none( engaged = [(view, name, mode) for view, name, mode in results if mode is not None] if engaged and len(engaged) < len(results): missing = ", ".join(name for _, name, mode in results if mode is None) + # Roll the engaged expert(s) back; a rollback that ALSO fails leaves one expert cached + # while state/status would report the whole pipeline uncached -- a silent inconsistency, + # so surface it as a hard reload-required error instead of returning a false "uncached". + rollback_failed: list[str] = [] for view, name, _mode in engaged: - _disengage_step_cache( + if not _disengage_step_cache( getattr(view, "transformer", None), reason = f"all-or-none rollback: cache did not engage on {missing}", logger = logger, + ): + rollback_failed.append(name) + if rollback_failed: + raise RuntimeError( + "step cache engagement was partial and rollback failed for " + + ", ".join(rollback_failed) + + "; reload the video model before generating" ) return None, ( f"step cache engaged on only {len(engaged)}/{len(results)} experts " @@ -1534,6 +1545,7 @@ class VideoBackend: compiled = "compiled" in speed_optims, attention_backend = attention_engaged, speed_active = effective_speed != SPEED_OFF, + speed_mode = effective_speed, logger = logger, ) if cfg_parallel_proxy is not None: @@ -2174,20 +2186,27 @@ class VideoBackend: # CONFIGURED step count: a clip at a different step count re-engages # (marker carries "#s{steps}") to keep skips aligned. This only # re-sizes the already-engaged cache; the on choice is preserved. - for view, expert_name in zip( - _views_for(pipe, fam), _transformer_names(pipe, fam) - ): + # Transactional across MoE experts (like the load / AUTO paths): refuse to + # stack a fresh cache over one whose removal failed, and roll back a mixed + # resize so status never reports MagCache over an asymmetric pair. + def _resize_explicit_magcache( + view: Any, expert_name: str + ) -> Optional[str]: transformer = getattr(view, "transformer", None) marker = getattr(transformer, "_unsloth_step_cache", None) # endswith, not substring: "#s5" would match inside "#s50". if not marker or str(marker).endswith(f"#s{int(steps)}"): - continue - _disengage_step_cache( + return TC_MAGCACHE # already sized for these steps + if not _disengage_step_cache( transformer, reason = f"explicit magcache re-interpolating for {steps} steps", logger = logger, - ) - apply_step_cache( + ): + raise RuntimeError( + "could not disable the existing MagCache before resizing it " + f"for {steps} steps; reload the video model before generating" + ) + return apply_step_cache( view, mode = TC_MAGCACHE, threshold = state.cache_threshold, @@ -2198,6 +2217,19 @@ class VideoBackend: expert = expert_name, logger = logger, ) + + resized, resize_reason = _step_cache_all_or_none( + pipe, fam, _resize_explicit_magcache, logger = logger + ) + object.__setattr__(state, "transformer_cache", resized) + entry = (state.resolved or {}).get("transformer_cache") + if isinstance(entry, dict): + entry["value"] = resized or "off" + entry["reason"] = resize_reason or ( + f"explicit MagCache resized for {steps} steps" + if resized + else f"MagCache could not be resized for {steps} steps" + ) if state.transformer_cache: self._reset_step_cache(pipe) # Dual-GPU CFG parallelism: resolve this generation's routing AFTER the diff --git a/studio/backend/tests/test_diffusion_attention.py b/studio/backend/tests/test_diffusion_attention.py index a67688d842..57d575a11d 100644 --- a/studio/backend/tests/test_diffusion_attention.py +++ b/studio/backend/tests/test_diffusion_attention.py @@ -514,3 +514,51 @@ def test_kernels_hub_compatible_reads_hub_version(monkeypatch): # Undeterminable hub -> keep the previous (permissive) behaviour. monkeypatch.setattr(importlib.metadata, "version", _boom) assert att._kernels_hub_compatible() is True + + +# ── per-device backend guard (CFG-parallel heterogeneous replica) ───────────────── +def _stub_cuda_capability(monkeypatch, caps): + """Stub torch.cuda.get_device_capability(idx) from a {idx: (major, minor)} map.""" + torch = types.ModuleType("torch") + torch.cuda = types.SimpleNamespace( + get_device_capability = lambda idx: caps[idx], + ) + monkeypatch.setitem(__import__("sys").modules, "torch", torch) + + +def test_backend_supported_on_device_none_is_always_ok(monkeypatch): + # None = native: nothing to arch-gate, so any device is fine (even unqueryable). + assert att.attention_backend_supported_on_device(None, 0) is True + + +def test_backend_supported_on_device_flash3_hopper_only(monkeypatch): + # FA3 is SM90 (Hopper) only: supported on the Hopper primary, NOT on a Blackwell replica. + _stub_cuda_capability(monkeypatch, {0: (9, 0), 1: (10, 0)}) + assert att.attention_backend_supported_on_device("_flash_3_hub", 0) is True + assert att.attention_backend_supported_on_device("_flash_3_hub", 1) is False + + +def test_backend_supported_on_device_flash4_blackwell_only(monkeypatch): + # FA4 needs SM100 (Blackwell): rejected on a Hopper replica. + _stub_cuda_capability(monkeypatch, {0: (10, 0), 1: (9, 0)}) + assert att.attention_backend_supported_on_device("flash_4_hub", 0) is True + assert att.attention_backend_supported_on_device("flash_4_hub", 1) is False + + +def test_backend_supported_on_device_cudnn_needs_ampere(monkeypatch): + # cuDNN fused SDPA needs Ampere+ (SM80): rejected on a pre-Ampere (T4/SM75) replica. + _stub_cuda_capability(monkeypatch, {0: (9, 0), 1: (7, 5)}) + assert att.attention_backend_supported_on_device("_native_cudnn", 0) is True + assert att.attention_backend_supported_on_device("_native_cudnn", 1) is False + + +def test_backend_supported_on_device_unqueryable_is_permissive(monkeypatch): + # An unqueryable device must not block on a guess (best-effort, like _backend_arch_supported). + torch = types.ModuleType("torch") + + def _boom(_idx): + raise RuntimeError("no device props") + + torch.cuda = types.SimpleNamespace(get_device_capability = _boom) + monkeypatch.setitem(__import__("sys").modules, "torch", torch) + assert att.attention_backend_supported_on_device("_flash_3_hub", 3) is True diff --git a/studio/backend/tests/test_diffusion_attention_trim.py b/studio/backend/tests/test_diffusion_attention_trim.py index 355834ab64..0f816ae11f 100644 --- a/studio/backend/tests/test_diffusion_attention_trim.py +++ b/studio/backend/tests/test_diffusion_attention_trim.py @@ -204,3 +204,49 @@ def test_install_trim_noop_when_transformer_class_mismatch(): fam = types.SimpleNamespace(transformer_class = "HunyuanVideo15Transformer3DModel") pipe = types.SimpleNamespace(transformer = types.SimpleNamespace()) # class name mismatch assert att.install_hunyuan_attention_trim(pipe, fam) is False + + +# ── null-mask flag lifecycle (scoped to one hooked forward) ─────────────────────── +def test_set_and_post_hook_clear_null_mask_flag(): + # _set_hunyuan_null_mask flips every block's flag; the post-hook clears it and returns + # the output unchanged. + dit = _fake_dit() + att._set_hunyuan_null_mask(dit, True) + assert all(getattr(b.attn, att._NULL_ATTN_FLAG) is True for b in dit.transformer_blocks) + sentinel = object() + returned = att._hunyuan_trim_post_hook(dit, (), sentinel) + assert returned is sentinel + assert all(getattr(b.attn, att._NULL_ATTN_FLAG) is False for b in dit.transformer_blocks) + + +def test_post_hook_always_clears_flag_after_forward_and_on_exception(): + # Wire the pre+post hooks the way install_hunyuan_attention_trim does on a real module: the + # flag is only ever True DURING the forward its pre-hook set up. After the call it is False, + # so a later direct dit.forward(...) can never run unmasked over untrimmed padding -- and the + # always_call post-hook clears it even when the forward raises (no latch across exceptions). + class _DiT(torch.nn.Module): + def __init__(self): + super().__init__() + self.transformer_blocks = [ + types.SimpleNamespace(attn = types.SimpleNamespace()) for _ in range(2) + ] + self.boom = False + + def forward(self): + # The processor would read a True flag here (padding removed by the pre-hook). + assert all(getattr(b.attn, att._NULL_ATTN_FLAG) for b in self.transformer_blocks) + if self.boom: + raise RuntimeError("mid-forward boom") + return "ok" + + dit = _DiT() + dit.register_forward_pre_hook(lambda m, _a: att._set_hunyuan_null_mask(m, True)) + dit.register_forward_hook(att._hunyuan_trim_post_hook, always_call = True) + + assert dit() == "ok" + assert all(getattr(b.attn, att._NULL_ATTN_FLAG) is False for b in dit.transformer_blocks) + + dit.boom = True + with pytest.raises(RuntimeError): + dit() + assert all(getattr(b.attn, att._NULL_ATTN_FLAG) is False for b in dit.transformer_blocks) diff --git a/studio/backend/tests/test_diffusion_cfg_parallel.py b/studio/backend/tests/test_diffusion_cfg_parallel.py index 2df417db1c..6dd623387b 100644 --- a/studio/backend/tests/test_diffusion_cfg_parallel.py +++ b/studio/backend/tests/test_diffusion_cfg_parallel.py @@ -594,3 +594,142 @@ def test_teardown_restores_pipe_and_guider(monkeypatch): def test_teardown_tolerates_foreign_object(): teardown_cfg_parallel(types.SimpleNamespace(transformer = None), object()) + + +# ── secondary picker: viability before identity ─────────────────────────────────── +def test_pick_secondary_prefers_viable_over_unusable_match(monkeypatch): + # A matching GPU too small for the replica must NOT beat a viable heterogeneous GPU: the + # min_free_bytes filter ranks first, so explicit "on" still engages the usable device + # instead of failing the memory gate while a usable card sits idle. + _stub_torch( + monkeypatch, + device_count = 3, + free = {1: (10 << 30, 80 << 30), 2: (60 << 30, 80 << 30)}, + names = {0: "NVIDIA B200", 1: "NVIDIA B200", 2: "NVIDIA H100"}, + ) + idx, free, match = _pick_secondary_device(0, min_free_bytes = 40 << 30) + assert idx == 2 and free == 60 << 30 and match is False + + +def test_pick_secondary_still_prefers_match_when_both_viable(monkeypatch): + # When BOTH fit the replica, the identity match still wins (bit-identity beats headroom). + _stub_torch( + monkeypatch, + device_count = 3, + free = {1: (50 << 30, 80 << 30), 2: (60 << 30, 80 << 30)}, + names = {0: "NVIDIA B200", 1: "NVIDIA B200", 2: "NVIDIA H100"}, + ) + idx, _free, match = _pick_secondary_device(0, min_free_bytes = 40 << 30) + assert idx == 1 and match is True + + +# ── replica lever mirroring (F3 attention arch guard, F4 max tier) ──────────────── +class _LoadableFuseDiT(_FakeDiT): + """A DiT whose class can build a replica and record a direct QKV fuse.""" + + @classmethod + def from_pretrained(cls, *a, **k): + return cls(device_index = 1) + + def to(self, *a, **k): + return self + + def eval(self): + return self + + def fuse_qkv_projections(self): + self.fused = getattr(self, "fused", 0) + 1 + + +def _engage_stubs(monkeypatch): + import core.inference.diffusion_cache as cache_mod + import core.inference.diffusion_cfg_parallel as cp + import core.inference.diffusion_speed as speed + + monkeypatch.setattr(cp, "_install_threadsafe_cudnn_attention", lambda logger = None: True) + monkeypatch.setattr(cp, "_restore_threadsafe_cudnn_attention", lambda: None) + monkeypatch.setattr(cache_mod, "_ensure_block_metadata_registered", lambda *a, **k: None) + return speed + + +def test_replica_mirrors_max_tier_compile_and_fusion(monkeypatch): + # Under speed_mode="max" the replica must compile max-autotune AND fuse QKV directly, or it + # becomes the slower branch and throttles the whole parallel run. + _stub_torch(monkeypatch) + speed = _engage_stubs(monkeypatch) + compile_kwargs: list = [] + monkeypatch.setattr( + speed, "_compile_repeated_blocks", + lambda view, logger, **kw: compile_kwargs.append(kw) or True, + ) + pipe = _CtxPipe(_LoadableFuseDiT()) + proxy, reason = _gate( + monkeypatch, pipe, _fam(), requested = "on", compiled = True, + speed_mode = "max", attention_backend = None, + ) + assert proxy is not None, reason + try: + assert compile_kwargs and compile_kwargs[0].get("max_autotune") is True + assert getattr(proxy._replica, "fused", 0) == 1 + finally: + teardown_cfg_parallel(pipe, proxy) + + +def test_replica_default_tier_no_max_autotune_no_fuse(monkeypatch): + # speed_mode="default": the replica compiles dynamic (max_autotune False) and is NOT fused, + # mirroring the primary's default tier. + _stub_torch(monkeypatch) + speed = _engage_stubs(monkeypatch) + compile_kwargs: list = [] + monkeypatch.setattr( + speed, "_compile_repeated_blocks", + lambda view, logger, **kw: compile_kwargs.append(kw) or True, + ) + pipe = _CtxPipe(_LoadableFuseDiT()) + proxy, reason = _gate( + monkeypatch, pipe, _fam(), requested = "on", compiled = True, + speed_mode = "default", attention_backend = None, + ) + assert proxy is not None, reason + try: + assert compile_kwargs and compile_kwargs[0].get("max_autotune") is False + assert getattr(proxy._replica, "fused", 0) == 0 + finally: + teardown_cfg_parallel(pipe, proxy) + + +def test_replica_pins_native_when_backend_unsupported_on_secondary(monkeypatch): + # The primary-resolved attention backend is arch-gated against the PRIMARY; on a heterogeneous + # replica it must be re-validated and, when unsupported, downgraded to native there rather than + # installed to crash on the replica's first attention kernel. + import core.inference.diffusion_attention as attn + + _stub_torch(monkeypatch) + _engage_stubs(monkeypatch) + monkeypatch.setattr(attn, "attention_backend_supported_on_device", lambda backend, idx: False) + applied: list = [] + monkeypatch.setattr( + attn, "apply_attention_backend", + lambda view, backend, logger = None: applied.append(backend) or backend, + ) + pipe = _CtxPipe(_LoadableFuseDiT()) + proxy, reason = _gate( + monkeypatch, pipe, _fam(), requested = "on", attention_backend = "_flash_3_hub", + ) + assert proxy is not None, reason + try: + assert applied == [None] # native pinned on the replica, not the unsupported FA3 + finally: + teardown_cfg_parallel(pipe, proxy) + + +def test_const_cache_cleared_each_plan_generation(monkeypatch): + # Prompt/conditioning constants are reusable only within one generation; plan_generation must + # release the previous generation's replica-side copies up front (no cross-generation VRAM pin). + proxy, _primary, _replica, guider = _make_proxy(monkeypatch) + guider.num_conditions = 2 + proxy._const_cache[123] = ("v", "moved") + assert proxy._const_cache + proxy.plan_generation(cache_engaged = True, steps = 20, width = 512, height = 512, frames = 17) + assert proxy._const_cache == {} + proxy.shutdown() diff --git a/studio/backend/tests/test_diffusion_precision.py b/studio/backend/tests/test_diffusion_precision.py index 29f9bbc24a..601c279d27 100644 --- a/studio/backend/tests/test_diffusion_precision.py +++ b/studio/backend/tests/test_diffusion_precision.py @@ -749,3 +749,38 @@ def test_quantize_partial_cast_failure_fails_load(monkeypatch): pipe = types.SimpleNamespace(text_encoder = _PartiallyCastEncoder()) with pytest.raises(RuntimeError, match = "partially quantized"): quantize_text_encoders(pipe, _target(), mode = "fp8") + + +# ── layerwise fp8 partial mutation on the text encoder (F7, mirrors the VAE path) ── +class _LayerwiseCastEncoder: + """A text encoder an apply_layerwise_casting pass mutated (installed an fp8-storage upcast hook) + before raising. No torchao params, so the torchao detector is blind to the partial state.""" + + def __init__(self): + registry = types.SimpleNamespace( + get_hook = lambda name: object() if name == "layerwise_casting" else None + ) + self._sub = types.SimpleNamespace(_diffusers_hook = registry) + + def modules(self): + return [self, self._sub] + + def named_parameters(self): + return iter(()) + + +def test_quantize_te_layerwise_partial_cast_fails_load(monkeypatch): + _stub_torch(monkeypatch) + hooks = types.ModuleType("diffusers.hooks") + casting = types.ModuleType("diffusers.hooks.layerwise_casting") + casting.DEFAULT_SKIP_MODULES_PATTERN = ("norm",) + + def _boom(module, **kwargs): + raise RuntimeError("encoder layerwise cast failed mid-pass") + + hooks.apply_layerwise_casting = _boom + monkeypatch.setitem(sys.modules, "diffusers.hooks", hooks) + monkeypatch.setitem(sys.modules, "diffusers.hooks.layerwise_casting", casting) + pipe = types.SimpleNamespace(text_encoder = _LayerwiseCastEncoder()) + with pytest.raises(RuntimeError, match = "leftover fp8 hooks"): + quantize_text_encoders(pipe, _target(), mode = "fp8") diff --git a/studio/backend/tests/test_diffusion_vae_quant.py b/studio/backend/tests/test_diffusion_vae_quant.py index 2408d2e347..ae458c9813 100644 --- a/studio/backend/tests/test_diffusion_vae_quant.py +++ b/studio/backend/tests/test_diffusion_vae_quant.py @@ -572,3 +572,56 @@ def test_quantize_vae_partial_cast_failure_fails_load(monkeypatch): pipe = types.SimpleNamespace(vae = _PartiallyQuantizedVae()) with pytest.raises(RuntimeError, match = "partially quantized"): quantize_vae(pipe, _target(), mode = "fp8") + + +# ── layerwise fp8 partial mutation (torchao detector is blind to diffusers hooks) ── +class _LayerwiseCastVae: + """A VAE an apply_layerwise_casting pass mutated (installed an fp8-storage upcast hook on a + submodule) before the caster raised. It carries NO torchao params, so raise_if_partially_ + quantized would miss the partial state -- _has_layerwise_casting must catch it.""" + + def __init__(self): + registry = types.SimpleNamespace( + get_hook = lambda name: object() if name == "layerwise_casting" else None + ) + self._sub = types.SimpleNamespace(_diffusers_hook = registry) + + def modules(self): + return [self, self._sub] + + def named_parameters(self): + return iter(()) + + +def test_quantize_vae_layerwise_partial_cast_fails_load(monkeypatch): + _stub_torch(monkeypatch, cc = (10, 0)) + _allow_vae(monkeypatch, {VAE_QUANT_FP8}) + + def _boom(v, t): + raise RuntimeError("layerwise casting failed mid-pass") + + monkeypatch.setattr(vq, "_cast_vae_fp8", _boom) + pipe = types.SimpleNamespace(vae = _LayerwiseCastVae()) + with pytest.raises(RuntimeError, match = "leftover fp8 hooks"): + quantize_vae(pipe, _target(), mode = "fp8") + + +def test_quantize_vae_clean_layerwise_failure_stays_dense(monkeypatch): + # A failure with NO leftover hook (raised before mutating anything) still falls back to dense, + # preserving the storage-only fp8 contract -- the fail-closed path is scoped to real mutation. + _stub_torch(monkeypatch, cc = (10, 0)) + _allow_vae(monkeypatch, {VAE_QUANT_FP8}) + + class _CleanVae: + def modules(self): + return [self] + + def named_parameters(self): + return iter(()) + + def _boom(v, t): + raise RuntimeError("fp8 unsupported before any mutation") + + monkeypatch.setattr(vq, "_cast_vae_fp8", _boom) + pipe = types.SimpleNamespace(vae = _CleanVae()) + assert quantize_vae(pipe, _target(), mode = "fp8") is None diff --git a/studio/backend/tests/test_video_backend.py b/studio/backend/tests/test_video_backend.py index b8a83385e7..7ba097fc10 100644 --- a/studio/backend/tests/test_video_backend.py +++ b/studio/backend/tests/test_video_backend.py @@ -2138,3 +2138,46 @@ def test_step_cache_all_or_none_single_dit(monkeypatch): assert video._step_cache_all_or_none(pipe, fam, engage, logger = None) == ("magcache", None) assert calls == [(pipe, "transformer")] + + +def test_step_cache_all_or_none_raises_when_rollback_fails(monkeypatch): + # Partial engagement AND a failed rollback of the engaged expert leaves it cached while state + # would report the pipeline uncached -- a silent inconsistency, so raise a hard reload-required + # error instead of a false "uncached". + import core.inference.video as video + + pipe, fam, _t1, _t2 = _moe_pipe_and_fam() + monkeypatch.setattr( + video, "_disengage_step_cache", + lambda transformer, *, reason, logger = None: False, # rollback fails + ) + with pytest.raises(RuntimeError, match = "rollback failed"): + video._step_cache_all_or_none( + pipe, fam, + lambda view, expert_name: "fbcache" if expert_name == "transformer" else None, + logger = None, + ) + + +def test_explicit_magcache_hard_errors_when_disable_fails(fake_runtime, monkeypatch): + # An explicit MagCache resize must be transactional: if the existing cache cannot be disabled, + # refuse to stack a fresh cache over it (which would double-hook) and hard-error instead of + # silently re-applying while status still reports MagCache. + import core.inference.video as video + + backend = VideoBackend() + backend.load_pipeline( + "Wan-AI/Wan2.2-TI2V-5B-Diffusers", + model_kind = "pipeline", + transformer_cache = "magcache", + ) + reapplied: list = [] + monkeypatch.setattr(video, "_disengage_step_cache", lambda *a, **k: False) + monkeypatch.setattr( + video, "apply_step_cache", + lambda *a, **k: reapplied.append(k.get("steps")) or "magcache", + ) + with pytest.raises(RuntimeError, match = "reload the video model"): + backend.generate(prompt = "a sloth", steps = 30) + assert reapplied == [] # never stacked a new cache over the un-removable one + backend.unload() diff --git a/studio/frontend/src/features/video/api.ts b/studio/frontend/src/features/video/api.ts index 3466e40d00..a676068278 100644 --- a/studio/frontend/src/features/video/api.ts +++ b/studio/frontend/src/features/video/api.ts @@ -106,16 +106,30 @@ export interface VideoLoadRequest { | "sage" | "xformers" | "aiter"; - transformer_cache?: "off" | "fbcache" | "magcache"; + transformer_cache?: "off" | "auto" | "fbcache" | "magcache"; transformer_cache_threshold?: number; - // Step-cache speed/accuracy preset (omit for the backend default, "balanced"). - transformer_cache_quality?: "quality" | "balanced" | "fast"; + // Step-cache speed/accuracy preset (omit/"auto" for the family's measured default). + transformer_cache_quality?: "auto" | "quality" | "balanced" | "fast"; // Dual-GPU CFG branch parallelism (omit for auto: engages on measured families when a // second GPU with enough free VRAM is available; bit-identical with the step cache on). cfg_parallel?: "off" | "auto" | "on"; - // Dense DiT precision on full-pipeline loads (omit for the hardware-ladder auto; - // "none" pins plain bf16). GGUF / single-file checkpoints carry their own precision. - transformer_quant?: "none" | "fp8" | "int8" | "nvfp4" | "mxfp8"; + // Dense DiT precision on full-pipeline loads (omit/"auto" for the hardware-ladder auto; + // "none"/"off" pins plain bf16). GGUF / single-file checkpoints carry their own precision. + transformer_quant?: "auto" | "none" | "off" | "fp8" | "int8" | "nvfp4" | "mxfp8"; + // Companion text-encoder precision (Gemma3 / UMT5 / Qwen2.5-VL), loaded bf16 from the base + // repo regardless of how the DiT was sourced. Omit/"auto" for the measured scheme; + // "none"/"off" keeps it dense. + text_encoder_quant?: + | "auto" + | "none" + | "off" + | "fp8" + | "fp8_dynamic" + | "int8" + | "nvfp4"; + // VAE (video decoder) precision. Omit/"auto" engages layerwise fp8 where the family + // qualifies; fp8_dynamic is an explicit opt-in (never auto); "none"/"off" keeps it dense. + vae_quant?: "auto" | "none" | "off" | "fp8" | "fp8_dynamic"; } export interface VideoGenerateRequest {