diff --git a/megatron/core/transformer/transformer_block.py b/megatron/core/transformer/transformer_block.py index 142fba9a69f..921353947a8 100755 --- a/megatron/core/transformer/transformer_block.py +++ b/megatron/core/transformer/transformer_block.py @@ -611,15 +611,30 @@ def _checkpointed_forward( extract_layer_indices = set() intermediate_hidden_states: List[Tensor] = [] + # Unpack dual RoPE before checkpointing because autograd only accepts + # tensors (or None) in save_for_backward. + is_dual_rope = isinstance(rotary_pos_emb, (tuple, list)) + assert ( + not is_dual_rope or len(rotary_pos_emb) == 2 + ), "Dual RoPE input length is not equal to 2" + rotary_pos_emb = rotary_pos_emb if is_dual_rope else (None, rotary_pos_emb) + def custom(start: int, end: int): def custom_forward( hidden_states, attention_mask, context, context_mask, - rotary_pos_emb, + rotary_pos_emb_local, + rotary_pos_emb_global, padding_mask=None, ): + rotary_pos_emb = ( + (rotary_pos_emb_local, rotary_pos_emb_global) + if is_dual_rope + else rotary_pos_emb_global + ) + for index in range(start, end): layer = self._get_layer(index) @@ -669,7 +684,7 @@ def checkpoint_handler(forward_func): attention_mask, context, context_mask, - rotary_pos_emb, + *rotary_pos_emb, padding_mask, ) else: @@ -680,7 +695,7 @@ def checkpoint_handler(forward_func): attention_mask, context, context_mask, - rotary_pos_emb, + *rotary_pos_emb, padding_mask, ) @@ -725,7 +740,7 @@ def checkpoint_handler(forward_func): hidden_states, context = checkpoint_handler(custom(layer_idx, layer_idx + 1)) else: hidden_states, context = custom(layer_idx, layer_idx + 1)( - hidden_states, attention_mask, context, context_mask, rotary_pos_emb + hidden_states, attention_mask, context, context_mask, *rotary_pos_emb ) # Feature extraction: collect hidden states at specified global layer indices diff --git a/tests/unit_tests/transformer/test_transformer_block.py b/tests/unit_tests/transformer/test_transformer_block.py index 63add511bcd..0f19bc3dc95 100644 --- a/tests/unit_tests/transformer/test_transformer_block.py +++ b/tests/unit_tests/transformer/test_transformer_block.py @@ -14,6 +14,7 @@ from megatron.core.models.gpt.gpt_model import GPTModel from megatron.core.process_groups_config import ProcessGroupCollection from megatron.core.tensor_parallel.random import model_parallel_cuda_manual_seed +from megatron.core.transformer.attention import SelfAttention from megatron.core.transformer.enums import ModelType from megatron.core.transformer.pipeline_parallel_layer_layout import PipelineParallelLayerLayout from megatron.core.transformer.spec_utils import build_module @@ -77,13 +78,101 @@ def test_gpu_forward_full_checkpoint(self): def test_gpu_forward_full_checkpoint_fp8(self): self._run_full_checkpoint_test(fp8="e4m3") + def test_gpu_forward_full_checkpoint_dual_rope(self): + kv_channels = self.transformer_config.kv_channels + sequence_length = 32 + rotary_pos_emb = ( + torch.ones(sequence_length, 1, 1, kv_channels, device='cuda'), + torch.ones(sequence_length, 1, 1, kv_channels, device='cuda'), + ) + + def modify_arg_and_forward( + target_func, args, kwargs, target_name, target_index, arg_modifier_func + ): + args_list = list(args) + + if target_name in kwargs: + kwargs[target_name] = arg_modifier_func(kwargs[target_name]) + elif len(args_list) > target_index: + args_list[target_index] = arg_modifier_func(args_list[target_index]) + else: + raise RuntimeError( + f"Argument '{target_name}' at index {target_index} was not provided in args or kwargs." + ) + + return target_func(*args_list, **kwargs) + + class MockSelfAttentionWithDualRope(SelfAttention): + def forward(self, *args, **kwargs): + """Switch to either local or global RoPE embedding before forward.""" + + def arg_modifier_func(rotary_pos_emb): + assert isinstance(rotary_pos_emb, (tuple, list)) and len(rotary_pos_emb) == 2 + + if self.dual_rope_kind == "local_and_global": + assert rotary_pos_emb[0] is not None + assert rotary_pos_emb[1] is not None + + if self.layer_number % 2 == 0: + final_rotary_pos_emb = rotary_pos_emb[0] + else: + final_rotary_pos_emb = rotary_pos_emb[1] + elif self.dual_rope_kind == "local_only": + assert rotary_pos_emb[0] is not None + assert rotary_pos_emb[1] is None + + final_rotary_pos_emb = rotary_pos_emb[0] + elif self.dual_rope_kind == "global_only": + assert rotary_pos_emb[0] is None + assert rotary_pos_emb[1] is not None + + final_rotary_pos_emb = rotary_pos_emb[1] + else: + assert False, f"Unknown dual_rope_kind: {self.dual_rope_kind}" + + return final_rotary_pos_emb + + return modify_arg_and_forward( + super().forward, args, kwargs, "rotary_pos_emb", 5, arg_modifier_func + ) + + # Test non-Dual RoPE + self._run_full_checkpoint_test( + fp8=None, seq_len=sequence_length, rotary_pos_emb=rotary_pos_emb[0] + ) + + # Test Dual RoPE + self._run_full_checkpoint_test( + fp8=None, + seq_len=sequence_length, + attn_class=MockSelfAttentionWithDualRope, + rotary_pos_emb=rotary_pos_emb, + dual_rope_kind="local_and_global", + ) + self._run_full_checkpoint_test( + fp8=None, + seq_len=sequence_length, + attn_class=MockSelfAttentionWithDualRope, + rotary_pos_emb=(rotary_pos_emb[0], None), + dual_rope_kind="local_only", + ) + self._run_full_checkpoint_test( + fp8=None, + seq_len=sequence_length, + attn_class=MockSelfAttentionWithDualRope, + rotary_pos_emb=(None, rotary_pos_emb[1]), + dual_rope_kind="global_only", + ) + def test_gpu_forward_selective_checkpoint(self): self._run_selective_checkpoint_test(fp8=None) def test_gpu_forward_selective_checkpoint_fp8(self): self._run_selective_checkpoint_test(fp8="e4m3") - def _run_full_checkpoint_test(self, fp8): + def _run_full_checkpoint_test( + self, fp8, seq_len=None, attn_class=None, rotary_pos_emb=None, dual_rope_kind=None + ): transformer_config = self.transformer_config config = transformer_config config.recompute_granularity = 'full' @@ -93,11 +182,17 @@ def _run_full_checkpoint_test(self, fp8): full_transformer_block = TransformerBlock( config, get_gpt_layer_with_transformer_engine_spec() ) + if attn_class is not None: + for layer in full_transformer_block.layers: + layer.self_attention.__class__ = attn_class + assert not hasattr(layer.self_attention, "dual_rope_kind") + layer.self_attention.dual_rope_kind = dual_rope_kind + assert full_transformer_block.config.recompute_granularity == 'full' assert full_transformer_block.config.recompute_method == 'block' assert full_transformer_block.config.fp8 == fp8 - sequence_length = 32 + sequence_length = 32 if seq_len is None else seq_len micro_batch_size = 2 full_transformer_block.cuda() @@ -108,7 +203,9 @@ def _run_full_checkpoint_test(self, fp8): attention_mask = torch.ones((1, 1, sequence_length, sequence_length), dtype=bool).cuda() hidden_states = full_transformer_block( - hidden_states=hidden_states, attention_mask=attention_mask + hidden_states=hidden_states, + attention_mask=attention_mask, + rotary_pos_emb=rotary_pos_emb, ) assert hidden_states.shape[0] == sequence_length assert hidden_states.shape[1] == micro_batch_size