[train] support HyperParallel Expert-Parallel (#10811)

This commit is contained in:
Pangxz
2026-09-04 16:15:23 +08:00
committed by GitHub
parent 4451765a6b
commit dced5f8804
3 changed files with 43 additions and 4 deletions

View File

@@ -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.")

View File

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

View File

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