diff --git a/megatron/core/jit.py b/megatron/core/jit.py index b1aa3e0b611..b67810f2e34 100644 --- a/megatron/core/jit.py +++ b/megatron/core/jit.py @@ -7,12 +7,27 @@ jit_fuser = torch.jit.script # nvFuser is deprecated in PyTorch JIT starting from 2.2 -try: - if is_torch_min_version("2.2.0a0"): - jit_fuser = torch.compile -except ImportError: - def noop_decorator(func): - return func +def noop_decorator(func): + '''No-op decorator''' + return func + +def enable_jit_fuser(): + '''Enable the JIT fuser''' + global jit_fuser + try: + if is_torch_min_version("2.2.0a0"): + jit_fuser = torch.compile + except ImportError: + + jit_fuser = noop_decorator + + +def disable_jit_fuser(): + '''Disable the JIT fuser''' + global jit_fuser jit_fuser = noop_decorator + + +enable_jit_fuser() diff --git a/megatron/core/ssm/gated_delta_net.py b/megatron/core/ssm/gated_delta_net.py index e12dfd68062..4df1c0df294 100644 --- a/megatron/core/ssm/gated_delta_net.py +++ b/megatron/core/ssm/gated_delta_net.py @@ -18,6 +18,7 @@ from megatron.core.dist_checkpointing.mapping import ReplicaId, ShardedTensorFactory from megatron.core.fp8_utils import get_fp8_align_size from megatron.core.inference.contexts import BaseInferenceContext +from megatron.core.jit import jit_fuser from megatron.core.packed_seq_params import PackedSeqParams from megatron.core.process_groups_config import ProcessGroupCollection from megatron.core.tensor_parallel import get_cuda_rng_tracker @@ -384,7 +385,7 @@ def forward( # RMSNorm nvtx_range_push(suffix="gated_norm") - norm_out = self._torch_compiled_gated_norm(core_attn_out, gate) + norm_out = self._apply_gated_norm(core_attn_out, gate) nvtx_range_pop(suffix="gated_norm") # Transpose: b s x --> s b x @@ -399,8 +400,8 @@ def forward( return out, out_bias - @torch.compile - def _torch_compiled_gated_norm(self, x, gate): + @jit_fuser + def _apply_gated_norm(self, x, gate): # Output Norm x_dtype = x.dtype x = x.reshape(-1, x.shape[-1]) diff --git a/megatron/core/transformer/attention.py b/megatron/core/transformer/attention.py index c4738ea84b5..65bac394034 100644 --- a/megatron/core/transformer/attention.py +++ b/megatron/core/transformer/attention.py @@ -9,6 +9,7 @@ from megatron.core import tensor_parallel from megatron.core.inference.contexts import BaseInferenceContext +from megatron.core.jit import jit_fuser from megatron.core.models.common.embeddings.rope_utils import ( apply_rotary_pos_emb, apply_rotary_pos_emb_with_cos_sin, @@ -958,7 +959,7 @@ def forward( # Output gate if gate is not None: nvtx_range_push(suffix="output_gate") - core_attn_out = self._torch_compiled_output_gate(core_attn_out, gate) + core_attn_out = self._apply_output_gate(core_attn_out, gate) nvtx_range_pop(suffix="output_gate") # ================= @@ -978,8 +979,8 @@ def forward( return output, bias - @torch.compile - def _torch_compiled_output_gate(self, x, gate): + @jit_fuser + def _apply_output_gate(self, x, gate): x_dtype = x.dtype gate = gate.contiguous() gate = gate.view(*x.shape) diff --git a/megatron/core/transformer/transformer_block.py b/megatron/core/transformer/transformer_block.py index f189d5fa7a4..c3187350c43 100755 --- a/megatron/core/transformer/transformer_block.py +++ b/megatron/core/transformer/transformer_block.py @@ -801,6 +801,12 @@ def sharded_state_dict( elif isinstance(self.config.moe_layer_freq, list): non_homogeneous_layers = True + if isinstance(self.config.linear_attention_freq, int): + if self.config.linear_attention_freq > 1: + non_homogeneous_layers = True + elif isinstance(self.config.linear_attention_freq, list): + non_homogeneous_layers = True + if self.config.heterogeneous_block_specs: non_homogeneous_layers = True diff --git a/megatron/training/arguments.py b/megatron/training/arguments.py index 3413d1e1547..3056f2007f2 100644 --- a/megatron/training/arguments.py +++ b/megatron/training/arguments.py @@ -2349,6 +2349,8 @@ def _add_training_args(parser): help='The submodules to offload its input. Choices: "attn_norm", "core_attn", "attn_proj", "mlp_norm", "expert_fc1", "moe_act".') group.add_argument('--min-offloaded-tensor-size', type=int, default=1024*1024, help='The minimum size of the tensor to be offloaded.') + group.add_argument('--disable-jit-fuser', action='store_true', + help='Disable the JIT fuser.') return parser diff --git a/megatron/training/global_vars.py b/megatron/training/global_vars.py index 17bda9cbac5..ec402263d29 100644 --- a/megatron/training/global_vars.py +++ b/megatron/training/global_vars.py @@ -9,6 +9,7 @@ from megatron.core import Timers from megatron.core.config import set_experimental_flag from megatron.core.energy_monitor import EnergyMonitor +from megatron.core.jit import disable_jit_fuser from megatron.core.num_microbatches_calculator import init_num_microbatches_calculator, unset_num_microbatches_calculator from megatron.training import dist_signal_handler from megatron.training.tokenizer import build_tokenizer @@ -111,6 +112,9 @@ def set_global_variables(args, build_tokenizer=True): if args.exit_signal_handler: _set_signal_handler() + if args.disable_jit_fuser: + disable_jit_fuser() + def unset_global_variables(): """Unset global vars.