mirror of
https://github.com/hiyouga/LLaMA-Factory.git
synced 2026-09-27 01:45:42 +08:00
[train] support megatron-bridge for PT/SFT training (#10645)
This commit is contained in:
@@ -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",
|
||||
|
||||
@@ -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={
|
||||
|
||||
193
src/llamafactory/hparams/megatron_bridge_args.py
Normal file
193
src/llamafactory/hparams/megatron_bridge_args.py
Normal 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)
|
||||
@@ -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
|
||||
|
||||
|
||||
Reference in New Issue
Block a user