diff --git a/src/accelerate/accelerator.py b/src/accelerate/accelerator.py index 3b8829fabfb..59d40fb09c4 100755 --- a/src/accelerate/accelerator.py +++ b/src/accelerate/accelerator.py @@ -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 @@ -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(): diff --git a/src/accelerate/big_modeling.py b/src/accelerate/big_modeling.py index d4e579c24e7..98de9074de5 100644 --- a/src/accelerate/big_modeling.py +++ b/src/accelerate/big_modeling.py @@ -757,6 +757,33 @@ 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, ): @@ -764,9 +791,10 @@ def _attach_context_parallel_hooks( 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`): @@ -774,6 +802,29 @@ def _attach_context_parallel_hooks( """ + # 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 diff --git a/tests/test_big_modeling.py b/tests/test_big_modeling.py index e489c275eb3..2cc865b5e2d 100644 --- a/tests/test_big_modeling.py +++ b/tests/test_big_modeling.py @@ -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, @@ -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))