diff --git a/src/llamafactory/hparams/finetuning_args.py b/src/llamafactory/hparams/finetuning_args.py index f4bd4debd..007b9674e 100644 --- a/src/llamafactory/hparams/finetuning_args.py +++ b/src/llamafactory/hparams/finetuning_args.py @@ -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.") diff --git a/src/llamafactory/train/hyper_parallel/trainer.py b/src/llamafactory/train/hyper_parallel/trainer.py index 1410073a5..c05099622 100644 --- a/src/llamafactory/train/hyper_parallel/trainer.py +++ b/src/llamafactory/train/hyper_parallel/trainer.py @@ -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: diff --git a/src/llamafactory/train/hyper_parallel/workflow.py b/src/llamafactory/train/hyper_parallel/workflow.py index 4eb4cc99b..7f2efcd90 100644 --- a/src/llamafactory/train/hyper_parallel/workflow.py +++ b/src/llamafactory/train/hyper_parallel/workflow.py @@ -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