Skip to content
Merged
Show file tree
Hide file tree
Changes from 2 commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
16 changes: 11 additions & 5 deletions kempnerforge/distributed/parallel.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,6 +19,7 @@
import contextlib
import logging
from functools import partial
from typing import TYPE_CHECKING

import torch
from torch.distributed._composable.fsdp import MixedPrecisionPolicy, fully_shard
Expand All @@ -33,6 +34,10 @@
from kempnerforge.config.schema import ActivationCheckpointing
from kempnerforge.model.transformer import Transformer, TransformerBlock

if TYPE_CHECKING:
# Import-cycle-free: pipeline_parallel does not import this module.
from kempnerforge.distributed.pipeline_parallel import PipelineStageModule

logger = logging.getLogger(__name__)


Expand Down Expand Up @@ -182,7 +187,7 @@ def _has_ep_moe(module: torch.nn.Module) -> bool:


def _fsdp_wrap_transformer_blocks(
transformer: Transformer,
transformer: Transformer | PipelineStageModule,
Comment thread
amazloumi marked this conversation as resolved.
dp_mesh: DeviceMesh,
policy: MixedPrecisionPolicy,
reshard_after_forward: bool | int,
Expand Down Expand Up @@ -234,9 +239,10 @@ def _fsdp_wrap_transformer_blocks(
reshard_after_forward=reshard_after_forward,
)

# Cross-Attention blocks (Cross-Attention arch only). Empty dict
# for non-CA configs.
for ca_block in transformer.cross_attention_layers.values():
# Cross-Attention blocks (Cross-Attention arch only). Empty dict for
# non-CA configs, and absent entirely on a PipelineStageModule, which
# defines no cross_attention_layers.
for ca_block in getattr(transformer, "cross_attention_layers", {}).values():
Comment thread
amazloumi marked this conversation as resolved.
fully_shard(
ca_block,
mesh=dp_mesh,
Expand All @@ -248,7 +254,7 @@ def _fsdp_wrap_transformer_blocks(


def apply_fsdp2(
model: Transformer,
model: Transformer | PipelineStageModule,
device_mesh: DeviceMesh,
mp_policy: MixedPrecisionPolicy | None = None,
reshard_after_forward: bool | int = True,
Expand Down
74 changes: 74 additions & 0 deletions tests/unit/test_distributed.py
Original file line number Diff line number Diff line change
Expand Up @@ -508,3 +508,77 @@ def fake_fully_shard(mod, **kwargs): # noqa: ARG001
# Each captured object should be a TransformerBlock, not a sub-module.
for layer in captured:
assert isinstance(layer, TransformerBlock)

def test_pipeline_stage_module_gets_block_level_wrap(self, monkeypatch):
Comment thread
amazloumi marked this conversation as resolved.
Outdated
"""A PipelineStageModule defines no ``cross_attention_layers``.

The helper must wrap its blocks rather than raise AttributeError, which
is what every ``pp > 1`` run hit before the guard was added.
"""
from unittest.mock import MagicMock

import kempnerforge.distributed.parallel as parallel_mod
from kempnerforge.distributed.parallel import (
_fsdp_wrap_transformer_blocks,
default_mp_policy,
)
from kempnerforge.distributed.pipeline_parallel import build_stage_module

captured: list[object] = []

def fake_fully_shard(mod, **kwargs): # noqa: ARG001
captured.append(mod)

monkeypatch.setattr(parallel_mod, "fully_shard", fake_fully_shard)

cfg = ModelConfig(dim=64, n_layers=4, n_heads=4, n_kv_heads=4, vocab_size=256)
stage = build_stage_module(cfg, pp_rank=0, pp_size=2)
assert not hasattr(stage, "cross_attention_layers")

ep_sub = _fsdp_wrap_transformer_blocks(
stage, MagicMock(), default_mp_policy(), reshard_after_forward=False
)

assert ep_sub == 0
assert len(captured) == 2 # this stage owns 2 of the 4 blocks
for layer in captured:
assert isinstance(layer, TransformerBlock)

def test_cross_attention_blocks_are_wrapped(self, monkeypatch):
"""Cross-Attention configs populate ``cross_attention_layers``.

Covers the loop body the guard protects: without a CA model the loop
exits at zero iterations, so nothing ever wraps a CrossAttentionBlock.
"""
from unittest.mock import MagicMock

import kempnerforge.distributed.parallel as parallel_mod
from kempnerforge.config.vlm import CrossAttentionConfig
from kempnerforge.distributed.parallel import (
_fsdp_wrap_transformer_blocks,
default_mp_policy,
)
from kempnerforge.model.cross_attention import CrossAttentionBlock
from kempnerforge.model.transformer import Transformer

captured: list[object] = []

def fake_fully_shard(mod, **kwargs): # noqa: ARG001
captured.append(mod)

monkeypatch.setattr(parallel_mod, "fully_shard", fake_fully_shard)

cfg = ModelConfig(dim=64, n_layers=4, n_heads=4, n_kv_heads=4, vocab_size=256)
transformer = Transformer(
cfg, vlm_config=CrossAttentionConfig(cross_attention_every_n_layers=2)
)
assert len(transformer.cross_attention_layers) == 2

ep_sub = _fsdp_wrap_transformer_blocks(
transformer, MagicMock(), default_mp_policy(), reshard_after_forward=True
)

assert ep_sub == 0
# 4 transformer blocks, then the 2 cross-attention blocks.
assert len(captured) == 6
assert sum(isinstance(m, CrossAttentionBlock) for m in captured) == 2
Loading