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
4 changes: 3 additions & 1 deletion src/accelerate/accelerator.py
Original file line number Diff line number Diff line change
Expand Up @@ -34,7 +34,7 @@

from accelerate.utils.dataclasses import FP8BackendType

from .big_modeling import _attach_context_parallel_hooks
from .big_modeling import _attach_context_parallel_hooks, _refuse_recurrent_layers_under_sequence_parallelism
from .checkpointing import load_accelerator_state, load_custom_state, save_accelerator_state, save_custom_state
from .data_loader import DataLoaderDispatcher, prepare_data_loader, skip_first_batches
from .logging import get_logger
Expand Down Expand Up @@ -2405,6 +2405,8 @@ def _prepare_deepspeed(self, *args):
"UlyssesSPAttentionHF currently works with HF Transformers and expects the model object to have a config attribute but this model doesn't have one."
)

_refuse_recurrent_layers_under_sequence_parallelism(model)

kwagrs = {}
signature = inspect.signature(UlyssesSPAttentionHF.register_with_transformers)
if "disable_in_eval" in signature.parameters.keys():
Expand Down
57 changes: 54 additions & 3 deletions src/accelerate/big_modeling.py
Original file line number Diff line number Diff line change
Expand Up @@ -757,23 +757,74 @@ def _attach_layerwise_casting_hooks(
)


def _refuse_recurrent_layers_under_sequence_parallelism(model: nn.Module):
"""Refuse models whose layers carry a recurrent state across the sequence, under Ulysses.

Ulysses gathers the full sequence before attention, so ordinary attention layers are unaffected by the
sharding. Linear-attention layers are not: they carry a recurrent state along the sequence and are
computed inside the model's own layer code rather than through the attention interface Ulysses wraps,
so each rank restarts that state from zero and never exchanges it. The forward output of the first
shard is still correct, which makes the resulting gradients wrong in a way a loss curve does not show.

Args:
model (`nn.Module`):
The model about to be prepared for sequence parallelism.
"""
config = getattr(model, "config", None)
config = config.get_text_config() if hasattr(config, "get_text_config") else config
layer_types = getattr(config, "layer_types", None) or []
recurrent = sorted({layer_type for layer_type in layer_types if "linear_attention" in layer_type})
if recurrent:
raise ValueError(
f"Sequence parallelism does not support attention layers of type {recurrent} (model "
f"{model.__class__.__name__}). Those layers carry a recurrent state along the sequence and "
"bypass the attention interface, so sharding the sequence restarts the state on every rank "
"and produces wrong gradients without failing. Use a full-attention model, or disable "
"sequence parallelism."
)


def _attach_context_parallel_hooks(
model: nn.Module,
):
"""
Monkeypatch huggingface's `transformers` model to fix attention mask issues when using context parallelism.

This function attaches forward_pre_hooks to each self_attn module of the model, where each hook checks the
args/kwargs, if they contain an attention mask, if it does, it will remove this mask, check if it is a causal mask,
if yes, will add a kwarg `is_causal=True`, otherwise will raise an error. This is because context parallelism does
not support attention masks. This function modifies the model in place.
args/kwargs, if they contain an attention mask, if it does, it will remove this mask and add a kwarg
`is_causal=True`. This is because context parallelism does not support attention masks. Models whose layers use a
mask stricter than full causal (sliding-window or chunked attention) are rejected up front, since dropping their
mask would silently train them with full causal attention. This function modifies the model in place.

Args:
model (`nn.Module`):
The model to attach the hooks to.

"""

# The hook below discards the attention mask and forces `is_causal=True`. That is only
# equivalent to the model's own masking for plain causal attention. Models whose layers use
# a *stricter* mask (sliding-window or chunked attention) would otherwise be trained with
# full causal attention silently, so refuse them up front. Without this hook torch raises a
# shape error for such models, so nothing that works today starts failing here.
config = getattr(model, "config", None)
config = config.get_text_config() if hasattr(config, "get_text_config") else config
layer_types = getattr(config, "layer_types", None)
if layer_types is not None:
non_full = {layer_type for layer_type in layer_types if layer_type != "full_attention"}
else:
# Models that predate `layer_types` (Mistral, for one) apply a sliding window to every layer
# whenever `sliding_window` is set.
non_full = {"sliding_attention"} if getattr(config, "sliding_window", None) else set()
if non_full:
raise ValueError(
f"Context parallelism does not support attention layers of type {sorted(non_full)} "
f"(model {model.__class__.__name__}). Context parallelism can only express full causal "
"attention: the per-layer mask has to be dropped, so those layers would silently be "
"trained with full causal attention instead. Use a full-attention model, or disable "
"context parallelism."
)

def _self_attn_pre_forward_hook(_module, module_args, module_kwargs):
if "attention_mask" in module_kwargs:
module_kwargs["attention_mask"] = None
Expand Down
24 changes: 24 additions & 0 deletions tests/test_big_modeling.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,12 +19,14 @@
import unittest
from collections import OrderedDict
from tempfile import TemporaryDirectory
from types import SimpleNamespace

import torch
import torch.nn as nn
from transformers import AutoModelForCausalLM, AutoTokenizer

from accelerate.big_modeling import (
_refuse_recurrent_layers_under_sequence_parallelism,
cpu_offload,
cpu_offload_with_hook,
disk_offload,
Expand Down Expand Up @@ -1111,3 +1113,25 @@ def test_dipatch_model_fp4_simple(self):

assert model.h[0].self_attention.query_key_value.weight.dtype == torch.uint8
assert model.h[0].self_attention.query_key_value.weight.device.index == 0


class RefuseRecurrentLayersTester(unittest.TestCase):
"""Sequence parallelism must refuse layers that carry a recurrent state along the sequence."""

@staticmethod
def _model(layer_types):
config = SimpleNamespace(layer_types=layer_types, get_text_config=lambda: config)
return SimpleNamespace(config=config, __class__=type("FakeModel", (), {}))

def test_refuses_linear_attention(self):
model = self._model(["full_attention", "linear_attention"])
with self.assertRaises(ValueError) as raised:
_refuse_recurrent_layers_under_sequence_parallelism(model)
assert "linear_attention" in str(raised.exception)

def test_allows_full_attention(self):
_refuse_recurrent_layers_under_sequence_parallelism(self._model(["full_attention"]))

def test_allows_models_without_layer_types(self):
config = SimpleNamespace(get_text_config=lambda: config)
_refuse_recurrent_layers_under_sequence_parallelism(SimpleNamespace(config=config))
Loading