Compare commits
1
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
4297b0e776 |
@@ -20,7 +20,8 @@ from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines import ComposedPipelineBase
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.v1.training.training_utils import (
|
||||
compute_density_for_timestep_sampling, get_sigmas, normalize_dit_input)
|
||||
compute_density_for_timestep_sampling, get_sigmas, normalize_dit_input,
|
||||
apply_activation_checkpointing)
|
||||
|
||||
import wandb # isort: skip
|
||||
|
||||
@@ -52,6 +53,16 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
|
||||
self.transformer = self.get_module("transformer")
|
||||
assert self.transformer is not None
|
||||
|
||||
# if training_args.gradient_checkpointing:
|
||||
if True:
|
||||
# with set_forward_context(current_timestep=0,
|
||||
# attn_metadata=None):
|
||||
logger.info(f"Enabling activation-checkpointing ")
|
||||
self.transformer = apply_activation_checkpointing(
|
||||
self.transformer,
|
||||
checkpointing_type='full',
|
||||
)
|
||||
|
||||
self.transformer.requires_grad_(True)
|
||||
self.transformer.train()
|
||||
|
||||
|
||||
@@ -8,8 +8,13 @@ import torch.distributed as dist
|
||||
from torch.distributed.fsdp import FullStateDictConfig
|
||||
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
|
||||
from torch.distributed.fsdp import StateDictType
|
||||
|
||||
from fastvideo.v1.forward_context import get_forward_context, set_forward_context
|
||||
import contextlib
|
||||
from fastvideo.v1.logger import init_logger
|
||||
import collections
|
||||
from enum import Enum
|
||||
from torch.distributed.algorithms._checkpoint.checkpoint_wrapper import checkpoint_wrapper
|
||||
from torch.utils.checkpoint import CheckpointPolicy, create_selective_checkpoint_contexts
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@@ -323,3 +328,109 @@ def _has_foreach_support(tensors: List[torch.Tensor],
|
||||
device: torch.device) -> bool:
|
||||
return _device_has_foreach_support(device) and all(
|
||||
t is None or type(t) in [torch.Tensor] for t in tensors)
|
||||
|
||||
DIFFUSERS_TRANSFORMER_BLOCK_NAMES = [
|
||||
"transformer_blocks",
|
||||
"single_transformer_blocks",
|
||||
"temporal_transformer_blocks",
|
||||
"blocks",
|
||||
"layers",
|
||||
]
|
||||
|
||||
def debug_tensor_creation():
|
||||
original_new = torch.Tensor.__new__
|
||||
|
||||
def new_with_debug(cls, *args, **kwargs):
|
||||
tensor = original_new(cls, *args, **kwargs)
|
||||
if tensor.dtype == torch.float32:
|
||||
import traceback
|
||||
print(f"FP32 tensor created during recompute: {tensor.shape}")
|
||||
traceback.print_stack(limit=5)
|
||||
return tensor
|
||||
|
||||
torch.Tensor.__new__ = new_with_debug
|
||||
return lambda: setattr(torch.Tensor, '__new__', original_new)
|
||||
|
||||
def _ctx_fn():
|
||||
saved = get_forward_context() # ← runs while ctx is alive
|
||||
|
||||
@contextlib.contextmanager
|
||||
def _cm():
|
||||
with torch.amp.autocast('cuda', enabled=True, dtype=torch.bfloat16):
|
||||
if saved is None:
|
||||
yield
|
||||
else:
|
||||
# current_timestep = saved.current_timestep
|
||||
# if hasattr(current_timestep, 'dtype') and current_timestep.dtype == torch.float32:
|
||||
# current_timestep = current_timestep.to(torch.bfloat16)
|
||||
with set_forward_context(
|
||||
current_timestep=saved.current_timestep.to(torch.bfloat16),
|
||||
attn_metadata=saved.attn_metadata,
|
||||
forward_batch=saved.forward_batch,
|
||||
):
|
||||
yield
|
||||
|
||||
return _cm(), _cm()
|
||||
|
||||
class CheckpointType(str, Enum):
|
||||
FULL = "full"
|
||||
OPS = "ops"
|
||||
BLOCK_SKIP = "block_skip"
|
||||
|
||||
|
||||
_SELECTIVE_ACTIVATION_CHECKPOINTING_OPS = {
|
||||
torch.ops.aten.mm.default,
|
||||
torch.ops.aten._scaled_dot_product_efficient_attention.default,
|
||||
torch.ops.aten._scaled_dot_product_flash_attention.default,
|
||||
torch.ops._c10d_functional.reduce_scatter_tensor.default,
|
||||
}
|
||||
|
||||
|
||||
def apply_activation_checkpointing(
|
||||
module: torch.nn.Module, checkpointing_type: str = CheckpointType.FULL, n_layer: int = 1
|
||||
) -> torch.nn.Module:
|
||||
if checkpointing_type == CheckpointType.FULL:
|
||||
module = _apply_activation_checkpointing_blocks(module)
|
||||
elif checkpointing_type == CheckpointType.OPS:
|
||||
module = _apply_activation_checkpointing_ops(module, _SELECTIVE_ACTIVATION_CHECKPOINTING_OPS)
|
||||
elif checkpointing_type == CheckpointType.BLOCK_SKIP:
|
||||
module = _apply_activation_checkpointing_blocks(module, n_layer)
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Checkpointing type '{checkpointing_type}' not supported. Supported types are {CheckpointType.__members__.keys()}"
|
||||
)
|
||||
return module
|
||||
|
||||
|
||||
def _apply_activation_checkpointing_blocks(module: torch.nn.Module, n_layer: int = None) -> torch.nn.Module:
|
||||
for transformer_block_name in DIFFUSERS_TRANSFORMER_BLOCK_NAMES:
|
||||
blocks: torch.nn.Module = getattr(module, transformer_block_name, None)
|
||||
if blocks is None:
|
||||
continue
|
||||
print("here")
|
||||
for index, (layer_id, block) in enumerate(blocks.named_children()):
|
||||
if n_layer is None or index % n_layer == 0:
|
||||
block = checkpoint_wrapper(block, preserve_rng_state=False, context_fn=_ctx_fn)
|
||||
blocks.register_module(layer_id, block)
|
||||
return module
|
||||
|
||||
|
||||
def _apply_activation_checkpointing_ops(module: torch.nn.Module, ops) -> torch.nn.Module:
|
||||
|
||||
def _get_custom_policy(meta):
|
||||
def _custom_policy(ctx, func, *args, **kwargs):
|
||||
mode = "recompute" if ctx.is_recompute else "forward"
|
||||
mm_count_key = f"{mode}_mm_count"
|
||||
if func == torch.ops.aten.mm.default:
|
||||
meta[mm_count_key] += 1
|
||||
# Saves output of all compute ops, except every second mm
|
||||
to_save = func in ops and not (func == torch.ops.aten.mm.default and meta[mm_count_key] % 2 == 0)
|
||||
return CheckpointPolicy.MUST_SAVE if to_save else CheckpointPolicy.PREFER_RECOMPUTE
|
||||
|
||||
return _custom_policy
|
||||
|
||||
def selective_checkpointing_context_fn():
|
||||
meta = collections.defaultdict(int)
|
||||
return create_selective_checkpoint_contexts(_get_custom_policy(meta))
|
||||
|
||||
return checkpoint_wrapper(module, context_fn=selective_checkpointing_context_fn, preserve_rng_state=False)
|
||||
Reference in New Issue
Block a user