From 7157cdd48a0bf0c3baacd337cd711aff01675c22 Mon Sep 17 00:00:00 2001 From: Bubbliiiing <47347516+bubbliiiing@users.noreply.github.com> Date: Wed, 13 Aug 2025 19:20:11 +0800 Subject: [PATCH] Update 5b comfyui && Update Camera control && Update fun training and Readme (#284) Update 5b comfyui && Update Camera control && Update fun training and Readme --- README.md | 1 + README_ja-JP.md | 1 + README_zh-CN.md | 1 + comfyui/README.md | 1 + comfyui/comfyui_nodes.py | 7 +- .../wan2.1_fun_workflow_control_camera.json | 368 +-- ...an2.1_fun_workflow_control_trajectory.json | 528 ++--- .../v1/wan2.1_fun_workflow_v2v_control.json | 2 +- ...wan2.1_fun_workflow_v2v_control_canny.json | 2 +- ...wan2.1_fun_workflow_v2v_control_depth.json | 2 +- ....1_fun_workflow_v2v_control_depth_ref.json | 2 +- .../wan2.1_fun_workflow_v2v_control_pose.json | 2 +- ...2.1_fun_workflow_v2v_control_pose_ref.json | 2 +- .../wan2.1_fun_workflow_v2v_control_ref.json | 2 +- comfyui/wan2_2/nodes.py | 230 +- comfyui/wan2_2/v1/wan2.2_workflow_i2v_5b.json | 471 ++++ comfyui/wan2_2/v1/wan2.2_workflow_t2v_5b.json | 383 +++ comfyui/wan2_2_fun/nodes.py | 232 +- .../wan2.2_fun_workflow_control_camera.json | 670 ++++++ ...an2.2_fun_workflow_control_trajectory.json | 1045 +++++++++ .../v1/wan2.2_fun_workflow_t2v.json | 409 ++++ .../v1/wan2.2_fun_workflow_v2v_control.json | 2 +- ...wan2.2_fun_workflow_v2v_control_canny.json | 700 ++++++ ...wan2.2_fun_workflow_v2v_control_depth.json | 697 ++++++ ...2.2_fun_workflow_v2v_control_pose_ref.json | 753 ++++++ .../wan2.2_fun_workflow_v2v_control_ref.json | 2 +- examples/wan2.2/post_infer_queue.py | 5 +- examples/wan2.2/post_infer_queue_i2v.py | 5 +- examples/wan2.2/predict_ti2v.py | 2 +- examples/wan2.2_fun/app.py | 79 + examples/wan2.2_fun/launch_api.py | 91 + examples/wan2.2_fun/post_infer.py | 150 ++ examples/wan2.2_fun/post_infer_queue.py | 192 ++ examples/wan2.2_fun/post_infer_queue_i2v.py | 213 ++ .../post_infer_queue_v2v_control.py | 230 ++ examples/wan2.2_fun/predict_i2v.py | 83 +- examples/wan2.2_fun/predict_t2v.py | 336 +++ examples/wan2.2_fun/predict_v2v_control.py | 84 +- .../wan2.2_fun/predict_v2v_control_camera.py | 381 +++ .../wan2.2_fun/predict_v2v_control_ref.py | 86 +- scripts/wan2.1_fun/README_TRAIN.md | 4 +- scripts/wan2.2_fun/README_TRAIN.md | 227 ++ scripts/wan2.2_fun/README_TRAIN_CONTROL.md | 272 +++ .../wan2.2_fun/README_TRAIN_CONTROL_LORA.md | 262 +++ scripts/wan2.2_fun/README_TRAIN_LORA.md | 217 ++ scripts/wan2.2_fun/train.py | 1916 +++++++++++++++ scripts/wan2.2_fun/train.sh | 43 + scripts/wan2.2_fun/train_control.py | 2069 +++++++++++++++++ scripts/wan2.2_fun/train_control.sh | 45 + scripts/wan2.2_fun/train_lora.sh | 1 + .../pipeline/pipeline_wan2_2_fun_control.py | 2 +- videox_fun/pipeline/pipeline_wan2_2_ti2v.py | 4 +- videox_fun/ui/controller.py | 3 +- videox_fun/ui/wan2_2_fun_ui.py | 803 +++++++ videox_fun/ui/wan2_2_ui.py | 2 +- 55 files changed, 13550 insertions(+), 772 deletions(-) mode change 100755 => 100644 comfyui/wan2_1_fun/v1/wan2.1_fun_workflow_control_camera.json mode change 100755 => 100644 comfyui/wan2_1_fun/v1/wan2.1_fun_workflow_control_trajectory.json create mode 100644 comfyui/wan2_2/v1/wan2.2_workflow_i2v_5b.json create mode 100644 comfyui/wan2_2/v1/wan2.2_workflow_t2v_5b.json create mode 100644 comfyui/wan2_2_fun/v1/wan2.2_fun_workflow_control_camera.json create mode 100644 comfyui/wan2_2_fun/v1/wan2.2_fun_workflow_control_trajectory.json create mode 100644 comfyui/wan2_2_fun/v1/wan2.2_fun_workflow_t2v.json create mode 100644 comfyui/wan2_2_fun/v1/wan2.2_fun_workflow_v2v_control_canny.json create mode 100644 comfyui/wan2_2_fun/v1/wan2.2_fun_workflow_v2v_control_depth.json create mode 100644 comfyui/wan2_2_fun/v1/wan2.2_fun_workflow_v2v_control_pose_ref.json create mode 100755 examples/wan2.2_fun/app.py create mode 100755 examples/wan2.2_fun/launch_api.py create mode 100755 examples/wan2.2_fun/post_infer.py create mode 100755 examples/wan2.2_fun/post_infer_queue.py create mode 100755 examples/wan2.2_fun/post_infer_queue_i2v.py create mode 100755 examples/wan2.2_fun/post_infer_queue_v2v_control.py create mode 100644 examples/wan2.2_fun/predict_t2v.py create mode 100644 examples/wan2.2_fun/predict_v2v_control_camera.py create mode 100755 scripts/wan2.2_fun/README_TRAIN.md create mode 100755 scripts/wan2.2_fun/README_TRAIN_CONTROL.md create mode 100755 scripts/wan2.2_fun/README_TRAIN_CONTROL_LORA.md create mode 100755 scripts/wan2.2_fun/README_TRAIN_LORA.md create mode 100644 scripts/wan2.2_fun/train.py create mode 100644 scripts/wan2.2_fun/train.sh create mode 100644 scripts/wan2.2_fun/train_control.py create mode 100644 scripts/wan2.2_fun/train_control.sh create mode 100644 videox_fun/ui/wan2_2_fun_ui.py diff --git a/README.md b/README.md index 87d9d81..3d3edf4 100755 --- a/README.md +++ b/README.md @@ -550,6 +550,7 @@ CogVideoX-Fun can be found in [Readme Train](scripts/cogvideox_fun/README_TRAIN. |--|--|--|--|--| | Wan2.2-Fun-A14B-InP | 64.0 GB | [🤗Link](https://huggingface.co/alibaba-pai/Wan2.2-Fun-A14B-InP) | [😄Link](https://modelscope.cn/models/PAI/Wan2.2-Fun-A14B-InP) | Wan2.2-Fun-14B text-to-video generation weights, trained at multiple resolutions, supports start-end image prediction. | | Wan2.2-Fun-A14B-Control | 64.0 GB | [🤗Link](https://huggingface.co/alibaba-pai/Wan2.2-Fun-A14B-Control) | [😄Link](https://modelscope.cn/models/PAI/Wan2.2-Fun-A14B-Control)| Wan2.2-Fun-14B video control weights, supporting various control conditions such as Canny, Depth, Pose, MLSD, etc., and trajectory control. Supports multi-resolution (512, 768, 1024) video prediction at 81 frames, trained at 16 frames per second, with multilingual prediction support. | +| Wan2.2-Fun-A14B-Control-Camera | 64.0 GB | [🤗Link](https://huggingface.co/alibaba-pai/Wan2.2-Fun-A14B-Control-Camera) | [😄Link](https://modelscope.cn/models/PAI/Wan2.2-Fun-A14B-Control-Camera)| Wan2.2-Fun-14B camera lens control weights. Supports multi-resolution (512, 768, 1024) video prediction, trained with 81 frames at 16 FPS, supports multilingual prediction. | ## 2. Wan2.2 diff --git a/README_ja-JP.md b/README_ja-JP.md index b7f5472..4717de6 100755 --- a/README_ja-JP.md +++ b/README_ja-JP.md @@ -550,6 +550,7 @@ CogVideoX-Funは[Readme Train](scripts/cogvideox_fun/README_TRAIN.md)と[Readme |------|----------------|------------|-------------|------| | Wan2.2-Fun-A14B-InP | 64.0 GB | [🤗Link](https://huggingface.co/alibaba-pai/Wan2.2-Fun-A14B-InP) | [😄Link](https://modelscope.cn/models/PAI/Wan2.2-Fun-A14B-InP) | Wan2.2-Fun-14Bのテキスト・画像から動画を生成するモデルの重み。複数の解像度で学習されており、動画の最初と最後のフレームの予測をサポートしています。 | | Wan2.2-Fun-A14B-Control | 64.0 GB | [🤗Link](https://huggingface.co/alibaba-pai/Wan2.2-Fun-A14B-Control) | [😄Link](https://modelscope.cn/models/PAI/Wan2.2-Fun-A14B-Control) | Wan2.2-Fun-14Bの動画制御用重み。Canny、Depth、Pose、MLSDなどのさまざまな制御条件に対応しており、軌跡制御もサポートしています。512、768、1024の複数解像度での動画生成が可能で、81フレーム、16fpsで学習されています。多言語対応の予測もサポートしています。 | +| Wan2.2-Fun-A14B-Contro-Camera | 64.0 GB | [🤗リンク](https://huggingface.co/alibaba-pai/Wan2.2-Fun-A14B-Control-Camera) | [😄リンク](https://modelscope.cn/models/PAI/Wan2.2-Fun-A14B-Control-Camera)| Wan2.2-Fun-14Bのカメラレンズ制御重み。512、768、1024のマルチ解像度での動画予測をサポートし、81フレーム、毎秒16フレームで訓練されています。多言語予測に対応しています。 | ## 2. Wan2.2 diff --git a/README_zh-CN.md b/README_zh-CN.md index cf56d8f..de6ec9a 100755 --- a/README_zh-CN.md +++ b/README_zh-CN.md @@ -540,6 +540,7 @@ CogVideoX-Fun可以查看[Readme Train](scripts/cogvideox_fun/README_TRAIN.md) |--|--|--|--|--| | Wan2.2-Fun-A14B-InP | 64.0 GB | [🤗Link](https://huggingface.co/alibaba-pai/Wan2.2-Fun-A14B-InP) | [😄Link](https://modelscope.cn/models/PAI/Wan2.2-Fun-A14B-InP) | Wan2.2-Fun-14B文图生视频权重,以多分辨率训练,支持首尾图预测。 | | Wan2.2-Fun-A14B-Control | 64.0 GB | [🤗Link](https://huggingface.co/alibaba-pai/Wan2.2-Fun-A14B-Control) | [😄Link](https://modelscope.cn/models/PAI/Wan2.2-Fun-A14B-Control)| Wan2.2-Fun-14B视频控制权重,支持不同的控制条件,如Canny、Depth、Pose、MLSD等,同时支持使用轨迹控制。支持多分辨率(512,768,1024)的视频预测,支持多分辨率(512,768,1024)的视频预测,以81帧、每秒16帧进行训练,支持多语言预测 | +| Wan2.2-Fun-A14B-Control-Camera | 64.0 GB | [🤗Link](https://huggingface.co/alibaba-pai/Wan2.2-Fun-A14B-Control-Camera) | [😄Link](https://modelscope.cn/models/PAI/Wan2.2-Fun-A14B-Control-Camera)| Wan2.2-Fun-14B相机镜头控制权重。支持多分辨率(512,768,1024)的视频预测,支持多分辨率(512,768,1024)的视频预测,以81帧、每秒16帧进行训练,支持多语言预测 | ## 2. Wan2.2 diff --git a/comfyui/README.md b/comfyui/README.md index 95a0b5b..5cc4f54 100755 --- a/comfyui/README.md +++ b/comfyui/README.md @@ -45,6 +45,7 @@ remote_zoe= "https://huggingface.co/lllyasviel/Annotators/resolve/main/ZoeD_M12_ |--|--|--|--|--| | Wan2.2-Fun-A14B-InP | 64.0 GB | [🤗Link](https://huggingface.co/alibaba-pai/Wan2.2-Fun-A14B-InP) | [😄Link](https://modelscope.cn/models/PAI/Wan2.2-Fun-A14B-InP) | Wan2.2-Fun-14B text-to-video generation weights, trained at multiple resolutions, supports start-end image prediction. | | Wan2.2-Fun-A14B-Control | 64.0 GB | [🤗Link](https://huggingface.co/alibaba-pai/Wan2.2-Fun-A14B-Control) | [😄Link](https://modelscope.cn/models/PAI/Wan2.2-Fun-A14B-Control)| Wan2.2-Fun-14B video control weights, supporting various control conditions such as Canny, Depth, Pose, MLSD, etc., and trajectory control. Supports multi-resolution (512, 768, 1024) video prediction at 81 frames, trained at 16 frames per second, with multilingual prediction support. | +| Wan2.2-Fun-A14B-Control-Camera | 64.0 GB | [🤗Link](https://huggingface.co/alibaba-pai/Wan2.2-Fun-A14B-Control-Camera) | [😄Link](https://modelscope.cn/models/PAI/Wan2.2-Fun-A14B-Control-Camera)| Wan2.2-Fun-14B camera lens control weights. Supports multi-resolution (512, 768, 1024) video prediction, trained with 81 frames at 16 FPS, supports multilingual prediction. | #### ii. Wan2.2 diff --git a/comfyui/comfyui_nodes.py b/comfyui/comfyui_nodes.py index 31ad1f6..6fad69d 100755 --- a/comfyui/comfyui_nodes.py +++ b/comfyui/comfyui_nodes.py @@ -77,21 +77,22 @@ class FunCompile: for i in range(len(funmodels["pipeline"].transformer.blocks)): funmodels["pipeline"].transformer.blocks[i] = torch.compile(funmodels["pipeline"].transformer.blocks[i]) - if hasattr(funmodels["pipeline"], "transformer_2"): + if hasattr(funmodels["pipeline"], "transformer_2") and funmodels["pipeline"].transformer_2 is not None: for i in range(len(funmodels["pipeline"].transformer_2.blocks)): funmodels["pipeline"].transformer_2.blocks[i] = torch.compile(funmodels["pipeline"].transformer_2.blocks[i]) elif hasattr(funmodels["pipeline"].transformer, "transformer_blocks"): for i in range(len(funmodels["pipeline"].transformer.transformer_blocks)): funmodels["pipeline"].transformer.transformer_blocks[i] = torch.compile(funmodels["pipeline"].transformer.transformer_blocks[i]) - if hasattr(funmodels["pipeline"], "transformer_2"): + + if hasattr(funmodels["pipeline"], "transformer_2") and funmodels["pipeline"].transformer_2 is not None: for i in range(len(funmodels["pipeline"].transformer_2.transformer_blocks)): funmodels["pipeline"].transformer_2.transformer_blocks[i] = torch.compile(funmodels["pipeline"].transformer_2.transformer_blocks[i]) else: funmodels["pipeline"].transformer.forward = torch.compile(funmodels["pipeline"].transformer.forward) - if hasattr(funmodels["pipeline"], "transformer_2"): + if hasattr(funmodels["pipeline"], "transformer_2") and funmodels["pipeline"].transformer_2 is not None: funmodels["pipeline"].transformer_2.forward = torch.compile(funmodels["pipeline"].transformer_2.forward) print("Add Compile") diff --git a/comfyui/wan2_1_fun/v1/wan2.1_fun_workflow_control_camera.json b/comfyui/wan2_1_fun/v1/wan2.1_fun_workflow_control_camera.json old mode 100755 new mode 100644 index 16820c9..50ae041 --- a/comfyui/wan2_1_fun/v1/wan2.1_fun_workflow_control_camera.json +++ b/comfyui/wan2_1_fun/v1/wan2.1_fun_workflow_control_camera.json @@ -1,18 +1,20 @@ { + "id": "addf6f13-291c-4b88-a80c-cd5785fa8f42", + "revision": 0, "last_node_id": 132, "last_link_id": 292, "nodes": [ { "id": 107, "type": "Note", - "pos": { - "0": 4, - "1": 634 - }, - "size": { - "0": 210, - "1": 58 - }, + "pos": [ + 4, + 634 + ], + "size": [ + 210, + 88 + ], "flags": {}, "order": 0, "mode": 0, @@ -30,14 +32,14 @@ { "id": 108, "type": "Note", - "pos": { - "0": -110, - "1": 842 - }, - "size": { - "0": 326.1556091308594, - "1": 145.20904541015625 - }, + "pos": [ + -110, + 842 + ], + "size": [ + 326.1556091308594, + 145.20904541015625 + ], "flags": {}, "order": 1, "mode": 0, @@ -52,54 +54,29 @@ "color": "#432", "bgcolor": "#653" }, - { - "id": 112, - "type": "Note", - "pos": { - "0": -203, - "1": 252 - }, - "size": { - "0": 427.074951171875, - "1": 143.9142608642578 - }, - "flags": {}, - "order": 2, - "mode": 0, - "inputs": [], - "outputs": [], - "properties": { - "text": "" - }, - "widgets_values": [ - "Due to the large size of models from EasyAnimateV5 and above, when using the 12B model, if your graphics card has 24GB or less of VRAM, please set GPU_memory_mode to model_cpu_offload_and_qfloat8. This will load the model in float8 to reduce VRAM consumption, otherwise you may receive an out-of-memory error. \n(由于EasyAnimateV5以上的模型较大,当使用12B模型时,如果使用的显卡显存为24G及以下,请将GPU_memory_mode设置为model_cpu_offload_and_qfloat8,使得模型加载在float8上减少显存消耗,否则会提示显存不足。)" - ], - "color": "#432", - "bgcolor": "#653" - }, { "id": 122, "type": "FunTextBox", - "pos": { - "0": 238, - "1": 805 - }, - "size": { - "0": 400, - "1": 200 - }, + "pos": [ + 238, + 805 + ], + "size": [ + 400, + 200 + ], "flags": {}, - "order": 3, + "order": 2, "mode": 0, "inputs": [], "outputs": [ { "name": "prompt", "type": "STRING_PROMPT", + "slot_index": 0, "links": [ 289 - ], - "slot_index": 0 + ] } ], "title": "Negtive Prompt(反向提示词)", @@ -113,16 +90,16 @@ { "id": 129, "type": "CameraBasicFromChaoJie", - "pos": { - "0": 805.2059326171875, - "1": 1012.381103515625 - }, - "size": { - "0": 315, - "1": 106 - }, + "pos": [ + 805.2059326171875, + 1012.381103515625 + ], + "size": [ + 315, + 106 + ], "flags": {}, - "order": 4, + "order": 3, "mode": 0, "inputs": [], "outputs": [ @@ -144,14 +121,14 @@ { "id": 130, "type": "CameraTrajectoryFromChaoJie", - "pos": { - "0": 1170.206298828125, - "1": 763.3814697265625 - }, - "size": { - "0": 367.79998779296875, - "1": 150 - }, + "pos": [ + 1170.206298828125, + 763.3814697265625 + ], + "size": [ + 367.79998779296875, + 150 + ], "flags": {}, "order": 10, "mode": 0, @@ -166,10 +143,10 @@ { "name": "camera_trajectory", "type": "STRING", + "slot_index": 0, "links": [ 292 - ], - "slot_index": 0 + ] }, { "name": "video_length", @@ -190,55 +167,53 @@ { "id": 106, "type": "VHS_VideoCombine", - "pos": { - "0": 1408, - "1": 68 - }, + "pos": [ + 1408, + 68 + ], "size": [ 390, - 537.4615384615385 + 537.4615478515625 ], "flags": {}, "order": 12, "mode": 0, "inputs": [ { - "name": "images", - "type": "IMAGE", - "link": 291, - "slot_index": 0, "label": "图像", - "shape": 7 + "name": "images", + "shape": 7, + "type": "IMAGE", + "link": 291 }, { - "name": "audio", - "type": "AUDIO", - "link": null, "label": "音频", - "shape": 7 + "name": "audio", + "shape": 7, + "type": "AUDIO", + "link": null }, { - "name": "meta_batch", - "type": "VHS_BatchManager", - "link": null, "label": "批次管理", - "shape": 7 + "name": "meta_batch", + "shape": 7, + "type": "VHS_BatchManager", + "link": null }, { "name": "vae", + "shape": 7, "type": "VAE", - "link": null, - "shape": 7 + "link": null } ], "outputs": [ { + "label": "文件名", "name": "Filenames", "type": "VHS_FILENAMES", - "links": null, "slot_index": 0, - "shape": 3, - "label": "文件名" + "links": null } ], "properties": { @@ -270,16 +245,16 @@ { "id": 121, "type": "FunTextBox", - "pos": { - "0": 235, - "1": 539 - }, - "size": { - "0": 400, - "1": 200 - }, + "pos": [ + 235, + 539 + ], + "size": [ + 400, + 200 + ], "flags": {}, - "order": 5, + "order": 4, "mode": 0, "inputs": [], "outputs": [ @@ -302,35 +277,33 @@ { "id": 100, "type": "LoadImage", - "pos": { - "0": 237.59738159179688, - "1": 1164.597412109375 - }, - "size": { - "0": 378.07147216796875, - "1": 314 - }, + "pos": [ + 237.59738159179688, + 1164.597412109375 + ], + "size": [ + 378.07147216796875, + 314 + ], "flags": {}, - "order": 6, + "order": 5, "mode": 0, "inputs": [], "outputs": [ { + "label": "图像", "name": "IMAGE", "type": "IMAGE", + "slot_index": 0, "links": [ 290 - ], - "slot_index": 0, - "shape": 3, - "label": "图像" + ] }, { + "label": "遮罩", "name": "MASK", "type": "MASK", - "links": null, - "shape": 3, - "label": "遮罩" + "links": null } ], "title": "Start Image(图片到视频的开始图片)", @@ -345,16 +318,16 @@ { "id": 131, "type": "CameraCombineFromChaoJie", - "pos": { - "0": 814.2059326171875, - "1": 763.3814697265625 - }, - "size": { - "0": 315, - "1": 178 - }, + "pos": [ + 814.2059326171875, + 763.3814697265625 + ], + "size": [ + 315, + 178 + ], "flags": {}, - "order": 7, + "order": 6, "mode": 0, "inputs": [], "outputs": [ @@ -381,16 +354,16 @@ { "id": 110, "type": "Note", - "pos": { - "0": 1158.206298828125, - "1": 970.381103515625 - }, - "size": { - "0": 608.1410522460938, - "1": 188.2682342529297 - }, + "pos": [ + 1158.206298828125, + 970.381103515625 + ], + "size": [ + 608.1410522460938, + 188.2682342529297 + ], "flags": {}, - "order": 8, + "order": 7, "mode": 0, "inputs": [], "outputs": [], @@ -406,14 +379,14 @@ { "id": 132, "type": "WanFunV2VSampler", - "pos": { - "0": 899, - "1": 68 - }, - "size": { - "0": 428.4000244140625, - "1": 486 - }, + "pos": [ + 899, + 68 + ], + "size": [ + 428.4000244140625, + 506 + ], "flags": {}, "order": 11, "mode": 0, @@ -435,42 +408,39 @@ }, { "name": "validation_video", + "shape": 7, "type": "IMAGE", - "link": null, - "shape": 7 + "link": null }, { "name": "control_video", + "shape": 7, "type": "IMAGE", - "link": null, - "shape": 7 + "link": null }, { "name": "start_image", + "shape": 7, "type": "IMAGE", - "link": 290, - "shape": 7 + "link": 290 }, { "name": "ref_image", + "shape": 7, "type": "IMAGE", - "link": null, - "shape": 7 - }, - { - "name": "riflex_k", - "type": "RIFLEXT_ARGS", - "link": null, - "shape": 7 + "link": null }, { "name": "camera_conditions", + "shape": 7, "type": "STRING", - "link": 292, - "widget": { - "name": "camera_conditions" - }, - "shape": 7 + "link": 292 + }, + { + "name": "riflex_k", + "shape": 7, + "type": "RIFLEXT_ARGS", + "link": null } ], "outputs": [ @@ -492,7 +462,7 @@ "fixed", 50, 6, - 1.0, + 1, "Flow", 0.1, true, @@ -504,26 +474,26 @@ { "id": 123, "type": "LoadWanFunModel", - "pos": { - "0": 281, - "1": 251 - }, - "size": { - "0": 315, - "1": 154 - }, + "pos": [ + 281, + 251 + ], + "size": [ + 315, + 154 + ], "flags": {}, - "order": 9, + "order": 8, "mode": 0, "inputs": [], "outputs": [ { "name": "funmodels", "type": "FunModels", + "slot_index": 0, "links": [ 287 - ], - "slot_index": 0 + ] } ], "properties": { @@ -536,6 +506,31 @@ "wan2.1/wan_civitai.yaml", "bf16" ] + }, + { + "id": 112, + "type": "Note", + "pos": [ + -203, + 252 + ], + "size": [ + 427.074951171875, + 143.9142608642578 + ], + "flags": {}, + "order": 9, + "mode": 0, + "inputs": [], + "outputs": [], + "properties": { + "text": "" + }, + "widgets_values": [ + "When using the 1.3B model, you can set GPU_memory_mode to model_cpu_offload for faster generation. When using the 14B model, you can use sequential_cpu_offload to save GPU memory during generation.\n(在使用1.3B模型时,可以设置GPU_memory_mode为model_cpu_offload进行更快速度的生成,在使用14B模型时,可以使用sequential_cpu_offload节省显存,进行生成。)" + ], + "color": "#432", + "bgcolor": "#653" } ], "links": [ @@ -592,12 +587,13 @@ 130, 0, 132, - 8, + 7, "STRING" ] ], "groups": [ { + "id": 1, "title": "Generate Control Video", "bounding": [ 773, @@ -610,6 +606,7 @@ "flags": {} }, { + "id": 2, "title": "First Image", "bounding": [ 191, @@ -622,6 +619,7 @@ "flags": {} }, { + "id": 3, "title": "Prompts", "bounding": [ 191, @@ -634,7 +632,8 @@ "flags": {} }, { - "title": "Load EasyAnimate", + "id": 4, + "title": "Load Model", "bounding": [ 189, 160, @@ -649,17 +648,18 @@ "config": {}, "extra": { "ds": { - "scale": 0.8264462809917358, + "scale": 0.6830134553650709, "offset": [ - 28.17192681115923, - -3.293207324975433 + 236.8040025924094, + -27.611371387475245 ] }, "node_versions": { - "CogVideoX-Fun": "a7fa7028d52498f13e983eba012a81ebcae24977", + "CogVideoX-Fun": "e935795d7684aaba72be65d2ad5ef580c0526f7a", "ComfyUI-VideoHelperSuite": "70faa9bcef65932ab72e7404d6373fb300013a2e", - "comfy-core": "v0.2.7-3-g8afb97c" - } + "comfy-core": "0.3.44" + }, + "frontendVersion": "1.21.3" }, "version": 0.4 } \ No newline at end of file diff --git a/comfyui/wan2_1_fun/v1/wan2.1_fun_workflow_control_trajectory.json b/comfyui/wan2_1_fun/v1/wan2.1_fun_workflow_control_trajectory.json old mode 100755 new mode 100644 index 54b1cfe..d6e14b9 --- a/comfyui/wan2_1_fun/v1/wan2.1_fun_workflow_control_trajectory.json +++ b/comfyui/wan2_1_fun/v1/wan2.1_fun_workflow_control_trajectory.json @@ -1,18 +1,20 @@ { + "id": "34fdaa39-a397-46b5-84d1-8846799c8c09", + "revision": 0, "last_node_id": 126, "last_link_id": 287, "nodes": [ { "id": 107, "type": "Note", - "pos": { - "0": 4, - "1": 634 - }, - "size": { - "0": 210, - "1": 58 - }, + "pos": [ + 4, + 634 + ], + "size": [ + 210, + 88 + ], "flags": {}, "order": 0, "mode": 0, @@ -30,14 +32,14 @@ { "id": 108, "type": "Note", - "pos": { - "0": -110, - "1": 842 - }, - "size": { - "0": 326.1556091308594, - "1": 145.20904541015625 - }, + "pos": [ + -110, + 842 + ], + "size": [ + 326.1556091308594, + 145.20904541015625 + ], "flags": {}, "order": 1, "mode": 0, @@ -53,62 +55,35 @@ "bgcolor": "#653" }, { - "id": 112, - "type": "Note", - "pos": { - "0": -203, - "1": 252 - }, - "size": { - "0": 427.074951171875, - "1": 143.9142608642578 - }, + "id": 100, + "type": "LoadImage", + "pos": [ + 238, + 1165 + ], + "size": [ + 378.07147216796875, + 314 + ], "flags": {}, "order": 2, "mode": 0, "inputs": [], - "outputs": [], - "properties": { - "text": "" - }, - "widgets_values": [ - "Due to the large size of models from EasyAnimateV5 and above, when using the 12B model, if your graphics card has 24GB or less of VRAM, please set GPU_memory_mode to model_cpu_offload_and_qfloat8. This will load the model in float8 to reduce VRAM consumption, otherwise you may receive an out-of-memory error. \n(由于EasyAnimateV5以上的模型较大,当使用12B模型时,如果使用的显卡显存为24G及以下,请将GPU_memory_mode设置为model_cpu_offload_and_qfloat8,使得模型加载在float8上减少显存消耗,否则会提示显存不足。)" - ], - "color": "#432", - "bgcolor": "#653" - }, - { - "id": 100, - "type": "LoadImage", - "pos": { - "0": 238, - "1": 1165 - }, - "size": { - "0": 378.07147216796875, - "1": 314 - }, - "flags": {}, - "order": 3, - "mode": 0, - "inputs": [], "outputs": [ { + "label": "图像", "name": "IMAGE", "type": "IMAGE", + "slot_index": 0, "links": [ 285 - ], - "slot_index": 0, - "shape": 3, - "label": "图像" + ] }, { + "label": "遮罩", "name": "MASK", "type": "MASK", - "links": null, - "shape": 3, - "label": "遮罩" + "links": null } ], "title": "Start Image(图片到视频的开始图片)", @@ -123,16 +98,16 @@ { "id": 110, "type": "Note", - "pos": { - "0": 847, - "1": 613 - }, - "size": { - "0": 608.1410522460938, - "1": 188.2682342529297 - }, + "pos": [ + 847, + 613 + ], + "size": [ + 608.1410522460938, + 188.2682342529297 + ], "flags": {}, - "order": 4, + "order": 3, "mode": 0, "inputs": [], "outputs": [], @@ -148,14 +123,14 @@ { "id": 118, "type": "AppendStringsToList", - "pos": { - "0": 1140.1396484375, - "1": 909.9193115234375 - }, - "size": { - "0": 315, - "1": 82 - }, + "pos": [ + 1140.1396484375, + 909.9193115234375 + ], + "size": [ + 315, + 82 + ], "flags": { "collapsed": false }, @@ -165,51 +140,42 @@ { "name": "string1", "type": "STRING", - "link": 265, - "widget": { - "name": "string1" - } + "link": 265 }, { "name": "string2", "type": "STRING", - "link": 266, - "widget": { - "name": "string2" - } + "link": 266 } ], "outputs": [ { "name": "STRING", "type": "STRING", + "slot_index": 0, "links": [ 267 - ], - "slot_index": 0 + ] } ], "properties": { "Node name for S&R": "AppendStringsToList" }, - "widgets_values": [ - "", - "" - ] + "widgets_values": [] }, { "id": 121, "type": "FunTextBox", - "pos": { - "0": 235, - "1": 539 - }, - "size": { - "0": 400, - "1": 200 - }, + "pos": [ + 235, + 539 + ], + "size": [ + 400, + 200 + ], "flags": {}, - "order": 5, + "order": 4, "mode": 0, "inputs": [], "outputs": [ @@ -232,14 +198,14 @@ { "id": 114, "type": "ImageMaximumNode", - "pos": { - "0": 2074, - "1": 905 - }, - "size": { - "0": 210, - "1": 46 - }, + "pos": [ + 2074, + 905 + ], + "size": [ + 210, + 46 + ], "flags": {}, "order": 15, "mode": 0, @@ -272,74 +238,69 @@ { "id": 95, "type": "CreateTrajectoryBasedOnKJNodes", - "pos": { - "0": 1574.139404296875, - "1": 929.9193115234375 - }, - "size": { - "0": 428.4000244140625, - "1": 58 - }, + "pos": [ + 1574.139404296875, + 929.9193115234375 + ], + "size": [ + 428.4000244140625, + 58 + ], "flags": {}, "order": 11, "mode": 0, "inputs": [ + { + "name": "coordinates", + "type": "STRING", + "link": 267 + }, { "name": "masks", "type": "MASK", "link": 249 - }, - { - "name": "coordinates", - "type": "STRING", - "link": 267, - "widget": { - "name": "coordinates" - } } ], "outputs": [ { "name": "image", "type": "IMAGE", + "slot_index": 0, "links": [ 237, 262, 284 - ], - "slot_index": 0 + ] } ], "properties": { "Node name for S&R": "CreateTrajectoryBasedOnKJNodes" }, - "widgets_values": [ - "" - ] + "widgets_values": [] }, { "id": 122, "type": "FunTextBox", - "pos": { - "0": 238, - "1": 805 - }, - "size": { - "0": 400, - "1": 200 - }, + "pos": [ + 238, + 805 + ], + "size": [ + 400, + 200 + ], "flags": {}, - "order": 6, + "order": 5, "mode": 0, "inputs": [], "outputs": [ { "name": "prompt", "type": "STRING_PROMPT", + "slot_index": 0, "links": [ 283 - ], - "slot_index": 0 + ] } ], "title": "Negtive Prompt(反向提示词)", @@ -353,41 +314,41 @@ { "id": 97, "type": "SplineEditor", - "pos": { - "0": 855.1397705078125, - "1": 1058.91943359375 - }, - "size": { - "0": 645, - "1": 812 - }, + "pos": [ + 855.1397705078125, + 1058.91943359375 + ], + "size": [ + 645, + 832 + ], "flags": {}, - "order": 7, + "order": 6, "mode": 0, "inputs": [ { "name": "bg_image", + "shape": 7, "type": "IMAGE", - "link": null, - "shape": 7 + "link": null } ], "outputs": [ { "name": "mask", "type": "MASK", + "slot_index": 0, "links": [ 249 - ], - "slot_index": 0 + ] }, { "name": "coord_str", "type": "STRING", + "slot_index": 1, "links": [ 265 - ], - "slot_index": 1 + ] }, { "name": "float", @@ -411,8 +372,8 @@ "imgData": null }, "widgets_values": [ - "[{\"x\":236.74497000000005,\"y\":169.10355000000004},{\"x\":263.53799999999995,\"y\":230.59575},{\"x\":321.61075881587004,\"y\":229.72197058276433},{\"x\":355.3934015486295,\"y\":168.9132136637973},{\"x\":343.23165016483614,\"y\":82.42964826793309},{\"x\":275.6663646993172,\"y\":55.40353408172552},{\"x\":197.29063355931524,\"y\":85.13225968655384},{\"x\":177.02104791965957,\"y\":148.64362802414163},{\"x\":205.39846781517753,\"y\":232.4245820013851},{\"x\":259.45069618759265,\"y\":167.56190795448694}]", - "[{\"x\":236.74496459960938,\"y\":169.10354614257812},{\"x\":238.94464111328125,\"y\":177.68576049804688},{\"x\":241.3327178955078,\"y\":186.21737670898438},{\"x\":243.95693969726562,\"y\":194.67919921875},{\"x\":246.88812255859375,\"y\":203.03933715820312},{\"x\":250.2393035888672,\"y\":211.23939514160156},{\"x\":254.20826721191406,\"y\":219.15655517578125},{\"x\":259.18768310546875,\"y\":226.46995544433594},{\"x\":265.87689208984375,\"y\":232.22171020507812},{\"x\":273.4639892578125,\"y\":236.787353515625},{\"x\":281.6380920410156,\"y\":240.17092895507812},{\"x\":290.3527526855469,\"y\":241.61793518066406},{\"x\":299.1367492675781,\"y\":240.67857360839844},{\"x\":307.49658203125,\"y\":237.7835235595703},{\"x\":315.3465576171875,\"y\":233.68600463867188},{\"x\":322.80859375,\"y\":228.9126739501953},{\"x\":330.00958251953125,\"y\":223.75323486328125},{\"x\":336.7490234375,\"y\":218.0097198486328},{\"x\":342.5494689941406,\"y\":211.32901000976562},{\"x\":346.96185302734375,\"y\":203.6618194580078},{\"x\":350.05450439453125,\"y\":195.3665771484375},{\"x\":352.2560729980469,\"y\":186.787109375},{\"x\":353.9389343261719,\"y\":178.08937072753906},{\"x\":355.33001708984375,\"y\":169.3397216796875},{\"x\":356.59344482421875,\"y\":160.57052612304688},{\"x\":357.7420959472656,\"y\":151.78565979003906},{\"x\":358.6518249511719,\"y\":142.9732208251953},{\"x\":359.1504211425781,\"y\":134.1287384033203},{\"x\":359.0207214355469,\"y\":125.2725601196289},{\"x\":358.035400390625,\"y\":116.4721908569336},{\"x\":356.0333251953125,\"y\":107.84730529785156},{\"x\":352.9960632324219,\"y\":99.53003692626953},{\"x\":349.046142578125,\"y\":91.60392761230469},{\"x\":344.37164306640625,\"y\":84.08070373535156},{\"x\":339.17694091796875,\"y\":76.9051742553711},{\"x\":333.4718322753906,\"y\":70.13098907470703},{\"x\":326.9710693359375,\"y\":64.12535858154297},{\"x\":319.4295654296875,\"y\":59.51286315917969},{\"x\":311.02655029296875,\"y\":56.7620735168457},{\"x\":302.25640869140625,\"y\":55.55453872680664},{\"x\":293.40484619140625,\"y\":55.21137619018555},{\"x\":284.5453796386719,\"y\":55.24800491333008},{\"x\":275.68701171875,\"y\":55.40314865112305},{\"x\":266.83221435546875,\"y\":55.690216064453125},{\"x\":257.9945373535156,\"y\":56.30451965332031},{\"x\":249.2041473388672,\"y\":57.398460388183594},{\"x\":240.5244140625,\"y\":59.16102981567383},{\"x\":232.06307983398438,\"y\":61.773162841796875},{\"x\":223.95306396484375,\"y\":65.32772827148438},{\"x\":216.29017639160156,\"y\":69.7664566040039},{\"x\":209.08375549316406,\"y\":74.91584777832031},{\"x\":202.27122497558594,\"y\":80.57798767089844},{\"x\":195.76087951660156,\"y\":86.58625030517578},{\"x\":189.468505859375,\"y\":92.8224105834961},{\"x\":183.6200714111328,\"y\":99.47235107421875},{\"x\":178.9163055419922,\"y\":106.95783233642578},{\"x\":176.3991241455078,\"y\":115.42044830322266},{\"x\":175.83934020996094,\"y\":124.25370788574219},{\"x\":176.12339782714844,\"y\":133.10797119140625},{\"x\":176.63381958007812,\"y\":141.95301818847656},{\"x\":177.14378356933594,\"y\":150.79808044433594},{\"x\":177.73876953125,\"y\":159.63778686523438},{\"x\":178.5083770751953,\"y\":168.46385192871094},{\"x\":179.4958038330078,\"y\":177.26815795898438},{\"x\":180.75762939453125,\"y\":186.03709411621094},{\"x\":182.37294006347656,\"y\":194.7476043701172},{\"x\":184.4583282470703,\"y\":203.35687255859375},{\"x\":187.1986083984375,\"y\":211.7789306640625},{\"x\":190.9047393798828,\"y\":219.81732177734375},{\"x\":196.108154296875,\"y\":226.9572296142578},{\"x\":203.41249084472656,\"y\":231.83749389648438},{\"x\":211.85409545898438,\"y\":230.89816284179688},{\"x\":218.7886199951172,\"y\":225.41867065429688},{\"x\":224.8160400390625,\"y\":218.92974853515625},{\"x\":230.3641357421875,\"y\":212.0236358642578},{\"x\":235.608642578125,\"y\":204.8833770751953},{\"x\":240.63941955566406,\"y\":197.59080505371094},{\"x\":245.50965881347656,\"y\":190.1898651123047},{\"x\":250.25364685058594,\"y\":182.7073516845703},{\"x\":254.8948974609375,\"y\":175.16050720214844},{\"x\":259.45068359375,\"y\":167.56190490722656}]", + "[{\"points\":[{\"x\":236.74497000000005,\"y\":169.10355000000004},{\"x\":263.53799999999995,\"y\":230.59575},{\"x\":321.61075881587004,\"y\":229.72197058276433},{\"x\":355.3934015486295,\"y\":168.9132136637973},{\"x\":343.23165016483614,\"y\":82.42964826793309},{\"x\":275.6663646993172,\"y\":55.40353408172552},{\"x\":197.29063355931524,\"y\":85.13225968655384},{\"x\":177.02104791965957,\"y\":148.64362802414163},{\"x\":205.39846781517753,\"y\":232.4245820013851},{\"x\":259.45069618759265,\"y\":167.56190795448694}],\"color\":\"#1f77b4\",\"name\":\"Spline 1\"}]", + "[[{\"x\":236.74496459960938,\"y\":169.10354614257812},{\"x\":238.94464111328125,\"y\":177.68576049804688},{\"x\":241.3327178955078,\"y\":186.21737670898438},{\"x\":243.95693969726562,\"y\":194.67919921875},{\"x\":246.88812255859375,\"y\":203.03933715820312},{\"x\":250.2393035888672,\"y\":211.23939514160156},{\"x\":254.20826721191406,\"y\":219.15655517578125},{\"x\":259.18768310546875,\"y\":226.46995544433594},{\"x\":265.87689208984375,\"y\":232.22171020507812},{\"x\":273.4639892578125,\"y\":236.787353515625},{\"x\":281.6380920410156,\"y\":240.17092895507812},{\"x\":290.3527526855469,\"y\":241.61793518066406},{\"x\":299.1367492675781,\"y\":240.67857360839844},{\"x\":307.49658203125,\"y\":237.7835235595703},{\"x\":315.3465576171875,\"y\":233.68600463867188},{\"x\":322.80859375,\"y\":228.9126739501953},{\"x\":330.00958251953125,\"y\":223.75323486328125},{\"x\":336.7490234375,\"y\":218.0097198486328},{\"x\":342.5494689941406,\"y\":211.32901000976562},{\"x\":346.96185302734375,\"y\":203.6618194580078},{\"x\":350.05450439453125,\"y\":195.3665771484375},{\"x\":352.2560729980469,\"y\":186.787109375},{\"x\":353.9389343261719,\"y\":178.08937072753906},{\"x\":355.33001708984375,\"y\":169.3397216796875},{\"x\":356.59344482421875,\"y\":160.57052612304688},{\"x\":357.7420959472656,\"y\":151.78565979003906},{\"x\":358.6518249511719,\"y\":142.9732208251953},{\"x\":359.1504211425781,\"y\":134.1287384033203},{\"x\":359.0207214355469,\"y\":125.2725601196289},{\"x\":358.035400390625,\"y\":116.4721908569336},{\"x\":356.0333251953125,\"y\":107.84730529785156},{\"x\":352.9960632324219,\"y\":99.53003692626953},{\"x\":349.046142578125,\"y\":91.60392761230469},{\"x\":344.37164306640625,\"y\":84.08070373535156},{\"x\":339.17694091796875,\"y\":76.9051742553711},{\"x\":333.4718322753906,\"y\":70.13098907470703},{\"x\":326.9710693359375,\"y\":64.12535858154297},{\"x\":319.4295654296875,\"y\":59.51286315917969},{\"x\":311.02655029296875,\"y\":56.7620735168457},{\"x\":302.25640869140625,\"y\":55.55453872680664},{\"x\":293.40484619140625,\"y\":55.21137619018555},{\"x\":284.5453796386719,\"y\":55.24800491333008},{\"x\":275.68701171875,\"y\":55.40314865112305},{\"x\":266.83221435546875,\"y\":55.690216064453125},{\"x\":257.9945373535156,\"y\":56.30451965332031},{\"x\":249.2041473388672,\"y\":57.398460388183594},{\"x\":240.5244140625,\"y\":59.16102981567383},{\"x\":232.06307983398438,\"y\":61.773162841796875},{\"x\":223.95306396484375,\"y\":65.32772827148438},{\"x\":216.29017639160156,\"y\":69.7664566040039},{\"x\":209.08375549316406,\"y\":74.91584777832031},{\"x\":202.27122497558594,\"y\":80.57798767089844},{\"x\":195.76087951660156,\"y\":86.58625030517578},{\"x\":189.468505859375,\"y\":92.8224105834961},{\"x\":183.6200714111328,\"y\":99.47235107421875},{\"x\":178.9163055419922,\"y\":106.95783233642578},{\"x\":176.3991241455078,\"y\":115.42044830322266},{\"x\":175.83934020996094,\"y\":124.25370788574219},{\"x\":176.12339782714844,\"y\":133.10797119140625},{\"x\":176.63381958007812,\"y\":141.95301818847656},{\"x\":177.14378356933594,\"y\":150.79808044433594},{\"x\":177.73876953125,\"y\":159.63778686523438},{\"x\":178.5083770751953,\"y\":168.46385192871094},{\"x\":179.4958038330078,\"y\":177.26815795898438},{\"x\":180.75762939453125,\"y\":186.03709411621094},{\"x\":182.37294006347656,\"y\":194.7476043701172},{\"x\":184.4583282470703,\"y\":203.35687255859375},{\"x\":187.1986083984375,\"y\":211.7789306640625},{\"x\":190.9047393798828,\"y\":219.81732177734375},{\"x\":196.108154296875,\"y\":226.9572296142578},{\"x\":203.41249084472656,\"y\":231.83749389648438},{\"x\":211.85409545898438,\"y\":230.89816284179688},{\"x\":218.7886199951172,\"y\":225.41867065429688},{\"x\":224.8160400390625,\"y\":218.92974853515625},{\"x\":230.3641357421875,\"y\":212.0236358642578},{\"x\":235.608642578125,\"y\":204.8833770751953},{\"x\":240.63941955566406,\"y\":197.59080505371094},{\"x\":245.50965881347656,\"y\":190.1898651123047},{\"x\":250.25364685058594,\"y\":182.7073516845703},{\"x\":254.8948974609375,\"y\":175.16050720214844},{\"x\":259.45068359375,\"y\":167.56190490722656}]]", 600, 382, 81, @@ -423,47 +384,46 @@ "list", 0, 1, - null, - null, + "", null ] }, { "id": 119, "type": "SplineEditor", - "pos": { - "0": 1544.139404296875, - "1": 1047.919189453125 - }, - "size": { - "0": 645, - "1": 812 - }, + "pos": [ + 1544.139404296875, + 1047.919189453125 + ], + "size": [ + 645, + 832 + ], "flags": {}, - "order": 8, + "order": 7, "mode": 0, "inputs": [ { "name": "bg_image", + "shape": 7, "type": "IMAGE", - "link": null, - "shape": 7 + "link": null } ], "outputs": [ { "name": "mask", "type": "MASK", - "links": [], - "slot_index": 0 + "slot_index": 0, + "links": [] }, { "name": "coord_str", "type": "STRING", + "slot_index": 1, "links": [ 266 - ], - "slot_index": 1 + ] }, { "name": "float", @@ -487,8 +447,8 @@ "imgData": null }, "widgets_values": [ - "[{\"x\":63.916760050380844,\"y\":114.45559357858895},{\"x\":71.34894145158792,\"y\":114.45559357858895}]", - "[{\"x\":63.9167594909668,\"y\":114.45559692382812},{\"x\":64.00965881347656,\"y\":114.45559692382812},{\"x\":64.1025619506836,\"y\":114.45559692382812},{\"x\":64.19546508789062,\"y\":114.45559692382812},{\"x\":64.28836822509766,\"y\":114.45559692382812},{\"x\":64.38127136230469,\"y\":114.45559692382812},{\"x\":64.47417449951172,\"y\":114.45559692382812},{\"x\":64.56707763671875,\"y\":114.45559692382812},{\"x\":64.65998077392578,\"y\":114.45559692382812},{\"x\":64.75287628173828,\"y\":114.45559692382812},{\"x\":64.84577941894531,\"y\":114.45559692382812},{\"x\":64.93868255615234,\"y\":114.45559692382812},{\"x\":65.03158569335938,\"y\":114.45559692382812},{\"x\":65.1244888305664,\"y\":114.45559692382812},{\"x\":65.21739196777344,\"y\":114.45559692382812},{\"x\":65.31029510498047,\"y\":114.45559692382812},{\"x\":65.4031982421875,\"y\":114.45559692382812},{\"x\":65.49609375,\"y\":114.45559692382812},{\"x\":65.58899688720703,\"y\":114.45559692382812},{\"x\":65.68190002441406,\"y\":114.45559692382812},{\"x\":65.7748031616211,\"y\":114.45559692382812},{\"x\":65.86770629882812,\"y\":114.45559692382812},{\"x\":65.96060943603516,\"y\":114.45559692382812},{\"x\":66.05351257324219,\"y\":114.45559692382812},{\"x\":66.14641571044922,\"y\":114.45559692382812},{\"x\":66.23931884765625,\"y\":114.45559692382812},{\"x\":66.33221435546875,\"y\":114.45559692382812},{\"x\":66.42511749267578,\"y\":114.45559692382812},{\"x\":66.51802062988281,\"y\":114.45559692382812},{\"x\":66.61092376708984,\"y\":114.45559692382812},{\"x\":66.70382690429688,\"y\":114.45559692382812},{\"x\":66.7967300415039,\"y\":114.45559692382812},{\"x\":66.88963317871094,\"y\":114.45559692382812},{\"x\":66.98253631591797,\"y\":114.45559692382812},{\"x\":67.075439453125,\"y\":114.45559692382812},{\"x\":67.1683349609375,\"y\":114.45559692382812},{\"x\":67.26123809814453,\"y\":114.45559692382812},{\"x\":67.35414123535156,\"y\":114.45559692382812},{\"x\":67.4470443725586,\"y\":114.45559692382812},{\"x\":67.53994750976562,\"y\":114.45559692382812},{\"x\":67.63285064697266,\"y\":114.45559692382812},{\"x\":67.72575378417969,\"y\":114.45559692382812},{\"x\":67.81864929199219,\"y\":114.45559692382812},{\"x\":67.91155242919922,\"y\":114.45559692382812},{\"x\":68.00445556640625,\"y\":114.45559692382812},{\"x\":68.09735870361328,\"y\":114.45559692382812},{\"x\":68.19026184082031,\"y\":114.45559692382812},{\"x\":68.28316497802734,\"y\":114.45559692382812},{\"x\":68.37606811523438,\"y\":114.45559692382812},{\"x\":68.4689712524414,\"y\":114.45559692382812},{\"x\":68.56187438964844,\"y\":114.45559692382812},{\"x\":68.65476989746094,\"y\":114.45559692382812},{\"x\":68.74767303466797,\"y\":114.45559692382812},{\"x\":68.840576171875,\"y\":114.45559692382812},{\"x\":68.93347930908203,\"y\":114.45559692382812},{\"x\":69.02638244628906,\"y\":114.45559692382812},{\"x\":69.1192855834961,\"y\":114.45559692382812},{\"x\":69.21218872070312,\"y\":114.45559692382812},{\"x\":69.30509185791016,\"y\":114.45559692382812},{\"x\":69.39799499511719,\"y\":114.45559692382812},{\"x\":69.49089050292969,\"y\":114.45559692382812},{\"x\":69.58379364013672,\"y\":114.45559692382812},{\"x\":69.67669677734375,\"y\":114.45559692382812},{\"x\":69.76959991455078,\"y\":114.45559692382812},{\"x\":69.86250305175781,\"y\":114.45559692382812},{\"x\":69.95540618896484,\"y\":114.45559692382812},{\"x\":70.04830932617188,\"y\":114.45559692382812},{\"x\":70.1412124633789,\"y\":114.45559692382812},{\"x\":70.2341079711914,\"y\":114.45559692382812},{\"x\":70.32701110839844,\"y\":114.45559692382812},{\"x\":70.41991424560547,\"y\":114.45559692382812},{\"x\":70.5128173828125,\"y\":114.45559692382812},{\"x\":70.60572052001953,\"y\":114.45559692382812},{\"x\":70.69862365722656,\"y\":114.45559692382812},{\"x\":70.7915267944336,\"y\":114.45559692382812},{\"x\":70.88442993164062,\"y\":114.45559692382812},{\"x\":70.97732543945312,\"y\":114.45559692382812},{\"x\":71.07022857666016,\"y\":114.45559692382812},{\"x\":71.16313171386719,\"y\":114.45559692382812},{\"x\":71.25603485107422,\"y\":114.45559692382812},{\"x\":71.34893798828125,\"y\":114.45559692382812}]", + "[{\"points\":[{\"x\":63.916760050380844,\"y\":114.45559357858895},{\"x\":71.34894145158792,\"y\":114.45559357858895}],\"color\":\"#1f77b4\",\"name\":\"Spline 1\"}]", + "[[{\"x\":63.9167594909668,\"y\":114.45559692382812},{\"x\":64.00965881347656,\"y\":114.45559692382812},{\"x\":64.1025619506836,\"y\":114.45559692382812},{\"x\":64.19546508789062,\"y\":114.45559692382812},{\"x\":64.28836822509766,\"y\":114.45559692382812},{\"x\":64.38127136230469,\"y\":114.45559692382812},{\"x\":64.47417449951172,\"y\":114.45559692382812},{\"x\":64.56707763671875,\"y\":114.45559692382812},{\"x\":64.65998077392578,\"y\":114.45559692382812},{\"x\":64.75287628173828,\"y\":114.45559692382812},{\"x\":64.84577941894531,\"y\":114.45559692382812},{\"x\":64.93868255615234,\"y\":114.45559692382812},{\"x\":65.03158569335938,\"y\":114.45559692382812},{\"x\":65.1244888305664,\"y\":114.45559692382812},{\"x\":65.21739196777344,\"y\":114.45559692382812},{\"x\":65.31029510498047,\"y\":114.45559692382812},{\"x\":65.4031982421875,\"y\":114.45559692382812},{\"x\":65.49609375,\"y\":114.45559692382812},{\"x\":65.58899688720703,\"y\":114.45559692382812},{\"x\":65.68190002441406,\"y\":114.45559692382812},{\"x\":65.7748031616211,\"y\":114.45559692382812},{\"x\":65.86770629882812,\"y\":114.45559692382812},{\"x\":65.96060943603516,\"y\":114.45559692382812},{\"x\":66.05351257324219,\"y\":114.45559692382812},{\"x\":66.14641571044922,\"y\":114.45559692382812},{\"x\":66.23931884765625,\"y\":114.45559692382812},{\"x\":66.33221435546875,\"y\":114.45559692382812},{\"x\":66.42511749267578,\"y\":114.45559692382812},{\"x\":66.51802062988281,\"y\":114.45559692382812},{\"x\":66.61092376708984,\"y\":114.45559692382812},{\"x\":66.70382690429688,\"y\":114.45559692382812},{\"x\":66.7967300415039,\"y\":114.45559692382812},{\"x\":66.88963317871094,\"y\":114.45559692382812},{\"x\":66.98253631591797,\"y\":114.45559692382812},{\"x\":67.075439453125,\"y\":114.45559692382812},{\"x\":67.1683349609375,\"y\":114.45559692382812},{\"x\":67.26123809814453,\"y\":114.45559692382812},{\"x\":67.35414123535156,\"y\":114.45559692382812},{\"x\":67.4470443725586,\"y\":114.45559692382812},{\"x\":67.53994750976562,\"y\":114.45559692382812},{\"x\":67.63285064697266,\"y\":114.45559692382812},{\"x\":67.72575378417969,\"y\":114.45559692382812},{\"x\":67.81864929199219,\"y\":114.45559692382812},{\"x\":67.91155242919922,\"y\":114.45559692382812},{\"x\":68.00445556640625,\"y\":114.45559692382812},{\"x\":68.09735870361328,\"y\":114.45559692382812},{\"x\":68.19026184082031,\"y\":114.45559692382812},{\"x\":68.28316497802734,\"y\":114.45559692382812},{\"x\":68.37606811523438,\"y\":114.45559692382812},{\"x\":68.4689712524414,\"y\":114.45559692382812},{\"x\":68.56187438964844,\"y\":114.45559692382812},{\"x\":68.65476989746094,\"y\":114.45559692382812},{\"x\":68.74767303466797,\"y\":114.45559692382812},{\"x\":68.840576171875,\"y\":114.45559692382812},{\"x\":68.93347930908203,\"y\":114.45559692382812},{\"x\":69.02638244628906,\"y\":114.45559692382812},{\"x\":69.1192855834961,\"y\":114.45559692382812},{\"x\":69.21218872070312,\"y\":114.45559692382812},{\"x\":69.30509185791016,\"y\":114.45559692382812},{\"x\":69.39799499511719,\"y\":114.45559692382812},{\"x\":69.49089050292969,\"y\":114.45559692382812},{\"x\":69.58379364013672,\"y\":114.45559692382812},{\"x\":69.67669677734375,\"y\":114.45559692382812},{\"x\":69.76959991455078,\"y\":114.45559692382812},{\"x\":69.86250305175781,\"y\":114.45559692382812},{\"x\":69.95540618896484,\"y\":114.45559692382812},{\"x\":70.04830932617188,\"y\":114.45559692382812},{\"x\":70.1412124633789,\"y\":114.45559692382812},{\"x\":70.2341079711914,\"y\":114.45559692382812},{\"x\":70.32701110839844,\"y\":114.45559692382812},{\"x\":70.41991424560547,\"y\":114.45559692382812},{\"x\":70.5128173828125,\"y\":114.45559692382812},{\"x\":70.60572052001953,\"y\":114.45559692382812},{\"x\":70.69862365722656,\"y\":114.45559692382812},{\"x\":70.7915267944336,\"y\":114.45559692382812},{\"x\":70.88442993164062,\"y\":114.45559692382812},{\"x\":70.97732543945312,\"y\":114.45559692382812},{\"x\":71.07022857666016,\"y\":114.45559692382812},{\"x\":71.16313171386719,\"y\":114.45559692382812},{\"x\":71.25603485107422,\"y\":114.45559692382812},{\"x\":71.34893798828125,\"y\":114.45559692382812}]]", 600, 382, 81, @@ -499,21 +459,20 @@ "list", 0, 1, - null, - null, + "", null ] }, { "id": 44, "type": "VHS_VideoCombine", - "pos": { - "0": 2241.138427734375, - "1": 1051.91943359375 - }, + "pos": [ + 2241.138427734375, + 1051.91943359375 + ], "size": [ 530, - 650.4 + 650.4000244140625 ], "flags": {}, "order": 12, @@ -521,36 +480,35 @@ "inputs": [ { "name": "images", + "shape": 7, "type": "IMAGE", - "link": 237, - "shape": 7 + "link": 237 }, { "name": "audio", + "shape": 7, "type": "AUDIO", - "link": null, - "shape": 7 + "link": null }, { "name": "meta_batch", + "shape": 7, "type": "VHS_BatchManager", - "link": null, - "shape": 7 + "link": null }, { "name": "vae", + "shape": 7, "type": "VAE", - "link": null, - "shape": 7 + "link": null } ], "outputs": [ { "name": "Filenames", "type": "VHS_FILENAMES", - "links": null, "slot_index": 0, - "shape": 3 + "links": null } ], "title": "Trajectory Outputs", @@ -584,10 +542,10 @@ { "id": 115, "type": "VHS_VideoCombine", - "pos": { - "0": 2819, - "1": 1056 - }, + "pos": [ + 2819, + 1056 + ], "size": [ 530, 310 @@ -598,36 +556,35 @@ "inputs": [ { "name": "images", + "shape": 7, "type": "IMAGE", - "link": 264, - "shape": 7 + "link": 264 }, { "name": "audio", + "shape": 7, "type": "AUDIO", - "link": null, - "shape": 7 + "link": null }, { "name": "meta_batch", + "shape": 7, "type": "VHS_BatchManager", - "link": null, - "shape": 7 + "link": null }, { "name": "vae", + "shape": 7, "type": "VAE", - "link": null, - "shape": 7 + "link": null } ], "outputs": [ { "name": "Filenames", "type": "VHS_FILENAMES", - "links": null, "slot_index": 0, - "shape": 3 + "links": null } ], "title": "Video with Trajectory Outputs", @@ -661,13 +618,13 @@ { "id": 126, "type": "WanFunV2VSampler", - "pos": { - "0": 902, - "1": 60 - }, + "pos": [ + 902, + 60 + ], "size": [ 428.4000244140625, - 486 + 506 ], "flags": {}, "order": 13, @@ -690,42 +647,39 @@ }, { "name": "validation_video", + "shape": 7, "type": "IMAGE", - "link": null, - "shape": 7 + "link": null }, { "name": "control_video", + "shape": 7, "type": "IMAGE", - "link": 284, - "shape": 7 + "link": 284 }, { "name": "start_image", + "shape": 7, "type": "IMAGE", - "link": 285, - "shape": 7 + "link": 285 }, { "name": "ref_image", + "shape": 7, "type": "IMAGE", - "link": null, - "shape": 7 - }, - { - "name": "riflex_k", - "type": "RIFLEXT_ARGS", - "link": null, - "shape": 7 + "link": null }, { "name": "camera_conditions", + "shape": 7, "type": "STRING", - "link": null, - "widget": { - "name": "camera_conditions" - }, - "shape": 7 + "link": null + }, + { + "name": "riflex_k", + "shape": 7, + "type": "RIFLEXT_ARGS", + "link": null } ], "outputs": [ @@ -748,7 +702,7 @@ "fixed", 50, 6, - 1.0, + 1, "Flow", 0.1, true, @@ -760,10 +714,10 @@ { "id": 106, "type": "VHS_VideoCombine", - "pos": { - "0": 1390, - "1": 61 - }, + "pos": [ + 1390, + 61 + ], "size": [ 390, 310 @@ -773,42 +727,40 @@ "mode": 0, "inputs": [ { - "name": "images", - "type": "IMAGE", - "link": 286, - "slot_index": 0, "label": "图像", - "shape": 7 + "name": "images", + "shape": 7, + "type": "IMAGE", + "link": 286 }, { - "name": "audio", - "type": "AUDIO", - "link": null, "label": "音频", - "shape": 7 + "name": "audio", + "shape": 7, + "type": "AUDIO", + "link": null }, { - "name": "meta_batch", - "type": "VHS_BatchManager", - "link": null, "label": "批次管理", - "shape": 7 + "name": "meta_batch", + "shape": 7, + "type": "VHS_BatchManager", + "link": null }, { "name": "vae", + "shape": 7, "type": "VAE", - "link": null, - "shape": 7 + "link": null } ], "outputs": [ { + "label": "文件名", "name": "Filenames", "type": "VHS_FILENAMES", - "links": null, "slot_index": 0, - "shape": 3, - "label": "文件名" + "links": null } ], "properties": { @@ -840,26 +792,26 @@ { "id": 123, "type": "LoadWanFunModel", - "pos": { - "0": 281, - "1": 251 - }, - "size": { - "0": 315, - "1": 154 - }, + "pos": [ + 281, + 251 + ], + "size": [ + 315, + 154 + ], "flags": {}, - "order": 9, + "order": 8, "mode": 0, "inputs": [], "outputs": [ { "name": "funmodels", "type": "FunModels", + "slot_index": 0, "links": [ 281 - ], - "slot_index": 0 + ] } ], "properties": { @@ -872,6 +824,31 @@ "wan2.1/wan_civitai.yaml", "bf16" ] + }, + { + "id": 112, + "type": "Note", + "pos": [ + -203, + 252 + ], + "size": [ + 427.074951171875, + 143.9142608642578 + ], + "flags": {}, + "order": 9, + "mode": 0, + "inputs": [], + "outputs": [], + "properties": { + "text": "" + }, + "widgets_values": [ + "When using the 1.3B model, you can set GPU_memory_mode to model_cpu_offload for faster generation. When using the 14B model, you can use sequential_cpu_offload to save GPU memory during generation.\n(在使用1.3B模型时,可以设置GPU_memory_mode为model_cpu_offload进行更快速度的生成,在使用14B模型时,可以使用sequential_cpu_offload节省显存,进行生成。)" + ], + "color": "#432", + "bgcolor": "#653" } ], "links": [ @@ -888,7 +865,7 @@ 97, 0, 95, - 0, + 1, "MASK" ], [ @@ -928,7 +905,7 @@ 118, 0, 95, - 1, + 0, "STRING" ], [ @@ -990,7 +967,8 @@ ], "groups": [ { - "title": "Load EasyAnimate", + "id": 1, + "title": "Load Model", "bounding": [ 189, 160, @@ -1002,6 +980,7 @@ "flags": {} }, { + "id": 2, "title": "Prompts", "bounding": [ 191, @@ -1014,6 +993,7 @@ "flags": {} }, { + "id": 3, "title": "First Image of Trajectory", "bounding": [ 191, @@ -1026,6 +1006,7 @@ "flags": {} }, { + "id": 4, "title": "Generate Control Video", "bounding": [ 786, @@ -1041,18 +1022,19 @@ "config": {}, "extra": { "ds": { - "scale": 0.6830134553650709, + "scale": 0.7513148009015782, "offset": [ - 198.78755142053404, - 143.5866129875247 + 119.78785424088606, + -47.62140429778734 ] }, "node_versions": { - "comfy-core": "v0.2.7-3-g8afb97c", - "ComfyUI-KJNodes": "4c5c26a2c91de356212419ac8bc7fcf9869527e9", - "CogVideoX-Fun": "717f0629175ad192927dc51ec95c4376816a4212", + "comfy-core": "0.3.44", + "ComfyUI-KJNodes": "ff49e1b01f10a14496b08e21bb89b64d2b15f333", + "CogVideoX-Fun": "e935795d7684aaba72be65d2ad5ef580c0526f7a", "ComfyUI-VideoHelperSuite": "70faa9bcef65932ab72e7404d6373fb300013a2e" - } + }, + "frontendVersion": "1.21.3" }, "version": 0.4 } \ No newline at end of file diff --git a/comfyui/wan2_1_fun/v1/wan2.1_fun_workflow_v2v_control.json b/comfyui/wan2_1_fun/v1/wan2.1_fun_workflow_v2v_control.json index 36c7568..e7630e5 100755 --- a/comfyui/wan2_1_fun/v1/wan2.1_fun_workflow_v2v_control.json +++ b/comfyui/wan2_1_fun/v1/wan2.1_fun_workflow_v2v_control.json @@ -523,7 +523,7 @@ "flags": {} }, { - "title": "Load EasyAnimate", + "title": "Load Model", "bounding": [ 218, -387, diff --git a/comfyui/wan2_1_fun/v1/wan2.1_fun_workflow_v2v_control_canny.json b/comfyui/wan2_1_fun/v1/wan2.1_fun_workflow_v2v_control_canny.json index a6c0fbb..37f6e6d 100755 --- a/comfyui/wan2_1_fun/v1/wan2.1_fun_workflow_v2v_control_canny.json +++ b/comfyui/wan2_1_fun/v1/wan2.1_fun_workflow_v2v_control_canny.json @@ -653,7 +653,7 @@ "flags": {} }, { - "title": "Load EasyAnimate", + "title": "Load Model", "bounding": [ 218, -387, diff --git a/comfyui/wan2_1_fun/v1/wan2.1_fun_workflow_v2v_control_depth.json b/comfyui/wan2_1_fun/v1/wan2.1_fun_workflow_v2v_control_depth.json index f8a8972..5cdd1a7 100755 --- a/comfyui/wan2_1_fun/v1/wan2.1_fun_workflow_v2v_control_depth.json +++ b/comfyui/wan2_1_fun/v1/wan2.1_fun_workflow_v2v_control_depth.json @@ -651,7 +651,7 @@ "flags": {} }, { - "title": "Load EasyAnimate", + "title": "Load Model", "bounding": [ 218, -387, diff --git a/comfyui/wan2_1_fun/v1/wan2.1_fun_workflow_v2v_control_depth_ref.json b/comfyui/wan2_1_fun/v1/wan2.1_fun_workflow_v2v_control_depth_ref.json index af196c9..3063d4a 100755 --- a/comfyui/wan2_1_fun/v1/wan2.1_fun_workflow_v2v_control_depth_ref.json +++ b/comfyui/wan2_1_fun/v1/wan2.1_fun_workflow_v2v_control_depth_ref.json @@ -696,7 +696,7 @@ "flags": {} }, { - "title": "Load EasyAnimate", + "title": "Load Model", "bounding": [ 218, -387, diff --git a/comfyui/wan2_1_fun/v1/wan2.1_fun_workflow_v2v_control_pose.json b/comfyui/wan2_1_fun/v1/wan2.1_fun_workflow_v2v_control_pose.json index 777dcb3..cd2a01c 100755 --- a/comfyui/wan2_1_fun/v1/wan2.1_fun_workflow_v2v_control_pose.json +++ b/comfyui/wan2_1_fun/v1/wan2.1_fun_workflow_v2v_control_pose.json @@ -651,7 +651,7 @@ "flags": {} }, { - "title": "Load EasyAnimate", + "title": "Load Model", "bounding": [ 218, -387, diff --git a/comfyui/wan2_1_fun/v1/wan2.1_fun_workflow_v2v_control_pose_ref.json b/comfyui/wan2_1_fun/v1/wan2.1_fun_workflow_v2v_control_pose_ref.json index 7e19323..f080828 100755 --- a/comfyui/wan2_1_fun/v1/wan2.1_fun_workflow_v2v_control_pose_ref.json +++ b/comfyui/wan2_1_fun/v1/wan2.1_fun_workflow_v2v_control_pose_ref.json @@ -696,7 +696,7 @@ "flags": {} }, { - "title": "Load EasyAnimate", + "title": "Load Model", "bounding": [ 218, -387, diff --git a/comfyui/wan2_1_fun/v1/wan2.1_fun_workflow_v2v_control_ref.json b/comfyui/wan2_1_fun/v1/wan2.1_fun_workflow_v2v_control_ref.json index 9e5858e..3b9e1e0 100755 --- a/comfyui/wan2_1_fun/v1/wan2.1_fun_workflow_v2v_control_ref.json +++ b/comfyui/wan2_1_fun/v1/wan2.1_fun_workflow_v2v_control_ref.json @@ -568,7 +568,7 @@ "flags": {} }, { - "title": "Load EasyAnimate", + "title": "Load Model", "bounding": [ 218, -387, diff --git a/comfyui/wan2_2/nodes.py b/comfyui/wan2_2/nodes.py index 423fd13..7de8831 100755 --- a/comfyui/wan2_2/nodes.py +++ b/comfyui/wan2_2/nodes.py @@ -18,9 +18,9 @@ from PIL import Image from ...videox_fun.data.bucket_sampler import (ASPECT_RATIO_512, get_closest_ratio) -from ...videox_fun.models import (AutoencoderKLWan, AutoTokenizer, CLIPModel, +from ...videox_fun.models import (AutoencoderKLWan, AutoencoderKLWan3_8, AutoTokenizer, CLIPModel, WanT5EncoderModel, Wan2_2Transformer3DModel) -from ...videox_fun.pipeline import Wan2_2I2VPipeline, Wan2_2Pipeline +from ...videox_fun.pipeline import Wan2_2I2VPipeline, Wan2_2Pipeline, Wan2_2TI2VPipeline from ...videox_fun.ui.controller import all_cheduler_dict from ...videox_fun.utils.fp8_optimization import ( convert_model_weight_to_float8, convert_weight_dtype_wrapper, replace_parameters_by_name) @@ -54,6 +54,7 @@ class LoadWan2_2Model: [ 'Wan2.2-T2V-A14B', 'Wan2.2-I2V-A14B', + 'Wan2.2-TI2V-5B', ], { "default": 'Wan2.2-T2V-A14B', @@ -69,6 +70,7 @@ class LoadWan2_2Model: [ "wan2.2/wan_civitai_t2v.yaml", "wan2.2/wan_civitai_i2v.yaml", + "wan2.2/wan_civitai_5b.yaml", ], { "default": "wan2.2/wan_civitai_t2v.yaml", @@ -131,7 +133,12 @@ class LoadWan2_2Model: print(f"- {os.path.join(eas_cache_dir, folder)}") raise ValueError("Please download Fun model") - vae = AutoencoderKLWan.from_pretrained( + # Get Vae + Choosen_AutoencoderKL = { + "AutoencoderKLWan": AutoencoderKLWan, + "AutoencoderKLWan3_8": AutoencoderKLWan3_8 + }[config['vae_kwargs'].get('vae_type', 'AutoencoderKLWan')] + vae = Choosen_AutoencoderKL.from_pretrained( os.path.join(model_name, config['vae_kwargs'].get('vae_subpath', 'vae')), additional_kwargs=OmegaConf.to_container(config['vae_kwargs']), ).to(weight_dtype) @@ -153,13 +160,15 @@ class LoadWan2_2Model: low_cpu_mem_usage=True, torch_dtype=weight_dtype, ) - - transformer_2 = Wan2_2Transformer3DModel.from_pretrained( - os.path.join(model_name, config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer')), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - low_cpu_mem_usage=True, - torch_dtype=weight_dtype, - ) + if config['transformer_additional_kwargs'].get('transformer_combination_type', 'single') == "moe": + transformer_2 = Wan2_2Transformer3DModel.from_pretrained( + os.path.join(model_name, config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer')), + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), + low_cpu_mem_usage=True, + torch_dtype=weight_dtype, + ) + else: + transformer_2 = None # Update pbar pbar.update(1) @@ -180,46 +189,59 @@ class LoadWan2_2Model: # Get pipeline model_type = "Inpaint" if model_type == "Inpaint": - if transformer.config.in_channels != vae.config.latent_channels: - pipeline = Wan2_2I2VPipeline( - transformer=transformer, - transformer_2=transformer_2, + if "wan_civitai_5b" in config_path: + pipeline = Wan2_2TI2VPipeline( vae=vae, tokenizer=tokenizer, text_encoder=text_encoder, + transformer=transformer, + transformer_2=transformer_2, scheduler=scheduler, ) else: - pipeline = Wan2_2Pipeline( - transformer=transformer, - transformer_2=transformer_2, - vae=vae, - tokenizer=tokenizer, - text_encoder=text_encoder, - scheduler=scheduler, - ) + if transformer.config.in_channels != vae.config.latent_channels: + pipeline = Wan2_2I2VPipeline( + transformer=transformer, + transformer_2=transformer_2, + vae=vae, + tokenizer=tokenizer, + text_encoder=text_encoder, + scheduler=scheduler, + ) + else: + pipeline = Wan2_2Pipeline( + transformer=transformer, + transformer_2=transformer_2, + vae=vae, + tokenizer=tokenizer, + text_encoder=text_encoder, + scheduler=scheduler, + ) else: raise ValueError(f"Model type {model_type} not supported") if GPU_memory_mode == "sequential_cpu_offload": replace_parameters_by_name(transformer, ["modulation",], device=device) - replace_parameters_by_name(transformer_2, ["modulation",], device=device) transformer.freqs = transformer.freqs.to(device=device) - transformer_2.freqs = transformer_2.freqs.to(device=device) + if transformer_2 is not None: + replace_parameters_by_name(transformer_2, ["modulation",], device=device) + transformer_2.freqs = transformer_2.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",], device=device) - convert_model_weight_to_float8(transformer_2, exclude_module_name=["modulation",], device=device) convert_weight_dtype_wrapper(transformer, weight_dtype) - convert_weight_dtype_wrapper(transformer_2, weight_dtype) + if transformer_2 is not None: + convert_model_weight_to_float8(transformer_2, exclude_module_name=["modulation",], device=device) + convert_weight_dtype_wrapper(transformer_2, 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",], device=device) - convert_model_weight_to_float8(transformer_2, exclude_module_name=["modulation",], device=device) convert_weight_dtype_wrapper(transformer, weight_dtype) - convert_weight_dtype_wrapper(transformer_2, weight_dtype) + if transformer_2 is not None: + convert_model_weight_to_float8(transformer_2, exclude_module_name=["modulation",], device=device) + convert_weight_dtype_wrapper(transformer_2, weight_dtype) pipeline.to(device=device) else: pipeline.to(device=device) @@ -366,14 +388,18 @@ class Wan2_2T2VSampler: pipeline.transformer.enable_teacache( coefficients, steps, teacache_threshold, num_skip_start_steps=num_skip_start_steps, offload=teacache_offload ) - pipeline.transformer_2.share_teacache(transformer=pipeline.transformer) + if pipeline.transformer_2 is not None: + pipeline.transformer_2.share_teacache(transformer=pipeline.transformer) else: pipeline.transformer.disable_teacache() + if pipeline.transformer_2 is not None: + pipeline.transformer_2.disable_teacache() if cfg_skip_ratio is not None: print(f"Enable cfg_skip_ratio {cfg_skip_ratio}.") pipeline.transformer.enable_cfg_skip(cfg_skip_ratio, steps) - pipeline.transformer_2.share_cfg_skip(transformer=pipeline.transformer) + if pipeline.transformer_2 is not None: + pipeline.transformer_2.share_cfg_skip(transformer=pipeline.transformer) generator= torch.Generator(device).manual_seed(seed) @@ -384,7 +410,8 @@ class Wan2_2T2VSampler: if riflex_k > 0: latent_frames = (video_length - 1) // pipeline.vae.config.temporal_compression_ratio + 1 pipeline.transformer.enable_riflex(k = riflex_k, L_test = latent_frames) - pipeline.transformer_2.enable_riflex(k = riflex_k, L_test = latent_frames) + if pipeline.transformer_2 is not None: + pipeline.transformer_2.enable_riflex(k = riflex_k, L_test = latent_frames) # Apply lora if funmodels.get("lora_cache", False): @@ -395,13 +422,6 @@ class Wan2_2T2VSampler: transformer_state_dict = pipeline.transformer.state_dict() for key in transformer_state_dict: transformer_cpu_cache[key] = transformer_state_dict[key].clone().cpu() - - # Save the original weights to cpu - if len(transformer_high_cpu_cache) == 0: - print('Save transformer high state_dict to cpu memory') - transformer_high_state_dict = pipeline.transformer_2.state_dict() - for key in transformer_high_state_dict: - transformer_high_cpu_cache[key] = transformer_high_state_dict[key].clone().cpu() lora_path_now = str(funmodels.get("loras", []) + funmodels.get("strength_model", [])) if lora_path_now != lora_path_before: @@ -411,31 +431,43 @@ class Wan2_2T2VSampler: for _lora_path, _lora_weight in zip(funmodels.get("loras", []), funmodels.get("strength_model", [])): pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype) - lora_high_path_now = str(funmodels.get("loras_high", []) + funmodels.get("strength_model", [])) - if lora_high_path_now != lora_high_path_before: - print('Merge Lora High with Cache') - lora_high_path_before = copy.deepcopy(lora_high_path_now) - pipeline.transformer_2.load_state_dict(transformer_cpu_cache) - for _lora_path, _lora_weight in zip(funmodels.get("loras_high", []), funmodels.get("strength_model", [])): - pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype, sub_transformer_name="transformer_2") + if pipeline.transformer_2 is not None: + # Save the original weights to cpu + if len(transformer_high_cpu_cache) == 0: + print('Save transformer high state_dict to cpu memory') + transformer_high_state_dict = pipeline.transformer_2.state_dict() + for key in transformer_high_state_dict: + transformer_high_cpu_cache[key] = transformer_high_state_dict[key].clone().cpu() + + lora_high_path_now = str(funmodels.get("loras_high", []) + funmodels.get("strength_model", [])) + if lora_high_path_now != lora_high_path_before: + print('Merge Lora High with Cache') + lora_high_path_before = copy.deepcopy(lora_high_path_now) + pipeline.transformer_2.load_state_dict(transformer_cpu_cache) + for _lora_path, _lora_weight in zip(funmodels.get("loras_high", []), funmodels.get("strength_model", [])): + pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype, sub_transformer_name="transformer_2") else: + print('Merge Lora') # Clear lora when switch from lora_cache=True to lora_cache=False. if len(transformer_cpu_cache) != 0: pipeline.transformer.load_state_dict(transformer_cpu_cache) transformer_cpu_cache = {} lora_path_before = "" gc.collect() - # Clear lora when switch from lora_cache=True to lora_cache=False. - if len(transformer_high_cpu_cache) != 0: - pipeline.transformer.load_state_dict(transformer_high_cpu_cache) - transformer_high_cpu_cache = {} - lora_high_path_before = "" - gc.collect() - print('Merge Lora') + for _lora_path, _lora_weight in zip(funmodels.get("loras", []), funmodels.get("strength_model", [])): pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype) - for _lora_path, _lora_weight in zip(funmodels.get("loras_high", []), funmodels.get("strength_model", [])): - pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype, sub_transformer_name="transformer_2") + + # Clear lora when switch from lora_cache=True to lora_cache=False. + if pipeline.transformer_2 is not None: + if len(transformer_high_cpu_cache) != 0: + pipeline.transformer_2.load_state_dict(transformer_high_cpu_cache) + transformer_high_cpu_cache = {} + lora_high_path_before = "" + gc.collect() + + for _lora_path, _lora_weight in zip(funmodels.get("loras_high", []), funmodels.get("strength_model", [])): + pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype, sub_transformer_name="transformer_2") sample = pipeline( prompt, @@ -455,8 +487,9 @@ class Wan2_2T2VSampler: print('Unmerge Lora') for _lora_path, _lora_weight in zip(funmodels.get("loras", []), funmodels.get("strength_model", [])): pipeline = unmerge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype) - for _lora_path, _lora_weight in zip(funmodels.get("loras_high", []), funmodels.get("strength_model", [])): - pipeline = unmerge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype, sub_transformer_name="transformer_2") + if pipeline.transformer_2 is not None: + for _lora_path, _lora_weight in zip(funmodels.get("loras_high", []), funmodels.get("strength_model", [])): + pipeline = unmerge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype, sub_transformer_name="transformer_2") return (videos,) @@ -541,14 +574,6 @@ class Wan2_2I2VSampler: mm.soft_empty_cache() gc.collect() - - start_img = [to_pil(_start_img) for _start_img in start_img] if start_img is not None else None - end_img = [to_pil(_end_img) for _end_img in end_img] if end_img is not None else None - # Count most suitable height and width - aspect_ratio_sample_size = {key : [x / 512 * base_resolution for x in ASPECT_RATIO_512[key]] for key in ASPECT_RATIO_512.keys()} - original_width, original_height = start_img[0].size if type(start_img) is list else Image.open(start_img).size - closest_size, closest_ratio = get_closest_ratio(original_height, original_width, ratios=aspect_ratio_sample_size) - height, width = [int(x / 16) * 16 for x in closest_size] # Get Pipeline pipeline = funmodels['pipeline'] @@ -556,6 +581,15 @@ class Wan2_2I2VSampler: config = funmodels['config'] weight_dtype = funmodels['dtype'] + start_img = [to_pil(_start_img) for _start_img in start_img] if start_img is not None else None + end_img = [to_pil(_end_img) for _end_img in end_img] if end_img is not None else None + # Count most suitable height and width + spatial_compression_ratio = pipeline.vae.config.spatial_compression_ratio if hasattr(pipeline.vae.config, "spatial_compression_ratio") else 8 + aspect_ratio_sample_size = {key : [x / 512 * base_resolution for x in ASPECT_RATIO_512[key]] for key in ASPECT_RATIO_512.keys()} + original_width, original_height = start_img[0].size if type(start_img) is list else Image.open(start_img).size + closest_size, closest_ratio = get_closest_ratio(original_height, original_width, ratios=aspect_ratio_sample_size) + height, width = [int(x / spatial_compression_ratio / 2) * spatial_compression_ratio * 2 for x in closest_size] + # Get boundary for wan boundary = config['transformer_additional_kwargs'].get('boundary', 0.900) @@ -567,12 +601,18 @@ class Wan2_2I2VSampler: pipeline.transformer.enable_teacache( coefficients, steps, teacache_threshold, num_skip_start_steps=num_skip_start_steps, offload=teacache_offload ) + if pipeline.transformer_2 is not None: + pipeline.transformer_2.share_teacache(transformer=pipeline.transformer) else: pipeline.transformer.disable_teacache() + if pipeline.transformer_2 is not None: + pipeline.transformer_2.disable_teacache() if cfg_skip_ratio is not None: print(f"Enable cfg_skip_ratio {cfg_skip_ratio}.") pipeline.transformer.enable_cfg_skip(cfg_skip_ratio, steps) + if pipeline.transformer_2 is not None: + pipeline.transformer_2.share_cfg_skip(transformer=pipeline.transformer) generator= torch.Generator(device).manual_seed(seed) @@ -583,7 +623,8 @@ class Wan2_2I2VSampler: if riflex_k > 0: latent_frames = (video_length - 1) // pipeline.vae.config.temporal_compression_ratio + 1 pipeline.transformer.enable_riflex(k = riflex_k, L_test = latent_frames) - pipeline.transformer_2.enable_riflex(k = riflex_k, L_test = latent_frames) + if pipeline.transformer_2 is not None: + pipeline.transformer_2.enable_riflex(k = riflex_k, L_test = latent_frames) # Apply lora if funmodels.get("lora_cache", False): @@ -594,13 +635,6 @@ class Wan2_2I2VSampler: transformer_state_dict = pipeline.transformer.state_dict() for key in transformer_state_dict: transformer_cpu_cache[key] = transformer_state_dict[key].clone().cpu() - - # Save the original weights to cpu - if len(transformer_high_cpu_cache) == 0: - print('Save transformer high state_dict to cpu memory') - transformer_high_state_dict = pipeline.transformer_2.state_dict() - for key in transformer_high_state_dict: - transformer_high_cpu_cache[key] = transformer_high_state_dict[key].clone().cpu() lora_path_now = str(funmodels.get("loras", []) + funmodels.get("strength_model", [])) if lora_path_now != lora_path_before: @@ -610,31 +644,43 @@ class Wan2_2I2VSampler: for _lora_path, _lora_weight in zip(funmodels.get("loras", []), funmodels.get("strength_model", [])): pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype) - lora_high_path_now = str(funmodels.get("loras_high", []) + funmodels.get("strength_model", [])) - if lora_high_path_now != lora_high_path_before: - print('Merge Lora High with Cache') - lora_high_path_before = copy.deepcopy(lora_high_path_now) - pipeline.transformer_2.load_state_dict(transformer_cpu_cache) - for _lora_path, _lora_weight in zip(funmodels.get("loras_high", []), funmodels.get("strength_model", [])): - pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype, sub_transformer_name="transformer_2") + if pipeline.transformer_2 is not None: + # Save the original weights to cpu + if len(transformer_high_cpu_cache) == 0: + print('Save transformer high state_dict to cpu memory') + transformer_high_state_dict = pipeline.transformer_2.state_dict() + for key in transformer_high_state_dict: + transformer_high_cpu_cache[key] = transformer_high_state_dict[key].clone().cpu() + + lora_high_path_now = str(funmodels.get("loras_high", []) + funmodels.get("strength_model", [])) + if lora_high_path_now != lora_high_path_before: + print('Merge Lora High with Cache') + lora_high_path_before = copy.deepcopy(lora_high_path_now) + pipeline.transformer_2.load_state_dict(transformer_cpu_cache) + for _lora_path, _lora_weight in zip(funmodels.get("loras_high", []), funmodels.get("strength_model", [])): + pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype, sub_transformer_name="transformer_2") else: + print('Merge Lora') # Clear lora when switch from lora_cache=True to lora_cache=False. if len(transformer_cpu_cache) != 0: pipeline.transformer.load_state_dict(transformer_cpu_cache) transformer_cpu_cache = {} lora_path_before = "" gc.collect() - # Clear lora when switch from lora_cache=True to lora_cache=False. - if len(transformer_high_cpu_cache) != 0: - pipeline.transformer.load_state_dict(transformer_high_cpu_cache) - transformer_high_cpu_cache = {} - lora_high_path_before = "" - gc.collect() - print('Merge Lora') + for _lora_path, _lora_weight in zip(funmodels.get("loras", []), funmodels.get("strength_model", [])): pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype) - for _lora_path, _lora_weight in zip(funmodels.get("loras_high", []), funmodels.get("strength_model", [])): - pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype, sub_transformer_name="transformer_2") + + # Clear lora when switch from lora_cache=True to lora_cache=False. + if pipeline.transformer_2 is not None: + if len(transformer_high_cpu_cache) != 0: + pipeline.transformer_2.load_state_dict(transformer_high_cpu_cache) + transformer_high_cpu_cache = {} + lora_high_path_before = "" + gc.collect() + + for _lora_path, _lora_weight in zip(funmodels.get("loras_high", []), funmodels.get("strength_model", [])): + pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype, sub_transformer_name="transformer_2") sample = pipeline( prompt, @@ -657,7 +703,7 @@ class Wan2_2I2VSampler: print('Unmerge Lora') for _lora_path, _lora_weight in zip(funmodels.get("loras", []), funmodels.get("strength_model", [])): pipeline = unmerge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype) - for _lora_path, _lora_weight in zip(funmodels.get("loras_high", []), funmodels.get("strength_model", [])): - pipeline = unmerge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype, sub_transformer_name="transformer_2") - return (videos,) - + if pipeline.transformer_2 is not None: + for _lora_path, _lora_weight in zip(funmodels.get("loras_high", []), funmodels.get("strength_model", [])): + pipeline = unmerge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype, sub_transformer_name="transformer_2") + return (videos,) \ No newline at end of file diff --git a/comfyui/wan2_2/v1/wan2.2_workflow_i2v_5b.json b/comfyui/wan2_2/v1/wan2.2_workflow_i2v_5b.json new file mode 100644 index 0000000..661c7ea --- /dev/null +++ b/comfyui/wan2_2/v1/wan2.2_workflow_i2v_5b.json @@ -0,0 +1,471 @@ +{ + "id": "ca87b2cd-bd4a-4f31-82ec-e5e028f8848c", + "revision": 0, + "last_node_id": 103, + "last_link_id": 80, + "nodes": [ + { + "id": 87, + "type": "LoadImage", + "pos": [ + 306, + 495 + ], + "size": [ + 315, + 314 + ], + "flags": {}, + "order": 0, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "links": [ + 79 + ] + }, + { + "name": "MASK", + "type": "MASK", + "links": null + } + ], + "properties": { + "Node name for S&R": "LoadImage" + }, + "widgets_values": [ + "6.png", + "image" + ] + }, + { + "id": 78, + "type": "Note", + "pos": [ + 18, + -46 + ], + "size": [ + 210, + 88 + ], + "flags": {}, + "order": 1, + "mode": 0, + "inputs": [], + "outputs": [], + "properties": { + "text": "" + }, + "widgets_values": [ + "You can write prompt here\n(你可以在此填写提示词)" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 95, + "type": "Note", + "pos": [ + 34, + 550 + ], + "size": [ + 210, + 88 + ], + "flags": {}, + "order": 2, + "mode": 0, + "inputs": [], + "outputs": [], + "properties": { + "text": "" + }, + "widgets_values": [ + "You can upload image here\n(你可以在此上传图片)" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 73, + "type": "FunTextBox", + "pos": [ + 250, + 160 + ], + "size": [ + 383.7149963378906, + 183.83506774902344 + ], + "flags": {}, + "order": 3, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "prompt", + "type": "STRING_PROMPT", + "slot_index": 0, + "links": [ + 78 + ] + } + ], + "title": "Negtive Prompt(反向提示词)", + "properties": { + "Node name for S&R": "FunTextBox" + }, + "widgets_values": [ + "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走" + ] + }, + { + "id": 80, + "type": "Note", + "pos": [ + -75, + -297 + ], + "size": [ + 350.7127990722656, + 125.54820251464844 + ], + "flags": {}, + "order": 4, + "mode": 0, + "inputs": [], + "outputs": [], + "properties": { + "text": "" + }, + "widgets_values": [ + "When using the 14B model, you can use sequential_cpu_offload to save GPU memory during generation.\n(在使用14B模型时,可以使用sequential_cpu_offload节省显存,进行生成。)" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 17, + "type": "VHS_VideoCombine", + "pos": [ + 1277, + -70 + ], + "size": [ + 390, + 577.7142944335938 + ], + "flags": {}, + "order": 8, + "mode": 0, + "inputs": [ + { + "label": "图像", + "name": "images", + "shape": 7, + "type": "IMAGE", + "link": 80 + }, + { + "label": "音频", + "name": "audio", + "shape": 7, + "type": "AUDIO", + "link": null + }, + { + "label": "批次管理", + "name": "meta_batch", + "shape": 7, + "type": "VHS_BatchManager", + "link": null + }, + { + "name": "vae", + "shape": 7, + "type": "VAE", + "link": null + } + ], + "outputs": [ + { + "label": "文件名", + "name": "Filenames", + "type": "VHS_FILENAMES", + "slot_index": 0, + "links": null + } + ], + "properties": { + "Node name for S&R": "VHS_VideoCombine" + }, + "widgets_values": { + "frame_rate": 16, + "loop_count": 0, + "filename_prefix": "Fun", + "format": "video/h264-mp4", + "pix_fmt": "yuv420p", + "crf": 22, + "save_metadata": true, + "pingpong": false, + "save_output": true, + "videopreview": { + "hidden": false, + "paused": false, + "params": { + "filename": "Fun_00107.mp4", + "subfolder": "", + "type": "output", + "format": "video/h264-mp4", + "frame_rate": 16 + } + } + } + }, + { + "id": 99, + "type": "LoadWan2_2Model", + "pos": [ + 347.27996826171875, + -299.0150146484375 + ], + "size": [ + 276.705078125, + 130 + ], + "flags": {}, + "order": 5, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "funmodels", + "type": "FunModels", + "links": [ + 76 + ] + } + ], + "properties": { + "Node name for S&R": "LoadWan2_2Model" + }, + "widgets_values": [ + "Wan2.2-TI2V-5B", + "sequential_cpu_offload", + "wan2.2/wan_civitai_5b.yaml", + "bf16" + ] + }, + { + "id": 103, + "type": "Wan2_2I2VSampler", + "pos": [ + 819.5313110351562, + -63.450443267822266 + ], + "size": [ + 325.5747985839844, + 402 + ], + "flags": {}, + "order": 7, + "mode": 0, + "inputs": [ + { + "name": "funmodels", + "type": "FunModels", + "link": 76 + }, + { + "name": "prompt", + "type": "STRING_PROMPT", + "link": 77 + }, + { + "name": "negative_prompt", + "type": "STRING_PROMPT", + "link": 78 + }, + { + "name": "start_img", + "shape": 7, + "type": "IMAGE", + "link": 79 + }, + { + "name": "riflex_k", + "shape": 7, + "type": "RIFLEXT_ARGS", + "link": null + } + ], + "outputs": [ + { + "name": "images", + "type": "IMAGE", + "links": [ + 80 + ] + } + ], + "properties": { + "Node name for S&R": "Wan2_2I2VSampler" + }, + "widgets_values": [ + 81, + 960, + 43, + "fixed", + 50, + 6, + "Flow", + 0.1, + true, + 5, + true, + 0 + ] + }, + { + "id": 75, + "type": "FunTextBox", + "pos": [ + 250, + -50 + ], + "size": [ + 383.54010009765625, + 156.71620178222656 + ], + "flags": {}, + "order": 6, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "prompt", + "type": "STRING_PROMPT", + "slot_index": 0, + "links": [ + 77 + ] + } + ], + "title": "Positive Prompt(正向提示词)", + "properties": { + "Node name for S&R": "FunTextBox" + }, + "widgets_values": [ + "fireworks display over night city. The video is of high quality, and the view is very clear. High quality, masterpiece, best quality, highres, ultra-detailed, fantastic." + ] + } + ], + "links": [ + [ + 76, + 99, + 0, + 103, + 0, + "FunModels" + ], + [ + 77, + 75, + 0, + 103, + 1, + "STRING_PROMPT" + ], + [ + 78, + 73, + 0, + 103, + 2, + "STRING_PROMPT" + ], + [ + 79, + 87, + 0, + 103, + 3, + "IMAGE" + ], + [ + 80, + 103, + 0, + 17, + 0, + "IMAGE" + ] + ], + "groups": [ + { + "id": 1, + "title": "Load Model", + "bounding": [ + 220, + -380, + 472, + 232 + ], + "color": "#b06634", + "font_size": 24, + "flags": {} + }, + { + "id": 2, + "title": "Prompts", + "bounding": [ + 218, + -127, + 450, + 483 + ], + "color": "#3f789e", + "font_size": 24, + "flags": {} + }, + { + "id": 3, + "title": "Group", + "bounding": [ + 220, + 409, + 458, + 436 + ], + "color": "#a1309b", + "font_size": 24, + "flags": {} + } + ], + "config": {}, + "extra": { + "ds": { + "scale": 0.6830134553650705, + "offset": [ + 282.91743113564746, + 433.6498523638886 + ] + }, + "frontendVersion": "1.21.3", + "workspace_info": { + "id": "776b62b4-bd17-4ed3-9923-b7aad000b1ea" + }, + "node_versions": { + "comfy-core": "0.3.44", + "CogVideoX-Fun": "5f2a55692d0834e00d477939b818b9bb63c06535", + "ComfyUI-VideoHelperSuite": "70faa9bcef65932ab72e7404d6373fb300013a2e" + } + }, + "version": 0.4 +} \ No newline at end of file diff --git a/comfyui/wan2_2/v1/wan2.2_workflow_t2v_5b.json b/comfyui/wan2_2/v1/wan2.2_workflow_t2v_5b.json new file mode 100644 index 0000000..f1ff90b --- /dev/null +++ b/comfyui/wan2_2/v1/wan2.2_workflow_t2v_5b.json @@ -0,0 +1,383 @@ +{ + "id": "ca87b2cd-bd4a-4f31-82ec-e5e028f8848c", + "revision": 0, + "last_node_id": 105, + "last_link_id": 84, + "nodes": [ + { + "id": 78, + "type": "Note", + "pos": [ + 18, + -46 + ], + "size": [ + 210, + 88 + ], + "flags": {}, + "order": 0, + "mode": 0, + "inputs": [], + "outputs": [], + "properties": { + "text": "" + }, + "widgets_values": [ + "You can write prompt here\n(你可以在此填写提示词)" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 73, + "type": "FunTextBox", + "pos": [ + 250, + 160 + ], + "size": [ + 383.7149963378906, + 183.83506774902344 + ], + "flags": {}, + "order": 1, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "prompt", + "type": "STRING_PROMPT", + "slot_index": 0, + "links": [ + 83 + ] + } + ], + "title": "Negtive Prompt(反向提示词)", + "properties": { + "Node name for S&R": "FunTextBox" + }, + "widgets_values": [ + "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走" + ] + }, + { + "id": 80, + "type": "Note", + "pos": [ + -75, + -297 + ], + "size": [ + 350.7127990722656, + 125.54820251464844 + ], + "flags": {}, + "order": 2, + "mode": 0, + "inputs": [], + "outputs": [], + "properties": { + "text": "" + }, + "widgets_values": [ + "When using the 14B model, you can use sequential_cpu_offload to save GPU memory during generation.\n(在使用14B模型时,可以使用sequential_cpu_offload节省显存,进行生成。)" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 75, + "type": "FunTextBox", + "pos": [ + 250, + -50 + ], + "size": [ + 383.54010009765625, + 156.71620178222656 + ], + "flags": {}, + "order": 3, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "prompt", + "type": "STRING_PROMPT", + "slot_index": 0, + "links": [ + 82 + ] + } + ], + "title": "Positive Prompt(正向提示词)", + "properties": { + "Node name for S&R": "FunTextBox" + }, + "widgets_values": [ + "fireworks display over night city. The video is of high quality, and the view is very clear. High quality, masterpiece, best quality, highres, ultra-detailed, fantastic." + ] + }, + { + "id": 99, + "type": "LoadWan2_2Model", + "pos": [ + 347.27996826171875, + -299.0150146484375 + ], + "size": [ + 276.705078125, + 130 + ], + "flags": {}, + "order": 4, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "funmodels", + "type": "FunModels", + "links": [ + 81 + ] + } + ], + "properties": { + "Node name for S&R": "LoadWan2_2Model" + }, + "widgets_values": [ + "Wan2.2-TI2V-5B", + "sequential_cpu_offload", + "wan2.2/wan_civitai_5b.yaml", + "bf16" + ] + }, + { + "id": 17, + "type": "VHS_VideoCombine", + "pos": [ + 1277, + -70 + ], + "size": [ + 390, + 577.7142944335938 + ], + "flags": {}, + "order": 6, + "mode": 0, + "inputs": [ + { + "label": "图像", + "name": "images", + "shape": 7, + "type": "IMAGE", + "link": 84 + }, + { + "label": "音频", + "name": "audio", + "shape": 7, + "type": "AUDIO", + "link": null + }, + { + "label": "批次管理", + "name": "meta_batch", + "shape": 7, + "type": "VHS_BatchManager", + "link": null + }, + { + "name": "vae", + "shape": 7, + "type": "VAE", + "link": null + } + ], + "outputs": [ + { + "label": "文件名", + "name": "Filenames", + "type": "VHS_FILENAMES", + "slot_index": 0, + "links": null + } + ], + "properties": { + "Node name for S&R": "VHS_VideoCombine" + }, + "widgets_values": { + "frame_rate": 16, + "loop_count": 0, + "filename_prefix": "Fun", + "format": "video/h264-mp4", + "pix_fmt": "yuv420p", + "crf": 22, + "save_metadata": true, + "pingpong": false, + "save_output": true, + "videopreview": { + "hidden": false, + "paused": false, + "params": { + "filename": "Fun_00108.mp4", + "subfolder": "", + "type": "output", + "format": "video/h264-mp4", + "frame_rate": 16 + } + } + } + }, + { + "id": 105, + "type": "Wan2_2FunT2VSampler", + "pos": [ + 827.6821899414062, + -69.17436981201172 + ], + "size": [ + 340.3540954589844, + 430 + ], + "flags": {}, + "order": 5, + "mode": 0, + "inputs": [ + { + "name": "funmodels", + "type": "FunModels", + "link": 81 + }, + { + "name": "prompt", + "type": "STRING_PROMPT", + "link": 82 + }, + { + "name": "negative_prompt", + "type": "STRING_PROMPT", + "link": 83 + }, + { + "name": "riflex_k", + "shape": 7, + "type": "RIFLEXT_ARGS", + "link": null + } + ], + "outputs": [ + { + "name": "images", + "type": "IMAGE", + "links": [ + 84 + ] + } + ], + "properties": { + "Node name for S&R": "Wan2_2FunT2VSampler" + }, + "widgets_values": [ + 81, + 1280, + 704, + false, + 43, + "fixed", + 50, + 6, + "Flow", + 0.1, + true, + 5, + true, + 0 + ] + } + ], + "links": [ + [ + 81, + 99, + 0, + 105, + 0, + "FunModels" + ], + [ + 82, + 75, + 0, + 105, + 1, + "STRING_PROMPT" + ], + [ + 83, + 73, + 0, + 105, + 2, + "STRING_PROMPT" + ], + [ + 84, + 105, + 0, + 17, + 0, + "IMAGE" + ] + ], + "groups": [ + { + "id": 1, + "title": "Load Model", + "bounding": [ + 220, + -380, + 472, + 232 + ], + "color": "#b06634", + "font_size": 24, + "flags": {} + }, + { + "id": 2, + "title": "Prompts", + "bounding": [ + 218, + -127, + 450, + 483 + ], + "color": "#3f789e", + "font_size": 24, + "flags": {} + } + ], + "config": {}, + "extra": { + "ds": { + "scale": 0.6830134553650705, + "offset": [ + 313.88657762002225, + 431.4308258013887 + ] + }, + "frontendVersion": "1.21.3", + "workspace_info": { + "id": "776b62b4-bd17-4ed3-9923-b7aad000b1ea" + }, + "node_versions": { + "CogVideoX-Fun": "5f2a55692d0834e00d477939b818b9bb63c06535", + "ComfyUI-VideoHelperSuite": "70faa9bcef65932ab72e7404d6373fb300013a2e" + } + }, + "version": 0.4 +} \ No newline at end of file diff --git a/comfyui/wan2_2_fun/nodes.py b/comfyui/wan2_2_fun/nodes.py index a32955f..741578b 100755 --- a/comfyui/wan2_2_fun/nodes.py +++ b/comfyui/wan2_2_fun/nodes.py @@ -18,7 +18,7 @@ from PIL import Image from ...videox_fun.data.bucket_sampler import (ASPECT_RATIO_512, get_closest_ratio) -from ...videox_fun.models import (AutoencoderKLWan, AutoTokenizer, CLIPModel, +from ...videox_fun.models import (AutoencoderKLWan, AutoencoderKLWan3_8, AutoTokenizer, CLIPModel, WanT5EncoderModel, Wan2_2Transformer3DModel) from ...videox_fun.pipeline import Wan2_2FunInpaintPipeline, Wan2_2FunPipeline, Wan2_2FunControlPipeline from ...videox_fun.ui.controller import all_cheduler_dict @@ -48,6 +48,7 @@ class LoadWan2_2FunModel: [ 'Wan2.2-Fun-A14B-InP', 'Wan2.2-Fun-A14B-Control', + 'Wan2.2-Fun-A14B-Control-Camera', ], { "default": 'Wan2.2-Fun-A14B-InP', @@ -68,6 +69,7 @@ class LoadWan2_2FunModel: "config": ( [ "wan2.2/wan_civitai_i2v.yaml", + "wan2.2/wan_civitai_5b.yaml", ], { "default": "wan2.2/wan_civitai_i2v.yaml", @@ -129,7 +131,11 @@ class LoadWan2_2FunModel: print(f"- {os.path.join(eas_cache_dir, folder)}") raise ValueError("Please download Fun model") - vae = AutoencoderKLWan.from_pretrained( + Choosen_AutoencoderKL = { + "AutoencoderKLWan": AutoencoderKLWan, + "AutoencoderKLWan3_8": AutoencoderKLWan3_8 + }[config['vae_kwargs'].get('vae_type', 'AutoencoderKLWan')] + vae = Choosen_AutoencoderKL.from_pretrained( os.path.join(model_name, config['vae_kwargs'].get('vae_subpath', 'vae')), additional_kwargs=OmegaConf.to_container(config['vae_kwargs']), ).to(weight_dtype) @@ -357,21 +363,24 @@ class Wan2_2FunT2VSampler: # Load Sampler pipeline.scheduler = all_cheduler_dict[scheduler](**filter_kwargs(all_cheduler_dict[scheduler], OmegaConf.to_container(config['scheduler_kwargs']))) - coefficients = get_teacache_coefficients(model_name) if enable_teacache else None if coefficients is not None: print(f"Enable TeaCache with threshold {teacache_threshold} and skip the first {num_skip_start_steps} steps.") pipeline.transformer.enable_teacache( coefficients, steps, teacache_threshold, num_skip_start_steps=num_skip_start_steps, offload=teacache_offload ) - pipeline.transformer_2.share_teacache(transformer=pipeline.transformer) + if pipeline.transformer_2 is not None: + pipeline.transformer_2.share_teacache(transformer=pipeline.transformer) else: pipeline.transformer.disable_teacache() + if pipeline.transformer_2 is not None: + pipeline.transformer_2.disable_teacache() if cfg_skip_ratio is not None: print(f"Enable cfg_skip_ratio {cfg_skip_ratio}.") pipeline.transformer.enable_cfg_skip(cfg_skip_ratio, steps) - pipeline.transformer_2.share_cfg_skip(transformer=pipeline.transformer) + if pipeline.transformer_2 is not None: + pipeline.transformer_2.share_cfg_skip(transformer=pipeline.transformer) generator= torch.Generator(device).manual_seed(seed) @@ -382,7 +391,10 @@ class Wan2_2FunT2VSampler: if riflex_k > 0: latent_frames = (video_length - 1) // pipeline.vae.config.temporal_compression_ratio + 1 pipeline.transformer.enable_riflex(k = riflex_k, L_test = latent_frames) - pipeline.transformer_2.enable_riflex(k = riflex_k, L_test = latent_frames) + if pipeline.transformer_2 is not None: + pipeline.transformer_2.enable_riflex(k = riflex_k, L_test = latent_frames) + + input_video, input_video_mask, clip_image = get_image_to_video_latent(None, None, video_length=video_length, sample_size=(height, width)) # Apply lora if funmodels.get("lora_cache", False): @@ -393,13 +405,6 @@ class Wan2_2FunT2VSampler: transformer_state_dict = pipeline.transformer.state_dict() for key in transformer_state_dict: transformer_cpu_cache[key] = transformer_state_dict[key].clone().cpu() - - # Save the original weights to cpu - if len(transformer_high_cpu_cache) == 0: - print('Save transformer high state_dict to cpu memory') - transformer_high_state_dict = pipeline.transformer_2.state_dict() - for key in transformer_high_state_dict: - transformer_high_cpu_cache[key] = transformer_high_state_dict[key].clone().cpu() lora_path_now = str(funmodels.get("loras", []) + funmodels.get("strength_model", [])) if lora_path_now != lora_path_before: @@ -409,31 +414,43 @@ class Wan2_2FunT2VSampler: for _lora_path, _lora_weight in zip(funmodels.get("loras", []), funmodels.get("strength_model", [])): pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype) - lora_high_path_now = str(funmodels.get("loras_high", []) + funmodels.get("strength_model", [])) - if lora_high_path_now != lora_high_path_before: - print('Merge Lora High with Cache') - lora_high_path_before = copy.deepcopy(lora_high_path_now) - pipeline.transformer_2.load_state_dict(transformer_cpu_cache) - for _lora_path, _lora_weight in zip(funmodels.get("loras_high", []), funmodels.get("strength_model", [])): - pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype, sub_transformer_name="transformer_2") + if pipeline.transformer_2 is not None: + # Save the original weights to cpu + if len(transformer_high_cpu_cache) == 0: + print('Save transformer high state_dict to cpu memory') + transformer_high_state_dict = pipeline.transformer_2.state_dict() + for key in transformer_high_state_dict: + transformer_high_cpu_cache[key] = transformer_high_state_dict[key].clone().cpu() + + lora_high_path_now = str(funmodels.get("loras_high", []) + funmodels.get("strength_model", [])) + if lora_high_path_now != lora_high_path_before: + print('Merge Lora High with Cache') + lora_high_path_before = copy.deepcopy(lora_high_path_now) + pipeline.transformer_2.load_state_dict(transformer_cpu_cache) + for _lora_path, _lora_weight in zip(funmodels.get("loras_high", []), funmodels.get("strength_model", [])): + pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype, sub_transformer_name="transformer_2") else: + print('Merge Lora') # Clear lora when switch from lora_cache=True to lora_cache=False. if len(transformer_cpu_cache) != 0: pipeline.transformer.load_state_dict(transformer_cpu_cache) transformer_cpu_cache = {} lora_path_before = "" gc.collect() - # Clear lora when switch from lora_cache=True to lora_cache=False. - if len(transformer_high_cpu_cache) != 0: - pipeline.transformer.load_state_dict(transformer_high_cpu_cache) - transformer_high_cpu_cache = {} - lora_high_path_before = "" - gc.collect() - print('Merge Lora') + for _lora_path, _lora_weight in zip(funmodels.get("loras", []), funmodels.get("strength_model", [])): pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype) - for _lora_path, _lora_weight in zip(funmodels.get("loras_high", []), funmodels.get("strength_model", [])): - pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype, sub_transformer_name="transformer_2") + + # Clear lora when switch from lora_cache=True to lora_cache=False. + if pipeline.transformer_2 is not None: + if len(transformer_high_cpu_cache) != 0: + pipeline.transformer_2.load_state_dict(transformer_high_cpu_cache) + transformer_high_cpu_cache = {} + lora_high_path_before = "" + gc.collect() + + for _lora_path, _lora_weight in zip(funmodels.get("loras_high", []), funmodels.get("strength_model", [])): + pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype, sub_transformer_name="transformer_2") sample = pipeline( prompt, @@ -444,6 +461,9 @@ class Wan2_2FunT2VSampler: generator = generator, guidance_scale = cfg, num_inference_steps = steps, + + video = input_video, + mask_video = input_video_mask, boundary = boundary, comfyui_progressbar = True, ).videos @@ -453,8 +473,9 @@ class Wan2_2FunT2VSampler: print('Unmerge Lora') for _lora_path, _lora_weight in zip(funmodels.get("loras", []), funmodels.get("strength_model", [])): pipeline = unmerge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype) - for _lora_path, _lora_weight in zip(funmodels.get("loras_high", []), funmodels.get("strength_model", [])): - pipeline = unmerge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype, sub_transformer_name="transformer_2") + if pipeline.transformer_2 is not None: + for _lora_path, _lora_weight in zip(funmodels.get("loras_high", []), funmodels.get("strength_model", [])): + pipeline = unmerge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype, sub_transformer_name="transformer_2") return (videos,) @@ -541,6 +562,12 @@ class Wan2_2FunInpaintSampler: mm.soft_empty_cache() gc.collect() + + # Get Pipeline + pipeline = funmodels['pipeline'] + model_name = funmodels['model_name'] + config = funmodels['config'] + weight_dtype = funmodels['dtype'] start_img = [to_pil(_start_img) for _start_img in start_img] if start_img is not None else None end_img = [to_pil(_end_img) for _end_img in end_img] if end_img is not None else None @@ -549,33 +576,30 @@ class Wan2_2FunInpaintSampler: original_width, original_height = start_img[0].size if type(start_img) is list else Image.open(start_img).size closest_size, closest_ratio = get_closest_ratio(original_height, original_width, ratios=aspect_ratio_sample_size) height, width = [int(x / 16) * 16 for x in closest_size] - - # Get Pipeline - pipeline = funmodels['pipeline'] - model_name = funmodels['model_name'] - config = funmodels['config'] - weight_dtype = funmodels['dtype'] # Get boundary for wan boundary = config['transformer_additional_kwargs'].get('boundary', 0.900) # Load Sampler pipeline.scheduler = all_cheduler_dict[scheduler](**filter_kwargs(all_cheduler_dict[scheduler], OmegaConf.to_container(config['scheduler_kwargs']))) - coefficients = get_teacache_coefficients(model_name) if enable_teacache else None if coefficients is not None: print(f"Enable TeaCache with threshold {teacache_threshold} and skip the first {num_skip_start_steps} steps.") pipeline.transformer.enable_teacache( coefficients, steps, teacache_threshold, num_skip_start_steps=num_skip_start_steps, offload=teacache_offload ) - pipeline.transformer_2.share_teacache(transformer=pipeline.transformer) + if pipeline.transformer_2 is not None: + pipeline.transformer_2.share_teacache(transformer=pipeline.transformer) else: pipeline.transformer.disable_teacache() + if pipeline.transformer_2 is not None: + pipeline.transformer_2.disable_teacache() if cfg_skip_ratio is not None: print(f"Enable cfg_skip_ratio {cfg_skip_ratio}.") pipeline.transformer.enable_cfg_skip(cfg_skip_ratio, steps) - pipeline.transformer_2.share_cfg_skip(transformer=pipeline.transformer) + if pipeline.transformer_2 is not None: + pipeline.transformer_2.share_cfg_skip(transformer=pipeline.transformer) generator= torch.Generator(device).manual_seed(seed) @@ -585,7 +609,8 @@ class Wan2_2FunInpaintSampler: if riflex_k > 0: latent_frames = (video_length - 1) // pipeline.vae.config.temporal_compression_ratio + 1 pipeline.transformer.enable_riflex(k = riflex_k, L_test = latent_frames) - pipeline.transformer_2.enable_riflex(k = riflex_k, L_test = latent_frames) + if pipeline.transformer_2 is not None: + pipeline.transformer_2.enable_riflex(k = riflex_k, L_test = latent_frames) input_video, input_video_mask, clip_image = get_image_to_video_latent(start_img, end_img, video_length=video_length, sample_size=(height, width)) @@ -598,13 +623,6 @@ class Wan2_2FunInpaintSampler: transformer_state_dict = pipeline.transformer.state_dict() for key in transformer_state_dict: transformer_cpu_cache[key] = transformer_state_dict[key].clone().cpu() - - # Save the original weights to cpu - if len(transformer_high_cpu_cache) == 0: - print('Save transformer high state_dict to cpu memory') - transformer_high_state_dict = pipeline.transformer_2.state_dict() - for key in transformer_high_state_dict: - transformer_high_cpu_cache[key] = transformer_high_state_dict[key].clone().cpu() lora_path_now = str(funmodels.get("loras", []) + funmodels.get("strength_model", [])) if lora_path_now != lora_path_before: @@ -614,31 +632,43 @@ class Wan2_2FunInpaintSampler: for _lora_path, _lora_weight in zip(funmodels.get("loras", []), funmodels.get("strength_model", [])): pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype) - lora_high_path_now = str(funmodels.get("loras_high", []) + funmodels.get("strength_model", [])) - if lora_high_path_now != lora_high_path_before: - print('Merge Lora High with Cache') - lora_high_path_before = copy.deepcopy(lora_high_path_now) - pipeline.transformer_2.load_state_dict(transformer_cpu_cache) - for _lora_path, _lora_weight in zip(funmodels.get("loras_high", []), funmodels.get("strength_model", [])): - pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype, sub_transformer_name="transformer_2") + if pipeline.transformer_2 is not None: + # Save the original weights to cpu + if len(transformer_high_cpu_cache) == 0: + print('Save transformer high state_dict to cpu memory') + transformer_high_state_dict = pipeline.transformer_2.state_dict() + for key in transformer_high_state_dict: + transformer_high_cpu_cache[key] = transformer_high_state_dict[key].clone().cpu() + + lora_high_path_now = str(funmodels.get("loras_high", []) + funmodels.get("strength_model", [])) + if lora_high_path_now != lora_high_path_before: + print('Merge Lora High with Cache') + lora_high_path_before = copy.deepcopy(lora_high_path_now) + pipeline.transformer_2.load_state_dict(transformer_cpu_cache) + for _lora_path, _lora_weight in zip(funmodels.get("loras_high", []), funmodels.get("strength_model", [])): + pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype, sub_transformer_name="transformer_2") else: + print('Merge Lora') # Clear lora when switch from lora_cache=True to lora_cache=False. if len(transformer_cpu_cache) != 0: pipeline.transformer.load_state_dict(transformer_cpu_cache) transformer_cpu_cache = {} lora_path_before = "" gc.collect() - # Clear lora when switch from lora_cache=True to lora_cache=False. - if len(transformer_high_cpu_cache) != 0: - pipeline.transformer.load_state_dict(transformer_high_cpu_cache) - transformer_high_cpu_cache = {} - lora_high_path_before = "" - gc.collect() - print('Merge Lora') + for _lora_path, _lora_weight in zip(funmodels.get("loras", []), funmodels.get("strength_model", [])): pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype) - for _lora_path, _lora_weight in zip(funmodels.get("loras_high", []), funmodels.get("strength_model", [])): - pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype, sub_transformer_name="transformer_2") + + # Clear lora when switch from lora_cache=True to lora_cache=False. + if pipeline.transformer_2 is not None: + if len(transformer_high_cpu_cache) != 0: + pipeline.transformer_2.load_state_dict(transformer_high_cpu_cache) + transformer_high_cpu_cache = {} + lora_high_path_before = "" + gc.collect() + + for _lora_path, _lora_weight in zip(funmodels.get("loras_high", []), funmodels.get("strength_model", [])): + pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype, sub_transformer_name="transformer_2") sample = pipeline( prompt, @@ -661,8 +691,9 @@ class Wan2_2FunInpaintSampler: print('Unmerge Lora') for _lora_path, _lora_weight in zip(funmodels.get("loras", []), funmodels.get("strength_model", [])): pipeline = unmerge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype) - for _lora_path, _lora_weight in zip(funmodels.get("loras_high", []), funmodels.get("strength_model", [])): - pipeline = unmerge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype, sub_transformer_name="transformer_2") + if pipeline.transformer_2 is not None: + for _lora_path, _lora_weight in zip(funmodels.get("loras_high", []), funmodels.get("strength_model", [])): + pipeline = unmerge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype, sub_transformer_name="transformer_2") return (videos,) @@ -806,14 +837,18 @@ class Wan2_2FunV2VSampler: pipeline.transformer.enable_teacache( coefficients, steps, teacache_threshold, num_skip_start_steps=num_skip_start_steps, offload=teacache_offload ) - pipeline.transformer_2.share_teacache(transformer=pipeline.transformer) + if pipeline.transformer_2 is not None: + pipeline.transformer_2.share_teacache(transformer=pipeline.transformer) else: pipeline.transformer.disable_teacache() + if pipeline.transformer_2 is not None: + pipeline.transformer_2.disable_teacache() if cfg_skip_ratio is not None: print(f"Enable cfg_skip_ratio {cfg_skip_ratio}.") pipeline.transformer.enable_cfg_skip(cfg_skip_ratio, steps) - pipeline.transformer_2.share_cfg_skip(transformer=pipeline.transformer) + if pipeline.transformer_2 is not None: + pipeline.transformer_2.share_cfg_skip(transformer=pipeline.transformer) generator= torch.Generator(device).manual_seed(seed) @@ -823,7 +858,8 @@ class Wan2_2FunV2VSampler: if riflex_k > 0: latent_frames = (video_length - 1) // pipeline.vae.config.temporal_compression_ratio + 1 pipeline.transformer.enable_riflex(k = riflex_k, L_test = latent_frames) - pipeline.transformer_2.enable_riflex(k = riflex_k, L_test = latent_frames) + if pipeline.transformer_2 is not None: + pipeline.transformer_2.enable_riflex(k = riflex_k, L_test = latent_frames) if model_type == "Inpaint": input_video, input_video_mask, ref_image, clip_image = get_video_to_video_latent(validation_video, video_length=video_length, sample_size=(height, width), fps=16, ref_image=ref_image[0] if ref_image is not None else ref_image) @@ -860,13 +896,6 @@ class Wan2_2FunV2VSampler: transformer_state_dict = pipeline.transformer.state_dict() for key in transformer_state_dict: transformer_cpu_cache[key] = transformer_state_dict[key].clone().cpu() - - # Save the original weights to cpu - if len(transformer_high_cpu_cache) == 0: - print('Save transformer high state_dict to cpu memory') - transformer_high_state_dict = pipeline.transformer_2.state_dict() - for key in transformer_high_state_dict: - transformer_high_cpu_cache[key] = transformer_high_state_dict[key].clone().cpu() lora_path_now = str(funmodels.get("loras", []) + funmodels.get("strength_model", [])) if lora_path_now != lora_path_before: @@ -876,31 +905,43 @@ class Wan2_2FunV2VSampler: for _lora_path, _lora_weight in zip(funmodels.get("loras", []), funmodels.get("strength_model", [])): pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype) - lora_high_path_now = str(funmodels.get("loras_high", []) + funmodels.get("strength_model", [])) - if lora_high_path_now != lora_high_path_before: - print('Merge Lora High with Cache') - lora_high_path_before = copy.deepcopy(lora_high_path_now) - pipeline.transformer_2.load_state_dict(transformer_cpu_cache) - for _lora_path, _lora_weight in zip(funmodels.get("loras_high", []), funmodels.get("strength_model", [])): - pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype, sub_transformer_name="transformer_2") + if pipeline.transformer_2 is not None: + # Save the original weights to cpu + if len(transformer_high_cpu_cache) == 0: + print('Save transformer high state_dict to cpu memory') + transformer_high_state_dict = pipeline.transformer_2.state_dict() + for key in transformer_high_state_dict: + transformer_high_cpu_cache[key] = transformer_high_state_dict[key].clone().cpu() + + lora_high_path_now = str(funmodels.get("loras_high", []) + funmodels.get("strength_model", [])) + if lora_high_path_now != lora_high_path_before: + print('Merge Lora High with Cache') + lora_high_path_before = copy.deepcopy(lora_high_path_now) + pipeline.transformer_2.load_state_dict(transformer_cpu_cache) + for _lora_path, _lora_weight in zip(funmodels.get("loras_high", []), funmodels.get("strength_model", [])): + pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype, sub_transformer_name="transformer_2") else: + print('Merge Lora') # Clear lora when switch from lora_cache=True to lora_cache=False. if len(transformer_cpu_cache) != 0: pipeline.transformer.load_state_dict(transformer_cpu_cache) transformer_cpu_cache = {} lora_path_before = "" gc.collect() - # Clear lora when switch from lora_cache=True to lora_cache=False. - if len(transformer_high_cpu_cache) != 0: - pipeline.transformer.load_state_dict(transformer_high_cpu_cache) - transformer_high_cpu_cache = {} - lora_high_path_before = "" - gc.collect() - print('Merge Lora') + for _lora_path, _lora_weight in zip(funmodels.get("loras", []), funmodels.get("strength_model", [])): pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype) - for _lora_path, _lora_weight in zip(funmodels.get("loras_high", []), funmodels.get("strength_model", [])): - pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype, sub_transformer_name="transformer_2") + + # Clear lora when switch from lora_cache=True to lora_cache=False. + if pipeline.transformer_2 is not None: + if len(transformer_high_cpu_cache) != 0: + pipeline.transformer_2.load_state_dict(transformer_high_cpu_cache) + transformer_high_cpu_cache = {} + lora_high_path_before = "" + gc.collect() + + for _lora_path, _lora_weight in zip(funmodels.get("loras_high", []), funmodels.get("strength_model", [])): + pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype, sub_transformer_name="transformer_2") if model_type == "Inpaint": sample = pipeline( @@ -944,6 +985,7 @@ class Wan2_2FunV2VSampler: print('Unmerge Lora') for _lora_path, _lora_weight in zip(funmodels.get("loras", []), funmodels.get("strength_model", [])): pipeline = unmerge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype) - for _lora_path, _lora_weight in zip(funmodels.get("loras_high", []), funmodels.get("strength_model", [])): - pipeline = unmerge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype, sub_transformer_name="transformer_2") - return (videos,) + if pipeline.transformer_2 is not None: + for _lora_path, _lora_weight in zip(funmodels.get("loras_high", []), funmodels.get("strength_model", [])): + pipeline = unmerge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype, sub_transformer_name="transformer_2") + return (videos,) \ No newline at end of file diff --git a/comfyui/wan2_2_fun/v1/wan2.2_fun_workflow_control_camera.json b/comfyui/wan2_2_fun/v1/wan2.2_fun_workflow_control_camera.json new file mode 100644 index 0000000..723886f --- /dev/null +++ b/comfyui/wan2_2_fun/v1/wan2.2_fun_workflow_control_camera.json @@ -0,0 +1,670 @@ +{ + "id": "f2736b24-5147-427c-a9dd-036fd3053b5d", + "revision": 0, + "last_node_id": 135, + "last_link_id": 300, + "nodes": [ + { + "id": 107, + "type": "Note", + "pos": [ + 4, + 634 + ], + "size": [ + 210, + 88 + ], + "flags": {}, + "order": 0, + "mode": 0, + "inputs": [], + "outputs": [], + "properties": { + "text": "" + }, + "widgets_values": [ + "You can write prompt here\n(你可以在此填写提示词)" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 108, + "type": "Note", + "pos": [ + -110, + 842 + ], + "size": [ + 326.1556091308594, + 145.20904541015625 + ], + "flags": {}, + "order": 1, + "mode": 0, + "inputs": [], + "outputs": [], + "properties": { + "text": "" + }, + "widgets_values": [ + "Using longer neg prompt such as \"Blurring, mutation, deformation, distortion, dark and solid, comics.\" can increase stability. Adding words such as \"quiet, solid\" to the neg prompt can increase dynamism.\n(使用更长的neg prompt如\"模糊,突变,变形,失真,画面暗,画面固定,连环画,漫画,线稿,没有主体。\",可以增加稳定性。在neg prompt中添加\"安静,固定\"等词语可以增加动态性。)" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 122, + "type": "FunTextBox", + "pos": [ + 238, + 805 + ], + "size": [ + 400, + 200 + ], + "flags": {}, + "order": 2, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "prompt", + "type": "STRING_PROMPT", + "slot_index": 0, + "links": [ + 297 + ] + } + ], + "title": "Negtive Prompt(反向提示词)", + "properties": { + "Node name for S&R": "FunTextBox" + }, + "widgets_values": [ + "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走" + ] + }, + { + "id": 129, + "type": "CameraBasicFromChaoJie", + "pos": [ + 805.2059326171875, + 1012.381103515625 + ], + "size": [ + 315, + 106 + ], + "flags": {}, + "order": 3, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "CameraPose", + "type": "CameraPose", + "links": null + } + ], + "properties": { + "Node name for S&R": "CameraBasicFromChaoJie" + }, + "widgets_values": [ + "Static", + 1, + 16 + ] + }, + { + "id": 130, + "type": "CameraTrajectoryFromChaoJie", + "pos": [ + 1170.206298828125, + 763.3814697265625 + ], + "size": [ + 367.79998779296875, + 150 + ], + "flags": {}, + "order": 10, + "mode": 0, + "inputs": [ + { + "name": "camera_pose", + "type": "CameraPose", + "link": 285 + } + ], + "outputs": [ + { + "name": "camera_trajectory", + "type": "STRING", + "slot_index": 0, + "links": [ + 299 + ] + }, + { + "name": "video_length", + "type": "INT", + "links": null + } + ], + "properties": { + "Node name for S&R": "CameraTrajectoryFromChaoJie" + }, + "widgets_values": [ + 0.532139961, + 0.946026558, + 0.5, + 0.5 + ] + }, + { + "id": 106, + "type": "VHS_VideoCombine", + "pos": [ + 1408, + 68 + ], + "size": [ + 390, + 310 + ], + "flags": {}, + "order": 12, + "mode": 0, + "inputs": [ + { + "label": "图像", + "name": "images", + "shape": 7, + "type": "IMAGE", + "link": 300 + }, + { + "label": "音频", + "name": "audio", + "shape": 7, + "type": "AUDIO", + "link": null + }, + { + "label": "批次管理", + "name": "meta_batch", + "shape": 7, + "type": "VHS_BatchManager", + "link": null + }, + { + "name": "vae", + "shape": 7, + "type": "VAE", + "link": null + } + ], + "outputs": [ + { + "label": "文件名", + "name": "Filenames", + "type": "VHS_FILENAMES", + "slot_index": 0, + "links": null + } + ], + "properties": { + "Node name for S&R": "VHS_VideoCombine" + }, + "widgets_values": { + "frame_rate": 16, + "loop_count": 0, + "filename_prefix": "Fun", + "format": "video/h264-mp4", + "pix_fmt": "yuv420p", + "crf": 22, + "save_metadata": true, + "pingpong": false, + "save_output": true, + "videopreview": { + "hidden": false, + "paused": false, + "params": { + "filename": "Fun_00112.mp4", + "subfolder": "", + "type": "output", + "format": "video/h264-mp4", + "frame_rate": 16 + } + } + } + }, + { + "id": 121, + "type": "FunTextBox", + "pos": [ + 235, + 539 + ], + "size": [ + 400, + 200 + ], + "flags": {}, + "order": 4, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "prompt", + "type": "STRING_PROMPT", + "links": [ + 296 + ] + } + ], + "title": "Positive Prompt(正向提示词)", + "properties": { + "Node name for S&R": "FunTextBox" + }, + "widgets_values": [ + "Fireworks light up the evening sky over a sprawling cityscape with gothic-style buildings featuring pointed towers and clock faces. The city is lit by both artificial lights from the buildings and the colorful bursts of the fireworks. The scene is viewed from an elevated angle, showcasing a vibrant urban environment set against a backdrop of a dramatic, partially cloudy sky at dusk." + ] + }, + { + "id": 100, + "type": "LoadImage", + "pos": [ + 237.59738159179688, + 1164.597412109375 + ], + "size": [ + 378.07147216796875, + 314 + ], + "flags": {}, + "order": 5, + "mode": 0, + "inputs": [], + "outputs": [ + { + "label": "图像", + "name": "IMAGE", + "type": "IMAGE", + "slot_index": 0, + "links": [ + 298 + ] + }, + { + "label": "遮罩", + "name": "MASK", + "type": "MASK", + "links": null + } + ], + "title": "Start Image(图片到视频的开始图片)", + "properties": { + "Node name for S&R": "LoadImage" + }, + "widgets_values": [ + "5.png", + "image" + ] + }, + { + "id": 131, + "type": "CameraCombineFromChaoJie", + "pos": [ + 814.2059326171875, + 763.3814697265625 + ], + "size": [ + 315, + 178 + ], + "flags": {}, + "order": 6, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "CameraPose", + "type": "CameraPose", + "links": [ + 285 + ] + } + ], + "properties": { + "Node name for S&R": "CameraCombineFromChaoJie" + }, + "widgets_values": [ + "Pan Right", + "Pan Up", + "Static", + "Static", + 1, + 81 + ] + }, + { + "id": 110, + "type": "Note", + "pos": [ + 1158.206298828125, + 970.381103515625 + ], + "size": [ + 608.1410522460938, + 188.2682342529297 + ], + "flags": {}, + "order": 7, + "mode": 0, + "inputs": [], + "outputs": [], + "properties": { + "text": "" + }, + "widgets_values": [ + "CameraCombine is used to combine multiple camera movements, while CameraBasic produces a single camera movement. The nodes come from https://github.com/chaojie/ComfyUI-CameraCtrl-Wrapper/. Since ComfyUI-CameraCtrl-Wrapper requires a specific version of diffusers, the code has been copied into the current repository.\n(CameraCombine用于组合多个镜头运动,CameraBasic产出单个镜头运动;节点来自于https://github.com/chaojie/ComfyUI-CameraCtrl-Wrapper/,由于ComfyUI-CameraCtrl-Wrapper有具体diffusers版本要求,故复制代码到当前库中。)" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 112, + "type": "Note", + "pos": [ + -203, + 252 + ], + "size": [ + 427.074951171875, + 143.9142608642578 + ], + "flags": {}, + "order": 8, + "mode": 0, + "inputs": [], + "outputs": [], + "properties": { + "text": "" + }, + "widgets_values": [ + "When using the 14B model, you can use sequential_cpu_offload to save GPU memory during generation.\n(在使用14B模型时,可以使用sequential_cpu_offload节省显存,进行生成。)" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 134, + "type": "LoadWan2_2FunModel", + "pos": [ + 273.49224853515625, + 249.6970672607422 + ], + "size": [ + 350.6878662109375, + 155.20101928710938 + ], + "flags": {}, + "order": 9, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "funmodels", + "type": "FunModels", + "links": [ + 295 + ] + } + ], + "properties": { + "Node name for S&R": "LoadWan2_2FunModel" + }, + "widgets_values": [ + "Wan2.2-Fun-A14B-Control-Camera", + "Control", + "sequential_cpu_offload", + "wan2.2/wan_civitai_i2v.yaml", + "bf16" + ] + }, + { + "id": 135, + "type": "Wan2_2FunV2VSampler", + "pos": [ + 923.0663452148438, + 107.64505767822266 + ], + "size": [ + 350.2320251464844, + 526 + ], + "flags": {}, + "order": 11, + "mode": 0, + "inputs": [ + { + "name": "funmodels", + "type": "FunModels", + "link": 295 + }, + { + "name": "prompt", + "type": "STRING_PROMPT", + "link": 296 + }, + { + "name": "negative_prompt", + "type": "STRING_PROMPT", + "link": 297 + }, + { + "name": "validation_video", + "shape": 7, + "type": "IMAGE", + "link": null + }, + { + "name": "control_video", + "shape": 7, + "type": "IMAGE", + "link": null + }, + { + "name": "start_image", + "shape": 7, + "type": "IMAGE", + "link": 298 + }, + { + "name": "end_image", + "shape": 7, + "type": "IMAGE", + "link": null + }, + { + "name": "ref_image", + "shape": 7, + "type": "IMAGE", + "link": null + }, + { + "name": "camera_conditions", + "shape": 7, + "type": "STRING", + "link": 299 + }, + { + "name": "riflex_k", + "shape": 7, + "type": "RIFLEXT_ARGS", + "link": null + } + ], + "outputs": [ + { + "name": "images", + "type": "IMAGE", + "links": [ + 300 + ] + } + ], + "properties": { + "Node name for S&R": "Wan2_2FunV2VSampler" + }, + "widgets_values": [ + 81, + 640, + 43, + "fixed", + 50, + 6.000000000000001, + 1, + "Flow", + 0.1, + true, + 5, + true, + 0 + ] + } + ], + "links": [ + [ + 285, + 131, + 0, + 130, + 0, + "CameraPose" + ], + [ + 295, + 134, + 0, + 135, + 0, + "FunModels" + ], + [ + 296, + 121, + 0, + 135, + 1, + "STRING_PROMPT" + ], + [ + 297, + 122, + 0, + 135, + 2, + "STRING_PROMPT" + ], + [ + 298, + 100, + 0, + 135, + 5, + "IMAGE" + ], + [ + 299, + 130, + 0, + 135, + 8, + "STRING" + ], + [ + 300, + 135, + 0, + 106, + 0, + "IMAGE" + ] + ], + "groups": [ + { + "id": 1, + "title": "Generate Control Video", + "bounding": [ + 773, + 666, + 1025, + 531 + ], + "color": "#3f789e", + "font_size": 24, + "flags": {} + }, + { + "id": 2, + "title": "First Image", + "bounding": [ + 191, + 1068, + 475, + 456 + ], + "color": "#a1309b", + "font_size": 24, + "flags": {} + }, + { + "id": 3, + "title": "Prompts", + "bounding": [ + 191, + 456, + 475, + 587 + ], + "color": "#3f789e", + "font_size": 24, + "flags": {} + }, + { + "id": 4, + "title": "Load Model", + "bounding": [ + 189, + 160, + 475, + 269 + ], + "color": "#b06634", + "font_size": 24, + "flags": {} + } + ], + "config": {}, + "extra": { + "ds": { + "scale": 0.6830134553650709, + "offset": [ + 246.58945220178464, + 51.36424150314978 + ] + }, + "node_versions": { + "CogVideoX-Fun": "e935795d7684aaba72be65d2ad5ef580c0526f7a", + "ComfyUI-VideoHelperSuite": "70faa9bcef65932ab72e7404d6373fb300013a2e", + "comfy-core": "0.3.44" + }, + "frontendVersion": "1.21.3" + }, + "version": 0.4 +} \ No newline at end of file diff --git a/comfyui/wan2_2_fun/v1/wan2.2_fun_workflow_control_trajectory.json b/comfyui/wan2_2_fun/v1/wan2.2_fun_workflow_control_trajectory.json new file mode 100644 index 0000000..6280aa8 --- /dev/null +++ b/comfyui/wan2_2_fun/v1/wan2.2_fun_workflow_control_trajectory.json @@ -0,0 +1,1045 @@ +{ + "id": "a707e0b0-7fbd-4f91-9c23-1c5bdca19bad", + "revision": 0, + "last_node_id": 129, + "last_link_id": 294, + "nodes": [ + { + "id": 107, + "type": "Note", + "pos": [ + 4, + 634 + ], + "size": [ + 210, + 88 + ], + "flags": {}, + "order": 0, + "mode": 0, + "inputs": [], + "outputs": [], + "properties": { + "text": "" + }, + "widgets_values": [ + "You can write prompt here\n(你可以在此填写提示词)" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 108, + "type": "Note", + "pos": [ + -110, + 842 + ], + "size": [ + 326.1556091308594, + 145.20904541015625 + ], + "flags": {}, + "order": 1, + "mode": 0, + "inputs": [], + "outputs": [], + "properties": { + "text": "" + }, + "widgets_values": [ + "Using longer neg prompt such as \"Blurring, mutation, deformation, distortion, dark and solid, comics.\" can increase stability. Adding words such as \"quiet, solid\" to the neg prompt can increase dynamism.\n(使用更长的neg prompt如\"模糊,突变,变形,失真,画面暗,画面固定,连环画,漫画,线稿,没有主体。\",可以增加稳定性。在neg prompt中添加\"安静,固定\"等词语可以增加动态性。)" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 100, + "type": "LoadImage", + "pos": [ + 238, + 1165 + ], + "size": [ + 378.07147216796875, + 314 + ], + "flags": {}, + "order": 2, + "mode": 0, + "inputs": [], + "outputs": [ + { + "label": "图像", + "name": "IMAGE", + "type": "IMAGE", + "slot_index": 0, + "links": [ + 292 + ] + }, + { + "label": "遮罩", + "name": "MASK", + "type": "MASK", + "links": null + } + ], + "title": "Start Image(图片到视频的开始图片)", + "properties": { + "Node name for S&R": "LoadImage" + }, + "widgets_values": [ + "1.png", + "image" + ] + }, + { + "id": 118, + "type": "AppendStringsToList", + "pos": [ + 1140.1396484375, + 909.9193115234375 + ], + "size": [ + 315, + 82 + ], + "flags": { + "collapsed": false + }, + "order": 10, + "mode": 0, + "inputs": [ + { + "name": "string1", + "type": "STRING", + "link": 265 + }, + { + "name": "string2", + "type": "STRING", + "link": 266 + } + ], + "outputs": [ + { + "name": "STRING", + "type": "STRING", + "slot_index": 0, + "links": [ + 267 + ] + } + ], + "properties": { + "Node name for S&R": "AppendStringsToList" + }, + "widgets_values": [] + }, + { + "id": 121, + "type": "FunTextBox", + "pos": [ + 235, + 539 + ], + "size": [ + 400, + 200 + ], + "flags": {}, + "order": 3, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "prompt", + "type": "STRING_PROMPT", + "links": [ + 289 + ] + } + ], + "title": "Positive Prompt(正向提示词)", + "properties": { + "Node name for S&R": "FunTextBox" + }, + "widgets_values": [ + "一只棕褐色的狗正摇晃着脑袋,坐在一个舒适的房间里的浅色沙发上。沙发看起来柔软而宽敞,为这只活泼的狗狗提供了一个完美的休息地点。在狗的后面,靠墙摆放着一个架子,架子上挂着一幅精美的镶框画,画中描绘着一些美丽的风景或场景。画框周围装饰着粉红色的花朵,这些花朵不仅增添了房间的色彩,还带来了一丝自然和生机。房间里的灯光柔和而温暖,从天花板上的吊灯和角落里的台灯散发出来,营造出一种温馨舒适的氛围。整个空间给人一种宁静和谐的感觉,仿佛时间在这里变得缓慢而美好。" + ] + }, + { + "id": 95, + "type": "CreateTrajectoryBasedOnKJNodes", + "pos": [ + 1574.139404296875, + 929.9193115234375 + ], + "size": [ + 428.4000244140625, + 58 + ], + "flags": {}, + "order": 11, + "mode": 0, + "inputs": [ + { + "name": "coordinates", + "type": "STRING", + "link": 267 + }, + { + "name": "masks", + "type": "MASK", + "link": 249 + } + ], + "outputs": [ + { + "name": "image", + "type": "IMAGE", + "slot_index": 0, + "links": [ + 237, + 262, + 291 + ] + } + ], + "properties": { + "Node name for S&R": "CreateTrajectoryBasedOnKJNodes" + }, + "widgets_values": [] + }, + { + "id": 122, + "type": "FunTextBox", + "pos": [ + 238, + 805 + ], + "size": [ + 400, + 200 + ], + "flags": {}, + "order": 4, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "prompt", + "type": "STRING_PROMPT", + "slot_index": 0, + "links": [ + 290 + ] + } + ], + "title": "Negtive Prompt(反向提示词)", + "properties": { + "Node name for S&R": "FunTextBox" + }, + "widgets_values": [ + "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走" + ] + }, + { + "id": 97, + "type": "SplineEditor", + "pos": [ + 855.1397705078125, + 1058.91943359375 + ], + "size": [ + 645, + 832 + ], + "flags": {}, + "order": 5, + "mode": 0, + "inputs": [ + { + "name": "bg_image", + "shape": 7, + "type": "IMAGE", + "link": null + } + ], + "outputs": [ + { + "name": "mask", + "type": "MASK", + "slot_index": 0, + "links": [ + 249 + ] + }, + { + "name": "coord_str", + "type": "STRING", + "slot_index": 1, + "links": [ + 265 + ] + }, + { + "name": "float", + "type": "FLOAT", + "links": null + }, + { + "name": "count", + "type": "INT", + "links": null + }, + { + "name": "normalized_str", + "type": "STRING", + "links": null + } + ], + "properties": { + "Node name for S&R": "SplineEditor", + "points": "SplineEditor", + "imgData": null + }, + "widgets_values": [ + "[{\"points\":[{\"x\":236.74497000000005,\"y\":169.10355000000004},{\"x\":263.53799999999995,\"y\":230.59575},{\"x\":321.61075881587004,\"y\":229.72197058276433},{\"x\":355.3934015486295,\"y\":168.9132136637973},{\"x\":343.23165016483614,\"y\":82.42964826793309},{\"x\":275.6663646993172,\"y\":55.40353408172552},{\"x\":197.29063355931524,\"y\":85.13225968655384},{\"x\":177.02104791965957,\"y\":148.64362802414163},{\"x\":205.39846781517753,\"y\":232.4245820013851},{\"x\":259.45069618759265,\"y\":167.56190795448694}],\"color\":\"#1f77b4\",\"name\":\"Spline 1\"}]", + "[[{\"x\":236.74496459960938,\"y\":169.10354614257812},{\"x\":238.94464111328125,\"y\":177.68576049804688},{\"x\":241.3327178955078,\"y\":186.21737670898438},{\"x\":243.95693969726562,\"y\":194.67919921875},{\"x\":246.88812255859375,\"y\":203.03933715820312},{\"x\":250.2393035888672,\"y\":211.23939514160156},{\"x\":254.20826721191406,\"y\":219.15655517578125},{\"x\":259.18768310546875,\"y\":226.46995544433594},{\"x\":265.87689208984375,\"y\":232.22171020507812},{\"x\":273.4639892578125,\"y\":236.787353515625},{\"x\":281.6380920410156,\"y\":240.17092895507812},{\"x\":290.3527526855469,\"y\":241.61793518066406},{\"x\":299.1367492675781,\"y\":240.67857360839844},{\"x\":307.49658203125,\"y\":237.7835235595703},{\"x\":315.3465576171875,\"y\":233.68600463867188},{\"x\":322.80859375,\"y\":228.9126739501953},{\"x\":330.00958251953125,\"y\":223.75323486328125},{\"x\":336.7490234375,\"y\":218.0097198486328},{\"x\":342.5494689941406,\"y\":211.32901000976562},{\"x\":346.96185302734375,\"y\":203.6618194580078},{\"x\":350.05450439453125,\"y\":195.3665771484375},{\"x\":352.2560729980469,\"y\":186.787109375},{\"x\":353.9389343261719,\"y\":178.08937072753906},{\"x\":355.33001708984375,\"y\":169.3397216796875},{\"x\":356.59344482421875,\"y\":160.57052612304688},{\"x\":357.7420959472656,\"y\":151.78565979003906},{\"x\":358.6518249511719,\"y\":142.9732208251953},{\"x\":359.1504211425781,\"y\":134.1287384033203},{\"x\":359.0207214355469,\"y\":125.2725601196289},{\"x\":358.035400390625,\"y\":116.4721908569336},{\"x\":356.0333251953125,\"y\":107.84730529785156},{\"x\":352.9960632324219,\"y\":99.53003692626953},{\"x\":349.046142578125,\"y\":91.60392761230469},{\"x\":344.37164306640625,\"y\":84.08070373535156},{\"x\":339.17694091796875,\"y\":76.9051742553711},{\"x\":333.4718322753906,\"y\":70.13098907470703},{\"x\":326.9710693359375,\"y\":64.12535858154297},{\"x\":319.4295654296875,\"y\":59.51286315917969},{\"x\":311.02655029296875,\"y\":56.7620735168457},{\"x\":302.25640869140625,\"y\":55.55453872680664},{\"x\":293.40484619140625,\"y\":55.21137619018555},{\"x\":284.5453796386719,\"y\":55.24800491333008},{\"x\":275.68701171875,\"y\":55.40314865112305},{\"x\":266.83221435546875,\"y\":55.690216064453125},{\"x\":257.9945373535156,\"y\":56.30451965332031},{\"x\":249.2041473388672,\"y\":57.398460388183594},{\"x\":240.5244140625,\"y\":59.16102981567383},{\"x\":232.06307983398438,\"y\":61.773162841796875},{\"x\":223.95306396484375,\"y\":65.32772827148438},{\"x\":216.29017639160156,\"y\":69.7664566040039},{\"x\":209.08375549316406,\"y\":74.91584777832031},{\"x\":202.27122497558594,\"y\":80.57798767089844},{\"x\":195.76087951660156,\"y\":86.58625030517578},{\"x\":189.468505859375,\"y\":92.8224105834961},{\"x\":183.6200714111328,\"y\":99.47235107421875},{\"x\":178.9163055419922,\"y\":106.95783233642578},{\"x\":176.3991241455078,\"y\":115.42044830322266},{\"x\":175.83934020996094,\"y\":124.25370788574219},{\"x\":176.12339782714844,\"y\":133.10797119140625},{\"x\":176.63381958007812,\"y\":141.95301818847656},{\"x\":177.14378356933594,\"y\":150.79808044433594},{\"x\":177.73876953125,\"y\":159.63778686523438},{\"x\":178.5083770751953,\"y\":168.46385192871094},{\"x\":179.4958038330078,\"y\":177.26815795898438},{\"x\":180.75762939453125,\"y\":186.03709411621094},{\"x\":182.37294006347656,\"y\":194.7476043701172},{\"x\":184.4583282470703,\"y\":203.35687255859375},{\"x\":187.1986083984375,\"y\":211.7789306640625},{\"x\":190.9047393798828,\"y\":219.81732177734375},{\"x\":196.108154296875,\"y\":226.9572296142578},{\"x\":203.41249084472656,\"y\":231.83749389648438},{\"x\":211.85409545898438,\"y\":230.89816284179688},{\"x\":218.7886199951172,\"y\":225.41867065429688},{\"x\":224.8160400390625,\"y\":218.92974853515625},{\"x\":230.3641357421875,\"y\":212.0236358642578},{\"x\":235.608642578125,\"y\":204.8833770751953},{\"x\":240.63941955566406,\"y\":197.59080505371094},{\"x\":245.50965881347656,\"y\":190.1898651123047},{\"x\":250.25364685058594,\"y\":182.7073516845703},{\"x\":254.8948974609375,\"y\":175.16050720214844},{\"x\":259.45068359375,\"y\":167.56190490722656}]]", + 600, + 382, + 81, + "path", + "cardinal", + 0.5, + 1, + "list", + 0, + 1, + "", + null + ] + }, + { + "id": 119, + "type": "SplineEditor", + "pos": [ + 1544.139404296875, + 1047.919189453125 + ], + "size": [ + 645, + 832 + ], + "flags": {}, + "order": 6, + "mode": 0, + "inputs": [ + { + "name": "bg_image", + "shape": 7, + "type": "IMAGE", + "link": null + } + ], + "outputs": [ + { + "name": "mask", + "type": "MASK", + "slot_index": 0, + "links": [] + }, + { + "name": "coord_str", + "type": "STRING", + "slot_index": 1, + "links": [ + 266 + ] + }, + { + "name": "float", + "type": "FLOAT", + "links": null + }, + { + "name": "count", + "type": "INT", + "links": null + }, + { + "name": "normalized_str", + "type": "STRING", + "links": null + } + ], + "properties": { + "Node name for S&R": "SplineEditor", + "points": "SplineEditor", + "imgData": null + }, + "widgets_values": [ + "[{\"points\":[{\"x\":63.916760050380844,\"y\":114.45559357858895},{\"x\":71.34894145158792,\"y\":114.45559357858895}],\"color\":\"#1f77b4\",\"name\":\"Spline 1\"}]", + "[[{\"x\":63.9167594909668,\"y\":114.45559692382812},{\"x\":64.00965881347656,\"y\":114.45559692382812},{\"x\":64.1025619506836,\"y\":114.45559692382812},{\"x\":64.19546508789062,\"y\":114.45559692382812},{\"x\":64.28836822509766,\"y\":114.45559692382812},{\"x\":64.38127136230469,\"y\":114.45559692382812},{\"x\":64.47417449951172,\"y\":114.45559692382812},{\"x\":64.56707763671875,\"y\":114.45559692382812},{\"x\":64.65998077392578,\"y\":114.45559692382812},{\"x\":64.75287628173828,\"y\":114.45559692382812},{\"x\":64.84577941894531,\"y\":114.45559692382812},{\"x\":64.93868255615234,\"y\":114.45559692382812},{\"x\":65.03158569335938,\"y\":114.45559692382812},{\"x\":65.1244888305664,\"y\":114.45559692382812},{\"x\":65.21739196777344,\"y\":114.45559692382812},{\"x\":65.31029510498047,\"y\":114.45559692382812},{\"x\":65.4031982421875,\"y\":114.45559692382812},{\"x\":65.49609375,\"y\":114.45559692382812},{\"x\":65.58899688720703,\"y\":114.45559692382812},{\"x\":65.68190002441406,\"y\":114.45559692382812},{\"x\":65.7748031616211,\"y\":114.45559692382812},{\"x\":65.86770629882812,\"y\":114.45559692382812},{\"x\":65.96060943603516,\"y\":114.45559692382812},{\"x\":66.05351257324219,\"y\":114.45559692382812},{\"x\":66.14641571044922,\"y\":114.45559692382812},{\"x\":66.23931884765625,\"y\":114.45559692382812},{\"x\":66.33221435546875,\"y\":114.45559692382812},{\"x\":66.42511749267578,\"y\":114.45559692382812},{\"x\":66.51802062988281,\"y\":114.45559692382812},{\"x\":66.61092376708984,\"y\":114.45559692382812},{\"x\":66.70382690429688,\"y\":114.45559692382812},{\"x\":66.7967300415039,\"y\":114.45559692382812},{\"x\":66.88963317871094,\"y\":114.45559692382812},{\"x\":66.98253631591797,\"y\":114.45559692382812},{\"x\":67.075439453125,\"y\":114.45559692382812},{\"x\":67.1683349609375,\"y\":114.45559692382812},{\"x\":67.26123809814453,\"y\":114.45559692382812},{\"x\":67.35414123535156,\"y\":114.45559692382812},{\"x\":67.4470443725586,\"y\":114.45559692382812},{\"x\":67.53994750976562,\"y\":114.45559692382812},{\"x\":67.63285064697266,\"y\":114.45559692382812},{\"x\":67.72575378417969,\"y\":114.45559692382812},{\"x\":67.81864929199219,\"y\":114.45559692382812},{\"x\":67.91155242919922,\"y\":114.45559692382812},{\"x\":68.00445556640625,\"y\":114.45559692382812},{\"x\":68.09735870361328,\"y\":114.45559692382812},{\"x\":68.19026184082031,\"y\":114.45559692382812},{\"x\":68.28316497802734,\"y\":114.45559692382812},{\"x\":68.37606811523438,\"y\":114.45559692382812},{\"x\":68.4689712524414,\"y\":114.45559692382812},{\"x\":68.56187438964844,\"y\":114.45559692382812},{\"x\":68.65476989746094,\"y\":114.45559692382812},{\"x\":68.74767303466797,\"y\":114.45559692382812},{\"x\":68.840576171875,\"y\":114.45559692382812},{\"x\":68.93347930908203,\"y\":114.45559692382812},{\"x\":69.02638244628906,\"y\":114.45559692382812},{\"x\":69.1192855834961,\"y\":114.45559692382812},{\"x\":69.21218872070312,\"y\":114.45559692382812},{\"x\":69.30509185791016,\"y\":114.45559692382812},{\"x\":69.39799499511719,\"y\":114.45559692382812},{\"x\":69.49089050292969,\"y\":114.45559692382812},{\"x\":69.58379364013672,\"y\":114.45559692382812},{\"x\":69.67669677734375,\"y\":114.45559692382812},{\"x\":69.76959991455078,\"y\":114.45559692382812},{\"x\":69.86250305175781,\"y\":114.45559692382812},{\"x\":69.95540618896484,\"y\":114.45559692382812},{\"x\":70.04830932617188,\"y\":114.45559692382812},{\"x\":70.1412124633789,\"y\":114.45559692382812},{\"x\":70.2341079711914,\"y\":114.45559692382812},{\"x\":70.32701110839844,\"y\":114.45559692382812},{\"x\":70.41991424560547,\"y\":114.45559692382812},{\"x\":70.5128173828125,\"y\":114.45559692382812},{\"x\":70.60572052001953,\"y\":114.45559692382812},{\"x\":70.69862365722656,\"y\":114.45559692382812},{\"x\":70.7915267944336,\"y\":114.45559692382812},{\"x\":70.88442993164062,\"y\":114.45559692382812},{\"x\":70.97732543945312,\"y\":114.45559692382812},{\"x\":71.07022857666016,\"y\":114.45559692382812},{\"x\":71.16313171386719,\"y\":114.45559692382812},{\"x\":71.25603485107422,\"y\":114.45559692382812},{\"x\":71.34893798828125,\"y\":114.45559692382812}]]", + 600, + 382, + 81, + "path", + "cardinal", + 0.5, + 1, + "list", + 0, + 1, + "", + null + ] + }, + { + "id": 44, + "type": "VHS_VideoCombine", + "pos": [ + 2241.138427734375, + 1051.91943359375 + ], + "size": [ + 530, + 650.4000244140625 + ], + "flags": {}, + "order": 12, + "mode": 0, + "inputs": [ + { + "name": "images", + "shape": 7, + "type": "IMAGE", + "link": 237 + }, + { + "name": "audio", + "shape": 7, + "type": "AUDIO", + "link": null + }, + { + "name": "meta_batch", + "shape": 7, + "type": "VHS_BatchManager", + "link": null + }, + { + "name": "vae", + "shape": 7, + "type": "VAE", + "link": null + } + ], + "outputs": [ + { + "name": "Filenames", + "type": "VHS_FILENAMES", + "slot_index": 0, + "links": null + } + ], + "title": "Trajectory Outputs", + "properties": { + "Node name for S&R": "VHS_VideoCombine" + }, + "widgets_values": { + "frame_rate": 16, + "loop_count": 0, + "filename_prefix": "Fun-Trajectory", + "format": "video/h264-mp4", + "pix_fmt": "yuv420p", + "crf": 22, + "save_metadata": true, + "pingpong": false, + "save_output": true, + "videopreview": { + "hidden": false, + "paused": false, + "params": { + "filename": "Fun-Trajectory_00002.mp4", + "subfolder": "", + "type": "output", + "format": "video/h264-mp4", + "frame_rate": 16 + }, + "muted": false + } + } + }, + { + "id": 115, + "type": "VHS_VideoCombine", + "pos": [ + 2819, + 1056 + ], + "size": [ + 530, + 310 + ], + "flags": {}, + "order": 16, + "mode": 0, + "inputs": [ + { + "name": "images", + "shape": 7, + "type": "IMAGE", + "link": 264 + }, + { + "name": "audio", + "shape": 7, + "type": "AUDIO", + "link": null + }, + { + "name": "meta_batch", + "shape": 7, + "type": "VHS_BatchManager", + "link": null + }, + { + "name": "vae", + "shape": 7, + "type": "VAE", + "link": null + } + ], + "outputs": [ + { + "name": "Filenames", + "type": "VHS_FILENAMES", + "slot_index": 0, + "links": null + } + ], + "title": "Video with Trajectory Outputs", + "properties": { + "Node name for S&R": "VHS_VideoCombine" + }, + "widgets_values": { + "frame_rate": 16, + "loop_count": 0, + "filename_prefix": "Fun-Trajectory-Merge", + "format": "video/h264-mp4", + "pix_fmt": "yuv420p", + "crf": 22, + "save_metadata": true, + "pingpong": false, + "save_output": true, + "videopreview": { + "hidden": false, + "paused": false, + "params": { + "filename": "EasyAnimate-Trajectory-Merge_00009.mp4", + "subfolder": "", + "type": "output", + "format": "video/h264-mp4", + "frame_rate": 8 + }, + "muted": false + } + } + }, + { + "id": 112, + "type": "Note", + "pos": [ + -203, + 252 + ], + "size": [ + 427.074951171875, + 143.9142608642578 + ], + "flags": {}, + "order": 7, + "mode": 0, + "inputs": [], + "outputs": [], + "properties": { + "text": "" + }, + "widgets_values": [ + "When using the 14B model, you can use sequential_cpu_offload to save GPU memory during generation.\n(在使用14B模型时,可以使用sequential_cpu_offload节省显存,进行生成。)" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 128, + "type": "LoadWan2_2FunModel", + "pos": [ + 277.63623046875, + 254.2862548828125 + ], + "size": [ + 363.85333251953125, + 154 + ], + "flags": {}, + "order": 8, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "funmodels", + "type": "FunModels", + "links": [ + 288 + ] + } + ], + "properties": { + "Node name for S&R": "LoadWan2_2FunModel" + }, + "widgets_values": [ + "Wan2.2-Fun-A14B-Control", + "Control", + "sequential_cpu_offload", + "wan2.2/wan_civitai_i2v.yaml", + "bf16" + ] + }, + { + "id": 110, + "type": "Note", + "pos": [ + 836.5111083984375, + 639.8340454101562 + ], + "size": [ + 608.1410522460938, + 188.2682342529297 + ], + "flags": {}, + "order": 9, + "mode": 0, + "inputs": [], + "outputs": [], + "properties": { + "text": "" + }, + "widgets_values": [ + "Please set the mask height and mask width according to the height and width of the reference image. \nPlease set the video_length of the Spline Editor below to be the same as the video_length of the Sampler above. \nThe nodes are from KJNodes. For more details, please check https://github.com/kijai/ComfyUI-KJNodes/tree/main. \n\n请根据参考图片的高和宽设置mask height和mask width;\n请将下方Spline Editor的video_legnth置的与上方Sampler的video_legnth一样;\n部分节点来自于KJNodes,具体查看https://github.com/kijai/ComfyUI-KJNodes/tree/main;" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 106, + "type": "VHS_VideoCombine", + "pos": [ + 1386.0306396484375, + 42.99613952636719 + ], + "size": [ + 390, + 310 + ], + "flags": {}, + "order": 14, + "mode": 0, + "inputs": [ + { + "label": "图像", + "name": "images", + "shape": 7, + "type": "IMAGE", + "link": 293 + }, + { + "label": "音频", + "name": "audio", + "shape": 7, + "type": "AUDIO", + "link": null + }, + { + "label": "批次管理", + "name": "meta_batch", + "shape": 7, + "type": "VHS_BatchManager", + "link": null + }, + { + "name": "vae", + "shape": 7, + "type": "VAE", + "link": null + } + ], + "outputs": [ + { + "label": "文件名", + "name": "Filenames", + "type": "VHS_FILENAMES", + "slot_index": 0, + "links": null + } + ], + "properties": { + "Node name for S&R": "VHS_VideoCombine" + }, + "widgets_values": { + "frame_rate": 16, + "loop_count": 0, + "filename_prefix": "Fun", + "format": "video/h264-mp4", + "pix_fmt": "yuv420p", + "crf": 22, + "save_metadata": true, + "pingpong": false, + "save_output": true, + "videopreview": { + "hidden": false, + "paused": false, + "params": { + "filename": "EasyAnimate_00105.mp4", + "subfolder": "", + "type": "output", + "format": "video/h264-mp4", + "frame_rate": 8 + } + } + } + }, + { + "id": 114, + "type": "ImageMaximumNode", + "pos": [ + 2114.7646484375, + 884.5827026367188 + ], + "size": [ + 210, + 46 + ], + "flags": {}, + "order": 15, + "mode": 0, + "inputs": [ + { + "name": "video_1", + "type": "IMAGE", + "link": 294 + }, + { + "name": "video_2", + "type": "IMAGE", + "link": 262 + } + ], + "outputs": [ + { + "name": "image", + "type": "IMAGE", + "links": [ + 264 + ] + } + ], + "properties": { + "Node name for S&R": "ImageMaximumNode" + }, + "widgets_values": [] + }, + { + "id": 129, + "type": "Wan2_2FunV2VSampler", + "pos": [ + 835.7046508789062, + 41.22541809082031 + ], + "size": [ + 470.3740234375, + 526 + ], + "flags": {}, + "order": 13, + "mode": 0, + "inputs": [ + { + "name": "funmodels", + "type": "FunModels", + "link": 288 + }, + { + "name": "prompt", + "type": "STRING_PROMPT", + "link": 289 + }, + { + "name": "negative_prompt", + "type": "STRING_PROMPT", + "link": 290 + }, + { + "name": "validation_video", + "shape": 7, + "type": "IMAGE", + "link": null + }, + { + "name": "control_video", + "shape": 7, + "type": "IMAGE", + "link": 291 + }, + { + "name": "start_image", + "shape": 7, + "type": "IMAGE", + "link": 292 + }, + { + "name": "end_image", + "shape": 7, + "type": "IMAGE", + "link": null + }, + { + "name": "ref_image", + "shape": 7, + "type": "IMAGE", + "link": null + }, + { + "name": "camera_conditions", + "shape": 7, + "type": "STRING", + "link": null + }, + { + "name": "riflex_k", + "shape": 7, + "type": "RIFLEXT_ARGS", + "link": null + } + ], + "outputs": [ + { + "name": "images", + "type": "IMAGE", + "links": [ + 293, + 294 + ] + } + ], + "properties": { + "Node name for S&R": "Wan2_2FunV2VSampler" + }, + "widgets_values": [ + 81, + 640, + 43, + "fixed", + 50, + 6.000000000000001, + 1, + "Flow", + 0.1, + true, + 5, + true, + 0 + ] + } + ], + "links": [ + [ + 237, + 95, + 0, + 44, + 0, + "IMAGE" + ], + [ + 249, + 97, + 0, + 95, + 1, + "MASK" + ], + [ + 262, + 95, + 0, + 114, + 1, + "IMAGE" + ], + [ + 264, + 114, + 0, + 115, + 0, + "IMAGE" + ], + [ + 265, + 97, + 1, + 118, + 0, + "STRING" + ], + [ + 266, + 119, + 1, + 118, + 1, + "STRING" + ], + [ + 267, + 118, + 0, + 95, + 0, + "STRING" + ], + [ + 288, + 128, + 0, + 129, + 0, + "FunModels" + ], + [ + 289, + 121, + 0, + 129, + 1, + "STRING_PROMPT" + ], + [ + 290, + 122, + 0, + 129, + 2, + "STRING_PROMPT" + ], + [ + 291, + 95, + 0, + 129, + 4, + "IMAGE" + ], + [ + 292, + 100, + 0, + 129, + 5, + "IMAGE" + ], + [ + 293, + 129, + 0, + 106, + 0, + "IMAGE" + ], + [ + 294, + 129, + 0, + 114, + 0, + "IMAGE" + ] + ], + "groups": [ + { + "id": 1, + "title": "Load Model", + "bounding": [ + 189, + 160, + 475, + 269 + ], + "color": "#b06634", + "font_size": 24, + "flags": {} + }, + { + "id": 2, + "title": "Prompts", + "bounding": [ + 191, + 456, + 475, + 587 + ], + "color": "#3f789e", + "font_size": 24, + "flags": {} + }, + { + "id": 3, + "title": "First Image of Trajectory", + "bounding": [ + 191, + 1068, + 475, + 456 + ], + "color": "#a1309b", + "font_size": 24, + "flags": {} + }, + { + "id": 4, + "title": "Generate Control Video", + "bounding": [ + 786, + 841, + 2616, + 1056 + ], + "color": "#3f789e", + "font_size": 24, + "flags": {} + } + ], + "config": {}, + "extra": { + "ds": { + "scale": 0.6830134553650711, + "offset": [ + 109.96141399088638, + -28.91205721966216 + ] + }, + "frontendVersion": "1.21.3", + "node_versions": { + "comfy-core": "0.3.44", + "ComfyUI-KJNodes": "ff49e1b01f10a14496b08e21bb89b64d2b15f333", + "CogVideoX-Fun": "e935795d7684aaba72be65d2ad5ef580c0526f7a", + "ComfyUI-VideoHelperSuite": "70faa9bcef65932ab72e7404d6373fb300013a2e" + } + }, + "version": 0.4 +} \ No newline at end of file diff --git a/comfyui/wan2_2_fun/v1/wan2.2_fun_workflow_t2v.json b/comfyui/wan2_2_fun/v1/wan2.2_fun_workflow_t2v.json new file mode 100644 index 0000000..9a8fdc6 --- /dev/null +++ b/comfyui/wan2_2_fun/v1/wan2.2_fun_workflow_t2v.json @@ -0,0 +1,409 @@ +{ + "id": "8d9a378f-1cac-4610-8858-351d5982a6ab", + "revision": 0, + "last_node_id": 103, + "last_link_id": 73, + "nodes": [ + { + "id": 75, + "type": "FunTextBox", + "pos": [ + 250, + -50 + ], + "size": [ + 383.54010009765625, + 156.71620178222656 + ], + "flags": {}, + "order": 0, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "prompt", + "type": "STRING_PROMPT", + "slot_index": 0, + "links": [ + 72 + ] + } + ], + "title": "Positive Prompt(正向提示词)", + "properties": { + "Node name for S&R": "FunTextBox" + }, + "widgets_values": [ + "fireworks display over night city. The video is of high quality, and the view is very clear. High quality, masterpiece, best quality, highres, ultra-detailed, fantastic." + ] + }, + { + "id": 78, + "type": "Note", + "pos": [ + 18, + -46 + ], + "size": [ + 210, + 88 + ], + "flags": {}, + "order": 1, + "mode": 0, + "inputs": [], + "outputs": [], + "properties": { + "text": "" + }, + "widgets_values": [ + "You can write prompt here\n(你可以在此填写提示词)" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 94, + "type": "Note", + "pos": [ + 17, + -35 + ], + "size": [ + 210, + 88 + ], + "flags": {}, + "order": 2, + "mode": 0, + "inputs": [], + "outputs": [], + "properties": { + "text": "" + }, + "widgets_values": [ + "You can write prompt here\n(你可以在此填写提示词)" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 80, + "type": "Note", + "pos": [ + -95.01139068603516, + -334.30706787109375 + ], + "size": [ + 355.636474609375, + 132.4238739013672 + ], + "flags": {}, + "order": 3, + "mode": 0, + "inputs": [], + "outputs": [], + "properties": { + "text": "" + }, + "widgets_values": [ + "When using the 14B model, you can use sequential_cpu_offload to save GPU memory during generation.\n(在使用14B模型时,可以使用sequential_cpu_offload节省显存,进行生成。)" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 73, + "type": "FunTextBox", + "pos": [ + 250, + 160 + ], + "size": [ + 383.7149963378906, + 183.83506774902344 + ], + "flags": {}, + "order": 4, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "prompt", + "type": "STRING_PROMPT", + "slot_index": 0, + "links": [ + 73 + ] + } + ], + "title": "Negtive Prompt(反向提示词)", + "properties": { + "Node name for S&R": "FunTextBox" + }, + "widgets_values": [ + "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走" + ] + }, + { + "id": 17, + "type": "VHS_VideoCombine", + "pos": [ + 1257.373291015625, + -146.444091796875 + ], + "size": [ + 390, + 537.4615478515625 + ], + "flags": {}, + "order": 7, + "mode": 0, + "inputs": [ + { + "label": "图像", + "name": "images", + "shape": 7, + "type": "IMAGE", + "link": 70 + }, + { + "label": "音频", + "name": "audio", + "shape": 7, + "type": "AUDIO", + "link": null + }, + { + "label": "批次管理", + "name": "meta_batch", + "shape": 7, + "type": "VHS_BatchManager", + "link": null + }, + { + "name": "vae", + "shape": 7, + "type": "VAE", + "link": null + } + ], + "outputs": [ + { + "label": "文件名", + "name": "Filenames", + "type": "VHS_FILENAMES", + "slot_index": 0, + "links": null + } + ], + "properties": { + "Node name for S&R": "VHS_VideoCombine" + }, + "widgets_values": { + "frame_rate": 16, + "loop_count": 0, + "filename_prefix": "Fun", + "format": "video/h264-mp4", + "pix_fmt": "yuv420p", + "crf": 22, + "save_metadata": true, + "pingpong": false, + "save_output": true, + "videopreview": { + "hidden": false, + "paused": false, + "params": { + "filename": "Fun_00008.mp4", + "subfolder": "", + "type": "output", + "format": "video/h264-mp4", + "frame_rate": 16 + } + } + } + }, + { + "id": 99, + "type": "LoadWan2_2FunModel", + "pos": [ + 290, + -334 + ], + "size": [ + 315, + 154 + ], + "flags": {}, + "order": 5, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "funmodels", + "type": "FunModels", + "links": [ + 71 + ] + } + ], + "properties": { + "Node name for S&R": "LoadWan2_2FunModel" + }, + "widgets_values": [ + "Wan2.2-Fun-A14B-InP", + "Inpaint", + "sequential_cpu_offload", + "wan2.2/wan_civitai_i2v.yaml", + "bf16" + ] + }, + { + "id": 103, + "type": "Wan2_2FunT2VSampler", + "pos": [ + 819.128173828125, + -145.90399169921875 + ], + "size": [ + 340.3540954589844, + 430 + ], + "flags": {}, + "order": 6, + "mode": 0, + "inputs": [ + { + "name": "funmodels", + "type": "FunModels", + "link": 71 + }, + { + "name": "prompt", + "type": "STRING_PROMPT", + "link": 72 + }, + { + "name": "negative_prompt", + "type": "STRING_PROMPT", + "link": 73 + }, + { + "name": "riflex_k", + "shape": 7, + "type": "RIFLEXT_ARGS", + "link": null + } + ], + "outputs": [ + { + "name": "images", + "type": "IMAGE", + "links": [ + 70 + ] + } + ], + "properties": { + "Node name for S&R": "Wan2_2FunT2VSampler" + }, + "widgets_values": [ + 81, + 832, + 480, + false, + 43, + "fixed", + 50, + 6, + "Flow", + 0.1, + true, + 5, + true, + 0 + ] + } + ], + "links": [ + [ + 70, + 103, + 0, + 17, + 0, + "IMAGE" + ], + [ + 71, + 99, + 0, + 103, + 0, + "FunModels" + ], + [ + 72, + 75, + 0, + 103, + 1, + "STRING_PROMPT" + ], + [ + 73, + 73, + 0, + 103, + 2, + "STRING_PROMPT" + ] + ], + "groups": [ + { + "id": 2, + "title": "Prompts", + "bounding": [ + 218, + -127, + 450, + 483 + ], + "color": "#3f789e", + "font_size": 24, + "flags": {} + }, + { + "id": 3, + "title": "Load Model", + "bounding": [ + 227, + -416, + 469, + 256 + ], + "color": "#b06634", + "font_size": 24, + "flags": {} + } + ], + "config": {}, + "extra": { + "ds": { + "scale": 0.8264462809917354, + "offset": [ + 175.2685031365507, + 463.74646895767313 + ] + }, + "workspace_info": { + "id": "776b62b4-bd17-4ed3-9923-b7aad000b1ea" + }, + "node_versions": { + "CogVideoX-Fun": "5f2a55692d0834e00d477939b818b9bb63c06535", + "ComfyUI-VideoHelperSuite": "70faa9bcef65932ab72e7404d6373fb300013a2e" + }, + "frontendVersion": "1.21.3" + }, + "version": 0.4 +} \ No newline at end of file diff --git a/comfyui/wan2_2_fun/v1/wan2.2_fun_workflow_v2v_control.json b/comfyui/wan2_2_fun/v1/wan2.2_fun_workflow_v2v_control.json index 53de08e..308b51a 100755 --- a/comfyui/wan2_2_fun/v1/wan2.2_fun_workflow_v2v_control.json +++ b/comfyui/wan2_2_fun/v1/wan2.2_fun_workflow_v2v_control.json @@ -529,7 +529,7 @@ "flags": {} }, { - "title": "Load EasyAnimate", + "title": "Load Model", "bounding": [ 218, -387, diff --git a/comfyui/wan2_2_fun/v1/wan2.2_fun_workflow_v2v_control_canny.json b/comfyui/wan2_2_fun/v1/wan2.2_fun_workflow_v2v_control_canny.json new file mode 100644 index 0000000..d458e19 --- /dev/null +++ b/comfyui/wan2_2_fun/v1/wan2.2_fun_workflow_v2v_control_canny.json @@ -0,0 +1,700 @@ +{ + "id": "90329caf-a94f-48a6-80d7-d6e167a8b1e3", + "revision": 0, + "last_node_id": 102, + "last_link_id": 78, + "nodes": [ + { + "id": 78, + "type": "Note", + "pos": [ + 18, + -46 + ], + "size": [ + 210, + 88 + ], + "flags": {}, + "order": 0, + "mode": 0, + "inputs": [], + "outputs": [], + "properties": { + "text": "" + }, + "widgets_values": [ + "You can write prompt here\n(你可以在此填写提示词)" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 79, + "type": "Note", + "pos": [ + 15.739953994750977, + 462.38665771484375 + ], + "size": [ + 210, + 88 + ], + "flags": {}, + "order": 1, + "mode": 0, + "inputs": [], + "outputs": [], + "properties": { + "text": "" + }, + "widgets_values": [ + "You can upload video here\n(在此上传视频)" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 88, + "type": "Note", + "pos": [ + -99, + 197 + ], + "size": [ + 326.1556091308594, + 145.20904541015625 + ], + "flags": {}, + "order": 2, + "mode": 0, + "inputs": [], + "outputs": [], + "properties": { + "text": "" + }, + "widgets_values": [ + "Using longer neg prompt such as \"Blurring, mutation, deformation, distortion, dark and solid, comics.\" can increase stability. Adding words such as \"quiet, solid\" to the neg prompt can increase dynamism.\n(使用更长的neg prompt如\"模糊,突变,变形,失真,画面暗,画面固定,连环画,漫画,线稿,没有主体。\",可以增加稳定性。在neg prompt中添加\"安静,固定\"等词语可以增加动态性。)" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 89, + "type": "Note", + "pos": [ + -192, + -293 + ], + "size": [ + 427.074951171875, + 143.9142608642578 + ], + "flags": {}, + "order": 3, + "mode": 0, + "inputs": [], + "outputs": [], + "properties": { + "text": "" + }, + "widgets_values": [ + "When using the 14B model, you can use sequential_cpu_offload to save GPU memory during generation.\n(在使用14B模型时,可以使用sequential_cpu_offload节省显存,进行生成。)" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 92, + "type": "FunTextBox", + "pos": [ + 254, + -46 + ], + "size": [ + 380.845703125, + 157.68350219726562 + ], + "flags": {}, + "order": 4, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "prompt", + "type": "STRING_PROMPT", + "slot_index": 0, + "links": [ + 73 + ] + } + ], + "title": "Positive Prompt(正向提示词)", + "properties": { + "Node name for S&R": "FunTextBox" + }, + "widgets_values": [ + "在这个阳光明媚的户外花园里,美女身穿一袭及膝的白色无袖连衣裙,裙摆在她轻盈的舞姿中轻柔地摆动,宛如一只翩翩起舞的蝴蝶。阳光透过树叶间洒下斑驳的光影,映衬出她柔和的脸庞和清澈的眼眸,显得格外优雅。仿佛每一个动作都在诉说着青春与活力,她在草地上旋转,裙摆随之飞扬,仿佛整个花园都因她的舞动而欢愉。周围五彩缤纷的花朵在微风中摇曳,玫瑰、菊花、百合,各自释放出阵阵香气,营造出一种轻松而愉快的氛围。" + ] + }, + { + "id": 99, + "type": "VHS_VideoCombine", + "pos": [ + 1094, + 559 + ], + "size": [ + 315, + 851.4834594726562 + ], + "flags": {}, + "order": 9, + "mode": 0, + "inputs": [ + { + "name": "images", + "shape": 7, + "type": "IMAGE", + "link": 65 + }, + { + "name": "audio", + "shape": 7, + "type": "AUDIO", + "link": null + }, + { + "name": "meta_batch", + "shape": 7, + "type": "VHS_BatchManager", + "link": null + }, + { + "name": "vae", + "shape": 7, + "type": "VAE", + "link": null + } + ], + "outputs": [ + { + "name": "Filenames", + "type": "VHS_FILENAMES", + "links": null + } + ], + "properties": { + "Node name for S&R": "VHS_VideoCombine" + }, + "widgets_values": { + "frame_rate": 16, + "loop_count": 0, + "filename_prefix": "Fun-Preprocess-Video", + "format": "video/h264-mp4", + "pix_fmt": "yuv420p", + "crf": 19, + "save_metadata": true, + "pingpong": false, + "save_output": true, + "videopreview": { + "hidden": false, + "paused": false, + "params": { + "filename": "AnimateDiff_00001.gif", + "subfolder": "", + "type": "output", + "format": "image/gif", + "frame_rate": 8 + } + } + } + }, + { + "id": 94, + "type": "FunTextBox", + "pos": [ + 258, + 178 + ], + "size": [ + 368.5529479980469, + 159.4075927734375 + ], + "flags": {}, + "order": 5, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "prompt", + "type": "STRING_PROMPT", + "slot_index": 0, + "links": [ + 74 + ] + } + ], + "properties": { + "Node name for S&R": "FunTextBox" + }, + "widgets_values": [ + "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走" + ] + }, + { + "id": 97, + "type": "VideoToCanny", + "pos": [ + 729, + 566 + ], + "size": [ + 315, + 106 + ], + "flags": {}, + "order": 8, + "mode": 0, + "inputs": [ + { + "name": "input_video", + "type": "IMAGE", + "link": 66 + } + ], + "outputs": [ + { + "name": "images", + "type": "IMAGE", + "slot_index": 0, + "links": [ + 65, + 75 + ] + } + ], + "properties": { + "Node name for S&R": "VideoToCanny" + }, + "widgets_values": [ + 100, + 200, + 81 + ] + }, + { + "id": 101, + "type": "LoadWan2_2FunModel", + "pos": [ + 306.2298583984375, + -313.5478515625 + ], + "size": [ + 397.9565734863281, + 154 + ], + "flags": {}, + "order": 6, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "funmodels", + "type": "FunModels", + "links": [ + 78 + ] + } + ], + "properties": { + "Node name for S&R": "LoadWan2_2FunModel" + }, + "widgets_values": [ + "Wan2.2-Fun-A14B-Control", + "Control", + "sequential_cpu_offload", + "wan2.2/wan_civitai_i2v.yaml", + "bf16" + ] + }, + { + "id": 17, + "type": "VHS_VideoCombine", + "pos": [ + 1465.6671142578125, + -174.95152282714844 + ], + "size": [ + 390.9534912109375, + 966.9860229492188 + ], + "flags": {}, + "order": 11, + "mode": 0, + "inputs": [ + { + "label": "图像", + "name": "images", + "shape": 7, + "type": "IMAGE", + "link": 77 + }, + { + "label": "音频", + "name": "audio", + "shape": 7, + "type": "AUDIO", + "link": null + }, + { + "label": "批次管理", + "name": "meta_batch", + "shape": 7, + "type": "VHS_BatchManager", + "link": null + }, + { + "name": "vae", + "shape": 7, + "type": "VAE", + "link": null + } + ], + "outputs": [ + { + "label": "文件名", + "name": "Filenames", + "type": "VHS_FILENAMES", + "slot_index": 0, + "links": null + } + ], + "properties": { + "Node name for S&R": "VHS_VideoCombine" + }, + "widgets_values": { + "frame_rate": 16, + "loop_count": 0, + "filename_prefix": "Fun", + "format": "video/h264-mp4", + "pix_fmt": "yuv420p", + "crf": 22, + "save_metadata": true, + "pingpong": false, + "save_output": true, + "videopreview": { + "hidden": false, + "paused": false, + "params": { + "filename": "Fun_00004.mp4", + "subfolder": "", + "type": "output", + "format": "video/h264-mp4", + "frame_rate": 16 + } + } + } + }, + { + "id": 85, + "type": "VHS_LoadVideo", + "pos": [ + 335, + 476 + ], + "size": [ + 252.056640625, + 409.87884521484375 + ], + "flags": {}, + "order": 7, + "mode": 0, + "inputs": [ + { + "name": "meta_batch", + "shape": 7, + "type": "VHS_BatchManager", + "link": null + }, + { + "name": "vae", + "shape": 7, + "type": "VAE", + "link": null + } + ], + "outputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "slot_index": 0, + "links": [ + 66 + ] + }, + { + "name": "frame_count", + "type": "INT", + "links": null + }, + { + "name": "audio", + "type": "AUDIO", + "links": null + }, + { + "name": "video_info", + "type": "VHS_VIDEOINFO", + "links": null + } + ], + "properties": { + "Node name for S&R": "VHS_LoadVideo" + }, + "widgets_values": { + "video": "00000005.mp4", + "force_rate": 0, + "force_size": "Disabled", + "custom_width": 512, + "custom_height": 512, + "frame_load_cap": 0, + "skip_first_frames": 0, + "select_every_nth": 1, + "choose video to upload": "image", + "videopreview": { + "hidden": false, + "paused": false, + "params": { + "frame_load_cap": 0, + "skip_first_frames": 0, + "force_rate": 0, + "filename": "00000005.mp4", + "type": "input", + "format": "video/mp4", + "select_every_nth": 1 + } + } + } + }, + { + "id": 102, + "type": "Wan2_2FunV2VSampler", + "pos": [ + 848.4100341796875, + -180.23484802246094 + ], + "size": [ + 458.36383056640625, + 526 + ], + "flags": {}, + "order": 10, + "mode": 0, + "inputs": [ + { + "name": "funmodels", + "type": "FunModels", + "link": 78 + }, + { + "name": "prompt", + "type": "STRING_PROMPT", + "link": 73 + }, + { + "name": "negative_prompt", + "type": "STRING_PROMPT", + "link": 74 + }, + { + "name": "validation_video", + "shape": 7, + "type": "IMAGE", + "link": null + }, + { + "name": "control_video", + "shape": 7, + "type": "IMAGE", + "link": 75 + }, + { + "name": "start_image", + "shape": 7, + "type": "IMAGE", + "link": null + }, + { + "name": "end_image", + "shape": 7, + "type": "IMAGE", + "link": null + }, + { + "name": "ref_image", + "shape": 7, + "type": "IMAGE", + "link": null + }, + { + "name": "camera_conditions", + "shape": 7, + "type": "STRING", + "link": null + }, + { + "name": "riflex_k", + "shape": 7, + "type": "RIFLEXT_ARGS", + "link": null + } + ], + "outputs": [ + { + "name": "images", + "type": "IMAGE", + "links": [ + 77 + ] + } + ], + "properties": { + "Node name for S&R": "Wan2_2FunV2VSampler" + }, + "widgets_values": [ + 81, + 640, + 43, + "fixed", + 50, + 6.000000000000001, + 1, + "Flow", + 0.1, + true, + 5, + true, + 0 + ] + } + ], + "links": [ + [ + 65, + 97, + 0, + 99, + 0, + "IMAGE" + ], + [ + 66, + 85, + 0, + 97, + 0, + "IMAGE" + ], + [ + 73, + 92, + 0, + 102, + 1, + "STRING_PROMPT" + ], + [ + 74, + 94, + 0, + 102, + 2, + "STRING_PROMPT" + ], + [ + 75, + 97, + 0, + 102, + 4, + "IMAGE" + ], + [ + 77, + 102, + 0, + 17, + 0, + "IMAGE" + ], + [ + 78, + 101, + 0, + 102, + 0, + "FunModels" + ] + ], + "groups": [ + { + "id": 1, + "title": "Upload Your Video", + "bounding": [ + 218, + 385, + 487, + 789 + ], + "color": "#a1309b", + "font_size": 24, + "flags": {} + }, + { + "id": 2, + "title": "Load Model", + "bounding": [ + 218, + -387, + 542, + 248 + ], + "color": "#b06634", + "font_size": 24, + "flags": {} + }, + { + "id": 3, + "title": "Prompts", + "bounding": [ + 218, + -127, + 450, + 483 + ], + "color": "#3f789e", + "font_size": 24, + "flags": {} + } + ], + "config": {}, + "extra": { + "ds": { + "scale": 0.6830134553650709, + "offset": [ + 192.62803696741028, + 413.1926211906501 + ] + }, + "workspace_info": { + "id": "776b62b4-bd17-4ed3-9923-b7aad000b1ea" + }, + "node_versions": { + "CogVideoX-Fun": "e935795d7684aaba72be65d2ad5ef580c0526f7a", + "ComfyUI-VideoHelperSuite": "70faa9bcef65932ab72e7404d6373fb300013a2e" + }, + "frontendVersion": "1.21.3" + }, + "version": 0.4 +} \ No newline at end of file diff --git a/comfyui/wan2_2_fun/v1/wan2.2_fun_workflow_v2v_control_depth.json b/comfyui/wan2_2_fun/v1/wan2.2_fun_workflow_v2v_control_depth.json new file mode 100644 index 0000000..4ca7899 --- /dev/null +++ b/comfyui/wan2_2_fun/v1/wan2.2_fun_workflow_v2v_control_depth.json @@ -0,0 +1,697 @@ +{ + "id": "90329caf-a94f-48a6-80d7-d6e167a8b1e3", + "revision": 0, + "last_node_id": 103, + "last_link_id": 83, + "nodes": [ + { + "id": 78, + "type": "Note", + "pos": [ + 18, + -46 + ], + "size": [ + 210, + 88 + ], + "flags": {}, + "order": 0, + "mode": 0, + "inputs": [], + "outputs": [], + "properties": { + "text": "" + }, + "widgets_values": [ + "You can write prompt here\n(你可以在此填写提示词)" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 79, + "type": "Note", + "pos": [ + 15.739953994750977, + 462.38665771484375 + ], + "size": [ + 210, + 88 + ], + "flags": {}, + "order": 1, + "mode": 0, + "inputs": [], + "outputs": [], + "properties": { + "text": "" + }, + "widgets_values": [ + "You can upload video here\n(在此上传视频)" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 88, + "type": "Note", + "pos": [ + -99, + 197 + ], + "size": [ + 326.1556091308594, + 145.20904541015625 + ], + "flags": {}, + "order": 2, + "mode": 0, + "inputs": [], + "outputs": [], + "properties": { + "text": "" + }, + "widgets_values": [ + "Using longer neg prompt such as \"Blurring, mutation, deformation, distortion, dark and solid, comics.\" can increase stability. Adding words such as \"quiet, solid\" to the neg prompt can increase dynamism.\n(使用更长的neg prompt如\"模糊,突变,变形,失真,画面暗,画面固定,连环画,漫画,线稿,没有主体。\",可以增加稳定性。在neg prompt中添加\"安静,固定\"等词语可以增加动态性。)" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 89, + "type": "Note", + "pos": [ + -192, + -293 + ], + "size": [ + 427.074951171875, + 143.9142608642578 + ], + "flags": {}, + "order": 3, + "mode": 0, + "inputs": [], + "outputs": [], + "properties": { + "text": "" + }, + "widgets_values": [ + "When using the 14B model, you can use sequential_cpu_offload to save GPU memory during generation.\n(在使用14B模型时,可以使用sequential_cpu_offload节省显存,进行生成。)" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 92, + "type": "FunTextBox", + "pos": [ + 254, + -46 + ], + "size": [ + 380.845703125, + 157.68350219726562 + ], + "flags": {}, + "order": 4, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "prompt", + "type": "STRING_PROMPT", + "slot_index": 0, + "links": [ + 73 + ] + } + ], + "title": "Positive Prompt(正向提示词)", + "properties": { + "Node name for S&R": "FunTextBox" + }, + "widgets_values": [ + "在这个阳光明媚的户外花园里,美女身穿一袭及膝的白色无袖连衣裙,裙摆在她轻盈的舞姿中轻柔地摆动,宛如一只翩翩起舞的蝴蝶。阳光透过树叶间洒下斑驳的光影,映衬出她柔和的脸庞和清澈的眼眸,显得格外优雅。仿佛每一个动作都在诉说着青春与活力,她在草地上旋转,裙摆随之飞扬,仿佛整个花园都因她的舞动而欢愉。周围五彩缤纷的花朵在微风中摇曳,玫瑰、菊花、百合,各自释放出阵阵香气,营造出一种轻松而愉快的氛围。" + ] + }, + { + "id": 99, + "type": "VHS_VideoCombine", + "pos": [ + 1094, + 559 + ], + "size": [ + 315, + 851.4834594726562 + ], + "flags": {}, + "order": 9, + "mode": 0, + "inputs": [ + { + "name": "images", + "shape": 7, + "type": "IMAGE", + "link": 81 + }, + { + "name": "audio", + "shape": 7, + "type": "AUDIO", + "link": null + }, + { + "name": "meta_batch", + "shape": 7, + "type": "VHS_BatchManager", + "link": null + }, + { + "name": "vae", + "shape": 7, + "type": "VAE", + "link": null + } + ], + "outputs": [ + { + "name": "Filenames", + "type": "VHS_FILENAMES", + "links": null + } + ], + "properties": { + "Node name for S&R": "VHS_VideoCombine" + }, + "widgets_values": { + "frame_rate": 16, + "loop_count": 0, + "filename_prefix": "Fun-Preprocess-Video", + "format": "video/h264-mp4", + "pix_fmt": "yuv420p", + "crf": 19, + "save_metadata": true, + "pingpong": false, + "save_output": true, + "videopreview": { + "hidden": false, + "paused": false, + "params": { + "filename": "AnimateDiff_00001.gif", + "subfolder": "", + "type": "output", + "format": "image/gif", + "frame_rate": 8 + } + } + } + }, + { + "id": 94, + "type": "FunTextBox", + "pos": [ + 258, + 178 + ], + "size": [ + 368.5529479980469, + 159.4075927734375 + ], + "flags": {}, + "order": 5, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "prompt", + "type": "STRING_PROMPT", + "slot_index": 0, + "links": [ + 74 + ] + } + ], + "properties": { + "Node name for S&R": "FunTextBox" + }, + "widgets_values": [ + "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走" + ] + }, + { + "id": 101, + "type": "LoadWan2_2FunModel", + "pos": [ + 306.2298583984375, + -313.5478515625 + ], + "size": [ + 397.9565734863281, + 154 + ], + "flags": {}, + "order": 6, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "funmodels", + "type": "FunModels", + "links": [ + 78 + ] + } + ], + "properties": { + "Node name for S&R": "LoadWan2_2FunModel" + }, + "widgets_values": [ + "Wan2.2-Fun-A14B-Control", + "Control", + "sequential_cpu_offload", + "wan2.2/wan_civitai_i2v.yaml", + "bf16" + ] + }, + { + "id": 17, + "type": "VHS_VideoCombine", + "pos": [ + 1465.6671142578125, + -174.95152282714844 + ], + "size": [ + 390.9534912109375, + 966.9860229492188 + ], + "flags": {}, + "order": 11, + "mode": 0, + "inputs": [ + { + "label": "图像", + "name": "images", + "shape": 7, + "type": "IMAGE", + "link": 77 + }, + { + "label": "音频", + "name": "audio", + "shape": 7, + "type": "AUDIO", + "link": null + }, + { + "label": "批次管理", + "name": "meta_batch", + "shape": 7, + "type": "VHS_BatchManager", + "link": null + }, + { + "name": "vae", + "shape": 7, + "type": "VAE", + "link": null + } + ], + "outputs": [ + { + "label": "文件名", + "name": "Filenames", + "type": "VHS_FILENAMES", + "slot_index": 0, + "links": null + } + ], + "properties": { + "Node name for S&R": "VHS_VideoCombine" + }, + "widgets_values": { + "frame_rate": 16, + "loop_count": 0, + "filename_prefix": "Fun", + "format": "video/h264-mp4", + "pix_fmt": "yuv420p", + "crf": 22, + "save_metadata": true, + "pingpong": false, + "save_output": true, + "videopreview": { + "hidden": false, + "paused": false, + "params": { + "filename": "Fun_00004.mp4", + "subfolder": "", + "type": "output", + "format": "video/h264-mp4", + "frame_rate": 16 + } + } + } + }, + { + "id": 85, + "type": "VHS_LoadVideo", + "pos": [ + 335, + 476 + ], + "size": [ + 252.056640625, + 409.87884521484375 + ], + "flags": {}, + "order": 7, + "mode": 0, + "inputs": [ + { + "name": "meta_batch", + "shape": 7, + "type": "VHS_BatchManager", + "link": null + }, + { + "name": "vae", + "shape": 7, + "type": "VAE", + "link": null + } + ], + "outputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "slot_index": 0, + "links": [ + 82 + ] + }, + { + "name": "frame_count", + "type": "INT", + "links": null + }, + { + "name": "audio", + "type": "AUDIO", + "links": null + }, + { + "name": "video_info", + "type": "VHS_VIDEOINFO", + "links": null + } + ], + "properties": { + "Node name for S&R": "VHS_LoadVideo" + }, + "widgets_values": { + "video": "00000005.mp4", + "force_rate": 0, + "force_size": "Disabled", + "custom_width": 512, + "custom_height": 512, + "frame_load_cap": 0, + "skip_first_frames": 0, + "select_every_nth": 1, + "choose video to upload": "image", + "videopreview": { + "hidden": false, + "paused": false, + "params": { + "frame_load_cap": 0, + "skip_first_frames": 0, + "force_rate": 0, + "filename": "00000005.mp4", + "type": "input", + "format": "video/mp4", + "select_every_nth": 1 + } + } + } + }, + { + "id": 102, + "type": "Wan2_2FunV2VSampler", + "pos": [ + 848.4100341796875, + -180.23484802246094 + ], + "size": [ + 458.36383056640625, + 526 + ], + "flags": {}, + "order": 10, + "mode": 0, + "inputs": [ + { + "name": "funmodels", + "type": "FunModels", + "link": 78 + }, + { + "name": "prompt", + "type": "STRING_PROMPT", + "link": 73 + }, + { + "name": "negative_prompt", + "type": "STRING_PROMPT", + "link": 74 + }, + { + "name": "validation_video", + "shape": 7, + "type": "IMAGE", + "link": null + }, + { + "name": "control_video", + "shape": 7, + "type": "IMAGE", + "link": 83 + }, + { + "name": "start_image", + "shape": 7, + "type": "IMAGE", + "link": null + }, + { + "name": "end_image", + "shape": 7, + "type": "IMAGE", + "link": null + }, + { + "name": "ref_image", + "shape": 7, + "type": "IMAGE", + "link": null + }, + { + "name": "camera_conditions", + "shape": 7, + "type": "STRING", + "link": null + }, + { + "name": "riflex_k", + "shape": 7, + "type": "RIFLEXT_ARGS", + "link": null + } + ], + "outputs": [ + { + "name": "images", + "type": "IMAGE", + "links": [ + 77 + ] + } + ], + "properties": { + "Node name for S&R": "Wan2_2FunV2VSampler" + }, + "widgets_values": [ + 81, + 640, + 43, + "fixed", + 50, + 6.000000000000001, + 1, + "Flow", + 0.1, + true, + 5, + true, + 0 + ] + }, + { + "id": 103, + "type": "VideoToDepth", + "pos": [ + 734.0958862304688, + 474.7727355957031 + ], + "size": [ + 270, + 58 + ], + "flags": {}, + "order": 8, + "mode": 0, + "inputs": [ + { + "name": "input_video", + "type": "IMAGE", + "link": 82 + } + ], + "outputs": [ + { + "name": "images", + "type": "IMAGE", + "links": [ + 81, + 83 + ] + } + ], + "properties": { + "Node name for S&R": "VideoToDepth" + }, + "widgets_values": [ + 81 + ] + } + ], + "links": [ + [ + 73, + 92, + 0, + 102, + 1, + "STRING_PROMPT" + ], + [ + 74, + 94, + 0, + 102, + 2, + "STRING_PROMPT" + ], + [ + 77, + 102, + 0, + 17, + 0, + "IMAGE" + ], + [ + 78, + 101, + 0, + 102, + 0, + "FunModels" + ], + [ + 81, + 103, + 0, + 99, + 0, + "IMAGE" + ], + [ + 82, + 85, + 0, + 103, + 0, + "IMAGE" + ], + [ + 83, + 103, + 0, + 102, + 4, + "IMAGE" + ] + ], + "groups": [ + { + "id": 1, + "title": "Upload Your Video", + "bounding": [ + 218, + 385, + 487, + 789 + ], + "color": "#a1309b", + "font_size": 24, + "flags": {} + }, + { + "id": 2, + "title": "Load Model", + "bounding": [ + 218, + -387, + 542, + 248 + ], + "color": "#b06634", + "font_size": 24, + "flags": {} + }, + { + "id": 3, + "title": "Prompts", + "bounding": [ + 218, + -127, + 450, + 483 + ], + "color": "#3f789e", + "font_size": 24, + "flags": {} + } + ], + "config": {}, + "extra": { + "ds": { + "scale": 0.6830134553650709, + "offset": [ + 192.62803696741028, + 413.1926211906501 + ] + }, + "workspace_info": { + "id": "776b62b4-bd17-4ed3-9923-b7aad000b1ea" + }, + "node_versions": { + "CogVideoX-Fun": "e935795d7684aaba72be65d2ad5ef580c0526f7a", + "ComfyUI-VideoHelperSuite": "70faa9bcef65932ab72e7404d6373fb300013a2e" + }, + "frontendVersion": "1.21.3" + }, + "version": 0.4 +} \ No newline at end of file diff --git a/comfyui/wan2_2_fun/v1/wan2.2_fun_workflow_v2v_control_pose_ref.json b/comfyui/wan2_2_fun/v1/wan2.2_fun_workflow_v2v_control_pose_ref.json new file mode 100644 index 0000000..7bde836 --- /dev/null +++ b/comfyui/wan2_2_fun/v1/wan2.2_fun_workflow_v2v_control_pose_ref.json @@ -0,0 +1,753 @@ +{ + "id": "c5a53a91-88dc-41d4-b414-12e402c59e4c", + "revision": 0, + "last_node_id": 105, + "last_link_id": 91, + "nodes": [ + { + "id": 78, + "type": "Note", + "pos": [ + 18, + -46 + ], + "size": [ + 210, + 88 + ], + "flags": {}, + "order": 0, + "mode": 0, + "inputs": [], + "outputs": [], + "properties": { + "text": "" + }, + "widgets_values": [ + "You can write prompt here\n(你可以在此填写提示词)" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 79, + "type": "Note", + "pos": [ + -111.46612548828125, + 460.2178955078125 + ], + "size": [ + 210, + 88 + ], + "flags": {}, + "order": 1, + "mode": 0, + "inputs": [], + "outputs": [], + "properties": { + "text": "" + }, + "widgets_values": [ + "You can upload video here\n(在此上传视频)" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 88, + "type": "Note", + "pos": [ + -99, + 197 + ], + "size": [ + 326.1556091308594, + 145.20904541015625 + ], + "flags": {}, + "order": 2, + "mode": 0, + "inputs": [], + "outputs": [], + "properties": { + "text": "" + }, + "widgets_values": [ + "Using longer neg prompt such as \"Blurring, mutation, deformation, distortion, dark and solid, comics.\" can increase stability. Adding words such as \"quiet, solid\" to the neg prompt can increase dynamism.\n(使用更长的neg prompt如\"模糊,突变,变形,失真,画面暗,画面固定,连环画,漫画,线稿,没有主体。\",可以增加稳定性。在neg prompt中添加\"安静,固定\"等词语可以增加动态性。)" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 89, + "type": "Note", + "pos": [ + -192, + -293 + ], + "size": [ + 427.074951171875, + 143.9142608642578 + ], + "flags": {}, + "order": 3, + "mode": 0, + "inputs": [], + "outputs": [], + "properties": { + "text": "" + }, + "widgets_values": [ + "When using the 14B model, you can use sequential_cpu_offload to save GPU memory during generation.\n(在使用14B模型时,可以使用sequential_cpu_offload节省显存,进行生成。)" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 99, + "type": "VHS_VideoCombine", + "pos": [ + 1005, + 554 + ], + "size": [ + 315, + 849.46875 + ], + "flags": {}, + "order": 10, + "mode": 0, + "inputs": [ + { + "name": "images", + "shape": 7, + "type": "IMAGE", + "link": 73 + }, + { + "name": "audio", + "shape": 7, + "type": "AUDIO", + "link": null + }, + { + "name": "meta_batch", + "shape": 7, + "type": "VHS_BatchManager", + "link": null + }, + { + "name": "vae", + "shape": 7, + "type": "VAE", + "link": null + } + ], + "outputs": [ + { + "name": "Filenames", + "type": "VHS_FILENAMES", + "links": null + } + ], + "properties": { + "Node name for S&R": "VHS_VideoCombine" + }, + "widgets_values": { + "frame_rate": 16, + "loop_count": 0, + "filename_prefix": "Fun-Preprocess-Video", + "format": "video/h264-mp4", + "pix_fmt": "yuv420p", + "crf": 19, + "save_metadata": true, + "pingpong": false, + "save_output": true, + "videopreview": { + "hidden": false, + "paused": false, + "params": { + "filename": "Fun-Preprocess-Video_00007.mp4", + "subfolder": "", + "type": "output", + "format": "video/h264-mp4", + "frame_rate": 16 + } + } + } + }, + { + "id": 85, + "type": "VHS_LoadVideo", + "pos": [ + 207.79391479492188, + 473.83123779296875 + ], + "size": [ + 252.056640625, + 688.545166015625 + ], + "flags": {}, + "order": 4, + "mode": 0, + "inputs": [ + { + "name": "meta_batch", + "shape": 7, + "type": "VHS_BatchManager", + "link": null + }, + { + "name": "vae", + "shape": 7, + "type": "VAE", + "link": null + } + ], + "outputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "slot_index": 0, + "links": [ + 72 + ] + }, + { + "name": "frame_count", + "type": "INT", + "links": null + }, + { + "name": "audio", + "type": "AUDIO", + "links": null + }, + { + "name": "video_info", + "type": "VHS_VIDEOINFO", + "links": null + } + ], + "properties": { + "Node name for S&R": "VHS_LoadVideo" + }, + "widgets_values": { + "video": "000007.mp4", + "force_rate": 16, + "force_size": "Disabled", + "custom_width": 512, + "custom_height": 512, + "frame_load_cap": 0, + "skip_first_frames": 0, + "select_every_nth": 1, + "choose video to upload": "image", + "videopreview": { + "hidden": false, + "paused": false, + "params": { + "frame_load_cap": 0, + "skip_first_frames": 0, + "force_rate": 16, + "filename": "000007.mp4", + "type": "input", + "format": "video/mp4", + "select_every_nth": 1 + } + } + } + }, + { + "id": 17, + "type": "VHS_VideoCombine", + "pos": [ + 1488, + 8 + ], + "size": [ + 390.9534912109375, + 942.2557983398438 + ], + "flags": {}, + "order": 12, + "mode": 0, + "inputs": [ + { + "label": "图像", + "name": "images", + "shape": 7, + "type": "IMAGE", + "link": 84 + }, + { + "label": "音频", + "name": "audio", + "shape": 7, + "type": "AUDIO", + "link": null + }, + { + "label": "批次管理", + "name": "meta_batch", + "shape": 7, + "type": "VHS_BatchManager", + "link": null + }, + { + "name": "vae", + "shape": 7, + "type": "VAE", + "link": null + } + ], + "outputs": [ + { + "label": "文件名", + "name": "Filenames", + "type": "VHS_FILENAMES", + "slot_index": 0, + "links": null + } + ], + "properties": { + "Node name for S&R": "VHS_VideoCombine" + }, + "widgets_values": { + "frame_rate": 16, + "loop_count": 0, + "filename_prefix": "Fun", + "format": "video/h264-mp4", + "pix_fmt": "yuv420p", + "crf": 22, + "save_metadata": true, + "pingpong": false, + "save_output": true, + "videopreview": { + "hidden": false, + "paused": false, + "params": { + "filename": "Fun_00049.mp4", + "subfolder": "", + "type": "output", + "format": "video/h264-mp4", + "frame_rate": 16 + } + } + } + }, + { + "id": 102, + "type": "LoadImage", + "pos": [ + 553, + 598 + ], + "size": [ + 315, + 314 + ], + "flags": {}, + "order": 5, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "links": [ + 86 + ] + }, + { + "name": "MASK", + "type": "MASK", + "links": null + } + ], + "properties": { + "Node name for S&R": "LoadImage" + }, + "widgets_values": [ + "9.png", + "image" + ] + }, + { + "id": 101, + "type": "VideoToOpenpose", + "pos": [ + 558, + 474 + ], + "size": [ + 315, + 58 + ], + "flags": {}, + "order": 9, + "mode": 0, + "inputs": [ + { + "name": "input_video", + "type": "IMAGE", + "link": 72 + } + ], + "outputs": [ + { + "name": "images", + "type": "IMAGE", + "slot_index": 0, + "links": [ + 73, + 87, + 91 + ] + } + ], + "properties": { + "Node name for S&R": "VideoToOpenpose" + }, + "widgets_values": [ + 81 + ] + }, + { + "id": 104, + "type": "LoadWan2_2FunModel", + "pos": [ + 303.9479064941406, + -309.996337890625 + ], + "size": [ + 390.1385192871094, + 154 + ], + "flags": {}, + "order": 6, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "funmodels", + "type": "FunModels", + "links": [ + 88 + ] + } + ], + "properties": { + "Node name for S&R": "LoadWan2_2FunModel" + }, + "widgets_values": [ + "Wan2.2-Fun-A14B-Control", + "Control", + "sequential_cpu_offload", + "wan2.2/wan_civitai_i2v.yaml", + "bf16" + ] + }, + { + "id": 92, + "type": "FunTextBox", + "pos": [ + 254, + -46 + ], + "size": [ + 380.845703125, + 157.68350219726562 + ], + "flags": {}, + "order": 7, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "prompt", + "type": "STRING_PROMPT", + "slot_index": 0, + "links": [ + 89 + ] + } + ], + "title": "Positive Prompt(正向提示词)", + "properties": { + "Node name for S&R": "FunTextBox" + }, + "widgets_values": [ + "一位动漫风格的女孩。她有着紫色的短发,头上戴着一个黑色和金色相间的蝴蝶结。她的表情显得有些严肃或沉思,眼睛大而有神。女孩穿着一件白色衬衫,外面搭配了一件深蓝色的背心,背心上有一个粉色的蝴蝶结装饰。她的裙子是白色的,裙摆蓬松,整体造型非常可爱且精致。背景是一个简单的圆形图案,颜色为粉红色和灰色相间,给人一种柔和的感觉。整个画面色调柔和,人物形象生动鲜明。" + ] + }, + { + "id": 94, + "type": "FunTextBox", + "pos": [ + 258, + 178 + ], + "size": [ + 368.5529479980469, + 159.4075927734375 + ], + "flags": {}, + "order": 8, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "prompt", + "type": "STRING_PROMPT", + "slot_index": 0, + "links": [ + 90 + ] + } + ], + "properties": { + "Node name for S&R": "FunTextBox" + }, + "widgets_values": [ + "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走" + ] + }, + { + "id": 105, + "type": "Wan2_2FunV2VSampler", + "pos": [ + 987.7456665039062, + -76.86135864257812 + ], + "size": [ + 434.938232421875, + 526 + ], + "flags": {}, + "order": 11, + "mode": 0, + "inputs": [ + { + "name": "funmodels", + "type": "FunModels", + "link": 88 + }, + { + "name": "prompt", + "type": "STRING_PROMPT", + "link": 89 + }, + { + "name": "negative_prompt", + "type": "STRING_PROMPT", + "link": 90 + }, + { + "name": "validation_video", + "shape": 7, + "type": "IMAGE", + "link": 91 + }, + { + "name": "control_video", + "shape": 7, + "type": "IMAGE", + "link": 87 + }, + { + "name": "start_image", + "shape": 7, + "type": "IMAGE", + "link": null + }, + { + "name": "end_image", + "shape": 7, + "type": "IMAGE", + "link": null + }, + { + "name": "ref_image", + "shape": 7, + "type": "IMAGE", + "link": 86 + }, + { + "name": "camera_conditions", + "shape": 7, + "type": "STRING", + "link": null + }, + { + "name": "riflex_k", + "shape": 7, + "type": "RIFLEXT_ARGS", + "link": null + } + ], + "outputs": [ + { + "name": "images", + "type": "IMAGE", + "links": [ + 84 + ] + } + ], + "properties": { + "Node name for S&R": "Wan2_2FunV2VSampler" + }, + "widgets_values": [ + 81, + 640, + 43, + "fixed", + 50, + 6.000000000000001, + 1, + "Flow", + 0.1, + true, + 5, + true, + 0 + ] + } + ], + "links": [ + [ + 72, + 85, + 0, + 101, + 0, + "IMAGE" + ], + [ + 73, + 101, + 0, + 99, + 0, + "IMAGE" + ], + [ + 84, + 105, + 0, + 17, + 0, + "IMAGE" + ], + [ + 86, + 102, + 0, + 105, + 7, + "IMAGE" + ], + [ + 87, + 101, + 0, + 105, + 4, + "IMAGE" + ], + [ + 88, + 104, + 0, + 105, + 0, + "FunModels" + ], + [ + 89, + 92, + 0, + 105, + 1, + "STRING_PROMPT" + ], + [ + 90, + 94, + 0, + 105, + 2, + "STRING_PROMPT" + ], + [ + 91, + 101, + 0, + 105, + 3, + "IMAGE" + ] + ], + "groups": [ + { + "id": 1, + "title": "Upload Your Video And Reference Image", + "bounding": [ + 91, + 383, + 859, + 841 + ], + "color": "#a1309b", + "font_size": 24, + "flags": {} + }, + { + "id": 2, + "title": "Load Model", + "bounding": [ + 218, + -387, + 542, + 248 + ], + "color": "#b06634", + "font_size": 24, + "flags": {} + }, + { + "id": 3, + "title": "Prompts", + "bounding": [ + 218, + -127, + 450, + 483 + ], + "color": "#3f789e", + "font_size": 24, + "flags": {} + } + ], + "config": {}, + "extra": { + "ds": { + "scale": 0.7513148009015781, + "offset": [ + 235.41928184428536, + 445.7655228500244 + ] + }, + "workspace_info": { + "id": "776b62b4-bd17-4ed3-9923-b7aad000b1ea" + }, + "node_versions": { + "ComfyUI-VideoHelperSuite": "70faa9bcef65932ab72e7404d6373fb300013a2e", + "comfy-core": "0.3.44", + "CogVideoX-Fun": "e935795d7684aaba72be65d2ad5ef580c0526f7a" + }, + "frontendVersion": "1.21.3" + }, + "version": 0.4 +} \ No newline at end of file diff --git a/comfyui/wan2_2_fun/v1/wan2.2_fun_workflow_v2v_control_ref.json b/comfyui/wan2_2_fun/v1/wan2.2_fun_workflow_v2v_control_ref.json index fc17c04..e0ba01a 100755 --- a/comfyui/wan2_2_fun/v1/wan2.2_fun_workflow_v2v_control_ref.json +++ b/comfyui/wan2_2_fun/v1/wan2.2_fun_workflow_v2v_control_ref.json @@ -574,7 +574,7 @@ "flags": {} }, { - "title": "Load EasyAnimate", + "title": "Load Model", "bounding": [ 218, -387, diff --git a/examples/wan2.2/post_infer_queue.py b/examples/wan2.2/post_infer_queue.py index bc6e868..2d4f991 100755 --- a/examples/wan2.2/post_infer_queue.py +++ b/examples/wan2.2/post_infer_queue.py @@ -106,9 +106,8 @@ if __name__ == '__main__': # Recommended to be set between 0.05 and 0.30. A larger threshold can cache more steps, speeding up the inference process, # but it may cause slight differences between the generated content and the original content. # # --------------------------------------------------------------------------------------------------- # - # | Model Name | threshold | Model Name | threshold | Model Name | threshold | - # | Wan2.1-T2V-1.3B | 0.05~0.10 | Wan2.1-T2V-14B | 0.10~0.15 | Wan2.1-I2V-14B-720P | 0.20~0.30 | - # | Wan2.1-I2V-14B-480P | 0.20~0.25 | Wan2.1-Fun-*-1.3B-* | 0.05~0.10 | Wan2.1-Fun-*-14B-* | 0.20~0.30 | + # | Model Name | threshold | Model Name | threshold | + # | Wan2.2-T2V-A14B | 0.10~0.15 | Wan2.2-I2V-A14B | 0.15~0.20 | # # --------------------------------------------------------------------------------------------------- # teacache_threshold = 0.10 # The number of steps to skip TeaCache at the beginning of the inference process, which can diff --git a/examples/wan2.2/post_infer_queue_i2v.py b/examples/wan2.2/post_infer_queue_i2v.py index d24d20b..9dd1f43 100755 --- a/examples/wan2.2/post_infer_queue_i2v.py +++ b/examples/wan2.2/post_infer_queue_i2v.py @@ -123,9 +123,8 @@ if __name__ == '__main__': # Recommended to be set between 0.05 and 0.30. A larger threshold can cache more steps, speeding up the inference process, # but it may cause slight differences between the generated content and the original content. # # --------------------------------------------------------------------------------------------------- # - # | Model Name | threshold | Model Name | threshold | Model Name | threshold | - # | Wan2.1-T2V-1.3B | 0.05~0.10 | Wan2.1-T2V-14B | 0.10~0.15 | Wan2.1-I2V-14B-720P | 0.20~0.30 | - # | Wan2.1-I2V-14B-480P | 0.20~0.25 | Wan2.1-Fun-*-1.3B-* | 0.05~0.10 | Wan2.1-Fun-*-14B-* | 0.20~0.30 | + # | Model Name | threshold | Model Name | threshold | + # | Wan2.2-T2V-A14B | 0.10~0.15 | Wan2.2-I2V-A14B | 0.15~0.20 | # # --------------------------------------------------------------------------------------------------- # teacache_threshold = 0.10 # The number of steps to skip TeaCache at the beginning of the inference process, which can diff --git a/examples/wan2.2/predict_ti2v.py b/examples/wan2.2/predict_ti2v.py index 526eae1..f227073 100755 --- a/examples/wan2.2/predict_ti2v.py +++ b/examples/wan2.2/predict_ti2v.py @@ -216,7 +216,7 @@ scheduler = Choosen_Scheduler( # Get Pipeline pipeline = Wan2_2TI2VPipeline( transformer=transformer, - transformer_2=transformer_2 , + transformer_2=transformer_2, vae=vae, tokenizer=tokenizer, text_encoder=text_encoder, diff --git a/examples/wan2.2_fun/app.py b/examples/wan2.2_fun/app.py new file mode 100755 index 0000000..1844be3 --- /dev/null +++ b/examples/wan2.2_fun/app.py @@ -0,0 +1,79 @@ +import os +import sys +import time + +import torch + +current_file_path = os.path.abspath(__file__) +project_roots = [os.path.dirname(current_file_path), os.path.dirname(os.path.dirname(current_file_path)), os.path.dirname(os.path.dirname(os.path.dirname(current_file_path)))] +for project_root in project_roots: + sys.path.insert(0, project_root) if project_root not in sys.path else None + +from videox_fun.api.api import (infer_forward_api, + update_diffusion_transformer_api) +from videox_fun.ui.controller import flow_scheduler_dict +from videox_fun.ui.wan2_2_fun_ui import ui, ui_client, ui_host + +if __name__ == "__main__": + # Choose the ui mode + # "normal" refers to the standard UI, which allows users to click to switch models, change model types, and more. + # "host" represents the hosting mode, where the model is loaded directly at startup and can be accessed via + # the API to return generation results. + # "client" represents the client mode, offering a simple UI that sends requests to a remote API for generation. + ui_mode = "normal" + + # GPU memory mode, which can be choosen in [model_full_load, model_cpu_offload, model_cpu_offload_and_qfloat8, sequential_cpu_offload]. + # model_full_load means that the entire model will be moved to the GPU. + # + # 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. + # + # 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 + weight_dtype = torch.bfloat16 + + # Server ip + server_name = "0.0.0.0" + server_port = 7860 + + # Config path + config_path = "config/wan2.2/wan_civitai_i2v.yaml" + # Params below is used when ui_mode = "host" + # Model path of the pretrained model + model_name = "models/Diffusion_Transformer/Wan2.2-Fun-A14B-InP" + # "Inpaint" or "Control" + model_type = "Inpaint" + + if ui_mode == "host": + 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, compile_dit, weight_dtype) + + def gr_launch(): + # launch gradio + app, _, _ = demo.queue(status_update_rate=1).launch( + server_name=server_name, + server_port=server_port, + prevent_thread_lock=True + ) + + # launch api + infer_forward_api(None, app, controller) + update_diffusion_transformer_api(None, app, controller) + + gr_launch() + + # not close the python + while True: + time.sleep(5) \ No newline at end of file diff --git a/examples/wan2.2_fun/launch_api.py b/examples/wan2.2_fun/launch_api.py new file mode 100755 index 0000000..e1261dc --- /dev/null +++ b/examples/wan2.2_fun/launch_api.py @@ -0,0 +1,91 @@ +import argparse +import os +import sys +import time + +import gradio as gr +import ray +import torch + +current_file_path = os.path.abspath(__file__) +project_roots = [os.path.dirname(current_file_path), os.path.dirname(os.path.dirname(current_file_path)), os.path.dirname(os.path.dirname(os.path.dirname(current_file_path)))] +for project_root in project_roots: + sys.path.insert(0, project_root) if project_root not in sys.path else None + +from videox_fun.api.api_multi_nodes import (MultiNodesEngine, + multi_nodes_infer_forward_api) +from videox_fun.ui.controller import flow_scheduler_dict +from videox_fun.ui.wan2_2_fun_ui import Wan2_2_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, 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') + parser.add_argument('--config_path', type=str, default="config/wan2.1/wan_civitai.yaml", help='Path to config file') + parser.add_argument('--model_name', type=str, default="models/Diffusion_Transformer/Wan2.1-Fun-V1.1-1.3B-InP", help='Model path') + parser.add_argument('--model_type', type=str, default="Inpaint", help='Model type (Inpaint/Control)') + parser.add_argument('--savedir_sample', type=str, default=None, help='The save directory for samples') + args = parser.parse_args() + + weight_dtype = torch.float32 + if args.weight_dtype == "bf16": + weight_dtype = torch.bfloat16 + elif args.weight_dtype == "fp16": + weight_dtype = torch.float16 + + engine = MultiNodesEngine( + world_size=args.world_size, Controller=Wan2_2_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, + 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(): + # launch gradio + with gr.Blocks() as demo: + gr.Markdown("") + app, _, _ = demo.queue(status_update_rate=1).launch( + server_name=args.server_name, + server_port=args.server_port, + prevent_thread_lock=True + ) + + # launch api + multi_nodes_infer_forward_api(None, app, engine) + + gr_launch() + + # not close the python + while True: + time.sleep(5) + +if __name__ == "__main__": + main() \ No newline at end of file diff --git a/examples/wan2.2_fun/post_infer.py b/examples/wan2.2_fun/post_infer.py new file mode 100755 index 0000000..0493c9b --- /dev/null +++ b/examples/wan2.2_fun/post_infer.py @@ -0,0 +1,150 @@ +import base64 +import json +import time +from datetime import datetime + +import requests +import base64 + + +def post_diffusion_transformer(diffusion_transformer_path, url='http://127.0.0.1:7860'): + datas = json.dumps({ + "diffusion_transformer_path": diffusion_transformer_path + }) + r = requests.post(f'{url}/videox_fun/update_diffusion_transformer', data=datas, timeout=1500) + data = r.content.decode('utf-8') + return data + +def post_update_edition(edition, url='http://0.0.0.0:7860'): + datas = json.dumps({ + "edition": edition + }) + r = requests.post(f'{url}/videox_fun/update_edition', data=datas, timeout=1500) + data = r.content.decode('utf-8') + return data + + +def post_infer( + generation_method, + length_slider, + url='http://127.0.0.1:7860', + POST_TOKEN="", + timeout=5000, + base_model_path="none", + lora_model_path="none", + lora_alpha_slider=0.55, + prompt_textbox="A young woman with beautiful and clear eyes and blonde hair standing and white dress in a forest wearing a crown. She seems to be lost in thought, and the camera focuses on her face. The video is of high quality, and the view is very clear. High quality, masterpiece, best quality, highres, ultra-detailed, fantastic.", + negative_prompt_textbox="The video is not of a high quality, it has a low resolution. Watermark present in each frame. The background is solid. Strange body and strange trajectory. Distortion.", + sampler_dropdown="Flow", + sample_step_slider=50, + width_slider=672, + height_slider=384, + cfg_scale_slider=6, + seed_textbox=43 +): + # Prepare the data payload + datas = json.dumps({ + "base_model_path": base_model_path, + "lora_model_path": lora_model_path, + "lora_alpha_slider": lora_alpha_slider, + "prompt_textbox": prompt_textbox, + "negative_prompt_textbox": negative_prompt_textbox, + "sampler_dropdown": sampler_dropdown, + "sample_step_slider": sample_step_slider, + "width_slider": width_slider, + "height_slider": height_slider, + "generation_method": generation_method, + "length_slider": length_slider, + "cfg_scale_slider": cfg_scale_slider, + "seed_textbox": seed_textbox, + }) + + # Initialize session and set headers + session = requests.session() + session.headers.update({"Authorization": POST_TOKEN}) + + # Send POST request + if url[-1] == "/": + url = url[:-1] + post_r = session.post(f'{url}/videox_fun/infer_forward', data=datas, timeout=timeout) + + data = post_r.content.decode('utf-8') + return data + +if __name__ == '__main__': + # initiate time + time_start = time.time() + + # The Url you want to post + POST_URL = 'http://0.0.0.0:7860' + # Used in EAS. If you don't need Authorization, please set it to empty string. + TOKEN = '' + + # -------------------------- # + # Step 1: update edition + # -------------------------- # + # diffusion_transformer_path = "models/Diffusion_Transformer/Wan2.1-Fun-1.3B-InP" + # outputs = post_diffusion_transformer(diffusion_transformer_path) + # print('Output update edition: ', outputs) + + # -------------------------- # + # Step 2: infer + # -------------------------- # + # "Video Generation" and "Image Generation" + generation_method = "Video Generation" + # Video length + length_slider = 49 + # Used in Lora models + lora_model_path = "none" + lora_alpha_slider = 0.55 + # Prompts + prompt_textbox = "A young woman with beautiful and clear eyes and blonde hair standing and white dress in a forest wearing a crown. She seems to be lost in thought, and the camera focuses on her face. The video is of high quality, and the view is very clear. High quality, masterpiece, best quality, highres, ultra-detailed, fantastic." + negative_prompt_textbox = "The video is not of a high quality, it has a low resolution. Watermark present in each frame. The background is solid. Strange body and strange trajectory. Distortion." + # Sampler name + sampler_dropdown = "Flow" + # Sampler steps + sample_step_slider = 50 + # height and width + width_slider = 832 + height_slider = 480 + # cfg scale + cfg_scale_slider = 6 + seed_textbox = 43 + + outputs = post_infer( + generation_method, + length_slider, + lora_model_path=lora_model_path, + lora_alpha_slider=lora_alpha_slider, + prompt_textbox=prompt_textbox, + negative_prompt_textbox=negative_prompt_textbox, + sampler_dropdown=sampler_dropdown, + sample_step_slider=sample_step_slider, + width_slider=width_slider, + height_slider=height_slider, + cfg_scale_slider=cfg_scale_slider, + seed_textbox=seed_textbox, + url=POST_URL, + POST_TOKEN=TOKEN + ) + + # Get decoded data + outputs = json.loads(outputs) + base64_encoding = outputs["base64_encoding"] + decoded_data = base64.b64decode(base64_encoding) + + is_image = True if generation_method == "Image Generation" else False + if is_image or length_slider == 1: + file_path = "1.png" + else: + file_path = "1.mp4" + with open(file_path, "wb") as file: + file.write(decoded_data) + + # End of record time + # The calculated time difference is the execution time of the program, expressed in seconds / s + time_end = time.time() + time_sum = (time_end - time_start) + print('# --------------------------------------------------------- #') + print(f'# Total expenditure: {time_sum}s') + print('# --------------------------------------------------------- #') \ No newline at end of file diff --git a/examples/wan2.2_fun/post_infer_queue.py b/examples/wan2.2_fun/post_infer_queue.py new file mode 100755 index 0000000..1ece062 --- /dev/null +++ b/examples/wan2.2_fun/post_infer_queue.py @@ -0,0 +1,192 @@ +import base64 +import json +import time +import urllib.parse +import requests + + +def post_infer( + generation_method, + length_slider, + url='http://127.0.0.1:7860', + POST_TOKEN="", + timeout=5, + base_model_path="none", + lora_model_path="none", + lora_alpha_slider=0.55, + prompt_textbox="A young woman with beautiful and clear eyes and blonde hair standing and white dress in a forest wearing a crown. She seems to be lost in thought, and the camera focuses on her face. The video is of high quality, and the view is very clear. High quality, masterpiece, best quality, highres, ultra-detailed, fantastic.", + negative_prompt_textbox="色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走", + sampler_dropdown="Flow", + sample_step_slider=50, + width_slider=672, + height_slider=384, + cfg_scale_slider=6, + seed_textbox=43, + enable_teacache = None, + teacache_threshold = None, + num_skip_start_steps = None, + teacache_offload = None, + cfg_skip_ratio = None, + enable_riflex = None, + riflex_k = None, +): + # Prepare the data payload + datas = json.dumps({ + "base_model_path": base_model_path, + "lora_model_path": lora_model_path, + "lora_alpha_slider": lora_alpha_slider, + "prompt_textbox": prompt_textbox, + "negative_prompt_textbox": negative_prompt_textbox, + "sampler_dropdown": sampler_dropdown, + "sample_step_slider": sample_step_slider, + "width_slider": width_slider, + "height_slider": height_slider, + "generation_method": generation_method, + "length_slider": length_slider, + "cfg_scale_slider": cfg_scale_slider, + "seed_textbox": seed_textbox, + + "enable_teacache": enable_teacache, + "teacache_threshold": teacache_threshold, + "num_skip_start_steps": num_skip_start_steps, + "teacache_offload": teacache_offload, + "cfg_skip_ratio": cfg_skip_ratio, + "enable_riflex": enable_riflex, + "riflex_k": riflex_k, + }) + + # Initialize session and set headers + session = requests.session() + session.headers.update({"Authorization": POST_TOKEN}) + + # Send POST request + if url[-1] == "/": + url = url[:-1] + post_r = session.post(f'{url}/videox_fun/infer_forward', data=datas, timeout=timeout) + + # Extract request ID from POST response headers + request_id = post_r.headers.get("X-Eas-Queueservice-Request-Id") + + # Prepare query parameters for GET request + query = { + '_index_': '0', + '_length_': '1', + '_timeout_': str(timeout), + '_raw_': 'false', + '_auto_delete_': 'true', + } + if request_id: + query['requestId'] = request_id + + query_str = urllib.parse.urlencode(query) + + # Polling GET request until status code is not 204 + status_code = 204 + while status_code == 204: + if query_str: + get_r = session.get(f'{url}/sink?{query_str}', timeout=timeout) + else: + get_r = session.get(f'{url}/sink', timeout=timeout) + status_code = get_r.status_code + # Decode and return the response content + data = get_r.content.decode('utf-8') + return data + +if __name__ == '__main__': + # initiate time + time_start = time.time() + + # EAS队列配置 + EAS_URL = 'http://17xxxxxxxxx.pai-eas.aliyuncs.com/api/predict/xxxxxxxx' + # Use in EAS Queue + TOKEN = 'xxxxxxxx' + + # Support TeaCache. + enable_teacache = True + # Recommended to be set between 0.05 and 0.30. A larger threshold can cache more steps, speeding up the inference process, + # but it may cause slight differences between the generated content and the original content. + # # --------------------------------------------------------------------------------------------------- # + # | Model Name | threshold | Model Name | threshold | + # | Wan2.2-T2V-A14B | 0.10~0.15 | Wan2.2-I2V-A14B | 0.15~0.20 | + # | Wan2.2-Fun-A14B-* | 0.15~0.20 | + # # --------------------------------------------------------------------------------------------------- # + teacache_threshold = 0.10 + # The number of steps to skip TeaCache at the beginning of the inference process, which can + # reduce the impact of TeaCache on generated video quality. + num_skip_start_steps = 5 + # Whether to offload TeaCache tensors to cpu to save a little bit of GPU memory. + teacache_offload = False + + # Skip some cfg steps in inference + # Recommended to be set between 0.00 and 0.25 + cfg_skip_ratio = 0 + + # Riflex config + enable_riflex = False + # Index of intrinsic frequency + riflex_k = 6 + + # "Video Generation" and "Image Generation" + generation_method = "Video Generation" + # Video length + length_slider = 81 + # Used in Lora models + lora_model_path = "none" + lora_alpha_slider = 0.55 + # Prompts + prompt_textbox = "A young woman with beautiful and clear eyes and blonde hair standing and white dress in a forest wearing a crown. She seems to be lost in thought, and the camera focuses on her face. The video is of high quality, and the view is very clear. High quality, masterpiece, best quality, highres, ultra-detailed, fantastic." + negative_prompt_textbox = "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走" + # Sampler name + sampler_dropdown = "Flow" + # Sampler steps + sample_step_slider = 50 + # height and width + width_slider = 832 + height_slider = 480 + # cfg scale + cfg_scale_slider = 6 + seed_textbox = 43 + + outputs = post_infer( + generation_method, + length_slider, + lora_model_path=lora_model_path, + lora_alpha_slider=lora_alpha_slider, + prompt_textbox=prompt_textbox, + negative_prompt_textbox=negative_prompt_textbox, + sampler_dropdown=sampler_dropdown, + sample_step_slider=sample_step_slider, + width_slider=width_slider, + height_slider=height_slider, + cfg_scale_slider=cfg_scale_slider, + seed_textbox=seed_textbox, + enable_teacache = enable_teacache, + teacache_threshold = teacache_threshold, + num_skip_start_steps = num_skip_start_steps, + teacache_offload = teacache_offload, + cfg_skip_ratio = cfg_skip_ratio, + enable_riflex = enable_riflex, + riflex_k = riflex_k, + url=EAS_URL, + POST_TOKEN=TOKEN + ) + # Get decoded data + outputs = json.loads(base64.b64decode(json.loads(outputs)[0]['data'])) + base64_encoding = outputs["base64_encoding"] + decoded_data = base64.b64decode(base64_encoding) + + is_image = True if generation_method == "Image Generation" else False + if is_image or length_slider == 1: + file_path = "1.png" + else: + file_path = "1.mp4" + with open(file_path, "wb") as file: + file.write(decoded_data) + + # End of record time + # The calculated time difference is the execution time of the program, expressed in seconds / s + time_end = time.time() + time_sum = (time_end - time_start) + print('# --------------------------------------------------------- #') + print(f'# Total expenditure: {time_sum}s') + print('# --------------------------------------------------------- #') \ No newline at end of file diff --git a/examples/wan2.2_fun/post_infer_queue_i2v.py b/examples/wan2.2_fun/post_infer_queue_i2v.py new file mode 100755 index 0000000..0e54e90 --- /dev/null +++ b/examples/wan2.2_fun/post_infer_queue_i2v.py @@ -0,0 +1,213 @@ +import base64 +import json +import time +import urllib.parse +import requests +from PIL import Image +from io import BytesIO + + +def post_infer( + generation_method, + length_slider, + url='http://127.0.0.1:7860', + POST_TOKEN="", + timeout=5, + base_model_path="none", + lora_model_path="none", + lora_alpha_slider=0.55, + prompt_textbox="A young woman with beautiful and clear eyes and blonde hair standing and white dress in a forest wearing a crown. She seems to be lost in thought, and the camera focuses on her face. The video is of high quality, and the view is very clear. High quality, masterpiece, best quality, highres, ultra-detailed, fantastic.", + negative_prompt_textbox="色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走", + sampler_dropdown="Flow", + sample_step_slider=50, + width_slider=672, + height_slider=384, + cfg_scale_slider=6, + seed_textbox=43, + enable_teacache = None, + teacache_threshold = None, + num_skip_start_steps = None, + teacache_offload = None, + cfg_skip_ratio = None, + enable_riflex = None, + riflex_k = None, + start_image = None +): + if start_image: + try: + if not start_image.startswith("http"): + image = Image.open(start_image).convert("RGB") + # 将图片转换为 Base64 编码 + buffered = BytesIO() + image.save(buffered, format="JPEG") + start_image = base64.b64encode(buffered.getvalue()).decode('utf-8') + except Exception as e: + print(f"Error processing start_image: {e}") + raise + + # Prepare the data payload + datas = json.dumps({ + "base_model_path": base_model_path, + "lora_model_path": lora_model_path, + "lora_alpha_slider": lora_alpha_slider, + "prompt_textbox": prompt_textbox, + "negative_prompt_textbox": negative_prompt_textbox, + "sampler_dropdown": sampler_dropdown, + "sample_step_slider": sample_step_slider, + "width_slider": width_slider, + "height_slider": height_slider, + "generation_method": generation_method, + "length_slider": length_slider, + "cfg_scale_slider": cfg_scale_slider, + "seed_textbox": seed_textbox, + + "enable_teacache": enable_teacache, + "teacache_threshold": teacache_threshold, + "num_skip_start_steps": num_skip_start_steps, + "teacache_offload": teacache_offload, + "cfg_skip_ratio": cfg_skip_ratio, + "enable_riflex": enable_riflex, + "riflex_k": riflex_k, + "start_image": start_image + }) + + # Initialize session and set headers + session = requests.session() + session.headers.update({"Authorization": POST_TOKEN}) + + # Send POST request + if url[-1] == "/": + url = url[:-1] + post_r = session.post(f'{url}/videox_fun/infer_forward', data=datas, timeout=timeout) + + # Extract request ID from POST response headers + request_id = post_r.headers.get("X-Eas-Queueservice-Request-Id") + + # Prepare query parameters for GET request + query = { + '_index_': '0', + '_length_': '1', + '_timeout_': str(timeout), + '_raw_': 'false', + '_auto_delete_': 'true', + } + if request_id: + query['requestId'] = request_id + + query_str = urllib.parse.urlencode(query) + + # Polling GET request until status code is not 204 + status_code = 204 + while status_code == 204: + if query_str: + get_r = session.get(f'{url}/sink?{query_str}', timeout=timeout) + else: + get_r = session.get(f'{url}/sink', timeout=timeout) + status_code = get_r.status_code + # Decode and return the response content + data = get_r.content.decode('utf-8') + return data + + +if __name__ == '__main__': + # initiate time + time_start = time.time() + + # EAS队列配置 + EAS_URL = 'http://17xxxxxxxxx.pai-eas.aliyuncs.com/api/predict/xxxxxxxx' + # Use in EAS Queue + TOKEN = 'xxxxxxxx' + + # Support TeaCache. + enable_teacache = True + # Recommended to be set between 0.05 and 0.30. A larger threshold can cache more steps, speeding up the inference process, + # but it may cause slight differences between the generated content and the original content. + # # --------------------------------------------------------------------------------------------------- # + # | Model Name | threshold | Model Name | threshold | + # | Wan2.2-T2V-A14B | 0.10~0.15 | Wan2.2-I2V-A14B | 0.15~0.20 | + # | Wan2.2-Fun-A14B-* | 0.15~0.20 | + # # --------------------------------------------------------------------------------------------------- # + teacache_threshold = 0.10 + # The number of steps to skip TeaCache at the beginning of the inference process, which can + # reduce the impact of TeaCache on generated video quality. + num_skip_start_steps = 5 + # Whether to offload TeaCache tensors to cpu to save a little bit of GPU memory. + teacache_offload = False + + # Skip some cfg steps in inference + # Recommended to be set between 0.00 and 0.25 + cfg_skip_ratio = 0 + + # Riflex config + enable_riflex = False + # Index of intrinsic frequency + riflex_k = 6 + + # "Video Generation" and "Image Generation" + generation_method = "Video Generation" + # Video length + length_slider = 81 + # Used in Lora models + lora_model_path = "none" + lora_alpha_slider = 0.55 + # Prompts + prompt_textbox = "一只棕色的狗摇着头,坐在舒适房间里的浅色沙发上。在狗的后面,架子上有一幅镶框的画,周围是粉红色的花朵。房间里柔和温暖的灯光营造出舒适的氛围。" + negative_prompt_textbox = "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走" + # Sampler name + sampler_dropdown = "Flow" + # Sampler steps + sample_step_slider = 50 + # height and width + width_slider = 832 + height_slider = 480 + # cfg scale + cfg_scale_slider = 6 + seed_textbox = 43 + + # 起始图片路径 + start_image_path = "asset/1.png" # 替换为实际的图片路径 + + outputs = post_infer( + generation_method, + length_slider, + lora_model_path=lora_model_path, + lora_alpha_slider=lora_alpha_slider, + prompt_textbox=prompt_textbox, + negative_prompt_textbox=negative_prompt_textbox, + sampler_dropdown=sampler_dropdown, + sample_step_slider=sample_step_slider, + width_slider=width_slider, + height_slider=height_slider, + cfg_scale_slider=cfg_scale_slider, + seed_textbox=seed_textbox, + enable_teacache = enable_teacache, + teacache_threshold = teacache_threshold, + num_skip_start_steps = num_skip_start_steps, + teacache_offload = teacache_offload, + cfg_skip_ratio = cfg_skip_ratio, + enable_riflex = enable_riflex, + riflex_k = riflex_k, + url=EAS_URL, + POST_TOKEN=TOKEN, + start_image=start_image_path # 传递起始图片路径 + ) + # Get decoded data + outputs = json.loads(base64.b64decode(json.loads(outputs)[0]['data'])) + base64_encoding = outputs["base64_encoding"] + decoded_data = base64.b64decode(base64_encoding) + + is_image = True if generation_method == "Image Generation" else False + if is_image or length_slider == 1: + file_path = "1.png" + else: + file_path = "1.mp4" + with open(file_path, "wb") as file: + file.write(decoded_data) + + # End of record time + # The calculated time difference is the execution time of the program, expressed in seconds / s + time_end = time.time() + time_sum = (time_end - time_start) + print('# --------------------------------------------------------- #') + print(f'# Total expenditure: {time_sum}s') + print('# --------------------------------------------------------- #') diff --git a/examples/wan2.2_fun/post_infer_queue_v2v_control.py b/examples/wan2.2_fun/post_infer_queue_v2v_control.py new file mode 100755 index 0000000..a97a6e6 --- /dev/null +++ b/examples/wan2.2_fun/post_infer_queue_v2v_control.py @@ -0,0 +1,230 @@ +import base64 +import json +import time +import urllib.parse +from io import BytesIO + +import requests +from PIL import Image + + +def post_infer( + generation_method, + length_slider, + url='http://127.0.0.1:7860', + POST_TOKEN="", + timeout=5, + base_model_path="none", + lora_model_path="none", + lora_alpha_slider=0.55, + prompt_textbox="A young woman with beautiful and clear eyes and blonde hair standing and white dress in a forest wearing a crown. She seems to be lost in thought, and the camera focuses on her face. The video is of high quality, and the view is very clear. High quality, masterpiece, best quality, highres, ultra-detailed, fantastic.", + negative_prompt_textbox="色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走", + sampler_dropdown="Flow", + sample_step_slider=50, + width_slider=672, + height_slider=384, + cfg_scale_slider=6, + seed_textbox=43, + enable_teacache = None, + teacache_threshold = None, + num_skip_start_steps = None, + teacache_offload = None, + cfg_skip_ratio = None, + enable_riflex = None, + riflex_k = None, + control_video = None, + ref_image = None +): + if control_video: + try: + if not control_video.startswith("http"): + with open(control_video, "rb") as file: + video_data = file.read() + + control_video = base64.b64encode(video_data).decode('utf-8') + except Exception as e: + print(f"Error processing control_video: {e}") + raise + + if ref_image: + try: + if not ref_image.startswith("http"): + image = Image.open(ref_image).convert("RGB") + # 将图片转换为 Base64 编码 + buffered = BytesIO() + image.save(buffered, format="JPEG") + ref_image = base64.b64encode(buffered.getvalue()).decode('utf-8') + except Exception as e: + print(f"Error processing ref_image: {e}") + raise + + # Prepare the data payload + datas = json.dumps({ + "base_model_path": base_model_path, + "lora_model_path": lora_model_path, + "lora_alpha_slider": lora_alpha_slider, + "prompt_textbox": prompt_textbox, + "negative_prompt_textbox": negative_prompt_textbox, + "sampler_dropdown": sampler_dropdown, + "sample_step_slider": sample_step_slider, + "width_slider": width_slider, + "height_slider": height_slider, + "generation_method": generation_method, + "length_slider": length_slider, + "cfg_scale_slider": cfg_scale_slider, + "seed_textbox": seed_textbox, + + "ref_image": ref_image, + "enable_teacache": enable_teacache, + "teacache_threshold": teacache_threshold, + "num_skip_start_steps": num_skip_start_steps, + "teacache_offload": teacache_offload, + "cfg_skip_ratio": cfg_skip_ratio, + "enable_riflex": enable_riflex, + "riflex_k": riflex_k, + "control_video": control_video + }) + + # Initialize session and set headers + session = requests.session() + session.headers.update({"Authorization": POST_TOKEN}) + + # Send POST request + if url[-1] == "/": + url = url[:-1] + post_r = session.post(f'{url}/videox_fun/infer_forward', data=datas, timeout=timeout) + + # Extract request ID from POST response headers + request_id = post_r.headers.get("X-Eas-Queueservice-Request-Id") + + # Prepare query parameters for GET request + query = { + '_index_': '0', + '_length_': '1', + '_timeout_': str(timeout), + '_raw_': 'false', + '_auto_delete_': 'true', + } + if request_id: + query['requestId'] = request_id + + query_str = urllib.parse.urlencode(query) + + # Polling GET request until status code is not 204 + status_code = 204 + while status_code == 204: + if query_str: + get_r = session.get(f'{url}/sink?{query_str}', timeout=timeout) + else: + get_r = session.get(f'{url}/sink', timeout=timeout) + status_code = get_r.status_code + # Decode and return the response content + data = get_r.content.decode('utf-8') + return data + + +if __name__ == '__main__': + # initiate time + time_start = time.time() + + # EAS队列配置 + EAS_URL = 'http://17xxxxxxxxx.pai-eas.aliyuncs.com/api/predict/xxxxxxxx' + # Use in EAS Queue + TOKEN = 'xxxxxxxx' + + # Support TeaCache. + enable_teacache = True + # Recommended to be set between 0.05 and 0.30. A larger threshold can cache more steps, speeding up the inference process, + # but it may cause slight differences between the generated content and the original content. + # # --------------------------------------------------------------------------------------------------- # + # | Model Name | threshold | Model Name | threshold | + # | Wan2.2-T2V-A14B | 0.10~0.15 | Wan2.2-I2V-A14B | 0.15~0.20 | + # | Wan2.2-Fun-A14B-* | 0.15~0.20 | + # # --------------------------------------------------------------------------------------------------- # + teacache_threshold = 0.10 + # The number of steps to skip TeaCache at the beginning of the inference process, which can + # reduce the impact of TeaCache on generated video quality. + num_skip_start_steps = 5 + # Whether to offload TeaCache tensors to cpu to save a little bit of GPU memory. + teacache_offload = False + + # Skip some cfg steps in inference + # Recommended to be set between 0.00 and 0.25 + cfg_skip_ratio = 0 + + # Riflex config + enable_riflex = False + # Index of intrinsic frequency + riflex_k = 6 + + # "Video Generation" and "Image Generation" + generation_method = "Video Generation" + # Video length + length_slider = 81 + # Used in Lora models + lora_model_path = "none" + lora_alpha_slider = 0.55 + # Prompts + prompt_textbox = "在这个阳光明媚的户外花园里,美女身穿一袭及膝的白色无袖连衣裙,裙摆在她轻盈的舞姿中轻柔地摆动,宛如一只翩翩起舞的蝴蝶。阳光透过树叶间洒下斑驳的光影,映衬出她柔和的脸庞和清澈的眼眸,显得格外优雅。仿佛每一个动作都在诉说着青春与活力,她在草地上旋转,裙摆随之飞扬,仿佛整个花园都因她的舞动而欢愉。周围五彩缤纷的花朵在微风中摇曳,玫瑰、菊花、百合,各自释放出阵阵香气,营造出一种轻松而愉快的氛围。" + negative_prompt_textbox = "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走" + # Sampler name + sampler_dropdown = "Flow" + # Sampler steps + sample_step_slider = 50 + # height and width + width_slider = 480 + height_slider = 832 + # cfg scale + cfg_scale_slider = 6 + seed_textbox = 43 + + # 控制视频路径(可以是本地路径或 URL) + control_video_path = "asset/000000.mp4" # 替换为实际的视频路径 + # 参考图片路径 + ref_image_path = None # 替换为实际的图片路径 + + outputs = post_infer( + generation_method, + length_slider, + lora_model_path=lora_model_path, + lora_alpha_slider=lora_alpha_slider, + prompt_textbox=prompt_textbox, + negative_prompt_textbox=negative_prompt_textbox, + sampler_dropdown=sampler_dropdown, + sample_step_slider=sample_step_slider, + width_slider=width_slider, + height_slider=height_slider, + cfg_scale_slider=cfg_scale_slider, + seed_textbox=seed_textbox, + enable_teacache = enable_teacache, + teacache_threshold = teacache_threshold, + num_skip_start_steps = num_skip_start_steps, + teacache_offload = teacache_offload, + cfg_skip_ratio = cfg_skip_ratio, + enable_riflex = enable_riflex, + riflex_k = riflex_k, + url=EAS_URL, + POST_TOKEN=TOKEN, + control_video=control_video_path, # 传递控制视频路径 + ref_image=ref_image_path # 传递参考图片路径 + ) + # Get decoded data + outputs = json.loads(base64.b64decode(json.loads(outputs)[0]['data'])) + base64_encoding = outputs["base64_encoding"] + decoded_data = base64.b64decode(base64_encoding) + + is_image = True if generation_method == "Image Generation" else False + if is_image or length_slider == 1: + file_path = "1.png" + else: + file_path = "1.mp4" + with open(file_path, "wb") as file: + file.write(decoded_data) + + # End of record time + # The calculated time difference is the execution time of the program, expressed in seconds / s + time_end = time.time() + time_sum = (time_end - time_start) + print('# --------------------------------------------------------- #') + print(f'# Total expenditure: {time_sum}s') + print('# --------------------------------------------------------- #') diff --git a/examples/wan2.2_fun/predict_i2v.py b/examples/wan2.2_fun/predict_i2v.py index d9bafdc..0df872e 100644 --- a/examples/wan2.2_fun/predict_i2v.py +++ b/examples/wan2.2_fun/predict_i2v.py @@ -13,7 +13,7 @@ for project_root in project_roots: sys.path.insert(0, project_root) if project_root not in sys.path else None from videox_fun.dist import set_multi_gpus_devices, shard_model -from videox_fun.models import (AutoencoderKLWan, AutoTokenizer, CLIPModel, +from videox_fun.models import (AutoencoderKLWan, AutoencoderKLWan3_8, AutoTokenizer, CLIPModel, WanT5EncoderModel, Wan2_2Transformer3DModel) from videox_fun.models.cache_utils import get_teacache_coefficients from videox_fun.pipeline import Wan2_2I2VPipeline @@ -120,7 +120,7 @@ num_inference_steps = 50 # The lora_weight is used for low noise model, the lora_high_weight is used for high noise model. lora_weight = 0.55 lora_high_weight = 0.55 -save_path = "samples/wan-fun-videos-i2v" +save_path = "samples/wan-videos-fun-i2v" device = set_multi_gpus_devices(ulysses_degree, ring_degree) config = OmegaConf.load(config_path) @@ -132,13 +132,15 @@ transformer = Wan2_2Transformer3DModel.from_pretrained( low_cpu_mem_usage=True, torch_dtype=weight_dtype, ) - -transformer_2 = Wan2_2Transformer3DModel.from_pretrained( - os.path.join(model_name, config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer')), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - low_cpu_mem_usage=True, - torch_dtype=weight_dtype, -) +if config['transformer_additional_kwargs'].get('transformer_combination_type', 'single') == "moe": + transformer_2 = Wan2_2Transformer3DModel.from_pretrained( + os.path.join(model_name, config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer')), + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), + low_cpu_mem_usage=True, + torch_dtype=weight_dtype, + ) +else: + transformer_2 = None if transformer_path is not None: print(f"From checkpoint: {transformer_path}") @@ -152,21 +154,23 @@ if transformer_path is not None: m, u = transformer.load_state_dict(state_dict, strict=False) print(f"missing keys: {len(m)}, unexpected keys: {len(u)}") -if transformer_high_path is not None: - print(f"From checkpoint: {transformer_high_path}") - if transformer_high_path.endswith("safetensors"): - from safetensors.torch import load_file, safe_open - state_dict = load_file(transformer_high_path) - else: - state_dict = torch.load(transformer_high_path, map_location="cpu") - state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict +if transformer_2 is not None: + if transformer_high_path is not None: + print(f"From checkpoint: {transformer_high_path}") + if transformer_high_path.endswith("safetensors"): + from safetensors.torch import load_file, safe_open + state_dict = load_file(transformer_high_path) + else: + state_dict = torch.load(transformer_high_path, map_location="cpu") + state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict - m, u = transformer_2.load_state_dict(state_dict, strict=False) - print(f"missing keys: {len(m)}, unexpected keys: {len(u)}") + m, u = transformer_2.load_state_dict(state_dict, strict=False) + print(f"missing keys: {len(m)}, unexpected keys: {len(u)}") # Get Vae Choosen_AutoencoderKL = { "AutoencoderKLWan": AutoencoderKLWan, + "AutoencoderKLWan3_8": AutoencoderKLWan3_8 }[config['vae_kwargs'].get('vae_type', 'AutoencoderKLWan')] vae = Choosen_AutoencoderKL.from_pretrained( os.path.join(model_name, config['vae_kwargs'].get('vae_subpath', 'vae')), @@ -223,11 +227,13 @@ pipeline = Wan2_2I2VPipeline( if ulysses_degree > 1 or ring_degree > 1: from functools import partial transformer.enable_multi_gpus_inference() - transformer_2.enable_multi_gpus_inference() + if transformer_2 is not None: + transformer_2.enable_multi_gpus_inference() if fsdp_dit: shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype) pipeline.transformer = shard_fn(pipeline.transformer) - pipeline.transformer_2 = shard_fn(pipeline.transformer_2) + if transformer_2 is not None: + pipeline.transformer_2 = shard_fn(pipeline.transformer_2) print("Add FSDP DIT") if fsdp_text_encoder: shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype) @@ -237,29 +243,33 @@ if ulysses_degree > 1 or ring_degree > 1: if compile_dit: for i in range(len(pipeline.transformer.blocks)): pipeline.transformer.blocks[i] = torch.compile(pipeline.transformer.blocks[i]) - for i in range(len(pipeline.transformer_2.blocks)): - pipeline.transformer_2.blocks[i] = torch.compile(pipeline.transformer_2.blocks[i]) + if transformer_2 is not None: + for i in range(len(pipeline.transformer_2.blocks)): + pipeline.transformer_2.blocks[i] = torch.compile(pipeline.transformer_2.blocks[i]) print("Add Compile") if GPU_memory_mode == "sequential_cpu_offload": replace_parameters_by_name(transformer, ["modulation",], device=device) - replace_parameters_by_name(transformer_2, ["modulation",], device=device) transformer.freqs = transformer.freqs.to(device=device) - transformer_2.freqs = transformer_2.freqs.to(device=device) + if transformer_2 is not None: + replace_parameters_by_name(transformer_2, ["modulation",], device=device) + transformer_2.freqs = transformer_2.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",], device=device) - convert_model_weight_to_float8(transformer_2, exclude_module_name=["modulation",], device=device) convert_weight_dtype_wrapper(transformer, weight_dtype) - convert_weight_dtype_wrapper(transformer_2, weight_dtype) + if transformer_2 is not None: + convert_model_weight_to_float8(transformer_2, exclude_module_name=["modulation",], device=device) + convert_weight_dtype_wrapper(transformer_2, 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",], device=device) - convert_model_weight_to_float8(transformer_2, exclude_module_name=["modulation",], device=device) convert_weight_dtype_wrapper(transformer, weight_dtype) - convert_weight_dtype_wrapper(transformer_2, weight_dtype) + if transformer_2 is not None: + convert_model_weight_to_float8(transformer_2, exclude_module_name=["modulation",], device=device) + convert_weight_dtype_wrapper(transformer_2, weight_dtype) pipeline.to(device=device) else: pipeline.to(device=device) @@ -270,18 +280,21 @@ if coefficients is not None: pipeline.transformer.enable_teacache( coefficients, num_inference_steps, teacache_threshold, num_skip_start_steps=num_skip_start_steps, offload=teacache_offload ) - pipeline.transformer_2.share_teacache(transformer=pipeline.transformer) + if transformer_2 is not None: + pipeline.transformer_2.share_teacache(transformer=pipeline.transformer) 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) - pipeline.transformer_2.share_cfg_skip(transformer=pipeline.transformer) + if transformer_2 is not None: + pipeline.transformer_2.share_cfg_skip(transformer=pipeline.transformer) generator = torch.Generator(device=device).manual_seed(seed) if lora_path is not None: pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device) - pipeline = merge_lora(pipeline, lora_high_path, lora_high_weight, device=device, sub_transformer_name="transformer_2") + if transformer_2 is not None: + pipeline = merge_lora(pipeline, lora_high_path, lora_high_weight, device=device, sub_transformer_name="transformer_2") with torch.no_grad(): video_length = int((video_length - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if video_length != 1 else 1 @@ -289,7 +302,8 @@ with torch.no_grad(): if enable_riflex: pipeline.transformer.enable_riflex(k = riflex_k, L_test = latent_frames) - pipeline.transformer_2.enable_riflex(k = riflex_k, L_test = latent_frames) + if transformer_2 is not None: + pipeline.transformer_2.enable_riflex(k = riflex_k, L_test = latent_frames) input_video, input_video_mask, clip_image = get_image_to_video_latent(validation_image_start, validation_image_end, video_length=video_length, sample_size=sample_size) @@ -311,7 +325,8 @@ with torch.no_grad(): if lora_path is not None: pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device) - pipeline = unmerge_lora(pipeline, lora_high_path, lora_high_weight, device=device, sub_transformer_name="transformer_2") + if transformer_2 is not None: + pipeline = unmerge_lora(pipeline, lora_high_path, lora_high_weight, device=device, sub_transformer_name="transformer_2") def save_results(): if not os.path.exists(save_path): diff --git a/examples/wan2.2_fun/predict_t2v.py b/examples/wan2.2_fun/predict_t2v.py new file mode 100644 index 0000000..69431fa --- /dev/null +++ b/examples/wan2.2_fun/predict_t2v.py @@ -0,0 +1,336 @@ +import os +import sys + +import numpy as np +import torch +from diffusers import FlowMatchEulerDiscreteScheduler +from omegaconf import OmegaConf +from PIL import Image + +current_file_path = os.path.abspath(__file__) +project_roots = [os.path.dirname(current_file_path), os.path.dirname(os.path.dirname(current_file_path)), os.path.dirname(os.path.dirname(os.path.dirname(current_file_path)))] +for project_root in project_roots: + sys.path.insert(0, project_root) if project_root not in sys.path else None + +from videox_fun.dist import set_multi_gpus_devices, shard_model +from videox_fun.models import (AutoencoderKLWan, AutoTokenizer, CLIPModel, + WanT5EncoderModel, Wan2_2Transformer3DModel) +from videox_fun.models.cache_utils import get_teacache_coefficients +from videox_fun.pipeline import Wan2_2I2VPipeline +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 (filter_kwargs, get_image_to_video_latent, + save_videos_grid) +from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler +from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler + +# GPU memory mode, which can be choosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8, sequential_cpu_offload]. +# 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. +# +# 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" +# Multi GPUs config +# Please ensure that the product of ulysses_degree and ring_degree equals the number of GPUs used. +# For example, if you are using 8 GPUs, you can set ulysses_degree = 2 and ring_degree = 4. +# 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 + +# TeaCache config +enable_teacache = True +# Recommended to be set between 0.05 and 0.30. A larger threshold can cache more steps, speeding up the inference process, +# but it may cause slight differences between the generated content and the original content. +# # --------------------------------------------------------------------------------------------------- # +# | Model Name | threshold | Model Name | threshold | +# | Wan2.2-T2V-A14B | 0.10~0.15 | Wan2.2-I2V-A14B | 0.15~0.20 | +# | Wan2.2-Fun-A14B-* | 0.15~0.20 | +# # --------------------------------------------------------------------------------------------------- # +teacache_threshold = 0.10 +# The number of steps to skip TeaCache at the beginning of the inference process, which can +# reduce the impact of TeaCache on generated video quality. +num_skip_start_steps = 5 +# Whether to offload TeaCache tensors to cpu to save a little bit of GPU memory. +teacache_offload = False + +# Skip some cfg steps in inference +# Recommended to be set between 0.00 and 0.25 +cfg_skip_ratio = 0 + +# Riflex config +enable_riflex = False +# Index of intrinsic frequency +riflex_k = 6 + +# Config and model path +config_path = "config/wan2.2/wan_civitai_i2v.yaml" +# model path +model_name = "models/Diffusion_Transformer/Wan2.2-Fun-A14B-InP" + +# Choose the sampler in "Flow", "Flow_Unipc", "Flow_DPM++" +sampler_name = "Flow" +# [NOTE]: Noise schedule shift parameter. Affects temporal dynamics. +# Used when the sampler is in "Flow_Unipc", "Flow_DPM++". +shift = 5 + +# Load pretrained model if need +# The transformer_path is used for low noise model, the transformer_high_path is used for high noise model. +transformer_path = None +transformer_high_path = None +vae_path = None +# Load lora model if need +# The lora_path is used for low noise model, the lora_high_path is used for high noise model. +lora_path = None +lora_high_path = None + +# Other params +sample_size = [480, 832] +video_length = 81 +fps = 16 + +# Use torch.float16 if GPU does not support torch.bfloat16 +# ome graphics cards, such as v100, 2080ti, do not support torch.bfloat16 +weight_dtype = torch.bfloat16 +# 使用更长的neg prompt如"模糊,突变,变形,失真,画面暗,文本字幕,画面固定,连环画,漫画,线稿,没有主体。",可以增加稳定性 +# 在neg prompt中添加"安静,固定"等词语可以增加动态性。 +prompt = "一只棕色的狗摇着头,坐在舒适房间里的浅色沙发上。在狗的后面,架子上有一幅镶框的画,周围是粉红色的花朵。房间里柔和温暖的灯光营造出舒适的氛围。" +negative_prompt = "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走" +guidance_scale = 6.0 +seed = 43 +num_inference_steps = 50 +# The lora_weight is used for low noise model, the lora_high_weight is used for high noise model. +lora_weight = 0.55 +lora_high_weight = 0.55 +save_path = "samples/wan-fun-videos-i2v" + +device = set_multi_gpus_devices(ulysses_degree, ring_degree) +config = OmegaConf.load(config_path) +boundary = config['transformer_additional_kwargs'].get('boundary', 0.900) + +transformer = Wan2_2Transformer3DModel.from_pretrained( + os.path.join(model_name, config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer')), + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), + low_cpu_mem_usage=True, + torch_dtype=weight_dtype, +) + +transformer_2 = Wan2_2Transformer3DModel.from_pretrained( + os.path.join(model_name, config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer')), + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), + low_cpu_mem_usage=True, + torch_dtype=weight_dtype, +) + +if transformer_path is not None: + print(f"From checkpoint: {transformer_path}") + if transformer_path.endswith("safetensors"): + from safetensors.torch import load_file, safe_open + state_dict = load_file(transformer_path) + else: + state_dict = torch.load(transformer_path, map_location="cpu") + state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict + + m, u = transformer.load_state_dict(state_dict, strict=False) + print(f"missing keys: {len(m)}, unexpected keys: {len(u)}") + +if transformer_high_path is not None: + print(f"From checkpoint: {transformer_high_path}") + if transformer_high_path.endswith("safetensors"): + from safetensors.torch import load_file, safe_open + state_dict = load_file(transformer_high_path) + else: + state_dict = torch.load(transformer_high_path, map_location="cpu") + state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict + + m, u = transformer_2.load_state_dict(state_dict, strict=False) + print(f"missing keys: {len(m)}, unexpected keys: {len(u)}") + +# Get Vae +Choosen_AutoencoderKL = { + "AutoencoderKLWan": AutoencoderKLWan, + "AutoencoderKLWan3_8": AutoencoderKLWan3_8 +}[config['vae_kwargs'].get('vae_type', 'AutoencoderKLWan')] +vae = Choosen_AutoencoderKL.from_pretrained( + os.path.join(model_name, config['vae_kwargs'].get('vae_subpath', 'vae')), + additional_kwargs=OmegaConf.to_container(config['vae_kwargs']), +).to(weight_dtype) + +if vae_path is not None: + print(f"From checkpoint: {vae_path}") + if vae_path.endswith("safetensors"): + from safetensors.torch import load_file, safe_open + state_dict = load_file(vae_path) + else: + state_dict = torch.load(vae_path, map_location="cpu") + state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict + + m, u = vae.load_state_dict(state_dict, strict=False) + print(f"missing keys: {len(m)}, unexpected keys: {len(u)}") + +# Get Tokenizer +tokenizer = AutoTokenizer.from_pretrained( + os.path.join(model_name, config['text_encoder_kwargs'].get('tokenizer_subpath', 'tokenizer')), +) + +# Get Text encoder +text_encoder = WanT5EncoderModel.from_pretrained( + os.path.join(model_name, config['text_encoder_kwargs'].get('text_encoder_subpath', 'text_encoder')), + additional_kwargs=OmegaConf.to_container(config['text_encoder_kwargs']), + low_cpu_mem_usage=True, + torch_dtype=weight_dtype, +) +text_encoder = text_encoder.eval() + +# Get Scheduler +Choosen_Scheduler = scheduler_dict = { + "Flow": FlowMatchEulerDiscreteScheduler, + "Flow_Unipc": FlowUniPCMultistepScheduler, + "Flow_DPM++": FlowDPMSolverMultistepScheduler, +}[sampler_name] +if sampler_name == "Flow_Unipc" or sampler_name == "Flow_DPM++": + config['scheduler_kwargs']['shift'] = 1 +scheduler = Choosen_Scheduler( + **filter_kwargs(Choosen_Scheduler, OmegaConf.to_container(config['scheduler_kwargs'])) +) + +# Get Pipeline +pipeline = Wan2_2I2VPipeline( + transformer=transformer, + transformer_2=transformer_2, + vae=vae, + tokenizer=tokenizer, + text_encoder=text_encoder, + scheduler=scheduler, +) +if ulysses_degree > 1 or ring_degree > 1: + from functools import partial + transformer.enable_multi_gpus_inference() + transformer_2.enable_multi_gpus_inference() + if fsdp_dit: + shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype) + pipeline.transformer = shard_fn(pipeline.transformer) + pipeline.transformer_2 = shard_fn(pipeline.transformer_2) + 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]) + for i in range(len(pipeline.transformer_2.blocks)): + pipeline.transformer_2.blocks[i] = torch.compile(pipeline.transformer_2.blocks[i]) + print("Add Compile") + +if GPU_memory_mode == "sequential_cpu_offload": + replace_parameters_by_name(transformer, ["modulation",], device=device) + replace_parameters_by_name(transformer_2, ["modulation",], device=device) + transformer.freqs = transformer.freqs.to(device=device) + transformer_2.freqs = transformer_2.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",], device=device) + convert_model_weight_to_float8(transformer_2, exclude_module_name=["modulation",], device=device) + convert_weight_dtype_wrapper(transformer, weight_dtype) + convert_weight_dtype_wrapper(transformer_2, 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",], device=device) + convert_model_weight_to_float8(transformer_2, exclude_module_name=["modulation",], device=device) + convert_weight_dtype_wrapper(transformer, weight_dtype) + convert_weight_dtype_wrapper(transformer_2, weight_dtype) + pipeline.to(device=device) +else: + pipeline.to(device=device) + +coefficients = get_teacache_coefficients(model_name) if enable_teacache else None +if coefficients is not None: + print(f"Enable TeaCache with threshold {teacache_threshold} and skip the first {num_skip_start_steps} steps.") + pipeline.transformer.enable_teacache( + coefficients, num_inference_steps, teacache_threshold, num_skip_start_steps=num_skip_start_steps, offload=teacache_offload + ) + pipeline.transformer_2.share_teacache(transformer=pipeline.transformer) + +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) + pipeline.transformer_2.share_cfg_skip(transformer=pipeline.transformer) + +generator = torch.Generator(device=device).manual_seed(seed) + +if lora_path is not None: + pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device) + pipeline = merge_lora(pipeline, lora_high_path, lora_high_weight, device=device, sub_transformer_name="transformer_2") + +with torch.no_grad(): + video_length = int((video_length - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if video_length != 1 else 1 + latent_frames = (video_length - 1) // vae.config.temporal_compression_ratio + 1 + + if enable_riflex: + pipeline.transformer.enable_riflex(k = riflex_k, L_test = latent_frames) + pipeline.transformer_2.enable_riflex(k = riflex_k, L_test = latent_frames) + + input_video, input_video_mask, _ = get_image_to_video_latent(None, None, video_length=video_length, sample_size=sample_size) + + sample = pipeline( + prompt, + num_frames = video_length, + negative_prompt = negative_prompt, + height = sample_size[0], + width = sample_size[1], + generator = generator, + guidance_scale = guidance_scale, + num_inference_steps = num_inference_steps, + + video = input_video, + mask_video = input_video_mask, + boundary = boundary, + shift = shift, + ).videos + +if lora_path is not None: + pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device) + pipeline = unmerge_lora(pipeline, lora_high_path, lora_high_weight, device=device, sub_transformer_name="transformer_2") + +def save_results(): + if not os.path.exists(save_path): + os.makedirs(save_path, exist_ok=True) + + index = len([path for path in os.listdir(save_path)]) + 1 + prefix = str(index).zfill(8) + if video_length == 1: + video_path = os.path.join(save_path, prefix + ".png") + + image = sample[0, :, 0] + image = image.transpose(0, 1).transpose(1, 2) + image = (image * 255).numpy().astype(np.uint8) + image = Image.fromarray(image) + image.save(video_path) + else: + video_path = os.path.join(save_path, prefix + ".mp4") + save_videos_grid(sample, video_path, fps=fps) + +if ulysses_degree * ring_degree > 1: + import torch.distributed as dist + if dist.get_rank() == 0: + save_results() +else: + save_results() \ No newline at end of file diff --git a/examples/wan2.2_fun/predict_v2v_control.py b/examples/wan2.2_fun/predict_v2v_control.py index f24c953..b7a5dd1 100644 --- a/examples/wan2.2_fun/predict_v2v_control.py +++ b/examples/wan2.2_fun/predict_v2v_control.py @@ -14,7 +14,7 @@ for project_root in project_roots: sys.path.insert(0, project_root) if project_root not in sys.path else None from videox_fun.dist import set_multi_gpus_devices, shard_model -from videox_fun.models import (AutoencoderKLWan, AutoTokenizer, CLIPModel, +from videox_fun.models import (AutoencoderKLWan, AutoencoderKLWan3_8, AutoTokenizer, CLIPModel, WanT5EncoderModel, Wan2_2Transformer3DModel) from videox_fun.data.dataset_image_video import process_pose_file from videox_fun.models.cache_utils import get_teacache_coefficients @@ -128,7 +128,7 @@ negative_prompt = "色调艳丽,过曝,静态,细节模糊不清,字 # prompt = "A young woman with beautiful, clear eyes and blonde hair stands in the forest, wearing a white dress and a crown. Her expression is serene, reminiscent of a movie star, with fair and youthful skin. Her brown long hair flows in the wind. The video quality is very high, with a clear view. High quality, masterpiece, best quality, high resolution, ultra-fine, fantastical." # negative_prompt = "Twisted body, limb deformities, text captions, comic, static, ugly, error, messy code." guidance_scale = 6.0 -seed = 42 +seed = 43 num_inference_steps = 50 # The lora_weight is used for low noise model, the lora_high_weight is used for high noise model. lora_weight = 0.55 @@ -145,13 +145,15 @@ transformer = Wan2_2Transformer3DModel.from_pretrained( low_cpu_mem_usage=True, torch_dtype=weight_dtype, ) - -transformer_2 = Wan2_2Transformer3DModel.from_pretrained( - os.path.join(model_name, config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer')), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - low_cpu_mem_usage=True, - torch_dtype=weight_dtype, -) +if config['transformer_additional_kwargs'].get('transformer_combination_type', 'single') == "moe": + transformer_2 = Wan2_2Transformer3DModel.from_pretrained( + os.path.join(model_name, config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer')), + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), + low_cpu_mem_usage=True, + torch_dtype=weight_dtype, + ) +else: + transformer_2 = None if transformer_path is not None: print(f"From checkpoint: {transformer_path}") @@ -165,21 +167,23 @@ if transformer_path is not None: m, u = transformer.load_state_dict(state_dict, strict=False) print(f"missing keys: {len(m)}, unexpected keys: {len(u)}") -if transformer_high_path is not None: - print(f"From checkpoint: {transformer_high_path}") - if transformer_high_path.endswith("safetensors"): - from safetensors.torch import load_file, safe_open - state_dict = load_file(transformer_high_path) - else: - state_dict = torch.load(transformer_high_path, map_location="cpu") - state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict +if transformer_2 is not None: + if transformer_high_path is not None: + print(f"From checkpoint: {transformer_high_path}") + if transformer_high_path.endswith("safetensors"): + from safetensors.torch import load_file, safe_open + state_dict = load_file(transformer_high_path) + else: + state_dict = torch.load(transformer_high_path, map_location="cpu") + state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict - m, u = transformer_2.load_state_dict(state_dict, strict=False) - print(f"missing keys: {len(m)}, unexpected keys: {len(u)}") + m, u = transformer_2.load_state_dict(state_dict, strict=False) + print(f"missing keys: {len(m)}, unexpected keys: {len(u)}") # Get Vae Choosen_AutoencoderKL = { "AutoencoderKLWan": AutoencoderKLWan, + "AutoencoderKLWan3_8": AutoencoderKLWan3_8 }[config['vae_kwargs'].get('vae_type', 'AutoencoderKLWan')] vae = Choosen_AutoencoderKL.from_pretrained( os.path.join(model_name, config['vae_kwargs'].get('vae_subpath', 'vae')), @@ -236,11 +240,13 @@ pipeline = Wan2_2FunControlPipeline( if ulysses_degree > 1 or ring_degree > 1: from functools import partial transformer.enable_multi_gpus_inference() - transformer_2.enable_multi_gpus_inference() + if transformer_2 is not None: + transformer_2.enable_multi_gpus_inference() if fsdp_dit: shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype) pipeline.transformer = shard_fn(pipeline.transformer) - pipeline.transformer_2 = shard_fn(pipeline.transformer_2) + if transformer_2 is not None: + pipeline.transformer_2 = shard_fn(pipeline.transformer_2) print("Add FSDP DIT") if fsdp_text_encoder: shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype) @@ -250,29 +256,33 @@ if ulysses_degree > 1 or ring_degree > 1: if compile_dit: for i in range(len(pipeline.transformer.blocks)): pipeline.transformer.blocks[i] = torch.compile(pipeline.transformer.blocks[i]) - for i in range(len(pipeline.transformer_2.blocks)): - pipeline.transformer_2.blocks[i] = torch.compile(pipeline.transformer_2.blocks[i]) + if transformer_2 is not None: + for i in range(len(pipeline.transformer_2.blocks)): + pipeline.transformer_2.blocks[i] = torch.compile(pipeline.transformer_2.blocks[i]) print("Add Compile") if GPU_memory_mode == "sequential_cpu_offload": replace_parameters_by_name(transformer, ["modulation",], device=device) - replace_parameters_by_name(transformer_2, ["modulation",], device=device) transformer.freqs = transformer.freqs.to(device=device) - transformer_2.freqs = transformer_2.freqs.to(device=device) + if transformer_2 is not None: + replace_parameters_by_name(transformer_2, ["modulation",], device=device) + transformer_2.freqs = transformer_2.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",], device=device) - convert_model_weight_to_float8(transformer_2, exclude_module_name=["modulation",], device=device) convert_weight_dtype_wrapper(transformer, weight_dtype) - convert_weight_dtype_wrapper(transformer_2, weight_dtype) + if transformer_2 is not None: + convert_model_weight_to_float8(transformer_2, exclude_module_name=["modulation",], device=device) + convert_weight_dtype_wrapper(transformer_2, 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=["modulation",], device=device) - convert_model_weight_to_float8(transformer_2, exclude_module_name=["modulation",], device=device) convert_weight_dtype_wrapper(transformer, weight_dtype) - convert_weight_dtype_wrapper(transformer_2, weight_dtype) + if transformer_2 is not None: + convert_model_weight_to_float8(transformer_2, exclude_module_name=["modulation",], device=device) + convert_weight_dtype_wrapper(transformer_2, weight_dtype) pipeline.to(device=device) else: pipeline.to(device=device) @@ -283,18 +293,21 @@ if coefficients is not None: pipeline.transformer.enable_teacache( coefficients, num_inference_steps, teacache_threshold, num_skip_start_steps=num_skip_start_steps, offload=teacache_offload ) - pipeline.transformer_2.share_teacache(transformer=pipeline.transformer) + if transformer_2 is not None: + pipeline.transformer_2.share_teacache(transformer=pipeline.transformer) 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) - pipeline.transformer_2.share_cfg_skip(transformer=pipeline.transformer) + if transformer_2 is not None: + pipeline.transformer_2.share_cfg_skip(transformer=pipeline.transformer) generator = torch.Generator(device=device).manual_seed(seed) if lora_path is not None: pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device) - pipeline = merge_lora(pipeline, lora_high_path, lora_high_weight, device=device, sub_transformer_name="transformer_2") + if transformer_2 is not None: + pipeline = merge_lora(pipeline, lora_high_path, lora_high_weight, device=device, sub_transformer_name="transformer_2") with torch.no_grad(): video_length = int((video_length - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if video_length != 1 else 1 @@ -302,6 +315,8 @@ with torch.no_grad(): if enable_riflex: pipeline.transformer.enable_riflex(k = riflex_k, L_test = latent_frames) + if transformer_2 is not None: + pipeline.transformer_2.enable_riflex(k = riflex_k, L_test = latent_frames) inpaint_video, inpaint_video_mask, clip_image = get_image_to_video_latent(start_image, end_image, video_length=video_length, sample_size=sample_size) @@ -337,7 +352,8 @@ with torch.no_grad(): if lora_path is not None: pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device) - pipeline = unmerge_lora(pipeline, lora_high_path, lora_high_weight, device=device, sub_transformer_name="transformer_2") + if transformer_2 is not None: + pipeline = unmerge_lora(pipeline, lora_high_path, lora_high_weight, device=device, sub_transformer_name="transformer_2") def save_results(): if not os.path.exists(save_path): diff --git a/examples/wan2.2_fun/predict_v2v_control_camera.py b/examples/wan2.2_fun/predict_v2v_control_camera.py new file mode 100644 index 0000000..3b0e2e7 --- /dev/null +++ b/examples/wan2.2_fun/predict_v2v_control_camera.py @@ -0,0 +1,381 @@ +import os +import sys + +import numpy as np +import torch +from diffusers import FlowMatchEulerDiscreteScheduler +from omegaconf import OmegaConf +from PIL import Image +from transformers import AutoTokenizer + +current_file_path = os.path.abspath(__file__) +project_roots = [os.path.dirname(current_file_path), os.path.dirname(os.path.dirname(current_file_path)), os.path.dirname(os.path.dirname(os.path.dirname(current_file_path)))] +for project_root in project_roots: + sys.path.insert(0, project_root) if project_root not in sys.path else None + +from videox_fun.dist import set_multi_gpus_devices, shard_model +from videox_fun.models import (AutoencoderKLWan, AutoencoderKLWan3_8, AutoTokenizer, CLIPModel, + WanT5EncoderModel, Wan2_2Transformer3DModel) +from videox_fun.data.dataset_image_video import process_pose_file +from videox_fun.models.cache_utils import get_teacache_coefficients +from videox_fun.pipeline import Wan2_2FunControlPipeline, WanPipeline +from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8, + convert_weight_dtype_wrapper, + replace_parameters_by_name) +from videox_fun.utils.lora_utils import merge_lora, unmerge_lora +from videox_fun.utils.utils import (filter_kwargs, get_image_latent, get_image_to_video_latent, + get_video_to_video_latent, + save_videos_grid) +from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler +from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler + +# GPU memory mode, which can be choosen in [model_full_load, model_cpu_offload_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8, sequential_cpu_offload]. +# 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. +# +# 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" +# Multi GPUs config +# Please ensure that the product of ulysses_degree and ring_degree equals the number of GPUs used. +# For example, if you are using 8 GPUs, you can set ulysses_degree = 2 and ring_degree = 4. +# 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 + +# Support TeaCache. +enable_teacache = True +# Recommended to be set between 0.05 and 0.30. A larger threshold can cache more steps, speeding up the inference process, +# but it may cause slight differences between the generated content and the original content. +# # --------------------------------------------------------------------------------------------------- # +# | Model Name | threshold | Model Name | threshold | +# | Wan2.2-T2V-A14B | 0.10~0.15 | Wan2.2-I2V-A14B | 0.15~0.20 | +# | Wan2.2-Fun-A14B-* | 0.15~0.20 | +# # --------------------------------------------------------------------------------------------------- # +teacache_threshold = 0.10 +# The number of steps to skip TeaCache at the beginning of the inference process, which can +# reduce the impact of TeaCache on generated video quality. +num_skip_start_steps = 5 +# Whether to offload TeaCache tensors to cpu to save a little bit of GPU memory. +teacache_offload = False + +# Skip some cfg steps in inference +# Recommended to be set between 0.00 and 0.25 +cfg_skip_ratio = 0 + +# Riflex config +enable_riflex = False +# Index of intrinsic frequency +riflex_k = 6 + +# Config and model path +config_path = "config/wan2.2/wan_civitai_i2v.yaml" +# model path +model_name = "models/Diffusion_Transformer/Wan2.2-Fun-A14B-Control-Camera" + +# Choose the sampler in "Flow", "Flow_Unipc", "Flow_DPM++" +sampler_name = "Flow" +# [NOTE]: Noise schedule shift parameter. Affects temporal dynamics. +# Used when the sampler is in "Flow_Unipc", "Flow_DPM++". +# If you want to generate a 480p video, it is recommended to set the shift value to 3.0. +# If you want to generate a 720p video, it is recommended to set the shift value to 5.0. +shift = 5 + +# Load pretrained model if need +# The transformer_path is used for low noise model, the transformer_high_path is used for high noise model. +transformer_path = None +transformer_high_path = None +vae_path = None +# Load lora model if need +# The lora_path is used for low noise model, the lora_high_path is used for high noise model. +lora_path = None +lora_high_path = None + +# Other params +sample_size = [480, 832] +video_length = 81 +fps = 16 + +# Use torch.float16 if GPU does not support torch.bfloat16 +# ome graphics cards, such as v100, 2080ti, do not support torch.bfloat16 +weight_dtype = torch.bfloat16 +control_video = None +control_camera_txt = "asset/Zoom_In.txt" +start_image = "asset/7.png" +end_image = None +ref_image = None + +# 使用更长的neg prompt如"模糊,突变,变形,失真,画面暗,文本字幕,画面固定,连环画,漫画,线稿,没有主体。",可以增加稳定性 +# 在neg prompt中添加"安静,固定"等词语可以增加动态性。 +prompt = "一个小女孩正在户外玩耍。她穿着一件蓝色的短袖上衣和粉色的短裤,头发扎成一个可爱的辫子。她的脚上没有穿鞋,显得非常自然和随意。她正用一把红色的小铲子在泥土里挖土,似乎在进行某种有趣的活动,可能是种花或是挖掘宝藏。地上有一根长长的水管,可能是用来浇水的。背景是一片草地和一些绿色植物,阳光明媚,整个场景充满了童趣和生机。小女孩专注的表情和认真的动作让人感受到她的快乐和好奇心。" +negative_prompt = "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走" + +# Using longer neg prompt such as "Blurring, mutation, deformation, distortion, dark and solid, comics, text subtitles, line art." can increase stability +# Adding words such as "quiet, solid" to the neg prompt can increase dynamism. +# prompt = "A young woman with beautiful, clear eyes and blonde hair stands in the forest, wearing a white dress and a crown. Her expression is serene, reminiscent of a movie star, with fair and youthful skin. Her brown long hair flows in the wind. The video quality is very high, with a clear view. High quality, masterpiece, best quality, high resolution, ultra-fine, fantastical." +# negative_prompt = "Twisted body, limb deformities, text captions, comic, static, ugly, error, messy code." +guidance_scale = 6.0 +seed = 42 +num_inference_steps = 50 +# The lora_weight is used for low noise model, the lora_high_weight is used for high noise model. +lora_weight = 0.55 +lora_high_weight = 0.55 +save_path = "samples/wan-videos-fun-control" + +device = set_multi_gpus_devices(ulysses_degree, ring_degree) +config = OmegaConf.load(config_path) +boundary = config['transformer_additional_kwargs'].get('boundary', 0.875) + +transformer = Wan2_2Transformer3DModel.from_pretrained( + os.path.join(model_name, config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer')), + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), + low_cpu_mem_usage=True, + torch_dtype=weight_dtype, +) +if config['transformer_additional_kwargs'].get('transformer_combination_type', 'single') == "moe": + transformer_2 = Wan2_2Transformer3DModel.from_pretrained( + os.path.join(model_name, config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer')), + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), + low_cpu_mem_usage=True, + torch_dtype=weight_dtype, + ) +else: + transformer_2 = None + +if transformer_path is not None: + print(f"From checkpoint: {transformer_path}") + if transformer_path.endswith("safetensors"): + from safetensors.torch import load_file, safe_open + state_dict = load_file(transformer_path) + else: + state_dict = torch.load(transformer_path, map_location="cpu") + state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict + + m, u = transformer.load_state_dict(state_dict, strict=False) + print(f"missing keys: {len(m)}, unexpected keys: {len(u)}") + +if transformer_2 is not None: + if transformer_high_path is not None: + print(f"From checkpoint: {transformer_high_path}") + if transformer_high_path.endswith("safetensors"): + from safetensors.torch import load_file, safe_open + state_dict = load_file(transformer_high_path) + else: + state_dict = torch.load(transformer_high_path, map_location="cpu") + state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict + + m, u = transformer_2.load_state_dict(state_dict, strict=False) + print(f"missing keys: {len(m)}, unexpected keys: {len(u)}") + +# Get Vae +Choosen_AutoencoderKL = { + "AutoencoderKLWan": AutoencoderKLWan, + "AutoencoderKLWan3_8": AutoencoderKLWan3_8 +}[config['vae_kwargs'].get('vae_type', 'AutoencoderKLWan')] +vae = Choosen_AutoencoderKL.from_pretrained( + os.path.join(model_name, config['vae_kwargs'].get('vae_subpath', 'vae')), + additional_kwargs=OmegaConf.to_container(config['vae_kwargs']), +).to(weight_dtype) + +if vae_path is not None: + print(f"From checkpoint: {vae_path}") + if vae_path.endswith("safetensors"): + from safetensors.torch import load_file, safe_open + state_dict = load_file(vae_path) + else: + state_dict = torch.load(vae_path, map_location="cpu") + state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict + + m, u = vae.load_state_dict(state_dict, strict=False) + print(f"missing keys: {len(m)}, unexpected keys: {len(u)}") + +# Get Tokenizer +tokenizer = AutoTokenizer.from_pretrained( + os.path.join(model_name, config['text_encoder_kwargs'].get('tokenizer_subpath', 'tokenizer')), +) + +# Get Text encoder +text_encoder = WanT5EncoderModel.from_pretrained( + os.path.join(model_name, config['text_encoder_kwargs'].get('text_encoder_subpath', 'text_encoder')), + additional_kwargs=OmegaConf.to_container(config['text_encoder_kwargs']), + low_cpu_mem_usage=True, + torch_dtype=weight_dtype, +) +text_encoder = text_encoder.eval() + +# Get Scheduler +Choosen_Scheduler = scheduler_dict = { + "Flow": FlowMatchEulerDiscreteScheduler, + "Flow_Unipc": FlowUniPCMultistepScheduler, + "Flow_DPM++": FlowDPMSolverMultistepScheduler, +}[sampler_name] +if sampler_name == "Flow_Unipc" or sampler_name == "Flow_DPM++": + config['scheduler_kwargs']['shift'] = 1 +scheduler = Choosen_Scheduler( + **filter_kwargs(Choosen_Scheduler, OmegaConf.to_container(config['scheduler_kwargs'])) +) + +# Get Pipeline +pipeline = Wan2_2FunControlPipeline( + transformer=transformer, + transformer_2=transformer_2, + vae=vae, + tokenizer=tokenizer, + text_encoder=text_encoder, + scheduler=scheduler, +) +if ulysses_degree > 1 or ring_degree > 1: + from functools import partial + transformer.enable_multi_gpus_inference() + if transformer_2 is not None: + transformer_2.enable_multi_gpus_inference() + if fsdp_dit: + shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype) + pipeline.transformer = shard_fn(pipeline.transformer) + if transformer_2 is not None: + pipeline.transformer_2 = shard_fn(pipeline.transformer_2) + 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]) + if transformer_2 is not None: + for i in range(len(pipeline.transformer_2.blocks)): + pipeline.transformer_2.blocks[i] = torch.compile(pipeline.transformer_2.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) + if transformer_2 is not None: + replace_parameters_by_name(transformer_2, ["modulation",], device=device) + transformer_2.freqs = transformer_2.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",], device=device) + convert_weight_dtype_wrapper(transformer, weight_dtype) + if transformer_2 is not None: + convert_model_weight_to_float8(transformer_2, exclude_module_name=["modulation",], device=device) + convert_weight_dtype_wrapper(transformer_2, 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",], device=device) + convert_weight_dtype_wrapper(transformer, weight_dtype) + if transformer_2 is not None: + convert_model_weight_to_float8(transformer_2, exclude_module_name=["modulation",], device=device) + convert_weight_dtype_wrapper(transformer_2, weight_dtype) + pipeline.to(device=device) +else: + pipeline.to(device=device) + +coefficients = get_teacache_coefficients(model_name) if enable_teacache else None +if coefficients is not None: + print(f"Enable TeaCache with threshold {teacache_threshold} and skip the first {num_skip_start_steps} steps.") + pipeline.transformer.enable_teacache( + coefficients, num_inference_steps, teacache_threshold, num_skip_start_steps=num_skip_start_steps, offload=teacache_offload + ) + if transformer_2 is not None: + pipeline.transformer_2.share_teacache(transformer=pipeline.transformer) + +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) + if transformer_2 is not None: + pipeline.transformer_2.share_cfg_skip(transformer=pipeline.transformer) + +generator = torch.Generator(device=device).manual_seed(seed) + +if lora_path is not None: + pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device) + if transformer_2 is not None: + pipeline = merge_lora(pipeline, lora_high_path, lora_high_weight, device=device, sub_transformer_name="transformer_2") + +with torch.no_grad(): + video_length = int((video_length - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if video_length != 1 else 1 + latent_frames = (video_length - 1) // vae.config.temporal_compression_ratio + 1 + + if enable_riflex: + pipeline.transformer.enable_riflex(k = riflex_k, L_test = latent_frames) + if transformer_2 is not None: + pipeline.transformer_2.enable_riflex(k = riflex_k, L_test = latent_frames) + + inpaint_video, inpaint_video_mask, clip_image = get_image_to_video_latent(start_image, end_image, video_length=video_length, sample_size=sample_size) + + if ref_image is not None: + ref_image = get_image_latent(ref_image, sample_size=sample_size) + + if control_camera_txt is not None: + input_video, input_video_mask = None, None + control_camera_video = process_pose_file(control_camera_txt, sample_size[1], sample_size[0]) + control_camera_video = control_camera_video[:video_length].permute([3, 0, 1, 2]).unsqueeze(0) + else: + input_video, input_video_mask, _, _ = get_video_to_video_latent(control_video, video_length=video_length, sample_size=sample_size, fps=fps, ref_image=None) + control_camera_video = None + + sample = pipeline( + prompt, + num_frames = video_length, + negative_prompt = negative_prompt, + height = sample_size[0], + width = sample_size[1], + generator = generator, + guidance_scale = guidance_scale, + num_inference_steps = num_inference_steps, + + video = inpaint_video, + mask_video = inpaint_video_mask, + control_video = input_video, + control_camera_video = control_camera_video, + ref_image = ref_image, + boundary = boundary, + shift = shift, + ).videos + +if lora_path is not None: + pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device) + if transformer_2 is not None: + pipeline = unmerge_lora(pipeline, lora_high_path, lora_high_weight, device=device, sub_transformer_name="transformer_2") + +def save_results(): + if not os.path.exists(save_path): + os.makedirs(save_path, exist_ok=True) + + index = len([path for path in os.listdir(save_path)]) + 1 + prefix = str(index).zfill(8) + if video_length == 1: + video_path = os.path.join(save_path, prefix + ".png") + + image = sample[0, :, 0] + image = image.transpose(0, 1).transpose(1, 2) + image = (image * 255).numpy().astype(np.uint8) + image = Image.fromarray(image) + image.save(video_path) + else: + video_path = os.path.join(save_path, prefix + ".mp4") + save_videos_grid(sample, video_path, fps=fps) + +if ulysses_degree * ring_degree > 1: + import torch.distributed as dist + if dist.get_rank() == 0: + save_results() +else: + save_results() \ No newline at end of file diff --git a/examples/wan2.2_fun/predict_v2v_control_ref.py b/examples/wan2.2_fun/predict_v2v_control_ref.py index ec18805..b201673 100644 --- a/examples/wan2.2_fun/predict_v2v_control_ref.py +++ b/examples/wan2.2_fun/predict_v2v_control_ref.py @@ -14,7 +14,7 @@ for project_root in project_roots: sys.path.insert(0, project_root) if project_root not in sys.path else None from videox_fun.dist import set_multi_gpus_devices, shard_model -from videox_fun.models import (AutoencoderKLWan, AutoTokenizer, CLIPModel, +from videox_fun.models import (AutoencoderKLWan, AutoencoderKLWan3_8, AutoTokenizer, CLIPModel, WanT5EncoderModel, Wan2_2Transformer3DModel) from videox_fun.data.dataset_image_video import process_pose_file from videox_fun.models.cache_utils import get_teacache_coefficients @@ -29,7 +29,7 @@ from videox_fun.utils.utils import (filter_kwargs, get_image_latent, get_image_t from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler -# GPU memory mode, which can be choosen in [model_full_load, model_cpu_offload_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8, sequential_cpu_offload]. +# GPU memory mode, which can be choosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8, sequential_cpu_offload]. # 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, @@ -128,7 +128,7 @@ negative_prompt = "色调艳丽,过曝,静态,细节模糊不清,字 # prompt = "A young woman with beautiful, clear eyes and blonde hair stands in the forest, wearing a white dress and a crown. Her expression is serene, reminiscent of a movie star, with fair and youthful skin. Her brown long hair flows in the wind. The video quality is very high, with a clear view. High quality, masterpiece, best quality, high resolution, ultra-fine, fantastical." # negative_prompt = "Twisted body, limb deformities, text captions, comic, static, ugly, error, messy code." guidance_scale = 6.0 -seed = 42 +seed = 43 num_inference_steps = 50 # The lora_weight is used for low noise model, the lora_high_weight is used for high noise model. lora_weight = 0.55 @@ -145,13 +145,15 @@ transformer = Wan2_2Transformer3DModel.from_pretrained( low_cpu_mem_usage=True, torch_dtype=weight_dtype, ) - -transformer_2 = Wan2_2Transformer3DModel.from_pretrained( - os.path.join(model_name, config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer')), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - low_cpu_mem_usage=True, - torch_dtype=weight_dtype, -) +if config['transformer_additional_kwargs'].get('transformer_combination_type', 'single') == "moe": + transformer_2 = Wan2_2Transformer3DModel.from_pretrained( + os.path.join(model_name, config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer')), + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), + low_cpu_mem_usage=True, + torch_dtype=weight_dtype, + ) +else: + transformer_2 = None if transformer_path is not None: print(f"From checkpoint: {transformer_path}") @@ -165,21 +167,23 @@ if transformer_path is not None: m, u = transformer.load_state_dict(state_dict, strict=False) print(f"missing keys: {len(m)}, unexpected keys: {len(u)}") -if transformer_high_path is not None: - print(f"From checkpoint: {transformer_high_path}") - if transformer_high_path.endswith("safetensors"): - from safetensors.torch import load_file, safe_open - state_dict = load_file(transformer_high_path) - else: - state_dict = torch.load(transformer_high_path, map_location="cpu") - state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict +if transformer_2 is not None: + if transformer_high_path is not None: + print(f"From checkpoint: {transformer_high_path}") + if transformer_high_path.endswith("safetensors"): + from safetensors.torch import load_file, safe_open + state_dict = load_file(transformer_high_path) + else: + state_dict = torch.load(transformer_high_path, map_location="cpu") + state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict - m, u = transformer_2.load_state_dict(state_dict, strict=False) - print(f"missing keys: {len(m)}, unexpected keys: {len(u)}") + m, u = transformer_2.load_state_dict(state_dict, strict=False) + print(f"missing keys: {len(m)}, unexpected keys: {len(u)}") # Get Vae Choosen_AutoencoderKL = { "AutoencoderKLWan": AutoencoderKLWan, + "AutoencoderKLWan3_8": AutoencoderKLWan3_8 }[config['vae_kwargs'].get('vae_type', 'AutoencoderKLWan')] vae = Choosen_AutoencoderKL.from_pretrained( os.path.join(model_name, config['vae_kwargs'].get('vae_subpath', 'vae')), @@ -236,11 +240,13 @@ pipeline = Wan2_2FunControlPipeline( if ulysses_degree > 1 or ring_degree > 1: from functools import partial transformer.enable_multi_gpus_inference() - transformer_2.enable_multi_gpus_inference() + if transformer_2 is not None: + transformer_2.enable_multi_gpus_inference() if fsdp_dit: shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype) pipeline.transformer = shard_fn(pipeline.transformer) - pipeline.transformer_2 = shard_fn(pipeline.transformer_2) + if transformer_2 is not None: + pipeline.transformer_2 = shard_fn(pipeline.transformer_2) print("Add FSDP DIT") if fsdp_text_encoder: shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype) @@ -250,29 +256,33 @@ if ulysses_degree > 1 or ring_degree > 1: if compile_dit: for i in range(len(pipeline.transformer.blocks)): pipeline.transformer.blocks[i] = torch.compile(pipeline.transformer.blocks[i]) - for i in range(len(pipeline.transformer_2.blocks)): - pipeline.transformer_2.blocks[i] = torch.compile(pipeline.transformer_2.blocks[i]) + if transformer_2 is not None: + for i in range(len(pipeline.transformer_2.blocks)): + pipeline.transformer_2.blocks[i] = torch.compile(pipeline.transformer_2.blocks[i]) print("Add Compile") if GPU_memory_mode == "sequential_cpu_offload": replace_parameters_by_name(transformer, ["modulation",], device=device) - replace_parameters_by_name(transformer_2, ["modulation",], device=device) transformer.freqs = transformer.freqs.to(device=device) - transformer_2.freqs = transformer_2.freqs.to(device=device) + if transformer_2 is not None: + replace_parameters_by_name(transformer_2, ["modulation",], device=device) + transformer_2.freqs = transformer_2.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",], device=device) - convert_model_weight_to_float8(transformer_2, exclude_module_name=["modulation",], device=device) convert_weight_dtype_wrapper(transformer, weight_dtype) - convert_weight_dtype_wrapper(transformer_2, weight_dtype) + if transformer_2 is not None: + convert_model_weight_to_float8(transformer_2, exclude_module_name=["modulation",], device=device) + convert_weight_dtype_wrapper(transformer_2, 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=["modulation",], device=device) - convert_model_weight_to_float8(transformer_2, exclude_module_name=["modulation",], device=device) convert_weight_dtype_wrapper(transformer, weight_dtype) - convert_weight_dtype_wrapper(transformer_2, weight_dtype) + if transformer_2 is not None: + convert_model_weight_to_float8(transformer_2, exclude_module_name=["modulation",], device=device) + convert_weight_dtype_wrapper(transformer_2, weight_dtype) pipeline.to(device=device) else: pipeline.to(device=device) @@ -283,18 +293,21 @@ if coefficients is not None: pipeline.transformer.enable_teacache( coefficients, num_inference_steps, teacache_threshold, num_skip_start_steps=num_skip_start_steps, offload=teacache_offload ) - pipeline.transformer_2.share_teacache(transformer=pipeline.transformer) + if transformer_2 is not None: + pipeline.transformer_2.share_teacache(transformer=pipeline.transformer) 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) - pipeline.transformer_2.share_cfg_skip(transformer=pipeline.transformer) + if transformer_2 is not None: + pipeline.transformer_2.share_cfg_skip(transformer=pipeline.transformer) generator = torch.Generator(device=device).manual_seed(seed) if lora_path is not None: pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device) - pipeline = merge_lora(pipeline, lora_high_path, lora_high_weight, device=device, sub_transformer_name="transformer_2") + if transformer_2 is not None: + pipeline = merge_lora(pipeline, lora_high_path, lora_high_weight, device=device, sub_transformer_name="transformer_2") with torch.no_grad(): video_length = int((video_length - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if video_length != 1 else 1 @@ -302,6 +315,8 @@ with torch.no_grad(): if enable_riflex: pipeline.transformer.enable_riflex(k = riflex_k, L_test = latent_frames) + if transformer_2 is not None: + pipeline.transformer_2.enable_riflex(k = riflex_k, L_test = latent_frames) inpaint_video, inpaint_video_mask, clip_image = get_image_to_video_latent(start_image, end_image, video_length=video_length, sample_size=sample_size) @@ -337,7 +352,8 @@ with torch.no_grad(): if lora_path is not None: pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device) - pipeline = unmerge_lora(pipeline, lora_high_path, lora_high_weight, device=device, sub_transformer_name="transformer_2") + if transformer_2 is not None: + pipeline = unmerge_lora(pipeline, lora_high_path, lora_high_weight, device=device, sub_transformer_name="transformer_2") def save_results(): if not os.path.exists(save_path): diff --git a/scripts/wan2.1_fun/README_TRAIN.md b/scripts/wan2.1_fun/README_TRAIN.md index d680a86..906e023 100755 --- a/scripts/wan2.1_fun/README_TRAIN.md +++ b/scripts/wan2.1_fun/README_TRAIN.md @@ -135,7 +135,7 @@ export DATASET_META_NAME="datasets/internal_datasets/metadata.json" # 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_stage3_config.json --deepspeed_multinode_launcher standard scripts/wan2.1/train.py \ +accelerate launch --zero_stage 3 --zero3_save_16bit_model true --zero3_init_flag true --use_deepspeed --deepspeed_config_file config/zero_stage3_config.json --deepspeed_multinode_launcher standard scripts/wan2.1_fun/train.py \ --config_path="config/wan2.1/wan_civitai.yaml" \ --pretrained_model_name_or_path=$MODEL_NAME \ --train_data_dir=$DATASET_NAME \ @@ -184,7 +184,7 @@ export DATASET_META_NAME="datasets/internal_datasets/metadata.json" # export NCCL_P2P_DISABLE=1 NCCL_DEBUG=INFO -accelerate launch --mixed_precision="bf16" --use_fsdp --fsdp_auto_wrap_policy TRANSFORMER_BASED_WRAP --fsdp_transformer_layer_cls_to_wrap=WanAttentionBlock --fsdp_sharding_strategy "FULL_SHARD" --fsdp_state_dict_type=SHARDED_STATE_DICT --fsdp_backward_prefetch "BACKWARD_PRE" --fsdp_cpu_ram_efficient_loading False scripts/wan2.1/train.py \ +accelerate launch --mixed_precision="bf16" --use_fsdp --fsdp_auto_wrap_policy TRANSFORMER_BASED_WRAP --fsdp_transformer_layer_cls_to_wrap=WanAttentionBlock --fsdp_sharding_strategy "FULL_SHARD" --fsdp_state_dict_type=SHARDED_STATE_DICT --fsdp_backward_prefetch "BACKWARD_PRE" --fsdp_cpu_ram_efficient_loading False scripts/wan2.1_fun/train.py \ --config_path="config/wan2.1/wan_civitai.yaml" \ --pretrained_model_name_or_path=$MODEL_NAME \ --train_data_dir=$DATASET_NAME \ diff --git a/scripts/wan2.2_fun/README_TRAIN.md b/scripts/wan2.2_fun/README_TRAIN.md new file mode 100755 index 0000000..a99faa3 --- /dev/null +++ b/scripts/wan2.2_fun/README_TRAIN.md @@ -0,0 +1,227 @@ +## Training Code + +We can choose whether to use deep speed in Wan, which can save a lot of video memory. + +Some parameters in the sh file can be confusing, and they are explained in this document: + +- `enable_bucket` is used to enable bucket training. When enabled, the model does not crop the images and videos at the center, but instead, it trains the entire images and videos after grouping them into buckets based on resolution. +- `random_frame_crop` is used for random cropping on video frames to simulate videos with different frame counts. +- `random_hw_adapt` is used to enable automatic height and width scaling for images and videos. When `random_hw_adapt` is enabled, the training images will have their height and width set to `image_sample_size` as the maximum and `min(video_sample_size, 512)` as the minimum. For training videos, the height and width will be set to `image_sample_size` as the maximum and `min(video_sample_size, 512)` as the minimum. + - For example, when `random_hw_adapt` is enabled, with `video_sample_n_frames=49`, `video_sample_size=1024`, and `image_sample_size=1024`, the resolution of image inputs for training is `512x512` to `1024x1024`, and the resolution of video inputs for training is `512x512x49` to `1024x1024x49`. + - For example, when `random_hw_adapt` is enabled, with `video_sample_n_frames=49`, `video_sample_size=1024`, and `image_sample_size=256`, the resolution of image inputs for training is `256x256` to `1024x1024`, and the resolution of video inputs for training is `256x256x49`. +- `training_with_video_token_length` specifies training the model according to token length. For training images and videos, the height and width will be set to `image_sample_size` as the maximum and `video_sample_size` as the minimum. + - For example, when `training_with_video_token_length` is enabled, with `video_sample_n_frames=49`, `token_sample_size=1024`, `video_sample_size=1024`, and `image_sample_size=256`, the resolution of image inputs for training is `256x256` to `1024x1024`, and the resolution of video inputs for training is `256x256x49` to `1024x1024x49`. + - For example, when `training_with_video_token_length` is enabled, with `video_sample_n_frames=49`, `token_sample_size=512`, `video_sample_size=1024`, and `image_sample_size=256`, the resolution of image inputs for training is `256x256` to `1024x1024`, and the resolution of video inputs for training is `256x256x49` to `1024x1024x9`. + - The token length for a video with dimensions 512x512 and 49 frames is 13,312. We need to set the `token_sample_size = 512`. + - At 512x512 resolution, the number of video frames is 49 (~= 512 * 512 * 49 / 512 / 512). + - 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. +- `train_mode` is used to specify the training mode, which can be either normal or inpaint. Since Wan uses the inpaint model to achieve image-to-video generation, the default is set to inpaint mode. If you only wish to achieve text-to-video generation, you can remove this line, and it will default to the text-to-video mode. +- `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. +- `boundary_type`: The Wan2.2 series includes two distinct models that handle different noise levels, specified via the `boundary_type` parameter. `low`: Corresponds to the **low noise model** (low_noise_model). `high`: Corresponds to the **high noise model**. (high_noise_model). `full`: Corresponds to the ti2v 5B model (single mode). + +Wan T2V without deepspeed: + +Training 14B Wan2.2 without DeepSpeed may result in insufficient GPU memory. +```sh +export MODEL_NAME="models/Diffusion_Transformer/Wan2.2-Fun-A14B-InP" +export DATASET_NAME="datasets/internal_datasets/" +export DATASET_META_NAME="datasets/internal_datasets/metadata.json" +# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA. +# export NCCL_IB_DISABLE=1 +# export NCCL_P2P_DISABLE=1 +NCCL_DEBUG=INFO + +accelerate launch --mixed_precision="bf16" scripts/wan2.2_fun/train.py \ + --config_path="config/wan2.2/wan_civitai_i2v.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 \ + --boundary_type="low" \ + --low_vram \ + --train_mode="normal" \ + --trainable_modules "." +``` + +Wan T2V with deepspeed zero-2: + +Wan with DeepSpeed Zero-2 is suitable for training 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.2-Fun-A14B-InP" +export DATASET_NAME="datasets/internal_datasets/" +export DATASET_META_NAME="datasets/internal_datasets/metadata.json" +# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA. +# 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.2_fun/train.py \ + --config_path="config/wan2.2/wan_civitai_i2v.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 \ + --boundary_type="low" \ + --low_vram \ + --use_deepspeed \ + --train_mode="inpaint" \ + --trainable_modules "." +``` + +Wan T2V 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.2-Fun-A14B-InP" +export DATASET_NAME="datasets/internal_datasets/" +export DATASET_META_NAME="datasets/internal_datasets/metadata.json" +# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA. +# 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_stage3_config.json --deepspeed_multinode_launcher standard scripts/wan2.2_fun/train.py \ + --config_path="config/wan2.2/wan_civitai_i2v.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 \ + --boundary_type="low" \ + --low_vram \ + --use_deepspeed \ + --train_mode="inpaint" \ + --trainable_modules "." +``` + +Wan T2V with FSDP: + +Wan with FSDP is suitable for 14B Wan at high resolutions. Training shell command is as follows: +```sh +export MODEL_NAME="models/Diffusion_Transformer/Wan2.2-Fun-A14B-InP" +export DATASET_NAME="datasets/internal_datasets/" +export DATASET_META_NAME="datasets/internal_datasets/metadata.json" +# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA. +# export NCCL_IB_DISABLE=1 +# export NCCL_P2P_DISABLE=1 +NCCL_DEBUG=INFO + +accelerate launch --mixed_precision="bf16" --use_fsdp --fsdp_auto_wrap_policy TRANSFORMER_BASED_WRAP --fsdp_transformer_layer_cls_to_wrap=WanAttentionBlock --fsdp_sharding_strategy "FULL_SHARD" --fsdp_state_dict_type=SHARDED_STATE_DICT --fsdp_backward_prefetch "BACKWARD_PRE" --fsdp_cpu_ram_efficient_loading False scripts/wan2.2_fun/train.py \ + --config_path="config/wan2.2/wan_civitai_i2v.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 \ + --boundary_type="low" \ + --low_vram \ + --use_deepspeed \ + --train_mode="inpaint" \ + --trainable_modules "." +``` \ No newline at end of file diff --git a/scripts/wan2.2_fun/README_TRAIN_CONTROL.md b/scripts/wan2.2_fun/README_TRAIN_CONTROL.md new file mode 100755 index 0000000..f67aecf --- /dev/null +++ b/scripts/wan2.2_fun/README_TRAIN_CONTROL.md @@ -0,0 +1,272 @@ +## Training Code + +We can choose whether to use deep speed in Wan-Fun, which can save a lot of video memory. + +The metadata_control.json is a little different from normal json in Wan-Fun, you need to add a control_file_path, and [DWPose](https://github.com/IDEA-Research/DWPose) is suggested as tool to generate control file. + +```json +[ + { + "file_path": "train/00000001.mp4", + "control_file_path": "control/00000001.mp4", + "text": "A group of young men in suits and sunglasses are walking down a city street.", + "type": "video" + }, + { + "file_path": "train/00000002.jpg", + "control_file_path": "control/00000002.jpg", + "text": "A group of young men in suits and sunglasses are walking down a city street.", + "type": "image" + }, + ..... +] +``` + +Some parameters in the sh file can be confusing, and they are explained in this document: + +- `enable_bucket` is used to enable bucket training. When enabled, the model does not crop the images and videos at the center, but instead, it trains the entire images and videos after grouping them into buckets based on resolution. +- `random_frame_crop` is used for random cropping on video frames to simulate videos with different frame counts. +- `random_hw_adapt` is used to enable automatic height and width scaling for images and videos. When `random_hw_adapt` is enabled, the training images will have their height and width set to `image_sample_size` as the maximum and `min(video_sample_size, 512)` as the minimum. For training videos, the height and width will be set to `image_sample_size` as the maximum and `min(video_sample_size, 512)` as the minimum. + - For example, when `random_hw_adapt` is enabled, with `video_sample_n_frames=49`, `video_sample_size=1024`, and `image_sample_size=1024`, the resolution of image inputs for training is `512x512` to `1024x1024`, and the resolution of video inputs for training is `512x512x49` to `1024x1024x49`. + - For example, when `random_hw_adapt` is enabled, with `video_sample_n_frames=49`, `video_sample_size=1024`, and `image_sample_size=256`, the resolution of image inputs for training is `256x256` to `1024x1024`, and the resolution of video inputs for training is `256x256x49`. +- `training_with_video_token_length` specifies training the model according to token length. For training images and videos, the height and width will be set to `image_sample_size` as the maximum and `video_sample_size` as the minimum. + - For example, when `training_with_video_token_length` is enabled, with `video_sample_n_frames=49`, `token_sample_size=1024`, `video_sample_size=1024`, and `image_sample_size=256`, the resolution of image inputs for training is `256x256` to `1024x1024`, and the resolution of video inputs for training is `256x256x49` to `1024x1024x49`. + - For example, when `training_with_video_token_length` is enabled, with `video_sample_n_frames=49`, `token_sample_size=512`, `video_sample_size=1024`, and `image_sample_size=256`, the resolution of image inputs for training is `256x256` to `1024x1024`, and the resolution of video inputs for training is `256x256x49` to `1024x1024x9`. + - The token length for a video with dimensions 512x512 and 49 frames is 13,312. We need to set the `token_sample_size = 512`. + - At 512x512 resolution, the number of video frames is 49 (~= 512 * 512 * 49 / 512 / 512). + - 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. +- `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. +- `add_inpaint_info` determines whether to incorporate inpaint information into the model training. When enabled, this allows the model to support specifying starting and ending images in the controls during generation. +- `boundary_type`: The Wan2.2 series includes two distinct models that handle different noise levels, specified via the `boundary_type` parameter. `low`: Corresponds to the **low noise model** (low_noise_model). `high`: Corresponds to the **high noise model**. (high_noise_model). `full`: Corresponds to the ti2v 5B model (single mode). + +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.2_fun/xxx.py +``` + +Wan-Fun-Control without deepspeed: +```sh +export MODEL_NAME="models/Diffusion_Transformer/Wan2.2-Fun-A14B-Control" +export DATASET_NAME="datasets/internal_datasets/" +export DATASET_META_NAME="datasets/internal_datasets/metadata.json" +# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA. +# export NCCL_IB_DISABLE=1 +# export NCCL_P2P_DISABLE=1 +NCCL_DEBUG=INFO + +accelerate launch --mixed_precision="bf16" scripts/wan2.2_fun/train_control.py \ + --config_path="config/wan2.2/wan_civitai_i2v.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 \ + --boundary_type="low" \ + --low_vram \ + --train_mode="control_ref" \ + --control_ref_image="random" \ + --add_inpaint_info \ + --add_full_ref_image_in_self_attention \ + --trainable_modules "." +``` + +Wan-Fun-Control with deepspeed: +```sh +export MODEL_NAME="models/Diffusion_Transformer/Wan2.2-Fun-A14B-Control" +export DATASET_NAME="datasets/internal_datasets/" +export DATASET_META_NAME="datasets/internal_datasets/metadata.json" +# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA. +# 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.2_fun/train_control.py \ + --config_path="config/wan2.2/wan_civitai_i2v.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 \ + --boundary_type="low" \ + --low_vram \ + --use_deepspeed \ + --train_mode="control_ref" \ + --control_ref_image="random" \ + --add_inpaint_info \ + --add_full_ref_image_in_self_attention \ + --trainable_modules "." +``` + +Wan-Fun-Control 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.2-Fun-A14B-Control" +export DATASET_NAME="datasets/internal_datasets/" +export DATASET_META_NAME="datasets/internal_datasets/metadata.json" +# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA. +# 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_stage3_config.json --deepspeed_multinode_launcher standard scripts/wan2.2_fun/train_control.py \ + --config_path="config/wan2.2/wan_civitai_i2v.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 \ + --boundary_type="low" \ + --low_vram \ + --use_deepspeed \ + --train_mode="control_ref" \ + --control_ref_image="random" \ + --add_inpaint_info \ + --add_full_ref_image_in_self_attention \ + --trainable_modules "." +``` + +Wan-Fun-Control with FSDP: + +Wan with FSDP is suitable for 14B Wan at high resolutions. Training shell command is as follows: +```sh +export MODEL_NAME="models/Diffusion_Transformer/Wan2.2-Fun-A14B-Control" +export DATASET_NAME="datasets/internal_datasets/" +export DATASET_META_NAME="datasets/internal_datasets/metadata.json" +# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA. +# export NCCL_IB_DISABLE=1 +# export NCCL_P2P_DISABLE=1 +NCCL_DEBUG=INFO + +accelerate launch --mixed_precision="bf16" --use_fsdp --fsdp_auto_wrap_policy TRANSFORMER_BASED_WRAP --fsdp_transformer_layer_cls_to_wrap=WanAttentionBlock --fsdp_sharding_strategy "FULL_SHARD" --fsdp_state_dict_type=SHARDED_STATE_DICT --fsdp_backward_prefetch "BACKWARD_PRE" --fsdp_cpu_ram_efficient_loading False scripts/wan2.2_fun/train_control.py \ + --config_path="config/wan2.2/wan_civitai_i2v.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 \ + --boundary_type="low" \ + --low_vram \ + --use_deepspeed \ + --train_mode="control_ref" \ + --control_ref_image="random" \ + --add_inpaint_info \ + --add_full_ref_image_in_self_attention \ + --trainable_modules "." +``` \ No newline at end of file diff --git a/scripts/wan2.2_fun/README_TRAIN_CONTROL_LORA.md b/scripts/wan2.2_fun/README_TRAIN_CONTROL_LORA.md new file mode 100755 index 0000000..cb734fe --- /dev/null +++ b/scripts/wan2.2_fun/README_TRAIN_CONTROL_LORA.md @@ -0,0 +1,262 @@ +## Training Code + +We can choose whether to use deep speed in Wan-Fun, which can save a lot of video memory. + +The metadata_control.json is a little different from normal json in Wan-Fun, you need to add a control_file_path, and [DWPose](https://github.com/IDEA-Research/DWPose) is suggested as tool to generate control file. + +```json +[ + { + "file_path": "train/00000001.mp4", + "control_file_path": "control/00000001.mp4", + "text": "A group of young men in suits and sunglasses are walking down a city street.", + "type": "video" + }, + { + "file_path": "train/00000002.jpg", + "control_file_path": "control/00000002.jpg", + "text": "A group of young men in suits and sunglasses are walking down a city street.", + "type": "image" + }, + ..... +] +``` + +Some parameters in the sh file can be confusing, and they are explained in this document: + +- `enable_bucket` is used to enable bucket training. When enabled, the model does not crop the images and videos at the center, but instead, it trains the entire images and videos after grouping them into buckets based on resolution. +- `random_frame_crop` is used for random cropping on video frames to simulate videos with different frame counts. +- `random_hw_adapt` is used to enable automatic height and width scaling for images and videos. When `random_hw_adapt` is enabled, the training images will have their height and width set to `image_sample_size` as the maximum and `min(video_sample_size, 512)` as the minimum. For training videos, the height and width will be set to `image_sample_size` as the maximum and `min(video_sample_size, 512)` as the minimum. + - For example, when `random_hw_adapt` is enabled, with `video_sample_n_frames=49`, `video_sample_size=1024`, and `image_sample_size=1024`, the resolution of image inputs for training is `512x512` to `1024x1024`, and the resolution of video inputs for training is `512x512x49` to `1024x1024x49`. + - For example, when `random_hw_adapt` is enabled, with `video_sample_n_frames=49`, `video_sample_size=1024`, and `image_sample_size=256`, the resolution of image inputs for training is `256x256` to `1024x1024`, and the resolution of video inputs for training is `256x256x49`. +- `training_with_video_token_length` specifies training the model according to token length. For training images and videos, the height and width will be set to `image_sample_size` as the maximum and `video_sample_size` as the minimum. + - For example, when `training_with_video_token_length` is enabled, with `video_sample_n_frames=49`, `token_sample_size=1024`, `video_sample_size=1024`, and `image_sample_size=256`, the resolution of image inputs for training is `256x256` to `1024x1024`, and the resolution of video inputs for training is `256x256x49` to `1024x1024x49`. + - For example, when `training_with_video_token_length` is enabled, with `video_sample_n_frames=49`, `token_sample_size=512`, `video_sample_size=1024`, and `image_sample_size=256`, the resolution of image inputs for training is `256x256` to `1024x1024`, and the resolution of video inputs for training is `256x256x49` to `1024x1024x9`. + - The token length for a video with dimensions 512x512 and 49 frames is 13,312. We need to set the `token_sample_size = 512`. + - At 512x512 resolution, the number of video frames is 49 (~= 512 * 512 * 49 / 512 / 512). + - 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 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. +- `add_inpaint_info` determines whether to incorporate inpaint information into the model training. When enabled, this allows the model to support specifying starting and ending images in the controls during generation. +- `boundary_type`: The Wan2.2 series includes two distinct models that handle different noise levels, specified via the `boundary_type` parameter. `low`: Corresponds to the **low noise model** (low_noise_model). `high`: Corresponds to the **high noise model**. (high_noise_model). `full`: Corresponds to the ti2v 5B model (single mode). + +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.2_fun/xxx.py +``` + +Wan-Fun-Control without deepspeed: +```sh +export MODEL_NAME="models/Diffusion_Transformer/Wan2.2-Fun-A14B-Control" +export DATASET_NAME="datasets/internal_datasets/" +export DATASET_META_NAME="datasets/internal_datasets/metadata.json" +# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA. +# export NCCL_IB_DISABLE=1 +# export NCCL_P2P_DISABLE=1 +NCCL_DEBUG=INFO + +accelerate launch --mixed_precision="bf16" scripts/wan2.2_fun/train_control_lora.py \ + --config_path="config/wan2.2/wan_civitai_i2v.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 \ + --boundary_type="low" \ + --train_mode="control_ref" \ + --control_ref_image="random" \ + --add_inpaint_info \ + --add_full_ref_image_in_self_attention \ + --low_vram +``` + +Wan-Fun-Control with deepspeed: +```sh +export MODEL_NAME="models/Diffusion_Transformer/Wan2.2-Fun-A14B-Control" +export DATASET_NAME="datasets/internal_datasets/" +export DATASET_META_NAME="datasets/internal_datasets/metadata.json" +# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA. +# 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.2_fun/train_control_lora.py \ + --config_path="config/wan2.2/wan_civitai_i2v.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 \ + --boundary_type="low" \ + --use_deepspeed \ + --train_mode="control_ref" \ + --control_ref_image="random" \ + --add_inpaint_info \ + --add_full_ref_image_in_self_attention \ + --low_vram +``` + +Wan-Fun-Control 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.2-Fun-A14B-Control" +export DATASET_NAME="datasets/internal_datasets/" +export DATASET_META_NAME="datasets/internal_datasets/metadata.json" +# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA. +# 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_stage3_config.json --deepspeed_multinode_launcher standard scripts/wan2.2_fun/train_control_lora.py \ + --config_path="config/wan2.2/wan_civitai_i2v.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 \ + --boundary_type="low" \ + --save_state \ + --use_deepspeed \ + --train_mode="control_ref" \ + --control_ref_image="random" \ + --add_inpaint_info \ + --add_full_ref_image_in_self_attention \ + --low_vram +``` + +Wan-Fun-Control with FSDP: + +Wan with FSDP is suitable for 14B Wan at high resolutions. Training shell command is as follows: +```sh +export MODEL_NAME="models/Diffusion_Transformer/Wan2.2-Fun-A14B-Control" +export DATASET_NAME="datasets/internal_datasets/" +export DATASET_META_NAME="datasets/internal_datasets/metadata.json" +# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA. +# export NCCL_IB_DISABLE=1 +# export NCCL_P2P_DISABLE=1 +NCCL_DEBUG=INFO + +accelerate launch --mixed_precision="bf16" --use_fsdp --fsdp_auto_wrap_policy TRANSFORMER_BASED_WRAP --fsdp_transformer_layer_cls_to_wrap=WanAttentionBlock --fsdp_sharding_strategy "FULL_SHARD" --fsdp_state_dict_type=SHARDED_STATE_DICT --fsdp_backward_prefetch "BACKWARD_PRE" --fsdp_cpu_ram_efficient_loading False scripts/wan2.2_fun/train_control_lora.py \ + --config_path="config/wan2.2/wan_civitai_i2v.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 \ + --boundary_type="low" \ + --save_state \ + --use_fsdp \ + --train_mode="control_ref" \ + --control_ref_image="random" \ + --add_inpaint_info \ + --add_full_ref_image_in_self_attention \ + --low_vram +``` \ No newline at end of file diff --git a/scripts/wan2.2_fun/README_TRAIN_LORA.md b/scripts/wan2.2_fun/README_TRAIN_LORA.md new file mode 100755 index 0000000..95132d2 --- /dev/null +++ b/scripts/wan2.2_fun/README_TRAIN_LORA.md @@ -0,0 +1,217 @@ +## Lora Training Code + +We can choose whether to use deep speed in Wan, which can save a lot of video memory. + +Some parameters in the sh file can be confusing, and they are explained in this document: + +- `enable_bucket` is used to enable bucket training. When enabled, the model does not crop the images and videos at the center, but instead, it trains the entire images and videos after grouping them into buckets based on resolution. +- `random_frame_crop` is used for random cropping on video frames to simulate videos with different frame counts. +- `random_hw_adapt` is used to enable automatic height and width scaling for images and videos. When `random_hw_adapt` is enabled, the training images will have their height and width set to `image_sample_size` as the maximum and `min(video_sample_size, 512)` as the minimum. For training videos, the height and width will be set to `image_sample_size` as the maximum and `min(video_sample_size, 512)` as the minimum. + - For example, when `random_hw_adapt` is enabled, with `video_sample_n_frames=49`, `video_sample_size=1024`, and `image_sample_size=1024`, the resolution of image inputs for training is `512x512` to `1024x1024`, and the resolution of video inputs for training is `512x512x49` to `1024x1024x49`. + - For example, when `random_hw_adapt` is enabled, with `video_sample_n_frames=49`, `video_sample_size=1024`, and `image_sample_size=256`, the resolution of image inputs for training is `256x256` to `1024x1024`, and the resolution of video inputs for training is `256x256x49`. +- `training_with_video_token_length` specifies training the model according to token length. For training images and videos, the height and width will be set to `image_sample_size` as the maximum and `video_sample_size` as the minimum. + - For example, when `training_with_video_token_length` is enabled, with `video_sample_n_frames=49`, `token_sample_size=1024`, `video_sample_size=1024`, and `image_sample_size=256`, the resolution of image inputs for training is `256x256` to `1024x1024`, and the resolution of video inputs for training is `256x256x49` to `1024x1024x49`. + - For example, when `training_with_video_token_length` is enabled, with `video_sample_n_frames=49`, `token_sample_size=512`, `video_sample_size=1024`, and `image_sample_size=256`, the resolution of image inputs for training is `256x256` to `1024x1024`, and the resolution of video inputs for training is `256x256x49` to `1024x1024x9`. + - The token length for a video with dimensions 512x512 and 49 frames is 13,312. We need to set the `token_sample_size = 512`. + - At 512x512 resolution, the number of video frames is 49 (~= 512 * 512 * 49 / 512 / 512). + - 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. +- `train_mode` is used to specify the training mode, which can be either normal or inpaint. Since Wan uses the inpaint model to achieve image-to-video generation, the default is set to inpaint mode. If you only wish to achieve text-to-video generation, you can remove this line, and it will default to the text-to-video mode. +- `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. +- `boundary_type`: The Wan2.2 series includes two distinct models that handle different noise levels, specified via the `boundary_type` parameter. `low`: Corresponds to the **low noise model** (low_noise_model). `high`: Corresponds to the **high noise model**. (high_noise_model). `full`: Corresponds to the ti2v 5B model (single mode). + +Wan T2V without deepspeed: + +Training 14B Wan2.2 without DeepSpeed may result in insufficient GPU memory. +```sh +export MODEL_NAME="models/Diffusion_Transformer/Wan2.2-Fun-A14B-InP" +export DATASET_NAME="datasets/internal_datasets/" +export DATASET_META_NAME="datasets/internal_datasets/metadata.json" +# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA. +# export NCCL_IB_DISABLE=1 +# export NCCL_P2P_DISABLE=1 +NCCL_DEBUG=INFO + +accelerate launch --mixed_precision="bf16" scripts/wan2.2_fun/train_lora.py \ + --config_path="config/wan2.2/wan_civitai_i2v.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 \ + --boundary_type="low" \ + --train_mode="inpaint" \ + --low_vram +``` + +Wan T2V with deepspeed zero-2: + +Wan with DeepSpeed Zero-2 is suitable for training 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.2-Fun-A14B-InP" +export DATASET_NAME="datasets/internal_datasets/" +export DATASET_META_NAME="datasets/internal_datasets/metadata.json" +# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA. +# 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.2_fun/train_lora.py \ + --config_path="config/wan2.2/wan_civitai_i2v.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 \ + --boundary_type="low" \ + --use_deepspeed \ + --train_mode="inpaint" \ + --low_vram +``` + +Wan T2V 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.2-Fun-A14B-InP" +export DATASET_NAME="datasets/internal_datasets/" +export DATASET_META_NAME="datasets/internal_datasets/metadata.json" +# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA. +# 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_stage3_config.json --deepspeed_multinode_launcher standard scripts/wan2.2_fun/train_lora.py \ + --config_path="config/wan2.2/wan_civitai_i2v.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 \ + --boundary_type="low" \ + --save_state \ + --use_deepspeed \ + --train_mode="inpaint" \ + --low_vram +``` + +Wan T2V with FSDP: + +Wan with FSDP is suitable for 14B Wan at high resolutions. Training shell command is as follows: +```sh +export MODEL_NAME="models/Diffusion_Transformer/Wan2.2-Fun-A14B-InP" +export DATASET_NAME="datasets/internal_datasets/" +export DATASET_META_NAME="datasets/internal_datasets/metadata.json" +# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA. +# export NCCL_IB_DISABLE=1 +# export NCCL_P2P_DISABLE=1 +NCCL_DEBUG=INFO + +accelerate launch --mixed_precision="bf16" --use_fsdp --fsdp_auto_wrap_policy TRANSFORMER_BASED_WRAP --fsdp_transformer_layer_cls_to_wrap=WanAttentionBlock --fsdp_sharding_strategy "FULL_SHARD" --fsdp_state_dict_type=SHARDED_STATE_DICT --fsdp_backward_prefetch "BACKWARD_PRE" --fsdp_cpu_ram_efficient_loading False scripts/wan2.2_fun/train_lora.py \ + --config_path="config/wan2.2/wan_civitai_i2v.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 \ + --boundary_type="low" \ + --save_state \ + --use_deepspeed \ + --train_mode="inpaint" \ + --low_vram +``` \ No newline at end of file diff --git a/scripts/wan2.2_fun/train.py b/scripts/wan2.2_fun/train.py new file mode 100644 index 0000000..f169f15 --- /dev/null +++ b/scripts/wan2.2_fun/train.py @@ -0,0 +1,1916 @@ +"""Modified from https://github.com/huggingface/diffusers/blob/main/examples/text_to_image/train_text_to_image.py +""" +#!/usr/bin/env python +# coding=utf-8 +# Copyright 2024 The HuggingFace Inc. team. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and + +import argparse +import gc +import logging +import math +import os +import pickle +import random +import shutil +import sys + +import accelerate +import diffusers +import numpy as np +import torch +import torch.nn.functional as F +import torch.utils.checkpoint +import torchvision.transforms.functional as TF +import transformers +from accelerate import Accelerator, FullyShardedDataParallelPlugin +from accelerate.logging import get_logger +from accelerate.state import AcceleratorState +from accelerate.utils import ProjectConfiguration, set_seed +from diffusers import DDIMScheduler, FlowMatchEulerDiscreteScheduler +from diffusers.optimization import get_scheduler +from diffusers.training_utils import (EMAModel, + compute_density_for_timestep_sampling, + compute_loss_weighting_for_sd3) +from diffusers.utils import check_min_version, deprecate, is_wandb_available +from diffusers.utils.torch_utils import is_compiled_module +from einops import rearrange +from omegaconf import OmegaConf +from packaging import version +from PIL import Image +from torch.distributed.fsdp.fully_sharded_data_parallel import ( + FullOptimStateDictConfig, FullStateDictConfig, ShardedOptimStateDictConfig, + ShardedStateDictConfig) +from torch.utils.data import RandomSampler +from torch.utils.tensorboard import SummaryWriter +from torchvision import transforms +from tqdm.auto import tqdm +from transformers import AutoTokenizer +from transformers.utils import ContextManagers + +import datasets + +current_file_path = os.path.abspath(__file__) +project_roots = [os.path.dirname(current_file_path), os.path.dirname(os.path.dirname(current_file_path)), os.path.dirname(os.path.dirname(os.path.dirname(current_file_path)))] +for project_root in project_roots: + sys.path.insert(0, project_root) if project_root not in sys.path else None + +from videox_fun.data.bucket_sampler import (ASPECT_RATIO_512, + ASPECT_RATIO_RANDOM_CROP_512, + ASPECT_RATIO_RANDOM_CROP_PROB, + AspectRatioBatchImageVideoSampler, + RandomSampler, get_closest_ratio) +from videox_fun.data.dataset_image_video import (ImageVideoDataset, + ImageVideoSampler, + get_random_mask) +from videox_fun.models import (AutoencoderKLWan, CLIPModel, WanT5EncoderModel, + Wan2_2Transformer3DModel) +from videox_fun.pipeline import WanFunInpaintPipeline, WanFunPipeline +from videox_fun.utils.discrete_sampler import DiscreteSampling +from videox_fun.utils.utils import get_image_to_video_latent, save_videos_grid + +if is_wandb_available(): + import wandb + + +def filter_kwargs(cls, kwargs): + import inspect + sig = inspect.signature(cls.__init__) + valid_params = set(sig.parameters.keys()) - {'self', 'cls'} + filtered_kwargs = {k: v for k, v in kwargs.items() if k in valid_params} + return filtered_kwargs + +def resize_mask(mask, latent, process_first_frame_only=True): + latent_size = latent.size() + batch_size, channels, num_frames, height, width = mask.shape + + if process_first_frame_only: + target_size = list(latent_size[2:]) + target_size[0] = 1 + first_frame_resized = F.interpolate( + mask[:, :, 0:1, :, :], + size=target_size, + mode='trilinear', + align_corners=False + ) + + target_size = list(latent_size[2:]) + target_size[0] = target_size[0] - 1 + if target_size[0] != 0: + remaining_frames_resized = F.interpolate( + mask[:, :, 1:, :, :], + size=target_size, + mode='trilinear', + align_corners=False + ) + resized_mask = torch.cat([first_frame_resized, remaining_frames_resized], dim=2) + else: + resized_mask = first_frame_resized + else: + target_size = list(latent_size[2:]) + resized_mask = F.interpolate( + mask, + size=target_size, + mode='trilinear', + align_corners=False + ) + return resized_mask + +def linear_decay(initial_value, final_value, total_steps, current_step): + if current_step >= total_steps: + return final_value + current_step = max(0, current_step) + step_size = (final_value - initial_value) / total_steps + current_value = initial_value + step_size * current_step + return current_value + +def generate_timestep_with_lognorm(low, high, shape, device="cpu", generator=None): + u = torch.normal(mean=0.0, std=1.0, size=shape, device=device, generator=generator) + t = 1 / (1 + torch.exp(-u)) * (high - low) + low + return torch.clip(t.to(torch.int32), low, high - 1) + +# Will error if the minimal version of diffusers is not installed. Remove at your own risks. +check_min_version("0.18.0.dev0") + +logger = get_logger(__name__, log_level="INFO") + +def log_validation(vae, text_encoder, tokenizer, transformer3d, args, config, accelerator, weight_dtype, global_step): + try: + logger.info("Running validation... ") + + transformer3d_val = Wan2_2Transformer3DModel.from_pretrained( + os.path.join(args.pretrained_model_name_or_path, config['transformer_additional_kwargs'].get('transformer_subpath', 'transformer')), + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), + ).to(weight_dtype) + transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) + scheduler = FlowMatchEulerDiscreteScheduler( + **filter_kwargs(FlowMatchEulerDiscreteScheduler, OmegaConf.to_container(config['scheduler_kwargs'])) + ) + + if args.train_mode != "normal": + pipeline = WanFunInpaintPipeline( + vae=accelerator.unwrap_model(vae).to(weight_dtype), + text_encoder=accelerator.unwrap_model(text_encoder), + tokenizer=tokenizer, + transformer=transformer3d_val, + scheduler=scheduler, + ) + else: + pipeline = WanFunPipeline( + vae=accelerator.unwrap_model(vae).to(weight_dtype), + text_encoder=accelerator.unwrap_model(text_encoder), + tokenizer=tokenizer, + transformer=transformer3d_val, + scheduler=scheduler, + ) + pipeline = pipeline.to(accelerator.device) + + if args.seed is None: + generator = None + else: + generator = torch.Generator(device=accelerator.device).manual_seed(args.seed) + + images = [] + for i in range(len(args.validation_prompts)): + with torch.no_grad(): + if args.train_mode != "normal": + with torch.autocast("cuda", dtype=weight_dtype): + video_length = int((args.video_sample_n_frames - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if args.video_sample_n_frames != 1 else 1 + input_video, input_video_mask, _ = get_image_to_video_latent(None, None, video_length=video_length, sample_size=[args.video_sample_size, args.video_sample_size]) + sample = pipeline( + args.validation_prompts[i], + num_frames = video_length, + negative_prompt = "bad detailed", + height = args.video_sample_size, + width = args.video_sample_size, + guidance_scale = 6.0, + generator = generator, + + video = input_video, + mask_video = input_video_mask, + ).videos + os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) + save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-{i}.gif")) + + video_length = 1 + input_video, input_video_mask, _ = get_image_to_video_latent(None, None, video_length=video_length, sample_size=[args.video_sample_size, args.video_sample_size]) + sample = pipeline( + args.validation_prompts[i], + num_frames = video_length, + negative_prompt = "bad detailed", + height = args.video_sample_size, + width = args.video_sample_size, + guidance_scale = 6.0, + generator = generator, + + video = input_video, + mask_video = input_video_mask, + ).videos + os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) + save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-image-{i}.gif")) + else: + with torch.autocast("cuda", dtype=weight_dtype): + sample = pipeline( + args.validation_prompts[i], + num_frames = args.video_sample_n_frames, + negative_prompt = "bad detailed", + height = args.video_sample_size, + width = args.video_sample_size, + generator = generator + ).videos + os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) + save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-{i}.gif")) + + sample = pipeline( + args.validation_prompts[i], + num_frames = args.video_sample_n_frames, + negative_prompt = "bad detailed", + height = args.video_sample_size, + width = args.video_sample_size, + generator = generator + ).videos + os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) + save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-image-{i}.gif")) + + del pipeline + del transformer3d_val + gc.collect() + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + + return images + except Exception as e: + gc.collect() + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + print(f"Eval error with info {e}") + return None + +def parse_args(): + parser = argparse.ArgumentParser(description="Simple example of a training script.") + parser.add_argument( + "--input_perturbation", type=float, default=0, help="The scale of input perturbation. Recommended 0.1." + ) + parser.add_argument( + "--pretrained_model_name_or_path", + type=str, + default=None, + required=True, + help="Path to pretrained model or model identifier from huggingface.co/models.", + ) + parser.add_argument( + "--revision", + type=str, + default=None, + required=False, + help="Revision of pretrained model identifier from huggingface.co/models.", + ) + parser.add_argument( + "--variant", + type=str, + default=None, + help="Variant of the model files of the pretrained model identifier from huggingface.co/models, 'e.g.' fp16", + ) + parser.add_argument( + "--train_data_dir", + type=str, + default=None, + help=( + "A folder containing the training data. " + ), + ) + parser.add_argument( + "--train_data_meta", + type=str, + default=None, + help=( + "A csv containing the training data. " + ), + ) + parser.add_argument( + "--max_train_samples", + type=int, + default=None, + help=( + "For debugging purposes or quicker training, truncate the number of training examples to this " + "value if set." + ), + ) + parser.add_argument( + "--validation_prompts", + type=str, + default=None, + nargs="+", + help=("A set of prompts evaluated every `--validation_epochs` and logged to `--report_to`."), + ) + parser.add_argument( + "--output_dir", + type=str, + default="sd-model-finetuned", + help="The output directory where the model predictions and checkpoints will be written.", + ) + parser.add_argument( + "--cache_dir", + type=str, + default=None, + help="The directory where the downloaded models and datasets will be stored.", + ) + parser.add_argument("--seed", type=int, default=None, help="A seed for reproducible training.") + parser.add_argument( + "--random_flip", + action="store_true", + help="whether to randomly flip images horizontally", + ) + parser.add_argument( + "--use_came", + action="store_true", + help="whether to use came", + ) + parser.add_argument( + "--multi_stream", + action="store_true", + help="whether to use cuda multi-stream", + ) + parser.add_argument( + "--train_batch_size", type=int, default=16, help="Batch size (per device) for the training dataloader." + ) + parser.add_argument( + "--vae_mini_batch", type=int, default=32, help="mini batch size for vae." + ) + parser.add_argument("--num_train_epochs", type=int, default=100) + parser.add_argument( + "--max_train_steps", + type=int, + default=None, + help="Total number of training steps to perform. If provided, overrides num_train_epochs.", + ) + parser.add_argument( + "--gradient_accumulation_steps", + type=int, + default=1, + help="Number of updates steps to accumulate before performing a backward/update pass.", + ) + parser.add_argument( + "--gradient_checkpointing", + action="store_true", + help="Whether or not to use gradient checkpointing to save memory at the expense of slower backward pass.", + ) + parser.add_argument( + "--learning_rate", + type=float, + default=1e-4, + help="Initial learning rate (after the potential warmup period) to use.", + ) + parser.add_argument( + "--scale_lr", + action="store_true", + default=False, + help="Scale the learning rate by the number of GPUs, gradient accumulation steps, and batch size.", + ) + parser.add_argument( + "--lr_scheduler", + type=str, + default="constant", + help=( + 'The scheduler type to use. Choose between ["linear", "cosine", "cosine_with_restarts", "polynomial",' + ' "constant", "constant_with_warmup"]' + ), + ) + parser.add_argument( + "--lr_warmup_steps", type=int, default=500, help="Number of steps for the warmup in the lr scheduler." + ) + parser.add_argument( + "--use_8bit_adam", action="store_true", help="Whether or not to use 8-bit Adam from bitsandbytes." + ) + parser.add_argument( + "--allow_tf32", + action="store_true", + help=( + "Whether or not to allow TF32 on Ampere GPUs. Can be used to speed up training. For more information, see" + " https://pytorch.org/docs/stable/notes/cuda.html#tensorfloat-32-tf32-on-ampere-devices" + ), + ) + parser.add_argument("--use_ema", action="store_true", help="Whether to use EMA model.") + parser.add_argument( + "--non_ema_revision", + type=str, + default=None, + required=False, + help=( + "Revision of pretrained non-ema model identifier. Must be a branch, tag or git identifier of the local or" + " remote repository specified with --pretrained_model_name_or_path." + ), + ) + parser.add_argument( + "--dataloader_num_workers", + type=int, + default=0, + help=( + "Number of subprocesses to use for data loading. 0 means that the data will be loaded in the main process." + ), + ) + parser.add_argument("--adam_beta1", type=float, default=0.9, help="The beta1 parameter for the Adam optimizer.") + parser.add_argument("--adam_beta2", type=float, default=0.999, help="The beta2 parameter for the Adam optimizer.") + parser.add_argument("--adam_weight_decay", type=float, default=1e-2, help="Weight decay to use.") + parser.add_argument("--adam_epsilon", type=float, default=1e-08, help="Epsilon value for the Adam optimizer") + parser.add_argument("--max_grad_norm", default=1.0, type=float, help="Max gradient norm.") + parser.add_argument("--push_to_hub", action="store_true", help="Whether or not to push the model to the Hub.") + parser.add_argument("--hub_token", type=str, default=None, help="The token to use to push to the Model Hub.") + parser.add_argument( + "--prediction_type", + type=str, + default=None, + help="The prediction_type that shall be used for training. Choose between 'epsilon' or 'v_prediction' or leave `None`. If left to `None` the default prediction type of the scheduler: `noise_scheduler.config.prediciton_type` is chosen.", + ) + parser.add_argument( + "--hub_model_id", + type=str, + default=None, + help="The name of the repository to keep in sync with the local `output_dir`.", + ) + parser.add_argument( + "--logging_dir", + type=str, + default="logs", + help=( + "[TensorBoard](https://www.tensorflow.org/tensorboard) log directory. Will default to" + " *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***." + ), + ) + parser.add_argument( + "--report_model_info", action="store_true", help="Whether or not to report more info about model (such as norm, grad)." + ) + parser.add_argument( + "--mixed_precision", + type=str, + default=None, + choices=["no", "fp16", "bf16"], + help=( + "Whether to use mixed precision. Choose between fp16 and bf16 (bfloat16). Bf16 requires PyTorch >=" + " 1.10.and an Nvidia Ampere GPU. Default to the value of accelerate config of the current system or the" + " flag passed with the `accelerate.launch` command. Use this argument to override the accelerate config." + ), + ) + parser.add_argument( + "--report_to", + type=str, + default="tensorboard", + help=( + 'The integration to report the results and logs to. Supported platforms are `"tensorboard"`' + ' (default), `"wandb"` and `"comet_ml"`. Use `"all"` to report to all integrations.' + ), + ) + parser.add_argument("--local_rank", type=int, default=-1, help="For distributed training: local_rank") + parser.add_argument( + "--checkpointing_steps", + type=int, + default=500, + help=( + "Save a checkpoint of the training state every X updates. These checkpoints are only suitable for resuming" + " training using `--resume_from_checkpoint`." + ), + ) + parser.add_argument( + "--checkpoints_total_limit", + type=int, + default=None, + help=("Max number of checkpoints to store."), + ) + parser.add_argument( + "--resume_from_checkpoint", + type=str, + default=None, + help=( + "Whether training should be resumed from a previous checkpoint. Use a path saved by" + ' `--checkpointing_steps`, or `"latest"` to automatically select the last available checkpoint.' + ), + ) + parser.add_argument("--noise_offset", type=float, default=0, help="The scale of noise offset.") + parser.add_argument( + "--validation_epochs", + type=int, + default=5, + help="Run validation every X epochs.", + ) + parser.add_argument( + "--validation_steps", + type=int, + default=2000, + help="Run validation every X steps.", + ) + parser.add_argument( + "--tracker_project_name", + type=str, + default="text2image-fine-tune", + help=( + "The `project_name` argument passed to Accelerator.init_trackers for" + " more information see https://huggingface.co/docs/accelerate/v0.17.0/en/package_reference/accelerator#accelerate.Accelerator" + ), + ) + + parser.add_argument( + "--snr_loss", action="store_true", help="Whether or not to use snr_loss." + ) + parser.add_argument( + "--uniform_sampling", action="store_true", help="Whether or not to use uniform_sampling." + ) + parser.add_argument( + "--enable_text_encoder_in_dataloader", action="store_true", help="Whether or not to use text encoder in dataloader." + ) + parser.add_argument( + "--enable_bucket", action="store_true", help="Whether enable bucket sample in datasets." + ) + parser.add_argument( + "--random_ratio_crop", action="store_true", help="Whether enable random ratio crop sample in datasets." + ) + parser.add_argument( + "--random_frame_crop", action="store_true", help="Whether enable random frame crop sample in datasets." + ) + parser.add_argument( + "--random_hw_adapt", action="store_true", help="Whether enable random adapt height and width in datasets." + ) + parser.add_argument( + "--training_with_video_token_length", action="store_true", help="The training stage of the model in training.", + ) + parser.add_argument( + "--auto_tile_batch_size", action="store_true", help="Whether to auto tile batch size.", + ) + parser.add_argument( + "--motion_sub_loss", action="store_true", help="Whether enable motion sub loss." + ) + parser.add_argument( + "--motion_sub_loss_ratio", type=float, default=0.25, help="The ratio of motion sub loss." + ) + parser.add_argument( + "--train_sampling_steps", + type=int, + default=1000, + help="Run train_sampling_steps.", + ) + parser.add_argument( + "--keep_all_node_same_token_length", + action="store_true", + help="Reference of the length token.", + ) + parser.add_argument( + "--token_sample_size", + type=int, + default=512, + help="Sample size of the token.", + ) + parser.add_argument( + "--video_sample_size", + type=int, + default=512, + help="Sample size of the video.", + ) + parser.add_argument( + "--image_sample_size", + type=int, + default=512, + help="Sample size of the image.", + ) + parser.add_argument( + "--fix_sample_size", + nargs=2, type=int, default=None, + help="Fix Sample size [height, width] when using bucket and collate_fn." + ) + parser.add_argument( + "--video_sample_stride", + type=int, + default=4, + help="Sample stride of the video.", + ) + parser.add_argument( + "--video_sample_n_frames", + type=int, + default=17, + help="Num frame of video.", + ) + parser.add_argument( + "--video_repeat", + type=int, + default=0, + help="Num of repeat video.", + ) + parser.add_argument( + "--config_path", + type=str, + default=None, + help=( + "The config of the model in training." + ), + ) + parser.add_argument( + "--transformer_path", + type=str, + default=None, + help=("If you want to load the weight from other transformers, input its path."), + ) + parser.add_argument( + "--vae_path", + type=str, + default=None, + help=("If you want to load the weight from other vaes, input its path."), + ) + + parser.add_argument( + '--trainable_modules', + nargs='+', + help='Enter a list of trainable modules' + ) + parser.add_argument( + '--trainable_modules_low_learning_rate', + nargs='+', + default=[], + help='Enter a list of trainable modules with lower learning rate' + ) + parser.add_argument( + '--tokenizer_max_length', + type=int, + default=512, + help='Max length of tokenizer' + ) + parser.add_argument( + "--use_deepspeed", action="store_true", help="Whether or not to use deepspeed." + ) + parser.add_argument( + "--use_fsdp", action="store_true", help="Whether or not to use fsdp." + ) + parser.add_argument( + "--low_vram", action="store_true", help="Whether enable low_vram mode." + ) + parser.add_argument( + "--boundary_type", + type=str, + default="low", + help=( + 'The format of training data. Support `"low"` and `"high"`' + ), + ) + parser.add_argument( + "--train_mode", + type=str, + default="normal", + help=( + 'The format of training data. Support `"normal"`' + ' (default), `"i2v"`.' + ), + ) + parser.add_argument( + "--abnormal_norm_clip_start", + type=int, + default=1000, + help=( + 'When do we start doing additional processing on abnormal gradients. ' + ), + ) + parser.add_argument( + "--initial_grad_norm_ratio", + type=int, + default=5, + help=( + 'The initial gradient is relative to the multiple of the max_grad_norm. ' + ), + ) + parser.add_argument( + "--weighting_scheme", + type=str, + default="none", + choices=["sigma_sqrt", "logit_normal", "mode", "cosmap", "none"], + help=('We default to the "none" weighting scheme for uniform sampling and uniform loss'), + ) + parser.add_argument( + "--logit_mean", type=float, default=0.0, help="mean to use when using the `'logit_normal'` weighting scheme." + ) + parser.add_argument( + "--logit_std", type=float, default=1.0, help="std to use when using the `'logit_normal'` weighting scheme." + ) + parser.add_argument( + "--mode_scale", + type=float, + default=1.29, + help="Scale of mode weighting scheme. Only effective when using the `'mode'` as the `weighting_scheme`.", + ) + + args = parser.parse_args() + env_local_rank = int(os.environ.get("LOCAL_RANK", -1)) + if env_local_rank != -1 and env_local_rank != args.local_rank: + args.local_rank = env_local_rank + + # default to using the same revision for the non-ema model if not specified + if args.non_ema_revision is None: + args.non_ema_revision = args.revision + + return args + + +def main(): + args = parse_args() + + if args.report_to == "wandb" and args.hub_token is not None: + raise ValueError( + "You cannot use both --report_to=wandb and --hub_token due to a security risk of exposing your token." + " Please use `huggingface-cli login` to authenticate with the Hub." + ) + + if args.non_ema_revision is not None: + deprecate( + "non_ema_revision!=None", + "0.15.0", + message=( + "Downloading 'non_ema' weights from revision branches of the Hub is deprecated. Please make sure to" + " use `--variant=non_ema` instead." + ), + ) + logging_dir = os.path.join(args.output_dir, args.logging_dir) + + config = OmegaConf.load(args.config_path) + accelerator_project_config = ProjectConfiguration(project_dir=args.output_dir, logging_dir=logging_dir) + + accelerator = Accelerator( + gradient_accumulation_steps=args.gradient_accumulation_steps, + mixed_precision=args.mixed_precision, + log_with=args.report_to, + project_config=accelerator_project_config, + ) + + deepspeed_plugin = accelerator.state.deepspeed_plugin if hasattr(accelerator.state, "deepspeed_plugin") else None + fsdp_plugin = accelerator.state.fsdp_plugin if hasattr(accelerator.state, "fsdp_plugin") else None + if deepspeed_plugin is not None: + zero_stage = int(deepspeed_plugin.zero_stage) + fsdp_stage = 0 + print(f"Using DeepSpeed Zero stage: {zero_stage}") + + args.use_deepspeed = True + if zero_stage == 3: + print(f"Auto set save_state to True because zero_stage == 3") + args.save_state = True + elif fsdp_plugin is not None: + from torch.distributed.fsdp import ShardingStrategy + zero_stage = 0 + if fsdp_plugin.sharding_strategy is ShardingStrategy.FULL_SHARD: + fsdp_stage = 3 + elif fsdp_plugin.sharding_strategy is None: # The fsdp_plugin.sharding_strategy is None in FSDP 2. + fsdp_stage = 3 + elif fsdp_plugin.sharding_strategy is ShardingStrategy.SHARD_GRAD_OP: + fsdp_stage = 2 + else: + fsdp_stage = 0 + print(f"Using FSDP stage: {fsdp_stage}") + + args.use_fsdp = True + if fsdp_stage == 3: + print(f"Auto set save_state to True because fsdp_stage == 3") + args.save_state = True + else: + zero_stage = 0 + fsdp_stage = 0 + print("DeepSpeed is not enabled.") + + if accelerator.is_main_process: + writer = SummaryWriter(log_dir=logging_dir) + + # Make one log on every process with the configuration for debugging. + logging.basicConfig( + format="%(asctime)s - %(levelname)s - %(name)s - %(message)s", + datefmt="%m/%d/%Y %H:%M:%S", + level=logging.INFO, + ) + logger.info(accelerator.state, main_process_only=False) + if accelerator.is_local_main_process: + datasets.utils.logging.set_verbosity_warning() + transformers.utils.logging.set_verbosity_warning() + diffusers.utils.logging.set_verbosity_info() + else: + datasets.utils.logging.set_verbosity_error() + transformers.utils.logging.set_verbosity_error() + diffusers.utils.logging.set_verbosity_error() + + # If passed along, set the training seed now. + if args.seed is not None: + set_seed(args.seed) + rng = np.random.default_rng(np.random.PCG64(args.seed + accelerator.process_index)) + torch_rng = torch.Generator(accelerator.device).manual_seed(args.seed + accelerator.process_index) + else: + rng = None + torch_rng = None + index_rng = np.random.default_rng(np.random.PCG64(43)) + print(f"Init rng with seed {args.seed + accelerator.process_index}. Process_index is {accelerator.process_index}") + + # Handle the repository creation + if accelerator.is_main_process: + if args.output_dir is not None: + os.makedirs(args.output_dir, exist_ok=True) + + # For mixed precision training we cast all non-trainable weigths (vae, non-lora text_encoder and non-lora transformer3d) to half-precision + # as these weights are only used for inference, keeping weights in full precision is not required. + weight_dtype = torch.float32 + if accelerator.mixed_precision == "fp16": + weight_dtype = torch.float16 + args.mixed_precision = accelerator.mixed_precision + elif accelerator.mixed_precision == "bf16": + weight_dtype = torch.bfloat16 + args.mixed_precision = accelerator.mixed_precision + + # Load scheduler, tokenizer and models. + noise_scheduler = FlowMatchEulerDiscreteScheduler( + **filter_kwargs(FlowMatchEulerDiscreteScheduler, OmegaConf.to_container(config['scheduler_kwargs'])) + ) + + # Get Tokenizer + tokenizer = AutoTokenizer.from_pretrained( + os.path.join(args.pretrained_model_name_or_path, config['text_encoder_kwargs'].get('tokenizer_subpath', 'tokenizer')), + ) + + def deepspeed_zero_init_disabled_context_manager(): + """ + returns either a context list that includes one that will disable zero.Init or an empty context list + """ + deepspeed_plugin = AcceleratorState().deepspeed_plugin if accelerate.state.is_initialized() else None + if deepspeed_plugin is None: + return [] + + return [deepspeed_plugin.zero3_init_context_manager(enable=False)] + + # Currently Accelerate doesn't know how to handle multiple models under Deepspeed ZeRO stage 3. + # For this to work properly all models must be run through `accelerate.prepare`. But accelerate + # will try to assign the same optimizer with the same weights to all models during + # `deepspeed.initialize`, which of course doesn't work. + # + # For now the following workaround will partially support Deepspeed ZeRO-3, by excluding the 2 + # frozen models from being partitioned during `zero.Init` which gets called during + # `from_pretrained` So CLIPTextModel and AutoencoderKL will not enjoy the parameter sharding + # across multiple gpus and only UNet2DConditionModel will get ZeRO sharded. + with ContextManagers(deepspeed_zero_init_disabled_context_manager()): + # Get Text encoder + text_encoder = WanT5EncoderModel.from_pretrained( + os.path.join(args.pretrained_model_name_or_path, config['text_encoder_kwargs'].get('text_encoder_subpath', 'text_encoder')), + additional_kwargs=OmegaConf.to_container(config['text_encoder_kwargs']), + low_cpu_mem_usage=True, + torch_dtype=weight_dtype, + ) + text_encoder = text_encoder.eval() + # Get Vae + vae = AutoencoderKLWan.from_pretrained( + os.path.join(args.pretrained_model_name_or_path, config['vae_kwargs'].get('vae_subpath', 'vae')), + additional_kwargs=OmegaConf.to_container(config['vae_kwargs']), + ) + vae.eval() + + # Get Transformer + sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer') \ + if args.boundary_type == "low" else config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer') + transformer3d = Wan2_2Transformer3DModel.from_pretrained( + os.path.join(args.pretrained_model_name_or_path, sub_path), + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), + ).to(weight_dtype) + + # Freeze vae and text_encoder and set transformer3d to trainable + vae.requires_grad_(False) + text_encoder.requires_grad_(False) + transformer3d.requires_grad_(False) + + if args.transformer_path is not None: + print(f"From checkpoint: {args.transformer_path}") + if args.transformer_path.endswith("safetensors"): + from safetensors.torch import load_file, safe_open + state_dict = load_file(args.transformer_path) + else: + state_dict = torch.load(args.transformer_path, map_location="cpu") + state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict + + m, u = transformer3d.load_state_dict(state_dict, strict=False) + print(f"missing keys: {len(m)}, unexpected keys: {len(u)}") + assert len(u) == 0 + + if args.vae_path is not None: + print(f"From checkpoint: {args.vae_path}") + if args.vae_path.endswith("safetensors"): + from safetensors.torch import load_file, safe_open + state_dict = load_file(args.vae_path) + else: + state_dict = torch.load(args.vae_path, map_location="cpu") + state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict + + m, u = vae.load_state_dict(state_dict, strict=False) + print(f"missing keys: {len(m)}, unexpected keys: {len(u)}") + assert len(u) == 0 + + # A good trainable modules is showed below now. + # For 3D Patch: trainable_modules = ['ff.net', 'pos_embed', 'attn2', 'proj_out', 'timepositionalencoding', 'h_position', 'w_position'] + # For 2D Patch: trainable_modules = ['ff.net', 'attn2', 'timepositionalencoding', 'h_position', 'w_position'] + transformer3d.train() + if accelerator.is_main_process: + accelerator.print( + f"Trainable modules '{args.trainable_modules}'." + ) + for name, param in transformer3d.named_parameters(): + for trainable_module_name in args.trainable_modules + args.trainable_modules_low_learning_rate: + if trainable_module_name in name: + param.requires_grad = True + break + + # Create EMA for the transformer3d. + if args.use_ema: + if zero_stage == 3: + raise NotImplementedError("FSDP does not support EMA.") + + ema_transformer3d = Wan2_2Transformer3DModel.from_pretrained( + os.path.join(args.pretrained_model_name_or_path, config['transformer_additional_kwargs'].get('transformer_subpath', 'transformer')), + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), + ).to(weight_dtype) + + ema_transformer3d = EMAModel(ema_transformer3d.parameters(), model_cls=Wan2_2Transformer3DModel, model_config=ema_transformer3d.config) + + # `accelerate` 0.16.0 will have better support for customized saving + if version.parse(accelerate.__version__) >= version.parse("0.16.0"): + # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format + if fsdp_stage != 0: + def save_model_hook(models, weights, output_dir): + accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) + if accelerator.is_main_process: + from safetensors.torch import save_file + + safetensor_save_path = os.path.join(output_dir, f"diffusion_pytorch_model.safetensors") + accelerate_state_dict = {k: v.to(dtype=weight_dtype) for k, v in accelerate_state_dict.items()} + save_file(accelerate_state_dict, safetensor_save_path, metadata={"format": "pt"}) + + with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: + pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) + + def load_model_hook(models, input_dir): + pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") + if os.path.exists(pkl_path): + with open(pkl_path, 'rb') as file: + loaded_number, _ = pickle.load(file) + batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) + print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") + + elif zero_stage == 3: + # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format + def save_model_hook(models, weights, output_dir): + if accelerator.is_main_process: + with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: + pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) + + def load_model_hook(models, input_dir): + pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") + if os.path.exists(pkl_path): + with open(pkl_path, 'rb') as file: + loaded_number, _ = pickle.load(file) + batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) + print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") + else: + # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format + def save_model_hook(models, weights, output_dir): + if accelerator.is_main_process: + if args.use_ema: + ema_transformer3d.save_pretrained(os.path.join(output_dir, "transformer_ema")) + + models[0].save_pretrained(os.path.join(output_dir, "transformer")) + if not args.use_deepspeed: + weights.pop() + + with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: + pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) + + def load_model_hook(models, input_dir): + if args.use_ema: + ema_path = os.path.join(input_dir, "transformer_ema") + _, ema_kwargs = Wan2_2Transformer3DModel.load_config(ema_path, return_unused_kwargs=True) + load_model = Wan2_2Transformer3DModel.from_pretrained( + input_dir, subfolder="transformer_ema", + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']) + ) + load_model = EMAModel(load_model.parameters(), model_cls=Wan2_2Transformer3DModel, model_config=load_model.config) + load_model.load_state_dict(ema_kwargs) + + ema_transformer3d.load_state_dict(load_model.state_dict()) + ema_transformer3d.to(accelerator.device) + del load_model + + for i in range(len(models)): + # pop models so that they are not loaded again + model = models.pop() + + # load diffusers style into model + load_model = Wan2_2Transformer3DModel.from_pretrained( + input_dir, subfolder="transformer" + ) + model.register_to_config(**load_model.config) + + model.load_state_dict(load_model.state_dict()) + del load_model + + pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") + if os.path.exists(pkl_path): + with open(pkl_path, 'rb') as file: + loaded_number, _ = pickle.load(file) + batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) + print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") + + accelerator.register_save_state_pre_hook(save_model_hook) + accelerator.register_load_state_pre_hook(load_model_hook) + + if args.gradient_checkpointing: + transformer3d.enable_gradient_checkpointing() + + # Enable TF32 for faster training on Ampere GPUs, + # cf https://pytorch.org/docs/stable/notes/cuda.html#tensorfloat-32-tf32-on-ampere-devices + if args.allow_tf32: + torch.backends.cuda.matmul.allow_tf32 = True + + if args.scale_lr: + args.learning_rate = ( + args.learning_rate * args.gradient_accumulation_steps * args.train_batch_size * accelerator.num_processes + ) + + # Initialize the optimizer + if args.use_8bit_adam: + try: + import bitsandbytes as bnb + except ImportError: + raise ImportError( + "Please install bitsandbytes to use 8-bit Adam. You can do so by running `pip install bitsandbytes`" + ) + + optimizer_cls = bnb.optim.AdamW8bit + elif args.use_came: + try: + from came_pytorch import CAME + except: + raise ImportError( + "Please install came_pytorch to use CAME. You can do so by running `pip install came_pytorch`" + ) + + optimizer_cls = CAME + else: + optimizer_cls = torch.optim.AdamW + + trainable_params = list(filter(lambda p: p.requires_grad, transformer3d.parameters())) + trainable_params_optim = [ + {'params': [], 'lr': args.learning_rate}, + {'params': [], 'lr': args.learning_rate / 2}, + ] + in_already = [] + for name, param in transformer3d.named_parameters(): + high_lr_flag = False + if name in in_already: + continue + for trainable_module_name in args.trainable_modules: + if trainable_module_name in name: + in_already.append(name) + high_lr_flag = True + trainable_params_optim[0]['params'].append(param) + if accelerator.is_main_process: + print(f"Set {name} to lr : {args.learning_rate}") + break + if high_lr_flag: + continue + for trainable_module_name in args.trainable_modules_low_learning_rate: + if trainable_module_name in name: + in_already.append(name) + trainable_params_optim[1]['params'].append(param) + if accelerator.is_main_process: + print(f"Set {name} to lr : {args.learning_rate / 2}") + break + + if args.use_came: + optimizer = optimizer_cls( + trainable_params_optim, + lr=args.learning_rate, + # weight_decay=args.adam_weight_decay, + betas=(0.9, 0.999, 0.9999), + eps=(1e-30, 1e-16) + ) + else: + optimizer = optimizer_cls( + trainable_params_optim, + lr=args.learning_rate, + betas=(args.adam_beta1, args.adam_beta2), + weight_decay=args.adam_weight_decay, + eps=args.adam_epsilon, + ) + + # Get the training dataset + sample_n_frames_bucket_interval = vae.config.temporal_compression_ratio + + if args.fix_sample_size is not None and args.enable_bucket: + args.video_sample_size = max(max(args.fix_sample_size), args.video_sample_size) + args.image_sample_size = max(max(args.fix_sample_size), args.image_sample_size) + args.training_with_video_token_length = False + args.random_hw_adapt = False + + # Get the dataset + train_dataset = ImageVideoDataset( + args.train_data_meta, args.train_data_dir, + video_sample_size=args.video_sample_size, video_sample_stride=args.video_sample_stride, video_sample_n_frames=args.video_sample_n_frames, + video_repeat=args.video_repeat, + image_sample_size=args.image_sample_size, + enable_bucket=args.enable_bucket, enable_inpaint=True if args.train_mode != "normal" else False, + ) + + def worker_init_fn(_seed): + _seed = _seed * 256 + def _worker_init_fn(worker_id): + print(f"worker_init_fn with {_seed + worker_id}") + np.random.seed(_seed + worker_id) + random.seed(_seed + worker_id) + return _worker_init_fn + + if args.enable_bucket: + aspect_ratio_sample_size = {key : [x / 512 * args.video_sample_size for x in ASPECT_RATIO_512[key]] for key in ASPECT_RATIO_512.keys()} + batch_sampler_generator = torch.Generator().manual_seed(args.seed) + batch_sampler = AspectRatioBatchImageVideoSampler( + sampler=RandomSampler(train_dataset, generator=batch_sampler_generator), dataset=train_dataset.dataset, + batch_size=args.train_batch_size, train_folder = args.train_data_dir, drop_last=True, + aspect_ratios=aspect_ratio_sample_size, + ) + + def collate_fn(examples): + def get_length_to_frame_num(token_length): + if args.image_sample_size > args.video_sample_size: + sample_sizes = list(range(args.video_sample_size, args.image_sample_size + 1, 128)) + + if sample_sizes[-1] != args.image_sample_size: + sample_sizes.append(args.image_sample_size) + else: + sample_sizes = [args.image_sample_size] + + length_to_frame_num = { + sample_size: min(token_length / sample_size / sample_size, args.video_sample_n_frames) // sample_n_frames_bucket_interval * sample_n_frames_bucket_interval + 1 for sample_size in sample_sizes + } + + return length_to_frame_num + + def get_random_downsample_ratio(sample_size, image_ratio=[], + all_choices=False, rng=None): + def _create_special_list(length): + if length == 1: + return [1.0] + if length >= 2: + first_element = 0.90 + remaining_sum = 1.0 - first_element + other_elements_value = remaining_sum / (length - 1) + special_list = [first_element] + [other_elements_value] * (length - 1) + return special_list + + if sample_size >= 1536: + number_list = [1, 1.25, 1.5, 2, 2.5, 3] + image_ratio + elif sample_size >= 1024: + number_list = [1, 1.25, 1.5, 2] + image_ratio + elif sample_size >= 768: + number_list = [1, 1.25, 1.5] + image_ratio + elif sample_size >= 512: + number_list = [1] + image_ratio + else: + number_list = [1] + + if all_choices: + return number_list + + number_list_prob = np.array(_create_special_list(len(number_list))) + if rng is None: + return np.random.choice(number_list, p = number_list_prob) + else: + return rng.choice(number_list, p = number_list_prob) + + # Get token length + target_token_length = args.video_sample_n_frames * args.token_sample_size * args.token_sample_size + length_to_frame_num = get_length_to_frame_num(target_token_length) + + # Create new output + new_examples = {} + new_examples["target_token_length"] = target_token_length + new_examples["pixel_values"] = [] + new_examples["text"] = [] + # Used in Inpaint mode + if args.train_mode != "normal": + new_examples["mask_pixel_values"] = [] + new_examples["mask"] = [] + new_examples["clip_pixel_values"] = [] + + # Get downsample ratio in image and videos + pixel_value = examples[0]["pixel_values"] + data_type = examples[0]["data_type"] + f, h, w, c = np.shape(pixel_value) + if data_type == 'image': + random_downsample_ratio = 1 if not args.random_hw_adapt else get_random_downsample_ratio(args.image_sample_size, image_ratio=[args.image_sample_size / args.video_sample_size]) + + aspect_ratio_sample_size = {key : [x / 512 * args.image_sample_size / random_downsample_ratio for x in ASPECT_RATIO_512[key]] for key in ASPECT_RATIO_512.keys()} + aspect_ratio_random_crop_sample_size = {key : [x / 512 * args.image_sample_size / random_downsample_ratio for x in ASPECT_RATIO_RANDOM_CROP_512[key]] for key in ASPECT_RATIO_RANDOM_CROP_512.keys()} + + batch_video_length = args.video_sample_n_frames + sample_n_frames_bucket_interval + else: + if args.random_hw_adapt: + if args.training_with_video_token_length: + local_min_size = np.min(np.array([np.mean(np.array([np.shape(example["pixel_values"])[1], np.shape(example["pixel_values"])[2]])) for example in examples])) + # The video will be resized to a lower resolution than its own. + choice_list = [length for length in list(length_to_frame_num.keys()) if length < local_min_size * 1.25] + if len(choice_list) == 0: + choice_list = list(length_to_frame_num.keys()) + local_video_sample_size = np.random.choice(choice_list) + batch_video_length = length_to_frame_num[local_video_sample_size] + random_downsample_ratio = args.video_sample_size / local_video_sample_size + else: + random_downsample_ratio = get_random_downsample_ratio(args.video_sample_size) + batch_video_length = args.video_sample_n_frames + sample_n_frames_bucket_interval + else: + random_downsample_ratio = 1 + batch_video_length = args.video_sample_n_frames + sample_n_frames_bucket_interval + + aspect_ratio_sample_size = {key : [x / 512 * args.video_sample_size / random_downsample_ratio for x in ASPECT_RATIO_512[key]] for key in ASPECT_RATIO_512.keys()} + aspect_ratio_random_crop_sample_size = {key : [x / 512 * args.video_sample_size / random_downsample_ratio for x in ASPECT_RATIO_RANDOM_CROP_512[key]] for key in ASPECT_RATIO_RANDOM_CROP_512.keys()} + + if args.fix_sample_size is not None: + fix_sample_size = [int(x / 16) * 16 for x in args.fix_sample_size] + elif args.random_ratio_crop: + if rng is None: + random_sample_size = aspect_ratio_random_crop_sample_size[ + np.random.choice(list(aspect_ratio_random_crop_sample_size.keys()), p = ASPECT_RATIO_RANDOM_CROP_PROB) + ] + else: + random_sample_size = aspect_ratio_random_crop_sample_size[ + rng.choice(list(aspect_ratio_random_crop_sample_size.keys()), p = ASPECT_RATIO_RANDOM_CROP_PROB) + ] + random_sample_size = [int(x / 16) * 16 for x in random_sample_size] + else: + closest_size, closest_ratio = get_closest_ratio(h, w, ratios=aspect_ratio_sample_size) + closest_size = [int(x / 16) * 16 for x in closest_size] + + for example in examples: + if args.fix_sample_size is not None: + # To 0~1 + pixel_values = torch.from_numpy(example["pixel_values"]).permute(0, 3, 1, 2).contiguous() + pixel_values = pixel_values / 255. + + # Get adapt hw for resize + fix_sample_size = list(map(lambda x: int(x), fix_sample_size)) + transform = transforms.Compose([ + transforms.Resize(fix_sample_size, interpolation=transforms.InterpolationMode.BILINEAR), # Image.BICUBIC + transforms.CenterCrop(fix_sample_size), + transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True), + ]) + elif args.random_ratio_crop: + # To 0~1 + pixel_values = torch.from_numpy(example["pixel_values"]).permute(0, 3, 1, 2).contiguous() + pixel_values = pixel_values / 255. + + # Get adapt hw for resize + b, c, h, w = pixel_values.size() + th, tw = random_sample_size + if th / tw > h / w: + nh = int(th) + nw = int(w / h * nh) + else: + nw = int(tw) + nh = int(h / w * nw) + + transform = transforms.Compose([ + transforms.Resize([nh, nw]), + transforms.CenterCrop([int(x) for x in random_sample_size]), + transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True), + ]) + else: + # To 0~1 + pixel_values = torch.from_numpy(example["pixel_values"]).permute(0, 3, 1, 2).contiguous() + pixel_values = pixel_values / 255. + + # Get adapt hw for resize + closest_size = list(map(lambda x: int(x), closest_size)) + if closest_size[0] / h > closest_size[1] / w: + resize_size = closest_size[0], int(w * closest_size[0] / h) + else: + resize_size = int(h * closest_size[1] / w), closest_size[1] + + transform = transforms.Compose([ + transforms.Resize(resize_size, interpolation=transforms.InterpolationMode.BILINEAR), # Image.BICUBIC + transforms.CenterCrop(closest_size), + transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True), + ]) + new_examples["pixel_values"].append(transform(pixel_values)) + new_examples["text"].append(example["text"]) + + batch_video_length = int(min(batch_video_length, len(pixel_values))) + + # Magvae needs the number of frames to be 4n + 1. + batch_video_length = (batch_video_length - 1) // sample_n_frames_bucket_interval * sample_n_frames_bucket_interval + 1 + + if batch_video_length <= 0: + batch_video_length = 1 + + if args.train_mode != "normal": + mask = get_random_mask(new_examples["pixel_values"][-1].size()) + mask_pixel_values = new_examples["pixel_values"][-1] * (1 - mask) + # Wan 2.1 use 0 for masked pixels + # + torch.ones_like(new_examples["pixel_values"][-1]) * -1 * mask + new_examples["mask_pixel_values"].append(mask_pixel_values) + new_examples["mask"].append(mask) + + clip_pixel_values = new_examples["pixel_values"][-1][0].permute(1, 2, 0).contiguous() + clip_pixel_values = (clip_pixel_values * 0.5 + 0.5) * 255 + new_examples["clip_pixel_values"].append(clip_pixel_values) + + # Limit the number of frames to the same + new_examples["pixel_values"] = torch.stack([example[:batch_video_length] for example in new_examples["pixel_values"]]) + if args.train_mode != "normal": + new_examples["mask_pixel_values"] = torch.stack([example[:batch_video_length] for example in new_examples["mask_pixel_values"]]) + new_examples["mask"] = torch.stack([example[:batch_video_length] for example in new_examples["mask"]]) + new_examples["clip_pixel_values"] = torch.stack([example for example in new_examples["clip_pixel_values"]]) + + # Encode prompts when enable_text_encoder_in_dataloader=True + if args.enable_text_encoder_in_dataloader: + prompt_ids = tokenizer( + new_examples['text'], + max_length=args.tokenizer_max_length, + padding="max_length", + add_special_tokens=True, + truncation=True, + return_tensors="pt" + ) + encoder_hidden_states = text_encoder( + prompt_ids.input_ids + )[0] + new_examples['encoder_attention_mask'] = prompt_ids.attention_mask + new_examples['encoder_hidden_states'] = encoder_hidden_states + + return new_examples + + # DataLoaders creation: + train_dataloader = torch.utils.data.DataLoader( + train_dataset, + batch_sampler=batch_sampler, + collate_fn=collate_fn, + persistent_workers=True if args.dataloader_num_workers != 0 else False, + num_workers=args.dataloader_num_workers, + worker_init_fn=worker_init_fn(args.seed + accelerator.process_index) + ) + else: + # DataLoaders creation: + batch_sampler_generator = torch.Generator().manual_seed(args.seed) + batch_sampler = ImageVideoSampler(RandomSampler(train_dataset, generator=batch_sampler_generator), train_dataset, args.train_batch_size) + train_dataloader = torch.utils.data.DataLoader( + train_dataset, + batch_sampler=batch_sampler, + persistent_workers=True if args.dataloader_num_workers != 0 else False, + num_workers=args.dataloader_num_workers, + worker_init_fn=worker_init_fn(args.seed + accelerator.process_index) + ) + + # Scheduler and math around the number of training steps. + overrode_max_train_steps = False + num_update_steps_per_epoch = math.ceil(len(train_dataloader) / args.gradient_accumulation_steps) + if args.max_train_steps is None: + args.max_train_steps = args.num_train_epochs * num_update_steps_per_epoch + overrode_max_train_steps = True + + lr_scheduler = get_scheduler( + args.lr_scheduler, + optimizer=optimizer, + num_warmup_steps=args.lr_warmup_steps * accelerator.num_processes, + num_training_steps=args.max_train_steps * accelerator.num_processes, + ) + + # Prepare everything with our `accelerator`. + transformer3d, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( + transformer3d, optimizer, train_dataloader, lr_scheduler + ) + + if fsdp_stage != 0: + from functools import partial + from videox_fun.dist import set_multi_gpus_devices, shard_model + shard_fn = partial(shard_model, device_id=accelerator.device, param_dtype=weight_dtype) + text_encoder = shard_fn(text_encoder) + + if args.use_ema: + ema_transformer3d.to(accelerator.device) + + # Move text_encode and vae to gpu and cast to weight_dtype + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu") + + # We need to recalculate our total training steps as the size of the training dataloader may have changed. + num_update_steps_per_epoch = math.ceil(len(train_dataloader) / args.gradient_accumulation_steps) + if overrode_max_train_steps: + args.max_train_steps = args.num_train_epochs * num_update_steps_per_epoch + # Afterwards we recalculate our number of training epochs + args.num_train_epochs = math.ceil(args.max_train_steps / num_update_steps_per_epoch) + + # We need to initialize the trackers we use, and also store our configuration. + # The trackers initializes automatically on the main process. + if accelerator.is_main_process: + tracker_config = dict(vars(args)) + tracker_config.pop("validation_prompts") + tracker_config.pop("trainable_modules") + tracker_config.pop("trainable_modules_low_learning_rate") + tracker_config.pop("fix_sample_size") + accelerator.init_trackers(args.tracker_project_name, tracker_config) + + # Function for unwrapping if model was compiled with `torch.compile`. + def unwrap_model(model): + model = accelerator.unwrap_model(model) + model = model._orig_mod if is_compiled_module(model) else model + return model + + # Train! + total_batch_size = args.train_batch_size * accelerator.num_processes * args.gradient_accumulation_steps + + logger.info("***** Running training *****") + logger.info(f" Num examples = {len(train_dataset)}") + logger.info(f" Num Epochs = {args.num_train_epochs}") + logger.info(f" Instantaneous batch size per device = {args.train_batch_size}") + logger.info(f" Total train batch size (w. parallel, distributed & accumulation) = {total_batch_size}") + logger.info(f" Gradient Accumulation steps = {args.gradient_accumulation_steps}") + logger.info(f" Total optimization steps = {args.max_train_steps}") + global_step = 0 + first_epoch = 0 + + # Potentially load in the weights and states from a previous save + if args.resume_from_checkpoint: + if args.resume_from_checkpoint != "latest": + path = os.path.basename(args.resume_from_checkpoint) + else: + # Get the most recent checkpoint + dirs = os.listdir(args.output_dir) + dirs = [d for d in dirs if d.startswith("checkpoint")] + dirs = sorted(dirs, key=lambda x: int(x.split("-")[1])) + path = dirs[-1] if len(dirs) > 0 else None + + if path is None: + accelerator.print( + f"Checkpoint '{args.resume_from_checkpoint}' does not exist. Starting a new training run." + ) + args.resume_from_checkpoint = None + initial_global_step = 0 + else: + global_step = int(path.split("-")[1]) + + initial_global_step = global_step + + pkl_path = os.path.join(os.path.join(args.output_dir, path), "sampler_pos_start.pkl") + if os.path.exists(pkl_path): + with open(pkl_path, 'rb') as file: + _, first_epoch = pickle.load(file) + else: + first_epoch = global_step // num_update_steps_per_epoch + print(f"Load pkl from {pkl_path}. Get first_epoch = {first_epoch}.") + + accelerator.print(f"Resuming from checkpoint {path}") + accelerator.load_state(os.path.join(args.output_dir, path)) + else: + initial_global_step = 0 + + progress_bar = tqdm( + range(0, args.max_train_steps), + initial=initial_global_step, + desc="Steps", + # Only show the progress bar once on each machine. + disable=not accelerator.is_local_main_process, + ) + + if args.multi_stream and args.train_mode != "normal": + # create extra cuda streams to speedup inpaint vae computation + vae_stream_1 = torch.cuda.Stream() + vae_stream_2 = torch.cuda.Stream() + else: + vae_stream_1 = None + vae_stream_2 = None + + # Calculate the index we need + boundary = config['transformer_additional_kwargs'].get('boundary', 0.900) + split_timesteps = args.train_sampling_steps * boundary + differences = torch.abs(noise_scheduler.timesteps - split_timesteps) + closest_index = torch.argmin(differences).item() + print(f"The boundary is {boundary} and the boundary_type is {args.boundary_type}. The closest_index we calculate is {closest_index}") + if args.boundary_type == "high": + start_num_idx = 0 + train_sampling_steps = closest_index + elif args.boundary_type == "low": + start_num_idx = closest_index + train_sampling_steps = args.train_sampling_steps - closest_index + else: + start_num_idx = 0 + train_sampling_steps = args.train_sampling_steps + idx_sampling = DiscreteSampling(train_sampling_steps, start_num_idx=start_num_idx, uniform_sampling=args.uniform_sampling) + + for epoch in range(first_epoch, args.num_train_epochs): + train_loss = 0.0 + batch_sampler.sampler.generator = torch.Generator().manual_seed(args.seed + epoch) + for step, batch in enumerate(train_dataloader): + # Data batch sanity check + if epoch == first_epoch and step == 0: + pixel_values, texts = batch['pixel_values'].cpu(), batch['text'] + pixel_values = rearrange(pixel_values, "b f c h w -> b c f h w") + os.makedirs(os.path.join(args.output_dir, "sanity_check"), exist_ok=True) + for idx, (pixel_value, text) in enumerate(zip(pixel_values, texts)): + pixel_value = pixel_value[None, ...] + gif_name = '-'.join(text.replace('/', '').split()[:10]) if not text == '' else f'{global_step}-{idx}' + save_videos_grid(pixel_value, f"{args.output_dir}/sanity_check/{gif_name[:10]}.gif", rescale=True) + if args.train_mode != "normal": + clip_pixel_values, mask_pixel_values, texts = batch['clip_pixel_values'].cpu(), batch['mask_pixel_values'].cpu(), batch['text'] + mask_pixel_values = rearrange(mask_pixel_values, "b f c h w -> b c f h w") + for idx, (clip_pixel_value, pixel_value, text) in enumerate(zip(clip_pixel_values, mask_pixel_values, texts)): + pixel_value = pixel_value[None, ...] + Image.fromarray(np.uint8(clip_pixel_value)).save(f"{args.output_dir}/sanity_check/clip_{gif_name[:10] if not text == '' else f'{global_step}-{idx}'}.png") + save_videos_grid(pixel_value, f"{args.output_dir}/sanity_check/mask_{gif_name[:10] if not text == '' else f'{global_step}-{idx}'}.gif", rescale=True) + + with accelerator.accumulate(transformer3d): + # Convert images to latent space + pixel_values = batch["pixel_values"].to(weight_dtype) + + # Increase the batch size when the length of the latent sequence of the current sample is small + if args.auto_tile_batch_size and args.training_with_video_token_length and zero_stage != 3: + if args.video_sample_n_frames * args.token_sample_size * args.token_sample_size // 16 >= pixel_values.size()[1] * pixel_values.size()[3] * pixel_values.size()[4]: + pixel_values = torch.tile(pixel_values, (4, 1, 1, 1, 1)) + if args.enable_text_encoder_in_dataloader: + batch['encoder_hidden_states'] = torch.tile(batch['encoder_hidden_states'], (4, 1, 1)) + batch['encoder_attention_mask'] = torch.tile(batch['encoder_attention_mask'], (4, 1)) + else: + batch['text'] = batch['text'] * 4 + elif args.video_sample_n_frames * args.token_sample_size * args.token_sample_size // 4 >= pixel_values.size()[1] * pixel_values.size()[3] * pixel_values.size()[4]: + pixel_values = torch.tile(pixel_values, (2, 1, 1, 1, 1)) + if args.enable_text_encoder_in_dataloader: + batch['encoder_hidden_states'] = torch.tile(batch['encoder_hidden_states'], (2, 1, 1)) + batch['encoder_attention_mask'] = torch.tile(batch['encoder_attention_mask'], (2, 1)) + else: + batch['text'] = batch['text'] * 2 + + if args.train_mode != "normal": + mask_pixel_values = batch["mask_pixel_values"].to(weight_dtype) + mask = batch["mask"].to(weight_dtype) + # Increase the batch size when the length of the latent sequence of the current sample is small + if args.auto_tile_batch_size and args.training_with_video_token_length and zero_stage != 3: + if args.video_sample_n_frames * args.token_sample_size * args.token_sample_size // 16 >= pixel_values.size()[1] * pixel_values.size()[3] * pixel_values.size()[4]: + mask_pixel_values = torch.tile(mask_pixel_values, (4, 1, 1, 1, 1)) + mask = torch.tile(mask, (4, 1, 1, 1, 1)) + elif args.video_sample_n_frames * args.token_sample_size * args.token_sample_size // 4 >= pixel_values.size()[1] * pixel_values.size()[3] * pixel_values.size()[4]: + mask_pixel_values = torch.tile(mask_pixel_values, (2, 1, 1, 1, 1)) + mask = torch.tile(mask, (2, 1, 1, 1, 1)) + + if args.random_frame_crop: + def _create_special_list(length): + if length == 1: + return [1.0] + if length >= 2: + last_element = 0.90 + remaining_sum = 1.0 - last_element + other_elements_value = remaining_sum / (length - 1) + special_list = [other_elements_value] * (length - 1) + [last_element] + return special_list + select_frames = [_tmp for _tmp in list(range(sample_n_frames_bucket_interval + 1, args.video_sample_n_frames + sample_n_frames_bucket_interval, sample_n_frames_bucket_interval))] + select_frames_prob = np.array(_create_special_list(len(select_frames))) + + if len(select_frames) != 0: + if rng is None: + temp_n_frames = np.random.choice(select_frames, p = select_frames_prob) + else: + temp_n_frames = rng.choice(select_frames, p = select_frames_prob) + else: + temp_n_frames = 1 + + # Magvae needs the number of frames to be 4n + 1. + temp_n_frames = (temp_n_frames - 1) // sample_n_frames_bucket_interval + 1 + + pixel_values = pixel_values[:, :temp_n_frames, :, :] + + if args.train_mode != "normal": + mask_pixel_values = mask_pixel_values[:, :temp_n_frames, :, :] + mask = mask[:, :temp_n_frames, :, :] + + # Keep all node same token length to accelerate the traning when resolution grows. + if args.keep_all_node_same_token_length: + if args.token_sample_size > 256: + numbers_list = list(range(256, args.token_sample_size + 1, 128)) + + if numbers_list[-1] != args.token_sample_size: + numbers_list.append(args.token_sample_size) + else: + numbers_list = [256] + numbers_list = [_number * _number * args.video_sample_n_frames for _number in numbers_list] + + actual_token_length = index_rng.choice(numbers_list) + actual_video_length = (min( + actual_token_length / pixel_values.size()[-1] / pixel_values.size()[-2], args.video_sample_n_frames + ) - 1) // sample_n_frames_bucket_interval * sample_n_frames_bucket_interval + 1 + actual_video_length = int(max(actual_video_length, 1)) + + # Magvae needs the number of frames to be 4n + 1. + actual_video_length = (actual_video_length - 1) // sample_n_frames_bucket_interval + 1 + + pixel_values = pixel_values[:, :actual_video_length, :, :] + if args.train_mode != "normal": + mask_pixel_values = mask_pixel_values[:, :actual_video_length, :, :] + mask = mask[:, :actual_video_length, :, :] + + # Make the inpaint latents to be zeros. + if args.train_mode != "normal": + t2v_flag = [(_mask == 1).all() for _mask in mask] + new_t2v_flag = [] + for _mask in t2v_flag: + if _mask and np.random.rand() < 0.90: + new_t2v_flag.append(0) + else: + new_t2v_flag.append(1) + t2v_flag = torch.from_numpy(np.array(new_t2v_flag)).to(accelerator.device, dtype=weight_dtype) + + if args.low_vram: + torch.cuda.empty_cache() + vae.to(accelerator.device) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to("cpu") + + with torch.no_grad(): + # This way is quicker when batch grows up + def _batch_encode_vae(pixel_values): + pixel_values = rearrange(pixel_values, "b f c h w -> b c f h w") + bs = args.vae_mini_batch + new_pixel_values = [] + for i in range(0, pixel_values.shape[0], bs): + pixel_values_bs = pixel_values[i : i + bs] + pixel_values_bs = vae.encode(pixel_values_bs)[0] + pixel_values_bs = pixel_values_bs.sample() + new_pixel_values.append(pixel_values_bs) + return torch.cat(new_pixel_values, dim = 0) + if vae_stream_1 is not None: + vae_stream_1.wait_stream(torch.cuda.current_stream()) + with torch.cuda.stream(vae_stream_1): + latents = _batch_encode_vae(pixel_values) + else: + latents = _batch_encode_vae(pixel_values) + + if args.train_mode != "normal": + mask = rearrange(mask, "b f c h w -> b c f h w") + mask = torch.concat( + [ + torch.repeat_interleave(mask[:, :, 0:1], repeats=4, dim=2), + mask[:, :, 1:] + ], dim=2 + ) + mask = mask.view(mask.shape[0], mask.shape[2] // 4, 4, mask.shape[3], mask.shape[4]) + mask = mask.transpose(1, 2) + mask = resize_mask(1 - mask, latents) + + # Encode inpaint latents. + mask_latents = _batch_encode_vae(mask_pixel_values) + if vae_stream_2 is not None: + torch.cuda.current_stream().wait_stream(vae_stream_2) + + inpaint_latents = torch.concat([mask, mask_latents], dim=1) + inpaint_latents = t2v_flag[:, None, None, None, None] * inpaint_latents + + # wait for latents = vae.encode(pixel_values) to complete + if vae_stream_1 is not None: + torch.cuda.current_stream().wait_stream(vae_stream_1) + + if args.low_vram: + vae.to('cpu') + torch.cuda.empty_cache() + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device) + + if args.enable_text_encoder_in_dataloader: + prompt_embeds = batch['encoder_hidden_states'].to(device=latents.device) + else: + with torch.no_grad(): + prompt_ids = tokenizer( + batch['text'], + padding="max_length", + max_length=args.tokenizer_max_length, + truncation=True, + add_special_tokens=True, + return_tensors="pt" + ) + text_input_ids = prompt_ids.input_ids + prompt_attention_mask = prompt_ids.attention_mask + + seq_lens = prompt_attention_mask.gt(0).sum(dim=1).long() + prompt_embeds = text_encoder(text_input_ids.to(latents.device), attention_mask=prompt_attention_mask.to(latents.device))[0] + prompt_embeds = [u[:v] for u, v in zip(prompt_embeds, seq_lens)] + + if args.low_vram and not args.enable_text_encoder_in_dataloader: + text_encoder.to('cpu') + torch.cuda.empty_cache() + + bsz, channel, num_frames, height, width = latents.size() + noise = torch.randn(latents.size(), device=latents.device, generator=torch_rng, dtype=weight_dtype) + + if not args.uniform_sampling: + u = compute_density_for_timestep_sampling( + weighting_scheme=args.weighting_scheme, + batch_size=bsz, + logit_mean=args.logit_mean, + logit_std=args.logit_std, + mode_scale=args.mode_scale, + ) + indices = (u * noise_scheduler.config.num_train_timesteps).long() + else: + # Sample a random timestep for each image + # timesteps = generate_timestep_with_lognorm(0, args.train_sampling_steps, (bsz,), device=latents.device, generator=torch_rng) + # timesteps = torch.randint(0, args.train_sampling_steps, (bsz,), device=latents.device, generator=torch_rng) + indices = idx_sampling(bsz, generator=torch_rng, device=latents.device) + indices = indices.long().cpu() + timesteps = noise_scheduler.timesteps[indices].to(device=latents.device) + + def get_sigmas(timesteps, n_dim=4, dtype=torch.float32): + sigmas = noise_scheduler.sigmas.to(device=accelerator.device, dtype=dtype) + schedule_timesteps = noise_scheduler.timesteps.to(accelerator.device) + timesteps = timesteps.to(accelerator.device) + step_indices = [(schedule_timesteps == t).nonzero().item() for t in timesteps] + + sigma = sigmas[step_indices].flatten() + while len(sigma.shape) < n_dim: + sigma = sigma.unsqueeze(-1) + return sigma + + # Add noise according to flow matching. + # zt = (1 - texp) * x + texp * z1 + sigmas = get_sigmas(timesteps, n_dim=latents.ndim, dtype=latents.dtype) + noisy_latents = (1.0 - sigmas) * latents + sigmas * noise + + # Add noise + target = noise - latents + + target_shape = (vae.latent_channels, num_frames, width, height) + seq_len = math.ceil( + (target_shape[2] * target_shape[3]) / + (accelerator.unwrap_model(transformer3d).config.patch_size[1] * accelerator.unwrap_model(transformer3d).config.patch_size[2]) * + target_shape[1] + ) + + # Predict the noise residual + with torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=accelerator.device): + noise_pred = transformer3d( + x=noisy_latents, + context=prompt_embeds, + t=timesteps, + seq_len=seq_len, + y=inpaint_latents if args.train_mode != "normal" else None, + ) + + def custom_mse_loss(noise_pred, target, weighting=None, threshold=50): + noise_pred = noise_pred.float() + target = target.float() + diff = noise_pred - target + mse_loss = F.mse_loss(noise_pred, target, reduction='none') + mask = (diff.abs() <= threshold).float() + masked_loss = mse_loss * mask + if weighting is not None: + masked_loss = masked_loss * weighting + final_loss = masked_loss.mean() + return final_loss + + weighting = compute_loss_weighting_for_sd3(weighting_scheme=args.weighting_scheme, sigmas=sigmas) + loss = custom_mse_loss(noise_pred.float(), target.float(), weighting.float()) + loss = loss.mean() + + if args.motion_sub_loss and noise_pred.size()[1] > 2: + gt_sub_noise = noise_pred[:, 1:, :].float() - noise_pred[:, :-1, :].float() + pre_sub_noise = target[:, 1:, :].float() - target[:, :-1, :].float() + sub_loss = F.mse_loss(gt_sub_noise, pre_sub_noise, reduction="mean") + loss = loss * (1 - args.motion_sub_loss_ratio) + sub_loss * args.motion_sub_loss_ratio + + # Gather the losses across all processes for logging (if we use distributed training). + avg_loss = accelerator.gather(loss.repeat(args.train_batch_size)).mean() + train_loss += avg_loss.item() / args.gradient_accumulation_steps + + # Backpropagate + accelerator.backward(loss) + if accelerator.sync_gradients: + if not args.use_deepspeed and not args.use_fsdp: + trainable_params_grads = [p.grad for p in trainable_params if p.grad is not None] + trainable_params_total_norm = torch.norm(torch.stack([torch.norm(g.detach(), 2) for g in trainable_params_grads]), 2) + max_grad_norm = linear_decay(args.max_grad_norm * args.initial_grad_norm_ratio, args.max_grad_norm, args.abnormal_norm_clip_start, global_step) + if trainable_params_total_norm / max_grad_norm > 5 and global_step > args.abnormal_norm_clip_start: + actual_max_grad_norm = max_grad_norm / min((trainable_params_total_norm / max_grad_norm), 10) + else: + actual_max_grad_norm = max_grad_norm + else: + actual_max_grad_norm = args.max_grad_norm + + if not args.use_deepspeed and not args.use_fsdp and args.report_model_info and accelerator.is_main_process: + if trainable_params_total_norm > 1 and global_step > args.abnormal_norm_clip_start: + for name, param in transformer3d.named_parameters(): + if param.requires_grad: + writer.add_scalar(f'gradients/before_clip_norm/{name}', param.grad.norm(), global_step=global_step) + + norm_sum = accelerator.clip_grad_norm_(trainable_params, actual_max_grad_norm) + if not args.use_deepspeed and not args.use_fsdp and args.report_model_info and accelerator.is_main_process: + writer.add_scalar(f'gradients/norm_sum', norm_sum, global_step=global_step) + writer.add_scalar(f'gradients/actual_max_grad_norm', actual_max_grad_norm, global_step=global_step) + optimizer.step() + lr_scheduler.step() + optimizer.zero_grad() + + # Checks if the accelerator has performed an optimization step behind the scenes + if accelerator.sync_gradients: + + if args.use_ema: + ema_transformer3d.step(transformer3d.parameters()) + progress_bar.update(1) + global_step += 1 + accelerator.log({"train_loss": train_loss}, step=global_step) + train_loss = 0.0 + + if global_step % args.checkpointing_steps == 0: + if args.use_deepspeed or args.use_fsdp or accelerator.is_main_process: + # _before_ saving state, check if this save would set us over the `checkpoints_total_limit` + if args.checkpoints_total_limit is not None: + checkpoints = os.listdir(args.output_dir) + checkpoints = [d for d in checkpoints if d.startswith("checkpoint")] + checkpoints = sorted(checkpoints, key=lambda x: int(x.split("-")[1])) + + # before we save the new checkpoint, we need to have at _most_ `checkpoints_total_limit - 1` checkpoints + if len(checkpoints) >= args.checkpoints_total_limit: + num_to_remove = len(checkpoints) - args.checkpoints_total_limit + 1 + removing_checkpoints = checkpoints[0:num_to_remove] + + logger.info( + f"{len(checkpoints)} checkpoints already exist, removing {len(removing_checkpoints)} checkpoints" + ) + logger.info(f"removing checkpoints: {', '.join(removing_checkpoints)}") + + for removing_checkpoint in removing_checkpoints: + removing_checkpoint = os.path.join(args.output_dir, removing_checkpoint) + shutil.rmtree(removing_checkpoint) + + save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}") + accelerator.save_state(save_path) + logger.info(f"Saved state to {save_path}") + + if accelerator.is_main_process: + if args.validation_prompts is not None and global_step % args.validation_steps == 0: + if args.use_ema: + # Store the UNet parameters temporarily and load the EMA parameters to perform inference. + ema_transformer3d.store(transformer3d.parameters()) + ema_transformer3d.copy_to(transformer3d.parameters()) + log_validation( + vae, + text_encoder, + tokenizer, + transformer3d, + args, + config, + accelerator, + weight_dtype, + global_step, + ) + if args.use_ema: + # Switch back to the original transformer3d parameters. + ema_transformer3d.restore(transformer3d.parameters()) + + logs = {"step_loss": loss.detach().item(), "lr": lr_scheduler.get_last_lr()[0]} + progress_bar.set_postfix(**logs) + + if global_step >= args.max_train_steps: + break + + if accelerator.is_main_process: + if args.validation_prompts is not None and epoch % args.validation_epochs == 0: + if args.use_ema: + # Store the UNet parameters temporarily and load the EMA parameters to perform inference. + ema_transformer3d.store(transformer3d.parameters()) + ema_transformer3d.copy_to(transformer3d.parameters()) + log_validation( + vae, + text_encoder, + tokenizer, + transformer3d, + args, + config, + accelerator, + weight_dtype, + global_step, + ) + if args.use_ema: + # Switch back to the original transformer3d parameters. + ema_transformer3d.restore(transformer3d.parameters()) + + # Create the pipeline using the trained modules and save it. + accelerator.wait_for_everyone() + if accelerator.is_main_process: + transformer3d = unwrap_model(transformer3d) + if args.use_ema: + ema_transformer3d.copy_to(transformer3d.parameters()) + + if args.use_deepspeed or args.use_fsdp or accelerator.is_main_process: + save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}") + accelerator.save_state(save_path) + logger.info(f"Saved state to {save_path}") + + accelerator.end_training() + + +if __name__ == "__main__": + main() diff --git a/scripts/wan2.2_fun/train.sh b/scripts/wan2.2_fun/train.sh new file mode 100644 index 0000000..0f6d27c --- /dev/null +++ b/scripts/wan2.2_fun/train.sh @@ -0,0 +1,43 @@ +export MODEL_NAME="models/Diffusion_Transformer/Wan2.2-Fun-A14B-InP" +export DATASET_NAME="datasets/internal_datasets/" +export DATASET_META_NAME="datasets/internal_datasets/metadata.json" +# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA. +# export NCCL_IB_DISABLE=1 +# export NCCL_P2P_DISABLE=1 +NCCL_DEBUG=INFO + +accelerate launch --mixed_precision="bf16" scripts/wan2.2_fun/train.py \ + --config_path="config/wan2.2/wan_civitai_i2v.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 \ + --boundary_type="low" \ + --train_mode="inpaint" \ + --trainable_modules "." \ No newline at end of file diff --git a/scripts/wan2.2_fun/train_control.py b/scripts/wan2.2_fun/train_control.py new file mode 100644 index 0000000..244b476 --- /dev/null +++ b/scripts/wan2.2_fun/train_control.py @@ -0,0 +1,2069 @@ +"""Modified from https://github.com/huggingface/diffusers/blob/main/examples/text_to_image/train_text_to_image.py +""" +#!/usr/bin/env python +# coding=utf-8 +# Copyright 2024 The HuggingFace Inc. team. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and + +import argparse +import gc +import logging +import math +import os +import pickle +import random +import shutil +import sys + +import accelerate +import diffusers +import numpy as np +import torch +import torch.nn.functional as F +import torch.utils.checkpoint +import torchvision.transforms.functional as TF +import transformers +from accelerate import Accelerator, FullyShardedDataParallelPlugin +from accelerate.logging import get_logger +from accelerate.state import AcceleratorState +from accelerate.utils import ProjectConfiguration, set_seed +from diffusers import DDIMScheduler, FlowMatchEulerDiscreteScheduler +from diffusers.optimization import get_scheduler +from diffusers.training_utils import (EMAModel, + compute_density_for_timestep_sampling, + compute_loss_weighting_for_sd3) +from diffusers.utils import check_min_version, deprecate, is_wandb_available +from diffusers.utils.torch_utils import is_compiled_module +from einops import rearrange +from omegaconf import OmegaConf +from packaging import version +from PIL import Image +from torch.utils.data import RandomSampler +from torch.utils.tensorboard import SummaryWriter +from torchvision import transforms +from tqdm.auto import tqdm +from transformers import AutoTokenizer +from transformers.utils import ContextManagers + +import datasets + +current_file_path = os.path.abspath(__file__) +project_roots = [os.path.dirname(current_file_path), os.path.dirname(os.path.dirname(current_file_path)), os.path.dirname(os.path.dirname(os.path.dirname(current_file_path)))] +for project_root in project_roots: + sys.path.insert(0, project_root) if project_root not in sys.path else None + +from videox_fun.data.bucket_sampler import (ASPECT_RATIO_512, + ASPECT_RATIO_RANDOM_CROP_512, + ASPECT_RATIO_RANDOM_CROP_PROB, + AspectRatioBatchImageVideoSampler, + RandomSampler, get_closest_ratio) +from videox_fun.data.dataset_image_video import (ImageVideoControlDataset, + ImageVideoDataset, + ImageVideoSampler, + get_random_mask, + process_pose_file, + process_pose_params) +from videox_fun.models import (AutoencoderKLWan, CLIPModel, WanT5EncoderModel, + Wan2_2Transformer3DModel) +from videox_fun.pipeline import WanFunControlPipeline +from videox_fun.utils.discrete_sampler import DiscreteSampling +from videox_fun.utils.lora_utils import (create_network, merge_lora, + unmerge_lora) +from videox_fun.utils.utils import (get_image_to_video_latent, + get_video_to_video_latent, + save_videos_grid) + +if is_wandb_available(): + import wandb + + +def filter_kwargs(cls, kwargs): + import inspect + sig = inspect.signature(cls.__init__) + valid_params = set(sig.parameters.keys()) - {'self', 'cls'} + filtered_kwargs = {k: v for k, v in kwargs.items() if k in valid_params} + return filtered_kwargs + +def linear_decay(initial_value, final_value, total_steps, current_step): + if current_step >= total_steps: + return final_value + current_step = max(0, current_step) + step_size = (final_value - initial_value) / total_steps + current_value = initial_value + step_size * current_step + return current_value + +def generate_timestep_with_lognorm(low, high, shape, device="cpu", generator=None): + u = torch.normal(mean=0.0, std=1.0, size=shape, device=device, generator=generator) + t = 1 / (1 + torch.exp(-u)) * (high - low) + low + return torch.clip(t.to(torch.int32), low, high - 1) + +def resize_mask(mask, latent, process_first_frame_only=True): + latent_size = latent.size() + batch_size, channels, num_frames, height, width = mask.shape + + if process_first_frame_only: + target_size = list(latent_size[2:]) + target_size[0] = 1 + first_frame_resized = F.interpolate( + mask[:, :, 0:1, :, :], + size=target_size, + mode='trilinear', + align_corners=False + ) + + target_size = list(latent_size[2:]) + target_size[0] = target_size[0] - 1 + if target_size[0] != 0: + remaining_frames_resized = F.interpolate( + mask[:, :, 1:, :, :], + size=target_size, + mode='trilinear', + align_corners=False + ) + resized_mask = torch.cat([first_frame_resized, remaining_frames_resized], dim=2) + else: + resized_mask = first_frame_resized + else: + target_size = list(latent_size[2:]) + resized_mask = F.interpolate( + mask, + size=target_size, + mode='trilinear', + align_corners=False + ) + return resized_mask + +# Will error if the minimal version of diffusers is not installed. Remove at your own risks. +check_min_version("0.18.0.dev0") + +logger = get_logger(__name__, log_level="INFO") + +def log_validation(vae, text_encoder, tokenizer, transformer3d, args, config, accelerator, weight_dtype, global_step): + try: + logger.info("Running validation... ") + + transformer3d_val = Wan2_2Transformer3DModel.from_pretrained( + os.path.join(args.pretrained_model_name_or_path, config['transformer_additional_kwargs'].get('transformer_subpath', 'transformer')), + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), + ).to(weight_dtype) + transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) + scheduler = FlowMatchEulerDiscreteScheduler( + **filter_kwargs(FlowMatchEulerDiscreteScheduler, OmegaConf.to_container(config['scheduler_kwargs'])) + ) + + pipeline = WanFunControlPipeline( + vae=accelerator.unwrap_model(vae).to(weight_dtype), + text_encoder=accelerator.unwrap_model(text_encoder), + tokenizer=tokenizer, + transformer=transformer3d_val, + scheduler=scheduler, + ) + pipeline = pipeline.to(accelerator.device) + + if args.seed is None: + generator = None + else: + generator = torch.Generator(device=accelerator.device).manual_seed(args.seed) + + images = [] + for i in range(len(args.validation_prompts)): + with torch.no_grad(): + with torch.autocast("cuda", dtype=weight_dtype): + video_length = int(args.video_sample_n_frames // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if args.video_sample_n_frames != 1 else 1 + input_video, input_video_mask, ref_image, clip_image = get_video_to_video_latent(args.validation_paths[i], video_length=video_length, sample_size=[args.video_sample_size, args.video_sample_size]) + sample = pipeline( + args.validation_prompts[i], + num_frames = video_length, + negative_prompt = "bad detailed", + height = args.video_sample_size, + width = args.video_sample_size, + generator = generator, + + control_video = input_video, + ).videos + os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) + save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-{i}.gif")) + + del pipeline + del transformer3d_val + gc.collect() + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + + return images + except Exception as e: + gc.collect() + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + print(f"Eval error with info {e}") + return None + +def parse_args(): + parser = argparse.ArgumentParser(description="Simple example of a training script.") + parser.add_argument( + "--input_perturbation", type=float, default=0, help="The scale of input perturbation. Recommended 0.1." + ) + parser.add_argument( + "--pretrained_model_name_or_path", + type=str, + default=None, + required=True, + help="Path to pretrained model or model identifier from huggingface.co/models.", + ) + parser.add_argument( + "--revision", + type=str, + default=None, + required=False, + help="Revision of pretrained model identifier from huggingface.co/models.", + ) + parser.add_argument( + "--variant", + type=str, + default=None, + help="Variant of the model files of the pretrained model identifier from huggingface.co/models, 'e.g.' fp16", + ) + parser.add_argument( + "--train_data_dir", + type=str, + default=None, + help=( + "A folder containing the training data. " + ), + ) + parser.add_argument( + "--train_data_meta", + type=str, + default=None, + help=( + "A csv containing the training data. " + ), + ) + parser.add_argument( + "--max_train_samples", + type=int, + default=None, + help=( + "For debugging purposes or quicker training, truncate the number of training examples to this " + "value if set." + ), + ) + parser.add_argument( + "--validation_prompts", + type=str, + default=None, + nargs="+", + help=("A set of prompts evaluated every `--validation_epochs` and logged to `--report_to`."), + ) + parser.add_argument( + "--output_dir", + type=str, + default="sd-model-finetuned", + help="The output directory where the model predictions and checkpoints will be written.", + ) + parser.add_argument( + "--cache_dir", + type=str, + default=None, + help="The directory where the downloaded models and datasets will be stored.", + ) + parser.add_argument("--seed", type=int, default=None, help="A seed for reproducible training.") + parser.add_argument( + "--random_flip", + action="store_true", + help="whether to randomly flip images horizontally", + ) + parser.add_argument( + "--use_came", + action="store_true", + help="whether to use came", + ) + parser.add_argument( + "--multi_stream", + action="store_true", + help="whether to use cuda multi-stream", + ) + parser.add_argument( + "--train_batch_size", type=int, default=16, help="Batch size (per device) for the training dataloader." + ) + parser.add_argument( + "--vae_mini_batch", type=int, default=32, help="mini batch size for vae." + ) + parser.add_argument("--num_train_epochs", type=int, default=100) + parser.add_argument( + "--max_train_steps", + type=int, + default=None, + help="Total number of training steps to perform. If provided, overrides num_train_epochs.", + ) + parser.add_argument( + "--gradient_accumulation_steps", + type=int, + default=1, + help="Number of updates steps to accumulate before performing a backward/update pass.", + ) + parser.add_argument( + "--gradient_checkpointing", + action="store_true", + help="Whether or not to use gradient checkpointing to save memory at the expense of slower backward pass.", + ) + parser.add_argument( + "--learning_rate", + type=float, + default=1e-4, + help="Initial learning rate (after the potential warmup period) to use.", + ) + parser.add_argument( + "--scale_lr", + action="store_true", + default=False, + help="Scale the learning rate by the number of GPUs, gradient accumulation steps, and batch size.", + ) + parser.add_argument( + "--lr_scheduler", + type=str, + default="constant", + help=( + 'The scheduler type to use. Choose between ["linear", "cosine", "cosine_with_restarts", "polynomial",' + ' "constant", "constant_with_warmup"]' + ), + ) + parser.add_argument( + "--lr_warmup_steps", type=int, default=500, help="Number of steps for the warmup in the lr scheduler." + ) + parser.add_argument( + "--use_8bit_adam", action="store_true", help="Whether or not to use 8-bit Adam from bitsandbytes." + ) + parser.add_argument( + "--allow_tf32", + action="store_true", + help=( + "Whether or not to allow TF32 on Ampere GPUs. Can be used to speed up training. For more information, see" + " https://pytorch.org/docs/stable/notes/cuda.html#tensorfloat-32-tf32-on-ampere-devices" + ), + ) + parser.add_argument("--use_ema", action="store_true", help="Whether to use EMA model.") + parser.add_argument( + "--non_ema_revision", + type=str, + default=None, + required=False, + help=( + "Revision of pretrained non-ema model identifier. Must be a branch, tag or git identifier of the local or" + " remote repository specified with --pretrained_model_name_or_path." + ), + ) + parser.add_argument( + "--dataloader_num_workers", + type=int, + default=0, + help=( + "Number of subprocesses to use for data loading. 0 means that the data will be loaded in the main process." + ), + ) + parser.add_argument("--adam_beta1", type=float, default=0.9, help="The beta1 parameter for the Adam optimizer.") + parser.add_argument("--adam_beta2", type=float, default=0.999, help="The beta2 parameter for the Adam optimizer.") + parser.add_argument("--adam_weight_decay", type=float, default=1e-2, help="Weight decay to use.") + parser.add_argument("--adam_epsilon", type=float, default=1e-08, help="Epsilon value for the Adam optimizer") + parser.add_argument("--max_grad_norm", default=1.0, type=float, help="Max gradient norm.") + parser.add_argument("--push_to_hub", action="store_true", help="Whether or not to push the model to the Hub.") + parser.add_argument("--hub_token", type=str, default=None, help="The token to use to push to the Model Hub.") + parser.add_argument( + "--prediction_type", + type=str, + default=None, + help="The prediction_type that shall be used for training. Choose between 'epsilon' or 'v_prediction' or leave `None`. If left to `None` the default prediction type of the scheduler: `noise_scheduler.config.prediciton_type` is chosen.", + ) + parser.add_argument( + "--hub_model_id", + type=str, + default=None, + help="The name of the repository to keep in sync with the local `output_dir`.", + ) + parser.add_argument( + "--logging_dir", + type=str, + default="logs", + help=( + "[TensorBoard](https://www.tensorflow.org/tensorboard) log directory. Will default to" + " *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***." + ), + ) + parser.add_argument( + "--report_model_info", action="store_true", help="Whether or not to report more info about model (such as norm, grad)." + ) + parser.add_argument( + "--mixed_precision", + type=str, + default=None, + choices=["no", "fp16", "bf16"], + help=( + "Whether to use mixed precision. Choose between fp16 and bf16 (bfloat16). Bf16 requires PyTorch >=" + " 1.10.and an Nvidia Ampere GPU. Default to the value of accelerate config of the current system or the" + " flag passed with the `accelerate.launch` command. Use this argument to override the accelerate config." + ), + ) + parser.add_argument( + "--report_to", + type=str, + default="tensorboard", + help=( + 'The integration to report the results and logs to. Supported platforms are `"tensorboard"`' + ' (default), `"wandb"` and `"comet_ml"`. Use `"all"` to report to all integrations.' + ), + ) + parser.add_argument("--local_rank", type=int, default=-1, help="For distributed training: local_rank") + parser.add_argument( + "--checkpointing_steps", + type=int, + default=500, + help=( + "Save a checkpoint of the training state every X updates. These checkpoints are only suitable for resuming" + " training using `--resume_from_checkpoint`." + ), + ) + parser.add_argument( + "--checkpoints_total_limit", + type=int, + default=None, + help=("Max number of checkpoints to store."), + ) + parser.add_argument( + "--resume_from_checkpoint", + type=str, + default=None, + help=( + "Whether training should be resumed from a previous checkpoint. Use a path saved by" + ' `--checkpointing_steps`, or `"latest"` to automatically select the last available checkpoint.' + ), + ) + parser.add_argument("--noise_offset", type=float, default=0, help="The scale of noise offset.") + parser.add_argument( + "--validation_epochs", + type=int, + default=5, + help="Run validation every X epochs.", + ) + parser.add_argument( + "--validation_steps", + type=int, + default=2000, + help="Run validation every X steps.", + ) + parser.add_argument( + "--tracker_project_name", + type=str, + default="text2image-fine-tune", + help=( + "The `project_name` argument passed to Accelerator.init_trackers for" + " more information see https://huggingface.co/docs/accelerate/v0.17.0/en/package_reference/accelerator#accelerate.Accelerator" + ), + ) + + parser.add_argument( + "--snr_loss", action="store_true", help="Whether or not to use snr_loss." + ) + parser.add_argument( + "--uniform_sampling", action="store_true", help="Whether or not to use uniform_sampling." + ) + parser.add_argument( + "--enable_text_encoder_in_dataloader", action="store_true", help="Whether or not to use text encoder in dataloader." + ) + parser.add_argument( + "--enable_bucket", action="store_true", help="Whether enable bucket sample in datasets." + ) + parser.add_argument( + "--random_ratio_crop", action="store_true", help="Whether enable random ratio crop sample in datasets." + ) + parser.add_argument( + "--random_frame_crop", action="store_true", help="Whether enable random frame crop sample in datasets." + ) + parser.add_argument( + "--random_hw_adapt", action="store_true", help="Whether enable random adapt height and width in datasets." + ) + parser.add_argument( + "--training_with_video_token_length", action="store_true", help="The training stage of the model in training.", + ) + parser.add_argument( + "--auto_tile_batch_size", action="store_true", help="Whether to auto tile batch size.", + ) + parser.add_argument( + "--motion_sub_loss", action="store_true", help="Whether enable motion sub loss." + ) + parser.add_argument( + "--motion_sub_loss_ratio", type=float, default=0.25, help="The ratio of motion sub loss." + ) + parser.add_argument( + "--train_sampling_steps", + type=int, + default=1000, + help="Run train_sampling_steps.", + ) + parser.add_argument( + "--keep_all_node_same_token_length", + action="store_true", + help="Reference of the length token.", + ) + parser.add_argument( + "--token_sample_size", + type=int, + default=512, + help="Sample size of the token.", + ) + parser.add_argument( + "--video_sample_size", + type=int, + default=512, + help="Sample size of the video.", + ) + parser.add_argument( + "--image_sample_size", + type=int, + default=512, + help="Sample size of the image.", + ) + parser.add_argument( + "--fix_sample_size", + nargs=2, type=int, default=None, + help="Fix Sample size [height, width] when using bucket and collate_fn." + ) + parser.add_argument( + "--video_sample_stride", + type=int, + default=4, + help="Sample stride of the video.", + ) + parser.add_argument( + "--video_sample_n_frames", + type=int, + default=17, + help="Num frame of video.", + ) + parser.add_argument( + "--video_repeat", + type=int, + default=0, + help="Num of repeat video.", + ) + parser.add_argument( + "--config_path", + type=str, + default=None, + help=( + "The config of the model in training." + ), + ) + parser.add_argument( + "--transformer_path", + type=str, + default=None, + help=("If you want to load the weight from other transformers, input its path."), + ) + parser.add_argument( + "--vae_path", + type=str, + default=None, + help=("If you want to load the weight from other vaes, input its path."), + ) + + parser.add_argument( + '--trainable_modules', + nargs='+', + help='Enter a list of trainable modules' + ) + parser.add_argument( + '--trainable_modules_low_learning_rate', + nargs='+', + default=[], + help='Enter a list of trainable modules with lower learning rate' + ) + parser.add_argument( + '--tokenizer_max_length', + type=int, + default=512, + help='Max length of tokenizer' + ) + parser.add_argument( + "--use_deepspeed", action="store_true", help="Whether or not to use deepspeed." + ) + parser.add_argument( + "--use_fsdp", action="store_true", help="Whether or not to use fsdp." + ) + parser.add_argument( + "--low_vram", action="store_true", help="Whether enable low_vram mode." + ) + parser.add_argument( + "--boundary_type", + type=str, + default="low", + help=( + 'The format of training data. Support `"low"` and `"high"`' + ), + ) + parser.add_argument( + "--abnormal_norm_clip_start", + type=int, + default=1000, + help=( + 'When do we start doing additional processing on abnormal gradients. ' + ), + ) + parser.add_argument( + "--initial_grad_norm_ratio", + type=int, + default=5, + help=( + 'The initial gradient is relative to the multiple of the max_grad_norm. ' + ), + ) + parser.add_argument( + "--train_mode", + type=str, + default="control", + help=( + 'The format of training data. Support `"control"`' + ' (default), `"control_ref"`, `"control_camera_ref"`.' + ), + ) + parser.add_argument( + "--control_ref_image", + type=str, + default="first_frame", + help=( + 'The format of training data. Support `"first_frame"`' + ' (default), `"random"`.' + ), + ) + parser.add_argument( + "--add_full_ref_image_in_self_attention", + action="store_true", + help=( + 'Whether enable add full ref image in self attention.' + ), + ) + parser.add_argument( + "--add_inpaint_info", + action="store_true", + help=( + 'Whether enable add inpaint info in self attention.' + ), + ) + parser.add_argument( + "--weighting_scheme", + type=str, + default="none", + choices=["sigma_sqrt", "logit_normal", "mode", "cosmap", "none"], + help=('We default to the "none" weighting scheme for uniform sampling and uniform loss'), + ) + parser.add_argument( + "--logit_mean", type=float, default=0.0, help="mean to use when using the `'logit_normal'` weighting scheme." + ) + parser.add_argument( + "--logit_std", type=float, default=1.0, help="std to use when using the `'logit_normal'` weighting scheme." + ) + parser.add_argument( + "--mode_scale", + type=float, + default=1.29, + help="Scale of mode weighting scheme. Only effective when using the `'mode'` as the `weighting_scheme`.", + ) + + args = parser.parse_args() + env_local_rank = int(os.environ.get("LOCAL_RANK", -1)) + if env_local_rank != -1 and env_local_rank != args.local_rank: + args.local_rank = env_local_rank + + # default to using the same revision for the non-ema model if not specified + if args.non_ema_revision is None: + args.non_ema_revision = args.revision + + return args + + +def main(): + args = parse_args() + + if args.report_to == "wandb" and args.hub_token is not None: + raise ValueError( + "You cannot use both --report_to=wandb and --hub_token due to a security risk of exposing your token." + " Please use `huggingface-cli login` to authenticate with the Hub." + ) + + if args.non_ema_revision is not None: + deprecate( + "non_ema_revision!=None", + "0.15.0", + message=( + "Downloading 'non_ema' weights from revision branches of the Hub is deprecated. Please make sure to" + " use `--variant=non_ema` instead." + ), + ) + logging_dir = os.path.join(args.output_dir, args.logging_dir) + + config = OmegaConf.load(args.config_path) + accelerator_project_config = ProjectConfiguration(project_dir=args.output_dir, logging_dir=logging_dir) + + accelerator = Accelerator( + gradient_accumulation_steps=args.gradient_accumulation_steps, + mixed_precision=args.mixed_precision, + log_with=args.report_to, + project_config=accelerator_project_config, + ) + + deepspeed_plugin = accelerator.state.deepspeed_plugin if hasattr(accelerator.state, "deepspeed_plugin") else None + fsdp_plugin = accelerator.state.fsdp_plugin if hasattr(accelerator.state, "fsdp_plugin") else None + if deepspeed_plugin is not None: + zero_stage = int(deepspeed_plugin.zero_stage) + fsdp_stage = 0 + print(f"Using DeepSpeed Zero stage: {zero_stage}") + + args.use_deepspeed = True + if zero_stage == 3: + print(f"Auto set save_state to True because zero_stage == 3") + args.save_state = True + elif fsdp_plugin is not None: + from torch.distributed.fsdp import ShardingStrategy + zero_stage = 0 + if fsdp_plugin.sharding_strategy is ShardingStrategy.FULL_SHARD: + fsdp_stage = 3 + elif fsdp_plugin.sharding_strategy is None: # The fsdp_plugin.sharding_strategy is None in FSDP 2. + fsdp_stage = 3 + elif fsdp_plugin.sharding_strategy is ShardingStrategy.SHARD_GRAD_OP: + fsdp_stage = 2 + else: + fsdp_stage = 0 + print(f"Using FSDP stage: {fsdp_stage}") + + args.use_fsdp = True + if fsdp_stage == 3: + print(f"Auto set save_state to True because fsdp_stage == 3") + args.save_state = True + else: + zero_stage = 0 + fsdp_stage = 0 + print("DeepSpeed is not enabled.") + + if accelerator.is_main_process: + writer = SummaryWriter(log_dir=logging_dir) + + # Make one log on every process with the configuration for debugging. + logging.basicConfig( + format="%(asctime)s - %(levelname)s - %(name)s - %(message)s", + datefmt="%m/%d/%Y %H:%M:%S", + level=logging.INFO, + ) + logger.info(accelerator.state, main_process_only=False) + if accelerator.is_local_main_process: + datasets.utils.logging.set_verbosity_warning() + transformers.utils.logging.set_verbosity_warning() + diffusers.utils.logging.set_verbosity_info() + else: + datasets.utils.logging.set_verbosity_error() + transformers.utils.logging.set_verbosity_error() + diffusers.utils.logging.set_verbosity_error() + + # If passed along, set the training seed now. + if args.seed is not None: + set_seed(args.seed) + rng = np.random.default_rng(np.random.PCG64(args.seed + accelerator.process_index)) + torch_rng = torch.Generator(accelerator.device).manual_seed(args.seed + accelerator.process_index) + else: + rng = None + torch_rng = None + index_rng = np.random.default_rng(np.random.PCG64(43)) + print(f"Init rng with seed {args.seed + accelerator.process_index}. Process_index is {accelerator.process_index}") + + # Handle the repository creation + if accelerator.is_main_process: + if args.output_dir is not None: + os.makedirs(args.output_dir, exist_ok=True) + + # For mixed precision training we cast all non-trainable weigths (vae, non-lora text_encoder and non-lora transformer3d) to half-precision + # as these weights are only used for inference, keeping weights in full precision is not required. + weight_dtype = torch.float32 + if accelerator.mixed_precision == "fp16": + weight_dtype = torch.float16 + args.mixed_precision = accelerator.mixed_precision + elif accelerator.mixed_precision == "bf16": + weight_dtype = torch.bfloat16 + args.mixed_precision = accelerator.mixed_precision + + # Load scheduler, tokenizer and models. + noise_scheduler = FlowMatchEulerDiscreteScheduler( + **filter_kwargs(FlowMatchEulerDiscreteScheduler, OmegaConf.to_container(config['scheduler_kwargs'])) + ) + + # Get Tokenizer + tokenizer = AutoTokenizer.from_pretrained( + os.path.join(args.pretrained_model_name_or_path, config['text_encoder_kwargs'].get('tokenizer_subpath', 'tokenizer')), + ) + + def deepspeed_zero_init_disabled_context_manager(): + """ + returns either a context list that includes one that will disable zero.Init or an empty context list + """ + deepspeed_plugin = AcceleratorState().deepspeed_plugin if accelerate.state.is_initialized() else None + if deepspeed_plugin is None: + return [] + + return [deepspeed_plugin.zero3_init_context_manager(enable=False)] + + # Currently Accelerate doesn't know how to handle multiple models under Deepspeed ZeRO stage 3. + # For this to work properly all models must be run through `accelerate.prepare`. But accelerate + # will try to assign the same optimizer with the same weights to all models during + # `deepspeed.initialize`, which of course doesn't work. + # + # For now the following workaround will partially support Deepspeed ZeRO-3, by excluding the 2 + # frozen models from being partitioned during `zero.Init` which gets called during + # `from_pretrained` So CLIPTextModel and AutoencoderKL will not enjoy the parameter sharding + # across multiple gpus and only UNet2DConditionModel will get ZeRO sharded. + with ContextManagers(deepspeed_zero_init_disabled_context_manager()): + # Get Text encoder + text_encoder = WanT5EncoderModel.from_pretrained( + os.path.join(args.pretrained_model_name_or_path, config['text_encoder_kwargs'].get('text_encoder_subpath', 'text_encoder')), + additional_kwargs=OmegaConf.to_container(config['text_encoder_kwargs']), + low_cpu_mem_usage=True, + torch_dtype=weight_dtype, + ) + text_encoder = text_encoder.eval() + # Get Vae + vae = AutoencoderKLWan.from_pretrained( + os.path.join(args.pretrained_model_name_or_path, config['vae_kwargs'].get('vae_subpath', 'vae')), + additional_kwargs=OmegaConf.to_container(config['vae_kwargs']), + ) + vae.eval() + + # Get Transformer + sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer') \ + if args.boundary_type == "low" else config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer') + transformer3d = Wan2_2Transformer3DModel.from_pretrained( + os.path.join(args.pretrained_model_name_or_path, sub_path), + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), + ).to(weight_dtype) + + # Freeze vae and text_encoder and set transformer3d to trainable + vae.requires_grad_(False) + text_encoder.requires_grad_(False) + transformer3d.requires_grad_(False) + + if args.transformer_path is not None: + print(f"From checkpoint: {args.transformer_path}") + if args.transformer_path.endswith("safetensors"): + from safetensors.torch import load_file, safe_open + state_dict = load_file(args.transformer_path) + else: + state_dict = torch.load(args.transformer_path, map_location="cpu") + state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict + + m, u = transformer3d.load_state_dict(state_dict, strict=False) + print(f"missing keys: {len(m)}, unexpected keys: {len(u)}") + + if args.vae_path is not None: + print(f"From checkpoint: {args.vae_path}") + if args.vae_path.endswith("safetensors"): + from safetensors.torch import load_file, safe_open + state_dict = load_file(args.vae_path) + else: + state_dict = torch.load(args.vae_path, map_location="cpu") + state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict + + m, u = vae.load_state_dict(state_dict, strict=False) + print(f"missing keys: {len(m)}, unexpected keys: {len(u)}") + + # A good trainable modules is showed below now. + # For 3D Patch: trainable_modules = ['ff.net', 'pos_embed', 'attn2', 'proj_out', 'timepositionalencoding', 'h_position', 'w_position'] + # For 2D Patch: trainable_modules = ['ff.net', 'attn2', 'timepositionalencoding', 'h_position', 'w_position'] + transformer3d.train() + if accelerator.is_main_process: + accelerator.print( + f"Trainable modules '{args.trainable_modules}'." + ) + for name, param in transformer3d.named_parameters(): + for trainable_module_name in args.trainable_modules + args.trainable_modules_low_learning_rate: + if trainable_module_name in name: + param.requires_grad = True + break + + # Create EMA for the transformer3d. + if args.use_ema: + if zero_stage == 3: + raise NotImplementedError("FSDP does not support EMA.") + + ema_transformer3d = Wan2_2Transformer3DModel.from_pretrained( + os.path.join(args.pretrained_model_name_or_path, config['transformer_additional_kwargs'].get('transformer_subpath', 'transformer')), + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), + ).to(weight_dtype) + + ema_transformer3d = EMAModel(ema_transformer3d.parameters(), model_cls=Wan2_2Transformer3DModel, model_config=ema_transformer3d.config) + + # `accelerate` 0.16.0 will have better support for customized saving + if version.parse(accelerate.__version__) >= version.parse("0.16.0"): + # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format + if fsdp_stage != 0: + def save_model_hook(models, weights, output_dir): + accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) + if accelerator.is_main_process: + from safetensors.torch import save_file + + safetensor_save_path = os.path.join(output_dir, f"diffusion_pytorch_model.safetensors") + accelerate_state_dict = {k: v.to(dtype=weight_dtype) for k, v in accelerate_state_dict.items()} + save_file(accelerate_state_dict, safetensor_save_path, metadata={"format": "pt"}) + + with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: + pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) + + def load_model_hook(models, input_dir): + pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") + if os.path.exists(pkl_path): + with open(pkl_path, 'rb') as file: + loaded_number, _ = pickle.load(file) + batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) + print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") + + elif zero_stage == 3: + # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format + def save_model_hook(models, weights, output_dir): + if accelerator.is_main_process: + with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: + pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) + + def load_model_hook(models, input_dir): + pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") + if os.path.exists(pkl_path): + with open(pkl_path, 'rb') as file: + loaded_number, _ = pickle.load(file) + batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) + print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") + else: + # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format + def save_model_hook(models, weights, output_dir): + if accelerator.is_main_process: + if args.use_ema: + ema_transformer3d.save_pretrained(os.path.join(output_dir, "transformer_ema")) + + models[0].save_pretrained(os.path.join(output_dir, "transformer")) + if not args.use_deepspeed: + weights.pop() + + with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: + pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) + + def load_model_hook(models, input_dir): + if args.use_ema: + ema_path = os.path.join(input_dir, "transformer_ema") + _, ema_kwargs = Wan2_2Transformer3DModel.load_config(ema_path, return_unused_kwargs=True) + load_model = Wan2_2Transformer3DModel.from_pretrained( + input_dir, subfolder="transformer_ema", + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']) + ) + load_model = EMAModel(load_model.parameters(), model_cls=Wan2_2Transformer3DModel, model_config=load_model.config) + load_model.load_state_dict(ema_kwargs) + + ema_transformer3d.load_state_dict(load_model.state_dict()) + ema_transformer3d.to(accelerator.device) + del load_model + + for i in range(len(models)): + # pop models so that they are not loaded again + model = models.pop() + + # load diffusers style into model + load_model = Wan2_2Transformer3DModel.from_pretrained( + input_dir, subfolder="transformer" + ) + model.register_to_config(**load_model.config) + + model.load_state_dict(load_model.state_dict()) + del load_model + + pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") + if os.path.exists(pkl_path): + with open(pkl_path, 'rb') as file: + loaded_number, _ = pickle.load(file) + batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) + print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") + + accelerator.register_save_state_pre_hook(save_model_hook) + accelerator.register_load_state_pre_hook(load_model_hook) + + if args.gradient_checkpointing: + transformer3d.enable_gradient_checkpointing() + + # Enable TF32 for faster training on Ampere GPUs, + # cf https://pytorch.org/docs/stable/notes/cuda.html#tensorfloat-32-tf32-on-ampere-devices + if args.allow_tf32: + torch.backends.cuda.matmul.allow_tf32 = True + + if args.scale_lr: + args.learning_rate = ( + args.learning_rate * args.gradient_accumulation_steps * args.train_batch_size * accelerator.num_processes + ) + + # Initialize the optimizer + if args.use_8bit_adam: + try: + import bitsandbytes as bnb + except ImportError: + raise ImportError( + "Please install bitsandbytes to use 8-bit Adam. You can do so by running `pip install bitsandbytes`" + ) + + optimizer_cls = bnb.optim.AdamW8bit + elif args.use_came: + try: + from came_pytorch import CAME + except: + raise ImportError( + "Please install came_pytorch to use CAME. You can do so by running `pip install came_pytorch`" + ) + + optimizer_cls = CAME + else: + optimizer_cls = torch.optim.AdamW + + trainable_params = list(filter(lambda p: p.requires_grad, transformer3d.parameters())) + trainable_params_optim = [ + {'params': [], 'lr': args.learning_rate}, + {'params': [], 'lr': args.learning_rate / 2}, + ] + in_already = [] + for name, param in transformer3d.named_parameters(): + high_lr_flag = False + if name in in_already: + continue + for trainable_module_name in args.trainable_modules: + if trainable_module_name in name: + in_already.append(name) + high_lr_flag = True + trainable_params_optim[0]['params'].append(param) + if accelerator.is_main_process: + print(f"Set {name} to lr : {args.learning_rate}") + break + if high_lr_flag: + continue + for trainable_module_name in args.trainable_modules_low_learning_rate: + if trainable_module_name in name: + in_already.append(name) + trainable_params_optim[1]['params'].append(param) + if accelerator.is_main_process: + print(f"Set {name} to lr : {args.learning_rate / 2}") + break + + if args.use_came: + optimizer = optimizer_cls( + trainable_params_optim, + lr=args.learning_rate, + # weight_decay=args.adam_weight_decay, + betas=(0.9, 0.999, 0.9999), + eps=(1e-30, 1e-16) + ) + else: + optimizer = optimizer_cls( + trainable_params_optim, + lr=args.learning_rate, + betas=(args.adam_beta1, args.adam_beta2), + weight_decay=args.adam_weight_decay, + eps=args.adam_epsilon, + ) + + # Get the training dataset + sample_n_frames_bucket_interval = vae.config.temporal_compression_ratio + + if args.fix_sample_size is not None and args.enable_bucket: + args.video_sample_size = max(max(args.fix_sample_size), args.video_sample_size) + args.image_sample_size = max(max(args.fix_sample_size), args.image_sample_size) + args.training_with_video_token_length = False + args.random_hw_adapt = False + + # Get the dataset + train_dataset = ImageVideoControlDataset( + args.train_data_meta, args.train_data_dir, + video_sample_size=args.video_sample_size, video_sample_stride=args.video_sample_stride, video_sample_n_frames=args.video_sample_n_frames, + video_repeat=args.video_repeat, + image_sample_size=args.image_sample_size, + enable_bucket=args.enable_bucket, + enable_camera_info=args.train_mode == "control_camera_ref" + ) + + def worker_init_fn(_seed): + _seed = _seed * 256 + def _worker_init_fn(worker_id): + print(f"worker_init_fn with {_seed + worker_id}") + np.random.seed(_seed + worker_id) + random.seed(_seed + worker_id) + return _worker_init_fn + + if args.enable_bucket: + aspect_ratio_sample_size = {key : [x / 512 * args.video_sample_size for x in ASPECT_RATIO_512[key]] for key in ASPECT_RATIO_512.keys()} + batch_sampler_generator = torch.Generator().manual_seed(args.seed) + batch_sampler = AspectRatioBatchImageVideoSampler( + sampler=RandomSampler(train_dataset, generator=batch_sampler_generator), dataset=train_dataset.dataset, + batch_size=args.train_batch_size, train_folder = args.train_data_dir, drop_last=True, + aspect_ratios=aspect_ratio_sample_size, + ) + + def collate_fn(examples): + def get_length_to_frame_num(token_length): + if args.image_sample_size > args.video_sample_size: + sample_sizes = list(range(args.video_sample_size, args.image_sample_size + 1, 128)) + + if sample_sizes[-1] != args.image_sample_size: + sample_sizes.append(args.image_sample_size) + else: + sample_sizes = [args.image_sample_size] + + length_to_frame_num = { + sample_size: min(token_length / sample_size / sample_size, args.video_sample_n_frames) // sample_n_frames_bucket_interval * sample_n_frames_bucket_interval + 1 for sample_size in sample_sizes + } + + return length_to_frame_num + + def get_random_downsample_ratio(sample_size, image_ratio=[], + all_choices=False, rng=None): + def _create_special_list(length): + if length == 1: + return [1.0] + if length >= 2: + first_element = 0.90 + remaining_sum = 1.0 - first_element + other_elements_value = remaining_sum / (length - 1) + special_list = [first_element] + [other_elements_value] * (length - 1) + return special_list + + if sample_size >= 1536: + number_list = [1, 1.25, 1.5, 2, 2.5, 3] + image_ratio + elif sample_size >= 1024: + number_list = [1, 1.25, 1.5, 2] + image_ratio + elif sample_size >= 768: + number_list = [1, 1.25, 1.5] + image_ratio + elif sample_size >= 512: + number_list = [1] + image_ratio + else: + number_list = [1] + + if all_choices: + return number_list + + number_list_prob = np.array(_create_special_list(len(number_list))) + if rng is None: + return np.random.choice(number_list, p = number_list_prob) + else: + return rng.choice(number_list, p = number_list_prob) + + # Get token length + target_token_length = args.video_sample_n_frames * args.token_sample_size * args.token_sample_size + length_to_frame_num = get_length_to_frame_num(target_token_length) + + # Create new output + new_examples = {} + new_examples["target_token_length"] = target_token_length + new_examples["pixel_values"] = [] + new_examples["text"] = [] + # Used in Control Mode + new_examples["control_pixel_values"] = [] + # Used in Control Ref Mode + if args.train_mode != "control": + new_examples["ref_pixel_values"] = [] + new_examples["clip_pixel_values"] = [] + new_examples["clip_idx"] = [] + # Used in Control Camera Ref Mode + if args.train_mode == "control_camera_ref": + new_examples["control_camera_values"] = [] + + # Used in Inpaint mode + if args.add_inpaint_info: + new_examples["mask_pixel_values"] = [] + new_examples["mask"] = [] + new_examples["clip_pixel_values"] = [] + + # Get downsample ratio in image and videos + pixel_value = examples[0]["pixel_values"] + data_type = examples[0]["data_type"] + f, h, w, c = np.shape(pixel_value) + if data_type == 'image': + random_downsample_ratio = 1 if not args.random_hw_adapt else get_random_downsample_ratio(args.image_sample_size, image_ratio=[args.image_sample_size / args.video_sample_size]) + + aspect_ratio_sample_size = {key : [x / 512 * args.image_sample_size / random_downsample_ratio for x in ASPECT_RATIO_512[key]] for key in ASPECT_RATIO_512.keys()} + aspect_ratio_random_crop_sample_size = {key : [x / 512 * args.image_sample_size / random_downsample_ratio for x in ASPECT_RATIO_RANDOM_CROP_512[key]] for key in ASPECT_RATIO_RANDOM_CROP_512.keys()} + + batch_video_length = args.video_sample_n_frames + sample_n_frames_bucket_interval + else: + if args.random_hw_adapt: + if args.training_with_video_token_length: + local_min_size = np.min(np.array([np.mean(np.array([np.shape(example["pixel_values"])[1], np.shape(example["pixel_values"])[2]])) for example in examples])) + + def get_random_downsample_probability(choice_list, token_sample_size): + length = len(choice_list) + if length == 1: + return [1.0] # If there's only one element, it gets all the probability + + # Find the index of the closest value to token_sample_size + closest_index = min(range(length), key=lambda i: abs(choice_list[i] - token_sample_size)) + + # Assign 50% to the closest index + first_element = 0.50 + remaining_sum = 1.0 - first_element + + # Distribute the remaining 50% evenly among the other elements + other_elements_value = remaining_sum / (length - 1) if length > 1 else 0.0 + + # Construct the probability distribution + probability_list = [other_elements_value] * length + probability_list[closest_index] = first_element + + return probability_list + + choice_list = [length for length in list(length_to_frame_num.keys()) if length < local_min_size * 1.25] + if len(choice_list) == 0: + choice_list = list(length_to_frame_num.keys()) + probabilities = get_random_downsample_probability(choice_list, args.token_sample_size) + local_video_sample_size = np.random.choice(choice_list, p=probabilities) + + random_downsample_ratio = args.video_sample_size / local_video_sample_size + batch_video_length = length_to_frame_num[local_video_sample_size] + else: + random_downsample_ratio = get_random_downsample_ratio(args.video_sample_size) + batch_video_length = args.video_sample_n_frames + sample_n_frames_bucket_interval + else: + random_downsample_ratio = 1 + batch_video_length = args.video_sample_n_frames + sample_n_frames_bucket_interval + + aspect_ratio_sample_size = {key : [x / 512 * args.video_sample_size / random_downsample_ratio for x in ASPECT_RATIO_512[key]] for key in ASPECT_RATIO_512.keys()} + aspect_ratio_random_crop_sample_size = {key : [x / 512 * args.video_sample_size / random_downsample_ratio for x in ASPECT_RATIO_RANDOM_CROP_512[key]] for key in ASPECT_RATIO_RANDOM_CROP_512.keys()} + + if args.fix_sample_size is not None: + fix_sample_size = [int(x / 16) * 16 for x in args.fix_sample_size] + elif args.random_ratio_crop: + if rng is None: + random_sample_size = aspect_ratio_random_crop_sample_size[ + np.random.choice(list(aspect_ratio_random_crop_sample_size.keys()), p = ASPECT_RATIO_RANDOM_CROP_PROB) + ] + else: + random_sample_size = aspect_ratio_random_crop_sample_size[ + rng.choice(list(aspect_ratio_random_crop_sample_size.keys()), p = ASPECT_RATIO_RANDOM_CROP_PROB) + ] + random_sample_size = [int(x / 16) * 16 for x in random_sample_size] + else: + closest_size, closest_ratio = get_closest_ratio(h, w, ratios=aspect_ratio_sample_size) + closest_size = [int(x / 16) * 16 for x in closest_size] + + for example in examples: + # To 0~1 + pixel_values = torch.from_numpy(example["pixel_values"]).permute(0, 3, 1, 2).contiguous() + pixel_values = pixel_values / 255. + + control_pixel_values = torch.from_numpy(example["control_pixel_values"]).permute(0, 3, 1, 2).contiguous() + control_pixel_values = control_pixel_values / 255. + + if args.fix_sample_size is not None: + # Get adapt hw for resize + fix_sample_size = list(map(lambda x: int(x), fix_sample_size)) + transform = transforms.Compose([ + transforms.Resize(fix_sample_size, interpolation=transforms.InterpolationMode.BILINEAR), # Image.BICUBIC + transforms.CenterCrop(fix_sample_size), + transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True), + ]) + + transform_no_normalize = transforms.Compose([ + transforms.Resize(fix_sample_size, interpolation=transforms.InterpolationMode.BILINEAR), # Image.BICUBIC + transforms.CenterCrop(fix_sample_size), + ]) + elif args.random_ratio_crop: + # Get adapt hw for resize + b, c, h, w = pixel_values.size() + th, tw = random_sample_size + if th / tw > h / w: + nh = int(th) + nw = int(w / h * nh) + else: + nw = int(tw) + nh = int(h / w * nw) + + transform = transforms.Compose([ + transforms.Resize([nh, nw]), + transforms.CenterCrop([int(x) for x in random_sample_size]), + transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True), + ]) + + transform_no_normalize = transforms.Compose([ + transforms.Resize([nh, nw]), + transforms.CenterCrop([int(x) for x in random_sample_size]), + ]) + else: + # Get adapt hw for resize + closest_size = list(map(lambda x: int(x), closest_size)) + if closest_size[0] / h > closest_size[1] / w: + resize_size = closest_size[0], int(w * closest_size[0] / h) + else: + resize_size = int(h * closest_size[1] / w), closest_size[1] + + transform = transforms.Compose([ + transforms.Resize(resize_size, interpolation=transforms.InterpolationMode.BILINEAR), # Image.BICUBIC + transforms.CenterCrop(closest_size), + transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True), + ]) + + transform_no_normalize = transforms.Compose([ + transforms.Resize(resize_size, interpolation=transforms.InterpolationMode.BILINEAR), # Image.BICUBIC + transforms.CenterCrop(closest_size), + ]) + + new_examples["pixel_values"].append(transform(pixel_values)) + new_examples["control_pixel_values"].append(transform(control_pixel_values)) + + if args.train_mode == "control_camera_ref": + control_camera_values = example.get("control_camera_values", None) + if control_camera_values is None: + control_camera_values_size = ( + new_examples["control_pixel_values"][-1].size()[0], + 6, + new_examples["control_pixel_values"][-1].size()[2], + new_examples["control_pixel_values"][-1].size()[3] + ) + local_control_camera_values = torch.zeros(control_camera_values_size) + new_examples["control_camera_values"].append(local_control_camera_values) + else: + local_control_camera_values = process_pose_params(example["control_camera_values"], height=resize_size[0], width=resize_size[1]).permute(0, 3, 1, 2).contiguous() + new_examples["control_camera_values"].append(transform_no_normalize(local_control_camera_values)) + + new_examples["text"].append(example["text"]) + # Magvae needs the number of frames to be 4n + 1. + batch_video_length = int( + min( + batch_video_length, + (len(pixel_values) - 1) // sample_n_frames_bucket_interval * sample_n_frames_bucket_interval + 1, + ) + ) + if batch_video_length == 0: + batch_video_length = 1 + + if args.train_mode != "control": + if args.control_ref_image == "first_frame": + clip_index = 0 + else: + def _create_special_list(length): + if length == 1: + return [1.0] + if length >= 2: + first_element = 0.40 + remaining_sum = 1.0 - first_element + other_elements_value = remaining_sum / (length - 1) + special_list = [first_element] + [other_elements_value] * (length - 1) + return special_list + number_list_prob = np.array(_create_special_list(len(new_examples["pixel_values"][-1]))) + clip_index = np.random.choice(list(range(len(new_examples["pixel_values"][-1]))), p = number_list_prob) + new_examples["clip_idx"].append(clip_index) + + ref_pixel_values = new_examples["pixel_values"][-1][clip_index].unsqueeze(0) + new_examples["ref_pixel_values"].append(ref_pixel_values) + + clip_pixel_values = new_examples["pixel_values"][-1][clip_index].permute(1, 2, 0).contiguous() + clip_pixel_values = (clip_pixel_values * 0.5 + 0.5) * 255 + new_examples["clip_pixel_values"].append(clip_pixel_values) + + if args.add_inpaint_info: + mask = get_random_mask(new_examples["pixel_values"][-1].size()) + mask_pixel_values = new_examples["pixel_values"][-1] * (1 - mask) + # Wan 2.1 use 0 for masked pixels + # + torch.ones_like(new_examples["pixel_values"][-1]) * -1 * mask + new_examples["mask_pixel_values"].append(mask_pixel_values) + new_examples["mask"].append(mask) + + # Limit the number of frames to the same + new_examples["pixel_values"] = torch.stack([example[:batch_video_length] for example in new_examples["pixel_values"]]) + new_examples["control_pixel_values"] = torch.stack([example[:batch_video_length] for example in new_examples["control_pixel_values"]]) + if args.train_mode != "control": + new_examples["ref_pixel_values"] = torch.stack([example[:batch_video_length] for example in new_examples["ref_pixel_values"]]) + new_examples["clip_pixel_values"] = torch.stack([example for example in new_examples["clip_pixel_values"]]) + new_examples["clip_idx"] = torch.tensor(new_examples["clip_idx"]) + if args.train_mode == "control_camera_ref": + new_examples["control_camera_values"] = torch.stack([example[:batch_video_length] for example in new_examples["control_camera_values"]]) + if args.add_inpaint_info: + new_examples["mask_pixel_values"] = torch.stack([example[:batch_video_length] for example in new_examples["mask_pixel_values"]]) + new_examples["mask"] = torch.stack([example[:batch_video_length] for example in new_examples["mask"]]) + + # Encode prompts when enable_text_encoder_in_dataloader=True + if args.enable_text_encoder_in_dataloader: + prompt_ids = tokenizer( + new_examples['text'], + max_length=args.tokenizer_max_length, + padding="max_length", + add_special_tokens=True, + truncation=True, + return_tensors="pt" + ) + encoder_hidden_states = text_encoder( + prompt_ids.input_ids + )[0] + new_examples['encoder_attention_mask'] = prompt_ids.attention_mask + new_examples['encoder_hidden_states'] = encoder_hidden_states + + return new_examples + + # DataLoaders creation: + train_dataloader = torch.utils.data.DataLoader( + train_dataset, + batch_sampler=batch_sampler, + collate_fn=collate_fn, + persistent_workers=True if args.dataloader_num_workers != 0 else False, + num_workers=args.dataloader_num_workers, + worker_init_fn=worker_init_fn(args.seed + accelerator.process_index) + ) + else: + # DataLoaders creation: + batch_sampler_generator = torch.Generator().manual_seed(args.seed) + batch_sampler = ImageVideoSampler(RandomSampler(train_dataset, generator=batch_sampler_generator), train_dataset, args.train_batch_size) + train_dataloader = torch.utils.data.DataLoader( + train_dataset, + batch_sampler=batch_sampler, + persistent_workers=True if args.dataloader_num_workers != 0 else False, + num_workers=args.dataloader_num_workers, + worker_init_fn=worker_init_fn(args.seed + accelerator.process_index) + ) + + # Scheduler and math around the number of training steps. + overrode_max_train_steps = False + num_update_steps_per_epoch = math.ceil(len(train_dataloader) / args.gradient_accumulation_steps) + if args.max_train_steps is None: + args.max_train_steps = args.num_train_epochs * num_update_steps_per_epoch + overrode_max_train_steps = True + + lr_scheduler = get_scheduler( + args.lr_scheduler, + optimizer=optimizer, + num_warmup_steps=args.lr_warmup_steps * accelerator.num_processes, + num_training_steps=args.max_train_steps * accelerator.num_processes, + ) + + # Prepare everything with our `accelerator`. + transformer3d, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( + transformer3d, optimizer, train_dataloader, lr_scheduler + ) + + if fsdp_stage != 0: + from functools import partial + from videox_fun.dist import set_multi_gpus_devices, shard_model + shard_fn = partial(shard_model, device_id=accelerator.device, param_dtype=weight_dtype) + text_encoder = shard_fn(text_encoder) + + if args.use_ema: + ema_transformer3d.to(accelerator.device) + + # Move text_encode and vae to gpu and cast to weight_dtype + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu") + + # We need to recalculate our total training steps as the size of the training dataloader may have changed. + num_update_steps_per_epoch = math.ceil(len(train_dataloader) / args.gradient_accumulation_steps) + if overrode_max_train_steps: + args.max_train_steps = args.num_train_epochs * num_update_steps_per_epoch + # Afterwards we recalculate our number of training epochs + args.num_train_epochs = math.ceil(args.max_train_steps / num_update_steps_per_epoch) + + # We need to initialize the trackers we use, and also store our configuration. + # The trackers initializes automatically on the main process. + if accelerator.is_main_process: + tracker_config = dict(vars(args)) + tracker_config.pop("validation_prompts") + tracker_config.pop("trainable_modules") + tracker_config.pop("trainable_modules_low_learning_rate") + tracker_config.pop("fix_sample_size") + accelerator.init_trackers(args.tracker_project_name, tracker_config) + + # Function for unwrapping if model was compiled with `torch.compile`. + def unwrap_model(model): + model = accelerator.unwrap_model(model) + model = model._orig_mod if is_compiled_module(model) else model + return model + + # Train! + total_batch_size = args.train_batch_size * accelerator.num_processes * args.gradient_accumulation_steps + + logger.info("***** Running training *****") + logger.info(f" Num examples = {len(train_dataset)}") + logger.info(f" Num Epochs = {args.num_train_epochs}") + logger.info(f" Instantaneous batch size per device = {args.train_batch_size}") + logger.info(f" Total train batch size (w. parallel, distributed & accumulation) = {total_batch_size}") + logger.info(f" Gradient Accumulation steps = {args.gradient_accumulation_steps}") + logger.info(f" Total optimization steps = {args.max_train_steps}") + global_step = 0 + first_epoch = 0 + + # Potentially load in the weights and states from a previous save + if args.resume_from_checkpoint: + if args.resume_from_checkpoint != "latest": + path = os.path.basename(args.resume_from_checkpoint) + else: + # Get the most recent checkpoint + dirs = os.listdir(args.output_dir) + dirs = [d for d in dirs if d.startswith("checkpoint")] + dirs = sorted(dirs, key=lambda x: int(x.split("-")[1])) + path = dirs[-1] if len(dirs) > 0 else None + + if path is None: + accelerator.print( + f"Checkpoint '{args.resume_from_checkpoint}' does not exist. Starting a new training run." + ) + args.resume_from_checkpoint = None + initial_global_step = 0 + else: + global_step = int(path.split("-")[1]) + + initial_global_step = global_step + + pkl_path = os.path.join(os.path.join(args.output_dir, path), "sampler_pos_start.pkl") + if os.path.exists(pkl_path): + with open(pkl_path, 'rb') as file: + _, first_epoch = pickle.load(file) + else: + first_epoch = global_step // num_update_steps_per_epoch + print(f"Load pkl from {pkl_path}. Get first_epoch = {first_epoch}.") + + accelerator.print(f"Resuming from checkpoint {path}") + accelerator.load_state(os.path.join(args.output_dir, path)) + else: + initial_global_step = 0 + + progress_bar = tqdm( + range(0, args.max_train_steps), + initial=initial_global_step, + desc="Steps", + # Only show the progress bar once on each machine. + disable=not accelerator.is_local_main_process, + ) + + if args.multi_stream and args.train_mode != "normal": + # create extra cuda streams to speedup inpaint vae computation + vae_stream_1 = torch.cuda.Stream() + vae_stream_2 = torch.cuda.Stream() + else: + vae_stream_1 = None + vae_stream_2 = None + + # Calculate the index we need + boundary = config['transformer_additional_kwargs'].get('boundary', 0.900) + split_timesteps = args.train_sampling_steps * boundary + differences = torch.abs(noise_scheduler.timesteps - split_timesteps) + closest_index = torch.argmin(differences).item() + print(f"The boundary is {boundary} and the boundary_type is {args.boundary_type}. The closest_index we calculate is {closest_index}") + if args.boundary_type == "high": + start_num_idx = 0 + train_sampling_steps = closest_index + elif args.boundary_type == "low": + start_num_idx = closest_index + train_sampling_steps = args.train_sampling_steps - closest_index + else: + start_num_idx = 0 + train_sampling_steps = args.train_sampling_steps + idx_sampling = DiscreteSampling(train_sampling_steps, start_num_idx=start_num_idx, uniform_sampling=args.uniform_sampling) + + for epoch in range(first_epoch, args.num_train_epochs): + train_loss = 0.0 + batch_sampler.sampler.generator = torch.Generator().manual_seed(args.seed + epoch) + for step, batch in enumerate(train_dataloader): + # Data batch sanity check + if epoch == first_epoch and step == 0: + pixel_values, texts = batch['pixel_values'].cpu(), batch['text'] + control_pixel_values = batch["control_pixel_values"].cpu() + pixel_values = rearrange(pixel_values, "b f c h w -> b c f h w") + control_pixel_values = rearrange(control_pixel_values, "b f c h w -> b c f h w") + os.makedirs(os.path.join(args.output_dir, "sanity_check"), exist_ok=True) + for idx, (pixel_value, control_pixel_value, text) in enumerate(zip(pixel_values, control_pixel_values, texts)): + pixel_value = pixel_value[None, ...] + control_pixel_value = control_pixel_value[None, ...] + gif_name = '-'.join(text.replace('/', '').split()[:10]) if not text == '' else f'{global_step}-{idx}' + save_videos_grid(pixel_value, f"{args.output_dir}/sanity_check/{gif_name[:10]}.gif", rescale=True) + save_videos_grid(control_pixel_value, f"{args.output_dir}/sanity_check/{gif_name[:10]}_control.gif", rescale=True) + + if args.train_mode != "control": + ref_pixel_values = batch["ref_pixel_values"].cpu() + ref_pixel_values = rearrange(ref_pixel_values, "b f c h w -> b c f h w") + for idx, (ref_pixel_value, text) in enumerate(zip(ref_pixel_values, texts)): + ref_pixel_value = ref_pixel_value[None, ...] + gif_name = '-'.join(text.replace('/', '').split()[:10]) if not text == '' else f'{global_step}-{idx}' + save_videos_grid(ref_pixel_value, f"{args.output_dir}/sanity_check/{gif_name[:10]}_ref.gif", rescale=True) + + if args.add_inpaint_info: + clip_pixel_values, mask_pixel_values, texts = batch['clip_pixel_values'].cpu(), batch['mask_pixel_values'].cpu(), batch['text'] + mask_pixel_values = rearrange(mask_pixel_values, "b f c h w -> b c f h w") + for idx, (clip_pixel_value, pixel_value, text) in enumerate(zip(clip_pixel_values, mask_pixel_values, texts)): + pixel_value = pixel_value[None, ...] + Image.fromarray(np.uint8(clip_pixel_value)).save(f"{args.output_dir}/sanity_check/clip_{gif_name[:10] if not text == '' else f'{global_step}-{idx}'}.png") + save_videos_grid(pixel_value, f"{args.output_dir}/sanity_check/mask_{gif_name[:10] if not text == '' else f'{global_step}-{idx}'}.gif", rescale=True) + + with accelerator.accumulate(transformer3d): + # Convert images to latent space + pixel_values = batch["pixel_values"].to(weight_dtype) + control_pixel_values = batch["control_pixel_values"].to(weight_dtype) + if args.train_mode == "control_camera_ref": + control_camera_values = batch["control_camera_values"].to(weight_dtype) + + # Increase the batch size when the length of the latent sequence of the current sample is small + if args.auto_tile_batch_size and args.training_with_video_token_length and zero_stage != 3: + if args.video_sample_n_frames * args.token_sample_size * args.token_sample_size // 16 >= pixel_values.size()[1] * pixel_values.size()[3] * pixel_values.size()[4]: + pixel_values = torch.tile(pixel_values, (4, 1, 1, 1, 1)) + control_pixel_values = torch.tile(control_pixel_values, (4, 1, 1, 1, 1)) + if args.train_mode == "control_camera_ref": + control_camera_values = torch.tile(control_camera_values, (4, 1, 1, 1, 1)) + if args.enable_text_encoder_in_dataloader: + batch['encoder_hidden_states'] = torch.tile(batch['encoder_hidden_states'], (4, 1, 1)) + batch['encoder_attention_mask'] = torch.tile(batch['encoder_attention_mask'], (4, 1)) + else: + batch['text'] = batch['text'] * 4 + elif args.video_sample_n_frames * args.token_sample_size * args.token_sample_size // 4 >= pixel_values.size()[1] * pixel_values.size()[3] * pixel_values.size()[4]: + pixel_values = torch.tile(pixel_values, (2, 1, 1, 1, 1)) + control_pixel_values = torch.tile(control_pixel_values, (2, 1, 1, 1, 1)) + if args.train_mode == "control_camera_ref": + control_camera_values = torch.tile(control_camera_values, (2, 1, 1, 1, 1)) + if args.enable_text_encoder_in_dataloader: + batch['encoder_hidden_states'] = torch.tile(batch['encoder_hidden_states'], (2, 1, 1)) + batch['encoder_attention_mask'] = torch.tile(batch['encoder_attention_mask'], (2, 1)) + else: + batch['text'] = batch['text'] * 2 + + if args.train_mode != "control": + ref_pixel_values = batch["ref_pixel_values"].to(weight_dtype) + clip_idx = batch["clip_idx"] + # Increase the batch size when the length of the latent sequence of the current sample is small + if args.auto_tile_batch_size and args.training_with_video_token_length and zero_stage != 3: + if args.video_sample_n_frames * args.token_sample_size * args.token_sample_size // 16 >= pixel_values.size()[1] * pixel_values.size()[3] * pixel_values.size()[4]: + ref_pixel_values = torch.tile(ref_pixel_values, (4, 1, 1, 1, 1)) + clip_idx = torch.tile(clip_idx, (4,)) + elif args.video_sample_n_frames * args.token_sample_size * args.token_sample_size // 4 >= pixel_values.size()[1] * pixel_values.size()[3] * pixel_values.size()[4]: + ref_pixel_values = torch.tile(ref_pixel_values, (2, 1, 1, 1, 1)) + clip_idx = torch.tile(clip_idx, (2,)) + + if args.add_inpaint_info: + mask_pixel_values = batch["mask_pixel_values"].to(weight_dtype) + mask = batch["mask"].to(weight_dtype) + # Increase the batch size when the length of the latent sequence of the current sample is small + if args.training_with_video_token_length and not zero_stage == 3: + if args.video_sample_n_frames * args.token_sample_size * args.token_sample_size // 16 >= pixel_values.size()[1] * pixel_values.size()[3] * pixel_values.size()[4]: + mask_pixel_values = torch.tile(mask_pixel_values, (4, 1, 1, 1, 1)) + mask = torch.tile(mask, (4, 1, 1, 1, 1)) + elif args.video_sample_n_frames * args.token_sample_size * args.token_sample_size // 4 >= pixel_values.size()[1] * pixel_values.size()[3] * pixel_values.size()[4]: + mask_pixel_values = torch.tile(mask_pixel_values, (2, 1, 1, 1, 1)) + mask = torch.tile(mask, (2, 1, 1, 1, 1)) + + if args.random_frame_crop: + def _create_special_list(length): + if length == 1: + return [1.0] + if length >= 2: + last_element = 0.90 + remaining_sum = 1.0 - last_element + other_elements_value = remaining_sum / (length - 1) + special_list = [other_elements_value] * (length - 1) + [last_element] + return special_list + select_frames = [_tmp for _tmp in list(range(sample_n_frames_bucket_interval + 1, args.video_sample_n_frames + sample_n_frames_bucket_interval, sample_n_frames_bucket_interval))] + select_frames_prob = np.array(_create_special_list(len(select_frames))) + + if len(select_frames) != 0: + if rng is None: + temp_n_frames = np.random.choice(select_frames, p = select_frames_prob) + else: + temp_n_frames = rng.choice(select_frames, p = select_frames_prob) + else: + temp_n_frames = 1 + + # Magvae needs the number of frames to be 4n + 1. + temp_n_frames = (temp_n_frames - 1) // sample_n_frames_bucket_interval + 1 + + pixel_values = pixel_values[:, :temp_n_frames, :, :] + control_pixel_values = control_pixel_values[:, :temp_n_frames, :, :] + + # Keep all node same token length to accelerate the traning when resolution grows. + if args.keep_all_node_same_token_length: + if args.token_sample_size > 256: + numbers_list = list(range(256, args.token_sample_size + 1, 128)) + + if numbers_list[-1] != args.token_sample_size: + numbers_list.append(args.token_sample_size) + else: + numbers_list = [256] + numbers_list = [_number * _number * args.video_sample_n_frames for _number in numbers_list] + + actual_token_length = index_rng.choice(numbers_list) + actual_video_length = (min( + actual_token_length / pixel_values.size()[-1] / pixel_values.size()[-2], args.video_sample_n_frames + ) - 1) // sample_n_frames_bucket_interval * sample_n_frames_bucket_interval + 1 + actual_video_length = int(max(actual_video_length, 1)) + + # Magvae needs the number of frames to be 4n + 1. + actual_video_length = (actual_video_length - 1) // sample_n_frames_bucket_interval + 1 + + pixel_values = pixel_values[:, :actual_video_length, :, :] + control_pixel_values = control_pixel_values[:, :actual_video_length, :, :] + + if args.low_vram: + torch.cuda.empty_cache() + vae.to(accelerator.device) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to("cpu") + + with torch.no_grad(): + # This way is quicker when batch grows up + def _batch_encode_vae(pixel_values): + pixel_values = rearrange(pixel_values, "b f c h w -> b c f h w") + bs = args.vae_mini_batch + new_pixel_values = [] + for i in range(0, pixel_values.shape[0], bs): + pixel_values_bs = pixel_values[i : i + bs] + pixel_values_bs = vae.encode(pixel_values_bs)[0] + pixel_values_bs = pixel_values_bs.sample() + new_pixel_values.append(pixel_values_bs) + return torch.cat(new_pixel_values, dim = 0) + if vae_stream_1 is not None: + vae_stream_1.wait_stream(torch.cuda.current_stream()) + with torch.cuda.stream(vae_stream_1): + latents = _batch_encode_vae(pixel_values) + else: + latents = _batch_encode_vae(pixel_values) + + if args.train_mode != "control_camera_ref": + control_latents = _batch_encode_vae(control_pixel_values) + # Make control latents to zero + for bs_index in range(control_latents.size()[0]): + if rng is None: + zero_init_control_latents_conv_in = np.random.choice([0, 1], p = [0.90, 0.10]) + else: + zero_init_control_latents_conv_in = rng.choice([0, 1], p = [0.90, 0.10]) + + if zero_init_control_latents_conv_in: + control_latents[bs_index] = control_latents[bs_index] * 0 + control_camera_latents = None + else: + control_latents = None + control_camera_latents = rearrange(control_camera_values, "b f c h w -> b c f h w") + control_camera_latents = torch.concat( + [ + torch.repeat_interleave(control_camera_latents[:, :, 0:1], repeats=4, dim=2), + control_camera_latents[:, :, 1:] + ], dim=2 + ).transpose(1, 2).contiguous() + control_camera_latents = control_camera_latents.view(control_camera_latents.shape[0], control_camera_latents.shape[1] // 4, 4, control_camera_latents.shape[2], control_camera_latents.shape[3], control_camera_latents.shape[4]) + control_camera_latents = control_camera_latents.transpose(2, 3).contiguous() + control_camera_latents = control_camera_latents.view(control_camera_latents.shape[0], control_camera_latents.shape[1], control_camera_latents.shape[2] * 4, control_camera_latents.shape[4], control_camera_latents.shape[5]) + control_camera_latents = control_camera_latents.transpose(1, 2) + + if args.train_mode != "control": + ref_latents = _batch_encode_vae(ref_pixel_values) + if args.add_full_ref_image_in_self_attention: + full_ref = ref_latents[:, :, 0].clone() + + ref_latents_conv_in = torch.zeros_like(latents).to(ref_latents.device, ref_latents.dtype) + ref_latents_conv_in[:, :, :1] = ref_latents + for bs_index in range(ref_latents.size()[0]): + if rng is None: + zero_init_ref_latents_conv_in = np.random.choice([0, 1], p = [0.90, 0.10]) + else: + zero_init_ref_latents_conv_in = rng.choice([0, 1], p = [0.90, 0.10]) + + if clip_idx[bs_index] != 0 or (zero_init_ref_latents_conv_in and latents.size()[1] != 1): + ref_latents_conv_in[bs_index, :, :1] = ref_latents_conv_in[bs_index, :, :1] * 0 + + if args.add_full_ref_image_in_self_attention: + if rng is None: + zero_init_full_ref_conv_in = np.random.choice([0, 1], p = [0.90, 0.10]) + else: + zero_init_full_ref_conv_in = rng.choice([0, 1], p = [0.90, 0.10]) + if clip_idx[bs_index] == 0 or zero_init_full_ref_conv_in: + full_ref[bs_index] = full_ref[bs_index] * 0 + + if args.add_inpaint_info: + t2v_flag = [(_mask == 1).all() for _mask in mask] + new_t2v_flag = [] + for _mask in t2v_flag: + if _mask and np.random.rand() < 0.90: + new_t2v_flag.append(0) + else: + new_t2v_flag.append(1) + t2v_flag = torch.from_numpy(np.array(new_t2v_flag)).to(accelerator.device, dtype=weight_dtype) + + mask = rearrange(mask, "b f c h w -> b c f h w") + mask = torch.concat( + [ + torch.repeat_interleave(mask[:, :, 0:1], repeats=4, dim=2), + mask[:, :, 1:] + ], dim=2 + ) + mask = mask.view(mask.shape[0], mask.shape[2] // 4, 4, mask.shape[3], mask.shape[4]) + mask = mask.transpose(1, 2) + mask = resize_mask(1 - mask, latents) + + # Encode inpaint latents. + mask_latents = _batch_encode_vae(mask_pixel_values) + + inpaint_latents = torch.concat([mask, mask_latents], dim=1) + inpaint_latents = t2v_flag[:, None, None, None, None] * inpaint_latents + else: + inpaint_latents = None + + if control_latents is None: + if inpaint_latents is None: + control_latents = ref_latents_conv_in + else: + control_latents = inpaint_latents + else: + if inpaint_latents is None: + control_latents = torch.cat([control_latents, ref_latents_conv_in], dim = 1) + else: + control_latents = torch.cat([control_latents, inpaint_latents], dim = 1) + + # wait for latents = vae.encode(pixel_values) to complete + if vae_stream_1 is not None: + torch.cuda.current_stream().wait_stream(vae_stream_1) + + if args.low_vram: + vae.to('cpu') + torch.cuda.empty_cache() + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device) + + if args.enable_text_encoder_in_dataloader: + prompt_embeds = batch['encoder_hidden_states'].to(device=latents.device) + else: + with torch.no_grad(): + prompt_ids = tokenizer( + batch['text'], + padding="max_length", + max_length=args.tokenizer_max_length, + truncation=True, + add_special_tokens=True, + return_tensors="pt" + ) + text_input_ids = prompt_ids.input_ids + prompt_attention_mask = prompt_ids.attention_mask + + seq_lens = prompt_attention_mask.gt(0).sum(dim=1).long() + prompt_embeds = text_encoder(text_input_ids.to(latents.device), attention_mask=prompt_attention_mask.to(latents.device))[0] + prompt_embeds = [u[:v] for u, v in zip(prompt_embeds, seq_lens)] + + if args.low_vram and not args.enable_text_encoder_in_dataloader: + text_encoder.to('cpu') + torch.cuda.empty_cache() + + bsz, channel, num_frames, height, width = latents.size() + noise = torch.randn(latents.size(), device=latents.device, generator=torch_rng, dtype=weight_dtype) + + if not args.uniform_sampling: + u = compute_density_for_timestep_sampling( + weighting_scheme=args.weighting_scheme, + batch_size=bsz, + logit_mean=args.logit_mean, + logit_std=args.logit_std, + mode_scale=args.mode_scale, + ) + indices = (u * noise_scheduler.config.num_train_timesteps).long() + else: + # Sample a random timestep for each image + # timesteps = generate_timestep_with_lognorm(0, args.train_sampling_steps, (bsz,), device=latents.device, generator=torch_rng) + # timesteps = torch.randint(0, args.train_sampling_steps, (bsz,), device=latents.device, generator=torch_rng) + indices = idx_sampling(bsz, generator=torch_rng, device=latents.device) + indices = indices.long().cpu() + timesteps = noise_scheduler.timesteps[indices].to(device=latents.device) + + def get_sigmas(timesteps, n_dim=4, dtype=torch.float32): + sigmas = noise_scheduler.sigmas.to(device=accelerator.device, dtype=dtype) + schedule_timesteps = noise_scheduler.timesteps.to(accelerator.device) + timesteps = timesteps.to(accelerator.device) + step_indices = [(schedule_timesteps == t).nonzero().item() for t in timesteps] + + sigma = sigmas[step_indices].flatten() + while len(sigma.shape) < n_dim: + sigma = sigma.unsqueeze(-1) + return sigma + + # Add noise according to flow matching. + # zt = (1 - texp) * x + texp * z1 + sigmas = get_sigmas(timesteps, n_dim=latents.ndim, dtype=latents.dtype) + noisy_latents = (1.0 - sigmas) * latents + sigmas * noise + + # Add noise + target = noise - latents + + target_shape = (vae.latent_channels, num_frames, width, height) + seq_len = math.ceil( + (target_shape[2] * target_shape[3]) / + (accelerator.unwrap_model(transformer3d).config.patch_size[1] * accelerator.unwrap_model(transformer3d).config.patch_size[2]) * + target_shape[1] + ) + + # Predict the noise residual + with torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=accelerator.device): + noise_pred = transformer3d( + x=noisy_latents, + context=prompt_embeds, + t=timesteps, + seq_len=seq_len, + y=control_latents if args.train_mode != "control" else None, + y_camera=control_camera_latents if args.train_mode == "control_camera_ref" else None, + full_ref=full_ref if args.add_full_ref_image_in_self_attention else None, + ) + + def custom_mse_loss(noise_pred, target, weighting=None, threshold=50): + noise_pred = noise_pred.float() + target = target.float() + diff = noise_pred - target + mse_loss = F.mse_loss(noise_pred, target, reduction='none') + mask = (diff.abs() <= threshold).float() + masked_loss = mse_loss * mask + if weighting is not None: + masked_loss = masked_loss * weighting + final_loss = masked_loss.mean() + return final_loss + + weighting = compute_loss_weighting_for_sd3(weighting_scheme=args.weighting_scheme, sigmas=sigmas) + loss = custom_mse_loss(noise_pred.float(), target.float(), weighting.float()) + loss = loss.mean() + + if args.motion_sub_loss and noise_pred.size()[1] > 2: + gt_sub_noise = noise_pred[:, 1:, :].float() - noise_pred[:, :-1, :].float() + pre_sub_noise = target[:, 1:, :].float() - target[:, :-1, :].float() + sub_loss = F.mse_loss(gt_sub_noise, pre_sub_noise, reduction="mean") + loss = loss * (1 - args.motion_sub_loss_ratio) + sub_loss * args.motion_sub_loss_ratio + + # Gather the losses across all processes for logging (if we use distributed training). + avg_loss = accelerator.gather(loss.repeat(args.train_batch_size)).mean() + train_loss += avg_loss.item() / args.gradient_accumulation_steps + + # Backpropagate + accelerator.backward(loss) + if accelerator.sync_gradients: + if not args.use_deepspeed and not args.use_fsdp: + trainable_params_grads = [p.grad for p in trainable_params if p.grad is not None] + trainable_params_total_norm = torch.norm(torch.stack([torch.norm(g.detach(), 2) for g in trainable_params_grads]), 2) + max_grad_norm = linear_decay(args.max_grad_norm * args.initial_grad_norm_ratio, args.max_grad_norm, args.abnormal_norm_clip_start, global_step) + if trainable_params_total_norm / max_grad_norm > 5 and global_step > args.abnormal_norm_clip_start: + actual_max_grad_norm = max_grad_norm / min((trainable_params_total_norm / max_grad_norm), 10) + else: + actual_max_grad_norm = max_grad_norm + else: + actual_max_grad_norm = args.max_grad_norm + + if not args.use_deepspeed and not args.use_fsdp and args.report_model_info and accelerator.is_main_process: + if trainable_params_total_norm > 1 and global_step > args.abnormal_norm_clip_start: + for name, param in transformer3d.named_parameters(): + if param.requires_grad: + writer.add_scalar(f'gradients/before_clip_norm/{name}', param.grad.norm(), global_step=global_step) + + norm_sum = accelerator.clip_grad_norm_(trainable_params, actual_max_grad_norm) + if not args.use_deepspeed and not args.use_fsdp and args.report_model_info and accelerator.is_main_process: + writer.add_scalar(f'gradients/norm_sum', norm_sum, global_step=global_step) + writer.add_scalar(f'gradients/actual_max_grad_norm', actual_max_grad_norm, global_step=global_step) + optimizer.step() + lr_scheduler.step() + optimizer.zero_grad() + + # Checks if the accelerator has performed an optimization step behind the scenes + if accelerator.sync_gradients: + + if args.use_ema: + ema_transformer3d.step(transformer3d.parameters()) + progress_bar.update(1) + global_step += 1 + accelerator.log({"train_loss": train_loss}, step=global_step) + train_loss = 0.0 + + if global_step % args.checkpointing_steps == 0: + if args.use_deepspeed or args.use_fsdp or accelerator.is_main_process: + # _before_ saving state, check if this save would set us over the `checkpoints_total_limit` + if args.checkpoints_total_limit is not None: + checkpoints = os.listdir(args.output_dir) + checkpoints = [d for d in checkpoints if d.startswith("checkpoint")] + checkpoints = sorted(checkpoints, key=lambda x: int(x.split("-")[1])) + + # before we save the new checkpoint, we need to have at _most_ `checkpoints_total_limit - 1` checkpoints + if len(checkpoints) >= args.checkpoints_total_limit: + num_to_remove = len(checkpoints) - args.checkpoints_total_limit + 1 + removing_checkpoints = checkpoints[0:num_to_remove] + + logger.info( + f"{len(checkpoints)} checkpoints already exist, removing {len(removing_checkpoints)} checkpoints" + ) + logger.info(f"removing checkpoints: {', '.join(removing_checkpoints)}") + + for removing_checkpoint in removing_checkpoints: + removing_checkpoint = os.path.join(args.output_dir, removing_checkpoint) + shutil.rmtree(removing_checkpoint) + + save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}") + accelerator.save_state(save_path) + logger.info(f"Saved state to {save_path}") + + if accelerator.is_main_process: + if args.validation_prompts is not None and global_step % args.validation_steps == 0: + if args.use_ema: + # Store the UNet parameters temporarily and load the EMA parameters to perform inference. + ema_transformer3d.store(transformer3d.parameters()) + ema_transformer3d.copy_to(transformer3d.parameters()) + log_validation( + vae, + text_encoder, + tokenizer, + transformer3d, + args, + config, + accelerator, + weight_dtype, + global_step, + ) + if args.use_ema: + # Switch back to the original transformer3d parameters. + ema_transformer3d.restore(transformer3d.parameters()) + + logs = {"step_loss": loss.detach().item(), "lr": lr_scheduler.get_last_lr()[0]} + progress_bar.set_postfix(**logs) + + if global_step >= args.max_train_steps: + break + + if accelerator.is_main_process: + if args.validation_prompts is not None and epoch % args.validation_epochs == 0: + if args.use_ema: + # Store the UNet parameters temporarily and load the EMA parameters to perform inference. + ema_transformer3d.store(transformer3d.parameters()) + ema_transformer3d.copy_to(transformer3d.parameters()) + log_validation( + vae, + text_encoder, + tokenizer, + transformer3d, + args, + config, + accelerator, + weight_dtype, + global_step, + ) + if args.use_ema: + # Switch back to the original transformer3d parameters. + ema_transformer3d.restore(transformer3d.parameters()) + + # Create the pipeline using the trained modules and save it. + accelerator.wait_for_everyone() + if accelerator.is_main_process: + transformer3d = unwrap_model(transformer3d) + if args.use_ema: + ema_transformer3d.copy_to(transformer3d.parameters()) + + if args.use_deepspeed or args.use_fsdp or accelerator.is_main_process: + save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}") + accelerator.save_state(save_path) + logger.info(f"Saved state to {save_path}") + + accelerator.end_training() + + +if __name__ == "__main__": + main() diff --git a/scripts/wan2.2_fun/train_control.sh b/scripts/wan2.2_fun/train_control.sh new file mode 100644 index 0000000..89abfbd --- /dev/null +++ b/scripts/wan2.2_fun/train_control.sh @@ -0,0 +1,45 @@ +export MODEL_NAME="models/Diffusion_Transformer/Wan2.2-Fun-A14B-Control" +export DATASET_NAME="datasets/internal_datasets/" +export DATASET_META_NAME="datasets/internal_datasets/metadata.json" +# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA. +# export NCCL_IB_DISABLE=1 +# export NCCL_P2P_DISABLE=1 +NCCL_DEBUG=INFO + +accelerate launch --mixed_precision="bf16" scripts/wan2.2_fun/train_control.py \ + --config_path="config/wan2.2/wan_civitai_i2v.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 \ + --train_mode="control_ref" \ + --control_ref_image="random" \ + --add_inpaint_info \ + --add_full_ref_image_in_self_attention \ + --low_vram \ + --trainable_modules "." \ No newline at end of file diff --git a/scripts/wan2.2_fun/train_lora.sh b/scripts/wan2.2_fun/train_lora.sh index fc3e06f..5c8464e 100644 --- a/scripts/wan2.2_fun/train_lora.sh +++ b/scripts/wan2.2_fun/train_lora.sh @@ -38,4 +38,5 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.2_fun/train_lora.py \ --train_mode="inpaint" \ --boundary_type="low" \ --lora_skip_name="ffn" \ + --boundary_type="low" \ --low_vram diff --git a/videox_fun/pipeline/pipeline_wan2_2_fun_control.py b/videox_fun/pipeline/pipeline_wan2_2_fun_control.py index 9bb832b..0041438 100644 --- a/videox_fun/pipeline/pipeline_wan2_2_fun_control.py +++ b/videox_fun/pipeline/pipeline_wan2_2_fun_control.py @@ -158,7 +158,7 @@ class Wan2_2FunControlPipeline(DiffusionPipeline): """ _optional_components = ["transformer_2"] - model_cpu_offload_seq = "text_encoder->transformer->transformer_2->vae" + model_cpu_offload_seq = "text_encoder->transformer_2->transformer->vae" _callback_tensor_inputs = [ "latents", diff --git a/videox_fun/pipeline/pipeline_wan2_2_ti2v.py b/videox_fun/pipeline/pipeline_wan2_2_ti2v.py index 24e90ed..12b3bd8 100644 --- a/videox_fun/pipeline/pipeline_wan2_2_ti2v.py +++ b/videox_fun/pipeline/pipeline_wan2_2_ti2v.py @@ -601,7 +601,7 @@ class Wan2_2TI2VPipeline(DiffusionPipeline): pbar.update(1) # Prepare mask latent variables - if init_video is not None: + if init_video is not None and not (mask_video == 255).all(): bs, _, video_length, height, width = video.size() mask_condition = self.mask_processor.preprocess(rearrange(mask_video, "b c f h w -> (b f) c h w"), height=height, width=width) mask_condition = mask_condition.to(dtype=torch.float32) @@ -632,6 +632,8 @@ class Wan2_2TI2VPipeline(DiffusionPipeline): mask = F.interpolate(mask_condition[:, :1], size=latents.size()[-3:], mode='trilinear', align_corners=True).to(device, weight_dtype) latents = (1 - mask) * masked_video_latents + mask * latents + else: + init_video = None if comfyui_progressbar: pbar.update(1) diff --git a/videox_fun/ui/controller.py b/videox_fun/ui/controller.py index 52f763c..6a540d6 100755 --- a/videox_fun/ui/controller.py +++ b/videox_fun/ui/controller.py @@ -256,6 +256,7 @@ class Fun_Controller: validation_video, control_video, ): + spatial_compression_ratio = self.vae.config.spatial_compression_ratio if hasattr(self.vae.config, "spatial_compression_ratio") else 8 aspect_ratio_sample_size = {key : [x / 512 * base_resolution for x in ASPECT_RATIO_512[key]] for key in ASPECT_RATIO_512.keys()} if self.model_type == "Inpaint": if validation_video is not None: @@ -265,7 +266,7 @@ class Fun_Controller: else: original_width, original_height = Image.fromarray(cv2.VideoCapture(control_video).read()[1]).size closest_size, closest_ratio = get_closest_ratio(original_height, original_width, ratios=aspect_ratio_sample_size) - height_slider, width_slider = [int(x / 16) * 16 for x in closest_size] + height_slider, width_slider = [int(x / spatial_compression_ratio / 2) * spatial_compression_ratio * 2 for x in closest_size] return height_slider, width_slider def save_outputs(self, is_image, length_slider, sample, fps): diff --git a/videox_fun/ui/wan2_2_fun_ui.py b/videox_fun/ui/wan2_2_fun_ui.py new file mode 100644 index 0000000..f526e33 --- /dev/null +++ b/videox_fun/ui/wan2_2_fun_ui.py @@ -0,0 +1,803 @@ +"""Modified from https://github.com/guoyww/AnimateDiff/blob/main/app.py +""" +import os +import random + +import cv2 +import gradio as gr +import numpy as np +import torch +from omegaconf import OmegaConf +from PIL import Image +from safetensors import safe_open + +from ..data.bucket_sampler import ASPECT_RATIO_512, get_closest_ratio +from ..dist import set_multi_gpus_devices, shard_model +from ..models import (AutoencoderKLWan, AutoencoderKLWan3_8, AutoTokenizer, + CLIPModel, Wan2_2Transformer3DModel, WanT5EncoderModel) +from ..models.cache_utils import get_teacache_coefficients +from ..pipeline import Wan2_2FunControlPipeline, Wan2_2FunPipeline, Wan2_2FunInpaintPipeline +from ..utils.fp8_optimization import (convert_model_weight_to_float8, + convert_weight_dtype_wrapper, + replace_parameters_by_name) +from ..utils.lora_utils import merge_lora, unmerge_lora +from ..utils.utils import (filter_kwargs, get_image_latent, + get_image_to_video_latent, + get_video_to_video_latent, save_videos_grid, timer) +from .controller import (Fun_Controller, Fun_Controller_Client, + all_cheduler_dict, css, ddpm_scheduler_dict, + flow_scheduler_dict, gradio_version, + gradio_version_is_above_4) +from .ui import (create_cfg_and_seedbox, create_cfg_riflex_k, + create_cfg_skip_params, create_config, + create_fake_finetune_models_checkpoints, + create_fake_height_width, create_fake_model_checkpoints, + create_fake_model_type, create_finetune_models_checkpoints, + create_generation_method, + create_generation_methods_and_video_length, + create_height_width, create_model_checkpoints, + create_model_type, create_prompts, create_samplers, + create_teacache_params, create_ui_outputs) + + +class Wan2_2_Fun_Controller(Fun_Controller): + def update_diffusion_transformer(self, diffusion_transformer_dropdown): + print(f"Update diffusion transformer: {diffusion_transformer_dropdown}") + self.model_name = diffusion_transformer_dropdown + self.diffusion_transformer_dropdown = diffusion_transformer_dropdown + if diffusion_transformer_dropdown == "none": + return gr.update() + Choosen_AutoencoderKL = { + "AutoencoderKLWan": AutoencoderKLWan, + "AutoencoderKLWan3_8": AutoencoderKLWan3_8 + }[self.config['vae_kwargs'].get('vae_type', 'AutoencoderKLWan')] + self.vae = Choosen_AutoencoderKL.from_pretrained( + os.path.join(diffusion_transformer_dropdown, self.config['vae_kwargs'].get('vae_subpath', 'vae')), + additional_kwargs=OmegaConf.to_container(self.config['vae_kwargs']), + ).to(self.weight_dtype) + + # Get Transformer + self.transformer = Wan2_2Transformer3DModel.from_pretrained( + os.path.join(diffusion_transformer_dropdown, self.config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer')), + transformer_additional_kwargs=OmegaConf.to_container(self.config['transformer_additional_kwargs']), + low_cpu_mem_usage=True, + torch_dtype=self.weight_dtype, + ) + if self.config['transformer_additional_kwargs'].get('transformer_combination_type', 'single') == "moe": + self.transformer_2 = Wan2_2Transformer3DModel.from_pretrained( + os.path.join(diffusion_transformer_dropdown, self.config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer')), + transformer_additional_kwargs=OmegaConf.to_container(self.config['transformer_additional_kwargs']), + low_cpu_mem_usage=True, + torch_dtype=self.weight_dtype, + ) + else: + self.transformer_2 = None + + # Get Tokenizer + self.tokenizer = AutoTokenizer.from_pretrained( + os.path.join(diffusion_transformer_dropdown, self.config['text_encoder_kwargs'].get('tokenizer_subpath', 'tokenizer')), + ) + + # Get Text encoder + self.text_encoder = WanT5EncoderModel.from_pretrained( + os.path.join(diffusion_transformer_dropdown, self.config['text_encoder_kwargs'].get('text_encoder_subpath', 'text_encoder')), + additional_kwargs=OmegaConf.to_container(self.config['text_encoder_kwargs']), + low_cpu_mem_usage=True, + torch_dtype=self.weight_dtype, + ) + self.text_encoder = self.text_encoder.eval() + + Choosen_Scheduler = self.scheduler_dict[list(self.scheduler_dict.keys())[0]] + self.scheduler = Choosen_Scheduler( + **filter_kwargs(Choosen_Scheduler, OmegaConf.to_container(self.config['scheduler_kwargs'])) + ) + + # Get pipeline + if self.model_type == "Inpaint": + if self.transformer.config.in_channels != self.vae.config.latent_channels: + self.pipeline = Wan2_2FunInpaintPipeline( + vae=self.vae, + tokenizer=self.tokenizer, + text_encoder=self.text_encoder, + transformer=self.transformer, + transformer_2=self.transformer_2, + scheduler=self.scheduler, + ) + else: + self.pipeline = Wan2_2FunPipeline( + vae=self.vae, + tokenizer=self.tokenizer, + text_encoder=self.text_encoder, + transformer=self.transformer, + transformer_2=self.transformer_2, + scheduler=self.scheduler, + ) + else: + self.pipeline = Wan2_2FunControlPipeline( + vae=self.vae, + tokenizer=self.tokenizer, + text_encoder=self.text_encoder, + transformer=self.transformer, + transformer_2=self.transformer_2, + scheduler=self.scheduler, + ) + + if self.ulysses_degree > 1 or self.ring_degree > 1: + from functools import partial + self.transformer.enable_multi_gpus_inference() + if self.transformer_2 is not None: + self.transformer_2.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) + if self.transformer_2 is not None: + self.pipeline.transformer_2 = shard_fn(self.pipeline.transformer_2) + 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]) + if self.transformer_2 is not None: + for i in range(len(self.pipeline.transformer_2.blocks)): + self.pipeline.transformer_2.blocks[i] = torch.compile(self.pipeline.transformer_2.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) + if self.transformer_2 is not None: + replace_parameters_by_name(self.transformer_2, ["modulation",], device=self.device) + self.transformer_2.freqs = self.transformer_2.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",], device=self.device) + convert_weight_dtype_wrapper(self.transformer, self.weight_dtype) + if self.transformer_2 is not None: + convert_model_weight_to_float8(self.transformer_2, exclude_module_name=["modulation",], device=self.device) + convert_weight_dtype_wrapper(self.transformer_2, 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",], device=self.device) + convert_weight_dtype_wrapper(self.transformer, self.weight_dtype) + if self.transformer_2 is not None: + convert_model_weight_to_float8(self.transformer_2, exclude_module_name=["modulation",], device=self.device) + convert_weight_dtype_wrapper(self.transformer_2, self.weight_dtype) + self.pipeline.to(self.device) + else: + self.pipeline.to(self.device) + print("Update diffusion transformer done") + return gr.update() + + @timer + def generate( + self, + diffusion_transformer_dropdown, + base_model_dropdown, + lora_model_dropdown, + lora_alpha_slider, + prompt_textbox, + negative_prompt_textbox, + sampler_dropdown, + sample_step_slider, + resize_method, + width_slider, + height_slider, + base_resolution, + generation_method, + length_slider, + overlap_video_length, + partial_video_length, + cfg_scale_slider, + start_image, + end_image, + validation_video, + validation_video_mask, + control_video, + denoise_strength, + seed_textbox, + ref_image = None, + enable_teacache = None, + teacache_threshold = None, + num_skip_start_steps = None, + teacache_offload = None, + cfg_skip_ratio = None, + enable_riflex = None, + riflex_k = None, + base_model_2_dropdown=None, + lora_model_2_dropdown=None, + fps = None, + is_api = False, + ): + self.clear_cache() + + print(f"Input checking.") + _, comment = self.input_check( + resize_method, generation_method, start_image, end_image, validation_video,control_video, is_api + ) + print(f"Input checking down") + if comment != "OK": + return "", comment + is_image = True if generation_method == "Image Generation" else False + + if self.base_model_path != base_model_dropdown: + self.update_base_model(base_model_dropdown) + if self.base_model_2_path != base_model_2_dropdown: + self.update_lora_model(base_model_2_dropdown, is_checkpoint_2=True) + + if self.lora_model_path != lora_model_dropdown: + self.update_lora_model(lora_model_dropdown) + if self.lora_model_2_path != lora_model_2_dropdown: + self.update_lora_model(lora_model_2_dropdown, is_checkpoint_2=True) + + print(f"Load scheduler.") + scheduler_config = self.pipeline.scheduler.config + if sampler_dropdown == "Flow_Unipc" or sampler_dropdown == "Flow_DPM++": + scheduler_config['shift'] = 1 + self.pipeline.scheduler = self.scheduler_dict[sampler_dropdown].from_config(scheduler_config) + print(f"Load scheduler down.") + + if resize_method == "Resize according to Reference": + print(f"Calculate height and width according to Reference.") + height_slider, width_slider = self.get_height_width_from_reference( + base_resolution, start_image, validation_video, control_video, + ) + + if self.lora_model_path != "none": + print(f"Merge Lora.") + self.pipeline = merge_lora(self.pipeline, self.lora_model_path, multiplier=lora_alpha_slider) + if self.transformer_2 is not None: + self.pipeline = merge_lora(self.pipeline, self.lora_model_2_path, multiplier=lora_alpha_slider, sub_transformer_name="transformer_2") + print(f"Merge Lora done.") + + coefficients = get_teacache_coefficients(self.diffusion_transformer_dropdown) if enable_teacache else None + if coefficients is not None: + print(f"Enable TeaCache with threshold {teacache_threshold} and skip the first {num_skip_start_steps} steps.") + self.pipeline.transformer.enable_teacache( + coefficients, sample_step_slider, teacache_threshold, num_skip_start_steps=num_skip_start_steps, offload=teacache_offload + ) + if self.transformer_2 is not None: + self.pipeline.transformer_2.share_teacache(self.pipeline.transformer) + else: + print(f"Disable TeaCache.") + self.pipeline.transformer.disable_teacache() + if self.transformer_2 is not None: + self.pipeline.transformer_2.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) + if self.transformer_2 is not None: + self.pipeline.transformer_2.share_cfg_skip(self.pipeline.transformer) + + 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) + generator = torch.Generator(device=self.device).manual_seed(int(seed_textbox)) + print(f"Generate seed done.") + + if fps is None: + fps = 16 + boundary = self.config['transformer_additional_kwargs'].get('boundary', 0.875) + + if enable_riflex: + print(f"Enable riflex") + latent_frames = (int(length_slider) - 1) // self.vae.config.temporal_compression_ratio + 1 + self.pipeline.transformer.enable_riflex(k = riflex_k, L_test = latent_frames if not is_image else 1) + if self.transformer_2 is not None: + self.pipeline.transformer_2.enable_riflex(k = riflex_k, L_test = latent_frames if not is_image else 1) + + try: + print(f"Generation.") + if self.model_type == "Inpaint": + if self.transformer.config.in_channels != self.vae.config.latent_channels: + if validation_video is not None: + input_video, input_video_mask, _, clip_image = get_video_to_video_latent(validation_video, length_slider if not is_image else 1, sample_size=(height_slider, width_slider), validation_video_mask=validation_video_mask, fps=fps) + else: + input_video, input_video_mask, clip_image = get_image_to_video_latent(start_image, end_image, length_slider if not is_image else 1, sample_size=(height_slider, width_slider)) + + sample = self.pipeline( + prompt_textbox, + negative_prompt = negative_prompt_textbox, + num_inference_steps = sample_step_slider, + guidance_scale = cfg_scale_slider, + width = width_slider, + height = height_slider, + num_frames = length_slider if not is_image else 1, + generator = generator, + + video = input_video, + mask_video = input_video_mask, + boundary = boundary + ).videos + else: + sample = self.pipeline( + prompt_textbox, + negative_prompt = negative_prompt_textbox, + num_inference_steps = sample_step_slider, + guidance_scale = cfg_scale_slider, + width = width_slider, + height = height_slider, + num_frames = length_slider if not is_image else 1, + generator = generator, + boundary = boundary + ).videos + else: + inpaint_video, inpaint_video_mask, clip_image = get_image_to_video_latent(start_image, end_image, video_length=length_slider if not is_image else 1, sample_size=(height_slider, width_slider)) + + if ref_image is not None: + ref_image = get_image_latent(ref_image, sample_size=(height_slider, width_slider)) + + input_video, input_video_mask, _, _ = get_video_to_video_latent(control_video, video_length=length_slider if not is_image else 1, sample_size=(height_slider, width_slider), fps=fps, ref_image=None) + + sample = self.pipeline( + prompt_textbox, + negative_prompt = negative_prompt_textbox, + num_inference_steps = sample_step_slider, + guidance_scale = cfg_scale_slider, + width = width_slider, + height = height_slider, + num_frames = length_slider if not is_image else 1, + generator = generator, + + video = inpaint_video, + mask_video = inpaint_video_mask, + control_video = input_video, + ref_image = ref_image, + boundary = boundary, + ).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) + if is_api: + return "", f"Error. error information is {str(e)}" + else: + return gr.update(), gr.update(), f"Error. error information is {str(e)}" + + self.clear_cache() + # lora part + if self.lora_model_path != "none": + print(f"Unmerge Lora.") + self.pipeline = unmerge_lora(self.pipeline, self.lora_model_path, multiplier=lora_alpha_slider) + print(f"Unmerge Lora done.") + + print(f"Saving outputs.") + save_sample_path = self.save_outputs( + is_image, length_slider, sample, fps=fps + ) + print(f"Saving outputs done.") + + if is_image or length_slider == 1: + if is_api: + return save_sample_path, "Success" + else: + if gradio_version_is_above_4: + return gr.Image(value=save_sample_path, visible=True), gr.Video(value=None, visible=False), "Success" + else: + return gr.Image.update(value=save_sample_path, visible=True), gr.Video.update(value=None, visible=False), "Success" + else: + if is_api: + return save_sample_path, "Success" + else: + if gradio_version_is_above_4: + return gr.Image(visible=False, value=None), gr.Video(value=save_sample_path, visible=True), "Success" + else: + return gr.Image.update(visible=False, value=None), gr.Video.update(value=save_sample_path, visible=True), "Success" + +Wan2_2_Fun_Controller_Host = Wan2_2_Fun_Controller +Wan2_2_Fun_Controller_Client = Fun_Controller_Client + +def ui(GPU_memory_mode, scheduler_dict, config_path, compile_dit, weight_dtype, savedir_sample=None): + controller = Wan2_2_Fun_Controller( + GPU_memory_mode, scheduler_dict, model_name=None, model_type="Inpaint", + config_path=config_path, compile_dit=compile_dit, + weight_dtype=weight_dtype, savedir_sample=savedir_sample, + ) + + with gr.Blocks(css=css) as demo: + gr.Markdown( + """ + # Wan2.2-Fun: + + A Wan with more flexible generation conditions, capable of producing videos of different resolutions, around 5 seconds, and fps 16 (frames 1 to 81), as well as image generated videos. + + [Github](https://github.com/aigc-apps/VideoX-Fun/) + """ + ) + with gr.Column(variant="panel"): + config_dropdown, config_refresh_button = create_config(controller) + model_type = create_model_type(visible=True) + diffusion_transformer_dropdown, diffusion_transformer_refresh_button = \ + create_model_checkpoints(controller, visible=True) + base_model_dropdown, lora_model_dropdown, lora_alpha_slider, personalized_refresh_button = \ + create_finetune_models_checkpoints(controller, visible=True, add_checkpoint_2=True) + base_model_dropdown, base_model_2_dropdown = base_model_dropdown + lora_model_dropdown, lora_model_2_dropdown = lora_model_dropdown + + with gr.Row(): + enable_teacache, teacache_threshold, num_skip_start_steps, teacache_offload = \ + create_teacache_params(True, 0.10, 1, False) + cfg_skip_ratio = create_cfg_skip_params(0) + enable_riflex, riflex_k = create_cfg_riflex_k(False, 6) + + with gr.Column(variant="panel"): + prompt_textbox, negative_prompt_textbox = create_prompts(negative_prompt="色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走") + + with gr.Row(): + with gr.Column(): + sampler_dropdown, sample_step_slider = create_samplers(controller) + + resize_method, width_slider, height_slider, base_resolution = create_height_width( + default_height = 480, default_width = 832, maximum_height = 1344, + maximum_width = 1344, + ) + generation_method, length_slider, overlap_video_length, partial_video_length = \ + create_generation_methods_and_video_length( + ["Video Generation", "Image Generation"], + default_video_length=81, + maximum_video_length=161, + ) + 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 Control (视频控制)"], prompt_textbox, support_ref_image=True + ) + cfg_scale_slider, seed_textbox, seed_button = create_cfg_and_seedbox(gradio_version_is_above_4) + + generate_button = gr.Button(value="Generate (生成)", variant='primary') + + result_image, result_video, infer_progress = create_ui_outputs() + + config_dropdown.change( + fn=controller.update_config, + inputs=[config_dropdown], + outputs=[] + ) + + model_type.change( + fn=controller.update_model_type, + inputs=[model_type], + outputs=[] + ) + + def upload_generation_method(generation_method): + if generation_method == "Video Generation": + return [gr.update(visible=True, maximum=161, value=81, interactive=True), gr.update(visible=False), gr.update(visible=False)] + elif generation_method == "Image Generation": + return [gr.update(minimum=1, maximum=1, value=1, interactive=False), gr.update(visible=False), gr.update(visible=False)] + else: + return [gr.update(visible=True, maximum=1344), gr.update(visible=True), gr.update(visible=True)] + generation_method.change( + upload_generation_method, generation_method, [length_slider, overlap_video_length, partial_video_length] + ) + + def upload_source_method(source_method): + if source_method == "Text to Video (文本到视频)": + return [gr.update(visible=False), gr.update(visible=False), gr.update(visible=False), gr.update(value=None), gr.update(value=None), gr.update(value=None), gr.update(value=None), gr.update(value=None)] + elif source_method == "Image to Video (图片到视频)": + return [gr.update(visible=True), gr.update(visible=False), gr.update(visible=False), gr.update(), gr.update(), gr.update(value=None), gr.update(value=None), gr.update(value=None)] + elif source_method == "Video to Video (视频到视频)": + return [gr.update(visible=False), gr.update(visible=True), gr.update(visible=False), gr.update(value=None), gr.update(value=None), gr.update(), gr.update(), gr.update(value=None)] + else: + return [gr.update(visible=False), gr.update(visible=False), gr.update(visible=True), gr.update(value=None), gr.update(value=None), gr.update(value=None), gr.update(value=None), gr.update()] + source_method.change( + upload_source_method, source_method, [ + image_to_video_col, video_to_video_col, control_video_col, start_image, end_image, + validation_video, validation_video_mask, control_video + ] + ) + + def upload_resize_method(resize_method): + if resize_method == "Generate by": + return [gr.update(visible=True), gr.update(visible=True), gr.update(visible=False)] + else: + return [gr.update(visible=False), gr.update(visible=False), gr.update(visible=True)] + resize_method.change( + upload_resize_method, resize_method, [width_slider, height_slider, base_resolution] + ) + + generate_button.click( + fn=controller.generate, + inputs=[ + diffusion_transformer_dropdown, + base_model_dropdown, + lora_model_dropdown, + lora_alpha_slider, + prompt_textbox, + negative_prompt_textbox, + sampler_dropdown, + sample_step_slider, + resize_method, + width_slider, + height_slider, + base_resolution, + generation_method, + length_slider, + overlap_video_length, + partial_video_length, + cfg_scale_slider, + start_image, + end_image, + validation_video, + validation_video_mask, + control_video, + denoise_strength, + seed_textbox, + ref_image, + enable_teacache, + teacache_threshold, + num_skip_start_steps, + teacache_offload, + cfg_skip_ratio, + enable_riflex, + riflex_k, + base_model_2_dropdown, + lora_model_2_dropdown + ], + outputs=[result_image, result_video, infer_progress] + ) + return demo, controller + +def ui_host(GPU_memory_mode, scheduler_dict, model_name, model_type, config_path, compile_dit, weight_dtype, savedir_sample=None): + controller = Wan2_2_Fun_Controller_Host( + GPU_memory_mode, scheduler_dict, model_name=model_name, model_type=model_type, + config_path=config_path, compile_dit=compile_dit, + weight_dtype=weight_dtype, savedir_sample=savedir_sample, + ) + + with gr.Blocks(css=css) as demo: + gr.Markdown( + """ + # Wan2.2-Fun: + + A Wan with more flexible generation conditions, capable of producing videos of different resolutions, around 5 seconds, and fps 16 (frames 1 to 81), as well as image generated videos. + + [Github](https://github.com/aigc-apps/VideoX-Fun/) + """ + ) + with gr.Column(variant="panel"): + model_type = create_fake_model_type(visible=False) + diffusion_transformer_dropdown = create_fake_model_checkpoints(model_name, visible=True) + base_model_dropdown, lora_model_dropdown, lora_alpha_slider = \ + create_fake_finetune_models_checkpoints(visible=True, add_checkpoint_2=True) + base_model_dropdown, base_model_2_dropdown = base_model_dropdown + lora_model_dropdown, lora_model_2_dropdown = lora_model_dropdown + + with gr.Row(): + enable_teacache, teacache_threshold, num_skip_start_steps, teacache_offload = \ + create_teacache_params(True, 0.10, 1, False) + cfg_skip_ratio = create_cfg_skip_params(0) + enable_riflex, riflex_k = create_cfg_riflex_k(False, 6) + + with gr.Column(variant="panel"): + prompt_textbox, negative_prompt_textbox = create_prompts(negative_prompt="色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走") + + with gr.Row(): + with gr.Column(): + sampler_dropdown, sample_step_slider = create_samplers(controller) + + resize_method, width_slider, height_slider, base_resolution = create_height_width( + default_height = 480, default_width = 832, maximum_height = 1344, + maximum_width = 1344, + ) + generation_method, length_slider, overlap_video_length, partial_video_length = \ + create_generation_methods_and_video_length( + ["Video Generation", "Image Generation"], + default_video_length=81, + maximum_video_length=161, + ) + 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 Control (视频控制)"], prompt_textbox, support_ref_image=True + ) + cfg_scale_slider, seed_textbox, seed_button = create_cfg_and_seedbox(gradio_version_is_above_4) + + generate_button = gr.Button(value="Generate (生成)", variant='primary') + + result_image, result_video, infer_progress = create_ui_outputs() + + def upload_generation_method(generation_method): + if generation_method == "Video Generation": + return gr.update(visible=True, minimum=1, maximum=161, value=81, interactive=True) + elif generation_method == "Image Generation": + return gr.update(minimum=1, maximum=1, value=1, interactive=False) + generation_method.change( + upload_generation_method, generation_method, [length_slider] + ) + + def upload_source_method(source_method): + if source_method == "Text to Video (文本到视频)": + return [gr.update(visible=False), gr.update(visible=False), gr.update(visible=False), gr.update(value=None), gr.update(value=None), gr.update(value=None), gr.update(value=None), gr.update(value=None)] + elif source_method == "Image to Video (图片到视频)": + return [gr.update(visible=True), gr.update(visible=False), gr.update(visible=False), gr.update(), gr.update(), gr.update(value=None), gr.update(value=None), gr.update(value=None)] + elif source_method == "Video to Video (视频到视频)": + return [gr.update(visible=False), gr.update(visible=True), gr.update(visible=False), gr.update(value=None), gr.update(value=None), gr.update(), gr.update(), gr.update(value=None)] + else: + return [gr.update(visible=False), gr.update(visible=False), gr.update(visible=True), gr.update(value=None), gr.update(value=None), gr.update(value=None), gr.update(value=None), gr.update()] + source_method.change( + upload_source_method, source_method, [ + image_to_video_col, video_to_video_col, control_video_col, start_image, end_image, + validation_video, validation_video_mask, control_video + ] + ) + + def upload_resize_method(resize_method): + if resize_method == "Generate by": + return [gr.update(visible=True), gr.update(visible=True), gr.update(visible=False)] + else: + return [gr.update(visible=False), gr.update(visible=False), gr.update(visible=True)] + resize_method.change( + upload_resize_method, resize_method, [width_slider, height_slider, base_resolution] + ) + + generate_button.click( + fn=controller.generate, + inputs=[ + diffusion_transformer_dropdown, + base_model_dropdown, + lora_model_dropdown, + lora_alpha_slider, + prompt_textbox, + negative_prompt_textbox, + sampler_dropdown, + sample_step_slider, + resize_method, + width_slider, + height_slider, + base_resolution, + generation_method, + length_slider, + overlap_video_length, + partial_video_length, + cfg_scale_slider, + start_image, + end_image, + validation_video, + validation_video_mask, + control_video, + denoise_strength, + seed_textbox, + ref_image, + enable_teacache, + teacache_threshold, + num_skip_start_steps, + teacache_offload, + cfg_skip_ratio, + enable_riflex, + riflex_k, + base_model_2_dropdown, + lora_model_2_dropdown + ], + outputs=[result_image, result_video, infer_progress] + ) + return demo, controller + +def ui_client(scheduler_dict, model_name, savedir_sample=None): + controller = Wan2_2_Fun_Controller_Client(scheduler_dict, savedir_sample) + + with gr.Blocks(css=css) as demo: + gr.Markdown( + """ + # Wan2.2-Fun: + + A Wan with more flexible generation conditions, capable of producing videos of different resolutions, around 5 seconds, and fps 16 (frames 1 to 81), as well as image generated videos. + + [Github](https://github.com/aigc-apps/VideoX-Fun/) + """ + ) + with gr.Column(variant="panel"): + diffusion_transformer_dropdown = create_fake_model_checkpoints(model_name, visible=True) + base_model_dropdown, lora_model_dropdown, lora_alpha_slider = \ + create_fake_finetune_models_checkpoints(visible=True, add_checkpoint_2=True) + base_model_dropdown, base_model_2_dropdown = base_model_dropdown + lora_model_dropdown, lora_model_2_dropdown = lora_model_dropdown + + with gr.Row(): + enable_teacache, teacache_threshold, num_skip_start_steps, teacache_offload = \ + create_teacache_params(True, 0.10, 1, False) + cfg_skip_ratio = create_cfg_skip_params(0) + enable_riflex, riflex_k = create_cfg_riflex_k(False, 6) + + with gr.Column(variant="panel"): + prompt_textbox, negative_prompt_textbox = create_prompts(negative_prompt="色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走") + + with gr.Row(): + with gr.Column(): + sampler_dropdown, sample_step_slider = create_samplers(controller, maximum_step=50) + + resize_method, width_slider, height_slider, base_resolution = create_fake_height_width( + default_height = 480, default_width = 832, maximum_height = 1344, + maximum_width = 1344, + ) + generation_method, length_slider, overlap_video_length, partial_video_length = \ + create_generation_methods_and_video_length( + ["Video Generation", "Image Generation"], + default_video_length=81, + maximum_video_length=161, + ) + 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 (图片到视频)"], prompt_textbox + ) + + cfg_scale_slider, seed_textbox, seed_button = create_cfg_and_seedbox(gradio_version_is_above_4) + + generate_button = gr.Button(value="Generate (生成)", variant='primary') + + result_image, result_video, infer_progress = create_ui_outputs() + + def upload_generation_method(generation_method): + if generation_method == "Video Generation": + return gr.update(visible=True, minimum=5, maximum=161, value=49, interactive=True) + elif generation_method == "Image Generation": + return gr.update(minimum=1, maximum=1, value=1, interactive=False) + generation_method.change( + upload_generation_method, generation_method, [length_slider] + ) + + def upload_source_method(source_method): + if source_method == "Text to Video (文本到视频)": + return [gr.update(visible=False), gr.update(visible=False), gr.update(value=None), gr.update(value=None), gr.update(value=None), gr.update(value=None)] + elif source_method == "Image to Video (图片到视频)": + return [gr.update(visible=True), gr.update(visible=False), gr.update(), gr.update(), gr.update(value=None), gr.update(value=None)] + else: + return [gr.update(visible=False), gr.update(visible=True), gr.update(value=None), gr.update(value=None), gr.update(), gr.update()] + source_method.change( + upload_source_method, source_method, [image_to_video_col, video_to_video_col, start_image, end_image, validation_video, validation_video_mask] + ) + + def upload_resize_method(resize_method): + if resize_method == "Generate by": + return [gr.update(visible=True), gr.update(visible=True), gr.update(visible=False)] + else: + return [gr.update(visible=False), gr.update(visible=False), gr.update(visible=True)] + resize_method.change( + upload_resize_method, resize_method, [width_slider, height_slider, base_resolution] + ) + + generate_button.click( + fn=controller.generate, + inputs=[ + diffusion_transformer_dropdown, + base_model_dropdown, + lora_model_dropdown, + lora_alpha_slider, + prompt_textbox, + negative_prompt_textbox, + sampler_dropdown, + sample_step_slider, + resize_method, + width_slider, + height_slider, + base_resolution, + generation_method, + length_slider, + cfg_scale_slider, + start_image, + end_image, + validation_video, + validation_video_mask, + denoise_strength, + seed_textbox, + ref_image, + enable_teacache, + teacache_threshold, + num_skip_start_steps, + teacache_offload, + cfg_skip_ratio, + enable_riflex, + riflex_k, + base_model_2_dropdown, + lora_model_2_dropdown + ], + outputs=[result_image, result_video, infer_progress] + ) + return demo, controller \ No newline at end of file diff --git a/videox_fun/ui/wan2_2_ui.py b/videox_fun/ui/wan2_2_ui.py index 995a455..46f8244 100644 --- a/videox_fun/ui/wan2_2_ui.py +++ b/videox_fun/ui/wan2_2_ui.py @@ -93,7 +93,7 @@ class Wan2_2_Controller(Fun_Controller): # Get pipeline if self.model_type == "Inpaint": - if "ti2v" in self.config_path: + if "wan_civitai_5b" in self.config_path: self.pipeline = Wan2_2TI2VPipeline( vae=self.vae, tokenizer=self.tokenizer,