Update Training Code

This commit is contained in:
bubbliiiing
2025-06-10 03:54:11 +00:00
parent 72c2792139
commit a4476549e4
7 changed files with 41 additions and 4 deletions
+4
View File
@@ -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
+4
View File
@@ -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
+4
View File
@@ -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
+4
View File
@@ -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
+4
View File
@@ -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
+4
View File
@@ -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
+17 -4
View File
@@ -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,