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,