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