From 74426601fa02f530b0feb61882650e39def3f0b5 Mon Sep 17 00:00:00 2001 From: Bubbliiiing <47347516+bubbliiiing@users.noreply.github.com> Date: Mon, 8 Sep 2025 11:43:02 +0800 Subject: [PATCH] Update import (#291) --- videox_fun/dist/__init__.py | 47 +++++++++++++++-------- videox_fun/dist/fuser.py | 23 +++++------ videox_fun/models/__init__.py | 67 +++++++++++++++++++++++++-------- videox_fun/pipeline/__init__.py | 7 +++- videox_fun/utils/__init__.py | 18 +++++++-- 5 files changed, 114 insertions(+), 48 deletions(-) diff --git a/videox_fun/dist/__init__.py b/videox_fun/dist/__init__.py index 881adfb..978390e 100755 --- a/videox_fun/dist/__init__.py +++ b/videox_fun/dist/__init__.py @@ -12,38 +12,55 @@ from .qwen_xfuser import QwenImageMultiGPUsAttnProcessor2_0 from .flux_xfuser import FluxMultiGPUsAttnProcessor2_0 # 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: - # The simple_wrapper is used to solve the problem about conflicts between cython and torch.compile +if importlib.util.find_spec("paifuser") is not None: + # --------------------------------------------------------------- # + # 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 pai_fuser.core import parallel_magvit_vae - from pai_fuser.core.attention import wan_usp_sparse_attention_wrapper + # --------------------------------------------------------------- # + # Sparse Attention Kernel + # --------------------------------------------------------------- # + from paifuser.models import parallel_magvit_vae + from paifuser.ops import wan_usp_sparse_attention_wrapper from . import wan_xfuser + # --------------------------------------------------------------- # + # Sparse Attention + # --------------------------------------------------------------- # 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 + # --------------------------------------------------------------- # + # Fast Rope Kernel + # --------------------------------------------------------------- # + import types + import torch + from paifuser.ops import (ENABLE_KERNEL, usp_fast_rope_apply_qk, + usp_rope_apply_real_qk) + + def deepcopy_function(f): + return types.FunctionType(f.__code__, f.__globals__, name=f.__name__, argdefs=f.__defaults__,closure=f.__closure__) + + local_rope_apply_qk = deepcopy_function(wan_xfuser.rope_apply_qk) if ENABLE_KERNEL: - import torch - import types - - def deepcopy_function(f): - return types.FunctionType(f.__code__, f.__globals__, name=f.__name__, argdefs=f.__defaults__,closure=f.__closure__) - - local_rope_apply_qk = deepcopy_function(wan_xfuser.rope_apply_qk) def adaptive_fast_usp_rope_apply_qk(q, k, grid_sizes, freqs): if torch.is_grad_enabled(): return local_rope_apply_qk(q, k, grid_sizes, freqs) else: return usp_fast_rope_apply_qk(q, k, grid_sizes, freqs) + + else: + def adaptive_fast_usp_rope_apply_qk(q, k, grid_sizes, freqs): + return usp_rope_apply_real_qk(q, k, grid_sizes, freqs) - wan_xfuser.rope_apply_qk = adaptive_fast_usp_rope_apply_qk - rope_apply_qk = adaptive_fast_usp_rope_apply_qk - print("Import PAI Fast rope") + wan_xfuser.rope_apply_qk = adaptive_fast_usp_rope_apply_qk + rope_apply_qk = adaptive_fast_usp_rope_apply_qk + print("Import PAI Fast rope") \ No newline at end of file diff --git a/videox_fun/dist/fuser.py b/videox_fun/dist/fuser.py index de12a2b..d29bc9b 100755 --- a/videox_fun/dist/fuser.py +++ b/videox_fun/dist/fuser.py @@ -5,13 +5,13 @@ import torch.distributed as dist try: # 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: - import pai_fuser - from pai_fuser.core.distributed import ( + if importlib.util.find_spec("paifuser") is not None: + import paifuser + from paifuser.xfuser.core.distributed import ( get_sequence_parallel_rank, get_sequence_parallel_world_size, get_sp_group, get_world_group, init_distributed_environment, initialize_model_parallel) - from pai_fuser.core.long_ctx_attention import \ + from paifuser.xfuser.core.long_ctx_attention import \ xFuserLongContextAttention print("Import PAI DiT Turbo") else: @@ -32,18 +32,19 @@ except Exception as ex: init_distributed_environment = None initialize_model_parallel = None -def set_multi_gpus_devices(ulysses_degree, ring_degree): - if ulysses_degree > 1 or ring_degree > 1: +def set_multi_gpus_devices(ulysses_degree, ring_degree, classifier_free_guidance_degree=1): + if ulysses_degree > 1 or ring_degree > 1 or classifier_free_guidance_degree > 1: if get_sp_group is None: raise RuntimeError("xfuser is not installed.") dist.init_process_group("nccl") - print('parallel inference enabled: ulysses_degree=%d ring_degree=%d rank=%d world_size=%d' % ( - ulysses_degree, ring_degree, dist.get_rank(), + print('parallel inference enabled: ulysses_degree=%d ring_degree=%d classifier_free_guidance_degree=% rank=%d world_size=%d' % ( + ulysses_degree, ring_degree, classifier_free_guidance_degree, dist.get_rank(), dist.get_world_size())) - assert dist.get_world_size() == ring_degree * ulysses_degree, \ - "number of GPUs(%d) should be equal to ring_degree * ulysses_degree." % dist.get_world_size() + assert dist.get_world_size() == ring_degree * ulysses_degree * classifier_free_guidance_degree, \ + "number of GPUs(%d) should be equal to ring_degree * ulysses_degree * classifier_free_guidance_degree." % dist.get_world_size() init_distributed_environment(rank=dist.get_rank(), world_size=dist.get_world_size()) - initialize_model_parallel(sequence_parallel_degree=dist.get_world_size(), + initialize_model_parallel(sequence_parallel_degree=ring_degree * ulysses_degree, + classifier_free_guidance_degree=classifier_free_guidance_degree, ring_degree=ring_degree, ulysses_degree=ulysses_degree) # device = torch.device("cuda:%d" % dist.get_rank()) diff --git a/videox_fun/models/__init__.py b/videox_fun/models/__init__.py index b74f88a..41503bf 100755 --- a/videox_fun/models/__init__.py +++ b/videox_fun/models/__init__.py @@ -26,49 +26,84 @@ from .wan_vae import AutoencoderKLWan, AutoencoderKLWan_ from .wan_vae3_8 import AutoencoderKLWan2_2_, AutoencoderKLWan3_8 # 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: - # The simple_wrapper is used to solve the problem about conflicts between cython and torch.compile +if importlib.util.find_spec("paifuser") is not None: + # --------------------------------------------------------------- # + # 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 + # --------------------------------------------------------------- # + # VAE Parallel Kernel + # --------------------------------------------------------------- # from ..dist import parallel_magvit_vae AutoencoderKLWan_.decode = simple_wrapper(parallel_magvit_vae(0.4, 8)(AutoencoderKLWan_.decode)) AutoencoderKLWan2_2_.decode = simple_wrapper(parallel_magvit_vae(0.4, 16)(AutoencoderKLWan2_2_.decode)) + # --------------------------------------------------------------- # + # Sparse Attention + # --------------------------------------------------------------- # import torch - from pai_fuser.core.attention import wan_sparse_attention_wrapper + from paifuser.ops 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) + # --------------------------------------------------------------- # + # CFG Skip Turbo + # --------------------------------------------------------------- # import os - from pai_fuser.core import (cfg_skip_turbo, disable_cfg_skip, - enable_cfg_skip) + + if importlib.util.find_spec("paifuser.accelerator") is not None: + from paifuser.accelerator import (cfg_skip_turbo, disable_cfg_skip, + enable_cfg_skip, share_cfg_skip) + else: + from paifuser import (cfg_skip_turbo, disable_cfg_skip, + enable_cfg_skip, share_cfg_skip) WanTransformer3DModel.enable_cfg_skip = enable_cfg_skip()(WanTransformer3DModel.enable_cfg_skip) WanTransformer3DModel.disable_cfg_skip = disable_cfg_skip()(WanTransformer3DModel.disable_cfg_skip) + WanTransformer3DModel.share_cfg_skip = share_cfg_skip()(WanTransformer3DModel.share_cfg_skip) print("Import CFG Skip Turbo") - from pai_fuser.core.rope import ENABLE_KERNEL, fast_rope_apply_qk + # --------------------------------------------------------------- # + # RMS Norm Kernel + # --------------------------------------------------------------- # + from paifuser.ops import rms_norm_forward + WanRMSNorm.forward = rms_norm_forward + print("Import PAI RMS Fuse") + + # --------------------------------------------------------------- # + # Fast Rope Kernel + # --------------------------------------------------------------- # + import types + + import torch + from paifuser.ops import (ENABLE_KERNEL, fast_rope_apply_qk, + rope_apply_real_qk) + + from . import wan_transformer3d + + def deepcopy_function(f): + return types.FunctionType(f.__code__, f.__globals__, name=f.__name__, argdefs=f.__defaults__,closure=f.__closure__) + + local_rope_apply_qk = deepcopy_function(wan_transformer3d.rope_apply_qk) if ENABLE_KERNEL: - import types - from . import wan_transformer3d - - def deepcopy_function(f): - return types.FunctionType(f.__code__, f.__globals__, name=f.__name__, argdefs=f.__defaults__,closure=f.__closure__) - - local_rope_apply_qk = deepcopy_function(wan_transformer3d.rope_apply_qk) def adaptive_fast_rope_apply_qk(q, k, grid_sizes, freqs): if torch.is_grad_enabled(): return local_rope_apply_qk(q, k, grid_sizes, freqs) else: return fast_rope_apply_qk(q, k, grid_sizes, freqs) + else: + def adaptive_fast_rope_apply_qk(q, k, grid_sizes, freqs): + return rope_apply_real_qk(q, k, grid_sizes, freqs) - wan_transformer3d.rope_apply_qk = adaptive_fast_rope_apply_qk - rope_apply_qk = adaptive_fast_rope_apply_qk - print("Import PAI Fast rope") + wan_transformer3d.rope_apply_qk = adaptive_fast_rope_apply_qk + rope_apply_qk = adaptive_fast_rope_apply_qk + print("Import PAI Fast rope") \ No newline at end of file diff --git a/videox_fun/pipeline/__init__.py b/videox_fun/pipeline/__init__.py index 9bfeb7f..ec8b7f6 100755 --- a/videox_fun/pipeline/__init__.py +++ b/videox_fun/pipeline/__init__.py @@ -21,8 +21,11 @@ Wan2_2I2VPipeline = Wan2_2FunInpaintPipeline import importlib.util -if importlib.util.find_spec("pai_fuser") is not None: - from pai_fuser.core import sparse_reset +if importlib.util.find_spec("paifuser") is not None: + # --------------------------------------------------------------- # + # Sparse Attention + # --------------------------------------------------------------- # + from paifuser.ops import sparse_reset # Wan2.1 WanFunInpaintPipeline.__call__ = sparse_reset(WanFunInpaintPipeline.__call__) diff --git a/videox_fun/utils/__init__.py b/videox_fun/utils/__init__.py index 3b272e4..cb5b7ec 100755 --- a/videox_fun/utils/__init__.py +++ b/videox_fun/utils/__init__.py @@ -14,8 +14,11 @@ from .discrete_sampler import DiscreteSampling # 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 (convert_model_weight_to_float8, +if importlib.util.find_spec("paifuser") is not None: + # --------------------------------------------------------------- # + # FP8 Linear Kernel + # --------------------------------------------------------------- # + from paifuser.ops import (convert_model_weight_to_float8, convert_weight_dtype_wrapper) from . import fp8_optimization fp8_optimization.convert_model_weight_to_float8 = convert_model_weight_to_float8 @@ -24,8 +27,15 @@ if importlib.util.find_spec("pai_fuser") is not None: convert_weight_dtype_wrapper = fp8_optimization.convert_weight_dtype_wrapper print("Import PAI Quantization Turbo") - from pai_fuser.core import (cfg_skip_turbo, enable_cfg_skip, - disable_cfg_skip) + # --------------------------------------------------------------- # + # CFG Skip Turbo + # --------------------------------------------------------------- # + if importlib.util.find_spec("paifuser.accelerator") is not None: + from paifuser.accelerator import (cfg_skip_turbo, disable_cfg_skip, + enable_cfg_skip, share_cfg_skip) + else: + from paifuser import (cfg_skip_turbo, disable_cfg_skip, + enable_cfg_skip, share_cfg_skip) from . import cfg_optimization cfg_optimization.cfg_skip = cfg_skip_turbo cfg_skip = cfg_skip_turbo