From a847f7cf372747bf13166cd4c8605107cf07c51b Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Quentin=20Gallou=C3=A9dec?= Date: Thu, 20 Aug 2026 02:17:27 +0000 Subject: [PATCH 1/3] Refuse context parallelism for models with sliding-window or chunked attention layers --- src/accelerate/big_modeling.py | 25 ++++++++++++++++++++++--- 1 file changed, 22 insertions(+), 3 deletions(-) diff --git a/src/accelerate/big_modeling.py b/src/accelerate/big_modeling.py index d4e579c24e7..0d24d3daf28 100644 --- a/src/accelerate/big_modeling.py +++ b/src/accelerate/big_modeling.py @@ -764,9 +764,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 +775,24 @@ 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) or [] + non_full = {layer_type for layer_type in layer_types if layer_type != "full_attention"} + 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 From 43ebaac4771ede9a04e999f3a9d7743886b4a9a8 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Quentin=20Gallou=C3=A9dec?= Date: Thu, 20 Aug 2026 05:45:06 +0000 Subject: [PATCH 2/3] Also refuse pre-layer_types models that set sliding_window (Mistral) --- src/accelerate/big_modeling.py | 9 +++++++-- 1 file changed, 7 insertions(+), 2 deletions(-) diff --git a/src/accelerate/big_modeling.py b/src/accelerate/big_modeling.py index 0d24d3daf28..fbc12319069 100644 --- a/src/accelerate/big_modeling.py +++ b/src/accelerate/big_modeling.py @@ -782,8 +782,13 @@ def _attach_context_parallel_hooks( # 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) or [] - non_full = {layer_type for layer_type in layer_types if layer_type != "full_attention"} + 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)} " From 4bc2823b50959271ace55772e7a169a27de983f5 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Quentin=20Gallou=C3=A9dec?= Date: Wed, 26 Aug 2026 03:13:09 +0000 Subject: [PATCH 3/3] Refuse linear-attention models under Ulysses, where the recurrent state is never exchanged --- src/accelerate/accelerator.py | 4 +++- src/accelerate/big_modeling.py | 27 +++++++++++++++++++++++++++ tests/test_big_modeling.py | 24 ++++++++++++++++++++++++ 3 files changed, 54 insertions(+), 1 deletion(-) 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 fbc12319069..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, ): 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))