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:
@@ -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
@@ -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
@@ -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 | 描述 |
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
|
||||
Vendored
+28
-59
@@ -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")
|
||||
Vendored
+1
-1
@@ -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)
|
||||
|
||||
+1
@@ -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,
|
||||
|
||||
Vendored
+54
@@ -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
|
||||
Vendored
+10
-4
@@ -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
@@ -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")
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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
@@ -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
@@ -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
@@ -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")
|
||||
@@ -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
|
||||
@@ -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)
|
||||
)
|
||||
@@ -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():
|
||||
|
||||
Reference in New Issue
Block a user