Skip to content
Open
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
8 changes: 8 additions & 0 deletions colossalai/shardformer/modeling/llama.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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
Expand Down
40 changes: 40 additions & 0 deletions tests/test_shardformer/test_model/test_shard_llama.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down