mirror of
https://github.com/NVIDIA/Model-Optimizer.git
synced 2026-10-02 03:14:52 +08:00
[5/n] Add VSA for Video Diffusion (#1053)
### What does this PR do?
Type of change: ? <!-- Use one of the following: Bug fix, new feature,
new example, new tests, documentation. -->
<!-- Details about the change. -->
New feature. Adds Video Sparse Attention (VSA) as a new sparse attention
method in ModelOpt. VSA implements a two-branch architecture
(compression + sparse) using 3D block tiling for video diffusion model.
VSA integrates with HuggingFace models by registering as
attn_implementation="modelopt_vsa" in HF's ALL_ATTENTION_FUNCTIONS,
which is the same pattern used by the existing Triton FA backend. After
`sparsify()`, HF dispatches Q, K, V directly to the VSA kernel with no
monkey-patching needed.
### Usage
```python
# Load any HuggingFace model
from transformers import AutoModelForCausalLM
model = AutoModelForCausalLM.from_pretrained("meta-llama/Llama-3.1-8B")
# Define VSA config
vsa_config = {
"sparse_cfg": {
"*attn*": {
"method": "vsa",
"block_size_3d": (4, 4, 4), # 3D tile dimensions (T, H, W)
"top_k_ratio": 0.5, # keep top 50% of blocks
"video_shape": (8, 16, 16), # video dims after patchification
"enable": True,
},
"default": {"enable": False},
},
}
# Apply — registers modelopt_vsa with HF automatically
model = mtsa.sparsify(model, vsa_config)
```
### Testing
<!-- Mention how have you tested your change if applicable. -->
`pytest tests/unit/torch/sparsity/attention_sparsity/test_vsa.py`
### Before your PR is "*Ready for review*"
Make sure you read and follow [Contributor
guidelines](https://github.com/NVIDIA/Model-Optimizer/blob/main/CONTRIBUTING.md)
and your commits are signed (`git commit -s -S`).
Make sure you read and follow the [Security Best
Practices](https://github.com/NVIDIA/Model-Optimizer/blob/main/SECURITY.md#security-coding-practices-for-contributors)
(e.g. avoiding hardcoded `trust_remote_code=True`, `torch.load(...,
weights_only=False)`, `pickle`, etc.).
- Is this change backward compatible?: ✅ / ❌ / N/A <!--- If ❌, explain
why. -->
- If you copied code from any other sources or added a new PIP
dependency, did you follow guidance in `CONTRIBUTING.md`: ✅ / ❌ / N/A
<!--- Mandatory -->
- Did you write any new necessary tests?: ✅ / ❌ / N/A <!--- Mandatory
for new features or examples. -->
- Did you update
[Changelog](https://github.com/NVIDIA/Model-Optimizer/blob/main/CHANGELOG.rst)?:
✅ / ❌ / N/A <!--- Only for new features, API changes, critical bug fixes
or backward incompatible changes. -->
### Additional Information
<!-- E.g. related issue. -->
<!-- This is an auto-generated comment: release notes by coderabbit.ai
-->
## Summary by CodeRabbit
* **New Features**
* Added Video Sparse Attention (VSA) as a new sparse attention method
for attention optimization.
* Introduced VSA configuration support with customizable block sizes and
sparsity ratios.
* Integrated VSA with HuggingFace transformers for enhanced model
compatibility.
<!-- end of auto-generated comment: release notes by coderabbit.ai -->
---------
Signed-off-by: Kai Xu <kaix@nvidia.com>
This commit is contained in:
@@ -101,6 +101,7 @@ repos:
|
||||
examples/speculative_decoding/server_generate.py|
|
||||
experimental/dms/models/qwen3/configuration_qwen3_dms.py|
|
||||
experimental/dms/models/qwen3/modeling_qwen3_dms.py|
|
||||
modelopt/torch/sparsity/attention_sparsity/methods/vsa_utils.py|
|
||||
)$
|
||||
|
||||
# Default hook for Apache 2.0 in c/c++/cuda files
|
||||
|
||||
@@ -9,6 +9,7 @@ NVIDIA Model Optimizer Changelog
|
||||
- Added iterator interface using CalibrationDataReader in ONNX quantization workflow.
|
||||
- Add N:M sparse softmax support to the Triton flash attention kernel (``modelopt.torch.kernels.triton_fa``). See `examples/llm_sparsity/attention_sparsity/README.md <https://github.com/NVIDIA/Model-Optimizer/tree/main/examples/llm_sparsity/attention_sparsity>`_ for usage.
|
||||
- Add skip-softmax skipping to the Triton flash attention kernel (``modelopt.torch.kernels.triton_fa``). See `examples/llm_sparsity/attention_sparsity/README.md <https://github.com/NVIDIA/Model-Optimizer/tree/main/examples/llm_sparsity/attention_sparsity>`_ for usage.
|
||||
- Add Video Sparse Attention (VSA) method for video diffusion models (``modelopt.torch.sparsity.attention_sparsity``). VSA uses 3D block tiling with a two-branch architecture for attention speedup.
|
||||
- Enable PTQ workflow for the Step3.5-Flash MoE model with NVFP4 W4A4 + FP8 KV cache quantization. See `modelopt_recipes/models/Step3.5-Flash/nvfp4-mlp-only.yaml <https://github.com/NVIDIA/Model-Optimizer/blob/main/modelopt_recipes/models/Step3.5-Flash/nvfp4-mlp-only.yaml>`_ for more details.
|
||||
- Add support for vLLM fakequant reload using ModelOpt state for HF models. See `examples/vllm_serve/README.md <https://github.com/NVIDIA/Model-Optimizer/tree/main/examples/vllm_serve#load-qatptq-model-and-serve-in-vllm-wip>`_ for more details.
|
||||
- [Early Testing] Add Claude Code PTQ skill (``.claude/skills/ptq/``) for agent-assisted post-training quantization. The skill guides the agent through environment detection, model support checking, format selection, and execution via the launcher or manual SLURM/Docker/bare GPU paths. Includes handling for unlisted models with custom module patching. This feature is in early testing — use with caution.
|
||||
|
||||
@@ -35,10 +35,9 @@ if torch.cuda.is_available():
|
||||
|
||||
attention = _attention
|
||||
IS_AVAILABLE = True
|
||||
with import_plugin("transformers"):
|
||||
from .hf_triton_attention import register_triton_attention as _register_triton_attention
|
||||
from .hf_triton_attention import register_triton_attention as _register_triton_attention
|
||||
|
||||
register_triton_attention = _register_triton_attention
|
||||
register_triton_attention = _register_triton_attention
|
||||
|
||||
__all__ = [
|
||||
"IS_AVAILABLE",
|
||||
|
||||
@@ -498,7 +498,7 @@ SKIP_SOFTMAX_DEFAULT = {
|
||||
|
||||
|
||||
# Configuration with RULER calibration
|
||||
# Note: threshold field is omitted - calibration determines dynamic threshold λ = a / length
|
||||
# Note: threshold field is omitted - calibration determines dynamic threshold lambda = a / length
|
||||
# The calibrated threshold adapts to sequence length for optimal sparsity
|
||||
SKIP_SOFTMAX_CALIB = {
|
||||
"sparse_cfg": {
|
||||
@@ -521,6 +521,136 @@ SKIP_SOFTMAX_CALIB = {
|
||||
}
|
||||
|
||||
|
||||
class VSAAttributeConfig(ModeloptBaseConfig):
|
||||
"""Video Sparse Attention (VSA) attribute configuration.
|
||||
|
||||
VSA uses a two-branch architecture optimized for video diffusion models:
|
||||
1. Compression branch: Block-averaged coarse attention
|
||||
2. Sparse branch: Top-K block selection for fine-grained attention
|
||||
"""
|
||||
|
||||
method: str = ModeloptField(
|
||||
default="vsa",
|
||||
title="Sparse attention method.",
|
||||
description="Must be 'vsa' for Video Sparse Attention.",
|
||||
)
|
||||
|
||||
enable: bool = ModeloptField(
|
||||
default=True,
|
||||
title="Enable VSA.",
|
||||
description="If True, enables Video Sparse Attention. If False, bypasses sparsity.",
|
||||
)
|
||||
|
||||
block_size_3d: tuple[int, int, int] | list[int] = ModeloptField(
|
||||
default=(4, 4, 4),
|
||||
title="3D block size.",
|
||||
description=(
|
||||
"Video block dimensions (T, H, W) for spatial-temporal tiling. "
|
||||
"Default (4, 4, 4) creates 64-token blocks."
|
||||
),
|
||||
)
|
||||
|
||||
top_k_ratio: float = ModeloptField(
|
||||
default=0.5,
|
||||
title="Top-K selection ratio.",
|
||||
description=(
|
||||
"Ratio of blocks to keep in sparse branch (0.0 to 1.0). "
|
||||
"Lower values mean more sparsity. Default 0.5 keeps 50% of blocks."
|
||||
),
|
||||
)
|
||||
|
||||
video_shape: tuple[int, int, int] | list[int] | None = ModeloptField(
|
||||
default=None,
|
||||
title="Video shape.",
|
||||
description=(
|
||||
"Video dimensions (T, H, W) after patchification. "
|
||||
"Required for VSA — set via config or call set_video_shape() at runtime."
|
||||
),
|
||||
)
|
||||
|
||||
collect_stats: bool = ModeloptField(
|
||||
default=False,
|
||||
title="Collect statistics.",
|
||||
description="Whether to collect sparsity statistics during forward pass.",
|
||||
)
|
||||
|
||||
@field_validator("method")
|
||||
@classmethod
|
||||
def validate_vsa_method(cls, v):
|
||||
"""Validate method is 'vsa'."""
|
||||
if v != "vsa":
|
||||
raise ValueError(f"VSAAttributeConfig method must be 'vsa', got '{v}'")
|
||||
return v
|
||||
|
||||
@field_validator("block_size_3d")
|
||||
@classmethod
|
||||
def validate_block_size_3d(cls, v):
|
||||
"""Validate 3D block size."""
|
||||
if isinstance(v, list):
|
||||
v = tuple(v)
|
||||
if len(v) != 3:
|
||||
raise ValueError(f"block_size_3d must have 3 elements (T, H, W), got {len(v)}")
|
||||
if any(x <= 0 for x in v):
|
||||
raise ValueError(f"All block_size_3d values must be positive, got {v}")
|
||||
return v
|
||||
|
||||
@field_validator("top_k_ratio")
|
||||
@classmethod
|
||||
def validate_top_k_ratio(cls, v):
|
||||
"""Validate top-K ratio is in valid range."""
|
||||
if not 0.0 < v <= 1.0:
|
||||
raise ValueError(f"top_k_ratio must be in range (0, 1], got {v}")
|
||||
return v
|
||||
|
||||
@field_validator("video_shape")
|
||||
@classmethod
|
||||
def validate_video_shape(cls, v):
|
||||
"""Validate video shape if provided."""
|
||||
if v is None:
|
||||
return v
|
||||
if isinstance(v, list):
|
||||
v = tuple(v)
|
||||
if len(v) != 3:
|
||||
raise ValueError(f"video_shape must have 3 elements (T, H, W), got {len(v)}")
|
||||
if any(x <= 0 for x in v):
|
||||
raise ValueError(f"All video_shape values must be positive, got {v}")
|
||||
return v
|
||||
|
||||
|
||||
class VSAConfig(SparseAttentionConfig):
|
||||
"""Configuration for Video Sparse Attention optimization."""
|
||||
|
||||
sparse_cfg: SparseAttentionCfgType = ModeloptField(
|
||||
default={
|
||||
"*attn*": {
|
||||
"method": "vsa",
|
||||
"block_size_3d": (4, 4, 4),
|
||||
"top_k_ratio": 0.5,
|
||||
"enable": True,
|
||||
},
|
||||
"default": {"enable": False},
|
||||
},
|
||||
title="VSA configuration",
|
||||
description="Pattern-based configuration for Video Sparse Attention.",
|
||||
validate_default=True,
|
||||
)
|
||||
|
||||
|
||||
# Pre-defined VSA Configuration for video diffusion models.
|
||||
# Pattern "*attn*" matches attention module names by convention.
|
||||
VSA_DEFAULT = {
|
||||
"sparse_cfg": {
|
||||
"*attn*": {
|
||||
"method": "vsa",
|
||||
"block_size_3d": (4, 4, 4),
|
||||
"top_k_ratio": 0.5,
|
||||
"enable": True,
|
||||
},
|
||||
"default": {"enable": False},
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
# Default N:M sparse softmax configuration
|
||||
SPARSE_SOFTMAX_DEFAULT = {
|
||||
"sparse_cfg": {
|
||||
@@ -557,10 +687,13 @@ __all__ = [
|
||||
"SKIP_SOFTMAX_DEFAULT",
|
||||
"SKIP_SOFTMAX_TRITON_DEFAULT",
|
||||
"SPARSE_SOFTMAX_DEFAULT",
|
||||
"VSA_DEFAULT",
|
||||
"CalibrationConfig",
|
||||
"FlashSkipSoftmaxConfig",
|
||||
"SparseAttentionAttributeConfig",
|
||||
"SparseAttentionCfgType",
|
||||
"SparseAttentionConfig",
|
||||
"SparseAttributeConfig",
|
||||
"VSAAttributeConfig",
|
||||
"VSAConfig",
|
||||
]
|
||||
|
||||
@@ -33,26 +33,44 @@ from .utils import get_named_sparse_attention_modules, get_sparse_attention_modu
|
||||
|
||||
|
||||
def _set_attn_implementation(model: nn.Module, config: SparseAttentionConfig) -> None:
|
||||
"""Set the correct attn_implementation based on the sparse attention backend.
|
||||
"""Set the correct attn_implementation based on the sparse attention method/backend.
|
||||
|
||||
- ``backend="triton"``: registers the Triton kernel with HF and sets
|
||||
``attn_implementation="modelopt_triton"``.
|
||||
- ``backend="pytorch"`` (default): sets ``attn_implementation="eager"`` so that
|
||||
softmax-patching methods (e.g. skip-softmax) work correctly. FlashAttention
|
||||
and SDPA bypass ``F.softmax``, so eager is required.
|
||||
- ``method="vsa"``: no-op. VSA patches ``F.scaled_dot_product_attention``
|
||||
directly in ``SparseAttentionModule.forward()``, so no ``attn_implementation``
|
||||
change is needed.
|
||||
|
||||
This is called automatically during ``mtsa.sparsify()`` so users never need
|
||||
to manually set ``attn_implementation``.
|
||||
"""
|
||||
sparse_cfg = config.sparse_cfg if hasattr(config, "sparse_cfg") else {}
|
||||
|
||||
# Collect backends only from layer configs (identified by having a "method" key).
|
||||
# Collect methods and backends only from layer configs (identified by having a "method" key).
|
||||
# Other dict entries (e.g. "calibration") are not layer configs.
|
||||
backends = {
|
||||
v.get("backend", "pytorch")
|
||||
for v in sparse_cfg.values()
|
||||
if isinstance(v, dict) and "method" in v
|
||||
}
|
||||
layer_cfgs = [v for v in sparse_cfg.values() if isinstance(v, dict) and "method" in v]
|
||||
methods = {v.get("method") for v in layer_cfgs}
|
||||
backends = {v.get("backend", "pytorch") for v in layer_cfgs}
|
||||
|
||||
# VSA patches F.scaled_dot_product_attention directly — it does not change
|
||||
# attn_implementation. Skip the rest for VSA-only configs.
|
||||
if methods == {"vsa"}:
|
||||
return
|
||||
|
||||
# Reject mixed VSA + non-VSA configs (VSA patches SDPA globally per-module,
|
||||
# while softmax-patching methods need attn_implementation="eager").
|
||||
non_vsa_methods = methods - {"vsa"}
|
||||
if "vsa" in methods and non_vsa_methods:
|
||||
raise ValueError(
|
||||
f"Cannot mix VSA with other sparse attention methods ({non_vsa_methods}). "
|
||||
f"VSA patches F.scaled_dot_product_attention, which is incompatible "
|
||||
f"with softmax-patching or triton methods."
|
||||
)
|
||||
|
||||
model_config = getattr(model, "config", None)
|
||||
|
||||
if "triton" in backends and "pytorch" in backends:
|
||||
raise ValueError(
|
||||
@@ -60,15 +78,12 @@ def _set_attn_implementation(model: nn.Module, config: SparseAttentionConfig) ->
|
||||
"supported. All sparse attention layers must use the same backend."
|
||||
)
|
||||
|
||||
model_config = getattr(model, "config", None)
|
||||
|
||||
if "triton" in backends:
|
||||
from .kernels import register_triton_attention
|
||||
|
||||
if register_triton_attention is None:
|
||||
raise ImportError(
|
||||
"Triton backend requires 'triton' and 'transformers' packages. "
|
||||
"Install with: pip install triton transformers"
|
||||
"Triton backend requires 'triton' package. Install with: pip install triton"
|
||||
)
|
||||
if not register_triton_attention():
|
||||
raise RuntimeError(
|
||||
@@ -83,7 +98,6 @@ def _set_attn_implementation(model: nn.Module, config: SparseAttentionConfig) ->
|
||||
model_config._attn_implementation = "modelopt_triton"
|
||||
elif model_config is not None:
|
||||
# For pytorch backend, force eager for softmax patching.
|
||||
# TODO: Add the triton backend support for skip-softmax.
|
||||
model_config._attn_implementation = "eager"
|
||||
|
||||
|
||||
|
||||
@@ -24,4 +24,5 @@ __all__ = [
|
||||
]
|
||||
|
||||
# Import method implementations to trigger registration
|
||||
from . import flash_skip_softmax, triton_skip_softmax, triton_sparse_softmax
|
||||
# Note: vsa imports no external deps at module level; fastvideo_kernel is imported lazily at runtime.
|
||||
from . import flash_skip_softmax, triton_skip_softmax, triton_sparse_softmax, vsa
|
||||
|
||||
@@ -37,6 +37,31 @@ class SparseAttentionMethod(ABC):
|
||||
self.calibration_params: dict[str, dict[str, float]] | None = None
|
||||
# Target sparsity ratio per phase: {"prefill": 0.5, "decode": 0.5}
|
||||
self.target_sparse_ratio: dict[str, float] | None = None
|
||||
# Video shape for VSA (T, H, W). None for non-VSA methods.
|
||||
self.video_shape: tuple[int, int, int] | None = None
|
||||
|
||||
def forward_attention(
|
||||
self,
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
value: torch.Tensor,
|
||||
**kwargs,
|
||||
) -> tuple[torch.Tensor, dict]:
|
||||
"""Compute full attention replacement (e.g. VSA).
|
||||
|
||||
Default: raises NotImplementedError. Override for methods that replace
|
||||
the entire attention computation rather than patching softmax.
|
||||
|
||||
Args:
|
||||
query: Query tensor [batch, heads, seq_len, dim].
|
||||
key: Key tensor [batch, heads, seq_len, dim].
|
||||
value: Value tensor [batch, heads, seq_len, dim].
|
||||
**kwargs: Method-specific arguments.
|
||||
|
||||
Returns:
|
||||
Tuple of (attention_output, stats_dict).
|
||||
"""
|
||||
raise NotImplementedError(f"{type(self).__name__} does not implement forward_attention.")
|
||||
|
||||
def calculate_sparsity(
|
||||
self,
|
||||
|
||||
@@ -0,0 +1,347 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
"""Video Sparse Attention (VSA) method for video diffusion models.
|
||||
|
||||
VSA implements a two-branch sparse attention architecture:
|
||||
1. Compression Branch: Averages tokens within 3D video blocks and computes coarse attention
|
||||
2. Sparse Branch: Selects top-K blocks based on importance and computes fine-grained attention
|
||||
|
||||
Uses the optimized Triton kernel from fastvideo_kernel.
|
||||
|
||||
Integration:
|
||||
After ``mtsa.sparsify(model, VSA_DEFAULT)``, each attention layer's
|
||||
``F.scaled_dot_product_attention`` call is intercepted and replaced by the VSA
|
||||
kernel. Cross-attention (Q/K have different seq_len) is automatically skipped.
|
||||
This works with HF transformers and diffusers.
|
||||
"""
|
||||
|
||||
import math
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
|
||||
from . import SparseAttentionMethod, register_sparse_method
|
||||
from .vsa_utils import (
|
||||
construct_variable_block_sizes,
|
||||
get_non_pad_index,
|
||||
get_reverse_tile_partition_indices,
|
||||
get_tile_partition_indices,
|
||||
)
|
||||
|
||||
|
||||
@register_sparse_method("vsa")
|
||||
class VSA(SparseAttentionMethod):
|
||||
"""Video Sparse Attention with two-branch architecture.
|
||||
|
||||
VSA combines a compression branch (coarse-grained block attention) with
|
||||
a sparse branch (fine-grained attention on top-K selected blocks).
|
||||
|
||||
The final output is: output = out_compression * gate_compress + out_sparse
|
||||
|
||||
where gate_compress is a learned parameter from the model layer that
|
||||
controls the balance between compression and sparse branches.
|
||||
|
||||
Configuration Parameters:
|
||||
- block_size_3d: 3D tile dimensions (T, H, W), default (4, 4, 4)
|
||||
- top_k_ratio: Ratio of blocks to keep (0.0-1.0), default 0.5
|
||||
- video_shape: Video dimensions (T, H, W) after patchification
|
||||
|
||||
Requirements:
|
||||
- Model must expose gate_compress parameter in attention layers
|
||||
- Input tensors must be 4D: [batch, heads, seq_len, dim]
|
||||
"""
|
||||
|
||||
def __init__(self, method_config: dict | None = None):
|
||||
"""Initialize VSA method.
|
||||
|
||||
Args:
|
||||
method_config: Configuration dict with VSA parameters.
|
||||
"""
|
||||
super().__init__()
|
||||
config = method_config or {}
|
||||
|
||||
# Block configuration
|
||||
block_size = config.get("block_size_3d", (4, 4, 4))
|
||||
if isinstance(block_size, list):
|
||||
block_size = tuple(block_size)
|
||||
if len(block_size) != 3 or any(x <= 0 for x in block_size):
|
||||
raise ValueError(f"block_size_3d must be 3 positive integers, got {block_size}")
|
||||
self.block_size_3d = block_size
|
||||
self.block_elements = block_size[0] * block_size[1] * block_size[2]
|
||||
|
||||
# Sparsity configuration
|
||||
top_k_ratio = config.get("top_k_ratio", 0.5)
|
||||
if not 0.0 < top_k_ratio <= 1.0:
|
||||
raise ValueError(f"top_k_ratio must be in (0, 1], got {top_k_ratio}")
|
||||
self.top_k_ratio = top_k_ratio
|
||||
|
||||
# Video shape (can be set dynamically via set_video_shape or at call time)
|
||||
video_shape = config.get("video_shape", None)
|
||||
if video_shape is not None:
|
||||
if isinstance(video_shape, list):
|
||||
video_shape = tuple(video_shape)
|
||||
if len(video_shape) != 3 or any(x <= 0 for x in video_shape):
|
||||
raise ValueError(f"video_shape must be 3 positive integers, got {video_shape}")
|
||||
self.video_shape = video_shape
|
||||
|
||||
# Track last computed statistics
|
||||
self._last_stats: dict = {}
|
||||
|
||||
# Metadata cache: avoids recomputing tile indices on every forward pass.
|
||||
# Matches FastVideo's @lru_cache on utility functions.
|
||||
self._cached_metadata: dict[str, Any] | None = None
|
||||
self._cached_metadata_key: tuple | None = None
|
||||
|
||||
def set_video_shape(self, video_shape: tuple[int, int, int]):
|
||||
"""Set video shape for current forward pass.
|
||||
|
||||
Args:
|
||||
video_shape: Video dimensions (T, H, W) after patchification.
|
||||
"""
|
||||
self.video_shape = video_shape
|
||||
|
||||
def _compute_metadata(self, seq_len: int, device: torch.device) -> dict[str, Any]:
|
||||
"""Compute block metadata from video shape.
|
||||
|
||||
Results are cached and reused when called with the same (seq_len, video_shape)
|
||||
to avoid recomputing tile indices on every denoising step, matching FastVideo's
|
||||
``@functools.lru_cache`` on the underlying utility functions.
|
||||
|
||||
Args:
|
||||
seq_len: Sequence length (should equal T * H * W).
|
||||
device: Device for tensors.
|
||||
|
||||
Returns:
|
||||
Metadata dict with tile indices, variable sizes, etc.
|
||||
"""
|
||||
if self.video_shape is None:
|
||||
raise ValueError(
|
||||
f"video_shape must be provided for VSA but is None (seq_len={seq_len}). "
|
||||
f"Set it via the VSA config ('video_shape' key), call set_video_shape(), "
|
||||
f"or use a model-specific plugin (e.g., LTX-2 plugin) that computes it "
|
||||
f"from the model's patchifier."
|
||||
)
|
||||
|
||||
# Return cached metadata if inputs haven't changed
|
||||
cache_key = (seq_len, self.video_shape, device)
|
||||
if self._cached_metadata is not None and self._cached_metadata_key == cache_key:
|
||||
return self._cached_metadata
|
||||
|
||||
vid_t, vid_h, vid_w = self.video_shape
|
||||
ts_t, ts_h, ts_w = self.block_size_3d
|
||||
|
||||
# Validate sequence length matches video shape
|
||||
expected_seq_len = vid_t * vid_h * vid_w
|
||||
if seq_len != expected_seq_len:
|
||||
raise ValueError(
|
||||
f"Sequence length {seq_len} does not match video shape {self.video_shape} "
|
||||
f"(expected {expected_seq_len})"
|
||||
)
|
||||
|
||||
# Calculate number of tiles
|
||||
num_tiles = (
|
||||
math.ceil(vid_t / ts_t),
|
||||
math.ceil(vid_h / ts_h),
|
||||
math.ceil(vid_w / ts_w),
|
||||
)
|
||||
total_tiles = num_tiles[0] * num_tiles[1] * num_tiles[2]
|
||||
|
||||
# Get partitioning indices
|
||||
tile_indices = get_tile_partition_indices(self.video_shape, self.block_size_3d, device)
|
||||
reverse_indices = get_reverse_tile_partition_indices(
|
||||
self.video_shape, self.block_size_3d, device
|
||||
)
|
||||
variable_sizes = construct_variable_block_sizes(
|
||||
self.video_shape, num_tiles, self.block_size_3d, device
|
||||
)
|
||||
non_pad_index = get_non_pad_index(variable_sizes, self.block_elements)
|
||||
|
||||
# Calculate padded sizes
|
||||
t_padded = num_tiles[0] * ts_t
|
||||
h_padded = num_tiles[1] * ts_h
|
||||
w_padded = num_tiles[2] * ts_w
|
||||
padded_seq_len = t_padded * h_padded * w_padded
|
||||
|
||||
metadata = {
|
||||
"video_shape": self.video_shape,
|
||||
"tile_size": self.block_size_3d,
|
||||
"num_tiles": num_tiles,
|
||||
"total_tiles": total_tiles,
|
||||
"tile_indices": tile_indices,
|
||||
"reverse_indices": reverse_indices,
|
||||
"variable_sizes": variable_sizes,
|
||||
"non_pad_index": non_pad_index,
|
||||
"padded_seq_len": padded_seq_len,
|
||||
}
|
||||
|
||||
# Cache for reuse across denoising steps
|
||||
self._cached_metadata = metadata
|
||||
self._cached_metadata_key = cache_key
|
||||
|
||||
return metadata
|
||||
|
||||
def _tile_tensor(self, tensor: torch.Tensor, metadata: dict) -> torch.Tensor:
|
||||
"""Rearrange tensor into tile layout with padding.
|
||||
|
||||
Args:
|
||||
tensor: Input tensor [batch, heads, seq_len, dim].
|
||||
metadata: Metadata from _compute_metadata.
|
||||
|
||||
Returns:
|
||||
Tiled tensor [batch, heads, padded_seq_len, dim].
|
||||
"""
|
||||
batch, heads, seq_len, dim = tensor.shape
|
||||
device = tensor.device
|
||||
dtype = tensor.dtype
|
||||
|
||||
tile_indices = metadata["tile_indices"]
|
||||
non_pad_index = metadata["non_pad_index"]
|
||||
padded_seq_len = metadata["padded_seq_len"]
|
||||
|
||||
# Create padded tensor
|
||||
padded = torch.zeros((batch, heads, padded_seq_len, dim), device=device, dtype=dtype)
|
||||
|
||||
# Rearrange to tile order and place in padded positions
|
||||
padded[:, :, non_pad_index] = tensor[:, :, tile_indices]
|
||||
|
||||
return padded
|
||||
|
||||
def _untile_tensor(self, tensor: torch.Tensor, metadata: dict, seq_len: int) -> torch.Tensor:
|
||||
"""Reverse tile layout back to original order.
|
||||
|
||||
Args:
|
||||
tensor: Tiled tensor [batch, heads, padded_seq_len, dim].
|
||||
metadata: Metadata from _compute_metadata.
|
||||
seq_len: Original sequence length.
|
||||
|
||||
Returns:
|
||||
Output tensor [batch, heads, seq_len, dim].
|
||||
"""
|
||||
non_pad_index = metadata["non_pad_index"]
|
||||
reverse_indices = metadata["reverse_indices"]
|
||||
|
||||
# Extract non-padded tokens and reverse order
|
||||
return tensor[:, :, non_pad_index][:, :, reverse_indices]
|
||||
|
||||
def forward_attention(
|
||||
self,
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
value: torch.Tensor,
|
||||
gate_compress: torch.Tensor | None = None,
|
||||
video_shape: tuple[int, int, int] | None = None,
|
||||
**kwargs,
|
||||
) -> tuple[torch.Tensor, dict]:
|
||||
"""Compute VSA two-branch sparse attention.
|
||||
|
||||
Data flow (mirrors FastVideo's VideoSparseAttentionImpl):
|
||||
1. Compute tile metadata from video_shape
|
||||
2. Tile Q, K, V, gate_compress into padded tile order
|
||||
3. Run Triton VSA kernel on tiled tensors
|
||||
4. Untile output back to original token order
|
||||
|
||||
Args:
|
||||
query: Query tensor [batch, heads, seq_len, dim].
|
||||
key: Key tensor [batch, heads, seq_len, dim].
|
||||
value: Value tensor [batch, heads, seq_len, dim].
|
||||
gate_compress: Learned gating weights [batch, heads, seq_len, dim].
|
||||
If None, uses equal weighting (0.5) for both branches.
|
||||
video_shape: Video dimensions (T, H, W). If None, uses self.video_shape.
|
||||
**kwargs: Additional arguments (ignored).
|
||||
|
||||
Returns:
|
||||
Tuple of (attention_output, stats) where:
|
||||
- attention_output: [batch, heads, seq_len, dim]
|
||||
- stats: Dict with sparsity statistics
|
||||
"""
|
||||
if video_shape is not None:
|
||||
self.video_shape = video_shape
|
||||
|
||||
batch, heads, seq_len, dim = query.shape
|
||||
device = query.device
|
||||
|
||||
# Compute block metadata (cached across denoising steps)
|
||||
metadata = self._compute_metadata(seq_len, device)
|
||||
total_tiles = metadata["total_tiles"]
|
||||
variable_sizes = metadata["variable_sizes"]
|
||||
|
||||
# Calculate top-K based on ratio
|
||||
top_k = max(1, int(self.top_k_ratio * total_tiles))
|
||||
|
||||
# ========== TILE: rearrange tokens into tile order ==========
|
||||
# Mirrors FastVideo's VideoSparseAttentionImpl.preprocess_qkv (tile)
|
||||
query_tiled = self._tile_tensor(query, metadata)
|
||||
key_tiled = self._tile_tensor(key, metadata)
|
||||
value_tiled = self._tile_tensor(value, metadata)
|
||||
gate_tiled = (
|
||||
self._tile_tensor(gate_compress, metadata) if gate_compress is not None else None
|
||||
)
|
||||
|
||||
# ========== TRITON VSA KERNEL ==========
|
||||
# Kernel operates on tiled tensors in [batch, heads, padded_seq, dim] format
|
||||
try:
|
||||
from fastvideo_kernel import video_sparse_attn as triton_vsa_kernel
|
||||
except ImportError as e:
|
||||
raise ImportError(
|
||||
"VSA requires the 'fastvideo_kernel' package for its Triton sparse attention "
|
||||
f"kernel. Install it with: pip install fastvideo_kernel (error: {e})"
|
||||
) from e
|
||||
output_tiled = triton_vsa_kernel(
|
||||
query_tiled,
|
||||
key_tiled,
|
||||
value_tiled,
|
||||
variable_sizes, # variable_block_sizes (KV)
|
||||
variable_sizes, # q_variable_block_sizes (Q)
|
||||
top_k,
|
||||
block_size=self.block_size_3d,
|
||||
compress_attn_weight=gate_tiled,
|
||||
)
|
||||
|
||||
# ========== UNTILE: restore original token order ==========
|
||||
# Mirrors FastVideo's VideoSparseAttentionImpl.postprocess_output (untile)
|
||||
output = self._untile_tensor(output_tiled, metadata, seq_len)
|
||||
|
||||
# Compute statistics
|
||||
actual_sparsity = 1.0 - (top_k / total_tiles)
|
||||
stats = {
|
||||
"sparsity": [actual_sparsity],
|
||||
"phase": "vsa_triton",
|
||||
"total_blocks": total_tiles,
|
||||
"sparse_blocks": [total_tiles - top_k],
|
||||
"top_k": top_k,
|
||||
"video_shape": self.video_shape,
|
||||
}
|
||||
self._last_stats = stats
|
||||
|
||||
return output, stats
|
||||
|
||||
def get_threshold_info(self) -> dict[str, Any]:
|
||||
"""Get VSA configuration info.
|
||||
|
||||
Returns:
|
||||
Dictionary with VSA configuration.
|
||||
"""
|
||||
return {
|
||||
"type": "vsa",
|
||||
"block_size_3d": self.block_size_3d,
|
||||
"top_k_ratio": self.top_k_ratio,
|
||||
"video_shape": self.video_shape,
|
||||
}
|
||||
|
||||
@property
|
||||
def name(self) -> str:
|
||||
"""Method identifier."""
|
||||
return "vsa"
|
||||
@@ -0,0 +1,169 @@
|
||||
# Adapted from: https://github.com/hao-ai-lab/FastVideo/blob/5789955/fastvideo/attention/backends/video_sparse_attn.py
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
"""Utility functions for Video Sparse Attention (VSA).
|
||||
|
||||
This module provides 3D block operations for video sparse attention,
|
||||
including reshaping tensors into video blocks and variable block size computation.
|
||||
"""
|
||||
|
||||
import functools
|
||||
import math
|
||||
|
||||
import torch
|
||||
|
||||
|
||||
@functools.lru_cache(maxsize=10)
|
||||
def get_tile_partition_indices(
|
||||
video_shape: tuple[int, int, int],
|
||||
tile_size: tuple[int, int, int],
|
||||
device: torch.device,
|
||||
) -> torch.LongTensor:
|
||||
"""Get indices to partition video tokens into tiles.
|
||||
|
||||
Args:
|
||||
video_shape: Video dimensions (T, H, W) after patchification.
|
||||
tile_size: Tile dimensions (tile_T, tile_H, tile_W).
|
||||
device: Device for the output tensor.
|
||||
|
||||
Returns:
|
||||
LongTensor of indices to rearrange tokens into tile order.
|
||||
"""
|
||||
vid_t, vid_h, vid_w = video_shape
|
||||
ts, hs, ws = tile_size
|
||||
indices = torch.arange(vid_t * vid_h * vid_w, device=device, dtype=torch.long).reshape(
|
||||
vid_t, vid_h, vid_w
|
||||
)
|
||||
|
||||
tiles = []
|
||||
for t in range(math.ceil(vid_t / ts)):
|
||||
for h in range(math.ceil(vid_h / hs)):
|
||||
for w in range(math.ceil(vid_w / ws)):
|
||||
tile = indices[
|
||||
t * ts : min(t * ts + ts, vid_t),
|
||||
h * hs : min(h * hs + hs, vid_h),
|
||||
w * ws : min(w * ws + ws, vid_w),
|
||||
]
|
||||
tiles.append(tile.flatten())
|
||||
|
||||
return torch.cat(tiles, dim=0)
|
||||
|
||||
|
||||
@functools.lru_cache(maxsize=10)
|
||||
def get_reverse_tile_partition_indices(
|
||||
video_shape: tuple[int, int, int],
|
||||
tile_size: tuple[int, int, int],
|
||||
device: torch.device,
|
||||
) -> torch.LongTensor:
|
||||
"""Get indices to reverse tile partitioning back to original order.
|
||||
|
||||
Args:
|
||||
video_shape: Video dimensions (T, H, W) after patchification.
|
||||
tile_size: Tile dimensions (tile_T, tile_H, tile_W).
|
||||
device: Device for the output tensor.
|
||||
|
||||
Returns:
|
||||
LongTensor of indices to reverse the tile rearrangement.
|
||||
"""
|
||||
forward_indices = get_tile_partition_indices(video_shape, tile_size, device)
|
||||
return torch.argsort(forward_indices)
|
||||
|
||||
|
||||
@functools.lru_cache(maxsize=10)
|
||||
def construct_variable_block_sizes(
|
||||
video_shape: tuple[int, int, int],
|
||||
num_tiles: tuple[int, int, int],
|
||||
tile_size: tuple[int, int, int],
|
||||
device: torch.device,
|
||||
) -> torch.LongTensor:
|
||||
"""Compute valid (non-padded) token count for each tile.
|
||||
|
||||
Since video dimensions may not divide evenly by tile size, edge tiles
|
||||
will have fewer valid tokens. This function computes the actual valid
|
||||
token count for each tile.
|
||||
|
||||
Args:
|
||||
video_shape: Video dimensions (T, H, W) after patchification.
|
||||
num_tiles: Number of tiles in each dimension (n_T, n_H, n_W).
|
||||
tile_size: Tile dimensions (tile_T, tile_H, tile_W).
|
||||
device: Device for the output tensor.
|
||||
|
||||
Returns:
|
||||
LongTensor of shape [num_tiles_total] with valid tokens per tile.
|
||||
"""
|
||||
t, h, w = video_shape
|
||||
ts_t, ts_h, ts_w = tile_size
|
||||
n_t, n_h, n_w = num_tiles
|
||||
|
||||
def _sizes(dim_len: int, tile: int, n_tiles: int) -> torch.LongTensor:
|
||||
"""Compute size of each tile along one dimension."""
|
||||
sizes = torch.full((n_tiles,), tile, dtype=torch.long, device=device)
|
||||
remainder = dim_len - (n_tiles - 1) * tile
|
||||
sizes[-1] = remainder if remainder > 0 else tile
|
||||
return sizes
|
||||
|
||||
t_sizes = _sizes(t, ts_t, n_t) # [n_t]
|
||||
h_sizes = _sizes(h, ts_h, n_h) # [n_h]
|
||||
w_sizes = _sizes(w, ts_w, n_w) # [n_w]
|
||||
|
||||
# Broadcast multiply to get tokens per tile
|
||||
block_sizes = (
|
||||
t_sizes[:, None, None] * h_sizes[None, :, None] * w_sizes[None, None, :]
|
||||
).reshape(-1)
|
||||
|
||||
return block_sizes
|
||||
|
||||
|
||||
@functools.lru_cache(maxsize=10)
|
||||
def get_non_pad_index(
|
||||
variable_block_sizes: torch.LongTensor,
|
||||
max_block_size: int,
|
||||
) -> torch.LongTensor:
|
||||
"""Get indices of non-padded tokens in the padded layout.
|
||||
|
||||
When tiles have variable sizes, we pad to max_block_size. This function
|
||||
returns indices to extract only valid (non-padded) tokens.
|
||||
|
||||
Args:
|
||||
variable_block_sizes: Tensor of valid token counts per tile.
|
||||
max_block_size: Maximum tile size (usually tile_T * tile_H * tile_W).
|
||||
|
||||
Returns:
|
||||
LongTensor of indices for valid tokens.
|
||||
"""
|
||||
n_win = variable_block_sizes.shape[0]
|
||||
device = variable_block_sizes.device
|
||||
|
||||
starts_pad = torch.arange(n_win, device=device) * max_block_size
|
||||
index_pad = starts_pad[:, None] + torch.arange(max_block_size, device=device)[None, :]
|
||||
index_mask = (
|
||||
torch.arange(max_block_size, device=device)[None, :] < variable_block_sizes[:, None]
|
||||
)
|
||||
|
||||
return index_pad[index_mask]
|
||||
@@ -61,9 +61,26 @@ class SparseAttentionModule(DynamicModule):
|
||||
Args:
|
||||
attribute_cfg: Sparse attention attribute configuration.
|
||||
"""
|
||||
from .config import VSAAttributeConfig
|
||||
|
||||
# Determine which config class to use based on method
|
||||
config_dict = attribute_cfg or {}
|
||||
if isinstance(attribute_cfg, dict):
|
||||
method = config_dict.get("method", "flash_skip_softmax")
|
||||
elif attribute_cfg is not None and hasattr(attribute_cfg, "method"):
|
||||
method = attribute_cfg.method
|
||||
else:
|
||||
method = "flash_skip_softmax"
|
||||
|
||||
# Select appropriate config class based on method
|
||||
if method == "vsa":
|
||||
config_class = VSAAttributeConfig
|
||||
else:
|
||||
config_class = SparseAttentionAttributeConfig
|
||||
|
||||
# Ensure config is validated through Pydantic
|
||||
if not isinstance(attribute_cfg, SparseAttentionAttributeConfig):
|
||||
attribute_cfg = SparseAttentionAttributeConfig(**(attribute_cfg or {}))
|
||||
if not isinstance(attribute_cfg, (SparseAttentionAttributeConfig, VSAAttributeConfig)):
|
||||
attribute_cfg = config_class(**(config_dict))
|
||||
|
||||
# Store raw config for method initialization
|
||||
self._method_config = {}
|
||||
@@ -80,10 +97,10 @@ class SparseAttentionModule(DynamicModule):
|
||||
|
||||
# Process each attribute from validated config
|
||||
for attribute, val in attribute_cfg.model_dump().items():
|
||||
# Validate attribute if using config class
|
||||
if hasattr(SparseAttentionAttributeConfig, "model_fields"):
|
||||
assert attribute in SparseAttentionAttributeConfig.model_fields, (
|
||||
f"{attribute} is not a valid SparseAttentionModule attribute"
|
||||
# Validate attribute against the appropriate config class
|
||||
if hasattr(config_class, "model_fields"):
|
||||
assert attribute in config_class.model_fields, (
|
||||
f"{attribute} is not a valid {config_class.__name__} attribute"
|
||||
)
|
||||
|
||||
if attribute in _module_attributes:
|
||||
@@ -159,14 +176,28 @@ class SparseAttentionModule(DynamicModule):
|
||||
def forward(self, *args, **kwargs):
|
||||
"""Forward with selected sparse attention method.
|
||||
|
||||
This method dispatches to the appropriate sparse attention implementation
|
||||
based on the configured method and backend.
|
||||
- VSA: patches ``F.scaled_dot_product_attention`` to intercept the SDPA
|
||||
call inside the original forward. Cross-attention is skipped.
|
||||
- Softmax-patching methods (e.g. ``flash_skip_softmax``): use the
|
||||
context manager path below.
|
||||
"""
|
||||
# Pass through if sparse attention is disabled
|
||||
if not self.is_enabled:
|
||||
return super().forward(*args, **kwargs)
|
||||
|
||||
# Get the appropriate context manager for this configuration
|
||||
# VSA: patch F.scaled_dot_product_attention so the VSA kernel intercepts
|
||||
# the SDPA call inside the original forward. This works for diffusers models
|
||||
# since SDPA is the common attention primitive.
|
||||
# Only self-attention is replaced. Cross-attention (Q/K have different seq_len) is skipped.
|
||||
if self._method == "vsa":
|
||||
result = self._forward_with_vsa_sdpa_patch(args, kwargs)
|
||||
|
||||
if self._stats_manager is not None and self._last_stats is not None:
|
||||
self._stats_manager.collect(self._last_stats)
|
||||
self._last_stats = None
|
||||
return result
|
||||
|
||||
# Standard path: softmax patching
|
||||
context = self._get_sparse_context()
|
||||
|
||||
# Apply sparse attention through the context
|
||||
@@ -180,6 +211,61 @@ class SparseAttentionModule(DynamicModule):
|
||||
|
||||
return result
|
||||
|
||||
def _forward_with_vsa_sdpa_patch(self, args, kwargs):
|
||||
"""Run forward with F.scaled_dot_product_attention patched for VSA.
|
||||
|
||||
Replaces SDPA with the VSA kernel for self-attention calls (Q and K/V
|
||||
have the same seq_len). Cross-attention calls fall through to the
|
||||
original SDPA. Warns if SDPA was never called.
|
||||
"""
|
||||
import torch.nn.functional as F
|
||||
|
||||
from modelopt.torch.quantization.utils import replace_function
|
||||
|
||||
vsa = self._sparse_method_instance
|
||||
original_sdpa = F.scaled_dot_product_attention
|
||||
self._vsa_sdpa_called = False
|
||||
|
||||
def _patched_sdpa(query, key, value, attn_mask=None, dropout_p=0.0, is_causal=False, **kw):
|
||||
self._vsa_sdpa_called = True
|
||||
|
||||
# Fall back to original SDPA when VSA cannot handle this call:
|
||||
# - Cross-attention: Q and K/V have different seq_len
|
||||
# - video_shape not set: VSA cannot compute tile metadata
|
||||
# - seq_len mismatch: input doesn't match the configured video shape
|
||||
can_apply_vsa = (
|
||||
vsa.video_shape is not None
|
||||
and query.shape[2] == key.shape[2]
|
||||
and query.shape[2] == vsa.video_shape[0] * vsa.video_shape[1] * vsa.video_shape[2]
|
||||
)
|
||||
if not can_apply_vsa:
|
||||
return original_sdpa(
|
||||
query,
|
||||
key,
|
||||
value,
|
||||
attn_mask=attn_mask,
|
||||
dropout_p=dropout_p,
|
||||
is_causal=is_causal,
|
||||
**kw,
|
||||
)
|
||||
output, stats = vsa.forward_attention(query, key, value)
|
||||
self._last_stats = stats
|
||||
return output
|
||||
|
||||
with replace_function(F, "scaled_dot_product_attention", _patched_sdpa):
|
||||
result = super().forward(*args, **kwargs)
|
||||
|
||||
if not self._vsa_sdpa_called:
|
||||
import warnings
|
||||
|
||||
warnings.warn(
|
||||
f"VSA: F.scaled_dot_product_attention was not called during "
|
||||
f"{type(self).__name__}.forward(). The attention layer may use a "
|
||||
f"custom kernel that bypasses SDPA. VSA had no effect on this layer.",
|
||||
)
|
||||
|
||||
return result
|
||||
|
||||
def _get_sparse_context(self):
|
||||
"""Get the context manager for applying sparse attention.
|
||||
|
||||
|
||||
@@ -0,0 +1,446 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
"""CPU-only unit tests for Video Sparse Attention (VSA).
|
||||
|
||||
Tests cover:
|
||||
- vsa_utils.py: tile/untile index logic, variable block sizes
|
||||
- vsa.py: VSA method init, metadata computation, validation, caching, forward_attention
|
||||
- config.py: VSAAttributeConfig validation
|
||||
- HF integration: registration, sparsify, forward dispatch
|
||||
"""
|
||||
|
||||
import math
|
||||
import sys
|
||||
from unittest.mock import patch
|
||||
|
||||
import pytest
|
||||
|
||||
# The attention_sparsity package transitively imports the HF plugin, which
|
||||
# requires transformers. Skip the entire module when it is not installed.
|
||||
pytest.importorskip("transformers")
|
||||
|
||||
import torch
|
||||
from pydantic import ValidationError
|
||||
|
||||
from modelopt.torch.sparsity.attention_sparsity.config import VSAAttributeConfig, VSAConfig
|
||||
from modelopt.torch.sparsity.attention_sparsity.methods.vsa import VSA
|
||||
from modelopt.torch.sparsity.attention_sparsity.methods.vsa_utils import (
|
||||
construct_variable_block_sizes,
|
||||
get_non_pad_index,
|
||||
get_reverse_tile_partition_indices,
|
||||
get_tile_partition_indices,
|
||||
)
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# vsa_utils: tile partition indices
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestTilePartitionIndices:
|
||||
"""Tests for get_tile_partition_indices."""
|
||||
|
||||
def test_evenly_divisible(self):
|
||||
"""Tiles cover full volume with no remainder."""
|
||||
video_shape = (8, 8, 8)
|
||||
tile_size = (4, 4, 4)
|
||||
idx = get_tile_partition_indices(video_shape, tile_size, torch.device("cpu"))
|
||||
assert idx.shape == (8 * 8 * 8,)
|
||||
# Every original index appears exactly once
|
||||
assert torch.equal(idx.sort().values, torch.arange(512))
|
||||
|
||||
def test_non_divisible(self):
|
||||
"""Edge tiles are smaller when dims don't divide evenly."""
|
||||
video_shape = (5, 6, 7)
|
||||
tile_size = (4, 4, 4)
|
||||
seq_len = 5 * 6 * 7
|
||||
idx = get_tile_partition_indices(video_shape, tile_size, torch.device("cpu"))
|
||||
assert idx.shape == (seq_len,)
|
||||
assert torch.equal(idx.sort().values, torch.arange(seq_len))
|
||||
|
||||
def test_round_trip(self):
|
||||
"""tile then reverse_tile is identity."""
|
||||
video_shape = (6, 10, 8)
|
||||
tile_size = (4, 4, 4)
|
||||
device = torch.device("cpu")
|
||||
fwd = get_tile_partition_indices(video_shape, tile_size, device)
|
||||
rev = get_reverse_tile_partition_indices(video_shape, tile_size, device)
|
||||
# Applying forward then reverse should yield the original order
|
||||
assert torch.equal(fwd[rev], torch.arange(6 * 10 * 8))
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# vsa_utils: variable block sizes
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestVariableBlockSizes:
|
||||
"""Tests for construct_variable_block_sizes."""
|
||||
|
||||
def test_evenly_divisible(self):
|
||||
"""All tiles have full size when dims divide evenly."""
|
||||
video_shape = (8, 8, 8)
|
||||
tile_size = (4, 4, 4)
|
||||
num_tiles = (2, 2, 2)
|
||||
sizes = construct_variable_block_sizes(
|
||||
video_shape, num_tiles, tile_size, torch.device("cpu")
|
||||
)
|
||||
assert sizes.shape == (8,) # 2*2*2 tiles
|
||||
assert (sizes == 64).all() # every tile is full 4*4*4
|
||||
|
||||
def test_non_divisible_sum(self):
|
||||
"""Sum of variable sizes equals original sequence length."""
|
||||
video_shape = (5, 6, 7)
|
||||
tile_size = (4, 4, 4)
|
||||
num_tiles = (
|
||||
math.ceil(5 / 4),
|
||||
math.ceil(6 / 4),
|
||||
math.ceil(7 / 4),
|
||||
)
|
||||
sizes = construct_variable_block_sizes(
|
||||
video_shape, num_tiles, tile_size, torch.device("cpu")
|
||||
)
|
||||
assert sizes.sum().item() == 5 * 6 * 7
|
||||
|
||||
def test_partial_tile_smaller(self):
|
||||
"""Last tile along a non-divisible dim should be smaller."""
|
||||
video_shape = (5, 4, 4)
|
||||
tile_size = (4, 4, 4)
|
||||
num_tiles = (2, 1, 1)
|
||||
sizes = construct_variable_block_sizes(
|
||||
video_shape, num_tiles, tile_size, torch.device("cpu")
|
||||
)
|
||||
# First tile: 4*4*4=64, second tile: 1*4*4=16
|
||||
assert sizes[0].item() == 64
|
||||
assert sizes[1].item() == 16
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# vsa_utils: non-pad index
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestNonPadIndex:
|
||||
"""Tests for get_non_pad_index."""
|
||||
|
||||
def test_full_blocks(self):
|
||||
"""All blocks full size -> non_pad covers everything."""
|
||||
sizes = torch.tensor([64, 64, 64])
|
||||
npi = get_non_pad_index(sizes, 64)
|
||||
assert npi.shape == (192,) # 3 * 64
|
||||
|
||||
def test_partial_blocks(self):
|
||||
"""Partial blocks -> non_pad skips padding positions."""
|
||||
sizes = torch.tensor([64, 16])
|
||||
npi = get_non_pad_index(sizes, 64)
|
||||
assert npi.shape == (80,) # 64 + 16
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# VSA: tile/untile round-trip
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestTileUntileRoundTrip:
|
||||
"""Test _tile_tensor / _untile_tensor preserve data."""
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"video_shape",
|
||||
[(8, 8, 8), (5, 6, 7), (4, 4, 4)],
|
||||
ids=["even", "non-divisible", "single-tile"],
|
||||
)
|
||||
def test_round_trip(self, video_shape):
|
||||
"""tile then untile recovers the original tensor."""
|
||||
seq_len = video_shape[0] * video_shape[1] * video_shape[2]
|
||||
vsa = VSA({"video_shape": video_shape})
|
||||
meta = vsa._compute_metadata(seq_len, torch.device("cpu"))
|
||||
|
||||
x = torch.randn(2, 4, seq_len, 16) # [batch, heads, seq, dim]
|
||||
tiled = vsa._tile_tensor(x, meta)
|
||||
recovered = vsa._untile_tensor(tiled, meta, seq_len)
|
||||
|
||||
assert recovered.shape == x.shape
|
||||
assert torch.allclose(recovered, x)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# VSA method: init and config
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestVSAInit:
|
||||
"""Tests for VSA.__init__ and basic properties."""
|
||||
|
||||
def test_defaults(self):
|
||||
vsa = VSA()
|
||||
assert vsa.block_size_3d == (4, 4, 4)
|
||||
assert vsa.block_elements == 64
|
||||
assert vsa.top_k_ratio == 0.5
|
||||
assert vsa.video_shape is None
|
||||
assert vsa.name == "vsa"
|
||||
|
||||
def test_custom_config(self):
|
||||
vsa = VSA({"block_size_3d": [2, 2, 2], "top_k_ratio": 0.3, "video_shape": (8, 8, 8)})
|
||||
assert vsa.block_size_3d == (2, 2, 2)
|
||||
assert vsa.block_elements == 8
|
||||
assert vsa.top_k_ratio == 0.3
|
||||
assert vsa.video_shape == (8, 8, 8)
|
||||
|
||||
def test_set_video_shape(self):
|
||||
vsa = VSA()
|
||||
vsa.set_video_shape((4, 8, 12))
|
||||
assert vsa.video_shape == (4, 8, 12)
|
||||
|
||||
def test_get_threshold_info(self):
|
||||
vsa = VSA({"top_k_ratio": 0.7, "video_shape": (4, 4, 4)})
|
||||
info = vsa.get_threshold_info()
|
||||
assert info["type"] == "vsa"
|
||||
assert info["top_k_ratio"] == 0.7
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# VSA method: metadata computation and validation
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestVSAMetadata:
|
||||
"""Tests for VSA._compute_metadata validation and caching."""
|
||||
|
||||
def test_no_video_shape_raises(self):
|
||||
vsa = VSA()
|
||||
with pytest.raises(ValueError, match="video_shape must be provided"):
|
||||
vsa._compute_metadata(100, torch.device("cpu"))
|
||||
|
||||
def test_seq_len_mismatch_raises(self):
|
||||
vsa = VSA({"video_shape": (4, 4, 4)})
|
||||
with pytest.raises(ValueError, match="does not match video shape"):
|
||||
vsa._compute_metadata(100, torch.device("cpu")) # expected 64
|
||||
|
||||
def test_valid_metadata(self):
|
||||
vsa = VSA({"video_shape": (8, 8, 8)})
|
||||
meta = vsa._compute_metadata(512, torch.device("cpu"))
|
||||
assert meta["video_shape"] == (8, 8, 8)
|
||||
assert meta["num_tiles"] == (2, 2, 2)
|
||||
assert meta["total_tiles"] == 8
|
||||
|
||||
def test_metadata_caching(self):
|
||||
vsa = VSA({"video_shape": (8, 8, 8)})
|
||||
m1 = vsa._compute_metadata(512, torch.device("cpu"))
|
||||
m2 = vsa._compute_metadata(512, torch.device("cpu"))
|
||||
assert m1 is m2 # same object, not recomputed
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# VSA: forward_attention (kernel import guard)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestVSAForwardAttention:
|
||||
"""Tests for VSA.forward_attention."""
|
||||
|
||||
def test_missing_kernel_raises(self):
|
||||
"""forward_attention raises ImportError when fastvideo_kernel is missing."""
|
||||
vsa = VSA({"video_shape": (4, 4, 4), "top_k_ratio": 0.5})
|
||||
seq_len = 4 * 4 * 4
|
||||
q = torch.randn(1, 2, seq_len, 16)
|
||||
k = torch.randn(1, 2, seq_len, 16)
|
||||
v = torch.randn(1, 2, seq_len, 16)
|
||||
with (
|
||||
patch.dict(sys.modules, {"fastvideo_kernel": None}),
|
||||
pytest.raises(ImportError, match="fastvideo_kernel"),
|
||||
):
|
||||
vsa.forward_attention(q, k, v)
|
||||
|
||||
def test_video_shape_override(self):
|
||||
"""forward_attention accepts video_shape kwarg to override instance shape."""
|
||||
vsa = VSA({"video_shape": (4, 4, 4), "top_k_ratio": 0.5})
|
||||
new_shape = (8, 8, 8)
|
||||
seq_len = 8 * 8 * 8
|
||||
q = torch.randn(1, 2, seq_len, 16)
|
||||
with (
|
||||
patch.dict(sys.modules, {"fastvideo_kernel": None}),
|
||||
pytest.raises(ImportError),
|
||||
):
|
||||
vsa.forward_attention(q, q, q, video_shape=new_shape)
|
||||
assert vsa.video_shape == new_shape
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# VSAAttributeConfig validation
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestVSAAttributeConfig:
|
||||
"""Tests for VSAAttributeConfig pydantic validation."""
|
||||
|
||||
def test_valid_defaults(self):
|
||||
cfg = VSAAttributeConfig()
|
||||
assert cfg.method == "vsa"
|
||||
assert cfg.block_size_3d == (4, 4, 4)
|
||||
assert cfg.top_k_ratio == 0.5
|
||||
|
||||
def test_top_k_ratio_out_of_range(self):
|
||||
with pytest.raises(ValidationError, match="top_k_ratio"):
|
||||
VSAAttributeConfig(top_k_ratio=0.0)
|
||||
with pytest.raises(ValidationError, match="top_k_ratio"):
|
||||
VSAAttributeConfig(top_k_ratio=1.5)
|
||||
|
||||
def test_video_shape_wrong_length(self):
|
||||
with pytest.raises(ValidationError, match="3 elements"):
|
||||
VSAAttributeConfig(video_shape=(4, 4))
|
||||
|
||||
def test_video_shape_negative(self):
|
||||
with pytest.raises(ValidationError, match="positive"):
|
||||
VSAAttributeConfig(video_shape=(4, -1, 4))
|
||||
|
||||
def test_video_shape_none_allowed(self):
|
||||
cfg = VSAAttributeConfig(video_shape=None)
|
||||
assert cfg.video_shape is None
|
||||
|
||||
def test_vsa_config_defaults(self):
|
||||
cfg = VSAConfig()
|
||||
assert "*attn*" in cfg.sparse_cfg
|
||||
assert cfg.sparse_cfg["*attn*"]["method"] == "vsa"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# ModelOpt integration: sparsify() with VSA config
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
from _test_utils.torch.sparsity.sparse_attention_common import SimpleAttentionModel
|
||||
|
||||
import modelopt.torch.opt as mto
|
||||
import modelopt.torch.sparsity.attention_sparsity as sparse_attn
|
||||
from modelopt.torch.sparsity.attention_sparsity.sparse_attention import SparseAttentionModule
|
||||
|
||||
VSA_TEST_CFG = {
|
||||
"sparse_cfg": {
|
||||
"*attention*": {
|
||||
"method": "vsa",
|
||||
"block_size_3d": (4, 4, 4),
|
||||
"top_k_ratio": 0.5,
|
||||
"enable": True,
|
||||
},
|
||||
"default": {"enable": False},
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
class TestVSASparsifyIntegration:
|
||||
"""Test VSA integration with modelopt sparsify() API."""
|
||||
|
||||
def test_sparsify_creates_sparse_modules(self):
|
||||
"""sparsify() with VSA config replaces attention modules."""
|
||||
model = SimpleAttentionModel()
|
||||
sparse_model = sparse_attn.sparsify(model, VSA_TEST_CFG)
|
||||
|
||||
sparse_modules = [m for m in sparse_model.modules() if isinstance(m, SparseAttentionModule)]
|
||||
assert len(sparse_modules) > 0
|
||||
|
||||
def test_sparse_module_has_vsa_method(self):
|
||||
"""Replaced modules are configured with VSA method."""
|
||||
model = SimpleAttentionModel()
|
||||
sparse_model = sparse_attn.sparsify(model, VSA_TEST_CFG)
|
||||
|
||||
for module in sparse_model.modules():
|
||||
if isinstance(module, SparseAttentionModule):
|
||||
assert module._method == "vsa"
|
||||
assert isinstance(module._sparse_method_instance, VSA)
|
||||
assert module._sparse_method_instance.block_size_3d == (4, 4, 4)
|
||||
assert module._sparse_method_instance.top_k_ratio == 0.5
|
||||
|
||||
def test_enable_disable(self):
|
||||
"""Enable/disable works on VSA sparse modules."""
|
||||
model = SimpleAttentionModel()
|
||||
sparse_model = sparse_attn.sparsify(model, VSA_TEST_CFG)
|
||||
|
||||
for module in sparse_model.modules():
|
||||
if isinstance(module, SparseAttentionModule):
|
||||
assert module.is_enabled
|
||||
module.disable()
|
||||
assert not module.is_enabled
|
||||
module.enable()
|
||||
assert module.is_enabled
|
||||
|
||||
def test_threshold_info(self):
|
||||
"""VSA sparse modules report correct threshold info."""
|
||||
model = SimpleAttentionModel()
|
||||
sparse_model = sparse_attn.sparsify(model, VSA_TEST_CFG)
|
||||
|
||||
for module in sparse_model.modules():
|
||||
if isinstance(module, SparseAttentionModule):
|
||||
info = module.get_threshold_info()
|
||||
assert info["type"] == "vsa"
|
||||
assert info["top_k_ratio"] == 0.5
|
||||
|
||||
def test_save_restore(self):
|
||||
"""VSA modelopt_state can be saved and restored."""
|
||||
model = SimpleAttentionModel()
|
||||
sparse_model = sparse_attn.sparsify(model, VSA_TEST_CFG)
|
||||
|
||||
state = mto.modelopt_state(sparse_model)
|
||||
|
||||
# Restore to a fresh model
|
||||
model_restored = SimpleAttentionModel()
|
||||
mto.restore_from_modelopt_state(model_restored, state)
|
||||
|
||||
# Verify VSA method is restored
|
||||
for module in model_restored.modules():
|
||||
if isinstance(module, SparseAttentionModule):
|
||||
assert module._method == "vsa"
|
||||
assert isinstance(module._sparse_method_instance, VSA)
|
||||
|
||||
def test_pattern_matching(self):
|
||||
"""Pattern-based config selectively applies VSA."""
|
||||
model = SimpleAttentionModel()
|
||||
|
||||
# Pattern that won't match anything
|
||||
config = {
|
||||
"sparse_cfg": {
|
||||
"*nonexistent*": {
|
||||
"method": "vsa",
|
||||
"enable": True,
|
||||
},
|
||||
"default": {"enable": False},
|
||||
},
|
||||
}
|
||||
sparse_model = sparse_attn.sparsify(model, config)
|
||||
|
||||
# No modules should have VSA enabled
|
||||
for module in sparse_model.modules():
|
||||
if isinstance(module, SparseAttentionModule):
|
||||
assert not module.is_enabled
|
||||
|
||||
def test_forward_patches_sdpa(self):
|
||||
"""VSA patches F.scaled_dot_product_attention during forward.
|
||||
|
||||
SimpleAttentionModel uses nn.MultiheadAttention which calls SDPA.
|
||||
VSA intercepts the SDPA call. Without fastvideo_kernel, this raises
|
||||
ImportError — proving the interception works.
|
||||
"""
|
||||
model = SimpleAttentionModel()
|
||||
sparse_model = sparse_attn.sparsify(model, VSA_TEST_CFG)
|
||||
|
||||
# Set video_shape so metadata can be computed.
|
||||
# seq_len=64, video_shape (4,4,4) -> T*H*W=64
|
||||
for module in sparse_model.modules():
|
||||
if isinstance(module, SparseAttentionModule) and module.is_enabled:
|
||||
module._sparse_method_instance.set_video_shape((4, 4, 4))
|
||||
|
||||
with (
|
||||
patch.dict(sys.modules, {"fastvideo_kernel": None}),
|
||||
pytest.raises(ImportError, match="fastvideo_kernel"),
|
||||
):
|
||||
sparse_model(torch.randn(1, 64, 256))
|
||||
Reference in New Issue
Block a user