From 5162e1057c525b66a057c86510c713e67e007c08 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Bj=C3=B6rn=20Buschk=C3=A4mper?= Date: Mon, 29 Jun 2026 12:40:13 +0000 Subject: [PATCH 1/4] Add support for packed thd in BERT language module. MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Signed-off-by: Björn Buschkämper --- megatron/core/models/bert/bert_model.py | 39 +++++++++++--- tests/unit_tests/models/test_bert_model.py | 62 ++++++++++++++++++++++ 2 files changed, 95 insertions(+), 6 deletions(-) diff --git a/megatron/core/models/bert/bert_model.py b/megatron/core/models/bert/bert_model.py index 3fd1e01f4a1..7577152386f 100644 --- a/megatron/core/models/bert/bert_model.py +++ b/megatron/core/models/bert/bert_model.py @@ -14,6 +14,7 @@ from megatron.core.models.common.embeddings.language_model_embedding import LanguageModelEmbedding from megatron.core.models.common.embeddings.rotary_pos_embedding import RotaryEmbedding from megatron.core.models.common.language_module.language_module import LanguageModule +from megatron.core.packed_seq_params import PackedSeqParams from megatron.core.process_groups_config import ProcessGroupCollection from megatron.core.transformer.attention import SelfAttentionSubmodules from megatron.core.transformer.dot_product_attention import ( @@ -269,8 +270,23 @@ def bert_extended_attention_mask(self, attention_mask: Tensor) -> Tensor: return extended_attention_mask - def bert_position_ids(self, token_ids): + def bert_position_ids( + self, token_ids: Tensor, packed_seq_params: PackedSeqParams | None = None + ) -> Tensor: """Position ids for bert model""" + if packed_seq_params is not None: + assert token_ids.size(0) == 1, 'Packed BERT input should use dummy batch size 1' + assert packed_seq_params.cu_seqlens_q is not None, ( + 'packed_seq_params.cu_seqlens_q must be provided for packed BERT input' + ) + cu_seqlens = packed_seq_params.cu_seqlens_q.to(device=token_ids.device) + seq_lengths = (cu_seqlens[1:] - cu_seqlens[:-1]).to(torch.long) + seq_starts = torch.repeat_interleave(cu_seqlens[:-1].to(torch.long), seq_lengths) + token_positions = torch.arange( + token_ids.size(1), dtype=torch.long, device=token_ids.device + ) + return (token_positions - seq_starts).unsqueeze(0) + # Create position ids seq_length = token_ids.size(1) position_ids = torch.arange(seq_length, dtype=torch.long, device=token_ids.device) @@ -297,10 +313,11 @@ def set_input_tensor(self, input_tensor: Tensor) -> None: def forward( self, input_ids: Tensor, - attention_mask: Tensor, + attention_mask: Tensor | None, tokentype_ids: Tensor = None, lm_labels: Tensor = None, inference_context=None, + packed_seq_params: PackedSeqParams | None = None, *, inference_params: Optional[BaseInferenceContext] = None, ): @@ -315,11 +332,15 @@ def forward( inference_context = deprecate_inference_params(inference_context, inference_params) - extended_attention_mask = self.bert_extended_attention_mask(attention_mask) + extended_attention_mask = ( + None + if packed_seq_params is not None + else self.bert_extended_attention_mask(attention_mask) + ) if parallel_state.is_pipeline_first_stage(): input_ids = input_ids - position_ids = self.bert_position_ids(input_ids) + position_ids = self.bert_position_ids(input_ids, packed_seq_params) else: position_ids = None input_ids = None @@ -338,9 +359,14 @@ def forward( rotary_pos_emb = None if self.position_embedding_type == 'rope': rotary_seq_len = self.rotary_pos_emb.get_rotary_seq_len( - inference_context, self.encoder, encoder_input, self.config + inference_context, self.encoder, encoder_input, self.config, packed_seq_params + ) + rotary_pos_emb = self.rotary_pos_emb( + rotary_seq_len, + packed_seq=packed_seq_params is not None + and packed_seq_params.qkv_format == 'thd', + cp_group=packed_seq_params.cp_group if packed_seq_params is not None else None, ) - rotary_pos_emb = self.rotary_pos_emb(rotary_seq_len) # Run encoder. hidden_states = self.encoder( @@ -348,6 +374,7 @@ def forward( attention_mask=extended_attention_mask, inference_context=inference_context, rotary_pos_emb=rotary_pos_emb, + packed_seq_params=packed_seq_params, ) if not self.post_process: return hidden_states diff --git a/tests/unit_tests/models/test_bert_model.py b/tests/unit_tests/models/test_bert_model.py index db7b8255776..782eed85e04 100644 --- a/tests/unit_tests/models/test_bert_model.py +++ b/tests/unit_tests/models/test_bert_model.py @@ -12,6 +12,7 @@ get_bert_layer_with_transformer_engine_submodules, ) from megatron.core.models.bert.bert_model import BertModel +from megatron.core.packed_seq_params import PackedSeqParams from megatron.core.tensor_parallel.random import model_parallel_cuda_manual_seed from megatron.core.transformer.enums import AttnBackend, AttnMaskType from megatron.core.transformer.spec_utils import ModuleSpec @@ -92,6 +93,67 @@ def test_post_process_forward(self): assert logits[0].shape[1] == sequence_length assert logits[0].shape[2] == self.bert_model.vocab_size + @pytest.mark.internal + def test_packed_forward_uses_cu_seqlens_positions_and_no_attention_mask(self, mocker): + config = TransformerConfig( + num_layers=2, + hidden_size=12, + num_attention_heads=4, + use_cpu_initialization=True, + perform_initialization=True, + attention_backend=AttnBackend.unfused, + ) + bert_model = BertModel( + config=config, + num_tokentypes=0, + transformer_layer_spec=get_bert_layer_with_transformer_engine_spec(), + vocab_size=100, + max_sequence_length=4, + position_embedding_type='rope', + post_process=False, + ) + sequence_length = 6 + input_ids = torch.arange(sequence_length, dtype=torch.int64).unsqueeze(0) + cu_seqlens = torch.tensor([0, 2, 6], dtype=torch.int32) + packed_seq_params = PackedSeqParams( + qkv_format='thd', + cu_seqlens_q=cu_seqlens, + cu_seqlens_kv=cu_seqlens, + max_seqlen_q=2, + max_seqlen_kv=4, + ) + encoder_input = torch.ones(sequence_length, 1, config.hidden_size) + hidden_states = torch.zeros_like(encoder_input) + rotary_pos_emb = torch.ones(4, 1, 1, config.hidden_size // config.num_attention_heads) + + extended_attention_mask = mocker.patch.object(bert_model, 'bert_extended_attention_mask') + embedding_forward = mocker.patch.object( + bert_model.embedding, 'forward', return_value=encoder_input + ) + rotary_forward = mocker.patch.object( + bert_model.rotary_pos_emb, 'forward', return_value=rotary_pos_emb + ) + encoder_forward = mocker.patch.object( + bert_model.encoder, 'forward', return_value=hidden_states + ) + + output = bert_model.forward( + input_ids=input_ids, attention_mask=None, packed_seq_params=packed_seq_params + ) + + extended_attention_mask.assert_not_called() + assert torch.equal( + embedding_forward.call_args.kwargs['position_ids'], + torch.tensor([[0, 1, 0, 1, 2, 3]], dtype=torch.int64), + ) + rotary_forward.assert_called_once_with( + 4, packed_seq=True, cp_group=packed_seq_params.cp_group + ) + assert encoder_forward.call_args.kwargs['attention_mask'] is None + assert encoder_forward.call_args.kwargs['packed_seq_params'] is packed_seq_params + assert encoder_forward.call_args.kwargs['rotary_pos_emb'] is rotary_pos_emb + assert output is hidden_states + class TestBertModelAttentionDimensions: From cb70081b75bdab0d4325ef0bfc710ecd53c22175 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Bj=C3=B6rn=20Buschk=C3=A4mper?= Date: Mon, 29 Jun 2026 13:11:59 +0000 Subject: [PATCH 2/4] Format files using autoformat.sh. MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Signed-off-by: Björn Buschkämper --- megatron/core/models/bert/bert_model.py | 9 ++++----- 1 file changed, 4 insertions(+), 5 deletions(-) diff --git a/megatron/core/models/bert/bert_model.py b/megatron/core/models/bert/bert_model.py index 7577152386f..8b9e7253987 100644 --- a/megatron/core/models/bert/bert_model.py +++ b/megatron/core/models/bert/bert_model.py @@ -276,9 +276,9 @@ def bert_position_ids( """Position ids for bert model""" if packed_seq_params is not None: assert token_ids.size(0) == 1, 'Packed BERT input should use dummy batch size 1' - assert packed_seq_params.cu_seqlens_q is not None, ( - 'packed_seq_params.cu_seqlens_q must be provided for packed BERT input' - ) + assert ( + packed_seq_params.cu_seqlens_q is not None + ), 'packed_seq_params.cu_seqlens_q must be provided for packed BERT input' cu_seqlens = packed_seq_params.cu_seqlens_q.to(device=token_ids.device) seq_lengths = (cu_seqlens[1:] - cu_seqlens[:-1]).to(torch.long) seq_starts = torch.repeat_interleave(cu_seqlens[:-1].to(torch.long), seq_lengths) @@ -363,8 +363,7 @@ def forward( ) rotary_pos_emb = self.rotary_pos_emb( rotary_seq_len, - packed_seq=packed_seq_params is not None - and packed_seq_params.qkv_format == 'thd', + packed_seq=packed_seq_params is not None and packed_seq_params.qkv_format == 'thd', cp_group=packed_seq_params.cp_group if packed_seq_params is not None else None, ) From ab4aefc631780ea5c08fdea1ae03d0d16d31f22d Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Bj=C3=B6rn=20Buschk=C3=A4mper?= Date: Mon, 29 Jun 2026 13:32:45 +0000 Subject: [PATCH 3/4] Tighten validation on attention mask and input tensor dim. MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Signed-off-by: Björn Buschkämper --- megatron/core/models/bert/bert_model.py | 35 +++++++++--- tests/unit_tests/models/test_bert_model.py | 65 ++++++++++++++++++++++ 2 files changed, 91 insertions(+), 9 deletions(-) diff --git a/megatron/core/models/bert/bert_model.py b/megatron/core/models/bert/bert_model.py index 8b9e7253987..d5e075a05f9 100644 --- a/megatron/core/models/bert/bert_model.py +++ b/megatron/core/models/bert/bert_model.py @@ -275,12 +275,24 @@ def bert_position_ids( ) -> Tensor: """Position ids for bert model""" if packed_seq_params is not None: - assert token_ids.size(0) == 1, 'Packed BERT input should use dummy batch size 1' - assert ( - packed_seq_params.cu_seqlens_q is not None - ), 'packed_seq_params.cu_seqlens_q must be provided for packed BERT input' + if token_ids.size(0) != 1: + raise ValueError('Packed BERT input should use dummy batch size 1') + if packed_seq_params.cu_seqlens_q is None: + raise ValueError( + 'packed_seq_params.cu_seqlens_q must be provided for packed BERT input' + ) cu_seqlens = packed_seq_params.cu_seqlens_q.to(device=token_ids.device) + if cu_seqlens.dim() != 1 or cu_seqlens.numel() < 2: + raise ValueError('packed_seq_params.cu_seqlens_q must be a 1D tensor') + if cu_seqlens[0].item() != 0: + raise ValueError('packed_seq_params.cu_seqlens_q must start at 0') + if cu_seqlens[-1].item() != token_ids.size(1): + raise ValueError( + 'packed_seq_params.cu_seqlens_q must end at the packed sequence length' + ) seq_lengths = (cu_seqlens[1:] - cu_seqlens[:-1]).to(torch.long) + if torch.any(seq_lengths < 0): + raise ValueError('packed_seq_params.cu_seqlens_q must be monotonically increasing') seq_starts = torch.repeat_interleave(cu_seqlens[:-1].to(torch.long), seq_lengths) token_positions = torch.arange( token_ids.size(1), dtype=torch.long, device=token_ids.device @@ -332,11 +344,16 @@ def forward( inference_context = deprecate_inference_params(inference_context, inference_params) - extended_attention_mask = ( - None - if packed_seq_params is not None - else self.bert_extended_attention_mask(attention_mask) - ) + if packed_seq_params is None: + if attention_mask is None: + raise ValueError('attention_mask must be provided when packed_seq_params is None') + extended_attention_mask = self.bert_extended_attention_mask(attention_mask) + else: + if packed_seq_params.qkv_format != 'thd': + raise ValueError("BERT packed sequence support requires qkv_format='thd'") + if attention_mask is not None: + raise ValueError('attention_mask must be None when using packed BERT input') + extended_attention_mask = None if parallel_state.is_pipeline_first_stage(): input_ids = input_ids diff --git a/tests/unit_tests/models/test_bert_model.py b/tests/unit_tests/models/test_bert_model.py index 782eed85e04..3455a287856 100644 --- a/tests/unit_tests/models/test_bert_model.py +++ b/tests/unit_tests/models/test_bert_model.py @@ -154,6 +154,71 @@ def test_packed_forward_uses_cu_seqlens_positions_and_no_attention_mask(self, mo assert encoder_forward.call_args.kwargs['rotary_pos_emb'] is rotary_pos_emb assert output is hidden_states + @pytest.mark.internal + def test_forward_validates_dense_attention_mask_and_packed_format(self): + sequence_length = self.bert_model.max_sequence_length + input_ids = torch.arange(sequence_length, dtype=torch.int64).unsqueeze(0) + attention_mask = torch.ones((1, sequence_length), dtype=bool) + cu_seqlens = torch.tensor([0, sequence_length], dtype=torch.int32) + + with pytest.raises(ValueError, match='attention_mask must be provided'): + self.bert_model.forward(input_ids=input_ids, attention_mask=None) + + with pytest.raises(ValueError, match='qkv_format'): + self.bert_model.forward( + input_ids=input_ids, + attention_mask=None, + packed_seq_params=PackedSeqParams( + qkv_format='sbhd', + cu_seqlens_q=cu_seqlens, + cu_seqlens_kv=cu_seqlens, + max_seqlen_q=sequence_length, + max_seqlen_kv=sequence_length, + ), + ) + + with pytest.raises(ValueError, match='attention_mask must be None'): + self.bert_model.forward( + input_ids=input_ids, + attention_mask=attention_mask, + packed_seq_params=PackedSeqParams( + qkv_format='thd', + cu_seqlens_q=cu_seqlens, + cu_seqlens_kv=cu_seqlens, + max_seqlen_q=sequence_length, + max_seqlen_kv=sequence_length, + ), + ) + + @pytest.mark.internal + @pytest.mark.parametrize( + ('token_ids', 'cu_seqlens', 'error_match'), + [ + ( + torch.arange(4, dtype=torch.int64).repeat((2, 1)), + torch.tensor([0, 4]), + 'dummy batch', + ), + ( + torch.arange(4, dtype=torch.int64).unsqueeze(0), + None, + 'cu_seqlens_q must be provided', + ), + (torch.arange(4, dtype=torch.int64).unsqueeze(0), torch.tensor([1, 4]), 'start at 0'), + (torch.arange(4, dtype=torch.int64).unsqueeze(0), torch.tensor([0, 3]), 'end at'), + ( + torch.arange(4, dtype=torch.int64).unsqueeze(0), + torch.tensor([0, 3, 2, 4]), + 'monotonically', + ), + ], + ) + def test_packed_position_ids_validate_cu_seqlens(self, token_ids, cu_seqlens, error_match): + with pytest.raises(ValueError, match=error_match): + self.bert_model.bert_position_ids( + token_ids, PackedSeqParams(qkv_format='thd', cu_seqlens_q=cu_seqlens) + ) + class TestBertModelAttentionDimensions: From 1513676aa50f15cfc32d2cce6ce79b10481fe267 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Bj=C3=B6rn=20Buschk=C3=A4mper?= Date: Mon, 27 Jul 2026 11:45:51 +0000 Subject: [PATCH 4/4] Fix reviewer suggestions for packed thd in bert language module: - Build position ids from physical padded boundaries - One logit for whole pack in binary head pool - return_embeddings crashes on packed input - .item() call forces sync, unsafe MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Signed-off-by: Björn Buschkämper --- megatron/core/models/bert/bert_model.py | 204 ++++++++++++++-- tests/unit_tests/models/test_bert_model.py | 262 ++++++++++++++++++++- 2 files changed, 439 insertions(+), 27 deletions(-) diff --git a/megatron/core/models/bert/bert_model.py b/megatron/core/models/bert/bert_model.py index d5e075a05f9..a0f041b7a68 100644 --- a/megatron/core/models/bert/bert_model.py +++ b/megatron/core/models/bert/bert_model.py @@ -26,7 +26,7 @@ from megatron.core.transformer.transformer_config import TransformerConfig from megatron.core.transformer.transformer_layer import TransformerLayerSubmodules from megatron.core.transformer.utils import get_linear_layer -from megatron.core.utils import deprecate_inference_params, is_te_min_version +from megatron.core.utils import deprecate_inference_params, get_pg_size, is_te_min_version class BertModel(LanguageModule): @@ -78,7 +78,7 @@ def __init__( log_config_to_disk(config, locals(), prefix=type(self).__name__) if return_embeddings: - assert self.post_process and self.add_binary_head + assert post_process and add_binary_head self.config: TransformerConfig = config self.transformer_layer_spec: ModuleSpec = transformer_layer_spec @@ -275,28 +275,18 @@ def bert_position_ids( ) -> Tensor: """Position ids for bert model""" if packed_seq_params is not None: - if token_ids.size(0) != 1: - raise ValueError('Packed BERT input should use dummy batch size 1') - if packed_seq_params.cu_seqlens_q is None: + _, physical_cu_seqlens = self._get_packed_cu_seqlens(token_ids, packed_seq_params) + if self._get_packed_cp_size(packed_seq_params) > 1: raise ValueError( - 'packed_seq_params.cu_seqlens_q must be provided for packed BERT input' + 'position_ids must be provided for packed BERT input with context parallelism' ) - cu_seqlens = packed_seq_params.cu_seqlens_q.to(device=token_ids.device) - if cu_seqlens.dim() != 1 or cu_seqlens.numel() < 2: - raise ValueError('packed_seq_params.cu_seqlens_q must be a 1D tensor') - if cu_seqlens[0].item() != 0: - raise ValueError('packed_seq_params.cu_seqlens_q must start at 0') - if cu_seqlens[-1].item() != token_ids.size(1): - raise ValueError( - 'packed_seq_params.cu_seqlens_q must end at the packed sequence length' - ) - seq_lengths = (cu_seqlens[1:] - cu_seqlens[:-1]).to(torch.long) - if torch.any(seq_lengths < 0): - raise ValueError('packed_seq_params.cu_seqlens_q must be monotonically increasing') - seq_starts = torch.repeat_interleave(cu_seqlens[:-1].to(torch.long), seq_lengths) - token_positions = torch.arange( - token_ids.size(1), dtype=torch.long, device=token_ids.device + + total_tokens = token_ids.size(1) + segment_lengths = self._get_packed_segment_lengths(physical_cu_seqlens, total_tokens) + seq_starts = torch.repeat_interleave( + physical_cu_seqlens, segment_lengths, output_size=total_tokens ) + token_positions = torch.arange(total_tokens, dtype=torch.long, device=token_ids.device) return (token_positions - seq_starts).unsqueeze(0) # Create position ids @@ -306,6 +296,145 @@ def bert_position_ids( return position_ids + def _get_packed_cu_seqlens( + self, token_ids: Tensor, packed_seq_params: PackedSeqParams + ) -> tuple[Tensor, Tensor]: + """Validate packed input metadata and return logical and physical boundaries.""" + if token_ids.size(0) != 1: + raise ValueError('Packed BERT input should use dummy batch size 1') + if packed_seq_params.cu_seqlens_q is None: + raise ValueError( + 'packed_seq_params.cu_seqlens_q must be provided for packed BERT input' + ) + + cu_seqlens = packed_seq_params.cu_seqlens_q.to(device=token_ids.device, dtype=torch.long) + physical_cu_seqlens = self._get_packed_physical_cu_seqlens( + packed_seq_params, token_ids.device + ) + if cu_seqlens.dim() != 1 or cu_seqlens.numel() < 2: + raise ValueError('packed_seq_params.cu_seqlens_q must be a 1D tensor') + if physical_cu_seqlens.dim() != 1 or physical_cu_seqlens.numel() != cu_seqlens.numel(): + raise ValueError('packed_seq_params.cu_seqlens_q_padded must match cu_seqlens_q') + + # Reading cumulative lengths that live on the accelerator would force a + # device-to-host synchronization on every forward and is illegal during CUDA graph + # capture, so metadata is only validated while it is still on the host. + if physical_cu_seqlens.device.type == 'cpu': + cu_seqlens_values = cu_seqlens.tolist() + physical_cu_seqlens_values = physical_cu_seqlens.tolist() + if cu_seqlens_values[0] != 0 or physical_cu_seqlens_values[0] != 0: + raise ValueError('packed sequence cumulative lengths must start at 0') + if any( + end < start for start, end in zip(cu_seqlens_values[:-1], cu_seqlens_values[1:]) + ) or any( + end < start + for start, end in zip( + physical_cu_seqlens_values[:-1], physical_cu_seqlens_values[1:] + ) + ): + raise ValueError( + 'packed sequence cumulative lengths must be monotonically increasing' + ) + if self._get_packed_cp_size(packed_seq_params) == 1 and physical_cu_seqlens_values[ + -1 + ] > token_ids.size(1): + raise ValueError( + 'packed physical cumulative lengths must not exceed the packed sequence length' + ) + + return cu_seqlens, physical_cu_seqlens + + @staticmethod + def _get_packed_physical_cu_seqlens( + packed_seq_params: PackedSeqParams, device: torch.device + ) -> Tensor: + """Return the boundaries of the physical (padded) packed token buffer.""" + cu_seqlens = ( + packed_seq_params.cu_seqlens_q_padded + if packed_seq_params.cu_seqlens_q_padded is not None + else packed_seq_params.cu_seqlens_q + ) + return cu_seqlens.to(device=device, dtype=torch.long) + + @staticmethod + def _get_packed_segment_lengths(physical_cu_seqlens: Tensor, total_tokens: int) -> Tensor: + """Split the physical token buffer into per-sequence segments. + + Tokens beyond the last padded sequence are returned as one extra trailing segment, + matching :meth:`PackedSeqParams.__post_init__`. The trailing segment is empty when + the buffer ends exactly at the last padded boundary. + """ + boundaries = torch.cat( + [ + physical_cu_seqlens, + torch.tensor( + [total_tokens], + dtype=physical_cu_seqlens.dtype, + device=physical_cu_seqlens.device, + ), + ] + ) + return (boundaries[1:] - boundaries[:-1]).clamp(min=0) + + def _packed_mean_pooled_embeddings( + self, hidden_states: Tensor, packed_seq_params: PackedSeqParams + ) -> Tensor: + """Mean-pool the interior tokens of every sequence in a packed buffer. + + Mirrors the dense ``return_embeddings`` path, which averages ``embedding[1 : mask - 1]`` + per sample, but derives the per-sequence spans from the cumulative sequence lengths. + Padding tokens are excluded because the logical lengths bound the interior mask. + """ + physical_cu_seqlens = self._get_packed_physical_cu_seqlens( + packed_seq_params, hidden_states.device + ) + total_tokens = hidden_states.size(0) + num_sequences = physical_cu_seqlens.numel() - 1 + segment_lengths = self._get_packed_segment_lengths(physical_cu_seqlens, total_tokens) + segment_ids = torch.repeat_interleave( + torch.arange(segment_lengths.numel(), device=hidden_states.device), + segment_lengths, + output_size=total_tokens, + ) + positions = torch.arange( + total_tokens, device=hidden_states.device + ) - physical_cu_seqlens.index_select(0, segment_ids) + + cu_seqlens = packed_seq_params.cu_seqlens_q.to( + device=hidden_states.device, dtype=torch.long + ) + # The trailing segment holds tokens past the last sequence, so its length is zero. + sequence_lengths = torch.cat( + [ + cu_seqlens[1:] - cu_seqlens[:-1], + torch.zeros(1, dtype=torch.long, device=hidden_states.device), + ] + ) + interior_mask = (positions > 0) & ( + positions < sequence_lengths.index_select(0, segment_ids) - 1 + ) + + embeddings = hidden_states[:, 0, :].float() + output = torch.zeros( + (segment_lengths.numel(), embeddings.size(1)), + dtype=torch.float32, + device=hidden_states.device, + ) + output.index_add_(0, segment_ids, embeddings * interior_mask.unsqueeze(1)) + counts = torch.zeros( + segment_lengths.numel(), dtype=torch.float32, device=hidden_states.device + ) + counts.index_add_(0, segment_ids, interior_mask.float()) + return (output / counts.clamp_min(1).unsqueeze(1))[:num_sequences] + + def _get_packed_cp_size(self, packed_seq_params: PackedSeqParams) -> int: + """Return the effective context-parallel size for packed input.""" + if packed_seq_params.cp_group is not None: + return get_pg_size(packed_seq_params.cp_group) + if packed_seq_params.local_cp_size is not None: + return packed_seq_params.local_cp_size + return get_pg_size(self.cp_group) + def set_input_tensor(self, input_tensor: Tensor) -> None: """Sets input tensor to the model. @@ -331,6 +460,7 @@ def forward( inference_context=None, packed_seq_params: PackedSeqParams | None = None, *, + position_ids: Tensor = None, inference_params: Optional[BaseInferenceContext] = None, ): """Forward function of BERT model @@ -353,11 +483,30 @@ def forward( raise ValueError("BERT packed sequence support requires qkv_format='thd'") if attention_mask is not None: raise ValueError('attention_mask must be None when using packed BERT input') + if ( + self.post_process + and (self.add_binary_head or self.return_embeddings) + and self._get_packed_cp_size(packed_seq_params) > 1 + ): + raise ValueError( + 'Packed BERT binary-head and embedding post-processing do not support ' + 'context parallelism' + ) + if self.post_process and self.return_embeddings and self.config.sequence_parallel: + raise ValueError( + 'Packed BERT embedding post-processing does not support sequence parallelism' + ) extended_attention_mask = None if parallel_state.is_pipeline_first_stage(): input_ids = input_ids - position_ids = self.bert_position_ids(input_ids, packed_seq_params) + if position_ids is None: + position_ids = self.bert_position_ids(input_ids, packed_seq_params) + else: + if packed_seq_params is not None: + self._get_packed_cu_seqlens(input_ids, packed_seq_params) + if position_ids.shape != input_ids.shape: + raise ValueError('position_ids must have the same shape as input_ids') else: position_ids = None input_ids = None @@ -396,11 +545,20 @@ def forward( return hidden_states if self.add_binary_head: - pooled_output = self.pooler(hidden_states, 0) + if packed_seq_params is None: + pooled_output = self.pooler(hidden_states, 0) + else: + sequence_starts = self._get_packed_physical_cu_seqlens( + packed_seq_params, hidden_states.device + )[:-1] + pooled_output = self.pooler(hidden_states, sequence_starts).squeeze(1) else: pooled_output = None # for pylint. if self.return_embeddings: + if packed_seq_params is not None: + return self._packed_mean_pooled_embeddings(hidden_states, packed_seq_params) + embeddings = torch.transpose(hidden_states, 0, 1) masks = torch.sum(attention_mask, dim=1) # Collect masked embeddings. diff --git a/tests/unit_tests/models/test_bert_model.py b/tests/unit_tests/models/test_bert_model.py index 878a0b95915..133a56cc9a7 100644 --- a/tests/unit_tests/models/test_bert_model.py +++ b/tests/unit_tests/models/test_bert_model.py @@ -148,14 +148,17 @@ def test_packed_forward_uses_cu_seqlens_positions_and_no_attention_mask(self, mo position_embedding_type='rope', post_process=False, ) - sequence_length = 6 + sequence_length = 8 input_ids = torch.arange(sequence_length, dtype=torch.int64).unsqueeze(0) cu_seqlens = torch.tensor([0, 2, 6], dtype=torch.int32) + cu_seqlens_padded = torch.tensor([0, 4, 8], dtype=torch.int32) packed_seq_params = PackedSeqParams( qkv_format='thd', cu_seqlens_q=cu_seqlens, cu_seqlens_kv=cu_seqlens, - max_seqlen_q=2, + cu_seqlens_q_padded=cu_seqlens_padded, + cu_seqlens_kv_padded=cu_seqlens_padded, + max_seqlen_q=4, max_seqlen_kv=4, ) encoder_input = torch.ones(sequence_length, 1, config.hidden_size) @@ -180,7 +183,7 @@ def test_packed_forward_uses_cu_seqlens_positions_and_no_attention_mask(self, mo extended_attention_mask.assert_not_called() assert torch.equal( embedding_forward.call_args.kwargs['position_ids'], - torch.tensor([[0, 1, 0, 1, 2, 3]], dtype=torch.int64), + torch.tensor([[0, 1, 2, 3, 0, 1, 2, 3]], dtype=torch.int64), ) rotary_forward.assert_called_once_with( 4, packed_seq=True, cp_group=packed_seq_params.cp_group @@ -190,6 +193,253 @@ def test_packed_forward_uses_cu_seqlens_positions_and_no_attention_mask(self, mo assert encoder_forward.call_args.kwargs['rotary_pos_emb'] is rotary_pos_emb assert output is hidden_states + @pytest.mark.internal + def test_packed_forward_accepts_precomputed_cp_position_ids(self, mocker): + config = TransformerConfig( + num_layers=2, + hidden_size=12, + num_attention_heads=4, + use_cpu_initialization=True, + perform_initialization=True, + attention_backend=AttnBackend.unfused, + ) + bert_model = BertModel( + config=config, + num_tokentypes=0, + transformer_layer_spec=get_bert_layer_with_transformer_engine_spec(), + vocab_size=100, + max_sequence_length=4, + post_process=False, + ) + cp_group = mocker.Mock() + cp_group.size.return_value = 2 + bert_model.cp_group = cp_group + input_ids = torch.arange(4, dtype=torch.int64).unsqueeze(0) + position_ids = torch.tensor([[0, 1, 2, 3]], dtype=torch.int64) + cu_seqlens = torch.tensor([0, 4, 8], dtype=torch.int32) + packed_seq_params = PackedSeqParams( + qkv_format='thd', + cu_seqlens_q=cu_seqlens, + cu_seqlens_kv=cu_seqlens, + cu_seqlens_q_padded=cu_seqlens, + cu_seqlens_kv_padded=cu_seqlens, + max_seqlen_q=4, + max_seqlen_kv=4, + ) + encoder_input = torch.ones(4, 1, config.hidden_size) + hidden_states = torch.zeros_like(encoder_input) + embedding_forward = mocker.patch.object( + bert_model.embedding, 'forward', return_value=encoder_input + ) + mocker.patch.object(bert_model.encoder, 'forward', return_value=hidden_states) + + output = bert_model.forward( + input_ids=input_ids, + attention_mask=None, + packed_seq_params=packed_seq_params, + position_ids=position_ids, + ) + + assert embedding_forward.call_args.kwargs['position_ids'] is position_ids + assert output is hidden_states + + with pytest.raises(ValueError, match='dummy batch'): + bert_model.forward( + input_ids=input_ids.repeat(2, 1), + attention_mask=None, + packed_seq_params=packed_seq_params, + position_ids=position_ids.repeat(2, 1), + ) + + @pytest.mark.internal + def test_packed_binary_head_pools_each_sequence(self, mocker): + sequence_length = 8 + input_ids = torch.arange(sequence_length, dtype=torch.int64).unsqueeze(0) + cu_seqlens = torch.tensor([0, 2, 6], dtype=torch.int32) + cu_seqlens_padded = torch.tensor([0, 4, 8], dtype=torch.int32) + packed_seq_params = PackedSeqParams( + qkv_format='thd', + cu_seqlens_q=cu_seqlens, + cu_seqlens_kv=cu_seqlens, + cu_seqlens_q_padded=cu_seqlens_padded, + cu_seqlens_kv_padded=cu_seqlens_padded, + max_seqlen_q=4, + max_seqlen_kv=4, + ) + hidden_states = torch.zeros(sequence_length, 1, self.bert_model.config.hidden_size) + pooled_output = torch.ones(2, self.bert_model.config.hidden_size) + binary_logits = torch.ones(2, 2) + + mocker.patch.object(self.bert_model.embedding, 'forward', return_value=hidden_states) + mocker.patch.object(self.bert_model.encoder, 'forward', return_value=hidden_states) + pooler_forward = mocker.patch.object( + self.bert_model.pooler, 'forward', return_value=pooled_output.unsqueeze(1) + ) + mocker.patch.object(self.bert_model.lm_head, 'forward', return_value=hidden_states) + mocker.patch.object( + self.bert_model.output_layer, 'forward', return_value=(hidden_states, None) + ) + binary_head_forward = mocker.patch.object( + self.bert_model.binary_head, 'forward', return_value=binary_logits + ) + + _, output_binary_logits = self.bert_model.forward( + input_ids=input_ids, attention_mask=None, packed_seq_params=packed_seq_params + ) + + assert torch.equal(pooler_forward.call_args.args[1], torch.tensor([0, 4], dtype=torch.long)) + binary_head_forward.assert_called_once() + assert torch.equal(binary_head_forward.call_args.args[0], pooled_output) + assert binary_head_forward.call_args.args[0].shape == pooled_output.shape + assert output_binary_logits is binary_logits + + @pytest.mark.internal + def test_packed_return_embeddings_aggregates_each_sequence(self, mocker): + config = TransformerConfig( + num_layers=2, + hidden_size=12, + num_attention_heads=4, + use_cpu_initialization=True, + perform_initialization=True, + attention_backend=AttnBackend.unfused, + ) + bert_model = BertModel( + config=config, + num_tokentypes=0, + transformer_layer_spec=get_bert_layer_with_transformer_engine_spec(), + vocab_size=100, + max_sequence_length=4, + return_embeddings=True, + ) + input_ids = torch.arange(8, dtype=torch.int64).unsqueeze(0) + cu_seqlens = torch.tensor([0, 3, 6], dtype=torch.int32) + cu_seqlens_padded = torch.tensor([0, 4, 8], dtype=torch.int32) + packed_seq_params = PackedSeqParams( + qkv_format='thd', + cu_seqlens_q=cu_seqlens, + cu_seqlens_kv=cu_seqlens, + cu_seqlens_q_padded=cu_seqlens_padded, + cu_seqlens_kv_padded=cu_seqlens_padded, + max_seqlen_q=4, + max_seqlen_kv=4, + ) + hidden_states = ( + torch.arange(8, dtype=torch.float32).view(8, 1, 1).expand(-1, 1, config.hidden_size) + ) + + mocker.patch.object(bert_model.embedding, 'forward', return_value=hidden_states) + mocker.patch.object(bert_model.encoder, 'forward', return_value=hidden_states) + mocker.patch.object( + bert_model.pooler, 'forward', return_value=torch.zeros(2, 1, config.hidden_size) + ) + + output = bert_model.forward( + input_ids=input_ids, attention_mask=None, packed_seq_params=packed_seq_params + ) + + expected = torch.tensor([[1.0], [5.0]]).expand(-1, config.hidden_size) + assert torch.equal(output, expected) + + @pytest.mark.internal + def test_packed_forward_handles_trailing_buffer_padding(self, mocker): + config = TransformerConfig( + num_layers=2, + hidden_size=12, + num_attention_heads=4, + use_cpu_initialization=True, + perform_initialization=True, + attention_backend=AttnBackend.unfused, + ) + bert_model = BertModel( + config=config, + num_tokentypes=0, + transformer_layer_spec=get_bert_layer_with_transformer_engine_spec(), + vocab_size=100, + max_sequence_length=4, + return_embeddings=True, + ) + # The packed buffer holds eight tokens while the last padded sequence ends at seven, + # so a trailing padding region has to be handled without rejecting the input. + input_ids = torch.arange(8, dtype=torch.int64).unsqueeze(0) + cu_seqlens = torch.tensor([0, 3, 6], dtype=torch.int32) + cu_seqlens_padded = torch.tensor([0, 4, 7], dtype=torch.int32) + packed_seq_params = PackedSeqParams( + qkv_format='thd', + cu_seqlens_q=cu_seqlens, + cu_seqlens_kv=cu_seqlens, + cu_seqlens_q_padded=cu_seqlens_padded, + cu_seqlens_kv_padded=cu_seqlens_padded, + max_seqlen_q=4, + max_seqlen_kv=4, + ) + hidden_states = ( + torch.arange(8, dtype=torch.float32).view(8, 1, 1).expand(-1, 1, config.hidden_size) + ) + + embedding_forward = mocker.patch.object( + bert_model.embedding, 'forward', return_value=hidden_states + ) + mocker.patch.object(bert_model.encoder, 'forward', return_value=hidden_states) + pooler_forward = mocker.patch.object( + bert_model.pooler, 'forward', return_value=torch.zeros(2, 1, config.hidden_size) + ) + + output = bert_model.forward( + input_ids=input_ids, attention_mask=None, packed_seq_params=packed_seq_params + ) + + assert torch.equal( + embedding_forward.call_args.kwargs['position_ids'], + torch.tensor([[0, 1, 2, 3, 0, 1, 2, 0]], dtype=torch.int64), + ) + assert torch.equal(pooler_forward.call_args.args[1], torch.tensor([0, 4], dtype=torch.long)) + expected = torch.tensor([[1.0], [5.0]]).expand(-1, config.hidden_size) + assert torch.equal(output, expected) + + @pytest.mark.internal + def test_packed_cp_post_processing_is_rejected(self, mocker): + cp_group = mocker.Mock() + cp_group.size.return_value = 2 + self.bert_model.cp_group = cp_group + input_ids = torch.arange(4, dtype=torch.int64).unsqueeze(0) + cu_seqlens = torch.tensor([0, 4, 8], dtype=torch.int32) + + with pytest.raises(ValueError, match='post-processing'): + self.bert_model.forward( + input_ids=input_ids, + attention_mask=None, + packed_seq_params=PackedSeqParams( + qkv_format='thd', + cu_seqlens_q=cu_seqlens, + cu_seqlens_kv=cu_seqlens, + cu_seqlens_q_padded=cu_seqlens, + cu_seqlens_kv_padded=cu_seqlens, + max_seqlen_q=4, + max_seqlen_kv=4, + ), + position_ids=torch.arange(4, dtype=torch.int64).unsqueeze(0), + ) + + @pytest.mark.internal + def test_packed_return_embeddings_with_sequence_parallelism_is_rejected(self): + self.bert_model.return_embeddings = True + self.bert_model.config.sequence_parallel = True + input_ids = torch.arange(4, dtype=torch.int64).unsqueeze(0) + cu_seqlens = torch.tensor([0, 4], dtype=torch.int32) + + with pytest.raises(ValueError, match='sequence parallelism'): + self.bert_model.forward( + input_ids=input_ids, + attention_mask=None, + packed_seq_params=PackedSeqParams( + qkv_format='thd', + cu_seqlens_q=cu_seqlens, + cu_seqlens_kv=cu_seqlens, + max_seqlen_q=4, + max_seqlen_kv=4, + ), + ) + @pytest.mark.internal def test_forward_validates_dense_attention_mask_and_packed_format(self): sequence_length = self.bert_model.max_sequence_length @@ -241,7 +491,11 @@ def test_forward_validates_dense_attention_mask_and_packed_format(self): 'cu_seqlens_q must be provided', ), (torch.arange(4, dtype=torch.int64).unsqueeze(0), torch.tensor([1, 4]), 'start at 0'), - (torch.arange(4, dtype=torch.int64).unsqueeze(0), torch.tensor([0, 3]), 'end at'), + ( + torch.arange(4, dtype=torch.int64).unsqueeze(0), + torch.tensor([0, 5]), + 'must not exceed', + ), ( torch.arange(4, dtype=torch.int64).unsqueeze(0), torch.tensor([0, 3, 2, 4]),