diff --git a/README.md b/README.md index e46756d..278593c 100755 --- a/README.md +++ b/README.md @@ -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 | diff --git a/README_ja-JP.md b/README_ja-JP.md index fbc8283..3bef24a 100755 --- a/README_ja-JP.md +++ b/README_ja-JP.md @@ -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: diff --git a/README_zh-CN.md b/README_zh-CN.md index 5f74c6a..e77ae28 100755 --- a/README_zh-CN.md +++ b/README_zh-CN.md @@ -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 | 描述 | diff --git a/examples/cogvideox_fun/predict_i2v.py b/examples/cogvideox_fun/predict_i2v.py index 9ff718d..0020080 100755 --- a/examples/cogvideox_fun/predict_i2v.py +++ b/examples/cogvideox_fun/predict_i2v.py @@ -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) diff --git a/examples/cogvideox_fun/predict_t2v.py b/examples/cogvideox_fun/predict_t2v.py index 3dcded5..c9030d3 100755 --- a/examples/cogvideox_fun/predict_t2v.py +++ b/examples/cogvideox_fun/predict_t2v.py @@ -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) diff --git a/examples/cogvideox_fun/predict_v2v.py b/examples/cogvideox_fun/predict_v2v.py index f8d27c1..9dc1bde 100755 --- a/examples/cogvideox_fun/predict_v2v.py +++ b/examples/cogvideox_fun/predict_v2v.py @@ -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) diff --git a/examples/cogvideox_fun/predict_v2v_control.py b/examples/cogvideox_fun/predict_v2v_control.py index ed391f8..2f4380b 100755 --- a/examples/cogvideox_fun/predict_v2v_control.py +++ b/examples/cogvideox_fun/predict_v2v_control.py @@ -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) diff --git a/examples/wan2.1/predict_i2v.py b/examples/wan2.1/predict_i2v.py index 3fec9fc..befe23c 100755 --- a/examples/wan2.1/predict_i2v.py +++ b/examples/wan2.1/predict_i2v.py @@ -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, ) diff --git a/examples/wan2.1/predict_t2v.py b/examples/wan2.1/predict_t2v.py index a00a8dc..97acaac 100755 --- a/examples/wan2.1/predict_t2v.py +++ b/examples/wan2.1/predict_t2v.py @@ -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, ) diff --git a/examples/wan2.1_fun/predict_i2v.py b/examples/wan2.1_fun/predict_i2v.py index 4706215..5c1e3de 100755 --- a/examples/wan2.1_fun/predict_i2v.py +++ b/examples/wan2.1_fun/predict_i2v.py @@ -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, ) diff --git a/examples/wan2.1_fun/predict_t2v.py b/examples/wan2.1_fun/predict_t2v.py index 8f7583f..124fbc1 100755 --- a/examples/wan2.1_fun/predict_t2v.py +++ b/examples/wan2.1_fun/predict_t2v.py @@ -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, ) diff --git a/examples/wan2.1_fun/predict_v2v_control.py b/examples/wan2.1_fun/predict_v2v_control.py index 515897a..67ba5b5 100755 --- a/examples/wan2.1_fun/predict_v2v_control.py +++ b/examples/wan2.1_fun/predict_v2v_control.py @@ -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, ) diff --git a/examples/wan2.1_fun/predict_v2v_control_camera.py b/examples/wan2.1_fun/predict_v2v_control_camera.py index 34953bd..f8d8507 100755 --- a/examples/wan2.1_fun/predict_v2v_control_camera.py +++ b/examples/wan2.1_fun/predict_v2v_control_camera.py @@ -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, ) diff --git a/examples/wan2.1_fun/predict_v2v_control_ref.py b/examples/wan2.1_fun/predict_v2v_control_ref.py index 8c269d9..8599274 100755 --- a/examples/wan2.1_fun/predict_v2v_control_ref.py +++ b/examples/wan2.1_fun/predict_v2v_control_ref.py @@ -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, ) diff --git a/videox_fun/dist/__init__.py b/videox_fun/dist/__init__.py index 1ae7b7d..a42b71e 100755 --- a/videox_fun/dist/__init__.py +++ b/videox_fun/dist/__init__.py @@ -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 \ No newline at end of file + 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") \ No newline at end of file diff --git a/videox_fun/dist/cogvideox_xfuser.py b/videox_fun/dist/cogvideox_xfuser.py index 7b29cb7..5ad8105 100755 --- a/videox_fun/dist/cogvideox_xfuser.py +++ b/videox_fun/dist/cogvideox_xfuser.py @@ -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) diff --git a/videox_fun/dist/fsdp.py b/videox_fun/dist/fsdp.py old mode 100644 new mode 100755 index 569621a..555479f --- a/videox_fun/dist/fsdp.py +++ b/videox_fun/dist/fsdp.py @@ -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, diff --git a/videox_fun/dist/fuser.py b/videox_fun/dist/fuser.py new file mode 100644 index 0000000..bb437a7 --- /dev/null +++ b/videox_fun/dist/fuser.py @@ -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 \ No newline at end of file diff --git a/videox_fun/dist/wan_xfuser.py b/videox_fun/dist/wan_xfuser.py index f80a9d8..8bc59c8 100755 --- a/videox_fun/dist/wan_xfuser.py +++ b/videox_fun/dist/wan_xfuser.py @@ -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() diff --git a/videox_fun/models/__init__.py b/videox_fun/models/__init__.py old mode 100644 new mode 100755 index c241355..241e212 --- a/videox_fun/models/__init__.py +++ b/videox_fun/models/__init__.py @@ -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") \ No newline at end of file diff --git a/videox_fun/models/cache_utils.py b/videox_fun/models/cache_utils.py index 8be69b5..198b127 100755 --- a/videox_fun/models/cache_utils.py +++ b/videox_fun/models/cache_utils.py @@ -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 diff --git a/videox_fun/models/wan_transformer3d.py b/videox_fun/models/wan_transformer3d.py index 3b3fb12..2bc5f3d 100755 --- a/videox_fun/models/wan_transformer3d.py +++ b/videox_fun/models/wan_transformer3d.py @@ -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) diff --git a/videox_fun/models/wan_vae.py b/videox_fun/models/wan_vae.py index 08e01d5..cd28cb9 100755 --- a/videox_fun/models/wan_vae.py +++ b/videox_fun/models/wan_vae.py @@ -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] diff --git a/videox_fun/pipeline/__init__.py b/videox_fun/pipeline/__init__.py old mode 100644 new mode 100755 index da98677..fc8afdc --- a/videox_fun/pipeline/__init__.py +++ b/videox_fun/pipeline/__init__.py @@ -6,4 +6,15 @@ from .pipeline_wan_fun_inpaint import WanFunInpaintPipeline from .pipeline_wan_fun_control import WanFunControlPipeline WanPipeline = WanFunPipeline -WanI2VPipeline = WanFunInpaintPipeline \ No newline at end of file +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__) \ No newline at end of file diff --git a/videox_fun/ui/ui.py b/videox_fun/ui/ui.py index 0f84479..e357716 100755 --- a/videox_fun/ui/ui.py +++ b/videox_fun/ui/ui.py @@ -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 diff --git a/videox_fun/utils/__init__.py b/videox_fun/utils/__init__.py old mode 100644 new mode 100755 index e69de29..3b272e4 --- a/videox_fun/utils/__init__.py +++ b/videox_fun/utils/__init__.py @@ -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") \ No newline at end of file diff --git a/videox_fun/utils/cfg_optimization.py b/videox_fun/utils/cfg_optimization.py new file mode 100644 index 0000000..2d409cc --- /dev/null +++ b/videox_fun/utils/cfg_optimization.py @@ -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 \ No newline at end of file diff --git a/videox_fun/utils/fp8_optimization.py b/videox_fun/utils/fp8_optimization.py index c51e54f..0c55cbf 100755 --- a/videox_fun/utils/fp8_optimization.py +++ b/videox_fun/utils/fp8_optimization.py @@ -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) - ) \ No newline at end of file +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) + ) \ No newline at end of file diff --git a/videox_fun/utils/lora_utils.py b/videox_fun/utils/lora_utils.py index f801fa8..585f0a8 100755 --- a/videox_fun/utils/lora_utils.py +++ b/videox_fun/utils/lora_utils.py @@ -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():