[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

@@ -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