[diffusion] feat: support SparseVideoGen2 attention backend (#17507)
Co-authored-by: Mick <mickjagger19@icloud.com>
This commit is contained in:
@@ -16,23 +16,23 @@ default parameters when initializing and generating videos.
|
|||||||
|
|
||||||
### Video Generation Models
|
### Video Generation Models
|
||||||
|
|
||||||
| Model Name | Hugging Face Model ID | Resolutions | TeaCache | Sliding Tile Attn | Sage Attn | Video Sparse Attention (VSA) | Sparse Linear Attention (SLA) | Sage Sparse Linear Attention (SageSLA) |
|
| Model Name | Hugging Face Model ID | Resolutions | TeaCache | Sliding Tile Attn | Sage Attn | Video Sparse Attention (VSA) | Sparse Linear Attention (SLA) | Sage Sparse Linear Attention (SageSLA) | Sparse Video Gen 2 (SVG2) |
|
||||||
|:-----------------------------|:--------------------------------------------------|:--------------------|:--------:|:-----------------:|:---------:|:----------------------------:|:----------------------------:|:-----------------------------------------------:|
|
|:-----------------------------|:--------------------------------------------------|:--------------------|:--------:|:-----------------:|:---------:|:----------------------------:|:----------------------------:|:-----------------------------------------------:|:----------------------------------:|
|
||||||
| FastWan2.1 T2V 1.3B | `FastVideo/FastWan2.1-T2V-1.3B-Diffusers` | 480p | ⭕ | ⭕ | ⭕ | ✅ | ❌ | ❌ |
|
| FastWan2.1 T2V 1.3B | `FastVideo/FastWan2.1-T2V-1.3B-Diffusers` | 480p | ⭕ | ⭕ | ⭕ | ✅ | ❌ | ❌ | ❌ |
|
||||||
| FastWan2.2 TI2V 5B Full Attn | `FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers` | 720p | ⭕ | ⭕ | ⭕ | ✅ | ❌ | ❌ |
|
| FastWan2.2 TI2V 5B Full Attn | `FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers` | 720p | ⭕ | ⭕ | ⭕ | ✅ | ❌ | ❌ | ❌ |
|
||||||
| Wan2.2 TI2V 5B | `Wan-AI/Wan2.2-TI2V-5B-Diffusers` | 720p | ⭕ | ⭕ | ✅ | ⭕ | ❌ | ❌ |
|
| Wan2.2 TI2V 5B | `Wan-AI/Wan2.2-TI2V-5B-Diffusers` | 720p | ⭕ | ⭕ | ✅ | ⭕ | ❌ | ❌ | ❌ |
|
||||||
| Wan2.2 T2V A14B | `Wan-AI/Wan2.2-T2V-A14B-Diffusers` | 480p<br>720p | ❌ | ❌ | ✅ | ⭕ | ❌ | ❌ |
|
| Wan2.2 T2V A14B | `Wan-AI/Wan2.2-T2V-A14B-Diffusers` | 480p<br>720p | ❌ | ❌ | ✅ | ⭕ | ❌ | ❌ | ❌ |
|
||||||
| Wan2.2 I2V A14B | `Wan-AI/Wan2.2-I2V-A14B-Diffusers` | 480p<br>720p | ❌ | ❌ | ✅ | ⭕ | ❌ | ❌ |
|
| Wan2.2 I2V A14B | `Wan-AI/Wan2.2-I2V-A14B-Diffusers` | 480p<br>720p | ❌ | ❌ | ✅ | ⭕ | ❌ | ❌ | ❌ |
|
||||||
| HunyuanVideo | `hunyuanvideo-community/HunyuanVideo` | 720×1280<br>544×960 | ❌ | ✅ | ✅ | ⭕ | ❌ | ❌ |
|
| HunyuanVideo | `hunyuanvideo-community/HunyuanVideo` | 720×1280<br>544×960 | ❌ | ✅ | ✅ | ⭕ | ❌ | ❌ | ✅ |
|
||||||
| FastHunyuan | `FastVideo/FastHunyuan-diffusers` | 720×1280<br>544×960 | ❌ | ✅ | ✅ | ⭕ | ❌ | ❌ |
|
| FastHunyuan | `FastVideo/FastHunyuan-diffusers` | 720×1280<br>544×960 | ❌ | ✅ | ✅ | ⭕ | ❌ | ❌ | ✅ |
|
||||||
| Wan2.1 T2V 1.3B | `Wan-AI/Wan2.1-T2V-1.3B-Diffusers` | 480p | ✅ | ✅ | ✅ | ⭕ | ❌ | ❌ |
|
| Wan2.1 T2V 1.3B | `Wan-AI/Wan2.1-T2V-1.3B-Diffusers` | 480p | ✅ | ✅ | ✅ | ⭕ | ❌ | ❌ | ✅ |
|
||||||
| Wan2.1 T2V 14B | `Wan-AI/Wan2.1-T2V-14B-Diffusers` | 480p, 720p | ✅ | ✅ | ✅ | ⭕ | ❌ | ❌ |
|
| Wan2.1 T2V 14B | `Wan-AI/Wan2.1-T2V-14B-Diffusers` | 480p, 720p | ✅ | ✅ | ✅ | ⭕ | ❌ | ❌ | ✅ |
|
||||||
| Wan2.1 I2V 480P | `Wan-AI/Wan2.1-I2V-14B-480P-Diffusers` | 480p | ✅ | ✅ | ✅ | ⭕ | ❌ | ❌ |
|
| Wan2.1 I2V 480P | `Wan-AI/Wan2.1-I2V-14B-480P-Diffusers` | 480p | ✅ | ✅ | ✅ | ⭕ | ❌ | ❌ | ✅ |
|
||||||
| Wan2.1 I2V 720P | `Wan-AI/Wan2.1-I2V-14B-720P-Diffusers` | 720p | ✅ | ✅ | ✅ | ⭕ | ❌ | ❌ |
|
| Wan2.1 I2V 720P | `Wan-AI/Wan2.1-I2V-14B-720P-Diffusers` | 720p | ✅ | ✅ | ✅ | ⭕ | ❌ | ❌ | ✅ |
|
||||||
| TurboWan2.1 T2V 1.3B | `IPostYellow/TurboWan2.1-T2V-1.3B-Diffusers` | 480p | ✅ | ❌ | ❌ | ❌ | ✅ | ✅ |
|
| TurboWan2.1 T2V 1.3B | `IPostYellow/TurboWan2.1-T2V-1.3B-Diffusers` | 480p | ✅ | ❌ | ❌ | ❌ | ✅ | ✅ | ⭕ |
|
||||||
| TurboWan2.1 T2V 14B | `IPostYellow/TurboWan2.1-T2V-14B-Diffusers` | 480p | ✅ | ❌ | ❌ | ❌ | ✅ | ✅ |
|
| TurboWan2.1 T2V 14B | `IPostYellow/TurboWan2.1-T2V-14B-Diffusers` | 480p | ✅ | ❌ | ❌ | ❌ | ✅ | ✅ | ⭕ |
|
||||||
| TurboWan2.1 T2V 14B 720P | `IPostYellow/TurboWan2.1-T2V-14B-720P-Diffusers` | 720p | ✅ | ❌ | ❌ | ❌ | ✅ | ✅ |
|
| TurboWan2.1 T2V 14B 720P | `IPostYellow/TurboWan2.1-T2V-14B-720P-Diffusers` | 720p | ✅ | ❌ | ❌ | ❌ | ✅ | ✅ | ⭕ |
|
||||||
| TurboWan2.2 I2V A14B | `IPostYellow/TurboWan2.2-I2V-A14B-Diffusers` | 720p | ✅ | ❌ | ❌ | ❌ | ✅ | ✅ |
|
| TurboWan2.2 I2V A14B | `IPostYellow/TurboWan2.2-I2V-A14B-Diffusers` | 720p | ✅ | ❌ | ❌ | ❌ | ✅ | ✅ | ⭕ |
|
||||||
|
|
||||||
**Note**:
|
**Note**:
|
||||||
1.Wan2.2 TI2V 5B has some quality issues when performing I2V generation. We are working on fixing this issue.
|
1.Wan2.2 TI2V 5B has some quality issues when performing I2V generation. We are working on fixing this issue.
|
||||||
|
|||||||
@@ -29,6 +29,7 @@ For SGLang-native pipelines, the CLI accepts the lowercase names of `AttentionBa
|
|||||||
| `video_sparse_attn` | `VIDEO_SPARSE_ATTN` | Requires `vsa`. Configure `sparsity` via `--attention-backend-config`. |
|
| `video_sparse_attn` | `VIDEO_SPARSE_ATTN` | Requires `vsa`. Configure `sparsity` via `--attention-backend-config`. |
|
||||||
| `vmoba_attn` | `VMOBA_ATTN` | Requires `kernel.attn.vmoba_attn.vmoba`. Configure via `--attention-backend-config`. |
|
| `vmoba_attn` | `VMOBA_ATTN` | Requires `kernel.attn.vmoba_attn.vmoba`. Configure via `--attention-backend-config`. |
|
||||||
| `aiter` | `AITER` | Requires `aiter`. |
|
| `aiter` | `AITER` | Requires `aiter`. |
|
||||||
|
| `sparse_video_gen_2_attn` | `SPARSE_VIDEO_GEN_2_ATTN` | Requires `svg`. See installation instructions at https://github.com/svg-project/Sparse-VideoGen. |
|
||||||
|
|
||||||
## Selection priority
|
## Selection priority
|
||||||
|
|
||||||
@@ -92,6 +93,7 @@ Some backends require additional configuration. You can pass these parameters vi
|
|||||||
| `video_sparse_attn` | ✅ | ❌ | ❌ | CUDA-only. Requires `vsa`. Configure `sparsity` via `--attention-backend-config`. |
|
| `video_sparse_attn` | ✅ | ❌ | ❌ | CUDA-only. Requires `vsa`. Configure `sparsity` via `--attention-backend-config`. |
|
||||||
| `vmoba_attn` | ✅ | ❌ | ❌ | CUDA-only. Requires `kernel.attn.vmoba_attn.vmoba`. Configure via `--attention-backend-config`. |
|
| `vmoba_attn` | ✅ | ❌ | ❌ | CUDA-only. Requires `kernel.attn.vmoba_attn.vmoba`. Configure via `--attention-backend-config`. |
|
||||||
| `aiter` | ✅ | ❌ | ❌ | Requires `aiter`. |
|
| `aiter` | ✅ | ❌ | ❌ | Requires `aiter`. |
|
||||||
|
| `sparse_video_gen_2_attn` | ✅ | ❌ | ❌ | CUDA-only. Requires `svg`. |
|
||||||
|
|
||||||
## Usage
|
## Usage
|
||||||
|
|
||||||
|
|||||||
@@ -31,6 +31,7 @@ class DiTArchConfig(ArchConfig):
|
|||||||
AttentionBackendEnum.AITER,
|
AttentionBackendEnum.AITER,
|
||||||
AttentionBackendEnum.TORCH_SDPA,
|
AttentionBackendEnum.TORCH_SDPA,
|
||||||
AttentionBackendEnum.VIDEO_SPARSE_ATTN,
|
AttentionBackendEnum.VIDEO_SPARSE_ATTN,
|
||||||
|
AttentionBackendEnum.SPARSE_VIDEO_GEN_2_ATTN,
|
||||||
AttentionBackendEnum.VMOBA_ATTN,
|
AttentionBackendEnum.VMOBA_ATTN,
|
||||||
AttentionBackendEnum.SAGE_ATTN_3,
|
AttentionBackendEnum.SAGE_ATTN_3,
|
||||||
}
|
}
|
||||||
|
|||||||
+562
@@ -0,0 +1,562 @@
|
|||||||
|
"""
|
||||||
|
Sparse Video Gen 2 (SAP) attention backend.
|
||||||
|
|
||||||
|
This is a baseline integration that wires the backend into the
|
||||||
|
attention framework.
|
||||||
|
|
||||||
|
Adapted from https://github.com/svg-project/Sparse-VideoGen/blob/main/svg/models/wan/attention.py
|
||||||
|
"""
|
||||||
|
|
||||||
|
from dataclasses import dataclass, field
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
import torch
|
||||||
|
import torch.nn.functional as F
|
||||||
|
from torch.nn.attention import SDPBackend, sdpa_kernel
|
||||||
|
|
||||||
|
try:
|
||||||
|
from svg.kernels.triton.permute import (
|
||||||
|
apply_inverse_permutation_triton,
|
||||||
|
permute_tensor_by_labels_triton,
|
||||||
|
)
|
||||||
|
from svg.kmeans_utils import (
|
||||||
|
batch_kmeans_Euclid,
|
||||||
|
dynamic_block_sparse_fwd_flashinfer,
|
||||||
|
identify_dynamic_map,
|
||||||
|
)
|
||||||
|
|
||||||
|
svg2_available = True
|
||||||
|
except ImportError:
|
||||||
|
svg2_available = False
|
||||||
|
|
||||||
|
from sglang.multimodal_gen.runtime.layers.attention.backends.attention_backend import (
|
||||||
|
AttentionBackend,
|
||||||
|
AttentionImpl,
|
||||||
|
AttentionMetadata,
|
||||||
|
AttentionMetadataBuilder,
|
||||||
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.platforms import AttentionBackendEnum
|
||||||
|
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||||
|
|
||||||
|
logger = init_logger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
class SparseVideoGen2AttentionBackend(AttentionBackend):
|
||||||
|
|
||||||
|
accept_output_buffer: bool = True
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def get_supported_head_sizes() -> list[int]:
|
||||||
|
return [64, 128, 256]
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def get_enum() -> AttentionBackendEnum:
|
||||||
|
return AttentionBackendEnum.SPARSE_VIDEO_GEN_2_ATTN
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def get_impl_cls() -> type["SparseVideoGen2AttentionImpl"]:
|
||||||
|
return SparseVideoGen2AttentionImpl
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def get_metadata_cls() -> type["SparseVideoGen2AttentionMetadata"]:
|
||||||
|
return SparseVideoGen2AttentionMetadata
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def get_builder_cls() -> type["SparseVideoGen2AttentionMetadataBuilder"]:
|
||||||
|
return SparseVideoGen2AttentionMetadataBuilder
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class Svg2LayerCache:
|
||||||
|
# centroids for kmeans clustering
|
||||||
|
q_centroids: torch.Tensor | None = None
|
||||||
|
k_centroids: torch.Tensor | None = None
|
||||||
|
centroids_initialized: bool = False
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class Svg2Cache:
|
||||||
|
layers: dict[int, Svg2LayerCache] = field(default_factory=dict)
|
||||||
|
|
||||||
|
def get_layer(self, layer_idx: int) -> Svg2LayerCache:
|
||||||
|
layer_cache = self.layers.get(layer_idx)
|
||||||
|
if layer_cache is None:
|
||||||
|
layer_cache = Svg2LayerCache()
|
||||||
|
self.layers[layer_idx] = layer_cache
|
||||||
|
return layer_cache
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class SparseVideoGen2AttentionMetadata(AttentionMetadata):
|
||||||
|
current_timestep: int
|
||||||
|
num_q_centroids: int
|
||||||
|
num_k_centroids: int
|
||||||
|
top_p_kmeans: float
|
||||||
|
min_kc_ratio: float
|
||||||
|
kmeans_iter_init: int
|
||||||
|
kmeans_iter_step: int
|
||||||
|
zero_step_kmeans_init: bool
|
||||||
|
first_layers_fp: float
|
||||||
|
first_times_fp: float
|
||||||
|
context_length: int
|
||||||
|
num_frame: int
|
||||||
|
frame_size: int
|
||||||
|
cache: Svg2Cache
|
||||||
|
prompt_length: int | None = None
|
||||||
|
max_seqlen_q: int | None = None
|
||||||
|
max_seqlen_k: int | None = None
|
||||||
|
|
||||||
|
|
||||||
|
def _require_kwarg(kwargs: dict[str, Any], name: str) -> Any:
|
||||||
|
if name not in kwargs:
|
||||||
|
raise ValueError(
|
||||||
|
f"Missing required argument for SparseVideoGen2Attention: {name}"
|
||||||
|
)
|
||||||
|
return kwargs[name]
|
||||||
|
|
||||||
|
|
||||||
|
class SparseVideoGen2AttentionMetadataBuilder(AttentionMetadataBuilder):
|
||||||
|
|
||||||
|
def __init__(self) -> None:
|
||||||
|
pass
|
||||||
|
|
||||||
|
def prepare(self) -> None:
|
||||||
|
pass
|
||||||
|
|
||||||
|
def build( # type: ignore[override]
|
||||||
|
self,
|
||||||
|
current_timestep: int,
|
||||||
|
raw_latent_shape: tuple[int, ...],
|
||||||
|
patch_size: tuple[int, int, int],
|
||||||
|
cache: Svg2Cache,
|
||||||
|
num_q_centroids: int,
|
||||||
|
num_k_centroids: int,
|
||||||
|
top_p_kmeans: float,
|
||||||
|
min_kc_ratio: float,
|
||||||
|
kmeans_iter_init: int,
|
||||||
|
kmeans_iter_step: int,
|
||||||
|
zero_step_kmeans_init: bool,
|
||||||
|
first_layers_fp: float,
|
||||||
|
first_times_fp: float,
|
||||||
|
context_length: int = 0,
|
||||||
|
prompt_length: int | None = None,
|
||||||
|
**kwargs: dict[str, Any],
|
||||||
|
) -> SparseVideoGen2AttentionMetadata:
|
||||||
|
raw_shape = tuple(raw_latent_shape)
|
||||||
|
if len(raw_shape) == 5:
|
||||||
|
t, h, w = raw_shape[2:5]
|
||||||
|
elif len(raw_shape) == 3:
|
||||||
|
t, h, w = raw_shape
|
||||||
|
else:
|
||||||
|
raise ValueError(
|
||||||
|
"raw_latent_shape must be (T, H, W) or (B, C, T, H, W) for SAP attention"
|
||||||
|
)
|
||||||
|
pt, ph, pw = patch_size
|
||||||
|
if t % pt != 0 or h % ph != 0 or w % pw != 0:
|
||||||
|
raise ValueError(
|
||||||
|
"raw_latent_shape must be divisible by patch_size for SAP attention"
|
||||||
|
)
|
||||||
|
|
||||||
|
num_frame = t // pt
|
||||||
|
frame_size = (h // ph) * (w // pw)
|
||||||
|
|
||||||
|
return SparseVideoGen2AttentionMetadata(
|
||||||
|
current_timestep=current_timestep,
|
||||||
|
num_q_centroids=num_q_centroids,
|
||||||
|
num_k_centroids=num_k_centroids,
|
||||||
|
top_p_kmeans=top_p_kmeans,
|
||||||
|
min_kc_ratio=min_kc_ratio,
|
||||||
|
kmeans_iter_init=kmeans_iter_init,
|
||||||
|
kmeans_iter_step=kmeans_iter_step,
|
||||||
|
zero_step_kmeans_init=zero_step_kmeans_init,
|
||||||
|
first_layers_fp=first_layers_fp,
|
||||||
|
first_times_fp=first_times_fp,
|
||||||
|
context_length=context_length,
|
||||||
|
prompt_length=prompt_length,
|
||||||
|
num_frame=num_frame,
|
||||||
|
frame_size=frame_size,
|
||||||
|
cache=cache,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class SparseVideoGen2AttentionImpl(AttentionImpl):
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
num_heads: int,
|
||||||
|
head_size: int,
|
||||||
|
causal: bool,
|
||||||
|
softmax_scale: float,
|
||||||
|
num_kv_heads: int | None = None,
|
||||||
|
prefix: str = "",
|
||||||
|
**extra_impl_args,
|
||||||
|
) -> None:
|
||||||
|
if causal:
|
||||||
|
raise ValueError(
|
||||||
|
"Sparse Video Gen 2 attention does not support causal attention"
|
||||||
|
)
|
||||||
|
if not svg2_available:
|
||||||
|
raise ImportError(
|
||||||
|
"Sparse Video Gen 2 attention backend requires svg package to be installed"
|
||||||
|
"Please install it by following the instructions at "
|
||||||
|
"https://github.com/svg-project/Sparse-VideoGen"
|
||||||
|
)
|
||||||
|
self.prefix = prefix
|
||||||
|
self.layer_idx = self._get_layer_idx(prefix)
|
||||||
|
|
||||||
|
def _get_layer_idx(self, prefix: str) -> int:
|
||||||
|
parts = prefix.split(".")
|
||||||
|
if len(parts) < 3:
|
||||||
|
raise ValueError(
|
||||||
|
f"Invalid prefix for SparseVideoGen2AttentionImpl: {prefix}"
|
||||||
|
)
|
||||||
|
return int(parts[-3])
|
||||||
|
|
||||||
|
def kmeans_init(
|
||||||
|
self,
|
||||||
|
query: torch.Tensor,
|
||||||
|
key: torch.Tensor,
|
||||||
|
attn_metadata: SparseVideoGen2AttentionMetadata,
|
||||||
|
):
|
||||||
|
cfg, num_heads, seq_len, dim = query.size()
|
||||||
|
qlabels, qcentroids, qcluster_sizes, qiter = batch_kmeans_Euclid(
|
||||||
|
query.reshape(cfg * num_heads, seq_len, dim),
|
||||||
|
n_clusters=attn_metadata.num_q_centroids,
|
||||||
|
max_iters=attn_metadata.kmeans_iter_init,
|
||||||
|
)
|
||||||
|
klabels, kcentroids, kcluster_sizes, kiter = batch_kmeans_Euclid(
|
||||||
|
key.reshape(cfg * num_heads, seq_len, dim),
|
||||||
|
n_clusters=attn_metadata.num_k_centroids,
|
||||||
|
max_iters=attn_metadata.kmeans_iter_init,
|
||||||
|
)
|
||||||
|
|
||||||
|
layer_cache = attn_metadata.cache.get_layer(self.layer_idx)
|
||||||
|
layer_cache.q_centroids = qcentroids
|
||||||
|
layer_cache.k_centroids = kcentroids
|
||||||
|
|
||||||
|
return (
|
||||||
|
qlabels,
|
||||||
|
qcentroids,
|
||||||
|
qcluster_sizes,
|
||||||
|
qiter,
|
||||||
|
klabels,
|
||||||
|
kcentroids,
|
||||||
|
kcluster_sizes,
|
||||||
|
kiter,
|
||||||
|
)
|
||||||
|
|
||||||
|
def kmeans_step(
|
||||||
|
self,
|
||||||
|
query: torch.Tensor,
|
||||||
|
key: torch.Tensor,
|
||||||
|
attn_metadata: SparseVideoGen2AttentionMetadata,
|
||||||
|
):
|
||||||
|
cfg, num_heads, seq_len, dim = query.size()
|
||||||
|
layer_cache = attn_metadata.cache.get_layer(self.layer_idx)
|
||||||
|
qlabels, qcentroids, qcluster_sizes, qiter = batch_kmeans_Euclid(
|
||||||
|
query.reshape(cfg * num_heads, seq_len, dim),
|
||||||
|
n_clusters=attn_metadata.num_q_centroids,
|
||||||
|
max_iters=attn_metadata.kmeans_iter_step,
|
||||||
|
init_centroids=layer_cache.q_centroids,
|
||||||
|
)
|
||||||
|
klabels, kcentroids, kcluster_sizes, kiter = batch_kmeans_Euclid(
|
||||||
|
key.reshape(cfg * num_heads, seq_len, dim),
|
||||||
|
n_clusters=attn_metadata.num_k_centroids,
|
||||||
|
max_iters=attn_metadata.kmeans_iter_step,
|
||||||
|
init_centroids=layer_cache.k_centroids,
|
||||||
|
)
|
||||||
|
|
||||||
|
layer_cache.q_centroids = qcentroids
|
||||||
|
layer_cache.k_centroids = kcentroids
|
||||||
|
|
||||||
|
return (
|
||||||
|
qlabels,
|
||||||
|
qcentroids,
|
||||||
|
qcluster_sizes,
|
||||||
|
qiter,
|
||||||
|
klabels,
|
||||||
|
kcentroids,
|
||||||
|
kcluster_sizes,
|
||||||
|
kiter,
|
||||||
|
)
|
||||||
|
|
||||||
|
def kmeans_clustering(
|
||||||
|
self,
|
||||||
|
query: torch.Tensor,
|
||||||
|
key: torch.Tensor,
|
||||||
|
attn_metadata: SparseVideoGen2AttentionMetadata,
|
||||||
|
):
|
||||||
|
layer_cache = attn_metadata.cache.get_layer(self.layer_idx)
|
||||||
|
if not layer_cache.centroids_initialized:
|
||||||
|
(
|
||||||
|
qlabels,
|
||||||
|
qcentroids,
|
||||||
|
qcluster_sizes,
|
||||||
|
qiter,
|
||||||
|
klabels,
|
||||||
|
kcentroids,
|
||||||
|
kcluster_sizes,
|
||||||
|
kiter,
|
||||||
|
) = self.kmeans_init(query, key, attn_metadata)
|
||||||
|
layer_cache.centroids_initialized = True
|
||||||
|
logger.debug(
|
||||||
|
"Centroids initialized at layer %s (init iters: %s).",
|
||||||
|
self.layer_idx,
|
||||||
|
attn_metadata.kmeans_iter_init,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
(
|
||||||
|
qlabels,
|
||||||
|
qcentroids,
|
||||||
|
qcluster_sizes,
|
||||||
|
qiter,
|
||||||
|
klabels,
|
||||||
|
kcentroids,
|
||||||
|
kcluster_sizes,
|
||||||
|
kiter,
|
||||||
|
) = self.kmeans_step(query, key, attn_metadata)
|
||||||
|
|
||||||
|
return (
|
||||||
|
qlabels,
|
||||||
|
qcentroids,
|
||||||
|
qcluster_sizes,
|
||||||
|
qiter,
|
||||||
|
klabels,
|
||||||
|
kcentroids,
|
||||||
|
kcluster_sizes,
|
||||||
|
kiter,
|
||||||
|
)
|
||||||
|
|
||||||
|
def semantic_aware_permutation(
|
||||||
|
self,
|
||||||
|
query: torch.Tensor,
|
||||||
|
key: torch.Tensor,
|
||||||
|
value: torch.Tensor,
|
||||||
|
attn_metadata: SparseVideoGen2AttentionMetadata,
|
||||||
|
):
|
||||||
|
cfg, num_heads, seq_len, dim = query.size()
|
||||||
|
|
||||||
|
# 1. Kmeans clustering
|
||||||
|
(
|
||||||
|
qlabels,
|
||||||
|
qcentroids,
|
||||||
|
qcluster_sizes,
|
||||||
|
qiter,
|
||||||
|
klabels,
|
||||||
|
kcentroids,
|
||||||
|
kcluster_sizes,
|
||||||
|
kiter,
|
||||||
|
) = self.kmeans_clustering(query, key, attn_metadata)
|
||||||
|
|
||||||
|
# 2. Identify dynamic map
|
||||||
|
q_cluster_sizes = qcluster_sizes.view(
|
||||||
|
cfg, num_heads, attn_metadata.num_q_centroids
|
||||||
|
)
|
||||||
|
k_cluster_sizes = kcluster_sizes.view(
|
||||||
|
cfg, num_heads, attn_metadata.num_k_centroids
|
||||||
|
)
|
||||||
|
|
||||||
|
dynamic_map = identify_dynamic_map(
|
||||||
|
qcentroids.view(cfg, num_heads, attn_metadata.num_q_centroids, dim),
|
||||||
|
kcentroids.view(cfg, num_heads, attn_metadata.num_k_centroids, dim),
|
||||||
|
q_cluster_sizes,
|
||||||
|
k_cluster_sizes,
|
||||||
|
attn_metadata.top_p_kmeans,
|
||||||
|
attn_metadata.min_kc_ratio,
|
||||||
|
)
|
||||||
|
|
||||||
|
# 3. Permute the query, key, value
|
||||||
|
q_permuted, q_sorted_indices = permute_tensor_by_labels_triton(
|
||||||
|
query, qlabels, dim=2
|
||||||
|
)
|
||||||
|
k_permuted, k_sorted_indices = permute_tensor_by_labels_triton(
|
||||||
|
key, klabels, dim=2
|
||||||
|
)
|
||||||
|
v_permuted, v_sorted_indices = permute_tensor_by_labels_triton(
|
||||||
|
value, klabels, dim=2, sorted_indices=k_sorted_indices
|
||||||
|
)
|
||||||
|
|
||||||
|
return (
|
||||||
|
q_permuted,
|
||||||
|
k_permuted,
|
||||||
|
v_permuted,
|
||||||
|
dynamic_map,
|
||||||
|
q_cluster_sizes,
|
||||||
|
k_cluster_sizes,
|
||||||
|
q_sorted_indices,
|
||||||
|
)
|
||||||
|
|
||||||
|
def _hunyuan_dynamic_map_post_processing(
|
||||||
|
self,
|
||||||
|
q_perm: torch.Tensor,
|
||||||
|
k_perm: torch.Tensor,
|
||||||
|
v_perm: torch.Tensor,
|
||||||
|
query: torch.Tensor,
|
||||||
|
key: torch.Tensor,
|
||||||
|
value: torch.Tensor,
|
||||||
|
dyn_map: torch.Tensor,
|
||||||
|
qc_sz_s: torch.Tensor,
|
||||||
|
kc_sz_s: torch.Tensor,
|
||||||
|
q_sorted_indices: torch.Tensor,
|
||||||
|
video_length: int,
|
||||||
|
context_length: int,
|
||||||
|
prompt_length: int,
|
||||||
|
unprompt_length: int,
|
||||||
|
) -> tuple[
|
||||||
|
torch.Tensor,
|
||||||
|
torch.Tensor,
|
||||||
|
torch.Tensor,
|
||||||
|
torch.Tensor,
|
||||||
|
torch.Tensor,
|
||||||
|
torch.Tensor,
|
||||||
|
torch.Tensor,
|
||||||
|
]:
|
||||||
|
# Place the permuted video tokens back and keep text tokens at the tail.
|
||||||
|
query[:, :, :-context_length, :] = q_perm
|
||||||
|
key[:, :, :-context_length, :] = k_perm
|
||||||
|
value[:, :, :-context_length, :] = v_perm
|
||||||
|
|
||||||
|
# Add prompt/unprompt clusters to the dynamic map.
|
||||||
|
dyn_map = F.pad(dyn_map, (0, 2, 0, 2), value=0)
|
||||||
|
dyn_map[:, :, -2, :-1] = True
|
||||||
|
dyn_map[:, :, :-1, -2] = True
|
||||||
|
dyn_map[:, :, -1, -1] = True
|
||||||
|
|
||||||
|
qc_sz_s = F.pad(qc_sz_s, (0, 2), value=0)
|
||||||
|
qc_sz_s[:, :, -2] = prompt_length
|
||||||
|
qc_sz_s[:, :, -1] = unprompt_length
|
||||||
|
kc_sz_s = F.pad(kc_sz_s, (0, 2), value=0)
|
||||||
|
kc_sz_s[:, :, -2] = prompt_length
|
||||||
|
kc_sz_s[:, :, -1] = unprompt_length
|
||||||
|
|
||||||
|
q_sorted_indices = F.pad(q_sorted_indices, (0, context_length), value=0)
|
||||||
|
q_sorted_indices[:, video_length:] = torch.arange(
|
||||||
|
video_length,
|
||||||
|
video_length + context_length,
|
||||||
|
device=q_sorted_indices.device,
|
||||||
|
)
|
||||||
|
return query, key, value, dyn_map, qc_sz_s, kc_sz_s, q_sorted_indices
|
||||||
|
|
||||||
|
def forward(
|
||||||
|
self,
|
||||||
|
query: torch.Tensor,
|
||||||
|
key: torch.Tensor,
|
||||||
|
value: torch.Tensor,
|
||||||
|
attn_metadata: SparseVideoGen2AttentionMetadata,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
torch.backends.cuda.preferred_linalg_library(backend="magma")
|
||||||
|
res = None
|
||||||
|
# bshd -> bhsd
|
||||||
|
query = query.transpose(1, 2).contiguous()
|
||||||
|
key = key.transpose(1, 2).contiguous()
|
||||||
|
value = value.transpose(1, 2).contiguous()
|
||||||
|
batch_size, num_heads, seq_len, dim = query.size()
|
||||||
|
|
||||||
|
context_length, num_frame, frame_size = (
|
||||||
|
attn_metadata.context_length,
|
||||||
|
attn_metadata.num_frame,
|
||||||
|
attn_metadata.frame_size,
|
||||||
|
)
|
||||||
|
prompt_length = attn_metadata.prompt_length
|
||||||
|
if prompt_length is None:
|
||||||
|
prompt_length = context_length
|
||||||
|
|
||||||
|
assert (
|
||||||
|
seq_len == context_length + num_frame * frame_size
|
||||||
|
), f"Query Shape: {seq_len} is not equivalent to {context_length} + {num_frame} * {frame_size}"
|
||||||
|
|
||||||
|
# Determine if we use Full Attention to calculate
|
||||||
|
full_attention_flag = False
|
||||||
|
|
||||||
|
if self.layer_idx < attn_metadata.first_layers_fp:
|
||||||
|
full_attention_flag = True
|
||||||
|
if attn_metadata.current_timestep > attn_metadata.first_times_fp:
|
||||||
|
full_attention_flag = True
|
||||||
|
|
||||||
|
if full_attention_flag:
|
||||||
|
if attn_metadata.zero_step_kmeans_init:
|
||||||
|
video_length = attn_metadata.num_frame * attn_metadata.frame_size
|
||||||
|
query_video = query[:, :, :video_length, :].contiguous()
|
||||||
|
key_video = key[:, :, :video_length, :].contiguous()
|
||||||
|
self.kmeans_clustering(query_video, key_video, attn_metadata)
|
||||||
|
|
||||||
|
with sdpa_kernel(
|
||||||
|
SDPBackend.CUDNN_ATTENTION
|
||||||
|
): # not sure why we need to force cudnn here, but it's faster than flash attention
|
||||||
|
output_hidden_states = torch.nn.functional.scaled_dot_product_attention(
|
||||||
|
query, key, value, dropout_p=0.0, is_causal=False
|
||||||
|
)
|
||||||
|
|
||||||
|
res = output_hidden_states.reshape(
|
||||||
|
batch_size, num_heads, seq_len, dim
|
||||||
|
).transpose(1, 2)
|
||||||
|
else:
|
||||||
|
if context_length > 0:
|
||||||
|
video_length = num_frame * frame_size
|
||||||
|
unprompt_length = max(context_length - prompt_length, 0)
|
||||||
|
query_video = query[:, :, :video_length, :].contiguous()
|
||||||
|
key_video = key[:, :, :video_length, :].contiguous()
|
||||||
|
value_video = value[:, :, :video_length, :].contiguous()
|
||||||
|
|
||||||
|
(
|
||||||
|
q_perm,
|
||||||
|
k_perm,
|
||||||
|
v_perm,
|
||||||
|
dyn_map,
|
||||||
|
qc_sz_s,
|
||||||
|
kc_sz_s,
|
||||||
|
q_sorted_indices,
|
||||||
|
) = self.semantic_aware_permutation(
|
||||||
|
query_video, key_video, value_video, attn_metadata
|
||||||
|
)
|
||||||
|
(
|
||||||
|
q_perm,
|
||||||
|
k_perm,
|
||||||
|
v_perm,
|
||||||
|
dyn_map,
|
||||||
|
qc_sz_s,
|
||||||
|
kc_sz_s,
|
||||||
|
q_sorted_indices,
|
||||||
|
) = self._hunyuan_dynamic_map_post_processing(
|
||||||
|
q_perm,
|
||||||
|
k_perm,
|
||||||
|
v_perm,
|
||||||
|
query,
|
||||||
|
key,
|
||||||
|
value,
|
||||||
|
dyn_map,
|
||||||
|
qc_sz_s,
|
||||||
|
kc_sz_s,
|
||||||
|
q_sorted_indices,
|
||||||
|
video_length,
|
||||||
|
context_length,
|
||||||
|
prompt_length,
|
||||||
|
unprompt_length,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
(
|
||||||
|
q_perm,
|
||||||
|
k_perm,
|
||||||
|
v_perm,
|
||||||
|
dyn_map,
|
||||||
|
qc_sz_s,
|
||||||
|
kc_sz_s,
|
||||||
|
q_sorted_indices,
|
||||||
|
) = self.semantic_aware_permutation(query, key, value, attn_metadata)
|
||||||
|
|
||||||
|
output_permuted = dynamic_block_sparse_fwd_flashinfer(
|
||||||
|
q_perm, k_perm, v_perm, dyn_map, qc_sz_s, kc_sz_s, is_cpu=False
|
||||||
|
)
|
||||||
|
|
||||||
|
attn_output = apply_inverse_permutation_triton(
|
||||||
|
output_permuted, q_sorted_indices, dim=2
|
||||||
|
)
|
||||||
|
|
||||||
|
res = attn_output.reshape(batch_size, num_heads, seq_len, dim).transpose(
|
||||||
|
1, 2
|
||||||
|
)
|
||||||
|
|
||||||
|
torch.backends.cuda.preferred_linalg_library(
|
||||||
|
backend="default"
|
||||||
|
) # reset to default
|
||||||
|
return res.contiguous()
|
||||||
@@ -314,6 +314,19 @@ class WanTransformerBlock(nn.Module):
|
|||||||
self.to_out = RowParallelLinear(dim, dim, bias=True, reduce_results=True)
|
self.to_out = RowParallelLinear(dim, dim, bias=True, reduce_results=True)
|
||||||
tp_size = get_tp_world_size()
|
tp_size = get_tp_world_size()
|
||||||
self.local_num_heads = divide(num_heads, tp_size)
|
self.local_num_heads = divide(num_heads, tp_size)
|
||||||
|
self_attn_backends = supported_attention_backends
|
||||||
|
cross_attn_backends = supported_attention_backends
|
||||||
|
if (
|
||||||
|
supported_attention_backends is not None
|
||||||
|
and AttentionBackendEnum.SPARSE_VIDEO_GEN_2_ATTN
|
||||||
|
in supported_attention_backends
|
||||||
|
):
|
||||||
|
cross_attn_backends = supported_attention_backends.copy()
|
||||||
|
cross_attn_backends.remove(AttentionBackendEnum.SPARSE_VIDEO_GEN_2_ATTN)
|
||||||
|
logger.warning_once(
|
||||||
|
"Sparse Video Gen 2 attention backend is not supported for cross-attention; "
|
||||||
|
"removing SPARSE_VIDEO_GEN_2_ATTN from cross-attention backends."
|
||||||
|
)
|
||||||
if attention_type in ("sla", "sagesla"):
|
if attention_type in ("sla", "sagesla"):
|
||||||
self.attn1 = MinimalA2AAttnOp(
|
self.attn1 = MinimalA2AAttnOp(
|
||||||
num_heads=self.local_num_heads,
|
num_heads=self.local_num_heads,
|
||||||
@@ -330,7 +343,7 @@ class WanTransformerBlock(nn.Module):
|
|||||||
num_heads=self.local_num_heads,
|
num_heads=self.local_num_heads,
|
||||||
head_size=dim // num_heads,
|
head_size=dim // num_heads,
|
||||||
causal=False,
|
causal=False,
|
||||||
supported_attention_backends=supported_attention_backends,
|
supported_attention_backends=self_attn_backends,
|
||||||
prefix=f"{prefix}.attn1",
|
prefix=f"{prefix}.attn1",
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -365,7 +378,7 @@ class WanTransformerBlock(nn.Module):
|
|||||||
num_heads,
|
num_heads,
|
||||||
qk_norm=qk_norm,
|
qk_norm=qk_norm,
|
||||||
eps=eps,
|
eps=eps,
|
||||||
supported_attention_backends=supported_attention_backends,
|
supported_attention_backends=cross_attn_backends,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
# T2V
|
# T2V
|
||||||
@@ -374,7 +387,7 @@ class WanTransformerBlock(nn.Module):
|
|||||||
num_heads,
|
num_heads,
|
||||||
qk_norm=qk_norm,
|
qk_norm=qk_norm,
|
||||||
eps=eps,
|
eps=eps,
|
||||||
supported_attention_backends=supported_attention_backends,
|
supported_attention_backends=cross_attn_backends,
|
||||||
)
|
)
|
||||||
self.cross_attn_residual_norm = ScaleResidualLayerNormScaleShift(
|
self.cross_attn_residual_norm = ScaleResidualLayerNormScaleShift(
|
||||||
dim,
|
dim,
|
||||||
|
|||||||
@@ -1056,7 +1056,13 @@ class DenoisingStage(PipelineStage):
|
|||||||
)
|
)
|
||||||
|
|
||||||
# Predict noise residual
|
# Predict noise residual
|
||||||
attn_metadata = self._build_attn_metadata(i, batch, server_args)
|
attn_metadata = self._build_attn_metadata(
|
||||||
|
i,
|
||||||
|
batch,
|
||||||
|
server_args,
|
||||||
|
timestep_value=t_int,
|
||||||
|
timesteps=timesteps_cpu,
|
||||||
|
)
|
||||||
noise_pred = self._predict_noise_with_cfg(
|
noise_pred = self._predict_noise_with_cfg(
|
||||||
current_model=current_model,
|
current_model=current_model,
|
||||||
latent_model_input=latent_model_input,
|
latent_model_input=latent_model_input,
|
||||||
@@ -1190,7 +1196,13 @@ class DenoisingStage(PipelineStage):
|
|||||||
return noise_cfg
|
return noise_cfg
|
||||||
|
|
||||||
def _build_attn_metadata(
|
def _build_attn_metadata(
|
||||||
self, i: int, batch: Req, server_args: ServerArgs
|
self,
|
||||||
|
i: int,
|
||||||
|
batch: Req,
|
||||||
|
server_args: ServerArgs,
|
||||||
|
*,
|
||||||
|
timestep_value: int | None = None,
|
||||||
|
timesteps: torch.Tensor | None = None,
|
||||||
) -> Any | None:
|
) -> Any | None:
|
||||||
"""
|
"""
|
||||||
Build attention metadata for custom attention backends.
|
Build attention metadata for custom attention backends.
|
||||||
@@ -1218,6 +1230,92 @@ class DenoisingStage(PipelineStage):
|
|||||||
VSA_sparsity=server_args.attention_backend_config.VSA_sparsity,
|
VSA_sparsity=server_args.attention_backend_config.VSA_sparsity,
|
||||||
device=get_local_torch_device(),
|
device=get_local_torch_device(),
|
||||||
)
|
)
|
||||||
|
elif (
|
||||||
|
self.attn_backend.get_enum() == AttentionBackendEnum.SPARSE_VIDEO_GEN_2_ATTN
|
||||||
|
):
|
||||||
|
if timestep_value is None or timesteps is None:
|
||||||
|
raise ValueError(
|
||||||
|
"timestep_value and timesteps must be provided for SVG2 attention metadata"
|
||||||
|
)
|
||||||
|
|
||||||
|
svg2_cfg = server_args.attention_backend_config or {}
|
||||||
|
num_layers = server_args.pipeline_config.dit_config.num_layers
|
||||||
|
if (
|
||||||
|
server_args.pipeline_config.dit_config.prefix.lower() == "hunyuan"
|
||||||
|
and hasattr(server_args.pipeline_config.dit_config, "num_single_layers")
|
||||||
|
):
|
||||||
|
num_layers += server_args.pipeline_config.dit_config.num_single_layers
|
||||||
|
first_layers_fp = svg2_cfg.get("svg2_first_layers_fp", 0.03)
|
||||||
|
if first_layers_fp <= 1.0:
|
||||||
|
first_layers_fp = math.floor(first_layers_fp * num_layers)
|
||||||
|
first_layers_fp = max(0, min(int(first_layers_fp), num_layers))
|
||||||
|
|
||||||
|
first_times_fp = svg2_cfg.get("svg2_first_times_fp", 0.2)
|
||||||
|
if first_times_fp <= 1.0:
|
||||||
|
num_fp_steps = math.floor(first_times_fp * len(timesteps))
|
||||||
|
if num_fp_steps > 0:
|
||||||
|
first_times_fp = float(timesteps[num_fp_steps - 1].item() - 1)
|
||||||
|
else:
|
||||||
|
first_times_fp = float(timesteps.max().item() + 1)
|
||||||
|
|
||||||
|
current_timestep = int(timestep_value)
|
||||||
|
|
||||||
|
cache = batch.extra.get("svg2_cache")
|
||||||
|
if cache is None:
|
||||||
|
from sglang.multimodal_gen.runtime.layers.attention.backends.sparse_video_gen_2_attn import (
|
||||||
|
Svg2Cache,
|
||||||
|
)
|
||||||
|
|
||||||
|
cache = Svg2Cache()
|
||||||
|
batch.extra["svg2_cache"] = cache
|
||||||
|
|
||||||
|
patch_size = server_args.pipeline_config.dit_config.patch_size
|
||||||
|
if isinstance(patch_size, list):
|
||||||
|
patch_size = tuple(patch_size)
|
||||||
|
if isinstance(patch_size, int):
|
||||||
|
patch_size_t = getattr(
|
||||||
|
server_args.pipeline_config.dit_config, "patch_size_t", None
|
||||||
|
)
|
||||||
|
if patch_size_t is not None:
|
||||||
|
patch_size = (patch_size_t, patch_size, patch_size)
|
||||||
|
|
||||||
|
context_length = 0
|
||||||
|
prompt_length = None
|
||||||
|
if server_args.pipeline_config.dit_config.prefix.lower() == "hunyuan":
|
||||||
|
prompt_embeds = server_args.pipeline_config.get_pos_prompt_embeds(batch)
|
||||||
|
if isinstance(prompt_embeds, list):
|
||||||
|
text_embeds = prompt_embeds[0] if prompt_embeds else None
|
||||||
|
else:
|
||||||
|
text_embeds = prompt_embeds
|
||||||
|
if isinstance(text_embeds, torch.Tensor) and text_embeds.ndim >= 2:
|
||||||
|
context_length = int(text_embeds.shape[1])
|
||||||
|
if context_length > 0 and batch.prompt_attention_mask:
|
||||||
|
mask = batch.prompt_attention_mask[0]
|
||||||
|
if isinstance(mask, torch.Tensor):
|
||||||
|
if mask.shape[-1] > context_length:
|
||||||
|
mask = mask[:, -context_length:]
|
||||||
|
prompt_length = int(mask[0].sum().item())
|
||||||
|
if prompt_length is None:
|
||||||
|
prompt_length = context_length
|
||||||
|
|
||||||
|
attn_metadata = self.attn_metadata_builder.build(
|
||||||
|
current_timestep=current_timestep,
|
||||||
|
raw_latent_shape=batch.raw_latent_shape,
|
||||||
|
patch_size=patch_size,
|
||||||
|
num_q_centroids=svg2_cfg.get("svg2_num_q_centroids", 300),
|
||||||
|
num_k_centroids=svg2_cfg.get("svg2_num_k_centroids", 1000),
|
||||||
|
top_p_kmeans=svg2_cfg.get("svg2_top_p_kmeans", 0.9),
|
||||||
|
min_kc_ratio=svg2_cfg.get("svg2_min_kc_ratio", 0.1),
|
||||||
|
kmeans_iter_init=svg2_cfg.get("svg2_kmeans_iter_init", 50),
|
||||||
|
kmeans_iter_step=svg2_cfg.get("svg2_kmeans_iter_step", 2),
|
||||||
|
zero_step_kmeans_init=svg2_cfg.get("svg2_zero_step_kmeans_init", False),
|
||||||
|
first_layers_fp=first_layers_fp,
|
||||||
|
first_times_fp=first_times_fp,
|
||||||
|
context_length=context_length,
|
||||||
|
prompt_length=prompt_length,
|
||||||
|
cache=cache,
|
||||||
|
calculate_density=False, # only need density when doing head load balancing
|
||||||
|
)
|
||||||
elif self.attn_backend.get_enum() == AttentionBackendEnum.VMOBA_ATTN:
|
elif self.attn_backend.get_enum() == AttentionBackendEnum.VMOBA_ATTN:
|
||||||
moba_params = server_args.attention_backend_config.moba_config.copy()
|
moba_params = server_args.attention_backend_config.moba_config.copy()
|
||||||
moba_params.update(
|
moba_params.update(
|
||||||
|
|||||||
@@ -224,6 +224,35 @@ class CudaPlatformBase(Platform):
|
|||||||
raise ImportError(
|
raise ImportError(
|
||||||
"Video Sparse Attention backend is not installed."
|
"Video Sparse Attention backend is not installed."
|
||||||
) from e
|
) from e
|
||||||
|
elif selected_backend == AttentionBackendEnum.SPARSE_VIDEO_GEN_2_ATTN:
|
||||||
|
try:
|
||||||
|
from svg.kernels.triton.permute import ( # noqa: F401
|
||||||
|
apply_inverse_permutation_triton,
|
||||||
|
permute_tensor_by_labels_triton,
|
||||||
|
)
|
||||||
|
from svg.kmeans_utils import ( # noqa: F401
|
||||||
|
batch_kmeans_Euclid,
|
||||||
|
density_calculation,
|
||||||
|
dynamic_block_sparse_fwd_flashinfer,
|
||||||
|
identify_dynamic_map,
|
||||||
|
)
|
||||||
|
|
||||||
|
from sglang.multimodal_gen.runtime.layers.attention.backends.sparse_video_gen_2_attn import ( # noqa: F401
|
||||||
|
SparseVideoGen2AttentionBackend,
|
||||||
|
)
|
||||||
|
|
||||||
|
logger.info("Using Sparse Video Gen 2 (SAP) Attention backend")
|
||||||
|
return "sglang.multimodal_gen.runtime.layers.attention.backends.sparse_video_gen_2_attn.SparseVideoGen2AttentionBackend"
|
||||||
|
except ImportError as e:
|
||||||
|
logger.error(
|
||||||
|
"Failed to import Sparse Video Gen 2 (SAP) Attention backend: %s",
|
||||||
|
str(e),
|
||||||
|
)
|
||||||
|
raise ImportError(
|
||||||
|
"Sparse Video Gen 2 (SAP) Attention backend is not installed. "
|
||||||
|
"Please install it by following the instructions at "
|
||||||
|
"https://github.com/svg-project/Sparse-VideoGen"
|
||||||
|
) from e
|
||||||
elif selected_backend == AttentionBackendEnum.VMOBA_ATTN:
|
elif selected_backend == AttentionBackendEnum.VMOBA_ATTN:
|
||||||
try:
|
try:
|
||||||
from kernel.attn.vmoba_attn.vmoba import moba_attn_varlen # noqa: F401
|
from kernel.attn.vmoba_attn.vmoba import moba_attn_varlen # noqa: F401
|
||||||
|
|||||||
@@ -31,6 +31,7 @@ class AttentionBackendEnum(enum.Enum):
|
|||||||
SAGE_ATTN = enum.auto()
|
SAGE_ATTN = enum.auto()
|
||||||
SAGE_ATTN_3 = enum.auto()
|
SAGE_ATTN_3 = enum.auto()
|
||||||
VIDEO_SPARSE_ATTN = enum.auto()
|
VIDEO_SPARSE_ATTN = enum.auto()
|
||||||
|
SPARSE_VIDEO_GEN_2_ATTN = enum.auto()
|
||||||
VMOBA_ATTN = enum.auto()
|
VMOBA_ATTN = enum.auto()
|
||||||
AITER = enum.auto()
|
AITER = enum.auto()
|
||||||
SLA_ATTN = enum.auto()
|
SLA_ATTN = enum.auto()
|
||||||
|
|||||||
@@ -148,7 +148,10 @@ def _log_process_aware(
|
|||||||
if should_log:
|
if should_log:
|
||||||
# stacklevel=3 to show the original caller's location,
|
# stacklevel=3 to show the original caller's location,
|
||||||
# as this function is called by the patched methods.
|
# as this function is called by the patched methods.
|
||||||
logger_self.log(level, msg, *args, stacklevel=3, **kwargs)
|
if "stacklevel" in kwargs:
|
||||||
|
logger_self.log(level, msg, *args, **kwargs)
|
||||||
|
else:
|
||||||
|
logger_self.log(level, msg, *args, stacklevel=3, **kwargs)
|
||||||
|
|
||||||
|
|
||||||
class _SGLDiffusionLogger(Logger):
|
class _SGLDiffusionLogger(Logger):
|
||||||
|
|||||||
Reference in New Issue
Block a user