[train] support megatron-bridge for PT/SFT training (#10645)

This commit is contained in:
sunyi0505
2026-07-27 18:45:18 +08:00
committed by GitHub
parent 2ebe7be611
commit 9ce6b663e9
20 changed files with 2560 additions and 4 deletions

View File

@@ -16,6 +16,7 @@ from .data_args import DataArguments
from .evaluation_args import EvaluationArguments
from .finetuning_args import FinetuningArguments
from .generating_args import GeneratingArguments
from .megatron_bridge_args import MegatronBridgeArguments
from .model_args import ModelArguments
from .parser import get_eval_args, get_infer_args, get_ray_args, get_train_args, read_args
from .training_args import RayArguments, TrainingArguments
@@ -26,6 +27,7 @@ __all__ = [
"EvaluationArguments",
"FinetuningArguments",
"GeneratingArguments",
"MegatronBridgeArguments",
"ModelArguments",
"RayArguments",
"TrainingArguments",

View File

@@ -482,6 +482,21 @@ class FinetuningArguments(
)
},
)
use_megatron_bridge: bool = field(
default=False,
metadata={
"help": (
"Whether or not to use Megatron Bridge training backend. "
"Controlled by USE_MEGATRON_BRIDGE environment variable."
)
},
)
megatron_bridge_args: Any = field(
default=None,
init=False,
repr=False,
metadata={"help": "Megatron Bridge specific arguments, set when USE_MEGATRON_BRIDGE=1."},
)
use_hyper_parallel: bool = field(
default=False,
metadata={

View File

@@ -0,0 +1,193 @@
# Copyright 2025 the LlamaFactory team.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
import json
import os
from dataclasses import dataclass, field
from typing import Literal, Optional
from transformers.training_args import _convert_str_dict
@dataclass
class MegatronBridgeArguments:
r"""Arguments for Megatron Bridge distributed training backend.
Parallelism, optimizer overlap, checkpoint conversion, and selected Megatron
model-provider knobs are exposed here because Megatron Bridge uses a
standalone workflow outside the Hugging Face Trainer.
"""
tensor_model_parallel_size: int = field(
default=1,
metadata={"help": "Tensor model parallel size for Megatron Bridge."},
)
pipeline_model_parallel_size: int = field(
default=1,
metadata={"help": "Pipeline model parallel size for Megatron Bridge."},
)
expert_model_parallel_size: int = field(
default=1,
metadata={"help": "Expert model parallel size for MoE models."},
)
context_parallel_size: int = field(
default=1,
metadata={"help": "Context parallel size for Megatron Bridge."},
)
virtual_pipeline_model_parallel_size: Optional[int] = field(
default=None,
metadata={"help": "Virtual pipeline (interleaved) parallel size. None keeps provider default."},
)
sequence_parallel: bool = field(
default=False,
metadata={"help": "Whether to enable sequence parallelism."},
)
recompute_granularity: Optional[str] = field(
default=None,
metadata={"help": "Activation recomputation granularity: 'full' or 'selective'."},
)
recompute_method: Optional[Literal["uniform", "block"]] = field(
default=None,
metadata={"help": "Activation recomputation method: 'uniform' or 'block'."},
)
recompute_num_layers: Optional[int] = field(
default=None,
metadata={"help": "Number of layers per recompute unit when recompute_method is set."},
)
account_for_embedding_in_pipeline_split: Optional[bool] = field(
default=None,
metadata={"help": "Whether pipeline split accounts for the embedding layer."},
)
account_for_loss_in_pipeline_split: Optional[bool] = field(
default=None,
metadata={"help": "Whether pipeline split accounts for the loss layer."},
)
bias_activation_fusion: Optional[bool] = field(
default=None,
metadata={"help": "Enable bias+activation fusion. None keeps Megatron provider default."},
)
apply_rope_fusion: Optional[bool] = field(
default=None,
metadata={"help": "Enable RoPE fusion kernel. None keeps Megatron provider default."},
)
masked_softmax_fusion: Optional[bool] = field(
default=None,
metadata={"help": "Enable masked softmax fusion. None keeps Megatron provider default."},
)
cross_entropy_loss_fusion: Optional[bool] = field(
default=None,
metadata={"help": "Enable cross-entropy loss fusion. None keeps Megatron provider default."},
)
moe_grouped_gemm: Optional[bool] = field(
default=None,
metadata={"help": "Enable grouped GEMM for MoE experts. None keeps provider default."},
)
moe_token_dispatcher_type: Optional[Literal["allgather", "alltoall", "flex"]] = field(
default=None,
metadata={"help": "MoE token dispatcher type: allgather, alltoall, or flex."},
)
calculate_per_token_loss: Optional[bool] = field(
default=None,
metadata={
"help": (
"Whether to compute per-token loss. When context_parallel_size > 1, "
"this is forced to True regardless of this setting."
)
},
)
use_distributed_optimizer: bool = field(
default=True,
metadata={"help": "Whether to use Megatron distributed optimizer."},
)
overlap_param_gather: bool = field(
default=True,
metadata={"help": "Whether to overlap parameter all-gather with forward compute."},
)
overlap_grad_reduce: bool = field(
default=True,
metadata={"help": "Whether to overlap gradient all-reduce with backward compute."},
)
use_packed_sequences: bool = field(
default=False,
metadata={"help": "Whether to use packed sequences for SFT efficiency."},
)
mixed_precision: str = field(
default="bf16_mixed",
metadata={"help": "Mixed precision mode for Megatron Bridge, e.g. bf16_mixed or fp8."},
)
megatron_pretrained_checkpoint: Optional[str] = field(
default=None,
metadata={
"help": (
"Path to a Megatron-format pretrained checkpoint. "
"If unset, HF weights are converted automatically before training."
)
},
)
export_hf_on_finish: bool = field(
default=False,
metadata={"help": "Whether to export the final checkpoint to Hugging Face format after training."},
)
extra_config: Optional[str] = field(
default=None,
metadata={
"help": (
"Optional JSON string or path to a JSON file with extra Megatron Bridge model/training overrides. "
"Dot-paths are supported (e.g. train.train_iters or checkpoint.save_interval)."
)
},
)
def __post_init__(self) -> None:
if self.tensor_model_parallel_size < 1:
raise ValueError("`tensor_model_parallel_size` must be >= 1.")
if self.pipeline_model_parallel_size < 1:
raise ValueError("`pipeline_model_parallel_size` must be >= 1.")
if self.expert_model_parallel_size < 1:
raise ValueError("`expert_model_parallel_size` must be >= 1.")
if self.context_parallel_size < 1:
raise ValueError("`context_parallel_size` must be >= 1.")
if self.virtual_pipeline_model_parallel_size is not None and self.virtual_pipeline_model_parallel_size < 1:
raise ValueError("`virtual_pipeline_model_parallel_size` must be >= 1 when set.")
if self.sequence_parallel and self.tensor_model_parallel_size <= 1:
raise ValueError("`sequence_parallel` requires `tensor_model_parallel_size` > 1.")
if self.recompute_granularity is not None and self.recompute_granularity not in ("full", "selective"):
raise ValueError("`recompute_granularity` must be 'full' or 'selective'.")
if self.recompute_method is not None and self.recompute_method not in ("uniform", "block"):
raise ValueError("`recompute_method` must be 'uniform' or 'block'.")
if self.recompute_num_layers is not None and self.recompute_num_layers < 1:
raise ValueError("`recompute_num_layers` must be >= 1 when set.")
if self.moe_token_dispatcher_type is not None and self.moe_token_dispatcher_type not in (
"allgather",
"alltoall",
"flex",
):
raise ValueError("`moe_token_dispatcher_type` must be 'allgather', 'alltoall', or 'flex'.")
if isinstance(self.extra_config, str):
config_str = self.extra_config.strip()
if config_str.startswith("{"):
self.extra_config = _convert_str_dict(json.loads(config_str))
else:
self.extra_config = config_str
def load_extra_config(self) -> dict:
if self.extra_config is None:
return {}
if isinstance(self.extra_config, dict):
return self.extra_config
if not os.path.isfile(self.extra_config):
raise ValueError(f"`extra_config` file not found: {self.extra_config}")
with open(self.extra_config, encoding="utf-8") as f:
return json.load(f)

View File

@@ -33,11 +33,12 @@ from transformers.utils import is_torch_bf16_gpu_available, is_torch_npu_availab
from ..extras import logging
from ..extras.constants import CHECKPOINT_NAMES, EngineName
from ..extras.misc import check_dependencies, check_version, get_current_device, is_env_enabled
from ..extras.packages import is_mcore_adapter_available
from ..extras.packages import is_mcore_adapter_available, is_megatron_bridge_available
from .data_args import DataArguments
from .evaluation_args import EvaluationArguments
from .finetuning_args import FinetuningArguments
from .generating_args import GeneratingArguments
from .megatron_bridge_args import MegatronBridgeArguments
from .model_args import ModelArguments
from .training_args import RayArguments, TrainingArguments
@@ -81,6 +82,23 @@ else:
_TRAIN_MCA_ARGS = []
_TRAIN_MCA_CLS = tuple()
_TRAIN_MBRIDGE_ARGS = [
ModelArguments,
DataArguments,
TrainingArguments,
FinetuningArguments,
MegatronBridgeArguments,
GeneratingArguments,
]
_TRAIN_MBRIDGE_CLS = tuple[
ModelArguments,
DataArguments,
TrainingArguments,
FinetuningArguments,
MegatronBridgeArguments,
GeneratingArguments,
]
def read_args(args: dict[str, Any] | list[str] | None = None) -> dict[str, Any] | list[str]:
r"""Get arguments from the command line or a config file."""
@@ -246,6 +264,9 @@ def _check_extra_dependencies(
if finetuning_args.plot_loss:
check_version("matplotlib", mandatory=True)
if finetuning_args.use_megatron_bridge:
check_version("megatron-bridge", mandatory=True)
if training_args is not None:
if training_args.deepspeed:
check_version("deepspeed", mandatory=True)
@@ -283,6 +304,39 @@ def _configure_mca_training_args(training_args, data_args, finetuning_args) -> N
finetuning_args.use_mca = True
def _validate_megatron_bridge_parallel_args(mb_args: MegatronBridgeArguments, world_size: int) -> None:
parallel_size = (
mb_args.tensor_model_parallel_size
* mb_args.pipeline_model_parallel_size
* mb_args.context_parallel_size
* mb_args.expert_model_parallel_size
)
if parallel_size > world_size:
raise ValueError(f"Total Megatron Bridge parallel size ({parallel_size}) exceeds `world_size` ({world_size}).")
if world_size % parallel_size != 0:
raise ValueError(
f"Total Megatron Bridge parallel size ({parallel_size}) must divide `world_size` ({world_size})."
)
def _parse_train_mbridge_args(args: dict[str, Any] | list[str] | None = None) -> _TRAIN_MBRIDGE_CLS:
parser = HfArgumentParser(_TRAIN_MBRIDGE_ARGS)
allow_extra_keys = is_env_enabled("ALLOW_EXTRA_ARGS")
model_args, data_args, training_args, finetuning_args, mb_args, generating_args = _parse_args(
parser, args, allow_extra_keys=allow_extra_keys
)
_configure_mbridge_training_args(training_args, data_args, finetuning_args)
return model_args, data_args, training_args, finetuning_args, mb_args, generating_args
def _configure_mbridge_training_args(training_args, data_args, finetuning_args) -> None:
"""Patch training args to avoid args checking errors and sync Megatron Bridge settings."""
training_args.predict_with_generate = False
training_args.generation_max_length = data_args.cutoff_len
training_args.generation_num_beams = 1
finetuning_args.use_megatron_bridge = True
def _parse_infer_args(args: dict[str, Any] | list[str] | None = None) -> _INFER_CLS:
parser = HfArgumentParser(_INFER_ARGS)
allow_extra_keys = is_env_enabled("ALLOW_EXTRA_ARGS")
@@ -302,11 +356,22 @@ def get_ray_args(args: dict[str, Any] | list[str] | None = None) -> RayArguments
def get_train_args(args: dict[str, Any] | list[str] | None = None) -> _TRAIN_CLS:
mb_args = None
if is_env_enabled("USE_MCA"):
model_args, data_args, training_args, finetuning_args, generating_args = _parse_train_mca_args(args)
elif is_env_enabled("USE_MEGATRON_BRIDGE"):
if not is_megatron_bridge_available():
raise ImportError(
"megatron-bridge is required when USE_MEGATRON_BRIDGE=1. "
"Please install `megatron-bridge` and its dependencies."
)
model_args, data_args, training_args, finetuning_args, mb_args, generating_args = _parse_train_mbridge_args(
args
)
else:
model_args, data_args, training_args, finetuning_args, generating_args = _parse_train_args(args)
finetuning_args.use_mca = False
finetuning_args.use_megatron_bridge = False
# Setup logging
if training_args.should_log:
@@ -326,6 +391,22 @@ def get_train_args(args: dict[str, Any] | list[str] | None = None) -> _TRAIN_CLS
if finetuning_args.stage == "sft" and training_args.do_predict and not training_args.predict_with_generate:
raise ValueError("Please enable `predict_with_generate` to save model predictions.")
if finetuning_args.use_megatron_bridge:
if finetuning_args.use_mca or finetuning_args.use_hyper_parallel:
raise ValueError("Megatron Bridge cannot be used together with MCA or HyperParallel.")
if finetuning_args.stage not in ["pt", "sft"]:
raise ValueError("Megatron Bridge only supports the `pt` and `sft` stages.")
if finetuning_args.finetuning_type not in ["full", "lora"]:
raise ValueError("Megatron Bridge only supports `full` and `lora` finetuning.")
if model_args.quantization_bit is not None:
raise ValueError("Quantized models are not supported with Megatron Bridge.")
if training_args.deepspeed is not None:
raise ValueError("Megatron Bridge is incompatible with DeepSpeed.")
if mb_args is None:
raise ValueError("Megatron Bridge arguments are missing. Please set USE_MEGATRON_BRIDGE=1.")
_validate_megatron_bridge_parallel_args(mb_args, training_args.world_size)
finetuning_args.megatron_bridge_args = mb_args
if finetuning_args.stage in ["rm", "ppo"] and training_args.load_best_model_at_end:
raise ValueError("RM and PPO stages do not support `load_best_model_at_end`.")
@@ -400,7 +481,12 @@ def get_train_args(args: dict[str, Any] | list[str] | None = None) -> _TRAIN_CLS
if training_args.deepspeed is not None and (finetuning_args.use_galore or finetuning_args.use_apollo):
raise ValueError("GaLore and APOLLO are incompatible with DeepSpeed yet.")
if not finetuning_args.use_mca and training_args.fp8 and model_args.quantization_bit is not None:
if (
not finetuning_args.use_mca
and not finetuning_args.use_megatron_bridge
and training_args.fp8
and model_args.quantization_bit is not None
):
raise ValueError("FP8 training is not compatible with quantization. Please disable one of them.")
if model_args.infer_backend != EngineName.HF:
@@ -417,7 +503,12 @@ def get_train_args(args: dict[str, Any] | list[str] | None = None) -> _TRAIN_CLS
_check_extra_dependencies(model_args, finetuning_args, training_args)
_verify_trackio_args(training_args)
if not finetuning_args.use_mca and training_args.fp8_enable_fsdp_float8_all_gather and not training_args.fp8:
if (
not finetuning_args.use_mca
and not finetuning_args.use_megatron_bridge
and training_args.fp8_enable_fsdp_float8_all_gather
and not training_args.fp8
):
logger.warning_rank0("fp8_enable_fsdp_float8_all_gather requires fp8=True. Setting fp8=True.")
model_args.fp8 = True