mirror of
https://github.com/hiyouga/LLaMA-Factory.git
synced 2026-08-17 13:35:44 +08:00
[train] support megatron-bridge for PT/SFT training (#10645)
This commit is contained in:
@@ -27,6 +27,7 @@ from ..extras.misc import find_available_port, get_device_name, get_torch_device
|
||||
from ..extras.packages import (
|
||||
is_hyper_parallel_available,
|
||||
is_mcore_adapter_available,
|
||||
is_megatron_bridge_available,
|
||||
is_ray_available,
|
||||
is_transformers_version_greater_than,
|
||||
)
|
||||
@@ -100,6 +101,24 @@ def _training_function(config: dict[str, Any]) -> None:
|
||||
|
||||
run_sft_hp(model_args, data_args, training_args, finetuning_args, generating_args, callbacks)
|
||||
|
||||
elif finetuning_args.stage in ["pt", "sft"] and finetuning_args.use_megatron_bridge:
|
||||
if not is_megatron_bridge_available():
|
||||
raise ImportError(
|
||||
"megatron-bridge is not installed. "
|
||||
"Please install it with `pip install --no-build-isolation megatron-bridge`."
|
||||
)
|
||||
mb_args = finetuning_args.megatron_bridge_args
|
||||
if mb_args is None:
|
||||
raise ValueError("Megatron Bridge arguments are missing. Please set USE_MEGATRON_BRIDGE=1.")
|
||||
if finetuning_args.stage == "pt":
|
||||
from .megatron_bridge import run_pt as run_pt_mb
|
||||
|
||||
run_pt_mb(model_args, data_args, training_args, finetuning_args, mb_args, callbacks)
|
||||
else:
|
||||
from .megatron_bridge import run_sft as run_sft_mb
|
||||
|
||||
run_sft_mb(model_args, data_args, training_args, finetuning_args, mb_args, callbacks)
|
||||
|
||||
elif finetuning_args.stage in ["pt", "sft", "dpo"] and finetuning_args.use_mca:
|
||||
if not is_mcore_adapter_available():
|
||||
raise ImportError("mcore_adapter is not installed. Please install it with `pip install mcore-adapter`.")
|
||||
|
||||
Reference in New Issue
Block a user