diff --git a/scripts/wan2.1/train.py b/scripts/wan2.1/train.py index 4d3a827..c5aa7e6 100755 --- a/scripts/wan2.1/train.py +++ b/scripts/wan2.1/train.py @@ -773,8 +773,12 @@ def main(): 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 diff --git a/scripts/wan2.1/train_lora.py b/scripts/wan2.1/train_lora.py index 50a1883..7c021dc 100755 --- a/scripts/wan2.1/train_lora.py +++ b/scripts/wan2.1/train_lora.py @@ -772,8 +772,12 @@ def main(): 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 diff --git a/scripts/wan2.1_fun/train.py b/scripts/wan2.1_fun/train.py index 1863178..cf89f6e 100755 --- a/scripts/wan2.1_fun/train.py +++ b/scripts/wan2.1_fun/train.py @@ -743,8 +743,12 @@ def main(): 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 diff --git a/scripts/wan2.1_fun/train_control.py b/scripts/wan2.1_fun/train_control.py index 23d01c6..ef471de 100755 --- a/scripts/wan2.1_fun/train_control.py +++ b/scripts/wan2.1_fun/train_control.py @@ -675,8 +675,12 @@ def main(): 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 diff --git a/scripts/wan2.1_fun/train_control_lora.py b/scripts/wan2.1_fun/train_control_lora.py index 9a79d25..b36c39b 100755 --- a/scripts/wan2.1_fun/train_control_lora.py +++ b/scripts/wan2.1_fun/train_control_lora.py @@ -673,8 +673,12 @@ def main(): 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 diff --git a/scripts/wan2.1_fun/train_lora.py b/scripts/wan2.1_fun/train_lora.py index e64e224..5d392fe 100755 --- a/scripts/wan2.1_fun/train_lora.py +++ b/scripts/wan2.1_fun/train_lora.py @@ -735,8 +735,12 @@ def main(): 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 diff --git a/videox_fun/models/wan_transformer3d.py b/videox_fun/models/wan_transformer3d.py index 16bb296..ebb55ba 100755 --- a/videox_fun/models/wan_transformer3d.py +++ b/videox_fun/models/wan_transformer3d.py @@ -2,6 +2,7 @@ # Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved. import glob +import importlib.metadata import json import math import os @@ -17,6 +18,7 @@ from diffusers.configuration_utils import ConfigMixin, register_to_config from diffusers.loaders.single_file_model import FromOriginalModelMixin from diffusers.models.modeling_utils import ModelMixin from diffusers.utils import is_torch_version, logging +from packaging import version from torch import nn from ..dist import (get_sequence_parallel_rank, @@ -60,6 +62,13 @@ except: sageattn = None SAGE_ATTENTION_AVAILABLE = False +try: + diffusers_version = importlib.metadata.version("diffusers") +except importlib.metadata.PackageNotFoundError: + diffusers_version = "0.0.0" + +USE_NEW_SIGNATURE = version.parse(diffusers_version) >= version.parse("0.33.1") + def flash_attention( q, k, @@ -843,7 +852,14 @@ class WanTransformer3DModel(ModelMixin, ConfigMixin, FromOriginalModelMixin): self.gradient_checkpointing = False self.sp_world_size = 1 self.sp_world_rank = 0 - + + if USE_NEW_SIGNATURE: + def _set_gradient_checkpointing(self, enable=False, gradient_checkpointing_func=None): + self.gradient_checkpointing = enable + else: + def _set_gradient_checkpointing(self, module, value=False): + self.gradient_checkpointing = value + def enable_teacache( self, coefficients, @@ -908,9 +924,6 @@ class WanTransformer3DModel(ModelMixin, ConfigMixin, FromOriginalModelMixin): block.self_attn.forward = types.MethodType( usp_attn_forward, block.self_attn) - def _set_gradient_checkpointing(self, module, value=False): - self.gradient_checkpointing = value - @cfg_skip() def forward( self,