mirror of
https://github.com/hiyouga/LLaMA-Factory.git
synced 2026-09-06 23:35:46 +08:00
[train] support HyperParallel Expert-Parallel (#10811)
This commit is contained in:
@@ -501,7 +501,7 @@ class FinetuningArguments(
|
|||||||
default=False,
|
default=False,
|
||||||
metadata={
|
metadata={
|
||||||
"help": (
|
"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."
|
"Only supported for the 'pt' and 'sft' stages with full fine-tuning."
|
||||||
)
|
)
|
||||||
},
|
},
|
||||||
@@ -511,7 +511,8 @@ class FinetuningArguments(
|
|||||||
metadata={
|
metadata={
|
||||||
"help": (
|
"help": (
|
||||||
"Path to a JSON file containing HyperParallel strategy arguments "
|
"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,
|
default=1,
|
||||||
metadata={"help": "Context parallel size used when `use_hyper_parallel=True`."},
|
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(
|
use_muon: bool = field(
|
||||||
default=False,
|
default=False,
|
||||||
metadata={"help": "Whether or not to use the Muon optimizer."},
|
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.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.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_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:
|
if self.stage == "ppo" and self.reward_model is None:
|
||||||
raise ValueError("`reward_model` is necessary for PPO training.")
|
raise ValueError("`reward_model` is necessary for PPO training.")
|
||||||
|
|||||||
@@ -42,6 +42,7 @@ from hyper_parallel.integration.llamafactory.context_parallel import (
|
|||||||
get_dp_rank,
|
get_dp_rank,
|
||||||
shard_inputs_for_cp,
|
shard_inputs_for_cp,
|
||||||
)
|
)
|
||||||
|
from hyper_parallel.integration.llamafactory.expert_parallel import ep_prepare_model
|
||||||
from hyper_parallel.platform import get_platform
|
from hyper_parallel.platform import get_platform
|
||||||
from torch import nn
|
from torch import nn
|
||||||
|
|
||||||
@@ -163,10 +164,11 @@ class HyperParallelTrainer(CustomSeq2SeqTrainer):
|
|||||||
raise ValueError("HyperParallel trainer requires Accelerate FSDP2 mode to be enabled.")
|
raise ValueError("HyperParallel trainer requires Accelerate FSDP2 mode to be enabled.")
|
||||||
|
|
||||||
self._cp_size = hp_args.cp_size
|
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._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()
|
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
|
self.ref_model = ref_model
|
||||||
if self.ref_model is not None:
|
if self.ref_model is not None:
|
||||||
self.ref_model = self._prepare_model_for_hyper_parallel(self.ref_model)
|
self.ref_model = self._prepare_model_for_hyper_parallel(self.ref_model)
|
||||||
@@ -176,9 +178,11 @@ class HyperParallelTrainer(CustomSeq2SeqTrainer):
|
|||||||
self._accelerator_patches_active = False
|
self._accelerator_patches_active = False
|
||||||
|
|
||||||
def _prepare_model_for_hyper_parallel(self, model: nn.Module) -> nn.Module:
|
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:
|
if self._cp_size > 1:
|
||||||
model = cp_prepare_model(model, self.accelerator, self._hp_args)
|
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)
|
return fsdp2_prepare_model(self.accelerator, model, self._hp_args)
|
||||||
|
|
||||||
def _activate_accelerator_patches(self) -> None:
|
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:
|
if getattr(hp_args, "cp_size", None) != finetuning_args.hyper_parallel_cp_size:
|
||||||
setattr(hp_args, "cp_size", 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":
|
if hp_args.activation_mode != "none":
|
||||||
model_args.disable_gradient_checkpointing = True
|
model_args.disable_gradient_checkpointing = True
|
||||||
return hp_args
|
return hp_args
|
||||||
|
|||||||
Reference in New Issue
Block a user