diff --git a/src/megatron/bridge/training/gpt_step.py b/src/megatron/bridge/training/gpt_step.py index 589235c9ea..9744c2827e 100644 --- a/src/megatron/bridge/training/gpt_step.py +++ b/src/megatron/bridge/training/gpt_step.py @@ -253,8 +253,12 @@ def _forward_step_common( "max_seqlen": max_seqlen, "cu_seqlens_unpadded": cu_seqlens_unpadded, "cu_seqlens_unpadded_argmin": cu_seqlens_unpadded_argmin, - "total_tokens": tokens.size(1) if tokens is not None else labels.size(1), } + # total_tokens drives seq_idx computation in PackedSeqParams.__post_init__, + # which is only needed for Mamba/hybrid SSM layers. Skip it for pure + # transformer models to avoid per-step CUDA overhead. + if getattr(config, "is_hybrid_model", False): + packed_seq_params["total_tokens"] = tokens.size(1) if tokens is not None else labels.size(1) forward_args["packed_seq_params"] = get_packed_seq_params(packed_seq_params) with straggler_timer: diff --git a/src/megatron/bridge/training/llava_step.py b/src/megatron/bridge/training/llava_step.py index f3f1019556..644b0a5031 100644 --- a/src/megatron/bridge/training/llava_step.py +++ b/src/megatron/bridge/training/llava_step.py @@ -198,8 +198,12 @@ def forward_step( "cu_seqlens": cu_seqlens, "cu_seqlens_argmin": cu_seqlens_argmin, "max_seqlen": max_seqlen, - "total_tokens": input_ids.size(1) if input_ids is not None else labels.size(1), } + # total_tokens drives seq_idx computation in PackedSeqParams.__post_init__, + # which is only needed for Mamba/hybrid SSM layers. Skip it for pure + # transformer models to avoid per-step CUDA overhead. + if getattr(config, "is_hybrid_model", False): + packed_seq_params["total_tokens"] = input_ids.size(1) if input_ids is not None else labels.size(1) forward_args["packed_seq_params"] = get_packed_seq_params(packed_seq_params) check_for_nan_in_loss = state.cfg.rerun_state_machine.check_for_nan_in_loss diff --git a/src/megatron/bridge/training/vlm_step.py b/src/megatron/bridge/training/vlm_step.py index 8c3fb6df6b..3137cc36df 100644 --- a/src/megatron/bridge/training/vlm_step.py +++ b/src/megatron/bridge/training/vlm_step.py @@ -448,8 +448,12 @@ def forward_step( "cu_seqlens": cu_seqlens, "max_seqlen": max_seqlen, "cu_seqlens_argmin": cu_seqlens_argmin, - "total_tokens": tokens.size(1) if tokens is not None else labels.size(1), } + # total_tokens drives seq_idx computation in PackedSeqParams.__post_init__, + # which is only needed for Mamba/hybrid SSM layers. Skip it for pure + # transformer models to avoid per-step CUDA overhead. + if getattr(config, "is_hybrid_model", False): + packed_seq_params["total_tokens"] = tokens.size(1) if tokens is not None else labels.size(1) forward_args["packed_seq_params"] = get_packed_seq_params(packed_seq_params) if loss_mask is not None: