From c2d0c7ae2106fc029a1a66e8c8cc25045d03da59 Mon Sep 17 00:00:00 2001 From: liubo0902 <38622806+liubo0902@users.noreply.github.com> Date: Thu, 26 Jun 2025 16:17:03 +0800 Subject: [PATCH] add ac (#225) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Co-authored-by: 宣源 --- scripts/wan2.1/train.py | 11 ++++++ videox_fun/utils/ac_handle.py | 64 +++++++++++++++++++++++++++++++++++ 2 files changed, 75 insertions(+) create mode 100644 videox_fun/utils/ac_handle.py diff --git a/scripts/wan2.1/train.py b/scripts/wan2.1/train.py index 4d51476..e148ae5 100755 --- a/scripts/wan2.1/train.py +++ b/scripts/wan2.1/train.py @@ -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 diff --git a/videox_fun/utils/ac_handle.py b/videox_fun/utils/ac_handle.py new file mode 100644 index 0000000..91df98a --- /dev/null +++ b/videox_fun/utils/ac_handle.py @@ -0,0 +1,64 @@ +from functools import partial + +from torch.distributed.algorithms._checkpoint.checkpoint_wrapper import ( + CheckpointImpl, + apply_activation_checkpointing, + checkpoint_wrapper, +) + + +non_reentrant_wrapper = partial( + checkpoint_wrapper, + checkpoint_impl=CheckpointImpl.NO_REENTRANT, +) + + +def apply_checkpointing(model, block, p): + """ + Apply selective activation checkpointing. + + Selectivity is defined as a percentage p, which means we apply ac + on p of the total blocks. p is a floating number in the range of + [0, 1]. + + Some examples: + p = 0: no ac for all blocks. same as `fsdp_activation_checkpointing=False` + p = 1: apply ac on every block. i.e. "full ac". + p = 1/2: [ac, no-ac, ac, no-ac, ...] + p = 1/3: [no-ac, ac, no-ac, no-ac, ac, no-ac, ...] + p = 2/3: [ac, no-ac, ac, ac, no-ac, ac, ...] + Since blocks are homogeneous, we make ac blocks evenly spaced among + all blocks. + + Implementation: + For a given ac ratio p, we should essentially apply ac on every "1/p" + blocks. The first ac block can be as early as the 0th block, or as + late as the "1/p"th block, and we pick the middle one: (0.5p)th block. + Therefore, we are essentially to apply ac on: + (0.5/p)th block, (1.5/p)th block, (2.5/p)th block, etc., and of course, + with these values rounding to integers. + Since ac is applied recursively, we can simply use the following math + in the code to apply ac on corresponding blocks. + """ + block_idx = 0 + cut_off = 1 / 2 + # when passing p as a fraction number (e.g. 1/3), it will be interpreted + # as a string in argv, thus we need eval("1/3") here for fractions. + p = eval(p) if isinstance(p, str) else p + + def selective_checkpointing(submodule): + nonlocal block_idx + nonlocal cut_off + + if isinstance(submodule, block): + block_idx += 1 + if block_idx * p >= cut_off: + cut_off += 1 + return True + return False + + apply_activation_checkpointing( + model, + checkpoint_wrapper_fn=non_reentrant_wrapper, + check_fn=selective_checkpointing, + )