[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:
kaix-nv
2026-04-03 22:17:48 +00:00
committed by GitHub
parent 18ce04f1ce
commit df80a0f7b2
11 changed files with 1248 additions and 26 deletions
+1
View File
@@ -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
+1
View File
@@ -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.
+2 -3
View File
@@ -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))