Compare commits

...
6 Commits
Author SHA1 Message Date
Edenzzzz b118043f55 add example to docs 2025-06-14 22:38:55 +00:00
Edenzzzz 83d45684b6 Merge branch 'main' into refactor_attn 2025-06-14 22:00:58 +00:00
Edenzzzz afb16377a8 fix 2025-06-10 23:48:22 +00:00
Edenzzzz b44c28749c Merge branch 'main' into refactor_attn 2025-06-10 23:40:24 +00:00
Edenzzzz a9b9bad9a2 merge main 2025-06-10 23:38:59 +00:00
Edenzzzz 1b2b97544d fix 2025-06-10 23:32:00 +00:00
22 changed files with 88 additions and 70 deletions
+1 -1
View File
@@ -288,7 +288,7 @@ Sequence parallelism splits sequences across devices:
```python
# Distributed attention for long sequences
from fastvideo.v1.attention import DistributedAttention
from fastvideo.v1.layers.attention import DistributedAttention
self.attn = DistributedAttention(
num_heads=num_heads,
+3 -3
View File
@@ -96,8 +96,8 @@ Replace standard attention with FastVideo's optimized attention:
```python
# Local attention patterns
from fastvideo.v1.attention import LocalAttention
from fastvideo.v1.attention.backends.abstract import _Backend
from fastvideo.v1.layers.attention import LocalAttention
from fastvideo.v1.layers.attention.backends.abstract import _Backend
self.attn = LocalAttention(
num_heads=num_heads,
head_size=head_dim,
@@ -108,7 +108,7 @@ self.attn = LocalAttention(
)
# Distributed attention for long sequences
from fastvideo.v1.attention import DistributedAttention
from fastvideo.v1.layers.attention import DistributedAttention
self.attn = DistributedAttention(
num_heads=num_heads,
head_size=head_dim,
+8 -1
View File
@@ -1,7 +1,7 @@
(sta-demo)=
# 🔍 Demo
There is a demo for 2D STA with window size (6,6) operating on a (10, 10) image.
This is is a demo for 2D STA with window size (6,6) operating on a (10, 10) image.
<div style="text-align: center;">
<video controls width="800">
@@ -9,3 +9,10 @@ There is a demo for 2D STA with window size (6,6) operating on a (10, 10) image.
Your browser does not support the video tag.
</video>
</div>
You can run STA using the following command:
```bash
huggingface-cli download hunyuanvideo-community/HunyuanVideo --local-dir data/hunyuan
bash scripts/inference/inference_hunyuan_STA.sh
```
-20
View File
@@ -1,20 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
from fastvideo.v1.attention.backends.abstract import (AttentionBackend,
AttentionMetadata,
AttentionMetadataBuilder)
from fastvideo.v1.attention.layer import (DistributedAttention,
DistributedAttention_VSA,
LocalAttention)
from fastvideo.v1.attention.selector import get_attn_backend
__all__ = [
"DistributedAttention",
"DistributedAttention_VSA",
"LocalAttention",
"AttentionBackend",
"AttentionMetadata",
"AttentionMetadataBuilder",
# "AttentionState",
"get_attn_backend",
]
+4 -3
View File
@@ -5,16 +5,17 @@ import time
from collections import defaultdict
from contextlib import contextmanager
from dataclasses import dataclass
from typing import Optional
from typing import TYPE_CHECKING, Optional
import torch
# if TYPE_CHECKING:
from fastvideo.v1.attention import AttentionMetadata
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
if TYPE_CHECKING:
from fastvideo.v1.layers.attention import AttentionMetadata
logger = init_logger(__name__)
# TODO(will): check if this is needed
+19
View File
@@ -0,0 +1,19 @@
# SPDX-License-Identifier: Apache-2.0
from fastvideo.v1.layers.attention.backends.abstract import (
AttentionBackend, AttentionMetadata, AttentionMetadataBuilder)
from fastvideo.v1.layers.attention.layer import (DistributedAttention,
DistributedAttention_VSA,
LocalAttention)
from fastvideo.v1.layers.attention.selector import get_attn_backend
__all__ = [
"DistributedAttention",
"LocalAttention",
"DistributedAttention_VSA",
"AttentionBackend",
"AttentionMetadata",
"AttentionMetadataBuilder",
# "AttentionState",
"get_attn_backend",
]
@@ -14,10 +14,9 @@ try:
except ImportError:
flash_attn_func = flash_attn_2_func
from fastvideo.v1.attention.backends.abstract import (AttentionBackend,
AttentionImpl,
AttentionMetadata,
AttentionMetadataBuilder)
from fastvideo.v1.layers.attention.backends.abstract import (
AttentionBackend, AttentionImpl, AttentionMetadata,
AttentionMetadataBuilder)
from fastvideo.v1.logger import init_logger
logger = init_logger(__name__)
@@ -4,10 +4,10 @@ from typing import List, Optional, Type
import torch
from sageattention import sageattn
from fastvideo.v1.attention.backends.abstract import (
from fastvideo.v1.layers.attention.backends.abstract import (
AttentionBackend) # FlashAttentionMetadata,
from fastvideo.v1.attention.backends.abstract import (AttentionImpl,
AttentionMetadata)
from fastvideo.v1.layers.attention.backends.abstract import (AttentionImpl,
AttentionMetadata)
from fastvideo.v1.logger import init_logger
logger = init_logger(__name__)
@@ -3,10 +3,10 @@ from typing import List, Optional, Type
import torch
from fastvideo.v1.attention.backends.abstract import (
from fastvideo.v1.layers.attention.backends.abstract import (
AttentionBackend) # FlashAttentionMetadata,
from fastvideo.v1.attention.backends.abstract import (AttentionImpl,
AttentionMetadata)
from fastvideo.v1.layers.attention.backends.abstract import (AttentionImpl,
AttentionMetadata)
from fastvideo.v1.logger import init_logger
logger = init_logger(__name__)
@@ -8,13 +8,12 @@ from einops import rearrange
from st_attn import sliding_tile_attention
import fastvideo.v1.envs as envs
from fastvideo.v1.attention.backends.abstract import (AttentionBackend,
AttentionImpl,
AttentionMetadata,
AttentionMetadataBuilder)
from fastvideo.v1.distributed import get_sp_group
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.forward_context import ForwardContext, get_forward_context
from fastvideo.v1.layers.attention.backends.abstract import (
AttentionBackend, AttentionImpl, AttentionMetadata,
AttentionMetadataBuilder)
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
@@ -11,12 +11,11 @@ try:
except ImportError:
video_sparse_attn = None
from fastvideo.v1.attention.backends.abstract import (AttentionBackend,
AttentionImpl,
AttentionMetadata,
AttentionMetadataBuilder)
from fastvideo.v1.distributed import get_sp_group
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.layers.attention.backends.abstract import (
AttentionBackend, AttentionImpl, AttentionMetadata,
AttentionMetadataBuilder)
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
@@ -5,13 +5,13 @@ from typing import Optional, Tuple
import torch
import torch.nn as nn
from fastvideo.v1.attention.selector import (backend_name_to_enum,
get_attn_backend)
from fastvideo.v1.distributed.communication_op import (
sequence_model_parallel_all_gather, sequence_model_parallel_all_to_all_4D)
from fastvideo.v1.distributed.parallel_state import (get_sp_parallel_rank,
get_sp_world_size)
from fastvideo.v1.forward_context import ForwardContext, get_forward_context
from fastvideo.v1.layers.attention.selector import (backend_name_to_enum,
get_attn_backend)
from fastvideo.v1.platforms import _Backend
from fastvideo.v1.utils import get_compute_dtype
@@ -9,7 +9,7 @@ from typing import Generator, Optional, Tuple, Type, cast
import torch
import fastvideo.v1.envs as envs
from fastvideo.v1.attention.backends.abstract import AttentionBackend
from fastvideo.v1.layers.attention.backends.abstract import AttentionBackend
from fastvideo.v1.logger import init_logger
from fastvideo.v1.platforms import _Backend, current_platform
from fastvideo.v1.utils import STR_BACKEND_ENV_VAR, resolve_obj_by_qualname
+1 -1
View File
@@ -6,11 +6,11 @@ import numpy as np
import torch
import torch.nn as nn
from fastvideo.v1.attention import DistributedAttention, LocalAttention
from fastvideo.v1.configs.models.dits import HunyuanVideoConfig
from fastvideo.v1.configs.sample.teacache import TeaCacheParams
from fastvideo.v1.distributed.parallel_state import get_sp_world_size
from fastvideo.v1.forward_context import get_forward_context
from fastvideo.v1.layers.attention import DistributedAttention, LocalAttention
from fastvideo.v1.layers.layernorm import (LayerNormScaleShift, ScaleResidual,
ScaleResidualLayerNormScaleShift)
from fastvideo.v1.layers.linear import ReplicatedLinear
+1 -1
View File
@@ -16,9 +16,9 @@ import torch
from einops import rearrange, repeat
from torch import nn
from fastvideo.v1.attention import DistributedAttention, LocalAttention
from fastvideo.v1.configs.models.dits import StepVideoConfig
from fastvideo.v1.distributed.parallel_state import get_sp_world_size
from fastvideo.v1.layers.attention import DistributedAttention, LocalAttention
from fastvideo.v1.layers.layernorm import LayerNormScaleShift
from fastvideo.v1.layers.linear import ReplicatedLinear
from fastvideo.v1.layers.mlp import MLP
+15 -2
View File
@@ -8,12 +8,13 @@ import torch
import torch.nn as nn
import fastvideo.v1.envs as envs
from fastvideo.v1.attention import (DistributedAttention,
DistributedAttention_VSA, LocalAttention)
from fastvideo.v1.configs.models.dits import WanVideoConfig
from fastvideo.v1.configs.sample.wan import WanTeaCacheParams
from fastvideo.v1.distributed.parallel_state import get_sp_world_size
from fastvideo.v1.forward_context import get_forward_context
from fastvideo.v1.layers.attention import (DistributedAttention,
DistributedAttention_VSA,
LocalAttention)
from fastvideo.v1.layers.layernorm import (LayerNormScaleShift, RMSNorm,
ScaleResidual,
ScaleResidualLayerNormScaleShift)
@@ -569,6 +570,18 @@ class WanTransformer3DModel(CachableDiT):
self.gradient_checkpointing = False
# Initialize cache-related attributes
self.previous_e0_even = None
self.previous_e0_odd = None
self.previous_residual_even = None
self.previous_residual_odd = None
self.is_even = True
self.should_calc_even = True
self.should_calc_odd = True
self.accumulated_rel_l1_distance_even = 0
self.accumulated_rel_l1_distance_odd = 0
self.cnt = 0
self.__post_init__()
def forward(self,
+2 -2
View File
@@ -8,13 +8,13 @@ from typing import Iterable, Optional, Set, Tuple, Union
import torch
import torch.nn as nn
# from transformers.modeling_attn_mask_utils import _create_4d_causal_attention_mask, _prepare_4d_attention_mask
from fastvideo.v1.attention import LocalAttention
from fastvideo.v1.configs.models.encoders import (BaseEncoderOutput,
CLIPTextConfig,
CLIPVisionConfig)
from fastvideo.v1.distributed import divide, get_tp_world_size
from fastvideo.v1.layers.activation import get_act_fn
# from transformers.modeling_attn_mask_utils import _create_4d_causal_attention_mask, _prepare_4d_attention_mask
from fastvideo.v1.layers.attention import LocalAttention
from fastvideo.v1.layers.linear import (ColumnParallelLinear, QKVParallelLinear,
RowParallelLinear)
from fastvideo.v1.layers.quantization import QuantizationConfig
+2 -2
View File
@@ -28,12 +28,12 @@ from typing import Any, Dict, Iterable, Optional, Set, Tuple
import torch
from torch import nn
# from vllm.model_executor.layers.quantization import QuantizationConfig
from fastvideo.v1.attention import LocalAttention
# from ..utils import (extract_layer_index)
from fastvideo.v1.configs.models.encoders import BaseEncoderOutput, LlamaConfig
from fastvideo.v1.distributed import get_tp_world_size
from fastvideo.v1.layers.activation import SiluAndMul
# from vllm.model_executor.layers.quantization import QuantizationConfig
from fastvideo.v1.layers.attention import LocalAttention
from fastvideo.v1.layers.layernorm import RMSNorm
from fastvideo.v1.layers.linear import (MergedColumnParallelLinear,
QKVParallelLinear, RowParallelLinear)
+4 -3
View File
@@ -11,13 +11,13 @@ import torch
from einops import rearrange
from tqdm.auto import tqdm
from fastvideo.v1.attention import get_attn_backend
from fastvideo.v1.distributed import (get_sp_parallel_rank, get_sp_world_size,
get_torch_device, get_world_group)
from fastvideo.v1.distributed.communication_op import (
sequence_model_parallel_all_gather)
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.forward_context import set_forward_context
from fastvideo.v1.layers.attention import get_attn_backend
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.pipelines.stages.base import PipelineStage
@@ -26,13 +26,14 @@ from fastvideo.v1.platforms import _Backend
st_attn_available = False
if importlib.util.find_spec("st_attn") is not None:
st_attn_available = True
from fastvideo.v1.attention.backends.sliding_tile_attn import (
from fastvideo.v1.layers.attention.backends.sliding_tile_attn import (
SlidingTileAttentionBackend)
vsa_available = False
if importlib.util.find_spec("vsa") is not None:
vsa_available = True
from fastvideo.v1.attention.backends.video_sparse_attn import (
from fastvideo.v1.layers.attention.backends.video_sparse_attn import (
VideoSparseAttentionBackend)
logger = init_logger(__name__)
+10 -10
View File
@@ -117,10 +117,10 @@ class CudaPlatformBase(Platform):
try:
from st_attn import sliding_tile_attention # noqa: F401
from fastvideo.v1.attention.backends.sliding_tile_attn import ( # noqa: F401
from fastvideo.v1.layers.attention.backends.sliding_tile_attn import ( # noqa: F401
SlidingTileAttentionBackend)
logger.info("Using Sliding Tile Attention backend.")
return "fastvideo.v1.attention.backends.sliding_tile_attn.SlidingTileAttentionBackend"
return "fastvideo.v1.layers.attention.backends.sliding_tile_attn.SlidingTileAttentionBackend"
except ImportError as e:
logger.info(e)
logger.info(
@@ -130,10 +130,10 @@ class CudaPlatformBase(Platform):
try:
from sageattention import sageattn # noqa: F401
from fastvideo.v1.attention.backends.sage_attn import ( # noqa: F401
from fastvideo.v1.layers.attention.backends.sage_attn import ( # noqa: F401
SageAttentionBackend)
logger.info("Using Sage Attention backend.")
return "fastvideo.v1.attention.backends.sage_attn.SageAttentionBackend"
return "fastvideo.v1.layers.attention.backends.sage_attn.SageAttentionBackend"
except ImportError as e:
logger.info(e)
logger.info(
@@ -143,10 +143,10 @@ class CudaPlatformBase(Platform):
try:
from vsa import block_sparse_attn # noqa: F401
from fastvideo.v1.attention.backends.video_sparse_attn import ( # noqa: F401
from fastvideo.v1.layers.attention.backends.video_sparse_attn import ( # noqa: F401
VideoSparseAttentionBackend)
logger.info("Using Video Sparse Attention backend.")
return "fastvideo.v1.attention.backends.video_sparse_attn.VideoSparseAttentionBackend"
return "fastvideo.v1.layers.attention.backends.video_sparse_attn.VideoSparseAttentionBackend"
except ImportError as e:
logger.info(e)
logger.info(
@@ -154,7 +154,7 @@ class CudaPlatformBase(Platform):
)
elif selected_backend == _Backend.TORCH_SDPA:
logger.info("Using Torch SDPA backend.")
return "fastvideo.v1.attention.backends.sdpa.SDPABackend"
return "fastvideo.v1.layers.attention.backends.sdpa.SDPABackend"
elif selected_backend == _Backend.FLASH_ATTN or selected_backend is None:
pass
elif selected_backend:
@@ -178,7 +178,7 @@ class CudaPlatformBase(Platform):
try:
import flash_attn # noqa: F401
from fastvideo.v1.attention.backends.flash_attn import ( # noqa: F401
from fastvideo.v1.layers.attention.backends.flash_attn import ( # noqa: F401
FlashAttentionBackend)
supported_sizes = \
@@ -197,10 +197,10 @@ class CudaPlatformBase(Platform):
if target_backend == _Backend.TORCH_SDPA:
logger.info("Using Torch SDPA backend.")
return "fastvideo.v1.attention.backends.sdpa.SDPABackend"
return "fastvideo.v1.layers.attention.backends.sdpa.SDPABackend"
logger.info("Using Flash Attention backend.")
return "fastvideo.v1.attention.backends.flash_attn.FlashAttentionBackend"
return "fastvideo.v1.layers.attention.backends.flash_attn.FlashAttentionBackend"
@classmethod
def get_device_communicator_cls(cls) -> str: