mirror of
https://github.com/hiyouga/LLaMA-Factory.git
synced 2026-09-06 15:25:43 +08:00
[train] support HyperParallel Expert-Parallel (#10811)
This commit is contained in:
@@ -501,7 +501,7 @@ class FinetuningArguments(
|
||||
default=False,
|
||||
metadata={
|
||||
"help": (
|
||||
"Whether or not to use HyperParallel distributed training backend (FSDP/TP). "
|
||||
"Whether or not to use HyperParallel distributed training backend (FSDP/CP/EP). "
|
||||
"Only supported for the 'pt' and 'sft' stages with full fine-tuning."
|
||||
)
|
||||
},
|
||||
@@ -511,7 +511,8 @@ class FinetuningArguments(
|
||||
metadata={
|
||||
"help": (
|
||||
"Path to a JSON file containing HyperParallel strategy arguments "
|
||||
"(e.g., tp_size, param_dtype). Used when use_hyper_parallel=True."
|
||||
"(e.g., cp_size, ep_size, efsdp_size, token_dispatcher, param_dtype). "
|
||||
"Used when use_hyper_parallel=True."
|
||||
)
|
||||
},
|
||||
)
|
||||
@@ -519,6 +520,28 @@ class FinetuningArguments(
|
||||
default=1,
|
||||
metadata={"help": "Context parallel size used when `use_hyper_parallel=True`."},
|
||||
)
|
||||
hyper_parallel_ep_size: int = field(
|
||||
default=1,
|
||||
metadata={"help": "Expert parallel size used when `use_hyper_parallel=True`."},
|
||||
)
|
||||
hyper_parallel_efsdp_size: int = field(
|
||||
default=1,
|
||||
metadata={
|
||||
"help": (
|
||||
"Expert FSDP shard size used when `use_hyper_parallel=True`. "
|
||||
"Defaults to world size divided by expert parallel size."
|
||||
)
|
||||
},
|
||||
)
|
||||
hyper_parallel_token_dispatcher: Literal["all_to_all"] = field(
|
||||
default="all_to_all",
|
||||
metadata={
|
||||
"help": (
|
||||
"Expert token dispatcher used when 'use_hyper_parallel=True'. "
|
||||
"Currently only 'all_to_all' is supported."
|
||||
)
|
||||
},
|
||||
)
|
||||
use_muon: bool = field(
|
||||
default=False,
|
||||
metadata={"help": "Whether or not to use the Muon optimizer."},
|
||||
@@ -596,6 +619,7 @@ class FinetuningArguments(
|
||||
assert self.ref_model_quantization_bit in [None, 8, 4], "We only accept 4-bit or 8-bit quantization."
|
||||
assert self.reward_model_quantization_bit in [None, 8, 4], "We only accept 4-bit or 8-bit quantization."
|
||||
assert self.hyper_parallel_cp_size > 0, "`hyper_parallel_cp_size` must be greater than 0."
|
||||
assert self.hyper_parallel_ep_size > 0, "`hyper_parallel_ep_size` must be greater than 0."
|
||||
|
||||
if self.stage == "ppo" and self.reward_model is None:
|
||||
raise ValueError("`reward_model` is necessary for PPO training.")
|
||||
|
||||
@@ -42,6 +42,7 @@ from hyper_parallel.integration.llamafactory.context_parallel import (
|
||||
get_dp_rank,
|
||||
shard_inputs_for_cp,
|
||||
)
|
||||
from hyper_parallel.integration.llamafactory.expert_parallel import ep_prepare_model
|
||||
from hyper_parallel.platform import get_platform
|
||||
from torch import nn
|
||||
|
||||
@@ -163,10 +164,11 @@ class HyperParallelTrainer(CustomSeq2SeqTrainer):
|
||||
raise ValueError("HyperParallel trainer requires Accelerate FSDP2 mode to be enabled.")
|
||||
|
||||
self._cp_size = hp_args.cp_size
|
||||
self._ep_size = hp_args.ep_size
|
||||
self._cp_rank = get_cp_rank(hp_args) if self._cp_size > 1 else 0
|
||||
self._dp_rank = get_dp_rank(hp_args) if self._cp_size > 1 else get_platform().get_rank()
|
||||
|
||||
# Prepare ref_model with the same CP + HSDP path as the train model.
|
||||
# Prepare ref_model with the same CP + EP + HSDP path as the train model.
|
||||
self.ref_model = ref_model
|
||||
if self.ref_model is not None:
|
||||
self.ref_model = self._prepare_model_for_hyper_parallel(self.ref_model)
|
||||
@@ -176,9 +178,11 @@ class HyperParallelTrainer(CustomSeq2SeqTrainer):
|
||||
self._accelerator_patches_active = False
|
||||
|
||||
def _prepare_model_for_hyper_parallel(self, model: nn.Module) -> nn.Module:
|
||||
"""Apply CP runtime hooks before delegating to HyperParallel FSDP2 preparation."""
|
||||
"""Apply CP/EP preparation before delegating to HyperParallel FSDP2."""
|
||||
if self._cp_size > 1:
|
||||
model = cp_prepare_model(model, self.accelerator, self._hp_args)
|
||||
if self._ep_size > 1:
|
||||
model = ep_prepare_model(model, self.accelerator, self._hp_args)
|
||||
return fsdp2_prepare_model(self.accelerator, model, self._hp_args)
|
||||
|
||||
def _activate_accelerator_patches(self) -> None:
|
||||
|
||||
@@ -54,6 +54,17 @@ def _prepare_hp_args(finetuning_args: "FinetuningArguments", model_args: "ModelA
|
||||
if getattr(hp_args, "cp_size", None) != finetuning_args.hyper_parallel_cp_size:
|
||||
setattr(hp_args, "cp_size", finetuning_args.hyper_parallel_cp_size)
|
||||
|
||||
if getattr(hp_args, "ep_size", None) != finetuning_args.hyper_parallel_ep_size:
|
||||
setattr(hp_args, "ep_size", finetuning_args.hyper_parallel_ep_size)
|
||||
|
||||
if getattr(hp_args, "efsdp_size", None) != finetuning_args.hyper_parallel_efsdp_size:
|
||||
setattr(hp_args, "efsdp_size", finetuning_args.hyper_parallel_efsdp_size)
|
||||
|
||||
if getattr(hp_args, "token_dispatcher", None) != finetuning_args.hyper_parallel_token_dispatcher:
|
||||
setattr(hp_args, "token_dispatcher", finetuning_args.hyper_parallel_token_dispatcher)
|
||||
|
||||
hp_args.validate()
|
||||
|
||||
if hp_args.activation_mode != "none":
|
||||
model_args.disable_gradient_checkpointing = True
|
||||
return hp_args
|
||||
|
||||
Reference in New Issue
Block a user