Update cfg skip to wrapper && Update Teacache && Update Reamde (#200)

This commit is contained in:
Bubbliiiing
2025-05-09 18:11:53 +08:00
committed by GitHub
parent d7a37ef884
commit f26f0a809b
41 changed files with 1008 additions and 194 deletions
+29
View File
@@ -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:
+29
View File
@@ -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を例として説明します。
+23
View File
@@ -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为例。
+5 -2
View File
@@ -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(
+26 -3
View File
@@ -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():
+20 -2
View File
@@ -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:
+20 -2
View File
@@ -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:
+21 -3
View File
@@ -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:
+20 -2
View File
@@ -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:
+5 -2
View File
@@ -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
+26 -2
View File
@@ -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():
+13 -5
View File
@@ -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
+13 -5
View File
@@ -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
+5 -2
View File
@@ -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
+26 -2
View File
@@ -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():
+13 -5
View File
@@ -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
+18 -5
View File
@@ -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
+13 -5
View File
@@ -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
@@ -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
+13 -5
View File
@@ -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
+3 -3
View File
@@ -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
+176 -6
View File
@@ -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 "."
```
<details>
<summary>(Obsolete) V1.0:</summary>
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 "."
```
```
</details>
+169 -6
View File
@@ -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
```
<details>
<summary>(Obsolete) V1.0:</summary>
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
```
```
</details>
+3 -3
View File
@@ -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
Regular → Executable
+2 -2
View File
@@ -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
+4 -4
View File
@@ -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 "."
+3 -2
View File
@@ -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
Regular → Executable
+2 -2
View File
@@ -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
+1 -1
View File
@@ -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"
+44 -15
View File
@@ -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)
]
+2
View File
@@ -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):
+63 -3
View File
@@ -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
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
+29 -8
View File
@@ -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:
+2 -4
View File
@@ -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
@@ -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
@@ -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
+33 -13
View File
@@ -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
)
+14 -1
View File
@@ -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,
+31 -11
View File
@@ -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,
)
+31 -11
View File
@@ -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,
)
+41 -34
View File
@@ -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)
)