Skip to content
Merged
Show file tree
Hide file tree
Changes from all 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
22 changes: 15 additions & 7 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 All @@ -199,8 +204,10 @@ def _fsdp_wrap_transformer_blocks(
fires after both complete.

Cross-Attention blocks (when present) are wrapped once each, like
dense ``TransformerBlock``s. The dict is empty for non-CA configs,
so iteration is a no-op on the JD / text-only paths.
dense ``TransformerBlock``s. ``cross_attention_layers`` has three
states, all handled: absent, on a ``PipelineStageModule``, which
defines no such attribute; empty, on the JD / text-only paths; and
populated, on Cross-Attention configs. Only the third wraps anything.

Shared by ``apply_fsdp2`` (for text Transformers) and
``_apply_fsdp_vlm`` (for the inner Transformer of a VLMWrapper) so
Expand Down Expand Up @@ -234,9 +241,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 +256,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
89 changes: 89 additions & 0 deletions tests/unit/test_distributed.py
Original file line number Diff line number Diff line change
Expand Up @@ -508,3 +508,92 @@ 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)


class TestApplyFsdp2GuardedAttributes:
"""Drive the public ``apply_fsdp2`` over the three states of
``cross_attention_layers``: absent, empty, populated.

``fully_shard`` is monkeypatched and the mesh is a mock, so these run on CPU
with no process group. Going through ``apply_fsdp2`` rather than the block
helper is deliberate: it is the call site every run takes, and it also covers
the top-level wrap the helper does not perform.
"""

def _capture(self, monkeypatch):
import kempnerforge.distributed.parallel as parallel_mod

captured: list[object] = []

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

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

def _dp_mesh(self):
from unittest.mock import MagicMock

mesh = MagicMock()
mesh.mesh_dim_names = ("dp_shard",)
return mesh

def test_pipeline_stage_module_is_wrapped(self, monkeypatch):
"""A PipelineStageModule defines no ``cross_attention_layers``.

Before the guard, every ``pp > 1`` run raised AttributeError here.
"""
from kempnerforge.distributed.parallel import apply_fsdp2
from kempnerforge.distributed.pipeline_parallel import build_stage_module

captured = self._capture(monkeypatch)
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")

apply_fsdp2(stage, self._dp_mesh(), reshard_after_forward=False)

# This stage owns 2 of the 4 blocks, then the top-level wrap.
assert len(captured) == 3
assert all(isinstance(m, TransformerBlock) for m in captured[:2])
assert captured[-1] is stage

def test_cross_attention_blocks_are_wrapped(self, monkeypatch):
"""Populated ``cross_attention_layers`` -- the loop body the guard protects.

Without a CA model the loop exits at zero iterations, so nothing ever
wraps a CrossAttentionBlock.
"""
from kempnerforge.config.vlm import CrossAttentionConfig
from kempnerforge.distributed.parallel import apply_fsdp2
from kempnerforge.model.cross_attention import CrossAttentionBlock
from kempnerforge.model.transformer import Transformer

captured = self._capture(monkeypatch)
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

apply_fsdp2(transformer, self._dp_mesh())

# 4 transformer blocks, 2 cross-attention blocks, then the top-level wrap.
assert len(captured) == 7
assert sum(isinstance(m, CrossAttentionBlock) for m in captured) == 2
assert captured[-1] is transformer

def test_empty_cross_attention_dict_wraps_nothing_extra(self, monkeypatch):
"""A plain Transformer has the attribute but it is empty."""
from kempnerforge.distributed.parallel import apply_fsdp2
from kempnerforge.model.transformer import Transformer

captured = self._capture(monkeypatch)
cfg = ModelConfig(dim=64, n_layers=2, n_heads=4, n_kv_heads=4, vocab_size=256)
transformer = Transformer(cfg)
assert len(transformer.cross_attention_layers) == 0

apply_fsdp2(transformer, self._dp_mesh())

assert len(captured) == 3 # 2 blocks + top-level, no CA wraps
assert captured[-1] is transformer
Loading