Co-authored-by: 宣源 <xuanyuan.lb@alibaba-inc.com>
This commit is contained in:
liubo0902
2025-06-26 16:17:03 +08:00
committed by GitHub
co-authored by 宣源
parent 097208e817
commit c2d0c7ae21
2 changed files with 75 additions and 0 deletions
+11
View File
@@ -396,6 +396,12 @@ def parse_args():
action="store_true",
help="Whether or not to use gradient checkpointing to save memory at the expense of slower backward pass.",
)
parser.add_argument(
"--selective_ac",
type=float,
default=0,
help="Rate for transformer block apply checkpointing.",
)
parser.add_argument(
"--learning_rate",
type=float,
@@ -1050,6 +1056,11 @@ def main():
if args.gradient_checkpointing:
transformer3d.enable_gradient_checkpointing()
elif args.selective_ac > 0:
from videox_fun.utils.ac_handle import apply_checkpointing, partial
from videox_fun.models.wan_transformer3d import WanAttentionBlock
apply_selective_ac = partial(apply_checkpointing, block=WanAttentionBlock)
apply_selective_ac(transformer3d, p=args.selective_ac)
# Enable TF32 for faster training on Ampere GPUs,
# cf https://pytorch.org/docs/stable/notes/cuda.html#tensorfloat-32-tf32-on-ampere-devices