Compare commits
6
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
b118043f55 | ||
|
|
83d45684b6 | ||
|
|
afb16377a8 | ||
|
|
b44c28749c | ||
|
|
a9b9bad9a2 | ||
|
|
1b2b97544d |
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
```
|
||||
|
||||
@@ -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",
|
||||
]
|
||||
@@ -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
|
||||
|
||||
@@ -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",
|
||||
]
|
||||
+3
-4
@@ -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__)
|
||||
+3
-3
@@ -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
-3
@@ -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__)
|
||||
+3
-4
@@ -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
|
||||
|
||||
+3
-4
@@ -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
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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__)
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user