Compare commits

...
Author SHA1 Message Date
SolitaryThinker 6828eacdfe [bugfix]: harden FastWan attention routing 2026-07-15 12:33:40 -07:00
SolitaryThinker 1e38fe47ed [bugfix]: align SDPA capability test with unrestricted heads 2026-07-15 11:44:07 -07:00
SolitaryThinker 2157309857 [bugfix]: don't hard-fail explicit backend pins on layers that never supported the backend
The capability validation added for #1254 raised whenever the resolved
backend was incompatible and any explicit selection was active, even when
the resolution had already substituted a fallback for the pinned backend.
That killed legitimate loads: matrixgame2's CLIP vision encoder has
head_size=80, so a pinned FLASH_ATTN is replaced by SDPA at the platform
level, and the validator then blamed the FLASH_ATTN pin for SDPA's
declared head-size limits, aborting the whole model.

Hard-fail only when the explicit selection was actually honored and that
very backend rejects the layer (a real user error); a resolution that
fell back to a different backend degrades to the one-per-cached-resolution
warning, matching auto-selection.
2026-07-15 11:23:50 -07:00
SolitaryThinkerandUtkarsh Ranjan 409a86d520 [feat]: validate resolved attention backend against declared capabilities (#1499)
Wires AttentionBackend.validate_compatibility into the selector's cached
resolution path. When the backend was explicitly selected (global force
or FASTVIDEO_ATTENTION_BACKEND, which a model's
required_attention_backend also requires) and its self-described
capabilities reject the layer's head_size/dtype, raise ValueError --
matching how the platform layer hard-fails on missing
explicitly-requested backends -- instead of failing later with an opaque
kernel error. When the backend was auto-selected, log a single warning
per cached resolution rather than blocking.

Redesigned from #1499, whose warn-and-proceed check predated the
hard-fail framework. The selector-cache test stubs now subclass
AttentionBackend so the capability hooks resolve to the permissive
base-class defaults.

Co-authored-by: Utkarsh Ranjan <utkarsh@socratic.co>
2026-07-15 11:23:50 -07:00
SolitaryThinkerandUtkarsh Ranjan 2c4cbdb0ee [feat]: add capability self-description API to AttentionBackend (#1489)
Adds additive capability hooks to the AttentionBackend base class so
backends declare what they support and the selector can validate a
request against a layer's needs with a clear reason on mismatch
(#1254): get_supported_head_sizes (formalizes an existing per-backend
hook), get_supported_dtypes, supports_attention_mask, supports_varlen,
and validate_compatibility.

Defaults are the least-restrictive behavior, so backends that do not
override them keep today's behavior. Concrete declarations are limited
to what is verifiable from the backend code: SDPA also runs in fp32;
VSA and VMoBA impose their own sparsity pattern and do not consume a
dense attention mask. Folded from #1489 with review fixes: the
is_available() hook is dropped (the platform layer already hard-fails
on missing explicitly-requested backends) and the supports_varlen=True
claims for VSA/VMoBA are dropped (neither exposes caller-facing
variable-length packing).

Co-authored-by: Utkarsh Ranjan <utkarsh@socratic.co>
2026-07-15 11:23:50 -07:00
SolitaryThinker c29c7163fc [bugfix]: fix attention backend routing and coverage (#1494) 2026-07-15 11:23:50 -07:00
SolitaryThinker d2751d27c4 [bugfix]: detect FullAttn snapshot paths 2026-07-15 11:23:50 -07:00
SolitaryThinker 6e4ef6a1e4 [misc]: apply pre-commit formatting 2026-07-15 11:23:50 -07:00
SolitaryThinker ab66ae9a89 [bugfix]: preserve attention backend config JSON roundtrip 2026-07-15 11:23:50 -07:00
SolitaryThinker c07e199cb9 [feat]: fail loudly on FastWan VSA requirement instead of auto-forcing
Replace the per-model auto-force of VSA with a loud check so the transformer
blocks and the denoising stage can no longer disagree on the backend.

- selector.py: add get_selected_attn_backend() as the single source of the
  "global force > FASTVIDEO_ATTENTION_BACKEND" precedence, and
  check_attn_backend_requirement(), which uses it and raises ValueError when a
  model's required backend is not the selected one. _cached_get_attn_backend
  reuses the same helper instead of re-deriving the precedence.
- wanvideo.py: WanTransformer3DModel uses the check instead of wrapping block
  construction in a transient global force. The force only reached block
  construction and left DenoisingStage resolving a different backend, so a
  FastWan load without the env var built VSA blocks but no VSA metadata and
  crashed at forward.
- DiTConfig: add required_attention_backend as a base field (None default)
  instead of narrowing dit_config to WanVideoConfig. Keeps the repo-wide
  `dit_config: DiTConfig` convention and makes the requirement model-agnostic,
  living next to _supported_attention_backends.
- wan.py/registry.py: add FastWan2_2_TI2V_5B_FullAttn_Config (no VSA
  requirement) and route FastWan2.2-TI2V-5B-FullAttn-Diffusers to it so the
  dense checkpoint is no longer forced onto VSA.
- tests: FullAttn config requires nothing; cover the new check's branches.
2026-07-15 11:23:50 -07:00
SolitaryThinker 85a477b80c [misc]: simplify attention-backend selection + fail loudly on unknown FASTVIDEO_ATTENTION_BACKEND
- wanvideo.py: collapse the redundant backend-resolution triple, drop the
  "VIDEO_SPARSE_ATTN" magic string, and generalize the force-attn context to
  any config-required backend (behavior unchanged for current VSA usage)
- cuda.py: DRY the four near-identical "backend not installed" ImportErrors
  into a _backend_not_installed_error() helper
- envs.py: validate FASTVIDEO_ATTENTION_BACKEND at the read chokepoint so a
  typo'd value raises with the list of valid backends instead of silently
  falling back to auto-selection
- add regression test for the env-var validation
2026-07-15 11:23:50 -07:00
SolitaryThinker 4fd21120de [feat]: hard-fail on missing attention backends + force VSA for FastWan
Problem
-------
When the user sets FASTVIDEO_ATTENTION_BACKEND to SAGE_ATTN, SAGE_ATTN_THREE,
or ATTN_QAT_INFER but the corresponding package is not installed, the dispatch
silently logs a warning and falls through to FlashAttn. The user believes they
are running the requested backend but they are not.

The FastWan model family is sparse-distilled with VSA and produces incorrect
outputs at inference without it. The README tells users to set
FASTVIDEO_ATTENTION_BACKEND=VIDEO_SPARSE_ATTN before from_pretrained — another
silent footgun.

Change
------
fastvideo/platforms/cuda.py:
  - SAGE_ATTN: warn+fallthrough -> hard ImportError with install hint.
  - SAGE_ATTN_THREE: warn+fallthrough -> hard ImportError with install hint.
  - ATTN_QAT_INFER: warn+fallthrough -> hard ImportError with install hint.
  - VIDEO_SPARSE_ATTN: unchanged (already raises).

fastvideo/configs/models/dits/wanvideo.py + fastvideo/configs/pipelines/wan.py +
fastvideo/models/dits/wanvideo.py:
  - FastWan VSA requirement is now carried on the model's dit_config (scoped
    per-model) instead of a process-global force, so base Wan models loaded in
    the same process are unaffected.
  - FastWan2_1_T2V_480P_Config + FastWan2_2_TI2V_5B_Config set
    required_attention_backend=VIDEO_SPARSE_ATTN on their WanVideoConfig.
  - WanTransformer3DModel chooses its VSA transformer block from the config
    requirement before falling back to the selector global or env var.
  - FastWan block construction scopes VIDEO_SPARSE_ATTN through the selector
    context manager so get_attn_backend resolves VSA and missing kernels still
    hard-raise, then restores the previous global force.
  - The selector cache is cleared when the global force changes so stale
    FlashAttn selections cannot survive scoped/backend force changes.

Test plan
---------
- pre-commit run --files fastvideo/configs/models/dits/wanvideo.py fastvideo/configs/pipelines/wan.py fastvideo/models/dits/wanvideo.py fastvideo/tests/attention/test_fastwan_vsa_scoping.py: clean.
- python -m py_compile fastvideo/configs/models/dits/wanvideo.py fastvideo/configs/pipelines/wan.py fastvideo/models/dits/wanvideo.py fastvideo/tests/attention/test_fastwan_vsa_scoping.py: clean.
- pytest fastvideo/tests/attention/test_fastwan_vsa_scoping.py -q: 2 passed.
- Manual smoke: simulated missing sageattention, sageattn3, fastvideo_kernel, and unavailable attn_qat_infer; each requested backend raised ImportError with the expected backend/package/install message.
2026-07-15 11:23:50 -07:00
23 changed files with 1071 additions and 62 deletions
+1 -1
View File
@@ -56,7 +56,7 @@ pipeline initialization and sampling.
| Model Name | HuggingFace Model ID | Resolutions | TeaCache | Sliding Tile Attn (Legacy Branch) | Sage Attn | VSA | BSA |
|------------|---------------------|-------------|----------|-------------------|-----------|-----|-----|
| FastWan2.1 T2V 1.3B | `FastVideo/FastWan2.1-T2V-1.3B-Diffusers` | 480P | ⭕ | ⭕ | ⭕ | ✅ | ⭕ |
| FastWan2.2 TI2V 5B Full Attn* | `FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers` | 720P | ⭕ | ⭕ | ⭕ | ✅ | ⭕ |
| FastWan2.2 TI2V 5B Full Attn* | `FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers` | 720P | ⭕ | ⭕ | ⭕ | ❌ | ⭕ |
| Wan2.2 TI2V 5B | `Wan-AI/Wan2.2-TI2V-5B-Diffusers` | 720P | ⭕ | ⭕ | ✅ | ⭕ | ⭕ |
| DreamX-World 5B Cam | `FastVideo/DreamX-World-5B-Cam-Diffusers` | 480P | ⭕ | ⭕ | ⭕ | ⭕ | ⭕ |
| DreamX-World 5B AR | `FastVideo/DreamX-World-5B-Diffusers` | 704px1280p | ⭕ | ⭕ | ⭕ | ⭕ | ⭕ |
@@ -127,11 +127,12 @@ def save_image_from_base64(image_data: str, output_dir: str) -> Optional[str]:
def setup_model_environment(model_path: str) -> None:
# if "fullattn" in model_path.lower():
# os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "FLASH_ATTN"
# else:
# os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "VIDEO_SPARSE_ATTN"
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "FLASH_ATTN"
from fastvideo.platforms import AttentionBackendEnum
from fastvideo.registry import get_pipeline_config_cls_from_name
pipeline_config = get_pipeline_config_cls_from_name(model_path)()
attention_backend = pipeline_config.dit_config.required_attention_backend or AttentionBackendEnum.FLASH_ATTN
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = attention_backend.name
os.environ["FASTVIDEO_STAGE_LOGGING"] = "1"
@@ -1,6 +1,8 @@
# FastWan 2.2 TI2V 5B FullAttn LoRA finetune with VSA.
# FastWan 2.2 TI2V 5B FullAttn LoRA finetune with dense attention.
#
# Notes:
# - The FullAttn checkpoint has no VSA gate-compression weights. Train its
# dense attention projections without enabling VIDEO_SPARSE_ATTN.
# - This model still uses T2V-style preprocessed parquet data in the new
# training stack.
# - Validation/inference should use WanPipeline rather than the I2V pipeline;
@@ -21,7 +23,6 @@ models:
- to_k
- to_v
- to_out
- to_gate_compress
method:
_target_: fastvideo.train.methods.fine_tuning.finetune.FineTuneMethod
@@ -59,21 +60,18 @@ training:
gradient_accumulation_steps: 5
checkpoint:
output_dir: outputs/fast_wan_ti2v_fullattn_lora_vsa
output_dir: outputs/fast_wan_ti2v_fullattn_lora
training_state_checkpointing_steps: 5
checkpoints_total_limit: 2
resume_from_checkpoint: null
tracker:
project_name: fastvideo
run_name: fast_wan_ti2v_fullattn_lora_vsa
run_name: fast_wan_ti2v_fullattn_lora
model:
enable_gradient_checkpointing_type: full
vsa:
sparsity: 0.8
callbacks:
grad_clip:
_target_: fastvideo.train.callbacks.grad_clip.GradNormClipCallback
+72
View File
@@ -47,6 +47,78 @@ class AttentionBackend(ABC):
def get_builder_cls() -> type["AttentionMetadataBuilder"]:
raise NotImplementedError
# ------------------------------------------------------------------
# Capability self-description
#
# These let the selector validate a requested backend against a
# layer's needs and emit a clear reason on mismatch, instead of
# failing later with an opaque kernel error (see #1254). They are
# additive: the defaults describe the least-restrictive behavior, so
# a backend that does not override them keeps today's behavior.
# Several backends already declare ``get_supported_head_sizes``; this
# formalizes that hook on the base class and adds the other
# capability axes. Only override with values verifiable from the
# backend's implementation.
# ------------------------------------------------------------------
@staticmethod
def get_supported_head_sizes() -> list[int] | None:
"""Attention head sizes this backend supports.
Return None for no restriction.
"""
return None
@classmethod
def get_supported_dtypes(cls) -> tuple[torch.dtype, ...]:
"""Floating-point dtypes this backend supports."""
return (torch.float16, torch.bfloat16)
@classmethod
def supports_attention_mask(cls) -> bool:
"""Whether this backend can consume an explicit attention mask/bias.
Dense backends generally can; tiled/sparse video backends generally
cannot, since they impose their own sparsity pattern.
"""
return True
@classmethod
def supports_varlen(cls) -> bool:
"""Whether this backend supports single-launch variable-length
sequence packing (``cu_seqlens``), as in ``flash_attn_varlen_func``.
"""
return False
@classmethod
def validate_compatibility(
cls,
head_size: int,
dtype: torch.dtype,
*,
needs_attention_mask: bool = False,
needs_varlen: bool = False,
) -> str | None:
"""Return None if this backend is compatible with the given
requirements, else a human-readable reason for the mismatch.
The selector surfaces this reason, so an unsupported request is
never silently dropped.
"""
supported_sizes = cls.get_supported_head_sizes()
if supported_sizes is not None and head_size not in supported_sizes:
return (f"{cls.get_name()} does not support head_size="
f"{head_size} (supported: {supported_sizes})")
if dtype not in cls.get_supported_dtypes():
return f"{cls.get_name()} does not support dtype={dtype}"
if needs_attention_mask and not cls.supports_attention_mask():
return (f"{cls.get_name()} cannot consume the attention mask this "
"layer requires")
if needs_varlen and not cls.supports_varlen():
return (f"{cls.get_name()} does not support variable-length "
"sequence packing (varlen)")
return None
@dataclass
class AttentionMetadata:
+5
View File
@@ -24,6 +24,11 @@ class SDPABackend(AttentionBackend):
# CLIP vision encoders). None means "no restriction".
return None
@classmethod
def get_supported_dtypes(cls) -> tuple[torch.dtype, ...]:
# torch.nn.functional.scaled_dot_product_attention also runs in fp32.
return (torch.float16, torch.bfloat16, torch.float32)
@staticmethod
def get_name() -> str:
return "TORCH_SDPA"
@@ -119,6 +119,12 @@ class VideoSparseAttentionBackend(AttentionBackend):
def get_supported_head_sizes() -> list[int]:
return [64, 128]
@classmethod
def supports_attention_mask(cls) -> bool:
# VSA imposes its own tiled sparsity pattern; it does not consume a
# general dense attention mask.
return False
@staticmethod
def get_name() -> str:
return "VIDEO_SPARSE_ATTN"
+6
View File
@@ -18,6 +18,12 @@ class VMOBAAttentionBackend(AttentionBackend):
accept_output_buffer: bool = True
@classmethod
def supports_attention_mask(cls) -> bool:
# VMoBA selects blocks via its own MoBA routing; it does not consume
# a general dense attention mask.
return False
@staticmethod
def get_name() -> str:
return "VMOBA_ATTN"
+85 -18
View File
@@ -69,6 +69,7 @@ def global_force_attn_backend(attn_backend: AttentionBackendEnum | None) -> None
'''
global forced_attn_backend
forced_attn_backend = attn_backend
_cached_get_attn_backend.cache_clear()
def get_global_forced_attn_backend() -> AttentionBackendEnum | None:
@@ -79,6 +80,63 @@ def get_global_forced_attn_backend() -> AttentionBackendEnum | None:
return forced_attn_backend
def get_selected_attn_backend() -> AttentionBackendEnum | None:
'''
The attention backend selected before any per-layer filtering.
A global force (see `global_force_attn_backend`) takes precedence over the
`FASTVIDEO_ATTENTION_BACKEND` env var; returns None when neither is set
(automatic selection). This is the single source of the force-over-env
precedence rule, shared by `_cached_get_attn_backend` and
`check_attn_backend_requirement`.
'''
forced_backend = get_global_forced_attn_backend()
if forced_backend is not None:
return forced_backend
backend_by_env_var = envs.FASTVIDEO_ATTENTION_BACKEND
return backend_name_to_enum(backend_by_env_var) if backend_by_env_var is not None else None
def check_attn_backend_requirement(
required_backend: AttentionBackendEnum | None,
*,
model_name: str = "This model",
) -> AttentionBackendEnum | None:
'''
Resolve the selected attention backend, enforcing a model's hard requirement.
Some model families are only numerically correct with a specific backend
(e.g. FastWan is sparse-distilled with VSA and produces wrong outputs
otherwise). Rather than silently forcing the backend -- which would only
reach some construction sites and leave others (e.g. the denoising stage)
resolving a different backend -- we require the user to select it explicitly
and fail loudly otherwise, mirroring how the platform layer hard-fails on
missing explicitly-requested backends.
Arguments:
* required_backend: backend the model requires, or None for no requirement
* model_name: name used in the error message
Returns:
* the selected attention backend (`get_selected_attn_backend`), or None if
unset and unrequired
Raises:
* ValueError if `required_backend` is set but is not the selected backend
'''
selected_backend = get_selected_attn_backend()
if required_backend is not None and selected_backend != required_backend:
selected_name = selected_backend.name if selected_backend is not None else "unset"
raise ValueError(f"{model_name} requires the {required_backend.name} attention backend, but the effective "
f"FASTVIDEO_ATTENTION_BACKEND is {selected_name}. This checkpoint is only correct with "
f"{required_backend.name}; set FASTVIDEO_ATTENTION_BACKEND={required_backend.name} before "
f"loading the model.")
return selected_backend
def get_attn_backend(
head_size: int,
dtype: torch.dtype,
@@ -86,33 +144,23 @@ def get_attn_backend(
| None = None,
default_backend: AttentionBackendEnum | None = None,
) -> type[AttentionBackend]:
return _cached_get_attn_backend(head_size, dtype, supported_attention_backends, default_backend)
# Resolve the global force / environment override before entering the
# cached function so backend changes participate in the cache key. This
# also validates the environment on every lookup, including cache hits.
selected_backend = get_selected_attn_backend()
return _cached_get_attn_backend(head_size, dtype, supported_attention_backends, selected_backend, default_backend)
@cache
def _cached_get_attn_backend(
head_size: int,
dtype: torch.dtype,
supported_attention_backends: tuple[AttentionBackendEnum, ...]
| None = None,
supported_attention_backends: tuple[AttentionBackendEnum, ...] | None,
selected_backend: AttentionBackendEnum | None,
default_backend: AttentionBackendEnum | None = None,
) -> type[AttentionBackend]:
# Check whether a particular choice of backend was
# previously forced.
#
# THIS SELECTION OVERRIDES THE FASTVIDEO_ATTENTION_BACKEND
# ENVIRONMENT VARIABLE.
if not supported_attention_backends:
raise ValueError("supported_attention_backends is empty")
selected_backend = None
backend_by_global_setting: AttentionBackendEnum | None = (get_global_forced_attn_backend())
if backend_by_global_setting is not None:
selected_backend = backend_by_global_setting
else:
# Check the environment variable and override if specified
backend_by_env_var: str | None = envs.FASTVIDEO_ATTENTION_BACKEND
if backend_by_env_var is not None:
selected_backend = backend_name_to_enum(backend_by_env_var)
# Layer-level default (e.g. a checkpoint that requires a specific sparse
# backend). Lower precedence than the global force and the env var, so
@@ -135,7 +183,26 @@ def _cached_get_attn_backend(
attention_cls = current_platform.get_attn_backend_cls(selected_backend, head_size, dtype)
if not attention_cls:
raise ValueError(f"Invalid attention backend for {current_platform.device_name}")
return cast(type[AttentionBackend], resolve_obj_by_qualname(attention_cls))
backend = cast(type[AttentionBackend], resolve_obj_by_qualname(attention_cls))
# Validate the resolved backend against its self-described capabilities
# (see AttentionBackend.validate_compatibility, #1254). An explicit
# selection hard-fails only when it was actually honored: resolution can
# substitute a fallback for the selected backend (e.g. a FLASH_ATTN pin on
# a CLIP layer whose head size only SDPA serves -- see
# CudaPlatformBase.get_attn_backend_cls -- or a pin outside the layer's
# supported set above). Such a layer never participated in the pin, so its
# fallback degrades to the auto-selection warning -- emitted once per
# resolution, since this function is cached.
incompatibility = backend.validate_compatibility(head_size, dtype)
if incompatibility is not None:
if selected_backend is not None and backend.get_name() == selected_backend.name:
raise ValueError(f"Attention backend {selected_backend.name} was explicitly selected but is incompatible "
f"with this layer: {incompatibility}. Select a compatible backend or unset "
"FASTVIDEO_ATTENTION_BACKEND.")
logger.warning("Resolved attention backend %s may be incompatible with this layer: %s", backend.get_name(),
incompatibility)
return backend
@contextmanager
+52
View File
@@ -48,6 +48,58 @@ class DiTConfig(ModelConfig):
# FastVideoDiT-specific parameters
prefix: str = ""
quant_config: QuantizationConfig | None = None
# Some model families are only numerically correct with a specific attention
# backend (e.g. FastWan is sparse-distilled with VSA). When set, the model
# fails loudly if the selected FASTVIDEO_ATTENTION_BACKEND does not match it.
# None = no requirement (the common case).
required_attention_backend: AttentionBackendEnum | None = None
# A checkpoint can support the model's ordinary attention implementation
# while still being incompatible with a particular backend-specific block
# layout. Keep those checkpoint constraints separate from the architecture's
# general _supported_attention_backends list.
incompatible_attention_backends: tuple[AttentionBackendEnum, ...] = ()
@staticmethod
def _parse_attention_backend(backend: AttentionBackendEnum | str, field_name: str) -> AttentionBackendEnum:
if isinstance(backend, AttentionBackendEnum):
return backend
if isinstance(backend, str):
try:
return AttentionBackendEnum[backend]
except KeyError as exc:
valid_backends = ", ".join(AttentionBackendEnum.__members__)
raise ValueError(
f"Unknown attention backend {backend!r} in {field_name}; expected one of: {valid_backends}"
) from exc
raise TypeError(f"{field_name} must contain AttentionBackendEnum values or names, got {type(backend).__name__}")
def update_model_config(self, source_model_dict: dict[str, Any]) -> None:
source_model_dict = source_model_dict.copy()
required_backend = source_model_dict.get("required_attention_backend")
if isinstance(required_backend, str):
source_model_dict["required_attention_backend"] = self._parse_attention_backend(
required_backend, "required_attention_backend")
incompatible_backends = source_model_dict.get("incompatible_attention_backends")
if incompatible_backends is not None:
if not isinstance(incompatible_backends, list | tuple):
raise TypeError("incompatible_attention_backends must be a list or tuple")
source_model_dict["incompatible_attention_backends"] = tuple(
self._parse_attention_backend(backend, "incompatible_attention_backends")
for backend in incompatible_backends)
super().update_model_config(source_model_dict)
def validate_attention_backend_compatibility(
self,
selected_backend: AttentionBackendEnum | None,
*,
model_name: str = "This model",
) -> None:
if selected_backend in self.incompatible_attention_backends:
assert selected_backend is not None
raise ValueError(f"{model_name} is incompatible with the {selected_backend.name} attention backend. "
"Select a compatible attention backend before loading the model.")
@staticmethod
def add_cli_args(parser: Any, prefix: str = "dit-config") -> Any:
+6
View File
@@ -284,6 +284,12 @@ class PipelineConfig:
model_dict = asdict(value)
# Model Arch Config should be hidden away from the users
model_dict.pop("arch_config")
if isinstance(value, DiTConfig):
if value.required_attention_backend is not None:
model_dict["required_attention_backend"] = value.required_attention_backend.name
model_dict["incompatible_attention_backends"] = [
backend.name for backend in value.incompatible_attention_backends
]
output_dict[key] = model_dict
elif isinstance(value, tuple) and all(isinstance(v, ModelConfig) for v in value):
model_dicts = []
+29 -1
View File
@@ -12,6 +12,9 @@ from fastvideo.configs.models.encoders import (BaseEncoderOutput, CLIPVisionConf
from fastvideo.configs.models.vaes import WanVAEConfig
from fastvideo.configs.models.vaes.wanvae import WanVAEArchConfig
from fastvideo.configs.pipelines.base import PipelineConfig
from fastvideo.platforms import AttentionBackendEnum
FASTWAN_REQUIRED_ATTENTION_BACKEND = AttentionBackendEnum.VIDEO_SPARSE_ATTN
def t5_postprocess_text(outputs: BaseEncoderOutput) -> torch.Tensor:
@@ -64,7 +67,7 @@ class WanT2V480PConfig(PipelineConfig):
# WanConfig-specific added parameters
def __post_init__(self):
def __post_init__(self) -> None:
self.vae_config.load_encoder = False
self.vae_config.load_decoder = True
@@ -123,6 +126,10 @@ class FastWan2_1_T2V_480P_Config(WanT2V480PConfig):
flow_shift: float | None = 8.0
dmd_denoising_steps: list[int] | None = field(default_factory=lambda: [1000, 757, 522])
def __post_init__(self) -> None:
super().__post_init__()
self.dit_config.required_attention_backend = FASTWAN_REQUIRED_ATTENTION_BACKEND
@dataclass
class Wan2_2_TI2V_5B_Config(WanT2V480PConfig):
@@ -277,6 +284,27 @@ class FastWan2_2_TI2V_5B_Config(Wan2_2_TI2V_5B_Config):
flow_shift: float | None = 5.0
dmd_denoising_steps: list[int] | None = field(default_factory=lambda: [1000, 757, 522])
def __post_init__(self) -> None:
super().__post_init__()
self.dit_config.required_attention_backend = FASTWAN_REQUIRED_ATTENTION_BACKEND
@dataclass
class FastWan2_2_TI2V_5B_FullAttn_Config(Wan2_2_TI2V_5B_Config):
"""FullAttn (dense-attention) DMD variant of FastWan 2.2 TI2V 5B.
Shares FastWan's DMD denoising schedule but uses the published dense
checkpoint layout, which has no VSA gate-compression weights.
"""
flow_shift: float | None = 5.0
dmd_denoising_steps: list[int] | None = field(default_factory=lambda: [1000, 757, 522])
def __post_init__(self) -> None:
super().__post_init__()
self.dit_config.required_attention_backend = None
self.dit_config.incompatible_attention_backends = (FASTWAN_REQUIRED_ATTENTION_BACKEND, )
@dataclass
class Wan2_2_T2V_A14B_Config(WanT2V480PConfig):
+19 -1
View File
@@ -68,6 +68,24 @@ def maybe_convert_int(value: str | None) -> int | None:
return int(value)
def get_attention_backend() -> str | None:
"""Read FASTVIDEO_ATTENTION_BACKEND, failing loudly on an unknown value.
Returns None when unset. A typo'd / invalid backend name raises immediately
with the list of valid backends, rather than silently falling back to
auto-selection.
"""
backend = os.getenv("FASTVIDEO_ATTENTION_BACKEND", None)
if backend is None:
return None
# Imported lazily to avoid a circular import at module load time.
from fastvideo.platforms.interface import AttentionBackendEnum
if backend not in AttentionBackendEnum.__members__:
valid = ", ".join(AttentionBackendEnum.__members__)
raise ValueError(f"Invalid FASTVIDEO_ATTENTION_BACKEND={backend!r}. Valid backends are: {valid}")
return backend
# The begin-* and end* here are used by the documentation generator
# to extract the used env vars.
@@ -211,7 +229,7 @@ environment_variables: dict[str, Callable[[], Any]] = {
# FLASH_ATTN uses FlashAttention-3/2; to run FlashAttention-4 set
# FASTVIDEO_FA4=1 as well (see below).
"FASTVIDEO_ATTENTION_BACKEND":
lambda: os.getenv("FASTVIDEO_ATTENTION_BACKEND", None),
get_attention_backend,
# If set (=1), the FLASH_ATTN backend uses FlashAttention-4
# (flash_attn.cute). FA4 is opt-in and never auto-selected just because it
+17 -3
View File
@@ -7,9 +7,9 @@ from typing import Any
import torch
import torch.nn as nn
import fastvideo.envs as envs
from fastvideo.attention import (DistributedAttention, DistributedAttention_VSA,
LocalAttention)
from fastvideo.attention.selector import check_attn_backend_requirement
from fastvideo.configs.models.dits import WanVideoConfig
from fastvideo.distributed.communication_op import (
sequence_model_parallel_all_gather_with_unpad,
@@ -607,6 +607,21 @@ class WanTransformerBlock_VSA(nn.Module):
return hidden_states
def _select_wan_transformer_block(
config: WanVideoConfig,
*,
model_name: str,
) -> type[WanTransformerBlock] | type[WanTransformerBlock_VSA]:
# Resolve checkpoint-level requirements before choosing a block layout.
# Sparse FastWan checkpoints require VSA's extra gate-compression weights;
# dense checkpoints explicitly reject that layout.
attn_backend = check_attn_backend_requirement(config.required_attention_backend, model_name=model_name)
config.validate_attention_backend_compatibility(attn_backend, model_name=model_name)
if attn_backend == AttentionBackendEnum.VIDEO_SPARSE_ATTN:
return WanTransformerBlock_VSA
return WanTransformerBlock
class WanTransformer3DModel(BaseDiT):
_fsdp_shard_conditions = WanVideoConfig()._fsdp_shard_conditions
_compile_conditions = WanVideoConfig()._compile_conditions
@@ -619,6 +634,7 @@ class WanTransformer3DModel(BaseDiT):
def __init__(self, config: WanVideoConfig, hf_config: dict[str,
Any]) -> None:
super().__init__(config=config, hf_config=hf_config)
transformer_block = _select_wan_transformer_block(config, model_name=type(self).__name__)
self.quant_config = config.quant_config
inner_dim = config.num_attention_heads * config.attention_head_dim
@@ -651,8 +667,6 @@ class WanTransformer3DModel(BaseDiT):
)
# 3. Transformer blocks
attn_backend = envs.FASTVIDEO_ATTENTION_BACKEND
transformer_block = WanTransformerBlock_VSA if attn_backend == "VIDEO_SPARSE_ATTN" else WanTransformerBlock
self.blocks = nn.ModuleList([
transformer_block(inner_dim,
config.ffn_dim,
+24 -7
View File
@@ -29,6 +29,13 @@ pynvml = import_pynvml() # type: ignore[no-untyped-call]
torch.backends.cuda.enable_cudnn_sdp(False)
def _backend_not_installed_error(backend_label: str, env_value: str, what_missing: str,
remediation: str) -> ImportError:
"""Build the error raised when an explicitly-requested attention backend is unavailable."""
return ImportError(f"{backend_label} backend was explicitly requested via "
f"FASTVIDEO_ATTENTION_BACKEND={env_value} but {what_missing}. {remediation}")
def device_id_to_physical_device_id(device_id: int) -> int:
if "CUDA_VISIBLE_DEVICES" in os.environ:
device_ids = os.environ["CUDA_VISIBLE_DEVICES"].split(",")
@@ -126,8 +133,9 @@ class CudaPlatformBase(Platform):
return "fastvideo.attention.backends.sage_attn.SageAttentionBackend"
except ImportError as e:
logger.info(e)
logger.info("Sage Attention backend is not installed. Fall back to Flash Attention.")
raise _backend_not_installed_error("SageAttention", "SAGE_ATTN",
"the `sageattention` package is not installed",
"Install it with: uv pip install sageattention") from e
elif selected_backend == AttentionBackendEnum.SAGE_ATTN_THREE:
try:
from sageattn3 import sageattn3_blackwell # noqa: F401
@@ -138,15 +146,24 @@ class CudaPlatformBase(Platform):
return "fastvideo.attention.backends.sage_attn3.SageAttention3Backend"
except ImportError as e:
logger.info(e)
logger.info("Sage Attention 3 backend is not installed. Fall back to Flash Attention.")
raise _backend_not_installed_error("SageAttention 3", "SAGE_ATTN_THREE",
"the `sageattn3` package is not installed",
"Install it with: uv pip install sageattn3") from e
elif selected_backend == AttentionBackendEnum.ATTN_QAT_INFER:
from fastvideo.attention.backends.attn_qat_infer import ( # noqa: F401
AttnQatInferBackend, is_attn_qat_infer_available)
try:
from fastvideo.attention.backends.attn_qat_infer import ( # noqa: F401
AttnQatInferBackend, is_attn_qat_infer_available)
except ImportError as e:
raise _backend_not_installed_error(
"Attn-QAT inference", "ATTN_QAT_INFER", "the `attn_qat_infer` kernel package is not installed",
"Install the FastVideo Attn-QAT inference kernel or see: https://hao-ai-lab.github.io/FastVideo/"
) from e
if is_attn_qat_infer_available():
logger.info("Using Attn-QAT inference (modified SageAttention3 FP4) backend.")
return "fastvideo.attention.backends.attn_qat_infer.AttnQatInferBackend"
logger.info("Attn-QAT inference kernel is not built. Fall back to Flash Attention.")
raise _backend_not_installed_error(
"Attn-QAT inference", "ATTN_QAT_INFER", "the `attn_qat_infer` kernel package is not built",
"Install the FastVideo Attn-QAT inference kernel or see: https://hao-ai-lab.github.io/FastVideo/")
elif selected_backend == AttentionBackendEnum.ATTN_QAT_TRAIN:
from fastvideo.attention.backends.attn_qat_train import ( # noqa: F401
AttnQatTrainBackend, is_attn_qat_train_available)
+34 -17
View File
@@ -45,7 +45,7 @@ from fastvideo.configs.pipelines.turbodiffusion import (
)
from fastvideo.configs.pipelines.wan import (
FastWan2_1_T2V_480P_Config,
FastWan2_2_TI2V_5B_Config,
FastWan2_2_TI2V_5B_FullAttn_Config,
LucyEditDevConfig,
SelfForcingWan2_2_T2V480PConfig,
SelfForcingWanT2V480PConfig,
@@ -137,6 +137,9 @@ _MODEL_HF_PATH_TO_NAME: dict[str, str] = {}
# Detectors to identify model families from paths or class names
_MODEL_NAME_DETECTORS: list[tuple[str, Callable[[str], bool]]] = []
# Detectors to identify model families from model_index.json metadata
_MODEL_INDEX_DETECTORS: list[tuple[str, Callable[[dict[str, Any]], bool]]] = []
def register_configs(
sampling_param_cls: type[SamplingParam] | None,
@@ -144,6 +147,7 @@ def register_configs(
workload_types: tuple[WorkloadType, ...],
hf_model_paths: list[str] | None = None,
model_detectors: list[Callable[[str], bool]] | None = None,
model_index_detector: Callable[[dict[str, Any]], bool] | None = None,
model_family: str | None = None,
default_preset: str | None = None,
pipeline_cls_name: str | None = None,
@@ -175,6 +179,9 @@ def register_configs(
for detector in model_detectors:
_MODEL_NAME_DETECTORS.append((model_id, detector))
if model_index_detector:
_MODEL_INDEX_DETECTORS.append((model_id, model_index_detector))
def get_model_short_name(model_id: str) -> str:
if "/" in model_id:
@@ -209,14 +216,19 @@ def _get_config_info(
else:
config = maybe_download_model_index(model_path)
pipeline_name = config.get("_class_name", "").lower()
matched_model_names: list[str] = []
for model_id, detector in _MODEL_NAME_DETECTORS:
if detector(model_path.lower()) or detector(pipeline_name):
logger.debug("Matched model name '%s' using a registered detector.", model_id)
for model_id, index_detector in _MODEL_INDEX_DETECTORS:
if index_detector(config):
logger.debug("Matched model name '%s' using model index metadata.", model_id)
matched_model_names.append(model_id)
if not matched_model_names:
pipeline_name = config.get("_class_name", "").lower()
for model_id, name_detector in _MODEL_NAME_DETECTORS:
if name_detector(model_path.lower()) or name_detector(pipeline_name):
logger.debug("Matched model name '%s' using a registered detector.", model_id)
matched_model_names.append(model_id)
if matched_model_names:
if len(matched_model_names) > 1:
logger.warning(
@@ -924,6 +936,22 @@ def _register_configs() -> None:
model_family="wan",
default_preset="wan_fun_1_3b_control",
)
# Register the dense FullAttn variant before the generic WanDMDPipeline
# detector below. Its model index is the reliable discriminator for local
# checkpoints whose directory name does not preserve the Hugging Face slug.
register_configs(
sampling_param_cls=None,
pipeline_config_cls=FastWan2_2_TI2V_5B_FullAttn_Config,
workload_types=(WorkloadType.T2V, WorkloadType.I2V),
hf_model_paths=[
"FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers",
"FastVideo/FastWan2.2-TI2V-5B-Diffusers",
],
model_index_detector=lambda config: config.get("_class_name", "").lower() == "wandmdpipeline" and config.get(
"expand_timesteps") is True,
model_family="wan",
default_preset="fast_wan_2_2_ti2v_5b",
)
register_configs(
sampling_param_cls=None,
pipeline_config_cls=FastWan2_1_T2V_480P_Config,
@@ -980,17 +1008,6 @@ def _register_configs() -> None:
model_family="dreamx_world",
default_preset="dreamx_world_5b_ar",
)
register_configs(
sampling_param_cls=None,
pipeline_config_cls=FastWan2_2_TI2V_5B_Config,
workload_types=(WorkloadType.T2V, WorkloadType.I2V),
hf_model_paths=[
"FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers",
"FastVideo/FastWan2.2-TI2V-5B-Diffusers",
],
model_family="wan",
default_preset="fast_wan_2_2_ti2v_5b",
)
register_configs(
sampling_param_cls=None,
pipeline_config_cls=LucyEditDevConfig,
@@ -0,0 +1,27 @@
import pytest
import fastvideo.envs as envs
from fastvideo.platforms.interface import AttentionBackendEnum
def test_unknown_attention_backend_env_fails_loudly(monkeypatch):
# Given: a typo'd backend name in the env var.
monkeypatch.setenv("FASTVIDEO_ATTENTION_BACKEND", "FLASH_ATN")
# Then: reading the env var raises instead of silently auto-selecting,
# and the error names the offending value plus the valid backends.
with pytest.raises(ValueError) as excinfo:
_ = envs.FASTVIDEO_ATTENTION_BACKEND
message = str(excinfo.value)
assert "FLASH_ATN" in message
assert AttentionBackendEnum.FLASH_ATTN.name in message
def test_valid_attention_backend_env_passes_through(monkeypatch):
monkeypatch.setenv("FASTVIDEO_ATTENTION_BACKEND", AttentionBackendEnum.VIDEO_SPARSE_ATTN.name)
assert envs.FASTVIDEO_ATTENTION_BACKEND == AttentionBackendEnum.VIDEO_SPARSE_ATTN.name
def test_unset_attention_backend_env_is_none(monkeypatch):
monkeypatch.delenv("FASTVIDEO_ATTENTION_BACKEND", raising=False)
assert envs.FASTVIDEO_ATTENTION_BACKEND is None
@@ -0,0 +1,93 @@
from types import SimpleNamespace
import pytest
import torch
import fastvideo.platforms as platforms
from fastvideo.attention import selector
from fastvideo.attention.backends.abstract import AttentionBackend
from fastvideo.platforms.interface import AttentionBackendEnum
# Subclass AttentionBackend so the selector's capability validation
# (validate_compatibility) sees the permissive base-class defaults.
class FlashBackend(AttentionBackend):
pass
class VideoSparseBackend(AttentionBackend):
pass
@pytest.fixture(autouse=True)
def reset_attention_backend_selector():
selector.global_force_attn_backend(None)
yield
selector.global_force_attn_backend(None)
@pytest.fixture
def stub_backend_resolution(monkeypatch):
requested_backends = []
def get_attn_backend_cls(selected_backend, head_size, dtype):
requested_backends.append(selected_backend)
return selected_backend.name
platform = SimpleNamespace(
device_name="test device",
get_attn_backend_cls=get_attn_backend_cls,
)
backend_classes = {
AttentionBackendEnum.FLASH_ATTN.name: FlashBackend,
AttentionBackendEnum.VIDEO_SPARSE_ATTN.name: VideoSparseBackend,
}
monkeypatch.setattr(platforms, "_current_platform", platform)
monkeypatch.setattr(selector, "resolve_obj_by_qualname", backend_classes.__getitem__)
return requested_backends
def test_changing_env_backend_cannot_reuse_cached_backend(monkeypatch, stub_backend_resolution):
supported_backends = (
AttentionBackendEnum.FLASH_ATTN,
AttentionBackendEnum.VIDEO_SPARSE_ATTN,
)
monkeypatch.setenv("FASTVIDEO_ATTENTION_BACKEND", AttentionBackendEnum.FLASH_ATTN.name)
first_backend = selector.get_attn_backend(64, torch.float16, supported_backends)
monkeypatch.setenv("FASTVIDEO_ATTENTION_BACKEND", AttentionBackendEnum.VIDEO_SPARSE_ATTN.name)
second_backend = selector.get_attn_backend(64, torch.float16, supported_backends)
assert first_backend is FlashBackend
assert second_backend is VideoSparseBackend
assert stub_backend_resolution == [
AttentionBackendEnum.FLASH_ATTN,
AttentionBackendEnum.VIDEO_SPARSE_ATTN,
]
def test_invalid_env_backend_cannot_be_hidden_by_cache_hit(monkeypatch, stub_backend_resolution):
supported_backends = (AttentionBackendEnum.FLASH_ATTN, )
monkeypatch.setenv("FASTVIDEO_ATTENTION_BACKEND", AttentionBackendEnum.FLASH_ATTN.name)
assert selector.get_attn_backend(64, torch.float16, supported_backends) is FlashBackend
monkeypatch.setenv("FASTVIDEO_ATTENTION_BACKEND", "FLASH_ATN")
with pytest.raises(ValueError, match="FLASH_ATN"):
selector.get_attn_backend(64, torch.float16, supported_backends)
assert stub_backend_resolution == [AttentionBackendEnum.FLASH_ATTN]
def test_global_force_still_takes_precedence_over_env(monkeypatch, stub_backend_resolution):
supported_backends = (
AttentionBackendEnum.FLASH_ATTN,
AttentionBackendEnum.VIDEO_SPARSE_ATTN,
)
monkeypatch.setenv("FASTVIDEO_ATTENTION_BACKEND", AttentionBackendEnum.VIDEO_SPARSE_ATTN.name)
selector.global_force_attn_backend(AttentionBackendEnum.FLASH_ATTN)
backend = selector.get_attn_backend(64, torch.float16, supported_backends)
assert backend is FlashBackend
assert stub_backend_resolution == [AttentionBackendEnum.FLASH_ATTN]
@@ -0,0 +1,97 @@
# SPDX-License-Identifier: Apache-2.0
"""Unit tests for the AttentionBackend capability self-description API.
These exercise the additive capability hooks added for #1254 using a small
dummy backend, so they run on CPU without GPU kernels or optional
third-party attention packages installed.
"""
import torch
from fastvideo.attention.backends.abstract import (AttentionBackend, AttentionImpl, AttentionMetadata,
AttentionMetadataBuilder)
class _DummyBackend(AttentionBackend):
"""Minimal concrete backend that relies on the base-class defaults."""
@staticmethod
def get_name() -> str:
return "DUMMY"
@staticmethod
def get_impl_cls() -> type[AttentionImpl]:
raise NotImplementedError
@staticmethod
def get_metadata_cls() -> type[AttentionMetadata]:
raise NotImplementedError
@staticmethod
def get_builder_cls() -> type[AttentionMetadataBuilder]:
raise NotImplementedError
class _RestrictedBackend(_DummyBackend):
"""Overrides the capability hooks to a restrictive set."""
@staticmethod
def get_name() -> str:
return "RESTRICTED"
@staticmethod
def get_supported_head_sizes() -> list[int]:
return [64, 128]
@classmethod
def supports_attention_mask(cls) -> bool:
return False
def test_defaults_are_permissive():
assert _DummyBackend.get_supported_head_sizes() is None
assert _DummyBackend.supports_varlen() is False
assert _DummyBackend.supports_attention_mask() is True
assert torch.float16 in _DummyBackend.get_supported_dtypes()
assert torch.bfloat16 in _DummyBackend.get_supported_dtypes()
def test_compatible_request_returns_none():
# No head-size restriction, supported dtype, no mask required.
assert _DummyBackend.validate_compatibility(123, torch.float16) is None
def test_unsupported_dtype_reports_reason():
reason = _DummyBackend.validate_compatibility(64, torch.float32)
assert reason is not None
assert "dtype" in reason
assert "DUMMY" in reason
def test_unsupported_head_size_reports_reason():
reason = _RestrictedBackend.validate_compatibility(96, torch.float16)
assert reason is not None
assert "head_size" in reason
def test_supported_head_size_passes():
assert _RestrictedBackend.validate_compatibility(128, torch.float16) is None
def test_mask_requirement_reports_reason():
reason = _RestrictedBackend.validate_compatibility(64, torch.float16, needs_attention_mask=True)
assert reason is not None
assert "mask" in reason
def test_varlen_requirement_reports_reason():
# _DummyBackend uses the default supports_varlen() == False.
reason = _DummyBackend.validate_compatibility(64, torch.float16, needs_varlen=True)
assert reason is not None
assert "varlen" in reason
def test_real_backends_declare_capabilities():
# SDPA declares unrestricted head sizes and its broader dtype support.
from fastvideo.attention.backends.sdpa import SDPABackend
assert SDPABackend.get_supported_head_sizes() is None
assert torch.float32 in SDPABackend.get_supported_dtypes()
@@ -0,0 +1,97 @@
from __future__ import annotations
import builtins
import sys
from types import ModuleType
import pytest
import torch
from fastvideo.platforms.cuda import CudaPlatformBase
from fastvideo.platforms.interface import AttentionBackendEnum
class _GpuIndependentCudaPlatform(CudaPlatformBase):
@classmethod
def has_device_capability(
cls,
capability: tuple[int, int] | int,
device_id: int = 0,
) -> bool:
return False
@pytest.mark.parametrize(
("backend", "missing_module", "package_name", "install_hint"),
[
(
AttentionBackendEnum.SAGE_ATTN,
"sageattention",
"sageattention",
"uv pip install sageattention",
),
(
AttentionBackendEnum.SAGE_ATTN_THREE,
"sageattn3",
"sageattn3",
"uv pip install sageattn3",
),
(
AttentionBackendEnum.ATTN_QAT_INFER,
"fastvideo.attention.backends.attn_qat_infer",
"attn_qat_infer",
"https://hao-ai-lab.github.io/FastVideo/",
),
],
)
def test_explicit_backend_import_failure_does_not_fallback(
monkeypatch,
backend: AttentionBackendEnum,
missing_module: str,
package_name: str,
install_hint: str,
) -> None:
monkeypatch.setenv("FASTVIDEO_ATTENTION_BACKEND", backend.name)
real_import = builtins.__import__
def import_with_missing_backend(name, globals=None, locals=None, fromlist=(), level=0):
if name == missing_module:
raise ImportError(f"No module named {name!r}")
return real_import(name, globals, locals, fromlist, level)
monkeypatch.setattr(builtins, "__import__", import_with_missing_backend)
with pytest.raises(ImportError) as excinfo:
_GpuIndependentCudaPlatform.get_attn_backend_cls(backend, head_size=64, dtype=torch.float16)
message = str(excinfo.value)
assert "explicitly requested" in message
assert backend.name in message
assert package_name in message
assert install_hint in message
assert excinfo.value.__cause__ is not None
def test_explicit_attn_qat_infer_unavailable_does_not_fallback(monkeypatch) -> None:
monkeypatch.setenv("FASTVIDEO_ATTENTION_BACKEND", AttentionBackendEnum.ATTN_QAT_INFER.name)
module_name = "fastvideo.attention.backends.attn_qat_infer"
unavailable_backend = ModuleType(module_name)
unavailable_backend.AttnQatInferBackend = object
unavailable_backend.is_attn_qat_infer_available = lambda: False
monkeypatch.setitem(sys.modules, module_name, unavailable_backend)
with pytest.raises(ImportError) as excinfo:
_GpuIndependentCudaPlatform.get_attn_backend_cls(
AttentionBackendEnum.ATTN_QAT_INFER,
head_size=64,
dtype=torch.float16,
)
message = str(excinfo.value)
assert "explicitly requested" in message
assert AttentionBackendEnum.ATTN_QAT_INFER.name in message
assert "attn_qat_infer" in message
assert "not built" in message
assert "Install the FastVideo Attn-QAT inference kernel" in message
assert "https://hao-ai-lab.github.io/FastVideo/" in message
@@ -0,0 +1,238 @@
import json
from pathlib import Path
import pytest
from fastvideo.attention.selector import (
check_attn_backend_requirement,
get_global_forced_attn_backend,
global_force_attn_backend,
)
from fastvideo.configs.models.dits.base import DiTConfig
from fastvideo.configs.pipelines.base import PipelineConfig
from fastvideo.configs.pipelines.wan import (
FastWan2_1_T2V_480P_Config,
FastWan2_2_TI2V_5B_Config,
FastWan2_2_TI2V_5B_FullAttn_Config,
WanT2V480PConfig,
)
from fastvideo.models.dits.wanvideo import (
WanTransformerBlock,
WanTransformerBlock_VSA,
_select_wan_transformer_block,
)
from fastvideo.platforms.interface import AttentionBackendEnum
from fastvideo.registry import get_pipeline_config_cls_from_name
VSA = AttentionBackendEnum.VIDEO_SPARSE_ATTN
@pytest.fixture(autouse=True)
def reset_forced_attn_backend():
global_force_attn_backend(None)
yield
global_force_attn_backend(None)
def test_fastwan_required_vsa_does_not_leak_to_base_wan_config():
# Given: no process-global attention backend force is active.
assert get_global_forced_attn_backend() is None
# When: a FastWan config is instantiated before a base Wan config.
fastwan_config = FastWan2_1_T2V_480P_Config()
base_wan_config = WanT2V480PConfig()
# Then: FastWan carries the VSA requirement on its own DiT config only.
assert (
fastwan_config.dit_config.required_attention_backend
== AttentionBackendEnum.VIDEO_SPARSE_ATTN
)
assert base_wan_config.dit_config.required_attention_backend is None
assert get_global_forced_attn_backend() is None
def test_fastwan_2_2_required_vsa_does_not_mutate_global_backend():
# Given: no process-global attention backend force is active.
assert get_global_forced_attn_backend() is None
# When: the FastWan 2.2 TI2V config is instantiated.
fastwan_config = FastWan2_2_TI2V_5B_Config()
# Then: the VSA requirement stays scoped to that DiT config.
assert (
fastwan_config.dit_config.required_attention_backend
== AttentionBackendEnum.VIDEO_SPARSE_ATTN
)
assert get_global_forced_attn_backend() is None
def test_fastwan_2_2_fullattn_config_does_not_require_vsa():
config = FastWan2_2_TI2V_5B_FullAttn_Config()
assert config.dit_config.required_attention_backend is None
assert config.dit_config.incompatible_attention_backends == (VSA, )
@pytest.mark.parametrize(
"model_id",
[
"FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers",
"FastVideo/FastWan2.2-TI2V-5B-Diffusers",
],
)
def test_fastwan_2_2_fullattn_hf_ids_resolve_dense_config(model_id: str) -> None:
assert get_pipeline_config_cls_from_name(model_id) is FastWan2_2_TI2V_5B_FullAttn_Config
def _write_minimal_wan_dmd_repo(model_dir: Path, *, expand_timesteps: bool | None = None) -> None:
model_dir.mkdir(parents=True)
(model_dir / "transformer").mkdir()
model_index = {
"_class_name": "WanDMDPipeline",
"_diffusers_version": "0.35.0.dev0",
"transformer": ["diffusers", "WanTransformer3DModel"],
}
if expand_timesteps is not None:
model_index["expand_timesteps"] = expand_timesteps
(model_dir / "model_index.json").write_text(
json.dumps(model_index),
encoding="utf-8",
)
@pytest.mark.parametrize(
"relative_model_path",
[
Path("models--FastVideo--FastWan2.2-TI2V-5B-FullAttn-Diffusers") / "snapshots" / "deadbeef",
Path("models--FastVideo--FastWan2.2-TI2V-5B-Diffusers") / "snapshots" / "deadbeef",
Path("renamed-checkpoint"),
],
)
def test_fastwan_2_2_fullattn_local_path_resolves_dense_config(tmp_path: Path,
relative_model_path: Path) -> None:
model_dir = tmp_path / relative_model_path
_write_minimal_wan_dmd_repo(model_dir, expand_timesteps=True)
resolved_cls = get_pipeline_config_cls_from_name(str(model_dir))
assert resolved_cls is FastWan2_2_TI2V_5B_FullAttn_Config
def test_fastwan_2_2_sparse_local_path_stays_on_sparse_config(tmp_path: Path) -> None:
model_dir = tmp_path / "renamed-sparse-checkpoint"
_write_minimal_wan_dmd_repo(model_dir)
resolved_cls = get_pipeline_config_cls_from_name(str(model_dir))
assert resolved_cls is FastWan2_1_T2V_480P_Config
def test_fastwan_2_2_fullattn_rejects_vsa_before_block_construction():
config = FastWan2_2_TI2V_5B_FullAttn_Config().dit_config
global_force_attn_backend(VSA)
with pytest.raises(ValueError, match=r"incompatible with the VIDEO_SPARSE_ATTN attention backend"):
_select_wan_transformer_block(config, model_name="FullAttn")
@pytest.mark.parametrize("backend", [AttentionBackendEnum.FLASH_ATTN, AttentionBackendEnum.TORCH_SDPA])
def test_fastwan_2_2_fullattn_accepts_dense_backends(backend):
config = FastWan2_2_TI2V_5B_FullAttn_Config().dit_config
global_force_attn_backend(backend)
assert _select_wan_transformer_block(config, model_name="FullAttn") is WanTransformerBlock
def test_sparse_fastwan_selects_vsa_block():
config = FastWan2_2_TI2V_5B_Config().dit_config
global_force_attn_backend(VSA)
assert _select_wan_transformer_block(config, model_name="FastWan") is WanTransformerBlock_VSA
@pytest.mark.parametrize("config_cls", [FastWan2_1_T2V_480P_Config, FastWan2_2_TI2V_5B_Config])
def test_sparse_fastwan_pipeline_config_json_roundtrip(tmp_path, config_cls):
config_path = tmp_path / "pipeline_config.json"
config = config_cls()
config.dump_to_json(str(config_path))
payload = json.loads(config_path.read_text(encoding="utf-8"))
assert payload["dit_config"]["required_attention_backend"] == VSA.name
restored = config_cls()
restored.load_from_json(str(config_path))
assert restored.dit_config.required_attention_backend is VSA
def test_fullattn_pipeline_config_json_roundtrip(tmp_path):
config_path = tmp_path / "pipeline_config.json"
config = FastWan2_2_TI2V_5B_FullAttn_Config()
config.dump_to_json(str(config_path))
payload = json.loads(config_path.read_text(encoding="utf-8"))
assert payload["dit_config"]["required_attention_backend"] is None
assert payload["dit_config"]["incompatible_attention_backends"] == [VSA.name]
restored = FastWan2_2_TI2V_5B_FullAttn_Config()
restored.load_from_json(str(config_path))
assert restored.dit_config.required_attention_backend is None
assert restored.dit_config.incompatible_attention_backends == (VSA, )
def test_fullattn_config_cannot_be_overridden_to_require_vsa():
config = FastWan2_2_TI2V_5B_FullAttn_Config()
config.update_pipeline_config({"dit_config": {"required_attention_backend": VSA.name}})
assert config.dit_config.required_attention_backend is None
assert config.dit_config.incompatible_attention_backends == (VSA, )
def test_pipeline_config_load_restores_attention_backend_constraints(tmp_path):
config_path = tmp_path / "pipeline_config.json"
config = PipelineConfig(
dit_config=DiTConfig(
required_attention_backend=VSA,
incompatible_attention_backends=(AttentionBackendEnum.FLASH_ATTN, ),
))
config.dump_to_json(str(config_path))
restored = PipelineConfig()
restored.load_from_json(str(config_path))
assert restored.dit_config.required_attention_backend is VSA
assert restored.dit_config.incompatible_attention_backends == (AttentionBackendEnum.FLASH_ATTN, )
def test_check_requirement_returns_none_when_unrequired_and_unset(monkeypatch):
monkeypatch.delenv("FASTVIDEO_ATTENTION_BACKEND", raising=False)
assert check_attn_backend_requirement(None) is None
def test_check_requirement_passes_when_env_matches(monkeypatch):
monkeypatch.setenv("FASTVIDEO_ATTENTION_BACKEND", VSA.name)
assert check_attn_backend_requirement(VSA) == VSA
def test_check_requirement_passes_when_global_force_matches(monkeypatch):
# Env unset, but a global force satisfies the requirement (force > env).
monkeypatch.delenv("FASTVIDEO_ATTENTION_BACKEND", raising=False)
global_force_attn_backend(VSA)
assert check_attn_backend_requirement(VSA) == VSA
def test_check_requirement_raises_when_env_unset(monkeypatch):
monkeypatch.delenv("FASTVIDEO_ATTENTION_BACKEND", raising=False)
with pytest.raises(ValueError) as excinfo:
check_attn_backend_requirement(VSA, model_name="FastWan")
message = str(excinfo.value)
assert VSA.name in message
assert "FastWan" in message
def test_check_requirement_raises_when_env_mismatches(monkeypatch):
monkeypatch.setenv("FASTVIDEO_ATTENTION_BACKEND", AttentionBackendEnum.FLASH_ATTN.name)
with pytest.raises(ValueError):
check_attn_backend_requirement(VSA)
@@ -0,0 +1,150 @@
# SPDX-License-Identifier: Apache-2.0
"""Tests for selector-side capability validation (#1254).
The selector checks the resolved backend against its self-described
capabilities (AttentionBackend.validate_compatibility): an explicitly
selected backend hard-fails only when the resolution honored the
selection and that very backend is incompatible, mirroring how the
platform layer hard-fails on missing explicitly-requested backends. A
resolution that fell back to a different backend (the pin is outside the
layer's supported set, or the platform substituted a fallback such as
SDPA for an unsupported head size) only warns -- once per cached
resolution -- as does auto-selection.
These use dummy backends and a stubbed platform, so they run on CPU.
"""
from types import SimpleNamespace
from unittest import mock
import pytest
import torch
import fastvideo.platforms as platforms
from fastvideo.attention import selector
from fastvideo.attention.backends.abstract import (AttentionBackend, AttentionImpl, AttentionMetadata,
AttentionMetadataBuilder)
from fastvideo.platforms.interface import AttentionBackendEnum
SUPPORTED_BACKENDS = (AttentionBackendEnum.FLASH_ATTN, )
class _RestrictedBackend(AttentionBackend):
"""Backend that only supports head sizes 64 and 128.
Named FLASH_ATTN so the stubbed platform models an honored explicit
selection: the resolved class reports the selected backend's name.
"""
@staticmethod
def get_name() -> str:
return "FLASH_ATTN"
@staticmethod
def get_supported_head_sizes() -> list[int]:
return [64, 128]
@staticmethod
def get_impl_cls() -> type[AttentionImpl]:
raise NotImplementedError
@staticmethod
def get_metadata_cls() -> type[AttentionMetadata]:
raise NotImplementedError
@staticmethod
def get_builder_cls() -> type[AttentionMetadataBuilder]:
raise NotImplementedError
class _FallbackBackend(_RestrictedBackend):
"""A platform-substituted fallback: same restrictions, different name."""
@staticmethod
def get_name() -> str:
return "TORCH_SDPA"
@pytest.fixture(autouse=True)
def reset_attention_backend_selector():
# Also clears the selector cache between tests.
selector.global_force_attn_backend(None)
yield
selector.global_force_attn_backend(None)
@pytest.fixture(autouse=True)
def stub_backend_resolution(monkeypatch):
monkeypatch.delenv("FASTVIDEO_ATTENTION_BACKEND", raising=False)
platform = SimpleNamespace(
device_name="test device",
get_attn_backend_cls=lambda selected_backend, head_size, dtype: "RESTRICTED",
)
monkeypatch.setattr(platforms, "_current_platform", platform)
monkeypatch.setattr(selector, "resolve_obj_by_qualname", {"RESTRICTED": _RestrictedBackend}.__getitem__)
def test_explicitly_selected_incompatible_backend_raises(monkeypatch):
monkeypatch.setenv("FASTVIDEO_ATTENTION_BACKEND", AttentionBackendEnum.FLASH_ATTN.name)
with pytest.raises(ValueError, match="head_size"):
selector.get_attn_backend(96, torch.float16, SUPPORTED_BACKENDS)
def test_explicitly_selected_incompatible_sdpa_backend_raises(monkeypatch):
monkeypatch.setenv("FASTVIDEO_ATTENTION_BACKEND", AttentionBackendEnum.TORCH_SDPA.name)
monkeypatch.setattr(selector, "resolve_obj_by_qualname", {"RESTRICTED": _FallbackBackend}.__getitem__)
with pytest.raises(ValueError, match="head_size"):
selector.get_attn_backend(96, torch.float16, (AttentionBackendEnum.TORCH_SDPA, ))
def test_auto_selected_incompatible_backend_warns_once():
with mock.patch.object(selector.logger, "warning") as mock_warn:
first = selector.get_attn_backend(96, torch.float16, SUPPORTED_BACKENDS)
# A second lookup is a cache hit and must not warn again.
second = selector.get_attn_backend(96, torch.float16, SUPPORTED_BACKENDS)
assert first is _RestrictedBackend
assert second is _RestrictedBackend
mock_warn.assert_called_once()
# Backend name and reason are passed as format args.
assert "FLASH_ATTN" in mock_warn.call_args.args
def test_compatible_explicit_selection_passes(monkeypatch):
monkeypatch.setenv("FASTVIDEO_ATTENTION_BACKEND", AttentionBackendEnum.FLASH_ATTN.name)
with mock.patch.object(selector.logger, "warning") as mock_warn:
backend = selector.get_attn_backend(64, torch.float16, SUPPORTED_BACKENDS)
assert backend is _RestrictedBackend
mock_warn.assert_not_called()
def test_pinned_backend_platform_fallback_warns_instead_of_raising(monkeypatch):
# matrixgame2's CLIP vision encoder under a pinned FLASH_ATTN: the layer's
# head size (80) makes the platform substitute SDPA, whose declared
# capabilities also reject the head size. The pin was never honored for
# this layer, so the load must warn (once) instead of hard-failing.
monkeypatch.setenv("FASTVIDEO_ATTENTION_BACKEND", AttentionBackendEnum.FLASH_ATTN.name)
platform = SimpleNamespace(
device_name="test device",
get_attn_backend_cls=lambda selected_backend, head_size, dtype: "FALLBACK",
)
monkeypatch.setattr(platforms, "_current_platform", platform)
monkeypatch.setattr(selector, "resolve_obj_by_qualname", {"FALLBACK": _FallbackBackend}.__getitem__)
with mock.patch.object(selector.logger, "warning") as mock_warn:
first = selector.get_attn_backend(96, torch.float16, SUPPORTED_BACKENDS)
# A second lookup is a cache hit and must not warn again.
second = selector.get_attn_backend(96, torch.float16, SUPPORTED_BACKENDS)
assert first is _FallbackBackend
assert second is _FallbackBackend
mock_warn.assert_called_once()
assert "TORCH_SDPA" in mock_warn.call_args.args
def test_pin_outside_layer_supported_set_falls_back_with_warning(monkeypatch):
# A layer whose declared supported set excludes the pinned backend (e.g. an
# SDPA-only aux encoder) never participates in the global pin: it falls
# back to automatic selection with a warning instead of raising.
monkeypatch.setenv("FASTVIDEO_ATTENTION_BACKEND", AttentionBackendEnum.FLASH_ATTN.name)
with mock.patch.object(selector.logger, "warning") as mock_warn:
backend = selector.get_attn_backend(64, torch.float16, (AttentionBackendEnum.TORCH_SDPA, ))
assert backend is _RestrictedBackend
mock_warn.assert_called_once()
+1 -1
View File
@@ -277,7 +277,7 @@ def run_self_forcing_tests():
@app.function(gpu="L40S:1", image=image, timeout=900, secrets=[ci_env_secret])
def run_unit_test():
run_test(
"pytest ./fastvideo/tests/api/ ./fastvideo/tests/contract/ ./fastvideo/tests/dataset/ ./fastvideo/tests/workflow/ ./fastvideo/tests/entrypoints/ ./fastvideo/tests/train/ ./fastvideo/tests/stages/ ./fastvideo/tests/ops/ ./fastvideo/tests/training/test_trackers.py ./fastvideo/tests/attention/test_sdpa_metadata_mask_contract.py --ignore=./fastvideo/tests/entrypoints/test_openai_api_integration.py --ignore=./fastvideo/tests/train/models --ignore=./fastvideo/tests/train/methods -vs"
"pytest ./fastvideo/tests/api/ ./fastvideo/tests/attention/ ./fastvideo/tests/contract/ ./fastvideo/tests/dataset/ ./fastvideo/tests/workflow/ ./fastvideo/tests/entrypoints/ ./fastvideo/tests/train/ ./fastvideo/tests/stages/ ./fastvideo/tests/ops/ ./fastvideo/tests/training/test_trackers.py --ignore=./fastvideo/tests/entrypoints/test_openai_api_integration.py --ignore=./fastvideo/tests/train/models --ignore=./fastvideo/tests/train/methods -vs"
)
@@ -1,5 +1,5 @@
# Run with:
# FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN fastvideo generate --config scripts/inference/inference_wan_VSA_DMD_5B_720P.yaml
# FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN fastvideo generate --config scripts/inference/inference_wan_FullAttn_DMD_5B_720P.yaml
generator:
model_path: FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers
engine: