diff --git a/configs/1.8B_MoE16_sft.py b/configs/1.8B_MoE16_sft.py index f85302778..47dbfa579 100644 --- a/configs/1.8B_MoE16_sft.py +++ b/configs/1.8B_MoE16_sft.py @@ -203,6 +203,7 @@ weight=dict(size=1, overlap=True), expert=dict(size=-1, no_tp=False), expert_weight=dict(size=1, overlap=True), + expert_zero1=dict(size=-1), ) cudnn_deterministic = False diff --git a/configs/7B_MoE4_sft.py b/configs/7B_MoE4_sft.py index c558427cc..f40ac7fe5 100644 --- a/configs/7B_MoE4_sft.py +++ b/configs/7B_MoE4_sft.py @@ -201,6 +201,7 @@ weight=dict(size=1, overlap=True), expert=dict(size=-1, no_tp=False), expert_weight=dict(size=1, overlap=True), + expert_zero1=dict(size=-1), ) cudnn_deterministic = False diff --git a/internlm/checkpoint/components.py b/internlm/checkpoint/components.py index eee92c9c5..d4d4b5d3f 100644 --- a/internlm/checkpoint/components.py +++ b/internlm/checkpoint/components.py @@ -14,7 +14,7 @@ from internlm.solver.optimizer import HybridZeroOptimizer, HybridZeroOptimizer_v2 from internlm.utils.common import get_current_device from internlm.utils.logger import get_logger -from internlm.utils.parallel import is_using_isp +from internlm.utils.parallel import is_using_isp, is_using_moe from internlm.utils.storage_manager import get_fns, llm_load, llm_save from .utils import ( @@ -310,28 +310,69 @@ def load_optimizer_checkpoint(folder, optim): fns = get_fns(folder) max_tp, max_wp, max_pp, max_zero = 0, 0, 0, 0 + is_moe_optim = False + max_ep, max_ewp, max_moe_zero = 0, 0, 0 for fn in fns: if fn.startswith("optimizer_") and not fn.endswith(".md5"): if is_using_isp(): - _, wp, pp, zero = os.path.splitext(fn)[0].split("_") - max_zero = max(max_zero, int(zero[2:])) - max_wp = max(max_wp, int(wp[2:])) + if fn.startswith("optimizer_ep"): + is_moe_optim = True + _, ep, ewp, pp, moe_zero = os.path.splitext(fn)[0].split("_") + else: + _, wp, pp, zero = os.path.splitext(fn)[0].split("_") + if is_moe_optim: + max_ep = max(max_ep, int(ep[2:])) + max_ewp = max(max_ewp, int(ewp[3:])) + max_moe_zero = max(max_moe_zero, int(moe_zero[2:])) + else: + max_zero = max(max_zero, int(zero[2:])) + max_wp = max(max_wp, int(wp[2:])) max_pp = max(max_pp, int(pp[2:])) else: - _, tp, pp, zero = os.path.splitext(fn)[0].split("_") - max_zero = max(max_zero, int(zero[2:])) + if fn.startswith("optimizer_ep"): + is_moe_optim = True + _, ep, tp, pp, moe_zero = os.path.splitext(fn)[0].split("_") + else: + _, tp, pp, zero = os.path.splitext(fn)[0].split("_") max_tp = max(max_tp, int(tp[2:])) max_pp = max(max_pp, int(pp[2:])) + if is_moe_optim: + max_ep = max(max_ep, int(ep[2:])) + max_moe_zero = max(max_moe_zero, int(moe_zero[2:])) + else: + max_zero = max(max_zero, int(zero[2:])) zero_size = gpc.get_world_size(ParallelMode.ZERO1) tp_size = gpc.get_world_size(ParallelMode.TENSOR) wp_size = gpc.get_world_size(ParallelMode.WEIGHT) pp_size = gpc.get_world_size(ParallelMode.PIPELINE) + ep_size = gpc.get_world_size(ParallelMode.EXPERT) + moe_zero_size = gpc.get_world_size(ParallelMode.EXPERT_ZERO1) + if is_using_isp(): + ewp_size = gpc.get_world_size(ParallelMode.EXPERT_WEIGHT) + ewp_rank = gpc.get_local_rank(ParallelMode.EXPERT_WEIGHT) - assert zero_size == max_zero + 1, ( - f"The optimizer states are save for {max_zero+1} zero parallel, " - f"while current has {zero_size} zero broadcast range." - ) + if is_moe_optim: + assert moe_zero_size == max_moe_zero + 1, ( + f"The optimizer states are save for {max_moe_zero+1} expert zero parallelism, " + f"while current has {moe_zero_size} expert zero broadcast range." + ) + assert ( + ep_size == max_ep + 1 + ), f"The optimizer states are save for {max_ep+1} parallelism, while current has {ep_size} weight parallelism" + if is_using_isp(): + assert ewp_size == max_ewp + 1, ( + f"The optimizer states are save for {max_ewp+1} expert weight parallelism, " + f"while current has {ewp_size} expert weight parallelism" + ) + else: + assert zero_size == max_zero + 1, ( + f"The optimizer states are save for {max_zero+1} zero parallel, " + f"while current has {zero_size} zero broadcast range." + ) + assert ( + wp_size == max_wp + 1 + ), f"The optimizer states are save for {max_wp+1} parallelism, while current has {wp_size} weight parallelism" assert ( pp_size == max_pp + 1 ), f"The optimizer states are save for {max_pp+1} pipelines, while current has {pp_size} pipelines" @@ -339,18 +380,23 @@ def load_optimizer_checkpoint(folder, optim): assert ( tp_size == max_tp + 1 ), f"The optimizer states are save for {max_tp+1} parallelism, while current has {tp_size} tensor parallelism" - assert ( - wp_size == max_wp + 1 - ), f"The optimizer states are save for {max_wp+1} parallelism, while current has {wp_size} weight parallelism" zero_rank = gpc.get_local_rank(ParallelMode.ZERO1) tp_rank = gpc.get_local_rank(ParallelMode.TENSOR) wp_rank = gpc.get_local_rank(ParallelMode.WEIGHT) pp_rank = gpc.get_local_rank(ParallelMode.PIPELINE) + ep_rank = gpc.get_local_rank(ParallelMode.EXPERT) + moe_zero_rank = gpc.get_local_rank(ParallelMode.EXPERT_ZERO1) if is_using_isp(): - fp = f"optimizer_wp{wp_rank}_pp{pp_rank}_zo{zero_rank}.pt" + if is_using_moe() and moe_zero_size * ep_size * ewp_size > zero_size * wp_size: + fp = f"optimizer_ep{ep_rank}_ewp{ewp_rank}_pp{pp_rank}_zo{moe_zero_rank}.pt" + else: + fp = f"optimizer_wp{wp_rank}_pp{pp_rank}_zo{zero_rank}.pt" else: - fp = f"optimizer_tp{tp_rank}_pp{pp_rank}_zo{zero_rank}.pt" + if is_using_moe() and moe_zero_size * ep_size > zero_size: + fp = f"optimizer_ep{ep_rank}_tp{tp_rank}_pp{pp_rank}_zo{moe_zero_rank}.pt" + else: + fp = f"optimizer_tp{tp_rank}_pp{pp_rank}_zo{zero_rank}.pt" states = llm_load(os.path.join(folder, fp), map_location=get_current_device()) @@ -400,17 +446,35 @@ def save_optimizer_checkpoint(optim, state_path): tp_size = gpc.get_world_size(ParallelMode.TENSOR) wp_size = gpc.get_world_size(ParallelMode.WEIGHT) dp_size = gpc.get_world_size(ParallelMode.DATA) + ep_size = gpc.get_world_size(ParallelMode.EXPERT) + ep_rank = gpc.get_local_rank(ParallelMode.EXPERT) + moe_data_size = gpc.get_world_size(ParallelMode.EXPERT_DATA) + moe_zero_size = gpc.get_world_size(ParallelMode.EXPERT_ZERO1) + moe_zero_rank = gpc.get_local_rank(ParallelMode.EXPERT_ZERO1) + if is_using_isp(): + ewp_rank = gpc.get_local_rank(ParallelMode.EXPERT_WEIGHT) + ewp_size = gpc.get_world_size(ParallelMode.EXPERT_WEIGHT) states = optim.state_dict() if isinstance(optim, (HybridZeroOptimizer, HybridZeroOptimizer_v2)): if is_using_isp(): - fp = f"optimizer_wp{wp_rank}_pp{pp_rank}_zo{zero_rank}.pt" - if (gpc.get_global_rank() % (tp_size * dp_size)) < zero_size * wp_size: - llm_save(os.path.join(state_path, fp), states) + if is_using_moe() and moe_zero_size * ep_size * ewp_size > zero_size * wp_size: + fp = f"optimizer_ep{ep_rank}_ewp{ewp_rank}_pp{pp_rank}_zo{moe_zero_rank}.pt" + if (gpc.get_global_rank() % (ewp_size * ep_size * moe_data_size)) < moe_zero_size * ewp_size * ep_size: + llm_save(os.path.join(state_path, fp), states) + else: + fp = f"optimizer_wp{wp_rank}_pp{pp_rank}_zo{zero_rank}.pt" + if (gpc.get_global_rank() % (tp_size * dp_size)) < zero_size * wp_size: + llm_save(os.path.join(state_path, fp), states) else: - fp = f"optimizer_tp{tp_rank}_pp{pp_rank}_zo{zero_rank}.pt" - if (gpc.get_global_rank() % (tp_size * dp_size)) < zero_size * tp_size: - llm_save(os.path.join(state_path, fp), states) + if is_using_moe() and moe_zero_size * ep_size > zero_size: + fp = f"optimizer_ep{ep_rank}_tp{tp_rank}_pp{pp_rank}_zo{moe_zero_rank}.pt" + if (gpc.get_global_rank() % (tp_size * ep_size * moe_data_size)) < moe_zero_size * tp_size * ep_size: + llm_save(os.path.join(state_path, fp), states) + else: + fp = f"optimizer_tp{tp_rank}_pp{pp_rank}_zo{zero_rank}.pt" + if (gpc.get_global_rank() % (tp_size * dp_size)) < zero_size * tp_size: + llm_save(os.path.join(state_path, fp), states) if "zero_devide_optim_plan" in states: params_per_rank_id_dict = states.pop("zero_devide_optim_plan") fp_meta = os.path.join(state_path, optim.rank_unique_id) diff --git a/internlm/core/context/parallel_context.py b/internlm/core/context/parallel_context.py index 2f34785ac..3c9a1603c 100644 --- a/internlm/core/context/parallel_context.py +++ b/internlm/core/context/parallel_context.py @@ -158,6 +158,10 @@ def __init__(self): self.zero1_parallel_size = -1 self.nettest_parallel_size = 1 self.expert_parallel_size = -1 + self.expert_tensor_parallel_size = 1 + self.expert_weight_parallel_size = -1 + self.expert_data_parallel_size = -1 + self.expert_zero1_parallel_size = -1 self.num_processes_on_current_node = -1 self.virtual_pipeline_parallel_size = None self.virtual_pipeline_parallel_rank = None @@ -509,6 +513,8 @@ def init_parallel_groups(self): parallel_config._add_item("weight", dict(size=1, overlap=False)) if "expert" not in parallel_config: parallel_config._add_item("expert", dict(size=-1, no_tp=False)) + if "expert_zero1" not in parallel_config: + parallel_config._add_item("expert_zero1", dict(size=-1)) if "expert_weight" not in parallel_config: parallel_config._add_item("expert_weight", dict(size=1, overlap=False)) # set default value for sequence_2D @@ -530,6 +536,7 @@ def init_parallel_groups(self): self._set_parallel_size_from_config(parallel_config, "pipeline", "pipeline_parallel_size") self._set_parallel_size_from_config(parallel_config, "zero1", "zero1_parallel_size") self._set_parallel_size_from_config(parallel_config, "expert", "expert_parallel_size") + self._set_parallel_size_from_config(parallel_config, "expert_zero1", "expert_zero1_parallel_size") self._set_parallel_size_from_config(parallel_config, "expert_weight", "expert_weight_parallel_size") # the user should not set the data parallel size manually @@ -576,6 +583,20 @@ def init_parallel_groups(self): // self.expert_tensor_parallel_size // self.expert_parallel_size, ) + + if self.expert_zero1_parallel_size == -1: + self.expert_zero1_parallel_size = self.expert_data_parallel_size + self.expert_zero1_parallel_size = max(1, self.expert_zero1_parallel_size) + assert self.expert_zero1_parallel_size <= self.expert_data_parallel_size, ( + f"expert_zero1_parallel_size:{self.expert_zero1_parallel_size} should be less than " + f"expert_data_parallel_size:{self.expert_data_parallel_size}" + ) + assert self.expert_data_parallel_size % self.expert_zero1_parallel_size == 0, ( + f"expert_data_parallel_size:{self.expert_data_parallel_size} % expert_zero1_parallel_size: " + f"{self.expert_zero1_parallel_size} != 0" + ) + assert self.expert_zero1_parallel_size >= 1 + if ( isinstance(parallel_config["tensor"], dict) and parallel_config["tensor"]["mode"] == TensorParallelMode.isp.name @@ -636,6 +657,7 @@ def init_parallel_groups(self): self.expert_tensor_parallel_size, self.expert_weight_parallel_size, self.expert_data_parallel_size, + self.expert_zero1_parallel_size, parallel_config.sequence_2D, ] @@ -661,10 +683,10 @@ def init_parallel_groups(self): if self.pipeline_parallel_size > 1: initializers.append(pgroup_initializer.Initializer_Pipeline(*initializer_args)) if self.config.model.get("num_experts", 1) > 1: - if isinstance(parallel_config["tensor"], dict) and parallel_config["tensor"]["mode"] == "isp": - initializers.append(pgroup_initializer.Initializer_Expert_Weight_Data(*initializer_args)) + if parallel_config["tensor"]["mode"] == TensorParallelMode.isp.name: + initializers.append(pgroup_initializer.Initializer_Expert_Weight_Data_Zero(*initializer_args)) else: - initializers.append(pgroup_initializer.Initializer_Expert_Data(*initializer_args)) + initializers.append(pgroup_initializer.Initializer_Expert_Data_Zero(*initializer_args)) if parallel_config.sequence_2D.get("enable", False) is True: initializers.append(pgroup_initializer.Initializer_2D_SEQUENCE_PARALLEL(*initializer_args)) diff --git a/internlm/core/context/process_group_initializer.py b/internlm/core/context/process_group_initializer.py index fbc3e07a1..0055237df 100644 --- a/internlm/core/context/process_group_initializer.py +++ b/internlm/core/context/process_group_initializer.py @@ -54,6 +54,9 @@ class ParallelMode(Enum): # expert weight parallel EXPERT_WEIGHT = "expert_weight" + # expert zero1 parallel + EXPERT_ZERO1 = "expert_zero1" + # dummy mode, only used during mode construction DUMMY = "dummy" @@ -114,6 +117,7 @@ def __init__( expert_tensor_parallel_size: int, expert_weight_parallel_size: int, expert_data_parallel_size: int, + expert_zero1_parallel_size, sequence_2D_parallel: dict, ): self.rank = rank @@ -130,6 +134,7 @@ def __init__( self.expert_tensor_parallel_size = expert_tensor_parallel_size self.expert_weight_parallel_size = expert_weight_parallel_size self.expert_data_parallel_size = expert_data_parallel_size + self.expert_zero1_parallel_size = expert_zero1_parallel_size self.sequence_2D_parallel = sequence_2D_parallel assert sequence_parallel_size == tensor_parallel_size @@ -508,7 +513,7 @@ def init_dist_group(self, use_cpu: bool = False): return local_rank, group_world_size, process_group, cpu_group, ranks_in_group, mode -class Initializer_Expert_Data(ProcessGroupInitializer): +class Initializer_Expert_Data_Zero(ProcessGroupInitializer): """A ProcessGroupInitializer for expert data parallelism. Args: @@ -528,16 +533,20 @@ def __init__(self, *args, **kwargs): self.real_data_parallel_size = ( self.data_parallel_size * self.tensor_parallel_size // self.expert_tensor_parallel_size ) + self.ranks_num_per_moe_zero = self.expert_data_parallel_size // self.expert_zero1_parallel_size + assert self.real_data_parallel_size % self.expert_parallel_size == 0 + assert self.expert_data_parallel_size % self.expert_zero1_parallel_size == 0 def _get_expert_parallel_ranks(self): """ Create expert and data parallel groups - Example: world_size = 8, tensor_parallel_size = 2, expert_parallel_size = 2 - model_parallel_group = [0,1], [2,3], [4,5], [6,7] - data_parallel_group = [0,2,4,6], [1,3,5,7] - expert_parallel_group = [0,2], [4,6], [1,3], [5,7] - expert_data_parallel_group = [0,4], [2,6], [1,5], [3,7] + Example: world_size = 8, pipeline_parallel_size = 2, expert_parallel_size = 2, expert_zero1_parallel_size = 1, + pipeline_parallel_group = [0,4], [1,5], [2,6], [3,7] + data_parallel_group = [0,1,2,3], [4,5,6,7] + expert_parallel_group = [0,1], [2,3], [4,5], [6,7] + expert_data_parallel_group = [0,2], [1,3], [4,5], [5,7] + expert_zero1_parallel_group = [0],[2], [1],[3], [4],[5], [5],[7] """ data_parallel_groups = [] for i in range(self.pipeline_parallel_size): @@ -551,6 +560,7 @@ def _get_expert_parallel_ranks(self): expert_parallel_groups = [] expert_data_parallel_groups = [] + expert_zero1_parallel_groups = [] for dp_ranks in data_parallel_groups: # partition of expert parallel group, e.g. [0,2], [4,6] part_ep_group = [] @@ -560,8 +570,13 @@ def _get_expert_parallel_ranks(self): for expert_dp_ranks in zip(*part_ep_group): expert_data_parallel_groups.append(list(expert_dp_ranks)) - - return expert_parallel_groups, expert_data_parallel_groups + expert_zero1_parallel_groups.extend( + [ + expert_dp_ranks[i * self.expert_zero1_parallel_size : (i + 1) * self.expert_zero1_parallel_size] + for i in range(self.ranks_num_per_moe_zero) + ] + ) + return expert_parallel_groups, expert_data_parallel_groups, expert_zero1_parallel_groups def init_dist_group(self, use_cpu: bool = False): """Initialize expert parallel and expert data groups, and assign local_ranks and groups to each gpu. @@ -570,14 +585,12 @@ def init_dist_group(self, use_cpu: bool = False): list: [(local_rank, group_world_size, process_group, ranks_in_group, mode), ...]: A length 2 list consists of expert parallelism's and expert data parallelism's information tuple. """ - local_rank = None - ranks_in_group = None - process_group = None - cpu_group = None - group_world_size = None - expert_parallel_groups, expert_data_parallel_groups = self._get_expert_parallel_ranks() - groups = [] + ( + expert_parallel_groups, + expert_data_parallel_groups, + expert_zero1_parallel_groups, + ) = self._get_expert_parallel_ranks() for ranks in expert_parallel_groups: group = dist.new_group(ranks, timeout=LLM_NCCL_TIMEOUT) if use_cpu: @@ -618,6 +631,26 @@ def init_dist_group(self, use_cpu: bool = False): (local_rank, group_world_size, process_group, cpu_group, ranks_in_group, ParallelMode.EXPERT_DATA) ) + for ranks in expert_zero1_parallel_groups: + group = dist.new_group(ranks, timeout=LLM_NCCL_TIMEOUT) + if use_cpu: + group_cpu = ( + dist.new_group(ranks, backend="gloo", timeout=LLM_NCCL_TIMEOUT) + if dist.get_backend() != "gloo" + else group + ) + else: + group_cpu = None + if self.rank in ranks: + local_rank = ranks.index(self.rank) + group_world_size = len(ranks) + process_group = group + cpu_group = group_cpu + ranks_in_group = ranks + groups.append( + (local_rank, group_world_size, process_group, cpu_group, ranks_in_group, ParallelMode.EXPERT_ZERO1) + ) + if self.expert_tensor_parallel_size == 1: ranks = [self.rank] group = dist.new_group(ranks, timeout=LLM_NCCL_TIMEOUT) @@ -642,7 +675,7 @@ def init_dist_group(self, use_cpu: bool = False): return groups -class Initializer_Expert_Weight_Data(ProcessGroupInitializer): +class Initializer_Expert_Weight_Data_Zero(ProcessGroupInitializer): """A ProcessGroupInitializer for common weight's data parallelism. Args: rank (int): The rank of current process. @@ -663,26 +696,23 @@ class Initializer_Expert_Weight_Data(ProcessGroupInitializer): def __init__(self, *args, **kwargs): super().__init__(*args, **kwargs) self.ranks_num_per_pp = self.world_size // self.pipeline_parallel_size + self.ranks_num_per_moe_zero = self.expert_data_parallel_size // self.expert_zero1_parallel_size self.ranks_num_per_dp = self.expert_weight_parallel_size * self.expert_parallel_size assert self.world_size % self.pipeline_parallel_size == 0 assert self.world_size % (self.pipeline_parallel_size * self.expert_data_parallel_size) == 0 - def init_dist_group(self, use_cpu: bool = False): - """Initialize expert parallel groups for isp, and assign local_ranks and groups to each gpu. - Returns: - list: [(local_rank, group_world_size, process_group, ranks_in_group, mode), ...]: - A length 3 list consists of expert parallelism's, expert weight parallelism's - and expert data parallelism's information tuple. + def _get_expert_parallel_ranks(self): + """ Example: n=16 ewp=2 ep=4 edp=2 with nopp expert weight groups: [0, 1], [2, 3], [4, 5], [6, 7], [8, 9], [10, 11], [12, 13], [14, 15] expert groups: [0, 2, 4, 6], [1, 3, 5, 7], [8, 10, 12, 14], [9, 11, 13, 15] expert (weight) data groups:[0, 8], [1, 9], [2, 10], [3, 11], [4, 12], [5, 13], [6, 14], [7, 15] """ - expert_parallel_groups = [] expert_weight_parallel_groups = [] expert_data_parallel_groups = [] + expert_zero1_parallel_groups = [] for i in range(self.pipeline_parallel_size): part_dp_group = [] for j in range(self.expert_data_parallel_size): @@ -694,22 +724,43 @@ def init_dist_group(self, use_cpu: bool = False): ) ) ) - # print(data_parallel_groups) for expert_dp_ranks in zip(*part_dp_group): expert_data_parallel_groups.append(list(expert_dp_ranks)) + expert_zero1_parallel_groups.extend( + [ + expert_dp_ranks[k * self.expert_zero1_parallel_size : (k + 1) * self.expert_zero1_parallel_size] + for k in range(self.ranks_num_per_moe_zero) + ] + ) for dp_groups in part_dp_group: part_wp_group = [] - for i in range(0, self.ranks_num_per_dp, self.expert_weight_parallel_size): - part_wp_group.append(dp_groups[i : i + self.expert_weight_parallel_size]) - # print(part_wp_group) + for k in range(0, self.ranks_num_per_dp, self.expert_weight_parallel_size): + part_wp_group.append(dp_groups[k : k + self.expert_weight_parallel_size]) expert_weight_parallel_groups.extend(part_wp_group) for ep_ranks in zip(*part_wp_group): expert_parallel_groups.append(list(ep_ranks)) - # print(expert_weight_parallel_groups) - # print(expert_parallel_groups) - # print(expert_data_parallel_groups) + return ( + expert_parallel_groups, + expert_weight_parallel_groups, + expert_data_parallel_groups, + expert_zero1_parallel_groups, + ) + + def init_dist_group(self, use_cpu: bool = False): + """Initialize expert parallel groups for isp, and assign local_ranks and groups to each gpu. + Returns: + list: [(local_rank, group_world_size, process_group, ranks_in_group, mode), ...]: + A length 3 list consists of expert parallelism's, expert weight parallelism's + and expert data parallelism's information tuple. + """ groups = [] + ( + expert_parallel_groups, + expert_weight_parallel_groups, + expert_data_parallel_groups, + expert_zero1_parallel_groups, + ) = self._get_expert_parallel_ranks() for ranks in expert_parallel_groups: group = dist.new_group(ranks, timeout=LLM_NCCL_TIMEOUT) if use_cpu: @@ -750,6 +801,26 @@ def init_dist_group(self, use_cpu: bool = False): (local_rank, group_world_size, process_group, cpu_group, ranks_in_group, ParallelMode.EXPERT_DATA) ) + for ranks in expert_zero1_parallel_groups: + group = dist.new_group(ranks, timeout=LLM_NCCL_TIMEOUT) + if use_cpu: + group_cpu = ( + dist.new_group(ranks, backend="gloo", timeout=LLM_NCCL_TIMEOUT) + if dist.get_backend() != "gloo" + else group + ) + else: + group_cpu = None + if self.rank in ranks: + local_rank = ranks.index(self.rank) + group_world_size = len(ranks) + process_group = group + cpu_group = group_cpu + ranks_in_group = ranks + groups.append( + (local_rank, group_world_size, process_group, cpu_group, ranks_in_group, ParallelMode.EXPERT_ZERO1) + ) + for ranks in expert_weight_parallel_groups: group = dist.new_group(ranks, timeout=LLM_NCCL_TIMEOUT) if use_cpu: diff --git a/internlm/initialize/launch.py b/internlm/initialize/launch.py index c85568ec4..f3b0e20e3 100644 --- a/internlm/initialize/launch.py +++ b/internlm/initialize/launch.py @@ -15,6 +15,7 @@ from internlm.utils.common import get_master_node from internlm.utils.gputest import warmup_process_group from internlm.utils.logger import get_logger +from internlm.utils.parallel import is_using_moe from internlm.utils.timeout import llm_timeout from internlm.utils.utils import DataType, ModelType, TensorParallelMode @@ -507,7 +508,7 @@ def args_sanity_check(): gpc.config._add_item("selective_checkpoint", False) # moe not support overlap and zero1.5 for now - if gpc.config.model.get("num_experts", 1) > 1: + if is_using_moe(): assert not gpc.config.parallel.zero1.fsdp, "FSDP does not support num_experts > 1" assert ( not optim_ckpt.overlap_sync_grad & optim_ckpt.overlap_sync_param @@ -616,7 +617,7 @@ def launch( f"data parallel size: {gpc.data_parallel_size}, pipeline parallel size: {gpc.pipeline_parallel_size}, " f"tensor parallel size: {gpc.tensor_parallel_size}, weight parallel size: {gpc.weight_parallel_size}", ) - if gpc.config.model.get("num_experts", 1) > 1: + if is_using_moe(): logger.info( f"Creating MoE with num_experts: {gpc.config.model.num_experts} | " f"expert parallel size: {gpc.expert_parallel_size} | " diff --git a/internlm/solver/optimizer/utils.py b/internlm/solver/optimizer/utils.py index a0180a596..b9eaee3a9 100644 --- a/internlm/solver/optimizer/utils.py +++ b/internlm/solver/optimizer/utils.py @@ -360,7 +360,7 @@ def compute_norm(gradients, parameters, norm_type=2, zero_mode=ParallelMode.ZERO # Need to allreduce(avg) the norms across different ranks because moe params will not be synced during allreduce # model and zero have been reduced!!! - if zero_mode == ParallelMode.EXPERT_DATA: + if zero_mode == ParallelMode.EXPERT_ZERO1: pg = gpc.get_group(ParallelMode.EXPERT) scaled_norm = total_norm * 1.0 / float(gpc.get_world_size(ParallelMode.EXPERT)) scaled_norm_tensor = torch.tensor(scaled_norm, device=get_current_device(), dtype=torch.float) diff --git a/internlm/train/pipeline.py b/internlm/train/pipeline.py index 53057128b..80a8101ff 100644 --- a/internlm/train/pipeline.py +++ b/internlm/train/pipeline.py @@ -86,6 +86,7 @@ is_tensor_expert_data_parallel_parameter, is_tensor_zero_parallel_parameter, is_using_isp, + is_using_moe, is_weight_expert_data_parallel_parameter, is_weight_zero_parallel_parameter, sync_model_param, @@ -337,7 +338,7 @@ def initialize_parallel_communicator(model: Union[nn.Module, nn.ModuleList]): ) _embedding_communicator = EmbeddingWeightParallelCommunicator(ParallelMode.WEIGHT) - if gpc.config.model.get("num_experts", 1) > 1: + if is_using_moe(): # register communicator for moe isp column parallel linear. # NOTE: this wil overwrite registed communicator moe_isp_communicator = ISPCommunicator( @@ -372,7 +373,7 @@ def initialize_parallel_communicator(model: Union[nn.Module, nn.ModuleList]): TensorParallelCommunicator(process_group=gpc.get_group(ParallelMode.TENSOR), role=LinearRole.ROW) ) - if gpc.config.model.get("num_experts", 1) > 1: + if is_using_moe(): GroupedColumnLinear.register_cls_communicator( TensorParallelCommunicator(process_group=gpc.get_group(ParallelMode.TENSOR), role=LinearRole.COLUMN) ) @@ -417,7 +418,7 @@ def initialize_parallel_communicator(model: Union[nn.Module, nn.ModuleList]): save_total_input_as_activation=save_total_input_as_activation, ) ) - if gpc.config.model.get("num_experts", 1) > 1: + if is_using_moe(): GroupedColumnLinear.register_cls_communicator( SequenceParallelCommunicator( process_group=gpc.get_group(ParallelMode.TENSOR), diff --git a/internlm/train/utils.py b/internlm/train/utils.py index d1bf4fe90..24e6e0818 100644 --- a/internlm/train/utils.py +++ b/internlm/train/utils.py @@ -8,6 +8,7 @@ from internlm.core.naive_amp import unwrap_naive_amp from internlm.model.modules.utils import is_moe_param from internlm.utils.logger import get_logger +from internlm.utils.parallel import is_using_moe logger = get_logger(__file__) @@ -45,9 +46,9 @@ def split_params_into_different_groups_for_optimizer( # create new groups for fp32 parameter group new_groups["fp32"] = {"name": "fp32", "params": [], "optimizer_mode": ParallelMode.ZERO1} - if gpc.config.model.get("num_experts", 1) > 1: + if is_using_moe(): for key in gpc.expert_parallel_group_names: - new_groups[key] = {"name": key, "moe": True, "params": [], "optimizer_mode": ParallelMode.EXPERT_DATA} + new_groups[key] = {"name": key, "moe": True, "params": [], "optimizer_mode": ParallelMode.EXPERT_ZERO1} for pgroup in param_groups: # copy attribute from origin group, we assume the input param_groups only diff --git a/internlm/utils/parallel.py b/internlm/utils/parallel.py index 665353070..5404467a6 100644 --- a/internlm/utils/parallel.py +++ b/internlm/utils/parallel.py @@ -31,6 +31,10 @@ def is_using_isp(): ) +def is_using_moe(): + return gpc.config.model.get("num_experts", 1) > 1 + + def is_replica_zero_parallel_parameter(p): return hasattr(p, IS_REPLICA_ZERO_PARALLEL) and getattr(p, IS_REPLICA_ZERO_PARALLEL)