From a17b35acfbeb8f1ed30db3ad1be84d1b22ab05da Mon Sep 17 00:00:00 2001 From: Bubbliiiing <47347516+bubbliiiing@users.noreply.github.com> Date: Mon, 21 Jul 2025 15:22:13 +0800 Subject: [PATCH] Update Torch Compile Node in Comfyui (#253) --- comfyui/comfyui_nodes.py | 32 ++++++++++++++++++++++++++ videox_fun/dist/__init__.py | 14 +++++------ videox_fun/models/__init__.py | 30 +++++++++++++----------- videox_fun/models/wan_transformer3d.py | 5 +++- videox_fun/ui/cogvideox_fun_ui.py | 9 ++++---- videox_fun/ui/wan_fun_ui.py | 9 ++++---- videox_fun/ui/wan_ui.py | 9 ++++---- 7 files changed, 74 insertions(+), 34 deletions(-) diff --git a/comfyui/comfyui_nodes.py b/comfyui/comfyui_nodes.py index 3633544..c91e0d6 100755 --- a/comfyui/comfyui_nodes.py +++ b/comfyui/comfyui_nodes.py @@ -51,6 +51,36 @@ class FunRiflex: def process(self, riflex_k): return (riflex_k, ) +class FunCompile: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "cache_size_limit": ("INT", {"default": 64, "min": 0, "max": 10086}), + "funmodels": ("FunModels",) + } + } + RETURN_TYPES = ("FunModels",) + RETURN_NAMES = ("funmodels",) + FUNCTION = "compile" + CATEGORY = "CogVideoXFUNWrapper" + + def compile(self, cache_size_limit, funmodels): + torch._dynamo.config.cache_size_limit = cache_size_limit + if hasattr(funmodels["pipeline"].transformer, "blocks"): + for i in range(len(funmodels["pipeline"].transformer.blocks)): + funmodels["pipeline"].transformer.blocks[i] = torch.compile(funmodels["pipeline"].transformer.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]) + + else: + funmodels["pipeline"].transformer.forward = torch.compile(funmodels["pipeline"].transformer.forward) + + print("Add Compile") + return (funmodels,) + def gen_gaussian_heatmap(imgSize=200): circle_img = np.zeros((imgSize, imgSize,), np.float32) circle_mask = cv2.circle(circle_img, (imgSize//2, imgSize//2), imgSize//2 - 1, 1, -1) @@ -270,6 +300,7 @@ class CameraTrajectoryFromChaoJie: NODE_CLASS_MAPPINGS = { "FunTextBox": FunTextBox, "FunRiflex": FunRiflex, + "FunCompile": FunCompile, "LoadCogVideoXFunModel": LoadCogVideoXFunModel, "LoadCogVideoXFunLora": LoadCogVideoXFunLora, @@ -304,6 +335,7 @@ NODE_CLASS_MAPPINGS = { NODE_DISPLAY_NAME_MAPPINGS = { "FunTextBox": "FunTextBox", "FunRiflex": "FunRiflex", + "FunCompile": "FunCompile", "LoadCogVideoXFunModel": "Load CogVideoX-Fun Model", "LoadCogVideoXFunLora": "Load CogVideoX-Fun Lora", diff --git a/videox_fun/dist/__init__.py b/videox_fun/dist/__init__.py index 267c0b5..f038f06 100755 --- a/videox_fun/dist/__init__.py +++ b/videox_fun/dist/__init__.py @@ -11,18 +11,19 @@ from .wan_xfuser import usp_attn_forward # The pai_fuser is an internally developed acceleration package, which can be used on PAI. if importlib.util.find_spec("pai_fuser") is not None: - from pai_fuser.core import parallel_magvit_vae - from pai_fuser.core.attention import wan_usp_sparse_attention_wrapper - from . import wan_xfuser - # The simple_wrapper is used to solve the problem about conflicts between cython and torch.compile def simple_wrapper(func): def inner(*args, **kwargs): return func(*args, **kwargs) return inner - wan_xfuser.usp_attn_forward = simple_wrapper(wan_usp_sparse_attention_wrapper()(wan_xfuser.usp_attn_forward)) - usp_attn_forward = simple_wrapper(wan_xfuser.usp_attn_forward) + from pai_fuser.core import parallel_magvit_vae + from pai_fuser.core.attention import wan_usp_sparse_attention_wrapper + from . import wan_xfuser + + usp_sparse_attn_wrap_forward = simple_wrapper(wan_usp_sparse_attention_wrapper()(wan_xfuser.usp_attn_forward)) + wan_xfuser.usp_attn_forward = usp_sparse_attn_wrap_forward + usp_attn_forward = usp_sparse_attn_wrap_forward print("Import PAI VAE Turbo and Sparse Attention") from pai_fuser.core.rope import ENABLE_KERNEL, usp_fast_rope_apply_qk @@ -30,7 +31,6 @@ if importlib.util.find_spec("pai_fuser") is not None: if ENABLE_KERNEL: import torch import types - from .wan_xfuser import rope_apply def deepcopy_function(f): return types.FunctionType(f.__code__, f.__globals__, name=f.__name__, argdefs=f.__defaults__,closure=f.__closure__) diff --git a/videox_fun/models/__init__.py b/videox_fun/models/__init__.py index 6dccf30..52ee607 100755 --- a/videox_fun/models/__init__.py +++ b/videox_fun/models/__init__.py @@ -1,34 +1,36 @@ +import importlib.util + from transformers import AutoTokenizer, T5EncoderModel, T5Tokenizer from .cogvideox_transformer3d import CogVideoXTransformer3DModel from .cogvideox_vae import AutoencoderKLCogVideoX from .wan_image_encoder import CLIPModel from .wan_text_encoder import WanT5EncoderModel -from .wan_transformer3d import WanTransformer3DModel, WanSelfAttention +from .wan_transformer3d import WanSelfAttention, WanTransformer3DModel from .wan_vae import AutoencoderKLWan, AutoencoderKLWan_ - -import importlib.util - # The pai_fuser is an internally developed acceleration package, which can be used on PAI. if importlib.util.find_spec("pai_fuser") is not None: - from ..dist import parallel_magvit_vae - AutoencoderKLWan_.decode = parallel_magvit_vae(0.2, 8)(AutoencoderKLWan_.decode) - - from pai_fuser.core.attention import wan_sparse_attention_wrapper - import torch - # The simple_wrapper is used to solve the problem about conflicts between cython and torch.compile def simple_wrapper(func): def inner(*args, **kwargs): return func(*args, **kwargs) return inner + + from ..dist import parallel_magvit_vae + AutoencoderKLWan_.decode = simple_wrapper(parallel_magvit_vae(0.2, 8)(AutoencoderKLWan_.decode)) + + import torch + from pai_fuser.core.attention import wan_sparse_attention_wrapper + WanSelfAttention.forward = simple_wrapper(wan_sparse_attention_wrapper()(WanSelfAttention.forward)) print("Import Sparse Attention") - + + WanTransformer3DModel.forward = simple_wrapper(WanTransformer3DModel.forward) + import os - from pai_fuser.core import (cfg_skip_turbo, enable_cfg_skip, - disable_cfg_skip) + from pai_fuser.core import (cfg_skip_turbo, disable_cfg_skip, + enable_cfg_skip) WanTransformer3DModel.enable_cfg_skip = enable_cfg_skip()(WanTransformer3DModel.enable_cfg_skip) WanTransformer3DModel.disable_cfg_skip = disable_cfg_skip()(WanTransformer3DModel.disable_cfg_skip) @@ -38,7 +40,7 @@ if importlib.util.find_spec("pai_fuser") is not None: if ENABLE_KERNEL: import types - from .wan_transformer3d import rope_apply + from . import wan_transformer3d def deepcopy_function(f): return types.FunctionType(f.__code__, f.__globals__, name=f.__name__, argdefs=f.__defaults__,closure=f.__closure__) diff --git a/videox_fun/models/wan_transformer3d.py b/videox_fun/models/wan_transformer3d.py index c03155a..d58ce97 100755 --- a/videox_fun/models/wan_transformer3d.py +++ b/videox_fun/models/wan_transformer3d.py @@ -49,9 +49,12 @@ try: elif f"{major}.{minor}" == "8.9": from sageattention_sm89 import sageattn SAGE_ATTENTION_AVAILABLE = True - elif major>=9: + elif f"{major}.{minor}" == "9.0": from sageattention_sm90 import sageattn SAGE_ATTENTION_AVAILABLE = True + elif major>9: + from sageattention_sm120 import sageattn + SAGE_ATTENTION_AVAILABLE = True except: try: from sageattention import sageattn diff --git a/videox_fun/ui/cogvideox_fun_ui.py b/videox_fun/ui/cogvideox_fun_ui.py index 8964a14..e73e60e 100755 --- a/videox_fun/ui/cogvideox_fun_ui.py +++ b/videox_fun/ui/cogvideox_fun_ui.py @@ -193,6 +193,9 @@ class CogVideoXFunController(Fun_Controller): self.pipeline = merge_lora(self.pipeline, self.lora_model_path, multiplier=lora_alpha_slider) print(f"Merge Lora done.") + if fps is None: + fps = 8 + 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) @@ -265,7 +268,7 @@ class CogVideoXFunController(Fun_Controller): last_frames = init_frames + _partial_video_length else: if validation_video is not None: - input_video, input_video_mask, ref_image, 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=8) + input_video, input_video_mask, ref_image, 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) strength = denoise_strength 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)) @@ -297,7 +300,7 @@ class CogVideoXFunController(Fun_Controller): generator = generator ).videos else: - input_video, input_video_mask, ref_image, clip_image = get_video_to_video_latent(control_video, length_slider if not is_image else 1, sample_size=(height_slider, width_slider), fps=8) + input_video, input_video_mask, ref_image, clip_image = get_video_to_video_latent(control_video, length_slider if not is_image else 1, sample_size=(height_slider, width_slider), fps=fps) sample = self.pipeline( prompt_textbox, @@ -333,8 +336,6 @@ class CogVideoXFunController(Fun_Controller): print(f"Unmerge Lora done.") print(f"Saving outputs.") - if fps == None: - fps = 16 save_sample_path = self.save_outputs( is_image, length_slider, sample, fps=fps ) diff --git a/videox_fun/ui/wan_fun_ui.py b/videox_fun/ui/wan_fun_ui.py index 84c1e1f..a89536c 100755 --- a/videox_fun/ui/wan_fun_ui.py +++ b/videox_fun/ui/wan_fun_ui.py @@ -245,6 +245,9 @@ class Wan_Fun_Controller(Fun_Controller): 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 if enable_riflex: print(f"Enable riflex") @@ -256,7 +259,7 @@ class Wan_Fun_Controller(Fun_Controller): 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=16) + 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)) @@ -299,7 +302,7 @@ class Wan_Fun_Controller(Fun_Controller): if start_image is not None: start_image = get_image_latent(start_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=16, ref_image=None) + 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, @@ -339,8 +342,6 @@ class Wan_Fun_Controller(Fun_Controller): print(f"Unmerge Lora done.") print(f"Saving outputs.") - if fps == None: - fps = 16 save_sample_path = self.save_outputs( is_image, length_slider, sample, fps=fps ) diff --git a/videox_fun/ui/wan_ui.py b/videox_fun/ui/wan_ui.py index c08f2d8..3613817 100755 --- a/videox_fun/ui/wan_ui.py +++ b/videox_fun/ui/wan_ui.py @@ -237,6 +237,9 @@ class Wan_Controller(Fun_Controller): 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 if enable_riflex: print(f"Enable riflex") @@ -248,7 +251,7 @@ class Wan_Controller(Fun_Controller): 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=16) + 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)) @@ -291,7 +294,7 @@ class Wan_Controller(Fun_Controller): if start_image is not None: start_image = get_image_latent(start_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=16, ref_image=None) + 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, @@ -331,8 +334,6 @@ class Wan_Controller(Fun_Controller): print(f"Unmerge Lora done.") print(f"Saving outputs.") - if fps == None: - fps = 16 save_sample_path = self.save_outputs( is_image, length_slider, sample, fps=fps )