Pai sparse attention (#211)

* pai fuser sparse test

* make pai fuser more clear

* make pai fuser more clear

* make pai fuser more clear

* make pai fuser more clear

* Update Readme

* disable compile in rope for less error info

* Fix sage attention backward bug

* Fix checkpoint bugs

* Fix Attention

* Update predict and fast rope

* Fix bug in Sparse

* Fix bug in Sparse

* Fix bug in loras load

* Delete useless import

* Fix bug in ui

* Fix bug in ui
This commit is contained in:
Bubbliiiing
2025-05-26 16:00:58 +08:00
committed by GitHub
parent 87125a4af0
commit a1e6ea6335
29 changed files with 298 additions and 203 deletions
+2 -2
View File
@@ -554,8 +554,8 @@ V1.1:
| Wan2.1-Fun-V1.1-14B-InP | 47.0 GB | [🤗Link](https://huggingface.co/alibaba-pai/Wan2.1-Fun-V1.1-14B-InP) | [😄Link](https://modelscope.cn/models/PAI/Wan2.1-Fun-V1.1-14B-InP) | Wan2.1-Fun-V1.1-14B text-to-video generation weights, trained at multiple resolutions, supports start-end image prediction. |
| Wan2.1-Fun-V1.1-1.3B-Control | 19.0 GB | [🤗Link](https://huggingface.co/alibaba-pai/Wan2.1-Fun-V1.1-1.3B-Control) | [😄Link](https://modelscope.cn/models/PAI/Wan2.1-Fun-V1.1-1.3B-Control) | Wan2.1-Fun-V1.1-1.3B video control weights support various control conditions such as Canny, Depth, Pose, MLSD, etc., supports reference image + control condition-based control, and trajectory control. Supports multi-resolution (512, 768, 1024) video prediction, trained with 81 frames at 16 FPS, supports multilingual prediction. |
| Wan2.1-Fun-V1.1-14B-Control | 47.0 GB | [🤗Link](https://huggingface.co/alibaba-pai/Wan2.1-Fun-V1.1-14B-Control) | [😄Link](https://modelscope.cn/models/PAI/Wan2.1-Fun-V1.1-14B-Control) | Wan2.1-Fun-V1.1-14B video control weights support various control conditions such as Canny, Depth, Pose, MLSD, etc., supports reference image + control condition-based control, and trajectory control. Supports multi-resolution (512, 768, 1024) video prediction, trained with 81 frames at 16 FPS, supports multilingual prediction. |
| Wan2.1-Fun-V1.1-1.3B-Control-Camera | 19.0 GB | [🤗Link](https://huggingface.co/alibaba-pai/Wan2.1-Fun-V1.1-1.3B-Control) | [😄Link](https://modelscope.cn/models/PAI/Wan2.1-Fun-V1.1-1.3B-Control) | Wan2.1-Fun-V1.1-1.3B camera lens control weights. Supports multi-resolution (512, 768, 1024) video prediction, trained with 81 frames at 16 FPS, supports multilingual prediction. |
| Wan2.1-Fun-V1.1-14B-Control-Camera | 47.0 GB | [🤗Link](https://huggingface.co/alibaba-pai/Wan2.1-Fun-V1.1-14B-Control) | [😄Link](https://modelscope.cn/models/PAI/Wan2.1-Fun-V1.1-14B-Control) | Wan2.1-Fun-V1.1-14B camera lens control weights. Supports multi-resolution (512, 768, 1024) video prediction, trained with 81 frames at 16 FPS, supports multilingual prediction. |
| Wan2.1-Fun-V1.1-1.3B-Control-Camera | 19.0 GB | [🤗Link](https://huggingface.co/alibaba-pai/Wan2.1-Fun-V1.1-1.3B-Control-Camera) | [😄Link](https://modelscope.cn/models/PAI/Wan2.1-Fun-V1.1-1.3B-Control-Camera) | Wan2.1-Fun-V1.1-1.3B camera lens control weights. Supports multi-resolution (512, 768, 1024) video prediction, trained with 81 frames at 16 FPS, supports multilingual prediction. |
| Wan2.1-Fun-V1.1-14B-Control-Camera | 47.0 GB | [🤗Link](https://huggingface.co/alibaba-pai/Wan2.1-Fun-V1.1-14B-Control-Camera) | [😄Link](https://modelscope.cn/models/PAI/Wan2.1-Fun-V1.1-14B-Control-Camera) | Wan2.1-Fun-V1.1-14B camera lens control weights. Supports multi-resolution (512, 768, 1024) video prediction, trained with 81 frames at 16 FPS, supports multilingual prediction. |
V1.0:
| Name | Storage Space | Hugging Face | Model Scope | Description |
+2 -2
View File
@@ -552,8 +552,8 @@ V1.1:
| Wan2.1-Fun-V1.1-14B-InP | 47.0 GB | [🤗リンク](https://huggingface.co/alibaba-pai/Wan2.1-Fun-V1.1-14B-InP) | [😄リンク](https://modelscope.cn/models/PAI/Wan2.1-Fun-V1.1-14B-InP) | Wan2.1-Fun-V1.1-14Bのテキスト・画像から動画生成の重み。マルチ解像度で訓練され、最初と最後の画像予測をサポートします。 |
| Wan2.1-Fun-V1.1-1.3B-Control | 19.0 GB | [🤗リンク](https://huggingface.co/alibaba-pai/Wan2.1-Fun-V1.1-1.3B-Control) | [😄リンク](https://modelscope.cn/models/PAI/Wan2.1-Fun-V1.1-1.3B-Control)| Wan2.1-Fun-V1.1-1.3Bのビデオ制御重み。Canny、Depth、Pose、MLSDなどの異なる制御条件に対応し、参照画像+制御条件を使用した制御や軌跡制御をサポートします。512、768、1024のマルチ解像度での動画予測をサポートし、81フレーム、毎秒16フレームで訓練されています。多言語予測に対応しています。 |
| Wan2.1-Fun-V1.1-14B-Control | 47.0 GB | [🤗リンク](https://huggingface.co/alibaba-pai/Wan2.1-Fun-V1.1-14B-Control) | [😄リンク](https://modelscope.cn/models/PAI/Wan2.1-Fun-V1.1-14B-Control)| Wan2.1-Fun-V1.1-14Bのビデオ制御重み。Canny、Depth、Pose、MLSDなどの異なる制御条件に対応し、参照画像+制御条件を使用した制御や軌跡制御をサポートします。512、768、1024のマルチ解像度での動画予測をサポートし、81フレーム、毎秒16フレームで訓練されています。多言語予測に対応しています。 |
| Wan2.1-Fun-V1.1-1.3B-Control-Camera | 19.0 GB | [🤗リンク](https://huggingface.co/alibaba-pai/Wan2.1-Fun-V1.1-1.3B-Control) | [😄リンク](https://modelscope.cn/models/PAI/Wan2.1-Fun-V1.1-1.3B-Control)| Wan2.1-Fun-V1.1-1.3Bのカメラレンズ制御重み。512、768、1024のマルチ解像度での動画予測をサポートし、81フレーム、毎秒16フレームで訓練されています。多言語予測に対応しています。 |
| Wan2.1-Fun-V1.1-14B-Control | 47.0 GB | [🤗リンク](https://huggingface.co/alibaba-pai/Wan2.1-Fun-V1.1-14B-Control) | [😄リンク](https://modelscope.cn/models/PAI/Wan2.1-Fun-V1.1-14B-Control)| Wan2.1-Fun-V1.1-14Bのカメラレンズ制御重み。512、768、1024のマルチ解像度での動画予測をサポートし、81フレーム、毎秒16フレームで訓練されています。多言語予測に対応しています。 |
| Wan2.1-Fun-V1.1-1.3B-Control-Camera | 19.0 GB | [🤗リンク](https://huggingface.co/alibaba-pai/Wan2.1-Fun-V1.1-1.3B-Control-Camera) | [😄リンク](https://modelscope.cn/models/PAI/Wan2.1-Fun-V1.1-1.3B-Control-Camera)| Wan2.1-Fun-V1.1-1.3Bのカメラレンズ制御重み。512、768、1024のマルチ解像度での動画予測をサポートし、81フレーム、毎秒16フレームで訓練されています。多言語予測に対応しています。 |
| Wan2.1-Fun-V1.1-14B-Control-Camera | 47.0 GB | [🤗リンク](https://huggingface.co/alibaba-pai/Wan2.1-Fun-V1.1-14B-Control-Camera) | [😄リンク](https://modelscope.cn/models/PAI/Wan2.1-Fun-V1.1-14B-Control-Camera)| Wan2.1-Fun-V1.1-14Bのカメラレンズ制御重み。512、768、1024のマルチ解像度での動画予測をサポートし、81フレーム、毎秒16フレームで訓練されています。多言語予測に対応しています。 |
V1.0:
+2 -2
View File
@@ -543,8 +543,8 @@ V1.1:
| Wan2.1-Fun-V1.1-14B-InP | 47.0 GB | [🤗Link](https://huggingface.co/alibaba-pai/Wan2.1-Fun-V1.1-14B-InP) | [😄Link](https://modelscope.cn/models/PAI/Wan2.1-Fun-V1.1-14B-InP) | Wan2.1-Fun-V1.1-14B文图生视频权重,以多分辨率训练,支持首尾图预测。 |
| Wan2.1-Fun-V1.1-1.3B-Control | 19.0 GB | [🤗Link](https://huggingface.co/alibaba-pai/Wan2.1-Fun-V1.1-1.3B-Control) | [😄Link](https://modelscope.cn/models/PAI/Wan2.1-Fun-V1.1-1.3B-Control)| Wan2.1-Fun-V1.1-1.3B视频控制权重支持不同的控制条件,如Canny、Depth、Pose、MLSD等,支持参考图 + 控制条件进行控制,支持使用轨迹控制。支持多分辨率(512,768,1024)的视频预测,支持多分辨率(512,768,1024)的视频预测,以81帧、每秒16帧进行训练,支持多语言预测 |
| Wan2.1-Fun-V1.1-14B-Control | 47.0 GB | [🤗Link](https://huggingface.co/alibaba-pai/Wan2.1-Fun-V1.1-14B-Control) | [😄Link](https://modelscope.cn/models/PAI/Wan2.1-Fun-V1.1-14B-Control)| Wan2.1-Fun-V1.1-14B视视频控制权重支持不同的控制条件,如Canny、Depth、Pose、MLSD等,支持参考图 + 控制条件进行控制,支持使用轨迹控制。支持多分辨率(512,768,1024)的视频预测,支持多分辨率(512,768,1024)的视频预测,以81帧、每秒16帧进行训练,支持多语言预测 |
| Wan2.1-Fun-V1.1-1.3B-Control-Camera | 19.0 GB | [🤗Link](https://huggingface.co/alibaba-pai/Wan2.1-Fun-V1.1-1.3B-Control) | [😄Link](https://modelscope.cn/models/PAI/Wan2.1-Fun-V1.1-1.3B-Control)| Wan2.1-Fun-V1.1-1.3B相机镜头控制权重。支持多分辨率(512,768,1024)的视频预测,支持多分辨率(512,768,1024)的视频预测,以81帧、每秒16帧进行训练,支持多语言预测 |
| Wan2.1-Fun-V1.1-14B-Control | 47.0 GB | [🤗Link](https://huggingface.co/alibaba-pai/Wan2.1-Fun-V1.1-14B-Control) | [😄Link](https://modelscope.cn/models/PAI/Wan2.1-Fun-V1.1-14B-Control)| Wan2.1-Fun-V1.1-14B相机镜头控制权重。支持多分辨率(512,768,1024)的视频预测,支持多分辨率(512,768,1024)的视频预测,以81帧、每秒16帧进行训练,支持多语言预测 |
| Wan2.1-Fun-V1.1-1.3B-Control-Camera | 19.0 GB | [🤗Link](https://huggingface.co/alibaba-pai/Wan2.1-Fun-V1.1-1.3B-Control-Camera) | [😄Link](https://modelscope.cn/models/PAI/Wan2.1-Fun-V1.1-1.3B-Control-Camera)| Wan2.1-Fun-V1.1-1.3B相机镜头控制权重。支持多分辨率(512,768,1024)的视频预测,支持多分辨率(512,768,1024)的视频预测,以81帧、每秒16帧进行训练,支持多语言预测 |
| Wan2.1-Fun-V1.1-14B-Control-Camera | 47.0 GB | [🤗Link](https://huggingface.co/alibaba-pai/Wan2.1-Fun-V1.1-14B-Control-Camera) | [😄Link](https://modelscope.cn/models/PAI/Wan2.1-Fun-V1.1-14B-Control-Camera)| Wan2.1-Fun-V1.1-14B相机镜头控制权重。支持多分辨率(512,768,1024)的视频预测,支持多分辨率(512,768,1024)的视频预测,以81帧、每秒16帧进行训练,支持多语言预测 |
V1.0:
| 名称 | 存储空间 | Hugging Face | Model Scope | 描述 |
+1 -1
View File
@@ -95,7 +95,7 @@ device = set_multi_gpus_devices(ulysses_degree, ring_degree)
transformer = CogVideoXTransformer3DModel.from_pretrained(
model_name,
subfolder="transformer",
low_cpu_mem_usage=True if not fsdp_dit else False,
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
).to(weight_dtype)
+1 -1
View File
@@ -87,7 +87,7 @@ device = set_multi_gpus_devices(ulysses_degree, ring_degree)
transformer = CogVideoXTransformer3DModel.from_pretrained(
model_name,
subfolder="transformer",
low_cpu_mem_usage=True if not fsdp_dit else False,
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
).to(weight_dtype)
+1 -1
View File
@@ -94,7 +94,7 @@ device = set_multi_gpus_devices(ulysses_degree, ring_degree)
transformer = CogVideoXTransformer3DModel.from_pretrained(
model_name,
subfolder="transformer",
low_cpu_mem_usage=True if not fsdp_dit else False,
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
).to(weight_dtype)
@@ -90,7 +90,7 @@ device = set_multi_gpus_devices(ulysses_degree, ring_degree)
transformer = CogVideoXTransformer3DModel.from_pretrained(
model_name,
subfolder="transformer",
low_cpu_mem_usage=True if not fsdp_dit else False,
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
).to(weight_dtype)
+1 -1
View File
@@ -122,7 +122,7 @@ config = OmegaConf.load(config_path)
transformer = WanTransformer3DModel.from_pretrained(
os.path.join(model_name, config['transformer_additional_kwargs'].get('transformer_subpath', 'transformer')),
transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']),
low_cpu_mem_usage=True if not fsdp_dit else False,
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
)
+1 -1
View File
@@ -117,7 +117,7 @@ config = OmegaConf.load(config_path)
transformer = WanTransformer3DModel.from_pretrained(
os.path.join(model_name, config['transformer_additional_kwargs'].get('transformer_subpath', 'transformer')),
transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']),
low_cpu_mem_usage=True if not fsdp_dit else False,
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
)
+1 -1
View File
@@ -123,7 +123,7 @@ config = OmegaConf.load(config_path)
transformer = WanTransformer3DModel.from_pretrained(
os.path.join(model_name, config['transformer_additional_kwargs'].get('transformer_subpath', 'transformer')),
transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']),
low_cpu_mem_usage=True if not fsdp_dit else False,
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
)
+1 -1
View File
@@ -118,7 +118,7 @@ config = OmegaConf.load(config_path)
transformer = WanTransformer3DModel.from_pretrained(
os.path.join(model_name, config['transformer_additional_kwargs'].get('transformer_subpath', 'transformer')),
transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']),
low_cpu_mem_usage=True if not fsdp_dit else False,
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
)
+1 -1
View File
@@ -133,7 +133,7 @@ config = OmegaConf.load(config_path)
transformer = WanTransformer3DModel.from_pretrained(
os.path.join(model_name, config['transformer_additional_kwargs'].get('transformer_subpath', 'transformer')),
transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']),
low_cpu_mem_usage=True if not fsdp_dit else False,
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
)
@@ -133,7 +133,7 @@ config = OmegaConf.load(config_path)
transformer = WanTransformer3DModel.from_pretrained(
os.path.join(model_name, config['transformer_additional_kwargs'].get('transformer_subpath', 'transformer')),
transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']),
low_cpu_mem_usage=True if not fsdp_dit else False,
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
)
@@ -133,7 +133,7 @@ config = OmegaConf.load(config_path)
transformer = WanTransformer3DModel.from_pretrained(
os.path.join(model_name, config['transformer_additional_kwargs'].get('transformer_subpath', 'transformer')),
transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']),
low_cpu_mem_usage=True if not fsdp_dit else False,
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
)
+28 -59
View File
@@ -1,64 +1,33 @@
import torch
import torch.distributed as dist
import importlib.util
from .cogvideox_xfuser import CogVideoXMultiGPUsAttnProcessor2_0
from .fsdp import shard_model
from .fuser import (get_sequence_parallel_rank,
get_sequence_parallel_world_size, get_sp_group,
get_world_group, init_distributed_environment,
initialize_model_parallel, set_multi_gpus_devices,
xFuserLongContextAttention)
from .wan_xfuser import usp_attn_forward
try:
try:
import pai_fuser
from pai_fuser.core.distributed import (
get_sequence_parallel_rank, get_sequence_parallel_world_size,
get_sp_group, get_world_group, init_distributed_environment,
initialize_model_parallel)
from pai_fuser.core.long_ctx_attention import \
xFuserLongContextAttention
print("Enable PAI DiT Turbo")
except Exception as ex:
import xfuser
from xfuser.core.distributed import (get_sequence_parallel_rank,
get_sequence_parallel_world_size,
get_sp_group, get_world_group,
init_distributed_environment,
initialize_model_parallel)
from xfuser.core.long_ctx_attention import xFuserLongContextAttention
except Exception as ex:
get_sequence_parallel_world_size = None
get_sequence_parallel_rank = None
xFuserLongContextAttention = None
get_sp_group = None
get_world_group = None
init_distributed_environment = None
initialize_model_parallel = None
try:
# The pai_fuser is an internally developed acceleration package, which can be used on PAI.
if importlib.util.find_spec("pai_fuser") is not None:
from pai_fuser.core import parallel_magvit_vae
print("Enable PAI VAE Turbo")
except:
def parallel_magvit_vae(multi_gpus_overlap_scale, spatial_compression_ratio):
def decorator(func):
def wrapper(self, z, *args, **kwargs):
decoded = func(self, z, *args, **kwargs)
return decoded
return wrapper
return decorator
from pai_fuser.core.attention import wan_usp_sparse_attention_wrapper
from . import wan_xfuser
# The simple_wrapper is used to solve the problem about conflicts between cython and torch.compile
def simple_wrapper(func):
def inner(*args, **kwargs):
return func(*args, **kwargs)
return inner
def set_multi_gpus_devices(ulysses_degree, ring_degree):
if ulysses_degree > 1 or ring_degree > 1:
if get_sp_group is None:
raise RuntimeError("xfuser is not installed.")
dist.init_process_group("nccl")
print('parallel inference enabled: ulysses_degree=%d ring_degree=%d rank=%d world_size=%d' % (
ulysses_degree, ring_degree, dist.get_rank(),
dist.get_world_size()))
assert dist.get_world_size() == ring_degree * ulysses_degree, \
"number of GPUs(%d) should be equal to ring_degree * ulysses_degree." % dist.get_world_size()
init_distributed_environment(rank=dist.get_rank(), world_size=dist.get_world_size())
initialize_model_parallel(sequence_parallel_degree=dist.get_world_size(),
ring_degree=ring_degree,
ulysses_degree=ulysses_degree)
# device = torch.device("cuda:%d" % dist.get_rank())
device = torch.device(f"cuda:{get_world_group().local_rank}")
print('rank=%d device=%s' % (get_world_group().rank, str(device)))
else:
device = "cuda"
return device
wan_xfuser.usp_attn_forward = simple_wrapper(wan_usp_sparse_attention_wrapper()(wan_xfuser.usp_attn_forward))
usp_attn_forward = simple_wrapper(wan_xfuser.usp_attn_forward)
print("Import PAI VAE Turbo and Sparse Attention")
from pai_fuser.core.rope import ENABLE_KERNEL, usp_fast_rope_apply_qk
if ENABLE_KERNEL:
wan_xfuser.rope_apply_qk = usp_fast_rope_apply_qk
rope_apply_qk = usp_fast_rope_apply_qk
print("Import PAI Fast rope")
+1 -1
View File
@@ -5,7 +5,7 @@ import torch.nn.functional as F
from diffusers.models.attention import Attention
from diffusers.models.embeddings import apply_rotary_emb
from ..dist import (get_sequence_parallel_rank,
from .fuser import (get_sequence_parallel_rank,
get_sequence_parallel_world_size, get_sp_group,
init_distributed_environment, initialize_model_parallel,
xFuserLongContextAttention)
Vendored Regular → Executable
+1
View File
@@ -9,6 +9,7 @@ from torch.distributed.fsdp import MixedPrecision, ShardingStrategy
from torch.distributed.fsdp.wrap import lambda_auto_wrap_policy
from torch.distributed.utils import _free_storage
def shard_model(
model,
device_id,
+54
View File
@@ -0,0 +1,54 @@
import importlib.util
import torch
import torch.distributed as dist
try:
# The pai_fuser is an internally developed acceleration package, which can be used on PAI.
if importlib.util.find_spec("pai_fuser") is not None:
import pai_fuser
from pai_fuser.core.distributed import (
get_sequence_parallel_rank, get_sequence_parallel_world_size,
get_sp_group, get_world_group, init_distributed_environment,
initialize_model_parallel)
from pai_fuser.core.long_ctx_attention import \
xFuserLongContextAttention
print("Enable PAI DiT Turbo")
else:
import xfuser
from xfuser.core.distributed import (get_sequence_parallel_rank,
get_sequence_parallel_world_size,
get_sp_group, get_world_group,
init_distributed_environment,
initialize_model_parallel)
from xfuser.core.long_ctx_attention import xFuserLongContextAttention
print("Xfuser import sucessful")
except Exception as ex:
get_sequence_parallel_world_size = None
get_sequence_parallel_rank = None
xFuserLongContextAttention = None
get_sp_group = None
get_world_group = None
init_distributed_environment = None
initialize_model_parallel = None
def set_multi_gpus_devices(ulysses_degree, ring_degree):
if ulysses_degree > 1 or ring_degree > 1:
if get_sp_group is None:
raise RuntimeError("xfuser is not installed.")
dist.init_process_group("nccl")
print('parallel inference enabled: ulysses_degree=%d ring_degree=%d rank=%d world_size=%d' % (
ulysses_degree, ring_degree, dist.get_rank(),
dist.get_world_size()))
assert dist.get_world_size() == ring_degree * ulysses_degree, \
"number of GPUs(%d) should be equal to ring_degree * ulysses_degree." % dist.get_world_size()
init_distributed_environment(rank=dist.get_rank(), world_size=dist.get_world_size())
initialize_model_parallel(sequence_parallel_degree=dist.get_world_size(),
ring_degree=ring_degree,
ulysses_degree=ulysses_degree)
# device = torch.device("cuda:%d" % dist.get_rank())
device = torch.device(f"cuda:{get_world_group().local_rank}")
print('rank=%d device=%s' % (get_world_group().rank, str(device)))
else:
device = "cuda"
return device
+10 -4
View File
@@ -1,7 +1,7 @@
import torch
import torch.cuda.amp as amp
from ..dist import (get_sequence_parallel_rank,
from .fuser import (get_sequence_parallel_rank,
get_sequence_parallel_world_size, get_sp_group,
init_distributed_environment, initialize_model_parallel,
xFuserLongContextAttention)
@@ -20,6 +20,7 @@ def pad_freqs(original_tensor, target_len):
return padded_tensor
@amp.autocast(enabled=False)
@torch.compiler.disable()
def rope_apply(x, grid_sizes, freqs):
"""
x: [B, L, N, C].
@@ -59,12 +60,18 @@ def rope_apply(x, grid_sizes, freqs):
output.append(x_i)
return torch.stack(output)
def rope_apply_qk(q, k, grid_sizes, freqs):
q = rope_apply(q, grid_sizes, freqs)
k = rope_apply(k, grid_sizes, freqs)
return q, k
def usp_attn_forward(self,
x,
seq_lens,
grid_sizes,
freqs,
dtype=torch.bfloat16):
dtype=torch.bfloat16,
t=0):
b, s, n, d = *x.shape[:2], self.num_heads, self.head_dim
half_dtypes = (torch.float16, torch.bfloat16)
@@ -79,8 +86,7 @@ def usp_attn_forward(self,
return q, k, v
q, k, v = qkv_fn(x)
q = rope_apply(q, grid_sizes, freqs)
k = rope_apply(k, grid_sizes, freqs)
q, k = rope_apply_qk(q, k, grid_sizes, freqs)
# TODO: We should use unpaded q,k,v for attention.
# k_lens = seq_lens // get_sequence_parallel_world_size()
Regular → Executable
+36 -2
View File
@@ -4,5 +4,39 @@ from .cogvideox_transformer3d import CogVideoXTransformer3DModel
from .cogvideox_vae import AutoencoderKLCogVideoX
from .wan_image_encoder import CLIPModel
from .wan_text_encoder import WanT5EncoderModel
from .wan_transformer3d import WanTransformer3DModel
from .wan_vae import AutoencoderKLWan
from .wan_transformer3d import WanTransformer3DModel, WanSelfAttention
from .wan_vae import AutoencoderKLWan, AutoencoderKLWan_
import importlib.util
# The pai_fuser is an internally developed acceleration package, which can be used on PAI.
if importlib.util.find_spec("pai_fuser") is not None:
from ..dist import parallel_magvit_vae
AutoencoderKLWan_.decode = parallel_magvit_vae(0.2, 8)(AutoencoderKLWan_.decode)
from pai_fuser.core.attention import wan_sparse_attention_wrapper
import torch
# The simple_wrapper is used to solve the problem about conflicts between cython and torch.compile
def simple_wrapper(func):
def inner(*args, **kwargs):
return func(*args, **kwargs)
return inner
WanSelfAttention.forward = simple_wrapper(wan_sparse_attention_wrapper()(WanSelfAttention.forward))
print("Import Sparse Attention")
import os
from pai_fuser.core import (cfg_skip_turbo, enable_cfg_skip,
disable_cfg_skip)
WanTransformer3DModel.enable_cfg_skip = enable_cfg_skip()(WanTransformer3DModel.enable_cfg_skip)
WanTransformer3DModel.disable_cfg_skip = disable_cfg_skip()(WanTransformer3DModel.disable_cfg_skip)
print("Import CFG Skip Turbo")
from pai_fuser.core.rope import ENABLE_KERNEL, fast_rope_apply_qk
if ENABLE_KERNEL:
wan_transformer3d.rope_apply_qk = fast_rope_apply_qk
rope_apply_qk = fast_rope_apply_qk
print("Import PAI Fast rope")
-61
View File
@@ -1,8 +1,6 @@
import numpy as np
import torch
import importlib.util
def get_teacache_coefficients(model_name):
if "wan2.1-t2v-1.3b" in model_name.lower() or "wan2.1-fun-1.3b" in model_name.lower() or "wan2.1-fun-v1.1-1.3b" in model_name.lower():
return [-5.21862437e+04, 9.23041404e+03, -5.28275948e+02, 1.36987616e+01, -4.99875664e-02]
@@ -73,62 +71,3 @@ class TeaCache():
self.previous_residual = None
self.previous_residual_cond = None
self.previous_residual_uncond = None
if importlib.util.find_spec("pai_fuser") is not None:
from pai_fuser.core import (cfg_skip_turbo, enable_cfg_skip,
disable_cfg_skip)
cfg_skip = cfg_skip_turbo
print("Enable CFG Skip Turbo")
else:
def cfg_skip():
def decorator(func):
def wrapper(self, x, *args, **kwargs):
if self.cfg_skip_ratio is not None and self.current_steps >= self.num_inference_steps * (1 - self.cfg_skip_ratio):
bs = len(x)
bs_half = int(bs // 2)
new_x = x[bs_half:]
new_args = []
for arg in args:
if isinstance(arg, (torch.Tensor, list, tuple, np.ndarray)):
new_args.append(arg[bs_half:])
else:
new_args.append(arg)
new_kwargs = {}
for key, content in kwargs.items():
if isinstance(content, (torch.Tensor, list, tuple, np.ndarray)):
new_kwargs[key] = content[bs_half:]
else:
new_kwargs[key] = content
else:
new_x = x
new_args = args
new_kwargs = kwargs
result = func(self, new_x, *new_args, **new_kwargs)
if self.cfg_skip_ratio is not None and self.current_steps >= self.num_inference_steps * (1 - self.cfg_skip_ratio):
result = torch.cat([result, result], dim=0)
return result
return wrapper
return decorator
def enable_cfg_skip():
def decorator(func):
def wrapper(self, cfg_skip_ratio, num_steps, *args, **kwargs):
func(self, cfg_skip_ratio, num_steps, *args, **kwargs)
return
return wrapper
return decorator
def disable_cfg_skip():
def decorator(func):
def wrapper(self, *args, **kwargs):
func(self, *args, **kwargs)
return
return wrapper
return decorator
+33 -15
View File
@@ -21,9 +21,9 @@ from torch import nn
from ..dist import (get_sequence_parallel_rank,
get_sequence_parallel_world_size, get_sp_group,
xFuserLongContextAttention)
from ..dist.wan_xfuser import usp_attn_forward
from .cache_utils import TeaCache, cfg_skip, disable_cfg_skip, enable_cfg_skip
usp_attn_forward, xFuserLongContextAttention)
from ..utils import cfg_skip
from .cache_utils import TeaCache
from .wan_camera_adapter import SimpleAdapter
try:
@@ -185,6 +185,9 @@ def attention(
fa_version=None,
):
attention_type = os.environ.get("VIDEOX_ATTENTION_TYPE", "FLASH_ATTENTION")
if torch.is_grad_enabled() and attention_type == "SAGE_ATTENTION":
attention_type = "FLASH_ATTENTION"
if attention_type == "SAGE_ATTENTION" and SAGE_ATTENTION_AVAILABLE:
if q_lens is not None or k_lens is not None:
warnings.warn(
@@ -336,6 +339,7 @@ def get_resize_crop_region_for_grid(src, tgt_width, tgt_height):
return (crop_top, crop_left), (crop_top + resize_height, crop_left + resize_width)
@amp.autocast(enabled=False)
@torch.compiler.disable()
def rope_apply(x, grid_sizes, freqs):
n, c = x.size(2), x.size(3) // 2
@@ -366,6 +370,12 @@ def rope_apply(x, grid_sizes, freqs):
return torch.stack(output).float()
def rope_apply_qk(q, k, grid_sizes, freqs):
q = rope_apply(q, grid_sizes, freqs)
k = rope_apply(k, grid_sizes, freqs)
return q, k
class WanRMSNorm(nn.Module):
def __init__(self, dim, eps=1e-5):
@@ -423,7 +433,7 @@ class WanSelfAttention(nn.Module):
self.norm_q = WanRMSNorm(dim, eps=eps) if qk_norm else nn.Identity()
self.norm_k = WanRMSNorm(dim, eps=eps) if qk_norm else nn.Identity()
def forward(self, x, seq_lens, grid_sizes, freqs, dtype):
def forward(self, x, seq_lens, grid_sizes, freqs, dtype, t):
r"""
Args:
x(Tensor): Shape [B, L, num_heads, C / num_heads]
@@ -442,9 +452,11 @@ class WanSelfAttention(nn.Module):
q, k, v = qkv_fn(x)
q, k = rope_apply_qk(q, k, grid_sizes, freqs)
x = attention(
q=rope_apply(q, grid_sizes, freqs).to(dtype),
k=rope_apply(k, grid_sizes, freqs).to(dtype),
q.to(dtype),
k.to(dtype),
v=v.to(dtype),
k_lens=seq_lens,
window_size=self.window_size)
@@ -458,7 +470,7 @@ class WanSelfAttention(nn.Module):
class WanT2VCrossAttention(WanSelfAttention):
def forward(self, x, context, context_lens, dtype):
def forward(self, x, context, context_lens, dtype, t):
r"""
Args:
x(Tensor): Shape [B, L1, C]
@@ -502,7 +514,7 @@ class WanI2VCrossAttention(WanSelfAttention):
# self.alpha = nn.Parameter(torch.zeros((1, )))
self.norm_k_img = WanRMSNorm(dim, eps=eps) if qk_norm else nn.Identity()
def forward(self, x, context, context_lens, dtype):
def forward(self, x, context, context_lens, dtype, t):
r"""
Args:
x(Tensor): Shape [B, L1, C]
@@ -599,7 +611,8 @@ class WanAttentionBlock(nn.Module):
freqs,
context,
context_lens,
dtype=torch.float32
dtype=torch.float32,
t=0,
):
r"""
Args:
@@ -615,13 +628,13 @@ class WanAttentionBlock(nn.Module):
temp_x = self.norm1(x) * (1 + e[1]) + e[0]
temp_x = temp_x.to(dtype)
y = self.self_attn(temp_x, seq_lens, grid_sizes, freqs, dtype)
y = self.self_attn(temp_x, seq_lens, grid_sizes, freqs, dtype, t=t)
x = x + y * e[2]
# cross-attention & ffn function
def cross_attn_ffn(x, context, context_lens, e):
# cross-attention
x = x + self.cross_attn(self.norm3(x), context, context_lens, dtype)
x = x + self.cross_attn(self.norm3(x), context, context_lens, dtype, t=t)
# ffn function
temp_x = self.norm2(x) * (1 + e[4]) + e[3]
@@ -789,6 +802,9 @@ class WanTransformer3DModel(ModelMixin, ConfigMixin, FromOriginalModelMixin):
window_size, qk_norm, cross_attn_norm, eps)
for _ in range(num_layers)
])
for layer_idx, block in enumerate(self.blocks):
block.self_attn.layer_idx = layer_idx
block.self_attn.num_layers = self.num_layers
# head
self.head = Head(dim, out_dim, patch_size, eps)
@@ -843,7 +859,6 @@ class WanTransformer3DModel(ModelMixin, ConfigMixin, FromOriginalModelMixin):
def disable_teacache(self):
self.teacache = None
@enable_cfg_skip()
def enable_cfg_skip(self, cfg_skip_ratio, num_steps):
if cfg_skip_ratio != 0:
self.cfg_skip_ratio = cfg_skip_ratio
@@ -854,7 +869,6 @@ class WanTransformer3DModel(ModelMixin, ConfigMixin, FromOriginalModelMixin):
self.current_steps = 0
self.num_inference_steps = None
@disable_cfg_skip()
def disable_cfg_skip(self):
self.cfg_skip_ratio = None
self.current_steps = 0
@@ -1046,6 +1060,7 @@ class WanTransformer3DModel(ModelMixin, ConfigMixin, FromOriginalModelMixin):
context,
context_lens,
dtype,
t,
**ckpt_kwargs,
)
else:
@@ -1057,7 +1072,8 @@ class WanTransformer3DModel(ModelMixin, ConfigMixin, FromOriginalModelMixin):
freqs=self.freqs,
context=context,
context_lens=context_lens,
dtype=dtype
dtype=dtype,
t=t
)
x = block(x, **kwargs)
@@ -1085,6 +1101,7 @@ class WanTransformer3DModel(ModelMixin, ConfigMixin, FromOriginalModelMixin):
context,
context_lens,
dtype,
t,
**ckpt_kwargs,
)
else:
@@ -1096,7 +1113,8 @@ class WanTransformer3DModel(ModelMixin, ConfigMixin, FromOriginalModelMixin):
freqs=self.freqs,
context=context,
context_lens=context_lens,
dtype=dtype
dtype=dtype,
t=t
)
x = block(x, **kwargs)
-3
View File
@@ -14,8 +14,6 @@ from diffusers.models.modeling_utils import ModelMixin
from diffusers.utils.accelerate_utils import apply_forward_hook
from einops import rearrange
from ..dist import parallel_magvit_vae
CACHE_T = 2
@@ -549,7 +547,6 @@ class AutoencoderKLWan_(nn.Module):
self.clear_cache()
return x
@parallel_magvit_vae(0.2, 8)
def decode(self, z, scale):
self.clear_cache()
# z: [b,c,t,h,w]
Regular → Executable
+12 -1
View File
@@ -6,4 +6,15 @@ from .pipeline_wan_fun_inpaint import WanFunInpaintPipeline
from .pipeline_wan_fun_control import WanFunControlPipeline
WanPipeline = WanFunPipeline
WanI2VPipeline = WanFunInpaintPipeline
WanI2VPipeline = WanFunInpaintPipeline
import importlib.util
if importlib.util.find_spec("pai_fuser") is not None:
from pai_fuser.core import sparse_reset
WanFunInpaintPipeline.__call__ = sparse_reset(WanFunInpaintPipeline.__call__)
WanFunPipeline.__call__ = sparse_reset(WanFunPipeline.__call__)
WanFunControlPipeline.__call__ = sparse_reset(WanFunControlPipeline.__call__)
WanI2VPipeline.__call__ = sparse_reset(WanI2VPipeline.__call__)
WanPipeline.__call__ = sparse_reset(WanPipeline.__call__)
+2 -2
View File
@@ -83,7 +83,7 @@ def create_finetune_models_checkpoints(controller, visible):
with gr.Row(visible=visible):
base_model_dropdown = gr.Dropdown(
label="Select base Dreambooth model (选择基模型[非必需])",
choices=controller.personalized_model_list,
choices=["none"] + controller.personalized_model_list,
value="none",
interactive=True,
)
@@ -143,7 +143,7 @@ def create_teacache_params(
):
enable_teacache = gr.Checkbox(label="Enable TeaCache", value=enable_teacache)
teacache_threshold = gr.Slider(0.00, 0.25, value=teacache_threshold, step=0.01, label="TeaCache Threshold")
num_skip_start_steps = gr.Slider(0, 10, value=num_skip_start_steps, step=1, label="Number of Skip Start Steps")
num_skip_start_steps = gr.Slider(0, 10, value=num_skip_start_steps, step=5, label="Number of Skip Start Steps")
teacache_offload = gr.Checkbox(label="Offload TeaCache to CPU", value=teacache_offload)
return enable_teacache, teacache_threshold, num_skip_start_steps, teacache_offload
Regular → Executable
+32
View File
@@ -0,0 +1,32 @@
import importlib.util
from .fm_solvers import FlowDPMSolverMultistepScheduler
from .fm_solvers_unipc import FlowUniPCMultistepScheduler
from .fp8_optimization import (autocast_model_forward,
convert_model_weight_to_float8,
convert_weight_dtype_wrapper,
replace_parameters_by_name)
from .lora_utils import merge_lora, unmerge_lora
from .utils import (filter_kwargs, get_image_latent, get_image_to_video_latent,
get_video_to_video_latent, save_videos_grid)
from .cfg_optimization import cfg_skip
from .discrete_sampler import DiscreteSampling
# The pai_fuser is an internally developed acceleration package, which can be used on PAI.
if importlib.util.find_spec("pai_fuser") is not None:
from pai_fuser.core import (convert_model_weight_to_float8,
convert_weight_dtype_wrapper)
from . import fp8_optimization
fp8_optimization.convert_model_weight_to_float8 = convert_model_weight_to_float8
fp8_optimization.convert_weight_dtype_wrapper = convert_weight_dtype_wrapper
convert_model_weight_to_float8 = fp8_optimization.convert_model_weight_to_float8
convert_weight_dtype_wrapper = fp8_optimization.convert_weight_dtype_wrapper
print("Import PAI Quantization Turbo")
from pai_fuser.core import (cfg_skip_turbo, enable_cfg_skip,
disable_cfg_skip)
from . import cfg_optimization
cfg_optimization.cfg_skip = cfg_skip_turbo
cfg_skip = cfg_skip_turbo
print("Import CFG Skip Turbo")
+39
View File
@@ -0,0 +1,39 @@
import numpy as np
import torch
def cfg_skip():
def decorator(func):
def wrapper(self, x, *args, **kwargs):
if self.cfg_skip_ratio is not None and self.current_steps >= self.num_inference_steps * (1 - self.cfg_skip_ratio):
bs = len(x)
bs_half = int(bs // 2)
new_x = x[bs_half:]
new_args = []
for arg in args:
if isinstance(arg, (torch.Tensor, list, tuple, np.ndarray)):
new_args.append(arg[bs_half:])
else:
new_args.append(arg)
new_kwargs = {}
for key, content in kwargs.items():
if isinstance(content, (torch.Tensor, list, tuple, np.ndarray)):
new_kwargs[key] = content[bs_half:]
else:
new_kwargs[key] = content
else:
new_x = x
new_args = args
new_kwargs = kwargs
result = func(self, new_x, *new_args, **new_kwargs)
if self.cfg_skip_ratio is not None and self.current_steps >= self.num_inference_steps * (1 - self.cfg_skip_ratio):
result = torch.cat([result, result], dim=0)
return result
return wrapper
return decorator
+31 -36
View File
@@ -16,48 +16,43 @@ def replace_parameters_by_name(module, name_keywords, device):
for child_name, child_module in module.named_children():
replace_parameters_by_name(child_module, name_keywords, device)
if importlib.util.find_spec("pai_fuser") is not None:
from pai_fuser.core import (convert_model_weight_to_float8,
convert_weight_dtype_wrapper)
print("Enable PAI Quantization Turbo")
else:
def convert_model_weight_to_float8(model, exclude_module_name=['embed_tokens'], device=None):
for name, module in model.named_modules():
def convert_model_weight_to_float8(model, exclude_module_name=['embed_tokens'], device=None):
for name, module in model.named_modules():
flag = False
for _exclude_module_name in exclude_module_name:
if _exclude_module_name in name:
flag = True
if flag:
continue
for param_name, param in module.named_parameters():
flag = False
for _exclude_module_name in exclude_module_name:
if _exclude_module_name in name:
if _exclude_module_name in param_name:
flag = True
if flag:
continue
for param_name, param in module.named_parameters():
flag = False
for _exclude_module_name in exclude_module_name:
if _exclude_module_name in param_name:
flag = True
if flag:
continue
param.data = param.data.to(torch.float8_e4m3fn)
param.data = param.data.to(torch.float8_e4m3fn)
def autocast_model_forward(cls, origin_dtype, *inputs, **kwargs):
weight_dtype = cls.weight.dtype
cls.to(origin_dtype)
def autocast_model_forward(cls, origin_dtype, *inputs, **kwargs):
weight_dtype = cls.weight.dtype
cls.to(origin_dtype)
# Convert all inputs to the original dtype
inputs = [input.to(origin_dtype) for input in inputs]
out = cls.original_forward(*inputs, **kwargs)
# Convert all inputs to the original dtype
inputs = [input.to(origin_dtype) for input in inputs]
out = cls.original_forward(*inputs, **kwargs)
cls.to(weight_dtype)
return out
cls.to(weight_dtype)
return out
def convert_weight_dtype_wrapper(module, origin_dtype):
for name, module in module.named_modules():
if name == "" or "embed_tokens" in name:
continue
original_forward = module.forward
if hasattr(module, "weight") and module.weight is not None:
setattr(module, "original_forward", original_forward)
setattr(
module,
"forward",
lambda *inputs, m=module, **kwargs: autocast_model_forward(m, origin_dtype, *inputs, **kwargs)
)
def convert_weight_dtype_wrapper(module, origin_dtype):
for name, module in module.named_modules():
if name == "" or "embed_tokens" in name:
continue
original_forward = module.forward
if hasattr(module, "weight") and module.weight is not None:
setattr(module, "original_forward", original_forward)
setattr(
module,
"forward",
lambda *inputs, m=module, **kwargs: autocast_model_forward(m, origin_dtype, *inputs, **kwargs)
)
+2 -2
View File
@@ -369,7 +369,7 @@ def merge_lora(pipeline, lora_path, multiplier, device='cpu', dtype=torch.float3
LORA_PREFIX_TRANSFORMER = "lora_unet"
LORA_PREFIX_TEXT_ENCODER = "lora_te"
if state_dict is None:
state_dict = load_file(lora_path, device=device)
state_dict = load_file(lora_path)
else:
state_dict = state_dict
updates = defaultdict(dict)
@@ -447,7 +447,7 @@ def unmerge_lora(pipeline, lora_path, multiplier=1, device="cpu", dtype=torch.fl
"""Unmerge state_dict in LoRANetwork from the pipeline in diffusers."""
LORA_PREFIX_UNET = "lora_unet"
LORA_PREFIX_TEXT_ENCODER = "lora_te"
state_dict = load_file(lora_path, device=device)
state_dict = load_file(lora_path)
updates = defaultdict(dict)
for key, value in state_dict.items():