Skip to content

activation_checkpoint

Activation checkpointing policies for the modular training framework.

The modular trainer owns these policies under fastvideo.train, which keeps model plugins within one training package.

Classes

fastvideo.train.utils.activation_checkpoint.CheckpointType

Bases: str, Enum

Supported activation checkpointing policies.

Functions:

fastvideo.train.utils.activation_checkpoint.apply_activation_checkpointing

apply_activation_checkpointing(module: Module, checkpointing_type: str = FULL, n_layer: int = 1) -> Module

Apply the selected activation checkpointing policy to a module.

Source code in fastvideo/train/utils/activation_checkpoint.py
def apply_activation_checkpointing(
    module: torch.nn.Module,
    checkpointing_type: str = CheckpointType.FULL,
    n_layer: int = 1,
) -> torch.nn.Module:
    """Apply the selected activation checkpointing policy to a module."""
    if checkpointing_type == CheckpointType.FULL:
        module = _apply_activation_checkpointing_blocks(module)
    elif checkpointing_type == CheckpointType.OPS:
        # Wrapping each block, not the transformer root, keeps one block's
        # activations live during recompute instead of the whole model's.
        module = _apply_activation_checkpointing_blocks(
            module,
            context_fn=_selective_checkpointing_context_fn,
        )
    elif checkpointing_type == CheckpointType.BLOCK_SKIP:
        module = _apply_activation_checkpointing_blocks(module, n_layer)
    else:
        raise ValueError(f"Checkpointing type '{checkpointing_type}' not supported. "
                         f"Supported types are {CheckpointType.__members__.keys()}")
    return module

fastvideo.train.utils.activation_checkpoint.is_activation_checkpointed

is_activation_checkpointed(module: Module) -> bool

Return whether any submodule runs inside an activation checkpoint.

Source code in fastvideo/train/utils/activation_checkpoint.py
def is_activation_checkpointed(module: torch.nn.Module) -> bool:
    """Return whether any submodule runs inside an activation checkpoint."""
    return any(isinstance(submodule, CheckpointWrapper) for submodule in module.modules())

fastvideo.train.utils.activation_checkpoint.resolve_checkpointing_type

resolve_checkpointing_type(checkpointing_type: str | None, training_config: Any) -> str | None

Return a role's checkpointing type, falling back to training.model.

Source code in fastvideo/train/utils/activation_checkpoint.py
def resolve_checkpointing_type(
    checkpointing_type: str | None,
    training_config: Any,
) -> str | None:
    """Return a role's checkpointing type, falling back to ``training.model``."""
    return checkpointing_type or getattr(
        getattr(training_config, "model", None),
        "enable_gradient_checkpointing_type",
        None,
    )