Update cfg skip to wrapper && Update Teacache && Update Reamde (#200)
This commit is contained in:
@@ -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:
|
||||
|
||||
|
||||
@@ -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を例として説明します。
|
||||
|
||||
@@ -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为例。
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -20,7 +20,29 @@ from videox_fun.ui.cogvideox_fun_ui import CogVideoXFunController
|
||||
def main():
|
||||
parser = argparse.ArgumentParser(description='xDiT HTTP Service')
|
||||
parser.add_argument('--world_size', type=int, default=8, help='Number of parallel workers')
|
||||
parser.add_argument('--gpu_memory_mode', type=str, default="model_full_load", help='GPU memory mode')
|
||||
parser.add_argument(
|
||||
'--gpu_memory_mode', type=str, default="model_cpu_offload", help='''
|
||||
GPU memory mode, which can be choosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8].
|
||||
model_full_load means that the entire model will be moved to the GPU.
|
||||
|
||||
model_full_load_and_qfloat8 means that the entire model will be moved to the GPU,
|
||||
and the transformer model has been quantized to float8, which can save more GPU memory.
|
||||
|
||||
model_cpu_offload means that the entire model will be moved to the CPU after use, which can save some GPU memory.
|
||||
|
||||
model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
|
||||
and the transformer model has been quantized to float8, which can save more GPU memory.
|
||||
'''
|
||||
)
|
||||
parser.add_argument(
|
||||
'--compile_dit', action='store_true', help='''
|
||||
Enable compile dit.
|
||||
Compile will give a speedup in fixed resolution and need a little GPU memory.
|
||||
The compile_dit is not compatible with the fsdp_dit and sequential_cpu_offload.
|
||||
'''
|
||||
)
|
||||
parser.add_argument('--fsdp_dit', action='store_true', help="Use DIT FSDP to save more GPU memory in multi gpus.")
|
||||
parser.add_argument('--fsdp_text_encoder', action='store_true', help="Use Text Encoder FSDP to save more GPU memory in multi gpus.")
|
||||
parser.add_argument('--ulysses_degree', type=int, default=4, help='Degree of Ulysses configuration')
|
||||
parser.add_argument('--ring_degree', type=int, default=2, help='Degree of Ring configuration')
|
||||
parser.add_argument('--weight_dtype', type=str, default='bf16', help='Weight data type')
|
||||
@@ -40,8 +62,9 @@ def main():
|
||||
engine = MultiNodesEngine(
|
||||
world_size=args.world_size, Controller=CogVideoXFunController,
|
||||
GPU_memory_mode=args.gpu_memory_mode, scheduler_dict=flow_scheduler_dict, model_name=args.model_name, model_type=args.model_type, config_path=None,
|
||||
ulysses_degree=args.ulysses_degree, ring_degree=args.ring_degree, enable_teacache=False, teacache_threshold=0.1, num_skip_start_steps=5,
|
||||
teacache_offload=False, weight_dtype=weight_dtype, savedir_sample=args.savedir_sample,
|
||||
ulysses_degree=args.ulysses_degree, ring_degree=args.ring_degree,
|
||||
fsdp_dit=args.fsdp_dit, fsdp_text_encoder=args.fsdp_text_encoder, compile_dit=args.compile_dit,
|
||||
weight_dtype=weight_dtype, savedir_sample=args.savedir_sample,
|
||||
)
|
||||
|
||||
def gr_launch():
|
||||
|
||||
@@ -20,7 +20,8 @@ from videox_fun.models import (AutoencoderKLCogVideoX,
|
||||
T5Tokenizer)
|
||||
from videox_fun.pipeline import (CogVideoXFunInpaintPipeline,
|
||||
CogVideoXFunPipeline)
|
||||
from videox_fun.utils.fp8_optimization import convert_weight_dtype_wrapper
|
||||
from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8, replace_parameters_by_name,
|
||||
convert_weight_dtype_wrapper)
|
||||
from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
|
||||
from videox_fun.utils.utils import get_image_to_video_latent, save_videos_grid
|
||||
|
||||
@@ -44,7 +45,12 @@ GPU_memory_mode = "model_cpu_offload_and_qfloat8"
|
||||
# If you are using 1 GPU, you can set ulysses_degree = 1 and ring_degree = 1.
|
||||
ulysses_degree = 1
|
||||
ring_degree = 1
|
||||
# Use FSDP to save more GPU memory in multi gpus.
|
||||
fsdp_dit = False
|
||||
fsdp_text_encoder = True
|
||||
# Compile will give a speedup in fixed resolution and need a little GPU memory.
|
||||
# The compile_dit is not compatible with the fsdp_dit and sequential_cpu_offload.
|
||||
compile_dit = False
|
||||
|
||||
# Config and model path
|
||||
model_name = "models/Diffusion_Transformer/CogVideoX-Fun-V1.1-2b-InP"
|
||||
@@ -90,7 +96,7 @@ transformer = CogVideoXTransformer3DModel.from_pretrained(
|
||||
model_name,
|
||||
subfolder="transformer",
|
||||
low_cpu_mem_usage=True if not fsdp_dit else False,
|
||||
torch_dtype=torch.float8_e4m3fn if GPU_memory_mode == "model_cpu_offload_and_qfloat8" else weight_dtype,
|
||||
torch_dtype=weight_dtype,
|
||||
).to(weight_dtype)
|
||||
|
||||
if transformer_path is not None:
|
||||
@@ -167,15 +173,27 @@ if ulysses_degree > 1 or ring_degree > 1:
|
||||
if fsdp_dit:
|
||||
shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype)
|
||||
pipeline.transformer = shard_fn(pipeline.transformer)
|
||||
print("Add FSDP DIT")
|
||||
if fsdp_text_encoder:
|
||||
shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype)
|
||||
pipeline.text_encoder = shard_fn(pipeline.text_encoder)
|
||||
print("Add FSDP TEXT ENCODER")
|
||||
|
||||
if compile_dit:
|
||||
for i in range(len(pipeline.transformer.transformer_blocks)):
|
||||
pipeline.transformer.transformer_blocks[i] = torch.compile(pipeline.transformer.transformer_blocks[i])
|
||||
print("Add Compile")
|
||||
|
||||
if GPU_memory_mode == "sequential_cpu_offload":
|
||||
pipeline.enable_sequential_cpu_offload(device=device)
|
||||
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
|
||||
convert_model_weight_to_float8(transformer, exclude_module_name=[], device=device)
|
||||
convert_weight_dtype_wrapper(transformer, weight_dtype)
|
||||
pipeline.enable_model_cpu_offload(device=device)
|
||||
elif GPU_memory_mode == "model_cpu_offload":
|
||||
pipeline.enable_model_cpu_offload(device=device)
|
||||
elif GPU_memory_mode == "model_full_load_and_qfloat8":
|
||||
convert_model_weight_to_float8(transformer, exclude_module_name=[], device=device)
|
||||
convert_weight_dtype_wrapper(transformer, weight_dtype)
|
||||
pipeline.to(device=device)
|
||||
else:
|
||||
|
||||
@@ -20,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:
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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():
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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():
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -51,8 +51,9 @@ ulysses_degree = 1
|
||||
ring_degree = 1
|
||||
# Use FSDP to save more GPU memory in multi gpus.
|
||||
fsdp_dit = False
|
||||
fsdp_text_encoder = True
|
||||
# Compile will give a speedup in fixed resolution and need a little GPU memory.
|
||||
# The compile_dit is not compatible with the fsdp_dit.
|
||||
# The compile_dit is not compatible with the fsdp_dit and sequential_cpu_offload.
|
||||
compile_dit = False
|
||||
|
||||
# Support TeaCache.
|
||||
@@ -208,7 +209,11 @@ if ulysses_degree > 1 or ring_degree > 1:
|
||||
if fsdp_dit:
|
||||
shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype)
|
||||
pipeline.transformer = shard_fn(pipeline.transformer)
|
||||
print("Add FSDP")
|
||||
print("Add FSDP DIT")
|
||||
if fsdp_text_encoder:
|
||||
shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype)
|
||||
pipeline.text_encoder = shard_fn(pipeline.text_encoder)
|
||||
print("Add FSDP TEXT ENCODER")
|
||||
|
||||
if compile_dit:
|
||||
for i in range(len(pipeline.transformer.blocks)):
|
||||
@@ -220,13 +225,13 @@ if GPU_memory_mode == "sequential_cpu_offload":
|
||||
transformer.freqs = transformer.freqs.to(device=device)
|
||||
pipeline.enable_sequential_cpu_offload(device=device)
|
||||
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
|
||||
convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",])
|
||||
convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
|
||||
convert_weight_dtype_wrapper(transformer, weight_dtype)
|
||||
pipeline.enable_model_cpu_offload(device=device)
|
||||
elif GPU_memory_mode == "model_cpu_offload":
|
||||
pipeline.enable_model_cpu_offload(device=device)
|
||||
elif GPU_memory_mode == "model_full_load_and_qfloat8":
|
||||
convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",])
|
||||
convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
|
||||
convert_weight_dtype_wrapper(transformer, weight_dtype)
|
||||
pipeline.to(device=device)
|
||||
else:
|
||||
@@ -239,6 +244,10 @@ if coefficients is not None:
|
||||
coefficients, num_inference_steps, teacache_threshold, num_skip_start_steps=num_skip_start_steps, offload=teacache_offload
|
||||
)
|
||||
|
||||
if cfg_skip_ratio is not None:
|
||||
print(f"Enable cfg_skip_ratio {cfg_skip_ratio}.")
|
||||
pipeline.transformer.enable_cfg_skip(cfg_skip_ratio, num_inference_steps)
|
||||
|
||||
generator = torch.Generator(device=device).manual_seed(seed)
|
||||
|
||||
if lora_path is not None:
|
||||
@@ -287,7 +296,6 @@ with torch.no_grad():
|
||||
ref_image = ref_image,
|
||||
start_image = start_image,
|
||||
clip_image = clip_image,
|
||||
cfg_skip_ratio = cfg_skip_ratio,
|
||||
shift = shift,
|
||||
).videos
|
||||
|
||||
|
||||
@@ -51,8 +51,9 @@ ulysses_degree = 1
|
||||
ring_degree = 1
|
||||
# Use FSDP to save more GPU memory in multi gpus.
|
||||
fsdp_dit = False
|
||||
fsdp_text_encoder = True
|
||||
# Compile will give a speedup in fixed resolution and need a little GPU memory.
|
||||
# The compile_dit is not compatible with the fsdp_dit.
|
||||
# The compile_dit is not compatible with the fsdp_dit and sequential_cpu_offload.
|
||||
compile_dit = False
|
||||
|
||||
# Support TeaCache.
|
||||
@@ -208,7 +209,11 @@ if ulysses_degree > 1 or ring_degree > 1:
|
||||
if fsdp_dit:
|
||||
shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype)
|
||||
pipeline.transformer = shard_fn(pipeline.transformer)
|
||||
print("Add FSDP")
|
||||
print("Add FSDP DIT")
|
||||
if fsdp_text_encoder:
|
||||
shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype)
|
||||
pipeline.text_encoder = shard_fn(pipeline.text_encoder)
|
||||
print("Add FSDP TEXT ENCODER")
|
||||
|
||||
if compile_dit:
|
||||
for i in range(len(pipeline.transformer.blocks)):
|
||||
@@ -220,13 +225,13 @@ if GPU_memory_mode == "sequential_cpu_offload":
|
||||
transformer.freqs = transformer.freqs.to(device=device)
|
||||
pipeline.enable_sequential_cpu_offload(device=device)
|
||||
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
|
||||
convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",])
|
||||
convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
|
||||
convert_weight_dtype_wrapper(transformer, weight_dtype)
|
||||
pipeline.enable_model_cpu_offload(device=device)
|
||||
elif GPU_memory_mode == "model_cpu_offload":
|
||||
pipeline.enable_model_cpu_offload(device=device)
|
||||
elif GPU_memory_mode == "model_full_load_and_qfloat8":
|
||||
convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",])
|
||||
convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
|
||||
convert_weight_dtype_wrapper(transformer, weight_dtype)
|
||||
pipeline.to(device=device)
|
||||
else:
|
||||
@@ -239,6 +244,10 @@ if coefficients is not None:
|
||||
coefficients, num_inference_steps, teacache_threshold, num_skip_start_steps=num_skip_start_steps, offload=teacache_offload
|
||||
)
|
||||
|
||||
if cfg_skip_ratio is not None:
|
||||
print(f"Enable cfg_skip_ratio {cfg_skip_ratio}.")
|
||||
pipeline.transformer.enable_cfg_skip(cfg_skip_ratio, num_inference_steps)
|
||||
|
||||
generator = torch.Generator(device=device).manual_seed(seed)
|
||||
|
||||
if lora_path is not None:
|
||||
@@ -287,7 +296,6 @@ with torch.no_grad():
|
||||
ref_image = ref_image,
|
||||
start_image = start_image,
|
||||
clip_image = clip_image,
|
||||
cfg_skip_ratio = cfg_skip_ratio,
|
||||
shift = shift,
|
||||
).videos
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
Regular → Executable
+176
-6
@@ -38,8 +38,179 @@ Some parameters in the sh file can be confusing, and they are explained in this
|
||||
- At 1024x1024 resolution, the number of video frames is 9 (~= 512 * 512 * 49 / 1024 / 1024).
|
||||
- These resolutions combined with their corresponding lengths allow the model to generate videos of different sizes.
|
||||
- `resume_from_checkpoint` is used to set the training should be resumed from a previous checkpoint. Use a path or `"latest"` to automatically select the last available checkpoint.
|
||||
- `train_mode` is used to set the training mode.
|
||||
- The models named `Wan2.1-Fun-*-Control` are trained in the `control_ref` mode.
|
||||
- The models named `Wan2.1-Fun-*-Control-Camera` are trained in the `control_ref_camera` mode.
|
||||
- `control_ref_image` is used to specify the type of control image. The available options are `first_frame` and `random`.
|
||||
- `first_frame` is used in V1.0 because V1.0 supports using a specified start frame as the control image. The Control-Camera models use the first frame as the control image.
|
||||
- `random` is used in V1.1 because V1.1 supports both using a specified start frame and a reference image as the control image.
|
||||
- `add_full_ref_image_in_self_attention` determines whether to include the reference image in self-attention. This option is used in V1.1, as it supports using a reference image as the control image. It should not be used in V1.0 and Control-Camera models.
|
||||
|
||||
Wan-Fun-Control without deepspeed:
|
||||
When train model with multi machines, please set the params as follows:
|
||||
```sh
|
||||
export MASTER_ADDR="your master address"
|
||||
export MASTER_PORT=10086
|
||||
export WORLD_SIZE=1 # The number of machines
|
||||
export NUM_PROCESS=8 # The number of processes, such as WORLD_SIZE * 8
|
||||
export RANK=0 # The rank of this machine
|
||||
|
||||
accelerate launch --mixed_precision="bf16" --main_process_ip=$MASTER_ADDR --main_process_port=$MASTER_PORT --num_machines=$WORLD_SIZE --num_processes=$NUM_PROCESS --machine_rank=$RANK scripts/wan2.1_fun/xxx.py
|
||||
```
|
||||
|
||||
Wan-Fun-Control-V1.1 without deepspeed:
|
||||
```sh
|
||||
export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-Fun-V1.1-14B-Control"
|
||||
export DATASET_NAME="datasets/internal_datasets/"
|
||||
export DATASET_META_NAME="datasets/internal_datasets/metadata.json"
|
||||
export NCCL_IB_DISABLE=1
|
||||
export NCCL_P2P_DISABLE=1
|
||||
NCCL_DEBUG=INFO
|
||||
|
||||
accelerate launch --mixed_precision="bf16" scripts/wan2.1_fun/train_control.py \
|
||||
--config_path="config/wan2.1/wan_civitai.yaml" \
|
||||
--pretrained_model_name_or_path=$MODEL_NAME \
|
||||
--train_data_dir=$DATASET_NAME \
|
||||
--train_data_meta=$DATASET_META_NAME \
|
||||
--image_sample_size=1024 \
|
||||
--video_sample_size=256 \
|
||||
--token_sample_size=512 \
|
||||
--video_sample_stride=2 \
|
||||
--video_sample_n_frames=81 \
|
||||
--train_batch_size=1 \
|
||||
--video_repeat=1 \
|
||||
--gradient_accumulation_steps=1 \
|
||||
--dataloader_num_workers=8 \
|
||||
--num_train_epochs=100 \
|
||||
--checkpointing_steps=50 \
|
||||
--learning_rate=2e-05 \
|
||||
--lr_scheduler="constant_with_warmup" \
|
||||
--lr_warmup_steps=100 \
|
||||
--seed=42 \
|
||||
--output_dir="output_dir" \
|
||||
--gradient_checkpointing \
|
||||
--mixed_precision="bf16" \
|
||||
--adam_weight_decay=3e-2 \
|
||||
--adam_epsilon=1e-10 \
|
||||
--vae_mini_batch=1 \
|
||||
--max_grad_norm=0.05 \
|
||||
--random_hw_adapt \
|
||||
--training_with_video_token_length \
|
||||
--enable_bucket \
|
||||
--uniform_sampling \
|
||||
--low_vram \
|
||||
--train_mode="control_ref" \
|
||||
--control_ref_image="random" \
|
||||
--add_full_ref_image_in_self_attention \
|
||||
--trainable_modules "."
|
||||
```
|
||||
|
||||
Wan-Fun-Control-V1.1 with deepspeed:
|
||||
```sh
|
||||
export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-Fun-V1.1-14B-Control"
|
||||
export DATASET_NAME="datasets/internal_datasets/"
|
||||
export DATASET_META_NAME="datasets/internal_datasets/metadata.json"
|
||||
export NCCL_IB_DISABLE=1
|
||||
export NCCL_P2P_DISABLE=1
|
||||
NCCL_DEBUG=INFO
|
||||
|
||||
accelerate launch --use_deepspeed --deepspeed_config_file config/zero_stage2_config.json --deepspeed_multinode_launcher standard scripts/wan2.1_fun/train_control.py \
|
||||
--config_path="config/wan2.1/wan_civitai.yaml" \
|
||||
--pretrained_model_name_or_path=$MODEL_NAME \
|
||||
--train_data_dir=$DATASET_NAME \
|
||||
--train_data_meta=$DATASET_META_NAME \
|
||||
--image_sample_size=1024 \
|
||||
--video_sample_size=256 \
|
||||
--token_sample_size=512 \
|
||||
--video_sample_stride=2 \
|
||||
--video_sample_n_frames=81 \
|
||||
--train_batch_size=1 \
|
||||
--video_repeat=1 \
|
||||
--gradient_accumulation_steps=1 \
|
||||
--dataloader_num_workers=8 \
|
||||
--num_train_epochs=100 \
|
||||
--checkpointing_steps=50 \
|
||||
--learning_rate=2e-05 \
|
||||
--lr_scheduler="constant_with_warmup" \
|
||||
--lr_warmup_steps=100 \
|
||||
--seed=42 \
|
||||
--output_dir="output_dir" \
|
||||
--gradient_checkpointing \
|
||||
--mixed_precision="bf16" \
|
||||
--adam_weight_decay=3e-2 \
|
||||
--adam_epsilon=1e-10 \
|
||||
--vae_mini_batch=1 \
|
||||
--max_grad_norm=0.05 \
|
||||
--random_hw_adapt \
|
||||
--training_with_video_token_length \
|
||||
--enable_bucket \
|
||||
--uniform_sampling \
|
||||
--low_vram \
|
||||
--use_deepspeed \
|
||||
--train_mode="control_ref" \
|
||||
--control_ref_image="random" \
|
||||
--add_full_ref_image_in_self_attention \
|
||||
--trainable_modules "."
|
||||
```
|
||||
|
||||
Wan-Fun-Control-V1.1 with deepspeed zero-3:
|
||||
|
||||
Wan with DeepSpeed Zero-3 is suitable for 14B Wan at high resolutions. After training, you can use the following command to get the final model:
|
||||
```sh
|
||||
python scripts/zero_to_bf16.py output_dir/checkpoint-{our-num-steps} output_dir/checkpoint-{your-num-steps}-outputs --max_shard_size 80GB --safe_serialization
|
||||
```
|
||||
|
||||
Training shell command is as follows:
|
||||
```sh
|
||||
export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-Fun-V1.1-14B-Control"
|
||||
export DATASET_NAME="datasets/internal_datasets/"
|
||||
export DATASET_META_NAME="datasets/internal_datasets/metadata.json"
|
||||
export NCCL_IB_DISABLE=1
|
||||
export NCCL_P2P_DISABLE=1
|
||||
NCCL_DEBUG=INFO
|
||||
|
||||
accelerate launch --zero_stage 3 --zero3_save_16bit_model true --zero3_init_flag true --use_deepspeed --deepspeed_config_file config/zero_stage2.1_config.json --deepspeed_multinode_launcherr standard scripts/wan2.1_fun/train_control.py \
|
||||
--config_path="config/wan2.1/wan_civitai.yaml" \
|
||||
--pretrained_model_name_or_path=$MODEL_NAME \
|
||||
--train_data_dir=$DATASET_NAME \
|
||||
--train_data_meta=$DATASET_META_NAME \
|
||||
--image_sample_size=1024 \
|
||||
--video_sample_size=256 \
|
||||
--token_sample_size=512 \
|
||||
--video_sample_stride=2 \
|
||||
--video_sample_n_frames=81 \
|
||||
--train_batch_size=1 \
|
||||
--video_repeat=1 \
|
||||
--gradient_accumulation_steps=1 \
|
||||
--dataloader_num_workers=8 \
|
||||
--num_train_epochs=100 \
|
||||
--checkpointing_steps=50 \
|
||||
--learning_rate=2e-05 \
|
||||
--lr_scheduler="constant_with_warmup" \
|
||||
--lr_warmup_steps=100 \
|
||||
--seed=42 \
|
||||
--output_dir="output_dir" \
|
||||
--gradient_checkpointing \
|
||||
--mixed_precision="bf16" \
|
||||
--adam_weight_decay=3e-2 \
|
||||
--adam_epsilon=1e-10 \
|
||||
--vae_mini_batch=1 \
|
||||
--max_grad_norm=0.05 \
|
||||
--random_hw_adapt \
|
||||
--training_with_video_token_length \
|
||||
--enable_bucket \
|
||||
--uniform_sampling \
|
||||
--low_vram \
|
||||
--use_deepspeed \
|
||||
--train_mode="control_ref" \
|
||||
--control_ref_image="random" \
|
||||
--add_full_ref_image_in_self_attention \
|
||||
--trainable_modules "."
|
||||
```
|
||||
|
||||
<details>
|
||||
<summary>(Obsolete) V1.0:</summary>
|
||||
|
||||
Wan-Fun-Control-V1.0 without deepspeed:
|
||||
```sh
|
||||
export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-Fun-14B-Control"
|
||||
export DATASET_NAME="datasets/internal_datasets/"
|
||||
@@ -48,7 +219,6 @@ export NCCL_IB_DISABLE=1
|
||||
export NCCL_P2P_DISABLE=1
|
||||
NCCL_DEBUG=INFO
|
||||
|
||||
# When train model with multi machines, use "--config_file accelerate.yaml" instead of "--mixed_precision='bf16'".
|
||||
accelerate launch --mixed_precision="bf16" scripts/wan2.1_fun/train_control.py \
|
||||
--config_path="config/wan2.1/wan_civitai.yaml" \
|
||||
--pretrained_model_name_or_path=$MODEL_NAME \
|
||||
@@ -86,7 +256,7 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.1_fun/train_control.py \
|
||||
--trainable_modules "."
|
||||
```
|
||||
|
||||
Wan-Fun with deepspeed:
|
||||
Wan-Fun-Control-V1.0 with deepspeed:
|
||||
```sh
|
||||
export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-Fun-14B-Control"
|
||||
export DATASET_NAME="datasets/internal_datasets/"
|
||||
@@ -95,7 +265,6 @@ export NCCL_IB_DISABLE=1
|
||||
export NCCL_P2P_DISABLE=1
|
||||
NCCL_DEBUG=INFO
|
||||
|
||||
# When train model with multi machines, use "--config_file accelerate.yaml" instead of "--mixed_precision='bf16'".
|
||||
accelerate launch --use_deepspeed --deepspeed_config_file config/zero_stage2_config.json --deepspeed_multinode_launcher standard scripts/wan2.1_fun/train_control.py \
|
||||
--config_path="config/wan2.1/wan_civitai.yaml" \
|
||||
--pretrained_model_name_or_path=$MODEL_NAME \
|
||||
@@ -134,7 +303,7 @@ accelerate launch --use_deepspeed --deepspeed_config_file config/zero_stage2_con
|
||||
--trainable_modules "."
|
||||
```
|
||||
|
||||
Wan T2V with deepspeed zero-3:
|
||||
Wan-Fun-Control-V1.0 with deepspeed zero-3:
|
||||
|
||||
Wan with DeepSpeed Zero-3 is suitable for 14B Wan at high resolutions. After training, you can use the following command to get the final model:
|
||||
```sh
|
||||
@@ -186,4 +355,5 @@ accelerate launch --zero_stage 3 --zero3_save_16bit_model true --zero3_init_flag
|
||||
--train_mode="control_ref" \
|
||||
--control_ref_image="first_frame" \
|
||||
--trainable_modules "."
|
||||
```
|
||||
```
|
||||
</details>
|
||||
|
||||
Regular → Executable
+169
-6
@@ -37,9 +37,172 @@ Some parameters in the sh file can be confusing, and they are explained in this
|
||||
- At 768x768 resolution, the number of video frames is 21 (~= 512 * 512 * 49 / 768 / 768).
|
||||
- At 1024x1024 resolution, the number of video frames is 9 (~= 512 * 512 * 49 / 1024 / 1024).
|
||||
- These resolutions combined with their corresponding lengths allow the model to generate videos of different sizes.
|
||||
- `resume_from_checkpoint` is used to set the training should be resumed from a previous checkpoint. Use a path or `"latest"` to automatically select the last available checkpoint.
|
||||
- `resume_from_checkpoint` is used to set the training should be resumed from a previous checkpoint. Use a path or `"latest"` to automatically select the last available checkpoint and set the `save_state` to `True`.
|
||||
- `train_mode` is used to set the training mode.
|
||||
- The models named `Wan2.1-Fun-*-Control` are trained in the `control_ref` mode.
|
||||
- The models named `Wan2.1-Fun-*-Control-Camera` are trained in the `control_ref_camera` mode.
|
||||
- `control_ref_image` is used to specify the type of control image. The available options are `first_frame` and `random`.
|
||||
- `first_frame` is used in V1.0 because V1.0 supports using a specified start frame as the control image. The Control-Camera models use the first frame as the control image.
|
||||
- `random` is used in V1.1 because V1.1 supports both using a specified start frame and a reference image as the control image.
|
||||
- `add_full_ref_image_in_self_attention` determines whether to include the reference image in self-attention. This option is used in V1.1, as it supports using a reference image as the control image. It should not be used in V1.0 and Control-Camera models.
|
||||
|
||||
Wan-Fun-Control without deepspeed:
|
||||
When train model with multi machines, please set the params as follows:
|
||||
```sh
|
||||
export MASTER_ADDR="your master address"
|
||||
export MASTER_PORT=10086
|
||||
export WORLD_SIZE=1 # The number of machines
|
||||
export NUM_PROCESS=8 # The number of processes, such as WORLD_SIZE * 8
|
||||
export RANK=0 # The rank of this machine
|
||||
|
||||
accelerate launch --mixed_precision="bf16" --main_process_ip=$MASTER_ADDR --main_process_port=$MASTER_PORT --num_machines=$WORLD_SIZE --num_processes=$NUM_PROCESS --machine_rank=$RANK scripts/wan2.1_fun/xxx.py
|
||||
```
|
||||
|
||||
Wan-Fun-Control-V1.1 without deepspeed:
|
||||
```sh
|
||||
export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-Fun-V1.1-14B-Control"
|
||||
export DATASET_NAME="datasets/internal_datasets/"
|
||||
export DATASET_META_NAME="datasets/internal_datasets/metadata.json"
|
||||
export NCCL_IB_DISABLE=1
|
||||
export NCCL_P2P_DISABLE=1
|
||||
NCCL_DEBUG=INFO
|
||||
|
||||
accelerate launch --mixed_precision="bf16" scripts/wan2.1_fun/train_control_lora.py \
|
||||
--config_path="config/wan2.1/wan_civitai.yaml" \
|
||||
--pretrained_model_name_or_path=$MODEL_NAME \
|
||||
--train_data_dir=$DATASET_NAME \
|
||||
--train_data_meta=$DATASET_META_NAME \
|
||||
--image_sample_size=1024 \
|
||||
--video_sample_size=256 \
|
||||
--token_sample_size=512 \
|
||||
--video_sample_stride=2 \
|
||||
--video_sample_n_frames=81 \
|
||||
--train_batch_size=1 \
|
||||
--video_repeat=1 \
|
||||
--gradient_accumulation_steps=1 \
|
||||
--dataloader_num_workers=8 \
|
||||
--num_train_epochs=100 \
|
||||
--checkpointing_steps=50 \
|
||||
--learning_rate=1e-04 \
|
||||
--seed=42 \
|
||||
--output_dir="output_dir" \
|
||||
--gradient_checkpointing \
|
||||
--mixed_precision="bf16" \
|
||||
--adam_weight_decay=3e-2 \
|
||||
--adam_epsilon=1e-10 \
|
||||
--vae_mini_batch=1 \
|
||||
--max_grad_norm=0.05 \
|
||||
--random_hw_adapt \
|
||||
--training_with_video_token_length \
|
||||
--enable_bucket \
|
||||
--uniform_sampling \
|
||||
--train_mode="control_ref" \
|
||||
--control_ref_image="random" \
|
||||
--add_full_ref_image_in_self_attention \
|
||||
--low_vram
|
||||
```
|
||||
|
||||
Wan-Fun-Control-V1.1 with deepspeed:
|
||||
```sh
|
||||
export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-Fun-V1.1-14B-Control"
|
||||
export DATASET_NAME="datasets/internal_datasets/"
|
||||
export DATASET_META_NAME="datasets/internal_datasets/metadata.json"
|
||||
export NCCL_IB_DISABLE=1
|
||||
export NCCL_P2P_DISABLE=1
|
||||
NCCL_DEBUG=INFO
|
||||
|
||||
accelerate launch --use_deepspeed --deepspeed_config_file config/zero_stage2_config.json --deepspeed_multinode_launcher standard scripts/wan2.1_fun/train_control_lora.py \
|
||||
--config_path="config/wan2.1/wan_civitai.yaml" \
|
||||
--pretrained_model_name_or_path=$MODEL_NAME \
|
||||
--train_data_dir=$DATASET_NAME \
|
||||
--train_data_meta=$DATASET_META_NAME \
|
||||
--image_sample_size=1024 \
|
||||
--video_sample_size=256 \
|
||||
--token_sample_size=512 \
|
||||
--video_sample_stride=2 \
|
||||
--video_sample_n_frames=81 \
|
||||
--train_batch_size=1 \
|
||||
--video_repeat=1 \
|
||||
--gradient_accumulation_steps=1 \
|
||||
--dataloader_num_workers=8 \
|
||||
--num_train_epochs=100 \
|
||||
--checkpointing_steps=50 \
|
||||
--learning_rate=1e-04 \
|
||||
--seed=42 \
|
||||
--output_dir="output_dir" \
|
||||
--gradient_checkpointing \
|
||||
--mixed_precision="bf16" \
|
||||
--adam_weight_decay=3e-2 \
|
||||
--adam_epsilon=1e-10 \
|
||||
--vae_mini_batch=1 \
|
||||
--max_grad_norm=0.05 \
|
||||
--random_hw_adapt \
|
||||
--training_with_video_token_length \
|
||||
--enable_bucket \
|
||||
--uniform_sampling \
|
||||
--use_deepspeed \
|
||||
--train_mode="control_ref" \
|
||||
--control_ref_image="random" \
|
||||
--add_full_ref_image_in_self_attention \
|
||||
--low_vram
|
||||
```
|
||||
|
||||
Wan-Fun-Control-V1.1 with deepspeed zero-3:
|
||||
|
||||
Wan with DeepSpeed Zero-3 is suitable for 14B Wan at high resolutions. You must set save_state to True to save the model. After training, you can use the following command to get the final model:
|
||||
```sh
|
||||
python scripts/zero_to_bf16.py output_dir/checkpoint-{our-num-steps} output_dir/checkpoint-{your-num-steps}-outputs --max_shard_size 80GB --safe_serialization
|
||||
```
|
||||
|
||||
Training shell command is as follows:
|
||||
```sh
|
||||
export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-Fun-V1.1-14B-Control"
|
||||
export DATASET_NAME="datasets/internal_datasets/"
|
||||
export DATASET_META_NAME="datasets/internal_datasets/metadata.json"
|
||||
export NCCL_IB_DISABLE=1
|
||||
export NCCL_P2P_DISABLE=1
|
||||
NCCL_DEBUG=INFO
|
||||
|
||||
accelerate launch --zero_stage 3 --zero3_save_16bit_model true --zero3_init_flag true --use_deepspeed --deepspeed_config_file config/zero_stage2.1_config.json --deepspeed_multinode_launcherr standard scripts/wan2.1_fun/train_control.py \
|
||||
--config_path="config/wan2.1/wan_civitai.yaml" \
|
||||
--pretrained_model_name_or_path=$MODEL_NAME \
|
||||
--train_data_dir=$DATASET_NAME \
|
||||
--train_data_meta=$DATASET_META_NAME \
|
||||
--image_sample_size=1024 \
|
||||
--video_sample_size=256 \
|
||||
--token_sample_size=512 \
|
||||
--video_sample_stride=2 \
|
||||
--video_sample_n_frames=81 \
|
||||
--train_batch_size=1 \
|
||||
--video_repeat=1 \
|
||||
--gradient_accumulation_steps=1 \
|
||||
--dataloader_num_workers=8 \
|
||||
--num_train_epochs=100 \
|
||||
--checkpointing_steps=50 \
|
||||
--learning_rate=1e-04 \
|
||||
--seed=42 \
|
||||
--output_dir="output_dir" \
|
||||
--gradient_checkpointing \
|
||||
--mixed_precision="bf16" \
|
||||
--adam_weight_decay=3e-2 \
|
||||
--adam_epsilon=1e-10 \
|
||||
--vae_mini_batch=1 \
|
||||
--max_grad_norm=0.05 \
|
||||
--random_hw_adapt \
|
||||
--training_with_video_token_length \
|
||||
--enable_bucket \
|
||||
--uniform_sampling \
|
||||
--save_state \
|
||||
--use_deepspeed \
|
||||
--train_mode="control_ref" \
|
||||
--control_ref_image="random" \
|
||||
--add_full_ref_image_in_self_attention \
|
||||
--low_vram
|
||||
```
|
||||
|
||||
<details>
|
||||
<summary>(Obsolete) V1.0:</summary>
|
||||
|
||||
Wan-Fun-Control-V1.0 without deepspeed:
|
||||
```sh
|
||||
export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-Fun-14B-Control"
|
||||
export DATASET_NAME="datasets/internal_datasets/"
|
||||
@@ -82,7 +245,7 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.1_fun/train_control_lora
|
||||
--low_vram
|
||||
```
|
||||
|
||||
Wan-Fun with deepspeed:
|
||||
Wan-Fun-Control-V1.0 with deepspeed:
|
||||
```sh
|
||||
export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-Fun-14B-Control"
|
||||
export DATASET_NAME="datasets/internal_datasets/"
|
||||
@@ -91,7 +254,6 @@ export NCCL_IB_DISABLE=1
|
||||
export NCCL_P2P_DISABLE=1
|
||||
NCCL_DEBUG=INFO
|
||||
|
||||
# When train model with multi machines, use "--config_file accelerate.yaml" instead of "--mixed_precision='bf16'".
|
||||
accelerate launch --use_deepspeed --deepspeed_config_file config/zero_stage2_config.json --deepspeed_multinode_launcher standard scripts/wan2.1_fun/train_control_lora.py \
|
||||
--config_path="config/wan2.1/wan_civitai.yaml" \
|
||||
--pretrained_model_name_or_path=$MODEL_NAME \
|
||||
@@ -127,7 +289,7 @@ accelerate launch --use_deepspeed --deepspeed_config_file config/zero_stage2_con
|
||||
--low_vram
|
||||
```
|
||||
|
||||
Wan T2V with deepspeed zero-3:
|
||||
Wan-Fun-Control-V1.0 with deepspeed zero-3:
|
||||
|
||||
Wan with DeepSpeed Zero-3 is suitable for 14B Wan at high resolutions. You must set save_state to True to save the model. After training, you can use the following command to get the final model:
|
||||
```sh
|
||||
@@ -177,4 +339,5 @@ accelerate launch --zero_stage 3 --zero3_save_16bit_model true --zero3_init_flag
|
||||
--train_mode="control_ref" \
|
||||
--control_ref_image="first_frame" \
|
||||
--low_vram
|
||||
```
|
||||
```
|
||||
</details>
|
||||
@@ -24,7 +24,7 @@ Wan T2V without deepspeed:
|
||||
|
||||
Wan without DeepSpeed is more suitable for 1.3B Wan, as using it with 14B Wan may result in insufficient GPU memory.
|
||||
```sh
|
||||
export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-Fun-14B-InP"
|
||||
export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-Fun-V1.1-14B-InP"
|
||||
export DATASET_NAME="datasets/internal_datasets/"
|
||||
export DATASET_META_NAME="datasets/internal_datasets/metadata.json"
|
||||
export NCCL_IB_DISABLE=1
|
||||
@@ -69,7 +69,7 @@ Wan T2V with deepspeed zero-2:
|
||||
Wan with DeepSpeed Zero-2 is suitable for training 1.3B Wan and 14B Wan at low resolutions, but training 14B Wan at high resolutions may still result in insufficient GPU memory.
|
||||
|
||||
```sh
|
||||
export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-Fun-14B-InP"
|
||||
export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-Fun-V1.1-14B-InP"
|
||||
export DATASET_NAME="datasets/internal_datasets/"
|
||||
export DATASET_META_NAME="datasets/internal_datasets/metadata.json"
|
||||
export NCCL_IB_DISABLE=1
|
||||
@@ -119,7 +119,7 @@ python scripts/zero_to_bf16.py output_dir/checkpoint-{our-num-steps} output_dir/
|
||||
|
||||
Training shell command is as follows:
|
||||
```sh
|
||||
export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-Fun-14B-InP"
|
||||
export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-Fun-V1.1-14B-InP"
|
||||
export DATASET_NAME="datasets/internal_datasets/"
|
||||
export DATASET_META_NAME="datasets/internal_datasets/metadata.json"
|
||||
export NCCL_IB_DISABLE=1
|
||||
|
||||
Regular → Executable
+2
-2
@@ -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
|
||||
|
||||
Regular → Executable
+4
-4
@@ -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 "."
|
||||
Regular → Executable
+3
-2
@@ -1,4 +1,4 @@
|
||||
export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-Fun-14B-Control"
|
||||
export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-Fun-V1.1-14B-Control"
|
||||
export DATASET_NAME="datasets/internal_datasets/"
|
||||
export DATASET_META_NAME="datasets/internal_datasets/metadata.json"
|
||||
export NCCL_IB_DISABLE=1
|
||||
@@ -35,5 +35,6 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.1_fun/train_control_lora
|
||||
--enable_bucket \
|
||||
--uniform_sampling \
|
||||
--train_mode="control_ref" \
|
||||
--control_ref_image="first_frame" \
|
||||
--control_ref_image="random" \
|
||||
--add_full_ref_image_in_self_attention \
|
||||
--low_vram
|
||||
Regular → Executable
+2
-2
@@ -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
|
||||
|
||||
Regular → Executable
+1
-1
@@ -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"
|
||||
|
||||
@@ -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)
|
||||
]
|
||||
|
||||
Vendored
+2
@@ -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):
|
||||
|
||||
@@ -1,15 +1,16 @@
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
import importlib.util
|
||||
|
||||
def get_teacache_coefficients(model_name):
|
||||
if "wan2.1-t2v-1.3b" in model_name.lower() or "wan2.1-fun-1.3b" in model_name.lower() or "wan2.1-fun-v1.1-1.3b" in model_name.lower():
|
||||
return [-5.21862437e+04, 9.23041404e+03, -5.28275948e+02, 1.36987616e+01, -4.99875664e-02]
|
||||
elif "wan2.1-t2v-14b" in model_name.lower() or "wan2.1-fun-v1.1-14b" in model_name.lower():
|
||||
elif "wan2.1-t2v-14b" in model_name.lower():
|
||||
return [-3.03318725e+05, 4.90537029e+04, -2.65530556e+03, 5.87365115e+01, -3.15583525e-01]
|
||||
elif "wan2.1-i2v-14b-480p" in model_name.lower():
|
||||
return [2.57151496e+05, -3.54229917e+04, 1.40286849e+03, -1.35890334e+01, 1.32517977e-01]
|
||||
elif "wan2.1-i2v-14b-720p" in model_name.lower() or "wan2.1-fun-14b" in model_name.lower():
|
||||
elif "wan2.1-i2v-14b-720p" in model_name.lower() or "wan2.1-fun-14b" in model_name.lower() or "wan2.1-fun-v1.1-14b" in model_name.lower():
|
||||
return [8.10705460e+03, 2.13393892e+03, -3.72934672e+02, 1.66203073e+01, -4.17769401e-02]
|
||||
else:
|
||||
print(f"The model {model_name} is not supported by TeaCache.")
|
||||
@@ -71,4 +72,63 @@ class TeaCache():
|
||||
self.previous_modulated_input = None
|
||||
self.previous_residual = None
|
||||
self.previous_residual_cond = None
|
||||
self.previous_residual_uncond = None
|
||||
self.previous_residual_uncond = None
|
||||
|
||||
|
||||
if importlib.util.find_spec("pai_fuser") is not None:
|
||||
from pai_fuser.core import (cfg_skip_turbo, enable_cfg_skip,
|
||||
disable_cfg_skip)
|
||||
cfg_skip = cfg_skip_turbo
|
||||
print("Enable CFG Skip Turbo")
|
||||
else:
|
||||
def cfg_skip():
|
||||
def decorator(func):
|
||||
def wrapper(self, x, *args, **kwargs):
|
||||
if self.cfg_skip_ratio is not None and self.current_steps >= self.num_inference_steps * (1 - self.cfg_skip_ratio):
|
||||
bs = len(x)
|
||||
bs_half = int(bs // 2)
|
||||
|
||||
new_x = x[bs_half:]
|
||||
|
||||
new_args = []
|
||||
for arg in args:
|
||||
if isinstance(arg, (torch.Tensor, list, tuple, np.ndarray)):
|
||||
new_args.append(arg[bs_half:])
|
||||
else:
|
||||
new_args.append(arg)
|
||||
|
||||
new_kwargs = {}
|
||||
for key, content in kwargs.items():
|
||||
if isinstance(content, (torch.Tensor, list, tuple, np.ndarray)):
|
||||
new_kwargs[key] = content[bs_half:]
|
||||
else:
|
||||
new_kwargs[key] = content
|
||||
else:
|
||||
new_x = x
|
||||
new_args = args
|
||||
new_kwargs = kwargs
|
||||
|
||||
result = func(self, new_x, *new_args, **new_kwargs)
|
||||
|
||||
if self.cfg_skip_ratio is not None and self.current_steps >= self.num_inference_steps * (1 - self.cfg_skip_ratio):
|
||||
result = torch.cat([result, result], dim=0)
|
||||
|
||||
return result
|
||||
return wrapper
|
||||
return decorator
|
||||
|
||||
def enable_cfg_skip():
|
||||
def decorator(func):
|
||||
def wrapper(self, cfg_skip_ratio, num_steps, *args, **kwargs):
|
||||
func(self, cfg_skip_ratio, num_steps, *args, **kwargs)
|
||||
return
|
||||
return wrapper
|
||||
return decorator
|
||||
|
||||
def disable_cfg_skip():
|
||||
def decorator(func):
|
||||
def wrapper(self, *args, **kwargs):
|
||||
func(self, *args, **kwargs)
|
||||
return
|
||||
return wrapper
|
||||
return decorator
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -408,7 +408,6 @@ class WanFunPipeline(DiffusionPipeline):
|
||||
callback_on_step_end_tensor_inputs: List[str] = ["latents"],
|
||||
max_sequence_length: int = 512,
|
||||
comfyui_progressbar: bool = False,
|
||||
cfg_skip_ratio: int = None,
|
||||
shift: int = 5,
|
||||
) -> Union[WanPipelineOutput, Tuple]:
|
||||
"""
|
||||
@@ -513,11 +512,10 @@ class WanFunPipeline(DiffusionPipeline):
|
||||
seq_len = math.ceil((target_shape[2] * target_shape[3]) / (self.transformer.config.patch_size[1] * self.transformer.config.patch_size[2]) * target_shape[1])
|
||||
# 7. Denoising loop
|
||||
num_warmup_steps = max(len(timesteps) - num_inference_steps * self.scheduler.order, 0)
|
||||
self.transformer.num_inference_steps = num_inference_steps
|
||||
with self.progress_bar(total=num_inference_steps) as progress_bar:
|
||||
for i, t in enumerate(timesteps):
|
||||
if cfg_skip_ratio is not None and i >= num_inference_steps * (1 - cfg_skip_ratio):
|
||||
do_classifier_free_guidance = False
|
||||
in_prompt_embeds = prompt_embeds
|
||||
self.transformer.current_steps = i
|
||||
|
||||
if self.interrupt:
|
||||
continue
|
||||
|
||||
@@ -497,7 +497,6 @@ class WanFunControlPipeline(DiffusionPipeline):
|
||||
clip_image: Image = None,
|
||||
max_sequence_length: int = 512,
|
||||
comfyui_progressbar: bool = False,
|
||||
cfg_skip_ratio: int = None,
|
||||
shift: int = 5,
|
||||
) -> Union[WanPipelineOutput, Tuple]:
|
||||
"""
|
||||
@@ -703,11 +702,10 @@ class WanFunControlPipeline(DiffusionPipeline):
|
||||
seq_len = math.ceil((target_shape[2] * target_shape[3]) / (self.transformer.config.patch_size[1] * self.transformer.config.patch_size[2]) * target_shape[1])
|
||||
# 7. Denoising loop
|
||||
num_warmup_steps = max(len(timesteps) - num_inference_steps * self.scheduler.order, 0)
|
||||
self.transformer.num_inference_steps = num_inference_steps
|
||||
with self.progress_bar(total=num_inference_steps) as progress_bar:
|
||||
for i, t in enumerate(timesteps):
|
||||
if cfg_skip_ratio is not None and i >= num_inference_steps * (1 - cfg_skip_ratio):
|
||||
do_classifier_free_guidance = False
|
||||
in_prompt_embeds = prompt_embeds
|
||||
self.transformer.current_steps = i
|
||||
|
||||
if self.interrupt:
|
||||
continue
|
||||
|
||||
@@ -496,7 +496,6 @@ class WanFunInpaintPipeline(DiffusionPipeline):
|
||||
clip_image: Image = None,
|
||||
max_sequence_length: int = 512,
|
||||
comfyui_progressbar: bool = False,
|
||||
cfg_skip_ratio: int = None,
|
||||
shift: int = 5,
|
||||
) -> Union[WanPipelineOutput, Tuple]:
|
||||
"""
|
||||
@@ -658,11 +657,10 @@ class WanFunInpaintPipeline(DiffusionPipeline):
|
||||
seq_len = math.ceil((target_shape[2] * target_shape[3]) / (self.transformer.config.patch_size[1] * self.transformer.config.patch_size[2]) * target_shape[1])
|
||||
# 7. Denoising loop
|
||||
num_warmup_steps = max(len(timesteps) - num_inference_steps * self.scheduler.order, 0)
|
||||
self.transformer.num_inference_steps = num_inference_steps
|
||||
with self.progress_bar(total=num_inference_steps) as progress_bar:
|
||||
for i, t in enumerate(timesteps):
|
||||
if cfg_skip_ratio is not None and i >= num_inference_steps * (1 - cfg_skip_ratio):
|
||||
do_classifier_free_guidance = False
|
||||
in_prompt_embeds = prompt_embeds
|
||||
self.transformer.current_steps = i
|
||||
|
||||
if self.interrupt:
|
||||
continue
|
||||
|
||||
@@ -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
|
||||
)
|
||||
|
||||
|
||||
@@ -59,7 +59,9 @@ all_cheduler_dict = {**ddpm_scheduler_dict, **flow_scheduler_dict}
|
||||
class Fun_Controller:
|
||||
def __init__(
|
||||
self, GPU_memory_mode, scheduler_dict, model_name=None, model_type="Inpaint",
|
||||
config_path=None, ulysses_degree=1, ring_degree=1, weight_dtype=None, savedir_sample=None,
|
||||
config_path=None, ulysses_degree=1, ring_degree=1,
|
||||
fsdp_dit=False, fsdp_text_encoder=False, compile_dit=False,
|
||||
weight_dtype=None, savedir_sample=None,
|
||||
):
|
||||
# config dirs
|
||||
self.basedir = os.getcwd()
|
||||
@@ -82,6 +84,9 @@ class Fun_Controller:
|
||||
self.config = OmegaConf.load(config_path)
|
||||
self.ulysses_degree = ulysses_degree
|
||||
self.ring_degree = ring_degree
|
||||
self.fsdp_dit = fsdp_dit
|
||||
self.fsdp_text_encoder = fsdp_text_encoder
|
||||
self.compile_dit = compile_dit
|
||||
self.weight_dtype = weight_dtype
|
||||
self.device = set_multi_gpus_devices(self.ulysses_degree, self.ring_degree)
|
||||
|
||||
@@ -148,6 +153,14 @@ class Fun_Controller:
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.ipc_collect()
|
||||
|
||||
def auto_model_clear_cache(self, model):
|
||||
origin_device = model.device
|
||||
model = model.to("cpu")
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.ipc_collect()
|
||||
model = model.to(origin_device)
|
||||
|
||||
def input_check(self,
|
||||
resize_method,
|
||||
|
||||
+31
-11
@@ -37,6 +37,7 @@ from .ui import (create_cfg_and_seedbox, create_cfg_riflex_k,
|
||||
create_height_width, create_model_checkpoints,
|
||||
create_model_type, create_prompts, create_samplers,
|
||||
create_teacache_params, create_ui_outputs)
|
||||
from ..dist import set_multi_gpus_devices, shard_model
|
||||
|
||||
|
||||
class Wan_Fun_Controller(Fun_Controller):
|
||||
@@ -117,22 +118,36 @@ class Wan_Fun_Controller(Fun_Controller):
|
||||
)
|
||||
|
||||
if self.ulysses_degree > 1 or self.ring_degree > 1:
|
||||
from functools import partial
|
||||
self.transformer.enable_multi_gpus_inference()
|
||||
if self.fsdp_dit:
|
||||
shard_fn = partial(shard_model, device_id=self.device, param_dtype=self.weight_dtype)
|
||||
self.pipeline.transformer = shard_fn(self.pipeline.transformer)
|
||||
print("Add FSDP DIT")
|
||||
if self.fsdp_text_encoder:
|
||||
shard_fn = partial(shard_model, device_id=self.device, param_dtype=self.weight_dtype)
|
||||
self.pipeline.text_encoder = shard_fn(self.pipeline.text_encoder)
|
||||
print("Add FSDP TEXT ENCODER")
|
||||
|
||||
if self.compile_dit:
|
||||
for i in range(len(self.pipeline.transformer.blocks)):
|
||||
self.pipeline.transformer.blocks[i] = torch.compile(self.pipeline.transformer.blocks[i])
|
||||
print("Add Compile")
|
||||
|
||||
if self.GPU_memory_mode == "sequential_cpu_offload":
|
||||
replace_parameters_by_name(self.transformer, ["modulation",], device=self.device)
|
||||
self.transformer.freqs = self.transformer.freqs.to(device=self.device)
|
||||
self.pipeline.enable_sequential_cpu_offload(device=self.device)
|
||||
elif self.GPU_memory_mode == "model_cpu_offload_and_qfloat8":
|
||||
convert_model_weight_to_float8(self.transformer, exclude_module_name=["modulation",])
|
||||
convert_model_weight_to_float8(self.transformer, exclude_module_name=["modulation",], device=self.device)
|
||||
convert_weight_dtype_wrapper(self.transformer, self.weight_dtype)
|
||||
self.pipeline.enable_model_cpu_offload(device=self.device)
|
||||
elif self.GPU_memory_mode == "model_cpu_offload":
|
||||
self.pipeline.enable_model_cpu_offload(device=self.device)
|
||||
elif self.GPU_memory_mode == "model_full_load_and_qfloat8":
|
||||
convert_model_weight_to_float8(self.transformer, exclude_module_name=["modulation",])
|
||||
convert_model_weight_to_float8(self.transformer, exclude_module_name=["modulation",], device=self.device)
|
||||
convert_weight_dtype_wrapper(self.transformer, self.weight_dtype)
|
||||
self.pipeline.to(device=self.device)
|
||||
self.pipeline.to(self.device)
|
||||
else:
|
||||
self.pipeline.to(self.device)
|
||||
print("Update diffusion transformer done")
|
||||
@@ -219,7 +234,11 @@ class Wan_Fun_Controller(Fun_Controller):
|
||||
else:
|
||||
print(f"Disable TeaCache.")
|
||||
self.pipeline.transformer.disable_teacache()
|
||||
|
||||
|
||||
if cfg_skip_ratio is not None and cfg_skip_ratio >= 0:
|
||||
print(f"Enable cfg_skip_ratio {cfg_skip_ratio}.")
|
||||
self.pipeline.transformer.enable_cfg_skip(cfg_skip_ratio, sample_step_slider)
|
||||
|
||||
print(f"Generate seed.")
|
||||
if int(seed_textbox) != -1 and seed_textbox != "": torch.manual_seed(int(seed_textbox))
|
||||
else: seed_textbox = np.random.randint(0, 1e10)
|
||||
@@ -253,7 +272,6 @@ class Wan_Fun_Controller(Fun_Controller):
|
||||
video = input_video,
|
||||
mask_video = input_video_mask,
|
||||
clip_image = clip_image,
|
||||
cfg_skip_ratio = cfg_skip_ratio,
|
||||
).videos
|
||||
else:
|
||||
sample = self.pipeline(
|
||||
@@ -265,7 +283,6 @@ class Wan_Fun_Controller(Fun_Controller):
|
||||
height = height_slider,
|
||||
num_frames = length_slider if not is_image else 1,
|
||||
generator = generator,
|
||||
cfg_skip_ratio = cfg_skip_ratio,
|
||||
).videos
|
||||
else:
|
||||
if ref_image is not None:
|
||||
@@ -297,11 +314,14 @@ class Wan_Fun_Controller(Fun_Controller):
|
||||
ref_image = ref_image,
|
||||
start_image = start_image,
|
||||
clip_image = clip_image,
|
||||
cfg_skip_ratio = cfg_skip_ratio,
|
||||
).videos
|
||||
print(f"Generation done.")
|
||||
except Exception as e:
|
||||
self.auto_model_clear_cache(self.pipeline.transformer)
|
||||
self.auto_model_clear_cache(self.pipeline.text_encoder)
|
||||
self.auto_model_clear_cache(self.pipeline.vae)
|
||||
self.clear_cache()
|
||||
|
||||
print(f"Error. error information is {str(e)}")
|
||||
if self.lora_model_path != "none":
|
||||
self.pipeline = unmerge_lora(self.pipeline, self.lora_model_path, multiplier=lora_alpha_slider)
|
||||
@@ -343,10 +363,10 @@ class Wan_Fun_Controller(Fun_Controller):
|
||||
Wan_Fun_Controller_Host = Wan_Fun_Controller
|
||||
Wan_Fun_Controller_Client = Fun_Controller_Client
|
||||
|
||||
def ui(GPU_memory_mode, scheduler_dict, config_path, ulysses_degree, ring_degree, weight_dtype, savedir_sample=None):
|
||||
def ui(GPU_memory_mode, scheduler_dict, config_path, compile_dit, weight_dtype, savedir_sample=None):
|
||||
controller = Wan_Fun_Controller(
|
||||
GPU_memory_mode, scheduler_dict, model_name=None, model_type="Inpaint",
|
||||
config_path=config_path, ulysses_degree=ulysses_degree, ring_degree=ring_degree,
|
||||
config_path=config_path, compile_dit=compile_dit,
|
||||
weight_dtype=weight_dtype, savedir_sample=savedir_sample,
|
||||
)
|
||||
|
||||
@@ -481,10 +501,10 @@ def ui(GPU_memory_mode, scheduler_dict, config_path, ulysses_degree, ring_degree
|
||||
)
|
||||
return demo, controller
|
||||
|
||||
def ui_host(GPU_memory_mode, scheduler_dict, model_name, model_type, config_path, ulysses_degree, ring_degree, weight_dtype, savedir_sample=None):
|
||||
def ui_host(GPU_memory_mode, scheduler_dict, model_name, model_type, config_path, compile_dit, weight_dtype, savedir_sample=None):
|
||||
controller = Wan_Fun_Controller_Host(
|
||||
GPU_memory_mode, scheduler_dict, model_name=model_name, model_type=model_type,
|
||||
config_path=config_path, ulysses_degree=ulysses_degree, ring_degree=ring_degree,
|
||||
config_path=config_path, compile_dit=compile_dit,
|
||||
weight_dtype=weight_dtype, savedir_sample=savedir_sample,
|
||||
)
|
||||
|
||||
|
||||
+31
-11
@@ -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,
|
||||
)
|
||||
|
||||
|
||||
Regular → Executable
+41
-34
@@ -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)
|
||||
)
|
||||
Reference in New Issue
Block a user