From f26f0a809b7b9c61a2f86ba3d89b54de784e8a1f Mon Sep 17 00:00:00 2001
From: Bubbliiiing <47347516+bubbliiiing@users.noreply.github.com>
Date: Fri, 9 May 2025 18:11:53 +0800
Subject: [PATCH] Update cfg skip to wrapper && Update Teacache && Update
Reamde (#200)
---
README.md | 29 +++
README_ja-JP.md | 29 +++
README_zh-CN.md | 23 +++
examples/cogvideox_fun/app.py | 7 +-
examples/cogvideox_fun/launch_api.py | 29 ++-
examples/cogvideox_fun/predict_i2v.py | 22 ++-
examples/cogvideox_fun/predict_t2v.py | 22 ++-
examples/cogvideox_fun/predict_v2v.py | 24 ++-
examples/cogvideox_fun/predict_v2v_control.py | 22 ++-
examples/wan2.1/app.py | 7 +-
examples/wan2.1/launch_api.py | 28 ++-
examples/wan2.1/predict_i2v.py | 18 +-
examples/wan2.1/predict_t2v.py | 18 +-
examples/wan2.1_fun/app.py | 7 +-
examples/wan2.1_fun/launch_api.py | 28 ++-
examples/wan2.1_fun/predict_i2v.py | 18 +-
examples/wan2.1_fun/predict_t2v.py | 23 ++-
examples/wan2.1_fun/predict_v2v_control.py | 18 +-
.../wan2.1_fun/predict_v2v_control_camera.py | 18 +-
.../wan2.1_fun/predict_v2v_control_ref.py | 18 +-
scripts/wan2.1_fun/README_TRAIN.md | 6 +-
scripts/wan2.1_fun/README_TRAIN_CONTROL.md | 182 +++++++++++++++++-
.../wan2.1_fun/README_TRAIN_CONTROL_LORA.md | 175 ++++++++++++++++-
scripts/wan2.1_fun/README_TRAIN_LORA.md | 6 +-
scripts/wan2.1_fun/train.sh | 4 +-
scripts/wan2.1_fun/train_control.sh | 8 +-
scripts/wan2.1_fun/train_control_lora.sh | 5 +-
scripts/wan2.1_fun/train_lora.sh | 4 +-
scripts/wan2.1_fun/train_reward_lora.sh | 2 +-
videox_fun/api/api_multi_nodes.py | 59 ++++--
videox_fun/dist/__init__.py | 2 +
videox_fun/models/cache_utils.py | 66 ++++++-
videox_fun/models/wan_transformer3d.py | 37 +++-
videox_fun/pipeline/pipeline_wan_fun.py | 6 +-
.../pipeline/pipeline_wan_fun_control.py | 6 +-
.../pipeline/pipeline_wan_fun_inpaint.py | 6 +-
videox_fun/ui/cogvideox_fun_ui.py | 46 +++--
videox_fun/ui/controller.py | 15 +-
videox_fun/ui/wan_fun_ui.py | 42 ++--
videox_fun/ui/wan_ui.py | 42 ++--
videox_fun/utils/fp8_optimization.py | 75 ++++----
41 files changed, 1008 insertions(+), 194 deletions(-)
mode change 100644 => 100755 scripts/wan2.1_fun/README_TRAIN_CONTROL.md
mode change 100644 => 100755 scripts/wan2.1_fun/README_TRAIN_CONTROL_LORA.md
mode change 100644 => 100755 scripts/wan2.1_fun/train.sh
mode change 100644 => 100755 scripts/wan2.1_fun/train_control.sh
mode change 100644 => 100755 scripts/wan2.1_fun/train_control_lora.sh
mode change 100644 => 100755 scripts/wan2.1_fun/train_lora.sh
mode change 100644 => 100755 scripts/wan2.1_fun/train_reward_lora.sh
mode change 100644 => 100755 videox_fun/utils/fp8_optimization.py
diff --git a/README.md b/README.md
index c9bac9c..e46756d 100755
--- a/README.md
+++ b/README.md
@@ -407,6 +407,9 @@ Since Wan2.1 has a very large number of parameters, we need to consider memory o
For details, refer to [ComfyUI README](comfyui/README.md).
#### c. Running Python Files
+
+##### i. Single-GPU Inference:
+
- **Step 1**: Download the corresponding [weights](#model-zoo) and place them in the `models` folder.
- **Step 2**: Use different files for prediction based on the weights and prediction goals. This library currently supports CogVideoX-Fun, Wan2.1, and Wan2.1-Fun. Different models are distinguished by folder names under the `examples` folder, and their supported features vary. Use them accordingly. Below is an example using CogVideoX-Fun:
- **Text-to-Video**:
@@ -426,6 +429,32 @@ For details, refer to [ComfyUI README](comfyui/README.md).
- Run the file `examples/cogvideox_fun/predict_v2v_control.py` and wait for the results. The generated videos will be saved in the folder `samples/cogvideox-fun-videos_v2v_control`.
- **Step 3**: If you want to integrate other backbones or Loras trained by yourself, modify `lora_path` and relevant paths in `examples/{model_name}/predict_t2v.py` or `examples/{model_name}/predict_i2v.py` as needed.
+##### ii. Multi-GPU Inference:
+When using multi-GPU inference, please make sure to install the xfuser. We recommend installing xfuser==0.4.2 and yunchang==0.6.2.
+```
+pip install xfuser==0.4.2 --progress-bar off -i https://mirrors.aliyun.com/pypi/simple/
+pip install yunchang==0.6.2 --progress-bar off -i https://mirrors.aliyun.com/pypi/simple/
+```
+
+Please ensure that the product of `ulysses_degree` and `ring_degree` equals the number of GPUs being used. For example, if you are using 8 GPUs, you can set `ulysses_degree=2` and `ring_degree=4`, or alternatively `ulysses_degree=4` and `ring_degree=2`.
+
+- `ulysses_degree` performs parallelization after splitting across the heads.
+- `ring_degree` performs parallelization after splitting across the sequence.
+
+Compared to `ulysses_degree`, `ring_degree` incurs higher communication costs. Therefore, when setting these parameters, you should take into account both the sequence length and the number of heads in the model.
+
+Let’s take 8-GPU parallel inference as an example:
+
+- **For Wan2.1-Fun-V1.1-14B-InP**, which has 40 heads, `ulysses_degree` should be set to a divisor of 40 (e.g., 2, 4, 8, etc.). Thus, when using 8 GPUs for parallel inference, you can set `ulysses_degree=8` and `ring_degree=1`.
+
+- **For Wan2.1-Fun-V1.1-1.3B-InP**, which has 12 heads, `ulysses_degree` should be set to a divisor of 12 (e.g., 2, 4, etc.). Thus, when using 8 GPUs for parallel inference, you can set `ulysses_degree=4` and `ring_degree=2`.
+
+After setting the parameters, run the following command for parallel inference:
+
+```sh
+torchrun --nproc-per-node=8 examples/wan2.1_fun/predict_t2v.py
+```
+
#### d. Using the Web UI
The web UI supports text-to-video, image-to-video, video-to-video, and controlled video generation (Canny, Pose, Depth, etc.). This library currently supports CogVideoX-Fun, Wan2.1, and Wan2.1-Fun. Different models are distinguished by folder names under the `examples` folder, and their supported features vary. Use them accordingly. Below is an example using CogVideoX-Fun:
diff --git a/README_ja-JP.md b/README_ja-JP.md
index 827b3cc..fbc8283 100755
--- a/README_ja-JP.md
+++ b/README_ja-JP.md
@@ -407,6 +407,9 @@ Wan2.1のパラメータが非常に大きいため、GPUメモリを節約し
詳細は[ComfyUI README](comfyui/README.md)をご覧ください。
#### c. Pythonファイルを実行する
+
+##### i. 単一GPUでの推論:
+
- ステップ1: 対応する[重み](#model-zoo)をダウンロードし、`models`フォルダに配置します。
- ステップ2: 異なる重みと予測目標に基づいて、異なるファイルを使用して予測を行います。現在、このライブラリはCogVideoX-Fun、Wan2.1、およびWan2.1-Funをサポートしています。`examples`フォルダ内のフォルダ名で区別され、異なるモデルがサポートする機能が異なりますので、状況に応じて区別してください。以下はCogVideoX-Funを例として説明します。
- テキストからビデオ:
@@ -426,6 +429,32 @@ Wan2.1のパラメータが非常に大きいため、GPUメモリを節約し
- 次に、`examples/cogvideox_fun/predict_v2v_control.py`ファイルを実行し、結果が生成されるのを待ちます。結果は`samples/cogvideox-fun-videos_v2v_control`フォルダに保存されます。
- ステップ3: 自分でトレーニングした他のバックボーンやLoraを組み合わせたい場合は、必要に応じて`examples/{model_name}/predict_t2v.py`や`examples/{model_name}/predict_i2v.py`、`lora_path`を修正します。
+##### ii. 複数GPUでの推論:
+多カードでの推論を行う際は、xfuserリポジトリのインストールに注意してください。xfuser==0.4.2 と yunchang==0.6.2 のインストールが推奨されます。
+```
+pip install xfuser==0.4.2 --progress-bar off -i https://mirrors.aliyun.com/pypi/simple/
+pip install yunchang==0.6.2 --progress-bar off -i https://mirrors.aliyun.com/pypi/simple/
+```
+
+`ulysses_degree` と `ring_degree` の積が使用する GPU 数と一致することを確認してください。たとえば、8つのGPUを使用する場合、`ulysses_degree=2` と `ring_degree=4`、または `ulysses_degree=4` と `ring_degree=2` を設定することができます。
+
+- `ulysses_degree` はヘッド(head)に分割した後の並列化を行います。
+- `ring_degree` はシーケンスに分割した後の並列化を行います。
+
+`ring_degree` は `ulysses_degree` よりも通信コストが高いため、これらのパラメータを設定する際には、シーケンス長とモデルのヘッド数を考慮する必要があります。
+
+8GPUでの並列推論を例に挙げます:
+
+- **Wan2.1-Fun-V1.1-14B-InP** はヘッド数が40あります。この場合、`ulysses_degree` は40で割り切れる値(例:2, 4, 8など)に設定する必要があります。したがって、8GPUを使用して並列推論を行う場合、`ulysses_degree=8` と `ring_degree=1` を設定できます。
+
+- **Wan2.1-Fun-V1.1-1.3B-InP** はヘッド数が12あります。この場合、`ulysses_degree` は12で割り切れる値(例:2, 4など)に設定する必要があります。したがって、8GPUを使用して並列推論を行う場合、`ulysses_degree=4` と `ring_degree=2` を設定できます。
+
+パラメータの設定が完了したら、以下のコマンドで並列推論を実行してください:
+
+```sh
+torchrun --nproc-per-node=8 examples/wan2.1_fun/predict_t2v.py
+```
+
#### d. UIインターフェースを使用する
WebUIは、テキストからビデオ、画像からビデオ、ビデオからビデオ、および通常の制御付きビデオ生成(Canny、Pose、Depthなど)をサポートします。現在、このライブラリはCogVideoX-Fun、Wan2.1、およびWan2.1-Funをサポートしており、`examples`フォルダ内のフォルダ名で区別されています。異なるモデルがサポートする機能が異なるため、状況に応じて区別してください。以下はCogVideoX-Funを例として説明します。
diff --git a/README_zh-CN.md b/README_zh-CN.md
index 29933c6..5f74c6a 100755
--- a/README_zh-CN.md
+++ b/README_zh-CN.md
@@ -405,6 +405,9 @@ qfloat8会部分降低模型的性能,但可以节省更多的显存。如果
具体查看[ComfyUI README](comfyui/README.md)。
#### c、运行python文件
+
+##### i、单卡运行:
+
- 步骤1:下载对应[权重](#model-zoo)放入models文件夹。
- 步骤2:根据不同的权重与预测目标使用不同的文件进行预测。当前该库支持CogVideoX-Fun、Wan2.1和Wan2.1-Fun,在examples文件夹下用文件夹名以区分,不同模型支持的功能不同,请视具体情况予以区分。以CogVideoX-Fun为例。
- 文生视频:
@@ -424,6 +427,26 @@ qfloat8会部分降低模型的性能,但可以节省更多的显存。如果
- 而后运行examples/cogvideox_fun/predict_v2v_control.py文件,等待生成结果,结果保存在samples/cogvideox-fun-videos_v2v_control文件夹中。
- 步骤3:如果想结合自己训练的其他backbone与Lora,则看情况修改examples/{model_name}/predict_t2v.py中的examples/{model_name}/predict_i2v.py和lora_path。
+##### ii、多卡运行:
+在使用多卡预测时请注意安装xfuser仓库,推荐安装xfuser==0.4.2和yunchang==0.6.2。
+```
+pip install xfuser==0.4.2 --progress-bar off -i https://mirrors.aliyun.com/pypi/simple/
+pip install yunchang==0.6.2 --progress-bar off -i https://mirrors.aliyun.com/pypi/simple/
+```
+
+请确保ulysses_degree和ring_degree的乘积等于使用的GPU数量。例如,如果您使用8个GPU,则可以设置ulysses_degree=2和ring_degree=4,也可以设置ulysses_degree=4和ring_degree=2。
+
+ulysses_degree是在head进行切分后并行生成,ring_degree是在sequence上进行切分后并行生成。ring_degree相比ulysses_degree有更大的通信成本,在设置参数时需要结合序列长度和模型的head数进行设置。
+
+以8卡并行预测为例。
+- 以Wan2.1-Fun-V1.1-14B-InP为例,其head数为40,ulysses_degree需要设置为其可以整除的数如2、4、8等。因此在使用8卡并行预测时,可以设置ulysses_degree=8和ring_degree=1.
+- 以Wan2.1-Fun-V1.1-1.3B-InP为例,其head数为12,ulysses_degree需要设置为其可以整除的数如2、4等。因此在使用8卡并行预测时,可以设置ulysses_degree=4和ring_degree=2.
+
+设置完成后,使用如下指令进行并行预测:
+```sh
+torchrun --nproc-per-node=8 examples/wan2.1_fun/predict_t2v.py
+```
+
#### d、通过ui界面
webui支持文生视频、图生视频、视频生视频和普通控制生视频(Canny、Pose、Depth等)。当前该库支持CogVideoX-Fun、Wan2.1和Wan2.1-Fun,在examples文件夹下用文件夹名以区分,不同模型支持的功能不同,请视具体情况予以区分。以CogVideoX-Fun为例。
diff --git a/examples/cogvideox_fun/app.py b/examples/cogvideox_fun/app.py
index 1e09ada..e520055 100755
--- a/examples/cogvideox_fun/app.py
+++ b/examples/cogvideox_fun/app.py
@@ -33,6 +33,9 @@ if __name__ == "__main__":
# sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
# resulting in slower speeds but saving a large amount of GPU memory.
GPU_memory_mode = "model_cpu_offload_and_qfloat8"
+ # Compile will give a speedup in fixed resolution and need a little GPU memory.
+ # The compile_dit is not compatible with the fsdp_dit and sequential_cpu_offload.
+ compile_dit = False
# Use torch.float16 if GPU does not support torch.bfloat16
# ome graphics cards, such as v100, 2080ti, do not support torch.bfloat16
@@ -48,11 +51,11 @@ if __name__ == "__main__":
model_type = "Inpaint"
if ui_mode == "host":
- demo, controller = ui_host(GPU_memory_mode, ddpm_scheduler_dict, model_name, model_type, 1, 1, weight_dtype)
+ demo, controller = ui_host(GPU_memory_mode, ddpm_scheduler_dict, model_name, model_type, compile_dit, weight_dtype)
elif ui_mode == "client":
demo, controller = ui_client(ddpm_scheduler_dict, model_name)
else:
- demo, controller = ui(GPU_memory_mode, ddpm_scheduler_dict, 1, 1, weight_dtype)
+ demo, controller = ui(GPU_memory_mode, ddpm_scheduler_dict, compile_dit, weight_dtype)
# launch gradio
app, _, _ = demo.queue(status_update_rate=1).launch(
diff --git a/examples/cogvideox_fun/launch_api.py b/examples/cogvideox_fun/launch_api.py
index b7b410f..5f74da9 100755
--- a/examples/cogvideox_fun/launch_api.py
+++ b/examples/cogvideox_fun/launch_api.py
@@ -20,7 +20,29 @@ from videox_fun.ui.cogvideox_fun_ui import CogVideoXFunController
def main():
parser = argparse.ArgumentParser(description='xDiT HTTP Service')
parser.add_argument('--world_size', type=int, default=8, help='Number of parallel workers')
- parser.add_argument('--gpu_memory_mode', type=str, default="model_full_load", help='GPU memory mode')
+ parser.add_argument(
+ '--gpu_memory_mode', type=str, default="model_cpu_offload", help='''
+GPU memory mode, which can be choosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8].
+model_full_load means that the entire model will be moved to the GPU.
+
+model_full_load_and_qfloat8 means that the entire model will be moved to the GPU,
+and the transformer model has been quantized to float8, which can save more GPU memory.
+
+model_cpu_offload means that the entire model will be moved to the CPU after use, which can save some GPU memory.
+
+model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
+and the transformer model has been quantized to float8, which can save more GPU memory.
+ '''
+ )
+ parser.add_argument(
+ '--compile_dit', action='store_true', help='''
+Enable compile dit.
+Compile will give a speedup in fixed resolution and need a little GPU memory.
+The compile_dit is not compatible with the fsdp_dit and sequential_cpu_offload.
+ '''
+ )
+ parser.add_argument('--fsdp_dit', action='store_true', help="Use DIT FSDP to save more GPU memory in multi gpus.")
+ parser.add_argument('--fsdp_text_encoder', action='store_true', help="Use Text Encoder FSDP to save more GPU memory in multi gpus.")
parser.add_argument('--ulysses_degree', type=int, default=4, help='Degree of Ulysses configuration')
parser.add_argument('--ring_degree', type=int, default=2, help='Degree of Ring configuration')
parser.add_argument('--weight_dtype', type=str, default='bf16', help='Weight data type')
@@ -40,8 +62,9 @@ def main():
engine = MultiNodesEngine(
world_size=args.world_size, Controller=CogVideoXFunController,
GPU_memory_mode=args.gpu_memory_mode, scheduler_dict=flow_scheduler_dict, model_name=args.model_name, model_type=args.model_type, config_path=None,
- ulysses_degree=args.ulysses_degree, ring_degree=args.ring_degree, enable_teacache=False, teacache_threshold=0.1, num_skip_start_steps=5,
- teacache_offload=False, weight_dtype=weight_dtype, savedir_sample=args.savedir_sample,
+ ulysses_degree=args.ulysses_degree, ring_degree=args.ring_degree,
+ fsdp_dit=args.fsdp_dit, fsdp_text_encoder=args.fsdp_text_encoder, compile_dit=args.compile_dit,
+ weight_dtype=weight_dtype, savedir_sample=args.savedir_sample,
)
def gr_launch():
diff --git a/examples/cogvideox_fun/predict_i2v.py b/examples/cogvideox_fun/predict_i2v.py
index 9c4f278..9ff718d 100755
--- a/examples/cogvideox_fun/predict_i2v.py
+++ b/examples/cogvideox_fun/predict_i2v.py
@@ -20,7 +20,8 @@ from videox_fun.models import (AutoencoderKLCogVideoX,
T5Tokenizer)
from videox_fun.pipeline import (CogVideoXFunInpaintPipeline,
CogVideoXFunPipeline)
-from videox_fun.utils.fp8_optimization import convert_weight_dtype_wrapper
+from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8, replace_parameters_by_name,
+ convert_weight_dtype_wrapper)
from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
from videox_fun.utils.utils import get_image_to_video_latent, save_videos_grid
@@ -44,7 +45,12 @@ GPU_memory_mode = "model_cpu_offload_and_qfloat8"
# If you are using 1 GPU, you can set ulysses_degree = 1 and ring_degree = 1.
ulysses_degree = 1
ring_degree = 1
+# Use FSDP to save more GPU memory in multi gpus.
fsdp_dit = False
+fsdp_text_encoder = True
+# Compile will give a speedup in fixed resolution and need a little GPU memory.
+# The compile_dit is not compatible with the fsdp_dit and sequential_cpu_offload.
+compile_dit = False
# Config and model path
model_name = "models/Diffusion_Transformer/CogVideoX-Fun-V1.1-2b-InP"
@@ -90,7 +96,7 @@ transformer = CogVideoXTransformer3DModel.from_pretrained(
model_name,
subfolder="transformer",
low_cpu_mem_usage=True if not fsdp_dit else False,
- torch_dtype=torch.float8_e4m3fn if GPU_memory_mode == "model_cpu_offload_and_qfloat8" else weight_dtype,
+ torch_dtype=weight_dtype,
).to(weight_dtype)
if transformer_path is not None:
@@ -167,15 +173,27 @@ if ulysses_degree > 1 or ring_degree > 1:
if fsdp_dit:
shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype)
pipeline.transformer = shard_fn(pipeline.transformer)
+ print("Add FSDP DIT")
+ if fsdp_text_encoder:
+ shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype)
+ pipeline.text_encoder = shard_fn(pipeline.text_encoder)
+ print("Add FSDP TEXT ENCODER")
+
+if compile_dit:
+ for i in range(len(pipeline.transformer.transformer_blocks)):
+ pipeline.transformer.transformer_blocks[i] = torch.compile(pipeline.transformer.transformer_blocks[i])
+ print("Add Compile")
if GPU_memory_mode == "sequential_cpu_offload":
pipeline.enable_sequential_cpu_offload(device=device)
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
+ convert_model_weight_to_float8(transformer, exclude_module_name=[], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_cpu_offload":
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_full_load_and_qfloat8":
+ convert_model_weight_to_float8(transformer, exclude_module_name=[], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
pipeline.to(device=device)
else:
diff --git a/examples/cogvideox_fun/predict_t2v.py b/examples/cogvideox_fun/predict_t2v.py
index f18ad95..3dcded5 100755
--- a/examples/cogvideox_fun/predict_t2v.py
+++ b/examples/cogvideox_fun/predict_t2v.py
@@ -20,7 +20,8 @@ from videox_fun.models import (AutoencoderKLCogVideoX,
T5Tokenizer)
from videox_fun.pipeline import (CogVideoXFunPipeline,
CogVideoXFunInpaintPipeline)
-from videox_fun.utils.fp8_optimization import convert_weight_dtype_wrapper
+from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8, replace_parameters_by_name,
+ convert_weight_dtype_wrapper)
from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
from videox_fun.utils.utils import get_image_to_video_latent, save_videos_grid
from videox_fun.dist import set_multi_gpus_devices, shard_model
@@ -45,7 +46,12 @@ GPU_memory_mode = "model_cpu_offload_and_qfloat8"
# If you are using 1 GPU, you can set ulysses_degree = 1 and ring_degree = 1.
ulysses_degree = 1
ring_degree = 1
+# Use FSDP to save more GPU memory in multi gpus.
fsdp_dit = False
+fsdp_text_encoder = True
+# Compile will give a speedup in fixed resolution and need a little GPU memory.
+# The compile_dit is not compatible with the fsdp_dit and sequential_cpu_offload.
+compile_dit = False
# model path
model_name = "models/Diffusion_Transformer/CogVideoX-Fun-V1.1-2b-InP"
@@ -82,7 +88,7 @@ transformer = CogVideoXTransformer3DModel.from_pretrained(
model_name,
subfolder="transformer",
low_cpu_mem_usage=True if not fsdp_dit else False,
- torch_dtype=torch.float8_e4m3fn if GPU_memory_mode == "model_cpu_offload_and_qfloat8" else weight_dtype,
+ torch_dtype=weight_dtype,
).to(weight_dtype)
if transformer_path is not None:
@@ -159,15 +165,27 @@ if ulysses_degree > 1 or ring_degree > 1:
if fsdp_dit:
shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype)
pipeline.transformer = shard_fn(pipeline.transformer)
+ print("Add FSDP DIT")
+ if fsdp_text_encoder:
+ shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype)
+ pipeline.text_encoder = shard_fn(pipeline.text_encoder)
+ print("Add FSDP TEXT ENCODER")
+
+if compile_dit:
+ for i in range(len(pipeline.transformer.transformer_blocks)):
+ pipeline.transformer.transformer_blocks[i] = torch.compile(pipeline.transformer.transformer_blocks[i])
+ print("Add Compile")
if GPU_memory_mode == "sequential_cpu_offload":
pipeline.enable_sequential_cpu_offload(device=device)
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
+ convert_model_weight_to_float8(transformer, exclude_module_name=[], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_cpu_offload":
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_full_load_and_qfloat8":
+ convert_model_weight_to_float8(transformer, exclude_module_name=[], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
pipeline.to(device=device)
else:
diff --git a/examples/cogvideox_fun/predict_v2v.py b/examples/cogvideox_fun/predict_v2v.py
index 391e27b..f8d27c1 100755
--- a/examples/cogvideox_fun/predict_v2v.py
+++ b/examples/cogvideox_fun/predict_v2v.py
@@ -20,7 +20,8 @@ from videox_fun.models import (AutoencoderKLCogVideoX,
from videox_fun.pipeline import (CogVideoXFunPipeline,
CogVideoXFunInpaintPipeline)
from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
-from videox_fun.utils.fp8_optimization import convert_weight_dtype_wrapper
+from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8, replace_parameters_by_name,
+ convert_weight_dtype_wrapper)
from videox_fun.utils.utils import get_video_to_video_latent, save_videos_grid
from videox_fun.dist import set_multi_gpus_devices, shard_model
@@ -44,7 +45,12 @@ GPU_memory_mode = "model_cpu_offload_and_qfloat8"
# If you are using 1 GPU, you can set ulysses_degree = 1 and ring_degree = 1.
ulysses_degree = 1
ring_degree = 1
+# Use FSDP to save more GPU memory in multi gpus.
fsdp_dit = False
+fsdp_text_encoder = True
+# Compile will give a speedup in fixed resolution and need a little GPU memory.
+# The compile_dit is not compatible with the fsdp_dit and sequential_cpu_offload.
+compile_dit = False
# model path
model_name = "models/Diffusion_Transformer/CogVideoX-Fun-V1.1-2b-InP"
@@ -89,7 +95,7 @@ transformer = CogVideoXTransformer3DModel.from_pretrained(
model_name,
subfolder="transformer",
low_cpu_mem_usage=True if not fsdp_dit else False,
- torch_dtype=torch.float8_e4m3fn if GPU_memory_mode == "model_cpu_offload_and_qfloat8" else weight_dtype,
+ torch_dtype=weight_dtype,
).to(weight_dtype)
if transformer_path is not None:
@@ -166,15 +172,27 @@ if ulysses_degree > 1 or ring_degree > 1:
if fsdp_dit:
shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype)
pipeline.transformer = shard_fn(pipeline.transformer)
+ print("Add FSDP DIT")
+ if fsdp_text_encoder:
+ shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype)
+ pipeline.text_encoder = shard_fn(pipeline.text_encoder)
+ print("Add FSDP TEXT ENCODER")
+
+if compile_dit:
+ for i in range(len(pipeline.transformer.transformer_blocks)):
+ pipeline.transformer.transformer_blocks[i] = torch.compile(pipeline.transformer.transformer_blocks[i])
+ print("Add Compile")
if GPU_memory_mode == "sequential_cpu_offload":
pipeline.enable_sequential_cpu_offload(device=device)
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
+ convert_model_weight_to_float8(transformer, exclude_module_name=[], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_cpu_offload":
pipeline.enable_model_cpu_offload(device=device)
-elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
+elif GPU_memory_mode == "model_full_load_and_qfloat8":
+ convert_model_weight_to_float8(transformer, exclude_module_name=[], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
pipeline.to(device=device)
else:
diff --git a/examples/cogvideox_fun/predict_v2v_control.py b/examples/cogvideox_fun/predict_v2v_control.py
index 4644b3b..ed391f8 100755
--- a/examples/cogvideox_fun/predict_v2v_control.py
+++ b/examples/cogvideox_fun/predict_v2v_control.py
@@ -21,7 +21,8 @@ from videox_fun.models import (AutoencoderKLCogVideoX,
T5Tokenizer)
from videox_fun.pipeline import (CogVideoXFunControlPipeline,
CogVideoXFunInpaintPipeline)
-from videox_fun.utils.fp8_optimization import convert_weight_dtype_wrapper
+from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8, replace_parameters_by_name,
+ convert_weight_dtype_wrapper)
from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
from videox_fun.utils.utils import get_video_to_video_latent, save_videos_grid
from videox_fun.dist import set_multi_gpus_devices, shard_model
@@ -46,7 +47,12 @@ GPU_memory_mode = "model_cpu_offload_and_qfloat8"
# If you are using 1 GPU, you can set ulysses_degree = 1 and ring_degree = 1.
ulysses_degree = 1
ring_degree = 1
+# Use FSDP to save more GPU memory in multi gpus.
fsdp_dit = False
+fsdp_text_encoder = True
+# Compile will give a speedup in fixed resolution and need a little GPU memory.
+# The compile_dit is not compatible with the fsdp_dit and sequential_cpu_offload.
+compile_dit = False
# model path
model_name = "models/Diffusion_Transformer/CogVideoX-Fun-V1.1-2b-Pose"
@@ -85,7 +91,7 @@ transformer = CogVideoXTransformer3DModel.from_pretrained(
model_name,
subfolder="transformer",
low_cpu_mem_usage=True if not fsdp_dit else False,
- torch_dtype=torch.float8_e4m3fn if GPU_memory_mode == "model_cpu_offload_and_qfloat8" else weight_dtype,
+ torch_dtype=weight_dtype,
).to(weight_dtype)
if transformer_path is not None:
@@ -153,15 +159,27 @@ if ulysses_degree > 1 or ring_degree > 1:
if fsdp_dit:
shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype)
pipeline.transformer = shard_fn(pipeline.transformer)
+ print("Add FSDP DIT")
+ if fsdp_text_encoder:
+ shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype)
+ pipeline.text_encoder = shard_fn(pipeline.text_encoder)
+ print("Add FSDP TEXT ENCODER")
+
+if compile_dit:
+ for i in range(len(pipeline.transformer.transformer_blocks)):
+ pipeline.transformer.transformer_blocks[i] = torch.compile(pipeline.transformer.transformer_blocks[i])
+ print("Add Compile")
if GPU_memory_mode == "sequential_cpu_offload":
pipeline.enable_sequential_cpu_offload(device=device)
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
+ convert_model_weight_to_float8(transformer, exclude_module_name=[], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_cpu_offload":
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_full_load_and_qfloat8":
+ convert_model_weight_to_float8(transformer, exclude_module_name=[], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
pipeline.to(device=device)
else:
diff --git a/examples/wan2.1/app.py b/examples/wan2.1/app.py
index 97b4816..a5634a2 100755
--- a/examples/wan2.1/app.py
+++ b/examples/wan2.1/app.py
@@ -33,6 +33,9 @@ if __name__ == "__main__":
# sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
# resulting in slower speeds but saving a large amount of GPU memory.
GPU_memory_mode = "sequential_cpu_offload"
+ # Compile will give a speedup in fixed resolution and need a little GPU memory.
+ # The compile_dit is not compatible with the fsdp_dit and sequential_cpu_offload.
+ compile_dit = False
# Use torch.float16 if GPU does not support torch.bfloat16
# ome graphics cards, such as v100, 2080ti, do not support torch.bfloat16
@@ -49,11 +52,11 @@ if __name__ == "__main__":
model_name = "models/Diffusion_Transformer/Wan2.1-I2V-14B-480P"
if ui_mode == "host":
- demo, controller = ui_host(GPU_memory_mode, flow_scheduler_dict, model_name, "Inpaint", config_path, 1, 1, weight_dtype)
+ demo, controller = ui_host(GPU_memory_mode, flow_scheduler_dict, model_name, "Inpaint", config_path, compile_dit, weight_dtype)
elif ui_mode == "client":
demo, controller = ui_client(flow_scheduler_dict, model_name)
else:
- demo, controller = ui(GPU_memory_mode, flow_scheduler_dict, config_path, 1, 1, weight_dtype)
+ demo, controller = ui(GPU_memory_mode, flow_scheduler_dict, config_path, compile_dit, weight_dtype)
def gr_launch():
# launch gradio
diff --git a/examples/wan2.1/launch_api.py b/examples/wan2.1/launch_api.py
index a632461..73dac80 100755
--- a/examples/wan2.1/launch_api.py
+++ b/examples/wan2.1/launch_api.py
@@ -20,9 +20,31 @@ from videox_fun.ui.wan_ui import Wan_Controller
def main():
parser = argparse.ArgumentParser(description='xDiT HTTP Service')
parser.add_argument('--world_size', type=int, default=8, help='Number of parallel workers')
- parser.add_argument('--gpu_memory_mode', type=str, default="model_full_load", help='GPU memory mode')
+ parser.add_argument(
+ '--gpu_memory_mode', type=str, default="model_cpu_offload", help='''
+GPU memory mode, which can be choosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8].
+model_full_load means that the entire model will be moved to the GPU.
+
+model_full_load_and_qfloat8 means that the entire model will be moved to the GPU,
+and the transformer model has been quantized to float8, which can save more GPU memory.
+
+model_cpu_offload means that the entire model will be moved to the CPU after use, which can save some GPU memory.
+
+model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
+and the transformer model has been quantized to float8, which can save more GPU memory.
+ '''
+ )
parser.add_argument('--ulysses_degree', type=int, default=4, help='Degree of Ulysses configuration')
parser.add_argument('--ring_degree', type=int, default=2, help='Degree of Ring configuration')
+ parser.add_argument(
+ '--compile_dit', action='store_true', help='''
+Enable compile dit.
+Compile will give a speedup in fixed resolution and need a little GPU memory.
+The compile_dit is not compatible with the fsdp_dit and sequential_cpu_offload.
+ '''
+ )
+ parser.add_argument('--fsdp_dit', action='store_true', help="Use DIT FSDP to save more GPU memory in multi gpus.")
+ parser.add_argument('--fsdp_text_encoder', action='store_true', help="Use Text Encoder FSDP to save more GPU memory in multi gpus.")
parser.add_argument('--weight_dtype', type=str, default='bf16', help='Weight data type')
parser.add_argument('--server_name', type=str, default="0.0.0.0", help='Server IP address')
parser.add_argument('--server_port', type=int, default=7860, help='Server Port')
@@ -41,7 +63,9 @@ def main():
engine = MultiNodesEngine(
world_size=args.world_size, Controller=Wan_Controller,
GPU_memory_mode=args.gpu_memory_mode, scheduler_dict=flow_scheduler_dict, model_name=args.model_name, model_type=args.model_type, config_path=args.config_path,
- ulysses_degree=args.ulysses_degree, ring_degree=args.ring_degree, weight_dtype=weight_dtype, savedir_sample=args.savedir_sample,
+ ulysses_degree=args.ulysses_degree, ring_degree=args.ring_degree,
+ fsdp_dit=args.fsdp_dit, fsdp_text_encoder=args.fsdp_text_encoder, compile_dit=args.compile_dit,
+ weight_dtype=weight_dtype, savedir_sample=args.savedir_sample,
)
def gr_launch():
diff --git a/examples/wan2.1/predict_i2v.py b/examples/wan2.1/predict_i2v.py
index 7b7abde..c3ab137 100755
--- a/examples/wan2.1/predict_i2v.py
+++ b/examples/wan2.1/predict_i2v.py
@@ -48,8 +48,9 @@ ulysses_degree = 1
ring_degree = 1
# Use FSDP to save more GPU memory in multi gpus.
fsdp_dit = False
+fsdp_text_encoder = True
# Compile will give a speedup in fixed resolution and need a little GPU memory.
-# The compile_dit is not compatible with the fsdp_dit.
+# The compile_dit is not compatible with the fsdp_dit and sequential_cpu_offload.
compile_dit = False
# Support TeaCache.
@@ -197,7 +198,11 @@ if ulysses_degree > 1 or ring_degree > 1:
if fsdp_dit:
shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype)
pipeline.transformer = shard_fn(pipeline.transformer)
- print("Add FSDP")
+ print("Add FSDP DIT")
+ if fsdp_text_encoder:
+ shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype)
+ pipeline.text_encoder = shard_fn(pipeline.text_encoder)
+ print("Add FSDP TEXT ENCODER")
if compile_dit:
for i in range(len(pipeline.transformer.blocks)):
@@ -209,13 +214,13 @@ if GPU_memory_mode == "sequential_cpu_offload":
transformer.freqs = transformer.freqs.to(device=device)
pipeline.enable_sequential_cpu_offload(device=device)
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
- convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",])
+ convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_cpu_offload":
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_full_load_and_qfloat8":
- convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",])
+ convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
pipeline.to(device=device)
else:
@@ -228,6 +233,10 @@ if coefficients is not None:
coefficients, num_inference_steps, teacache_threshold, num_skip_start_steps=num_skip_start_steps, offload=teacache_offload
)
+if cfg_skip_ratio is not None:
+ print(f"Enable cfg_skip_ratio {cfg_skip_ratio}.")
+ pipeline.transformer.enable_cfg_skip(cfg_skip_ratio, num_inference_steps)
+
generator = torch.Generator(device=device).manual_seed(seed)
if lora_path is not None:
@@ -255,7 +264,6 @@ with torch.no_grad():
video = input_video,
mask_video = input_video_mask,
clip_image = clip_image,
- cfg_skip_ratio = cfg_skip_ratio,
shift = shift,
).videos
diff --git a/examples/wan2.1/predict_t2v.py b/examples/wan2.1/predict_t2v.py
index 105b128..01a6dc4 100755
--- a/examples/wan2.1/predict_t2v.py
+++ b/examples/wan2.1/predict_t2v.py
@@ -47,8 +47,9 @@ ulysses_degree = 1
ring_degree = 1
# Use FSDP to save more GPU memory in multi gpus.
fsdp_dit = False
+fsdp_text_encoder = True
# Compile will give a speedup in fixed resolution and need a little GPU memory.
-# The compile_dit is not compatible with the fsdp_dit.
+# The compile_dit is not compatible with the fsdp_dit and sequential_cpu_offload.
compile_dit = False
# TeaCache config
@@ -184,7 +185,11 @@ if ulysses_degree > 1 or ring_degree > 1:
if fsdp_dit:
shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype)
pipeline.transformer = shard_fn(pipeline.transformer)
- print("Add FSDP")
+ print("Add FSDP DIT")
+ if fsdp_text_encoder:
+ shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype)
+ pipeline.text_encoder = shard_fn(pipeline.text_encoder)
+ print("Add FSDP TEXT ENCODER")
if compile_dit:
for i in range(len(pipeline.transformer.blocks)):
@@ -196,13 +201,13 @@ if GPU_memory_mode == "sequential_cpu_offload":
transformer.freqs = transformer.freqs.to(device=device)
pipeline.enable_sequential_cpu_offload(device=device)
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
- convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",])
+ convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_cpu_offload":
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_full_load_and_qfloat8":
- convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",])
+ convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
pipeline.to(device=device)
else:
@@ -215,6 +220,10 @@ if coefficients is not None:
coefficients, num_inference_steps, teacache_threshold, num_skip_start_steps=num_skip_start_steps, offload=teacache_offload
)
+if cfg_skip_ratio is not None:
+ print(f"Enable cfg_skip_ratio {cfg_skip_ratio}.")
+ pipeline.transformer.enable_cfg_skip(cfg_skip_ratio, num_inference_steps)
+
generator = torch.Generator(device=device).manual_seed(seed)
if lora_path is not None:
@@ -236,7 +245,6 @@ with torch.no_grad():
generator = generator,
guidance_scale = guidance_scale,
num_inference_steps = num_inference_steps,
- cfg_skip_ratio = cfg_skip_ratio,
shift = shift,
).videos
diff --git a/examples/wan2.1_fun/app.py b/examples/wan2.1_fun/app.py
index 5db251d..60de1a5 100755
--- a/examples/wan2.1_fun/app.py
+++ b/examples/wan2.1_fun/app.py
@@ -33,6 +33,9 @@ if __name__ == "__main__":
# sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
# resulting in slower speeds but saving a large amount of GPU memory.
GPU_memory_mode = "sequential_cpu_offload"
+ # Compile will give a speedup in fixed resolution and need a little GPU memory.
+ # The compile_dit is not compatible with the fsdp_dit and sequential_cpu_offload.
+ compile_dit = False
# Use torch.float16 if GPU does not support torch.bfloat16
# ome graphics cards, such as v100, 2080ti, do not support torch.bfloat16
@@ -51,11 +54,11 @@ if __name__ == "__main__":
model_type = "Inpaint"
if ui_mode == "host":
- demo, controller = ui_host(GPU_memory_mode, flow_scheduler_dict, model_name, model_type, config_path, 1, 1, weight_dtype)
+ demo, controller = ui_host(GPU_memory_mode, flow_scheduler_dict, model_name, model_type, config_path, compile_dit, weight_dtype)
elif ui_mode == "client":
demo, controller = ui_client(flow_scheduler_dict, model_name)
else:
- demo, controller = ui(GPU_memory_mode, flow_scheduler_dict, config_path, 1, 1, weight_dtype)
+ demo, controller = ui(GPU_memory_mode, flow_scheduler_dict, config_path, compile_dit, weight_dtype)
def gr_launch():
# launch gradio
diff --git a/examples/wan2.1_fun/launch_api.py b/examples/wan2.1_fun/launch_api.py
index e4b12f0..1ee8847 100755
--- a/examples/wan2.1_fun/launch_api.py
+++ b/examples/wan2.1_fun/launch_api.py
@@ -20,9 +20,31 @@ from videox_fun.ui.wan_fun_ui import Wan_Fun_Controller
def main():
parser = argparse.ArgumentParser(description='xDiT HTTP Service')
parser.add_argument('--world_size', type=int, default=8, help='Number of parallel workers')
- parser.add_argument('--gpu_memory_mode', type=str, default="model_full_load", help='GPU memory mode')
+ parser.add_argument(
+ '--gpu_memory_mode', type=str, default="model_full_load", help='''
+GPU memory mode, which can be choosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8].
+model_full_load means that the entire model will be moved to the GPU.
+
+model_full_load_and_qfloat8 means that the entire model will be moved to the GPU,
+and the transformer model has been quantized to float8, which can save more GPU memory.
+
+model_cpu_offload means that the entire model will be moved to the CPU after use, which can save some GPU memory.
+
+model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
+and the transformer model has been quantized to float8, which can save more GPU memory.
+ '''
+ )
parser.add_argument('--ulysses_degree', type=int, default=4, help='Degree of Ulysses configuration')
parser.add_argument('--ring_degree', type=int, default=2, help='Degree of Ring configuration')
+ parser.add_argument(
+ '--compile_dit', action='store_true', help='''
+Enable compile dit.
+Compile will give a speedup in fixed resolution and need a little GPU memory.
+The compile_dit is not compatible with the fsdp_dit and sequential_cpu_offload.
+ '''
+ )
+ parser.add_argument('--fsdp_dit', action='store_true', help="Use DIT FSDP to save more GPU memory in multi gpus.")
+ parser.add_argument('--fsdp_text_encoder', action='store_true', help="Use Text Encoder FSDP to save more GPU memory in multi gpus.")
parser.add_argument('--weight_dtype', type=str, default='bf16', help='Weight data type')
parser.add_argument('--server_name', type=str, default="0.0.0.0", help='Server IP address')
parser.add_argument('--server_port', type=int, default=7860, help='Server Port')
@@ -41,7 +63,9 @@ def main():
engine = MultiNodesEngine(
world_size=args.world_size, Controller=Wan_Fun_Controller,
GPU_memory_mode=args.gpu_memory_mode, scheduler_dict=flow_scheduler_dict, model_name=args.model_name, model_type=args.model_type, config_path=args.config_path,
- ulysses_degree=args.ulysses_degree, ring_degree=args.ring_degree, weight_dtype=weight_dtype, savedir_sample=args.savedir_sample,
+ ulysses_degree=args.ulysses_degree, ring_degree=args.ring_degree,
+ fsdp_dit=args.fsdp_dit, fsdp_text_encoder=args.fsdp_text_encoder, compile_dit=args.compile_dit,
+ weight_dtype=weight_dtype, savedir_sample=args.savedir_sample,
)
def gr_launch():
diff --git a/examples/wan2.1_fun/predict_i2v.py b/examples/wan2.1_fun/predict_i2v.py
index 70837cd..f2b0b71 100755
--- a/examples/wan2.1_fun/predict_i2v.py
+++ b/examples/wan2.1_fun/predict_i2v.py
@@ -48,8 +48,9 @@ ulysses_degree = 1
ring_degree = 1
# Use FSDP to save more GPU memory in multi gpus.
fsdp_dit = False
+fsdp_text_encoder = True
# Compile will give a speedup in fixed resolution and need a little GPU memory.
-# The compile_dit is not compatible with the fsdp_dit.
+# The compile_dit is not compatible with the fsdp_dit and sequential_cpu_offload.
compile_dit = False
# Support TeaCache.
@@ -198,7 +199,11 @@ if ulysses_degree > 1 or ring_degree > 1:
if fsdp_dit:
shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype)
pipeline.transformer = shard_fn(pipeline.transformer)
- print("Add FSDP")
+ print("Add FSDP DIT")
+ if fsdp_text_encoder:
+ shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype)
+ pipeline.text_encoder = shard_fn(pipeline.text_encoder)
+ print("Add FSDP TEXT ENCODER")
if compile_dit:
for i in range(len(pipeline.transformer.blocks)):
@@ -210,13 +215,13 @@ if GPU_memory_mode == "sequential_cpu_offload":
transformer.freqs = transformer.freqs.to(device=device)
pipeline.enable_sequential_cpu_offload(device=device)
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
- convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",])
+ convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_cpu_offload":
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_full_load_and_qfloat8":
- convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",])
+ convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
pipeline.to(device=device)
else:
@@ -229,6 +234,10 @@ if coefficients is not None:
coefficients, num_inference_steps, teacache_threshold, num_skip_start_steps=num_skip_start_steps, offload=teacache_offload
)
+if cfg_skip_ratio is not None:
+ print(f"Enable cfg_skip_ratio {cfg_skip_ratio}.")
+ pipeline.transformer.enable_cfg_skip(cfg_skip_ratio, num_inference_steps)
+
generator = torch.Generator(device=device).manual_seed(seed)
if lora_path is not None:
@@ -256,7 +265,6 @@ with torch.no_grad():
video = input_video,
mask_video = input_video_mask,
clip_image = clip_image,
- cfg_skip_ratio = cfg_skip_ratio,
shift = shift,
).videos
diff --git a/examples/wan2.1_fun/predict_t2v.py b/examples/wan2.1_fun/predict_t2v.py
index 9e1fd06..7a76ef4 100755
--- a/examples/wan2.1_fun/predict_t2v.py
+++ b/examples/wan2.1_fun/predict_t2v.py
@@ -48,8 +48,9 @@ ulysses_degree = 1
ring_degree = 1
# Use FSDP to save more GPU memory in multi gpus.
fsdp_dit = False
+fsdp_text_encoder = True
# Compile will give a speedup in fixed resolution and need a little GPU memory.
-# The compile_dit is not compatible with the fsdp_dit.
+# The compile_dit is not compatible with the fsdp_dit and sequential_cpu_offload.
compile_dit = False
# Support TeaCache.
@@ -205,19 +206,29 @@ if ulysses_degree > 1 or ring_degree > 1:
if fsdp_dit:
shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype)
pipeline.transformer = shard_fn(pipeline.transformer)
+ print("Add FSDP DIT")
+ if fsdp_text_encoder:
+ shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype)
+ pipeline.text_encoder = shard_fn(pipeline.text_encoder)
+ print("Add FSDP TEXT ENCODER")
+
+if compile_dit:
+ for i in range(len(pipeline.transformer.blocks)):
+ pipeline.transformer.blocks[i] = torch.compile(pipeline.transformer.blocks[i])
+ print("Add Compile")
if GPU_memory_mode == "sequential_cpu_offload":
replace_parameters_by_name(transformer, ["modulation",], device=device)
transformer.freqs = transformer.freqs.to(device=device)
pipeline.enable_sequential_cpu_offload(device=device)
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
- convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",])
+ convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_cpu_offload":
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_full_load_and_qfloat8":
- convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",])
+ convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
pipeline.to(device=device)
else:
@@ -230,6 +241,10 @@ if coefficients is not None:
coefficients, num_inference_steps, teacache_threshold, num_skip_start_steps=num_skip_start_steps, offload=teacache_offload
)
+if cfg_skip_ratio is not None:
+ print(f"Enable cfg_skip_ratio {cfg_skip_ratio}.")
+ pipeline.transformer.enable_cfg_skip(cfg_skip_ratio, num_inference_steps)
+
generator = torch.Generator(device=device).manual_seed(seed)
if lora_path is not None:
@@ -257,7 +272,6 @@ with torch.no_grad():
video = input_video,
mask_video = input_video_mask,
- cfg_skip_ratio = cfg_skip_ratio,
shift = shift,
).videos
else:
@@ -270,7 +284,6 @@ with torch.no_grad():
generator = generator,
guidance_scale = guidance_scale,
num_inference_steps = num_inference_steps,
- cfg_skip_ratio = cfg_skip_ratio,
shift = shift,
).videos
diff --git a/examples/wan2.1_fun/predict_v2v_control.py b/examples/wan2.1_fun/predict_v2v_control.py
index 5f6260d..c1a58e6 100755
--- a/examples/wan2.1_fun/predict_v2v_control.py
+++ b/examples/wan2.1_fun/predict_v2v_control.py
@@ -51,8 +51,9 @@ ulysses_degree = 1
ring_degree = 1
# Use FSDP to save more GPU memory in multi gpus.
fsdp_dit = False
+fsdp_text_encoder = True
# Compile will give a speedup in fixed resolution and need a little GPU memory.
-# The compile_dit is not compatible with the fsdp_dit.
+# The compile_dit is not compatible with the fsdp_dit and sequential_cpu_offload.
compile_dit = False
# Support TeaCache.
@@ -208,7 +209,11 @@ if ulysses_degree > 1 or ring_degree > 1:
if fsdp_dit:
shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype)
pipeline.transformer = shard_fn(pipeline.transformer)
- print("Add FSDP")
+ print("Add FSDP DIT")
+ if fsdp_text_encoder:
+ shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype)
+ pipeline.text_encoder = shard_fn(pipeline.text_encoder)
+ print("Add FSDP TEXT ENCODER")
if compile_dit:
for i in range(len(pipeline.transformer.blocks)):
@@ -220,13 +225,13 @@ if GPU_memory_mode == "sequential_cpu_offload":
transformer.freqs = transformer.freqs.to(device=device)
pipeline.enable_sequential_cpu_offload(device=device)
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
- convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",])
+ convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_cpu_offload":
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_full_load_and_qfloat8":
- convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",])
+ convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
pipeline.to(device=device)
else:
@@ -239,6 +244,10 @@ if coefficients is not None:
coefficients, num_inference_steps, teacache_threshold, num_skip_start_steps=num_skip_start_steps, offload=teacache_offload
)
+if cfg_skip_ratio is not None:
+ print(f"Enable cfg_skip_ratio {cfg_skip_ratio}.")
+ pipeline.transformer.enable_cfg_skip(cfg_skip_ratio, num_inference_steps)
+
generator = torch.Generator(device=device).manual_seed(seed)
if lora_path is not None:
@@ -287,7 +296,6 @@ with torch.no_grad():
ref_image = ref_image,
start_image = start_image,
clip_image = clip_image,
- cfg_skip_ratio = cfg_skip_ratio,
shift = shift,
).videos
diff --git a/examples/wan2.1_fun/predict_v2v_control_camera.py b/examples/wan2.1_fun/predict_v2v_control_camera.py
index 317259f..24f79f4 100755
--- a/examples/wan2.1_fun/predict_v2v_control_camera.py
+++ b/examples/wan2.1_fun/predict_v2v_control_camera.py
@@ -51,8 +51,9 @@ ulysses_degree = 1
ring_degree = 1
# Use FSDP to save more GPU memory in multi gpus.
fsdp_dit = False
+fsdp_text_encoder = True
# Compile will give a speedup in fixed resolution and need a little GPU memory.
-# The compile_dit is not compatible with the fsdp_dit.
+# The compile_dit is not compatible with the fsdp_dit and sequential_cpu_offload.
compile_dit = False
# Support TeaCache.
@@ -208,7 +209,11 @@ if ulysses_degree > 1 or ring_degree > 1:
if fsdp_dit:
shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype)
pipeline.transformer = shard_fn(pipeline.transformer)
- print("Add FSDP")
+ print("Add FSDP DIT")
+ if fsdp_text_encoder:
+ shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype)
+ pipeline.text_encoder = shard_fn(pipeline.text_encoder)
+ print("Add FSDP TEXT ENCODER")
if compile_dit:
for i in range(len(pipeline.transformer.blocks)):
@@ -220,13 +225,13 @@ if GPU_memory_mode == "sequential_cpu_offload":
transformer.freqs = transformer.freqs.to(device=device)
pipeline.enable_sequential_cpu_offload(device=device)
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
- convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",])
+ convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_cpu_offload":
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_full_load_and_qfloat8":
- convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",])
+ convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
pipeline.to(device=device)
else:
@@ -239,6 +244,10 @@ if coefficients is not None:
coefficients, num_inference_steps, teacache_threshold, num_skip_start_steps=num_skip_start_steps, offload=teacache_offload
)
+if cfg_skip_ratio is not None:
+ print(f"Enable cfg_skip_ratio {cfg_skip_ratio}.")
+ pipeline.transformer.enable_cfg_skip(cfg_skip_ratio, num_inference_steps)
+
generator = torch.Generator(device=device).manual_seed(seed)
if lora_path is not None:
@@ -287,7 +296,6 @@ with torch.no_grad():
ref_image = ref_image,
start_image = start_image,
clip_image = clip_image,
- cfg_skip_ratio = cfg_skip_ratio,
shift = shift,
).videos
diff --git a/examples/wan2.1_fun/predict_v2v_control_ref.py b/examples/wan2.1_fun/predict_v2v_control_ref.py
index f9cdda0..6442105 100755
--- a/examples/wan2.1_fun/predict_v2v_control_ref.py
+++ b/examples/wan2.1_fun/predict_v2v_control_ref.py
@@ -51,8 +51,9 @@ ulysses_degree = 1
ring_degree = 1
# Use FSDP to save more GPU memory in multi gpus.
fsdp_dit = False
+fsdp_text_encoder = True
# Compile will give a speedup in fixed resolution and need a little GPU memory.
-# The compile_dit is not compatible with the fsdp_dit.
+# The compile_dit is not compatible with the fsdp_dit and sequential_cpu_offload.
compile_dit = False
# Support TeaCache.
@@ -208,7 +209,11 @@ if ulysses_degree > 1 or ring_degree > 1:
if fsdp_dit:
shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype)
pipeline.transformer = shard_fn(pipeline.transformer)
- print("Add FSDP")
+ print("Add FSDP DIT")
+ if fsdp_text_encoder:
+ shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype)
+ pipeline.text_encoder = shard_fn(pipeline.text_encoder)
+ print("Add FSDP TEXT ENCODER")
if compile_dit:
for i in range(len(pipeline.transformer.blocks)):
@@ -220,13 +225,13 @@ if GPU_memory_mode == "sequential_cpu_offload":
transformer.freqs = transformer.freqs.to(device=device)
pipeline.enable_sequential_cpu_offload(device=device)
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
- convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",])
+ convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_cpu_offload":
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_full_load_and_qfloat8":
- convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",])
+ convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
pipeline.to(device=device)
else:
@@ -239,6 +244,10 @@ if coefficients is not None:
coefficients, num_inference_steps, teacache_threshold, num_skip_start_steps=num_skip_start_steps, offload=teacache_offload
)
+if cfg_skip_ratio is not None:
+ print(f"Enable cfg_skip_ratio {cfg_skip_ratio}.")
+ pipeline.transformer.enable_cfg_skip(cfg_skip_ratio, num_inference_steps)
+
generator = torch.Generator(device=device).manual_seed(seed)
if lora_path is not None:
@@ -287,7 +296,6 @@ with torch.no_grad():
ref_image = ref_image,
start_image = start_image,
clip_image = clip_image,
- cfg_skip_ratio = cfg_skip_ratio,
shift = shift,
).videos
diff --git a/scripts/wan2.1_fun/README_TRAIN.md b/scripts/wan2.1_fun/README_TRAIN.md
index 37108e6..fd036f7 100755
--- a/scripts/wan2.1_fun/README_TRAIN.md
+++ b/scripts/wan2.1_fun/README_TRAIN.md
@@ -24,7 +24,7 @@ Wan T2V without deepspeed:
Wan without DeepSpeed is more suitable for 1.3B Wan, as using it with 14B Wan may result in insufficient GPU memory.
```sh
-export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-Fun-14B-InP"
+export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-Fun-V1.1-14B-InP"
export DATASET_NAME="datasets/internal_datasets/"
export DATASET_META_NAME="datasets/internal_datasets/metadata.json"
export NCCL_IB_DISABLE=1
@@ -72,7 +72,7 @@ Wan T2V with deepspeed zero-2:
Wan with DeepSpeed Zero-2 is suitable for training 1.3B Wan and 14B Wan at low resolutions, but training 14B Wan at high resolutions may still result in insufficient GPU memory.
```sh
-export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-Fun-14B-InP"
+export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-Fun-V1.1-14B-InP"
export DATASET_NAME="datasets/internal_datasets/"
export DATASET_META_NAME="datasets/internal_datasets/metadata.json"
export NCCL_IB_DISABLE=1
@@ -125,7 +125,7 @@ python scripts/zero_to_bf16.py output_dir/checkpoint-{our-num-steps} output_dir/
Training shell command is as follows:
```sh
-export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-Fun-14B-InP"
+export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-Fun-V1.1-14B-InP"
export DATASET_NAME="datasets/internal_datasets/"
export DATASET_META_NAME="datasets/internal_datasets/metadata.json"
export NCCL_IB_DISABLE=1
diff --git a/scripts/wan2.1_fun/README_TRAIN_CONTROL.md b/scripts/wan2.1_fun/README_TRAIN_CONTROL.md
old mode 100644
new mode 100755
index 36feed9..e417b3b
--- a/scripts/wan2.1_fun/README_TRAIN_CONTROL.md
+++ b/scripts/wan2.1_fun/README_TRAIN_CONTROL.md
@@ -38,8 +38,179 @@ Some parameters in the sh file can be confusing, and they are explained in this
- At 1024x1024 resolution, the number of video frames is 9 (~= 512 * 512 * 49 / 1024 / 1024).
- These resolutions combined with their corresponding lengths allow the model to generate videos of different sizes.
- `resume_from_checkpoint` is used to set the training should be resumed from a previous checkpoint. Use a path or `"latest"` to automatically select the last available checkpoint.
+- `train_mode` is used to set the training mode.
+ - The models named `Wan2.1-Fun-*-Control` are trained in the `control_ref` mode.
+ - The models named `Wan2.1-Fun-*-Control-Camera` are trained in the `control_ref_camera` mode.
+- `control_ref_image` is used to specify the type of control image. The available options are `first_frame` and `random`.
+ - `first_frame` is used in V1.0 because V1.0 supports using a specified start frame as the control image. The Control-Camera models use the first frame as the control image.
+ - `random` is used in V1.1 because V1.1 supports both using a specified start frame and a reference image as the control image.
+- `add_full_ref_image_in_self_attention` determines whether to include the reference image in self-attention. This option is used in V1.1, as it supports using a reference image as the control image. It should not be used in V1.0 and Control-Camera models.
-Wan-Fun-Control without deepspeed:
+When train model with multi machines, please set the params as follows:
+```sh
+export MASTER_ADDR="your master address"
+export MASTER_PORT=10086
+export WORLD_SIZE=1 # The number of machines
+export NUM_PROCESS=8 # The number of processes, such as WORLD_SIZE * 8
+export RANK=0 # The rank of this machine
+
+accelerate launch --mixed_precision="bf16" --main_process_ip=$MASTER_ADDR --main_process_port=$MASTER_PORT --num_machines=$WORLD_SIZE --num_processes=$NUM_PROCESS --machine_rank=$RANK scripts/wan2.1_fun/xxx.py
+```
+
+Wan-Fun-Control-V1.1 without deepspeed:
+```sh
+export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-Fun-V1.1-14B-Control"
+export DATASET_NAME="datasets/internal_datasets/"
+export DATASET_META_NAME="datasets/internal_datasets/metadata.json"
+export NCCL_IB_DISABLE=1
+export NCCL_P2P_DISABLE=1
+NCCL_DEBUG=INFO
+
+accelerate launch --mixed_precision="bf16" scripts/wan2.1_fun/train_control.py \
+ --config_path="config/wan2.1/wan_civitai.yaml" \
+ --pretrained_model_name_or_path=$MODEL_NAME \
+ --train_data_dir=$DATASET_NAME \
+ --train_data_meta=$DATASET_META_NAME \
+ --image_sample_size=1024 \
+ --video_sample_size=256 \
+ --token_sample_size=512 \
+ --video_sample_stride=2 \
+ --video_sample_n_frames=81 \
+ --train_batch_size=1 \
+ --video_repeat=1 \
+ --gradient_accumulation_steps=1 \
+ --dataloader_num_workers=8 \
+ --num_train_epochs=100 \
+ --checkpointing_steps=50 \
+ --learning_rate=2e-05 \
+ --lr_scheduler="constant_with_warmup" \
+ --lr_warmup_steps=100 \
+ --seed=42 \
+ --output_dir="output_dir" \
+ --gradient_checkpointing \
+ --mixed_precision="bf16" \
+ --adam_weight_decay=3e-2 \
+ --adam_epsilon=1e-10 \
+ --vae_mini_batch=1 \
+ --max_grad_norm=0.05 \
+ --random_hw_adapt \
+ --training_with_video_token_length \
+ --enable_bucket \
+ --uniform_sampling \
+ --low_vram \
+ --train_mode="control_ref" \
+ --control_ref_image="random" \
+ --add_full_ref_image_in_self_attention \
+ --trainable_modules "."
+```
+
+Wan-Fun-Control-V1.1 with deepspeed:
+```sh
+export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-Fun-V1.1-14B-Control"
+export DATASET_NAME="datasets/internal_datasets/"
+export DATASET_META_NAME="datasets/internal_datasets/metadata.json"
+export NCCL_IB_DISABLE=1
+export NCCL_P2P_DISABLE=1
+NCCL_DEBUG=INFO
+
+accelerate launch --use_deepspeed --deepspeed_config_file config/zero_stage2_config.json --deepspeed_multinode_launcher standard scripts/wan2.1_fun/train_control.py \
+ --config_path="config/wan2.1/wan_civitai.yaml" \
+ --pretrained_model_name_or_path=$MODEL_NAME \
+ --train_data_dir=$DATASET_NAME \
+ --train_data_meta=$DATASET_META_NAME \
+ --image_sample_size=1024 \
+ --video_sample_size=256 \
+ --token_sample_size=512 \
+ --video_sample_stride=2 \
+ --video_sample_n_frames=81 \
+ --train_batch_size=1 \
+ --video_repeat=1 \
+ --gradient_accumulation_steps=1 \
+ --dataloader_num_workers=8 \
+ --num_train_epochs=100 \
+ --checkpointing_steps=50 \
+ --learning_rate=2e-05 \
+ --lr_scheduler="constant_with_warmup" \
+ --lr_warmup_steps=100 \
+ --seed=42 \
+ --output_dir="output_dir" \
+ --gradient_checkpointing \
+ --mixed_precision="bf16" \
+ --adam_weight_decay=3e-2 \
+ --adam_epsilon=1e-10 \
+ --vae_mini_batch=1 \
+ --max_grad_norm=0.05 \
+ --random_hw_adapt \
+ --training_with_video_token_length \
+ --enable_bucket \
+ --uniform_sampling \
+ --low_vram \
+ --use_deepspeed \
+ --train_mode="control_ref" \
+ --control_ref_image="random" \
+ --add_full_ref_image_in_self_attention \
+ --trainable_modules "."
+```
+
+Wan-Fun-Control-V1.1 with deepspeed zero-3:
+
+Wan with DeepSpeed Zero-3 is suitable for 14B Wan at high resolutions. After training, you can use the following command to get the final model:
+```sh
+python scripts/zero_to_bf16.py output_dir/checkpoint-{our-num-steps} output_dir/checkpoint-{your-num-steps}-outputs --max_shard_size 80GB --safe_serialization
+```
+
+Training shell command is as follows:
+```sh
+export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-Fun-V1.1-14B-Control"
+export DATASET_NAME="datasets/internal_datasets/"
+export DATASET_META_NAME="datasets/internal_datasets/metadata.json"
+export NCCL_IB_DISABLE=1
+export NCCL_P2P_DISABLE=1
+NCCL_DEBUG=INFO
+
+accelerate launch --zero_stage 3 --zero3_save_16bit_model true --zero3_init_flag true --use_deepspeed --deepspeed_config_file config/zero_stage2.1_config.json --deepspeed_multinode_launcherr standard scripts/wan2.1_fun/train_control.py \
+ --config_path="config/wan2.1/wan_civitai.yaml" \
+ --pretrained_model_name_or_path=$MODEL_NAME \
+ --train_data_dir=$DATASET_NAME \
+ --train_data_meta=$DATASET_META_NAME \
+ --image_sample_size=1024 \
+ --video_sample_size=256 \
+ --token_sample_size=512 \
+ --video_sample_stride=2 \
+ --video_sample_n_frames=81 \
+ --train_batch_size=1 \
+ --video_repeat=1 \
+ --gradient_accumulation_steps=1 \
+ --dataloader_num_workers=8 \
+ --num_train_epochs=100 \
+ --checkpointing_steps=50 \
+ --learning_rate=2e-05 \
+ --lr_scheduler="constant_with_warmup" \
+ --lr_warmup_steps=100 \
+ --seed=42 \
+ --output_dir="output_dir" \
+ --gradient_checkpointing \
+ --mixed_precision="bf16" \
+ --adam_weight_decay=3e-2 \
+ --adam_epsilon=1e-10 \
+ --vae_mini_batch=1 \
+ --max_grad_norm=0.05 \
+ --random_hw_adapt \
+ --training_with_video_token_length \
+ --enable_bucket \
+ --uniform_sampling \
+ --low_vram \
+ --use_deepspeed \
+ --train_mode="control_ref" \
+ --control_ref_image="random" \
+ --add_full_ref_image_in_self_attention \
+ --trainable_modules "."
+```
+
+
+ (Obsolete) V1.0:
+
+Wan-Fun-Control-V1.0 without deepspeed:
```sh
export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-Fun-14B-Control"
export DATASET_NAME="datasets/internal_datasets/"
@@ -48,7 +219,6 @@ export NCCL_IB_DISABLE=1
export NCCL_P2P_DISABLE=1
NCCL_DEBUG=INFO
-# When train model with multi machines, use "--config_file accelerate.yaml" instead of "--mixed_precision='bf16'".
accelerate launch --mixed_precision="bf16" scripts/wan2.1_fun/train_control.py \
--config_path="config/wan2.1/wan_civitai.yaml" \
--pretrained_model_name_or_path=$MODEL_NAME \
@@ -86,7 +256,7 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.1_fun/train_control.py \
--trainable_modules "."
```
-Wan-Fun with deepspeed:
+Wan-Fun-Control-V1.0 with deepspeed:
```sh
export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-Fun-14B-Control"
export DATASET_NAME="datasets/internal_datasets/"
@@ -95,7 +265,6 @@ export NCCL_IB_DISABLE=1
export NCCL_P2P_DISABLE=1
NCCL_DEBUG=INFO
-# When train model with multi machines, use "--config_file accelerate.yaml" instead of "--mixed_precision='bf16'".
accelerate launch --use_deepspeed --deepspeed_config_file config/zero_stage2_config.json --deepspeed_multinode_launcher standard scripts/wan2.1_fun/train_control.py \
--config_path="config/wan2.1/wan_civitai.yaml" \
--pretrained_model_name_or_path=$MODEL_NAME \
@@ -134,7 +303,7 @@ accelerate launch --use_deepspeed --deepspeed_config_file config/zero_stage2_con
--trainable_modules "."
```
-Wan T2V with deepspeed zero-3:
+Wan-Fun-Control-V1.0 with deepspeed zero-3:
Wan with DeepSpeed Zero-3 is suitable for 14B Wan at high resolutions. After training, you can use the following command to get the final model:
```sh
@@ -186,4 +355,5 @@ accelerate launch --zero_stage 3 --zero3_save_16bit_model true --zero3_init_flag
--train_mode="control_ref" \
--control_ref_image="first_frame" \
--trainable_modules "."
-```
\ No newline at end of file
+```
+
diff --git a/scripts/wan2.1_fun/README_TRAIN_CONTROL_LORA.md b/scripts/wan2.1_fun/README_TRAIN_CONTROL_LORA.md
old mode 100644
new mode 100755
index 494035c..53d798d
--- a/scripts/wan2.1_fun/README_TRAIN_CONTROL_LORA.md
+++ b/scripts/wan2.1_fun/README_TRAIN_CONTROL_LORA.md
@@ -37,9 +37,172 @@ Some parameters in the sh file can be confusing, and they are explained in this
- At 768x768 resolution, the number of video frames is 21 (~= 512 * 512 * 49 / 768 / 768).
- At 1024x1024 resolution, the number of video frames is 9 (~= 512 * 512 * 49 / 1024 / 1024).
- These resolutions combined with their corresponding lengths allow the model to generate videos of different sizes.
-- `resume_from_checkpoint` is used to set the training should be resumed from a previous checkpoint. Use a path or `"latest"` to automatically select the last available checkpoint.
+- `resume_from_checkpoint` is used to set the training should be resumed from a previous checkpoint. Use a path or `"latest"` to automatically select the last available checkpoint and set the `save_state` to `True`.
+- `train_mode` is used to set the training mode.
+ - The models named `Wan2.1-Fun-*-Control` are trained in the `control_ref` mode.
+ - The models named `Wan2.1-Fun-*-Control-Camera` are trained in the `control_ref_camera` mode.
+- `control_ref_image` is used to specify the type of control image. The available options are `first_frame` and `random`.
+ - `first_frame` is used in V1.0 because V1.0 supports using a specified start frame as the control image. The Control-Camera models use the first frame as the control image.
+ - `random` is used in V1.1 because V1.1 supports both using a specified start frame and a reference image as the control image.
+- `add_full_ref_image_in_self_attention` determines whether to include the reference image in self-attention. This option is used in V1.1, as it supports using a reference image as the control image. It should not be used in V1.0 and Control-Camera models.
-Wan-Fun-Control without deepspeed:
+When train model with multi machines, please set the params as follows:
+```sh
+export MASTER_ADDR="your master address"
+export MASTER_PORT=10086
+export WORLD_SIZE=1 # The number of machines
+export NUM_PROCESS=8 # The number of processes, such as WORLD_SIZE * 8
+export RANK=0 # The rank of this machine
+
+accelerate launch --mixed_precision="bf16" --main_process_ip=$MASTER_ADDR --main_process_port=$MASTER_PORT --num_machines=$WORLD_SIZE --num_processes=$NUM_PROCESS --machine_rank=$RANK scripts/wan2.1_fun/xxx.py
+```
+
+Wan-Fun-Control-V1.1 without deepspeed:
+```sh
+export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-Fun-V1.1-14B-Control"
+export DATASET_NAME="datasets/internal_datasets/"
+export DATASET_META_NAME="datasets/internal_datasets/metadata.json"
+export NCCL_IB_DISABLE=1
+export NCCL_P2P_DISABLE=1
+NCCL_DEBUG=INFO
+
+accelerate launch --mixed_precision="bf16" scripts/wan2.1_fun/train_control_lora.py \
+ --config_path="config/wan2.1/wan_civitai.yaml" \
+ --pretrained_model_name_or_path=$MODEL_NAME \
+ --train_data_dir=$DATASET_NAME \
+ --train_data_meta=$DATASET_META_NAME \
+ --image_sample_size=1024 \
+ --video_sample_size=256 \
+ --token_sample_size=512 \
+ --video_sample_stride=2 \
+ --video_sample_n_frames=81 \
+ --train_batch_size=1 \
+ --video_repeat=1 \
+ --gradient_accumulation_steps=1 \
+ --dataloader_num_workers=8 \
+ --num_train_epochs=100 \
+ --checkpointing_steps=50 \
+ --learning_rate=1e-04 \
+ --seed=42 \
+ --output_dir="output_dir" \
+ --gradient_checkpointing \
+ --mixed_precision="bf16" \
+ --adam_weight_decay=3e-2 \
+ --adam_epsilon=1e-10 \
+ --vae_mini_batch=1 \
+ --max_grad_norm=0.05 \
+ --random_hw_adapt \
+ --training_with_video_token_length \
+ --enable_bucket \
+ --uniform_sampling \
+ --train_mode="control_ref" \
+ --control_ref_image="random" \
+ --add_full_ref_image_in_self_attention \
+ --low_vram
+```
+
+Wan-Fun-Control-V1.1 with deepspeed:
+```sh
+export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-Fun-V1.1-14B-Control"
+export DATASET_NAME="datasets/internal_datasets/"
+export DATASET_META_NAME="datasets/internal_datasets/metadata.json"
+export NCCL_IB_DISABLE=1
+export NCCL_P2P_DISABLE=1
+NCCL_DEBUG=INFO
+
+accelerate launch --use_deepspeed --deepspeed_config_file config/zero_stage2_config.json --deepspeed_multinode_launcher standard scripts/wan2.1_fun/train_control_lora.py \
+ --config_path="config/wan2.1/wan_civitai.yaml" \
+ --pretrained_model_name_or_path=$MODEL_NAME \
+ --train_data_dir=$DATASET_NAME \
+ --train_data_meta=$DATASET_META_NAME \
+ --image_sample_size=1024 \
+ --video_sample_size=256 \
+ --token_sample_size=512 \
+ --video_sample_stride=2 \
+ --video_sample_n_frames=81 \
+ --train_batch_size=1 \
+ --video_repeat=1 \
+ --gradient_accumulation_steps=1 \
+ --dataloader_num_workers=8 \
+ --num_train_epochs=100 \
+ --checkpointing_steps=50 \
+ --learning_rate=1e-04 \
+ --seed=42 \
+ --output_dir="output_dir" \
+ --gradient_checkpointing \
+ --mixed_precision="bf16" \
+ --adam_weight_decay=3e-2 \
+ --adam_epsilon=1e-10 \
+ --vae_mini_batch=1 \
+ --max_grad_norm=0.05 \
+ --random_hw_adapt \
+ --training_with_video_token_length \
+ --enable_bucket \
+ --uniform_sampling \
+ --use_deepspeed \
+ --train_mode="control_ref" \
+ --control_ref_image="random" \
+ --add_full_ref_image_in_self_attention \
+ --low_vram
+```
+
+Wan-Fun-Control-V1.1 with deepspeed zero-3:
+
+Wan with DeepSpeed Zero-3 is suitable for 14B Wan at high resolutions. You must set save_state to True to save the model. After training, you can use the following command to get the final model:
+```sh
+python scripts/zero_to_bf16.py output_dir/checkpoint-{our-num-steps} output_dir/checkpoint-{your-num-steps}-outputs --max_shard_size 80GB --safe_serialization
+```
+
+Training shell command is as follows:
+```sh
+export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-Fun-V1.1-14B-Control"
+export DATASET_NAME="datasets/internal_datasets/"
+export DATASET_META_NAME="datasets/internal_datasets/metadata.json"
+export NCCL_IB_DISABLE=1
+export NCCL_P2P_DISABLE=1
+NCCL_DEBUG=INFO
+
+accelerate launch --zero_stage 3 --zero3_save_16bit_model true --zero3_init_flag true --use_deepspeed --deepspeed_config_file config/zero_stage2.1_config.json --deepspeed_multinode_launcherr standard scripts/wan2.1_fun/train_control.py \
+ --config_path="config/wan2.1/wan_civitai.yaml" \
+ --pretrained_model_name_or_path=$MODEL_NAME \
+ --train_data_dir=$DATASET_NAME \
+ --train_data_meta=$DATASET_META_NAME \
+ --image_sample_size=1024 \
+ --video_sample_size=256 \
+ --token_sample_size=512 \
+ --video_sample_stride=2 \
+ --video_sample_n_frames=81 \
+ --train_batch_size=1 \
+ --video_repeat=1 \
+ --gradient_accumulation_steps=1 \
+ --dataloader_num_workers=8 \
+ --num_train_epochs=100 \
+ --checkpointing_steps=50 \
+ --learning_rate=1e-04 \
+ --seed=42 \
+ --output_dir="output_dir" \
+ --gradient_checkpointing \
+ --mixed_precision="bf16" \
+ --adam_weight_decay=3e-2 \
+ --adam_epsilon=1e-10 \
+ --vae_mini_batch=1 \
+ --max_grad_norm=0.05 \
+ --random_hw_adapt \
+ --training_with_video_token_length \
+ --enable_bucket \
+ --uniform_sampling \
+ --save_state \
+ --use_deepspeed \
+ --train_mode="control_ref" \
+ --control_ref_image="random" \
+ --add_full_ref_image_in_self_attention \
+ --low_vram
+```
+
+
+ (Obsolete) V1.0:
+
+Wan-Fun-Control-V1.0 without deepspeed:
```sh
export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-Fun-14B-Control"
export DATASET_NAME="datasets/internal_datasets/"
@@ -82,7 +245,7 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.1_fun/train_control_lora
--low_vram
```
-Wan-Fun with deepspeed:
+Wan-Fun-Control-V1.0 with deepspeed:
```sh
export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-Fun-14B-Control"
export DATASET_NAME="datasets/internal_datasets/"
@@ -91,7 +254,6 @@ export NCCL_IB_DISABLE=1
export NCCL_P2P_DISABLE=1
NCCL_DEBUG=INFO
-# When train model with multi machines, use "--config_file accelerate.yaml" instead of "--mixed_precision='bf16'".
accelerate launch --use_deepspeed --deepspeed_config_file config/zero_stage2_config.json --deepspeed_multinode_launcher standard scripts/wan2.1_fun/train_control_lora.py \
--config_path="config/wan2.1/wan_civitai.yaml" \
--pretrained_model_name_or_path=$MODEL_NAME \
@@ -127,7 +289,7 @@ accelerate launch --use_deepspeed --deepspeed_config_file config/zero_stage2_con
--low_vram
```
-Wan T2V with deepspeed zero-3:
+Wan-Fun-Control-V1.0 with deepspeed zero-3:
Wan with DeepSpeed Zero-3 is suitable for 14B Wan at high resolutions. You must set save_state to True to save the model. After training, you can use the following command to get the final model:
```sh
@@ -177,4 +339,5 @@ accelerate launch --zero_stage 3 --zero3_save_16bit_model true --zero3_init_flag
--train_mode="control_ref" \
--control_ref_image="first_frame" \
--low_vram
-```
\ No newline at end of file
+```
+
\ No newline at end of file
diff --git a/scripts/wan2.1_fun/README_TRAIN_LORA.md b/scripts/wan2.1_fun/README_TRAIN_LORA.md
index c4b9de5..3e71008 100755
--- a/scripts/wan2.1_fun/README_TRAIN_LORA.md
+++ b/scripts/wan2.1_fun/README_TRAIN_LORA.md
@@ -24,7 +24,7 @@ Wan T2V without deepspeed:
Wan without DeepSpeed is more suitable for 1.3B Wan, as using it with 14B Wan may result in insufficient GPU memory.
```sh
-export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-Fun-14B-InP"
+export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-Fun-V1.1-14B-InP"
export DATASET_NAME="datasets/internal_datasets/"
export DATASET_META_NAME="datasets/internal_datasets/metadata.json"
export NCCL_IB_DISABLE=1
@@ -69,7 +69,7 @@ Wan T2V with deepspeed zero-2:
Wan with DeepSpeed Zero-2 is suitable for training 1.3B Wan and 14B Wan at low resolutions, but training 14B Wan at high resolutions may still result in insufficient GPU memory.
```sh
-export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-Fun-14B-InP"
+export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-Fun-V1.1-14B-InP"
export DATASET_NAME="datasets/internal_datasets/"
export DATASET_META_NAME="datasets/internal_datasets/metadata.json"
export NCCL_IB_DISABLE=1
@@ -119,7 +119,7 @@ python scripts/zero_to_bf16.py output_dir/checkpoint-{our-num-steps} output_dir/
Training shell command is as follows:
```sh
-export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-Fun-14B-InP"
+export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-Fun-V1.1-14B-InP"
export DATASET_NAME="datasets/internal_datasets/"
export DATASET_META_NAME="datasets/internal_datasets/metadata.json"
export NCCL_IB_DISABLE=1
diff --git a/scripts/wan2.1_fun/train.sh b/scripts/wan2.1_fun/train.sh
old mode 100644
new mode 100755
index ba3b09f..6ea7195
--- a/scripts/wan2.1_fun/train.sh
+++ b/scripts/wan2.1_fun/train.sh
@@ -1,4 +1,4 @@
-export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-Fun-14B-InP"
+export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-Fun-V1.1-14B-InP"
export DATASET_NAME="datasets/internal_datasets/"
export DATASET_META_NAME="datasets/internal_datasets/metadata.json"
export NCCL_IB_DISABLE=1
@@ -41,7 +41,7 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.1_fun/train.py \
--trainable_modules "."
# # Training command for T2V
-# export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-Fun-14B-InP"
+# export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-Fun-V1.1-14B-InP"
# export DATASET_NAME="datasets/internal_datasets/"
# export DATASET_META_NAME="datasets/internal_datasets/metadata.json"
# export NCCL_IB_DISABLE=1
diff --git a/scripts/wan2.1_fun/train_control.sh b/scripts/wan2.1_fun/train_control.sh
old mode 100644
new mode 100755
index 06e2ed6..e56e927
--- a/scripts/wan2.1_fun/train_control.sh
+++ b/scripts/wan2.1_fun/train_control.sh
@@ -1,11 +1,10 @@
-export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-Fun-14B-Control"
+export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-Fun-V1.1-14B-Control"
export DATASET_NAME="datasets/internal_datasets/"
export DATASET_META_NAME="datasets/internal_datasets/metadata.json"
export NCCL_IB_DISABLE=1
export NCCL_P2P_DISABLE=1
NCCL_DEBUG=INFO
-# When train model with multi machines, use "--config_file accelerate.yaml" instead of "--mixed_precision='bf16'".
accelerate launch --mixed_precision="bf16" scripts/wan2.1_fun/train_control.py \
--config_path="config/wan2.1/wan_civitai.yaml" \
--pretrained_model_name_or_path=$MODEL_NAME \
@@ -38,6 +37,7 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.1_fun/train_control.py \
--enable_bucket \
--uniform_sampling \
--low_vram \
- --train_mode="control_object" \
- --control_ref_image="first_frame" \
+ --train_mode="control_ref" \
+ --control_ref_image="random" \
+ --add_full_ref_image_in_self_attention \
--trainable_modules "."
\ No newline at end of file
diff --git a/scripts/wan2.1_fun/train_control_lora.sh b/scripts/wan2.1_fun/train_control_lora.sh
old mode 100644
new mode 100755
index b6392e3..3107a07
--- a/scripts/wan2.1_fun/train_control_lora.sh
+++ b/scripts/wan2.1_fun/train_control_lora.sh
@@ -1,4 +1,4 @@
-export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-Fun-14B-Control"
+export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-Fun-V1.1-14B-Control"
export DATASET_NAME="datasets/internal_datasets/"
export DATASET_META_NAME="datasets/internal_datasets/metadata.json"
export NCCL_IB_DISABLE=1
@@ -35,5 +35,6 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.1_fun/train_control_lora
--enable_bucket \
--uniform_sampling \
--train_mode="control_ref" \
- --control_ref_image="first_frame" \
+ --control_ref_image="random" \
+ --add_full_ref_image_in_self_attention \
--low_vram
\ No newline at end of file
diff --git a/scripts/wan2.1_fun/train_lora.sh b/scripts/wan2.1_fun/train_lora.sh
old mode 100644
new mode 100755
index 15ac536..399b89d
--- a/scripts/wan2.1_fun/train_lora.sh
+++ b/scripts/wan2.1_fun/train_lora.sh
@@ -1,4 +1,4 @@
-export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-Fun-14B-InP"
+export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-Fun-V1.1-14B-InP"
export DATASET_NAME="datasets/internal_datasets/"
export DATASET_META_NAME="datasets/internal_datasets/metadata.json"
export NCCL_IB_DISABLE=1
@@ -38,7 +38,7 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.1_fun/train_lora.py \
--low_vram
# # Training command for T2V
-# export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-Fun-14B-InP"
+# export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-Fun-V1.1-14B-InP"
# export DATASET_NAME="datasets/internal_datasets/"
# export DATASET_META_NAME="datasets/internal_datasets/metadata.json"
# export NCCL_IB_DISABLE=1
diff --git a/scripts/wan2.1_fun/train_reward_lora.sh b/scripts/wan2.1_fun/train_reward_lora.sh
old mode 100644
new mode 100755
index b8404c5..edda3aa
--- a/scripts/wan2.1_fun/train_reward_lora.sh
+++ b/scripts/wan2.1_fun/train_reward_lora.sh
@@ -1,4 +1,4 @@
-export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-Fun-1.3B-InP"
+export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-Fun-V1.1-1.3B-InP"
export TRAIN_PROMPT_PATH="MovieGenVideoBench_train.txt"
# Performing validation simultaneously with training will increase time and GPU memory usage.
export VALIDATION_PROMPT_PATH="MovieGenVideoBench_val.txt"
diff --git a/videox_fun/api/api_multi_nodes.py b/videox_fun/api/api_multi_nodes.py
index 9bf3785..72db091 100755
--- a/videox_fun/api/api_multi_nodes.py
+++ b/videox_fun/api/api_multi_nodes.py
@@ -78,8 +78,9 @@ if ray is not None:
def __init__(
self, rank: int, world_size: int, Controller,
GPU_memory_mode, scheduler_dict, model_name=None, model_type="Inpaint",
- config_path=None, ulysses_degree=1, ring_degree=1, weight_dtype=None,
- savedir_sample=None,
+ config_path=None, ulysses_degree=1, ring_degree=1,
+ fsdp_dit=False, fsdp_text_encoder=False, compile_dit=False,
+ weight_dtype=None, savedir_sample=None,
):
# Set PyTorch distributed environment variables
os.environ["RANK"] = str(rank)
@@ -90,7 +91,9 @@ if ray is not None:
self.rank = rank
self.controller = Controller(
GPU_memory_mode, scheduler_dict, model_name=model_name, model_type=model_type, config_path=config_path,
- ulysses_degree=ulysses_degree, ring_degree=ring_degree, weight_dtype=weight_dtype, savedir_sample=savedir_sample,
+ ulysses_degree=ulysses_degree, ring_degree=ring_degree,
+ fsdp_dit=fsdp_dit, fsdp_text_encoder=fsdp_text_encoder, compile_dit=compile_dit,
+ weight_dtype=weight_dtype, savedir_sample=savedir_sample,
)
def generate(self, datas):
@@ -215,18 +218,39 @@ if ray is not None:
torch.cuda.ipc_collect()
save_sample_path = ""
comment = f"Error. error information is {str(e)}"
- return {"message": comment, "save_sample_path": None, "base64_encoding": None}
-
- if dist.get_rank() == 0:
+ if dist.is_initialized():
+ if dist.get_rank() == 0:
+ return {"message": comment, "save_sample_path": None, "base64_encoding": None}
+ else:
+ return None
+ else:
+ return {"message": comment, "save_sample_path": None, "base64_encoding": None}
+
+
+ if dist.is_initialized():
+ if dist.get_rank() == 0:
+ if save_sample_path != "":
+ return {"message": comment, "save_sample_path": save_sample_path, "base64_encoding": encode_file_to_base64(save_sample_path)}
+ else:
+ return {"message": comment, "save_sample_path": None, "base64_encoding": None}
+ else:
+ return None
+ else:
if save_sample_path != "":
return {"message": comment, "save_sample_path": save_sample_path, "base64_encoding": encode_file_to_base64(save_sample_path)}
else:
- return {"message": comment, "save_sample_path": save_sample_path, "base64_encoding": None}
- return None
+ return {"message": comment, "save_sample_path": None, "base64_encoding": None}
except Exception as e:
- print(f"Error generating image: {str(e)}")
- raise HTTPException(status_code=500, detail=str(e))
+ print(f"Error generating: {str(e)}")
+ comment = f"Error generating: {str(e)}"
+ if dist.is_initialized():
+ if dist.get_rank() == 0:
+ return {"message": comment, "save_sample_path": None, "base64_encoding": None}
+ else:
+ return None
+ else:
+ return {"message": comment, "save_sample_path": None, "base64_encoding": None}
class MultiNodesEngine:
def __init__(
@@ -238,10 +262,13 @@ if ray is not None:
model_name,
model_type,
config_path,
- ulysses_degree,
- ring_degree,
- weight_dtype,
- savedir_sample
+ ulysses_degree=1,
+ ring_degree=1,
+ fsdp_dit=False,
+ fsdp_text_encoder=False,
+ compile_dit=False,
+ weight_dtype=torch.bfloat16,
+ savedir_sample="samples"
):
# Ensure Ray is initialized
if not ray.is_initialized():
@@ -252,7 +279,9 @@ if ray is not None:
MultiNodesGenerator.remote(
rank, world_size, Controller,
GPU_memory_mode, scheduler_dict, model_name=model_name, model_type=model_type, config_path=config_path,
- ulysses_degree=ulysses_degree, ring_degree=ring_degree, weight_dtype=weight_dtype, savedir_sample=savedir_sample,
+ ulysses_degree=ulysses_degree, ring_degree=ring_degree,
+ fsdp_dit=fsdp_dit, fsdp_text_encoder=fsdp_text_encoder, compile_dit=compile_dit,
+ weight_dtype=weight_dtype, savedir_sample=savedir_sample,
)
for rank in range(num_workers)
]
diff --git a/videox_fun/dist/__init__.py b/videox_fun/dist/__init__.py
index 31f1827..1ae7b7d 100755
--- a/videox_fun/dist/__init__.py
+++ b/videox_fun/dist/__init__.py
@@ -12,6 +12,7 @@ try:
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,
@@ -31,6 +32,7 @@ except Exception as ex:
try:
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):
diff --git a/videox_fun/models/cache_utils.py b/videox_fun/models/cache_utils.py
index 117e9fc..8be69b5 100755
--- a/videox_fun/models/cache_utils.py
+++ b/videox_fun/models/cache_utils.py
@@ -1,15 +1,16 @@
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]
- elif "wan2.1-t2v-14b" in model_name.lower() or "wan2.1-fun-v1.1-14b" in model_name.lower():
+ elif "wan2.1-t2v-14b" in model_name.lower():
return [-3.03318725e+05, 4.90537029e+04, -2.65530556e+03, 5.87365115e+01, -3.15583525e-01]
elif "wan2.1-i2v-14b-480p" in model_name.lower():
return [2.57151496e+05, -3.54229917e+04, 1.40286849e+03, -1.35890334e+01, 1.32517977e-01]
- elif "wan2.1-i2v-14b-720p" in model_name.lower() or "wan2.1-fun-14b" in model_name.lower():
+ elif "wan2.1-i2v-14b-720p" in model_name.lower() or "wan2.1-fun-14b" in model_name.lower() or "wan2.1-fun-v1.1-14b" in model_name.lower():
return [8.10705460e+03, 2.13393892e+03, -3.72934672e+02, 1.66203073e+01, -4.17769401e-02]
else:
print(f"The model {model_name} is not supported by TeaCache.")
@@ -71,4 +72,63 @@ class TeaCache():
self.previous_modulated_input = None
self.previous_residual = None
self.previous_residual_cond = None
- self.previous_residual_uncond = None
\ No newline at end of file
+ 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 77c3660..42d2b07 100755
--- a/videox_fun/models/wan_transformer3d.py
+++ b/videox_fun/models/wan_transformer3d.py
@@ -23,7 +23,7 @@ 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
+from .cache_utils import TeaCache, cfg_skip, disable_cfg_skip, enable_cfg_skip
from .wan_camera_adapter import SimpleAdapter
try:
@@ -821,6 +821,9 @@ class WanTransformer3DModel(ModelMixin, ConfigMixin, FromOriginalModelMixin):
self.ref_conv = None
self.teacache = None
+ self.cfg_skip_ratio = None
+ self.current_steps = 0
+ self.num_inference_steps = None
self.gradient_checkpointing = False
self.sp_world_size = 1
self.sp_world_rank = 0
@@ -840,6 +843,23 @@ 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
+ self.current_steps = 0
+ self.num_inference_steps = num_steps
+ else:
+ self.cfg_skip_ratio = None
+ 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
+ self.num_inference_steps = None
+
def enable_riflex(
self,
k = 6,
@@ -876,7 +896,8 @@ class WanTransformer3DModel(ModelMixin, ConfigMixin, FromOriginalModelMixin):
def _set_gradient_checkpointing(self, module, value=False):
self.gradient_checkpointing = value
-
+
+ @cfg_skip()
def forward(
self,
x,
@@ -982,25 +1003,25 @@ class WanTransformer3DModel(ModelMixin, ConfigMixin, FromOriginalModelMixin):
modulated_inp = e0
skip_flag = self.teacache.cnt < self.teacache.num_skip_start_steps
if skip_flag:
- should_calc = True
+ self.should_calc = True
self.teacache.accumulated_rel_l1_distance = 0
else:
if cond_flag:
rel_l1_distance = self.teacache.compute_rel_l1_distance(self.teacache.previous_modulated_input, modulated_inp)
self.teacache.accumulated_rel_l1_distance += self.teacache.rescale_func(rel_l1_distance)
if self.teacache.accumulated_rel_l1_distance < self.teacache.rel_l1_thresh:
- should_calc = False
+ self.should_calc = False
else:
- should_calc = True
+ self.should_calc = True
self.teacache.accumulated_rel_l1_distance = 0
self.teacache.previous_modulated_input = modulated_inp
- self.teacache.should_calc = should_calc
+ self.teacache.should_calc = self.should_calc
else:
- should_calc = self.teacache.should_calc
+ self.should_calc = self.teacache.should_calc
# TeaCache
if self.teacache is not None:
- if not should_calc:
+ if not self.should_calc:
previous_residual = self.teacache.previous_residual_cond if cond_flag else self.teacache.previous_residual_uncond
x = x + previous_residual.to(x.device)
else:
diff --git a/videox_fun/pipeline/pipeline_wan_fun.py b/videox_fun/pipeline/pipeline_wan_fun.py
index 2372ba6..13f4584 100755
--- a/videox_fun/pipeline/pipeline_wan_fun.py
+++ b/videox_fun/pipeline/pipeline_wan_fun.py
@@ -408,7 +408,6 @@ class WanFunPipeline(DiffusionPipeline):
callback_on_step_end_tensor_inputs: List[str] = ["latents"],
max_sequence_length: int = 512,
comfyui_progressbar: bool = False,
- cfg_skip_ratio: int = None,
shift: int = 5,
) -> Union[WanPipelineOutput, Tuple]:
"""
@@ -513,11 +512,10 @@ class WanFunPipeline(DiffusionPipeline):
seq_len = math.ceil((target_shape[2] * target_shape[3]) / (self.transformer.config.patch_size[1] * self.transformer.config.patch_size[2]) * target_shape[1])
# 7. Denoising loop
num_warmup_steps = max(len(timesteps) - num_inference_steps * self.scheduler.order, 0)
+ self.transformer.num_inference_steps = num_inference_steps
with self.progress_bar(total=num_inference_steps) as progress_bar:
for i, t in enumerate(timesteps):
- if cfg_skip_ratio is not None and i >= num_inference_steps * (1 - cfg_skip_ratio):
- do_classifier_free_guidance = False
- in_prompt_embeds = prompt_embeds
+ self.transformer.current_steps = i
if self.interrupt:
continue
diff --git a/videox_fun/pipeline/pipeline_wan_fun_control.py b/videox_fun/pipeline/pipeline_wan_fun_control.py
index 6b3d3ed..80bba08 100755
--- a/videox_fun/pipeline/pipeline_wan_fun_control.py
+++ b/videox_fun/pipeline/pipeline_wan_fun_control.py
@@ -497,7 +497,6 @@ class WanFunControlPipeline(DiffusionPipeline):
clip_image: Image = None,
max_sequence_length: int = 512,
comfyui_progressbar: bool = False,
- cfg_skip_ratio: int = None,
shift: int = 5,
) -> Union[WanPipelineOutput, Tuple]:
"""
@@ -703,11 +702,10 @@ class WanFunControlPipeline(DiffusionPipeline):
seq_len = math.ceil((target_shape[2] * target_shape[3]) / (self.transformer.config.patch_size[1] * self.transformer.config.patch_size[2]) * target_shape[1])
# 7. Denoising loop
num_warmup_steps = max(len(timesteps) - num_inference_steps * self.scheduler.order, 0)
+ self.transformer.num_inference_steps = num_inference_steps
with self.progress_bar(total=num_inference_steps) as progress_bar:
for i, t in enumerate(timesteps):
- if cfg_skip_ratio is not None and i >= num_inference_steps * (1 - cfg_skip_ratio):
- do_classifier_free_guidance = False
- in_prompt_embeds = prompt_embeds
+ self.transformer.current_steps = i
if self.interrupt:
continue
diff --git a/videox_fun/pipeline/pipeline_wan_fun_inpaint.py b/videox_fun/pipeline/pipeline_wan_fun_inpaint.py
index 48e30a2..916593b 100755
--- a/videox_fun/pipeline/pipeline_wan_fun_inpaint.py
+++ b/videox_fun/pipeline/pipeline_wan_fun_inpaint.py
@@ -496,7 +496,6 @@ class WanFunInpaintPipeline(DiffusionPipeline):
clip_image: Image = None,
max_sequence_length: int = 512,
comfyui_progressbar: bool = False,
- cfg_skip_ratio: int = None,
shift: int = 5,
) -> Union[WanPipelineOutput, Tuple]:
"""
@@ -658,11 +657,10 @@ class WanFunInpaintPipeline(DiffusionPipeline):
seq_len = math.ceil((target_shape[2] * target_shape[3]) / (self.transformer.config.patch_size[1] * self.transformer.config.patch_size[2]) * target_shape[1])
# 7. Denoising loop
num_warmup_steps = max(len(timesteps) - num_inference_steps * self.scheduler.order, 0)
+ self.transformer.num_inference_steps = num_inference_steps
with self.progress_bar(total=num_inference_steps) as progress_bar:
for i, t in enumerate(timesteps):
- if cfg_skip_ratio is not None and i >= num_inference_steps * (1 - cfg_skip_ratio):
- do_classifier_free_guidance = False
- in_prompt_embeds = prompt_embeds
+ self.transformer.current_steps = i
if self.interrupt:
continue
diff --git a/videox_fun/ui/cogvideox_fun_ui.py b/videox_fun/ui/cogvideox_fun_ui.py
index f5441ad..f1411e6 100755
--- a/videox_fun/ui/cogvideox_fun_ui.py
+++ b/videox_fun/ui/cogvideox_fun_ui.py
@@ -15,7 +15,8 @@ from ..models import (AutoencoderKLCogVideoX, CogVideoXTransformer3DModel,
T5EncoderModel, T5Tokenizer)
from ..pipeline import (CogVideoXFunControlPipeline,
CogVideoXFunInpaintPipeline, CogVideoXFunPipeline)
-from ..utils.fp8_optimization import convert_weight_dtype_wrapper
+from ..utils.fp8_optimization import (convert_model_weight_to_float8, replace_parameters_by_name,
+ convert_weight_dtype_wrapper)
from ..utils.lora_utils import merge_lora, unmerge_lora
from ..utils.utils import (filter_kwargs, get_image_to_video_latent, get_image_latent, timer,
get_video_to_video_latent, save_videos_grid)
@@ -32,6 +33,7 @@ from .ui import (create_cfg_and_seedbox,
create_height_width, create_model_checkpoints,
create_model_type, create_prompts, create_samplers,
create_ui_outputs)
+from ..dist import set_multi_gpus_devices, shard_model
class CogVideoXFunController(Fun_Controller):
@@ -88,18 +90,34 @@ class CogVideoXFunController(Fun_Controller):
)
if self.ulysses_degree > 1 or self.ring_degree > 1:
+ from functools import partial
self.transformer.enable_multi_gpus_inference()
+ if self.fsdp_dit:
+ shard_fn = partial(shard_model, device_id=self.device, param_dtype=self.weight_dtype)
+ self.pipeline.transformer = shard_fn(self.pipeline.transformer)
+ print("Add FSDP DIT")
+ if self.fsdp_text_encoder:
+ shard_fn = partial(shard_model, device_id=self.device, param_dtype=self.weight_dtype)
+ self.pipeline.text_encoder = shard_fn(self.pipeline.text_encoder)
+ print("Add FSDP TEXT ENCODER")
+
+ if self.compile_dit:
+ for i in range(len(self.pipeline.transformer.transformer_blocks)):
+ self.pipeline.transformer.transformer_blocks[i] = torch.compile(self.pipeline.transformer.transformer_blocks[i])
+ print("Add Compile")
if self.GPU_memory_mode == "sequential_cpu_offload":
self.pipeline.enable_sequential_cpu_offload(device=self.device)
elif self.GPU_memory_mode == "model_cpu_offload_and_qfloat8":
+ convert_model_weight_to_float8(self.pipeline.transformer, exclude_module_name=[], device=self.device)
convert_weight_dtype_wrapper(self.pipeline.transformer, self.weight_dtype)
self.pipeline.enable_model_cpu_offload(device=self.device)
elif self.GPU_memory_mode == "model_cpu_offload":
self.pipeline.enable_model_cpu_offload(device=self.device)
elif self.GPU_memory_mode == "model_full_load_and_qfloat8":
+ convert_model_weight_to_float8(self.pipeline.transformer, exclude_module_name=[], device=self.device)
convert_weight_dtype_wrapper(self.pipeline.transformer, self.weight_dtype)
- self.pipeline.enable_model_cpu_offload(device=self.device)
+ self.pipeline.to(self.device)
else:
self.pipeline.to(self.device)
print("Update diffusion transformer done")
@@ -293,7 +311,11 @@ class CogVideoXFunController(Fun_Controller):
control_video = input_video,
).videos
except Exception as e:
+ self.auto_model_clear_cache(self.pipeline.transformer)
+ self.auto_model_clear_cache(self.pipeline.text_encoder)
+ self.auto_model_clear_cache(self.pipeline.vae)
self.clear_cache()
+
print(f"Error. error information is {str(e)}")
if self.lora_model_path != "none":
self.pipeline = unmerge_lora(self.pipeline, self.lora_model_path, multiplier=lora_alpha_slider)
@@ -335,12 +357,11 @@ class CogVideoXFunController(Fun_Controller):
CogVideoXFunController_Host = CogVideoXFunController
CogVideoXFunController_Client = Fun_Controller_Client
-def ui(GPU_memory_mode, scheduler_dict, ulysses_degree, ring_degree, weight_dtype, savedir_sample=None):
+def ui(GPU_memory_mode, scheduler_dict, compile_dit, weight_dtype, savedir_sample=None):
controller = CogVideoXFunController(
GPU_memory_mode, scheduler_dict, model_name=None, model_type="Inpaint",
- ulysses_degree=ulysses_degree, ring_degree=ring_degree,
- config_path=None, enable_teacache=None, teacache_threshold=None, weight_dtype=weight_dtype,
- savedir_sample=savedir_sample,
+ compile_dit=compile_dit,
+ weight_dtype=weight_dtype, savedir_sample=savedir_sample,
)
with gr.Blocks(css=css) as demo:
@@ -383,7 +404,7 @@ def ui(GPU_memory_mode, scheduler_dict, ulysses_degree, ring_degree, weight_dtyp
default_video_length=49,
maximum_video_length=85,
)
- image_to_video_col, video_to_video_col, control_video_col, source_method, start_image, template_gallery, end_image, validation_video, validation_video_mask, denoise_strength, control_video = create_generation_method(
+ image_to_video_col, video_to_video_col, control_video_col, source_method, start_image, template_gallery, end_image, validation_video, validation_video_mask, denoise_strength, control_video, ref_image = create_generation_method(
["Text to Video (文本到视频)", "Image to Video (图片到视频)", "Video to Video (视频到视频)", "Video Control (视频控制)"], prompt_textbox
)
cfg_scale_slider, seed_textbox, seed_button = create_cfg_and_seedbox(gradio_version_is_above_4)
@@ -466,12 +487,11 @@ def ui(GPU_memory_mode, scheduler_dict, ulysses_degree, ring_degree, weight_dtyp
)
return demo, controller
-def ui_host(GPU_memory_mode, scheduler_dict, model_name, model_type, ulysses_degree, ring_degree, weight_dtype, savedir_sample=None):
+def ui_host(GPU_memory_mode, scheduler_dict, model_name, model_type, compile_dit, weight_dtype, savedir_sample=None):
controller = CogVideoXFunController_Host(
GPU_memory_mode, scheduler_dict, model_name=model_name, model_type=model_type,
- ulysses_degree=ulysses_degree, ring_degree=ring_degree,
- config_path=None, enable_teacache=None, teacache_threshold=None, weight_dtype=weight_dtype,
- savedir_sample=savedir_sample,
+ compile_dit=compile_dit,
+ weight_dtype=weight_dtype, savedir_sample=savedir_sample,
)
with gr.Blocks(css=css) as demo:
@@ -512,7 +532,7 @@ def ui_host(GPU_memory_mode, scheduler_dict, model_name, model_type, ulysses_deg
default_video_length=49,
maximum_video_length=85,
)
- image_to_video_col, video_to_video_col, control_video_col, source_method, start_image, template_gallery, end_image, validation_video, validation_video_mask, denoise_strength, control_video = create_generation_method(
+ image_to_video_col, video_to_video_col, control_video_col, source_method, start_image, template_gallery, end_image, validation_video, validation_video_mask, denoise_strength, control_video, ref_image = create_generation_method(
["Text to Video (文本到视频)", "Image to Video (图片到视频)", "Video to Video (视频到视频)", "Video Control (视频控制)"], prompt_textbox
)
cfg_scale_slider, seed_textbox, seed_button = create_cfg_and_seedbox(gradio_version_is_above_4)
@@ -627,7 +647,7 @@ def ui_client(scheduler_dict, model_name, savedir_sample=None):
default_video_length=49,
maximum_video_length=85,
)
- image_to_video_col, video_to_video_col, control_video_col, source_method, start_image, template_gallery, end_image, validation_video, validation_video_mask, denoise_strength, control_video = create_generation_method(
+ image_to_video_col, video_to_video_col, control_video_col, source_method, start_image, template_gallery, end_image, validation_video, validation_video_mask, denoise_strength, control_video, ref_image = create_generation_method(
["Text to Video (文本到视频)", "Image to Video (图片到视频)", "Video to Video (视频到视频)"], prompt_textbox
)
diff --git a/videox_fun/ui/controller.py b/videox_fun/ui/controller.py
index a7a7b55..cb160d2 100755
--- a/videox_fun/ui/controller.py
+++ b/videox_fun/ui/controller.py
@@ -59,7 +59,9 @@ all_cheduler_dict = {**ddpm_scheduler_dict, **flow_scheduler_dict}
class Fun_Controller:
def __init__(
self, GPU_memory_mode, scheduler_dict, model_name=None, model_type="Inpaint",
- config_path=None, ulysses_degree=1, ring_degree=1, weight_dtype=None, savedir_sample=None,
+ config_path=None, ulysses_degree=1, ring_degree=1,
+ fsdp_dit=False, fsdp_text_encoder=False, compile_dit=False,
+ weight_dtype=None, savedir_sample=None,
):
# config dirs
self.basedir = os.getcwd()
@@ -82,6 +84,9 @@ class Fun_Controller:
self.config = OmegaConf.load(config_path)
self.ulysses_degree = ulysses_degree
self.ring_degree = ring_degree
+ self.fsdp_dit = fsdp_dit
+ self.fsdp_text_encoder = fsdp_text_encoder
+ self.compile_dit = compile_dit
self.weight_dtype = weight_dtype
self.device = set_multi_gpus_devices(self.ulysses_degree, self.ring_degree)
@@ -148,6 +153,14 @@ class Fun_Controller:
gc.collect()
torch.cuda.empty_cache()
torch.cuda.ipc_collect()
+
+ def auto_model_clear_cache(self, model):
+ origin_device = model.device
+ model = model.to("cpu")
+ gc.collect()
+ torch.cuda.empty_cache()
+ torch.cuda.ipc_collect()
+ model = model.to(origin_device)
def input_check(self,
resize_method,
diff --git a/videox_fun/ui/wan_fun_ui.py b/videox_fun/ui/wan_fun_ui.py
index 180bfdc..604547d 100755
--- a/videox_fun/ui/wan_fun_ui.py
+++ b/videox_fun/ui/wan_fun_ui.py
@@ -37,6 +37,7 @@ from .ui import (create_cfg_and_seedbox, create_cfg_riflex_k,
create_height_width, create_model_checkpoints,
create_model_type, create_prompts, create_samplers,
create_teacache_params, create_ui_outputs)
+from ..dist import set_multi_gpus_devices, shard_model
class Wan_Fun_Controller(Fun_Controller):
@@ -117,22 +118,36 @@ class Wan_Fun_Controller(Fun_Controller):
)
if self.ulysses_degree > 1 or self.ring_degree > 1:
+ from functools import partial
self.transformer.enable_multi_gpus_inference()
+ if self.fsdp_dit:
+ shard_fn = partial(shard_model, device_id=self.device, param_dtype=self.weight_dtype)
+ self.pipeline.transformer = shard_fn(self.pipeline.transformer)
+ print("Add FSDP DIT")
+ if self.fsdp_text_encoder:
+ shard_fn = partial(shard_model, device_id=self.device, param_dtype=self.weight_dtype)
+ self.pipeline.text_encoder = shard_fn(self.pipeline.text_encoder)
+ print("Add FSDP TEXT ENCODER")
+
+ if self.compile_dit:
+ for i in range(len(self.pipeline.transformer.blocks)):
+ self.pipeline.transformer.blocks[i] = torch.compile(self.pipeline.transformer.blocks[i])
+ print("Add Compile")
if self.GPU_memory_mode == "sequential_cpu_offload":
replace_parameters_by_name(self.transformer, ["modulation",], device=self.device)
self.transformer.freqs = self.transformer.freqs.to(device=self.device)
self.pipeline.enable_sequential_cpu_offload(device=self.device)
elif self.GPU_memory_mode == "model_cpu_offload_and_qfloat8":
- convert_model_weight_to_float8(self.transformer, exclude_module_name=["modulation",])
+ convert_model_weight_to_float8(self.transformer, exclude_module_name=["modulation",], device=self.device)
convert_weight_dtype_wrapper(self.transformer, self.weight_dtype)
self.pipeline.enable_model_cpu_offload(device=self.device)
elif self.GPU_memory_mode == "model_cpu_offload":
self.pipeline.enable_model_cpu_offload(device=self.device)
elif self.GPU_memory_mode == "model_full_load_and_qfloat8":
- convert_model_weight_to_float8(self.transformer, exclude_module_name=["modulation",])
+ convert_model_weight_to_float8(self.transformer, exclude_module_name=["modulation",], device=self.device)
convert_weight_dtype_wrapper(self.transformer, self.weight_dtype)
- self.pipeline.to(device=self.device)
+ self.pipeline.to(self.device)
else:
self.pipeline.to(self.device)
print("Update diffusion transformer done")
@@ -219,7 +234,11 @@ class Wan_Fun_Controller(Fun_Controller):
else:
print(f"Disable TeaCache.")
self.pipeline.transformer.disable_teacache()
-
+
+ if cfg_skip_ratio is not None and cfg_skip_ratio >= 0:
+ print(f"Enable cfg_skip_ratio {cfg_skip_ratio}.")
+ self.pipeline.transformer.enable_cfg_skip(cfg_skip_ratio, sample_step_slider)
+
print(f"Generate seed.")
if int(seed_textbox) != -1 and seed_textbox != "": torch.manual_seed(int(seed_textbox))
else: seed_textbox = np.random.randint(0, 1e10)
@@ -253,7 +272,6 @@ class Wan_Fun_Controller(Fun_Controller):
video = input_video,
mask_video = input_video_mask,
clip_image = clip_image,
- cfg_skip_ratio = cfg_skip_ratio,
).videos
else:
sample = self.pipeline(
@@ -265,7 +283,6 @@ class Wan_Fun_Controller(Fun_Controller):
height = height_slider,
num_frames = length_slider if not is_image else 1,
generator = generator,
- cfg_skip_ratio = cfg_skip_ratio,
).videos
else:
if ref_image is not None:
@@ -297,11 +314,14 @@ class Wan_Fun_Controller(Fun_Controller):
ref_image = ref_image,
start_image = start_image,
clip_image = clip_image,
- cfg_skip_ratio = cfg_skip_ratio,
).videos
print(f"Generation done.")
except Exception as e:
+ self.auto_model_clear_cache(self.pipeline.transformer)
+ self.auto_model_clear_cache(self.pipeline.text_encoder)
+ self.auto_model_clear_cache(self.pipeline.vae)
self.clear_cache()
+
print(f"Error. error information is {str(e)}")
if self.lora_model_path != "none":
self.pipeline = unmerge_lora(self.pipeline, self.lora_model_path, multiplier=lora_alpha_slider)
@@ -343,10 +363,10 @@ class Wan_Fun_Controller(Fun_Controller):
Wan_Fun_Controller_Host = Wan_Fun_Controller
Wan_Fun_Controller_Client = Fun_Controller_Client
-def ui(GPU_memory_mode, scheduler_dict, config_path, ulysses_degree, ring_degree, weight_dtype, savedir_sample=None):
+def ui(GPU_memory_mode, scheduler_dict, config_path, compile_dit, weight_dtype, savedir_sample=None):
controller = Wan_Fun_Controller(
GPU_memory_mode, scheduler_dict, model_name=None, model_type="Inpaint",
- config_path=config_path, ulysses_degree=ulysses_degree, ring_degree=ring_degree,
+ config_path=config_path, compile_dit=compile_dit,
weight_dtype=weight_dtype, savedir_sample=savedir_sample,
)
@@ -481,10 +501,10 @@ def ui(GPU_memory_mode, scheduler_dict, config_path, ulysses_degree, ring_degree
)
return demo, controller
-def ui_host(GPU_memory_mode, scheduler_dict, model_name, model_type, config_path, ulysses_degree, ring_degree, weight_dtype, savedir_sample=None):
+def ui_host(GPU_memory_mode, scheduler_dict, model_name, model_type, config_path, compile_dit, weight_dtype, savedir_sample=None):
controller = Wan_Fun_Controller_Host(
GPU_memory_mode, scheduler_dict, model_name=model_name, model_type=model_type,
- config_path=config_path, ulysses_degree=ulysses_degree, ring_degree=ring_degree,
+ config_path=config_path, compile_dit=compile_dit,
weight_dtype=weight_dtype, savedir_sample=savedir_sample,
)
diff --git a/videox_fun/ui/wan_ui.py b/videox_fun/ui/wan_ui.py
index 1d2fd55..01fac3d 100755
--- a/videox_fun/ui/wan_ui.py
+++ b/videox_fun/ui/wan_ui.py
@@ -36,6 +36,7 @@ from .ui import (create_cfg_and_seedbox, create_cfg_riflex_k,
create_height_width, create_model_checkpoints,
create_model_type, create_prompts, create_samplers,
create_teacache_params, create_ui_outputs)
+from ..dist import set_multi_gpus_devices, shard_model
class Wan_Controller(Fun_Controller):
@@ -109,22 +110,36 @@ class Wan_Controller(Fun_Controller):
raise ValueError("Not support now")
if self.ulysses_degree > 1 or self.ring_degree > 1:
+ from functools import partial
self.transformer.enable_multi_gpus_inference()
+ if self.fsdp_dit:
+ shard_fn = partial(shard_model, device_id=self.device, param_dtype=self.weight_dtype)
+ self.pipeline.transformer = shard_fn(self.pipeline.transformer)
+ print("Add FSDP DIT")
+ if self.fsdp_text_encoder:
+ shard_fn = partial(shard_model, device_id=self.device, param_dtype=self.weight_dtype)
+ self.pipeline.text_encoder = shard_fn(self.pipeline.text_encoder)
+ print("Add FSDP TEXT ENCODER")
+
+ if self.compile_dit:
+ for i in range(len(self.pipeline.transformer.blocks)):
+ self.pipeline.transformer.blocks[i] = torch.compile(self.pipeline.transformer.blocks[i])
+ print("Add Compile")
if self.GPU_memory_mode == "sequential_cpu_offload":
replace_parameters_by_name(self.transformer, ["modulation",], device=self.device)
self.transformer.freqs = self.transformer.freqs.to(device=self.device)
self.pipeline.enable_sequential_cpu_offload(device=self.device)
elif self.GPU_memory_mode == "model_cpu_offload_and_qfloat8":
- convert_model_weight_to_float8(self.transformer, exclude_module_name=["modulation",])
+ convert_model_weight_to_float8(self.transformer, exclude_module_name=["modulation",], device=self.device)
convert_weight_dtype_wrapper(self.transformer, self.weight_dtype)
self.pipeline.enable_model_cpu_offload(device=self.device)
elif self.GPU_memory_mode == "model_cpu_offload":
self.pipeline.enable_model_cpu_offload(device=self.device)
elif self.GPU_memory_mode == "model_full_load_and_qfloat8":
- convert_model_weight_to_float8(self.transformer, exclude_module_name=["modulation",])
+ convert_model_weight_to_float8(self.transformer, exclude_module_name=["modulation",], device=self.device)
convert_weight_dtype_wrapper(self.transformer, self.weight_dtype)
- self.pipeline.to(device=self.device)
+ self.pipeline.to(self.device)
else:
self.pipeline.to(self.device)
print("Update diffusion transformer done")
@@ -211,7 +226,11 @@ class Wan_Controller(Fun_Controller):
else:
print(f"Disable TeaCache.")
self.pipeline.transformer.disable_teacache()
-
+
+ if cfg_skip_ratio is not None and cfg_skip_ratio >= 0:
+ print(f"Enable cfg_skip_ratio {cfg_skip_ratio}.")
+ self.pipeline.transformer.enable_cfg_skip(cfg_skip_ratio, sample_step_slider)
+
print(f"Generate seed.")
if int(seed_textbox) != -1 and seed_textbox != "": torch.manual_seed(int(seed_textbox))
else: seed_textbox = np.random.randint(0, 1e10)
@@ -245,7 +264,6 @@ class Wan_Controller(Fun_Controller):
video = input_video,
mask_video = input_video_mask,
clip_image = clip_image,
- cfg_skip_ratio = cfg_skip_ratio,
).videos
else:
sample = self.pipeline(
@@ -257,7 +275,6 @@ class Wan_Controller(Fun_Controller):
height = height_slider,
num_frames = length_slider if not is_image else 1,
generator = generator,
- cfg_skip_ratio = cfg_skip_ratio,
).videos
else:
if ref_image is not None:
@@ -289,11 +306,14 @@ class Wan_Controller(Fun_Controller):
ref_image = ref_image,
start_image = start_image,
clip_image = clip_image,
- cfg_skip_ratio = cfg_skip_ratio,
).videos
print(f"Generation done.")
except Exception as e:
+ self.auto_model_clear_cache(self.pipeline.transformer)
+ self.auto_model_clear_cache(self.pipeline.text_encoder)
+ self.auto_model_clear_cache(self.pipeline.vae)
self.clear_cache()
+
print(f"Error. error information is {str(e)}")
if self.lora_model_path != "none":
self.pipeline = unmerge_lora(self.pipeline, self.lora_model_path, multiplier=lora_alpha_slider)
@@ -335,10 +355,10 @@ class Wan_Controller(Fun_Controller):
Wan_Controller_Host = Wan_Controller
Wan_Controller_Client = Fun_Controller_Client
-def ui(GPU_memory_mode, scheduler_dict, config_path, ulysses_degree, ring_degree, weight_dtype, savedir_sample=None):
+def ui(GPU_memory_mode, scheduler_dict, config_path, compile_dit, weight_dtype, savedir_sample=None):
controller = Wan_Controller(
GPU_memory_mode, scheduler_dict, model_name=None, model_type="Inpaint",
- config_path=config_path, ulysses_degree=ulysses_degree, ring_degree=ring_degree,
+ config_path=config_path, compile_dit=compile_dit,
weight_dtype=weight_dtype, savedir_sample=savedir_sample,
)
@@ -469,10 +489,10 @@ def ui(GPU_memory_mode, scheduler_dict, config_path, ulysses_degree, ring_degree
)
return demo, controller
-def ui_host(GPU_memory_mode, scheduler_dict, model_name, model_type, config_path, ulysses_degree, ring_degree, weight_dtype, savedir_sample=None):
+def ui_host(GPU_memory_mode, scheduler_dict, model_name, model_type, config_path, compile_dit, weight_dtype, savedir_sample=None):
controller = Wan_Controller_Host(
GPU_memory_mode, scheduler_dict, model_name=model_name, model_type=model_type,
- config_path=config_path, ulysses_degree=ulysses_degree, ring_degree=ring_degree,
+ config_path=config_path, compile_dit=compile_dit,
weight_dtype=weight_dtype, savedir_sample=savedir_sample,
)
diff --git a/videox_fun/utils/fp8_optimization.py b/videox_fun/utils/fp8_optimization.py
old mode 100644
new mode 100755
index 1aa6d26..c51e54f
--- a/videox_fun/utils/fp8_optimization.py
+++ b/videox_fun/utils/fp8_optimization.py
@@ -1,19 +1,10 @@
"""Modified from https://github.com/kijai/ComfyUI-MochiWrapper
"""
+import importlib.util
+
import torch
import torch.nn as nn
-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)
-
- cls.to(weight_dtype)
- return out
-
def replace_parameters_by_name(module, name_keywords, device):
from torch import nn
for name, param in list(module.named_parameters(recurse=False)):
@@ -25,32 +16,48 @@ 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)
-def convert_model_weight_to_float8(model, exclude_module_name=['embed_tokens']):
- 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():
+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():
flag = False
for _exclude_module_name in exclude_module_name:
- if _exclude_module_name in param_name:
+ if _exclude_module_name in name:
flag = True
if flag:
continue
- param.data = param.data.to(torch.float8_e4m3fn)
+ 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)
-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 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)
+
+ 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