From 799bb5d5818ddbcad16d1112df1f7275bb433f6d Mon Sep 17 00:00:00 2001 From: RunguoLi Date: Tue, 22 Sep 2026 12:41:52 -0500 Subject: [PATCH] [shardformer] fix llama causal mask with sequence parallelism and no flash attention With sequence parallelism the policy replaces the llama attention forward, whose non-flash path adds the attention mask only if it is not None. The model forward built that mask with transformers' `_update_causal_mask`, which returns None for sdpa (the default) in training without padding, relying on sdpa's is_causal. The attention was therefore bidirectional: training ran, but outputs and gradients were wrong (checked for all_to_all and split_gather). With pipeline parallelism on top, later stages passed the local part of the sequence to `_update_causal_mask` together with the full-length cache_position, which failed with "The size of tensor a (12) must match the size of tensor b (24)". When sequence parallelism is on without flash attention, which is exactly when the attention forward is replaced, build the 4d causal mask over the full sequence with `_prepare_4d_causal_attention_mask`. Add Ulysses, Ulysses + PP and TP + split_gather configs without flash attention to the llama test; the existing sequence parallel configs all enable flash attention. --- colossalai/shardformer/modeling/llama.py | 8 ++++ .../test_model/test_shard_llama.py | 40 +++++++++++++++++++ 2 files changed, 48 insertions(+) diff --git a/colossalai/shardformer/modeling/llama.py b/colossalai/shardformer/modeling/llama.py index fe102eecf25a..f99a026fab2e 100644 --- a/colossalai/shardformer/modeling/llama.py +++ b/colossalai/shardformer/modeling/llama.py @@ -8,6 +8,7 @@ from torch import nn from torch.nn import BCEWithLogitsLoss, CrossEntropyLoss, MSELoss from transformers.cache_utils import Cache, DynamicCache +from transformers.modeling_attn_mask_utils import _prepare_4d_causal_attention_mask from transformers.modeling_outputs import ( BaseModelOutputWithPast, CausalLMOutputWithPast, @@ -139,6 +140,13 @@ def llama_model_forward( is_causal=True, invert=(sp_mode != "ring_attn"), ) + elif shard_config.enable_sequence_parallelism: + # the attention forward is replaced for sequence parallelism (see the policy), and it takes a 4d causal + # mask over the full sequence. `_update_causal_mask` returns None for sdpa without padding (relying on + # sdpa's is_causal), and on later pipeline stages `hidden_states` only holds the local part of the sequence + attn_kwargs: torch.Tensor = _prepare_4d_causal_attention_mask( + attention_mask, (batch_size, seq_length), hidden_states, past_seen_tokens + ) else: attn_kwargs: torch.Tensor = self._update_causal_mask( attention_mask, hidden_states, cache_position, past_key_values diff --git a/tests/test_shardformer/test_model/test_shard_llama.py b/tests/test_shardformer/test_model/test_shard_llama.py index b97846408868..78bf24ad5838 100644 --- a/tests/test_shardformer/test_model/test_shard_llama.py +++ b/tests/test_shardformer/test_model/test_shard_llama.py @@ -225,6 +225,46 @@ def check_forward_backward(model_fn, data_gen_fn, output_transform_fn, loss_fn, "precision": "fp16", "initial_scale": 1, }, + # sequence parallelism without flash attention, the replaced attention forward needs a 4d causal mask + { # Ulysess + "tp_size": 1, + "pp_size": 1, + "sp_size": 2, + "num_microbatches": 1, + "enable_sequence_parallelism": True, + "sequence_parallelism_mode": "all_to_all", + "enable_flash_attention": False, + "use_lazy_init": True, + "zero_stage": 0, + "precision": "fp32", + "initial_scale": 1, + }, + { # Ulysess + PP + "tp_size": 1, + "pp_size": 2, + "sp_size": 2, + "num_microbatches": 2, + "enable_sequence_parallelism": True, + "sequence_parallelism_mode": "all_to_all", + "enable_flash_attention": False, + "use_lazy_init": True, + "zero_stage": 1, + "precision": "fp16", + "initial_scale": 1, + }, + { # TP + SP split_gather + "tp_size": 2, + "pp_size": 1, + "sp_size": 1, + "num_microbatches": 1, + "enable_sequence_parallelism": True, + "sequence_parallelism_mode": "split_gather", + "enable_flash_attention": False, + "use_lazy_init": True, + "zero_stage": 0, + "precision": "fp32", + "initial_scale": 1, + }, { "tp_size": 2, "pp_size": 1,