Skip to content
Merged
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
23 changes: 19 additions & 4 deletions megatron/core/transformer/transformer_block.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)

Expand Down Expand Up @@ -669,7 +684,7 @@ def checkpoint_handler(forward_func):
attention_mask,
context,
context_mask,
rotary_pos_emb,
*rotary_pos_emb,
padding_mask,
)
else:
Expand All @@ -680,7 +695,7 @@ def checkpoint_handler(forward_func):
attention_mask,
context,
context_mask,
rotary_pos_emb,
*rotary_pos_emb,
padding_mask,
)

Expand Down Expand Up @@ -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
Expand Down
103 changes: 100 additions & 3 deletions tests/unit_tests/transformer/test_transformer_block.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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'
Expand All @@ -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()

Expand All @@ -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
Expand Down
Loading