diff --git a/ds_pynative.yaml b/ds_pynative.yaml index e76328464fe7c9c9269aaf7c62bdc139f7a5b654..1ab6a1ebec5f5f5da3fd39e3c509d6e5e30e388c 100644 --- a/ds_pynative.yaml +++ b/ds_pynative.yaml @@ -108,7 +108,7 @@ train_dataset: &train_dataset pad: -1 # The token id of `pad` in the dataset. data_path: # Megatron dataset sampling ratio and path. - '1' - - "/home/l00913161/data/deepseek-datasets/mmap_deepseekv3_datasets_text_document" + - "/home/w00932055/dsv4/deepseek-datasets/mmap_deepseekv3_datasets_text_document" input_columns: ["input_ids", "labels", "loss_mask", "position_ids"] construct_args_key: ["input_ids", "labels", "loss_mask", "position_ids"] num_parallel_workers: 8 @@ -234,15 +234,19 @@ model: n_routed_experts: 16 num_experts_per_tok: 8 n_shared_experts: 1 + num_copy_experts: 0 + use_topk_router_with_load_balancing: False + moe_expected_ffn_experts: 2.0 # Default best value: top-k * FFN/(FFN + COPY) + moe_router_bias_update_rate: 0.001 moe_shared_expert_intermediate_size: 2048 moe_grouped_gemm: True moe_router_load_balancing_type: 'seq_aux_loss' - moe_aux_loss_coeff: 0. # 0.001 + moe_aux_loss_coeff: 0.001 # 0.001 scoring_func: 'sigmoid' norm_topk_prob: True moe_token_drop_policy: probs moe_router_enable_expert_bias: True - moe_router_bias_update_rate: 0. # 0.001 + moe_router_bias_update_rate: 0.001 # 0.001 # callbacks callbacks: diff --git a/mindformers/parallel_core/transformer_config.py b/mindformers/parallel_core/transformer_config.py index cd0b6df0e2e207b5b59ad80d2afc236b172493d2..dabf8f5fff5049127caef7fac24faf118d966dd5 100644 --- a/mindformers/parallel_core/transformer_config.py +++ b/mindformers/parallel_core/transformer_config.py @@ -254,7 +254,7 @@ class TransformerConfig(ModelParallelConfig, MFModelConfig): The default is "sub_seq_aux_loss". """ - moe_router_topk: int = 2 + moe_router_topk: int = 4 """Number of experts to route to for each token.""" moe_router_num_groups: Optional[int] = None @@ -354,6 +354,14 @@ class TransformerConfig(ModelParallelConfig, MFModelConfig): moe_apply_probs_on_input: bool = False """Apply probs on input of experts instead of applying after activation and glu.""" + num_copy_experts: int = 0 + + use_topk_router_with_load_balancing: bool = True + + moe_expected_ffn_experts: float = 2.0 # Default best value: top-k * FFN/(FFN + COPY) + + moe_router_bias_update_rate: float = 0.001 + # MindFormers New shared_expert_num: int = 0 """Number of shared experts.""" diff --git a/mindformers/parallel_core/transformer_config_utils.py b/mindformers/parallel_core/transformer_config_utils.py index 3f00c66b54e7a3c4ee18951919d73b5557b0d35f..2d0889e4c9b223b4a472928fe221c76118943710 100644 --- a/mindformers/parallel_core/transformer_config_utils.py +++ b/mindformers/parallel_core/transformer_config_utils.py @@ -382,7 +382,11 @@ COMMON_CONFIG_MAPPING = { "enable_expert_relocation": "enable_expert_relocation", "expert_relocation_initial_iteration": "expert_relocation_initial_iteration", "expert_relocation_freq": "expert_relocation_freq", - + "num_copy_experts": "num_copy_experts", + "use_topk_router_with_load_balancing": "use_topk_router_with_load_balancing", + "moe_expected_ffn_experts": "moe_expected_ffn_experts", + "moe_router_bias_update_rate": "moe_router_bias_update_rate", + # Context Parallel # not changes "context_parallel_algo": ("cp_comm_type", get_cp_comm_type), diff --git a/mindformers/pynative/base_models/gpt/moe_module_specs.py b/mindformers/pynative/base_models/gpt/moe_module_specs.py index 9d540addb1105615113aeb276e7493b2eb851231..09e19ff8cdb7b5977f8b33feb8367df88af9c446 100644 --- a/mindformers/pynative/base_models/gpt/moe_module_specs.py +++ b/mindformers/pynative/base_models/gpt/moe_module_specs.py @@ -31,9 +31,6 @@ def get_moe_module_spec( if num_experts is None: raise ValueError("num_experts cannot be None.") - # experts spec - if not moe_grouped_gemm: - raise NotImplementedError("moe_grouped_gemm = 'False' is not supported now.") shared_experts = SharedExpertMLP diff --git a/mindformers/pynative/transformers/moe/experts.py b/mindformers/pynative/transformers/moe/experts.py index 54c452272500720fead31ce4b2531af81f1887c5..f1e31a4cb7bb71aab8bb6b8494bc8a9bf4bc19a5 100644 --- a/mindformers/pynative/transformers/moe/experts.py +++ b/mindformers/pynative/transformers/moe/experts.py @@ -39,6 +39,14 @@ class GroupedMLP(nn.Cell): self.config = config self.num_local_experts = self.config.num_moe_experts self.top_k = self.config.moe_router_topk + self.num_copy_experts = self.config.num_copy_experts + + if self.num_copy_experts < 0 or self.num_copy_experts > self.num_local_experts: + raise ValueError( + f"num_copy_experts must be in [0, num_moe_experts], but got {self.num_copy_experts}." + ) + self.copy_expert_start = self.num_local_experts - self.num_copy_experts + self.num_ffn_experts = self.copy_expert_start if self.config.moe_apply_probs_on_input: if self.config.moe_router_topk == 1: @@ -64,15 +72,15 @@ class GroupedMLP(nn.Cell): self.moe_ffn_hidden_size *= 2 self.mul = mint.mul self.moe_token_dispatcher_type = config.moe_token_dispatcher_type - self.moe_use_experts_for_loop = True # config.moe_use_experts_for_loop + self.moe_use_experts_for_loop = not config.moe_grouped_gemm self.init_method = config.init_method # parameters self.weight1 = Parameter( - self.init_method([self.num_local_experts * self.hidden_size, self.moe_ffn_hidden_size]), + self.init_method([self.num_ffn_experts * self.hidden_size, self.moe_ffn_hidden_size]), name='w1') self.weight2 = Parameter( - self.init_method([self.num_local_experts * self.config.moe_ffn_hidden_size, self.hidden_size]), + self.init_method([self.num_ffn_experts * self.config.moe_ffn_hidden_size, self.hidden_size]), name='w2') self.cast = ops.cast @@ -183,10 +191,20 @@ class GroupedMLP(nn.Cell): # Process each expert out_experts_splits = [] for expert_idx, x_expert in enumerate(x_splits): + # CopyExpert: identity mapping (no FFN), keep routing weights + if expert_idx >= self.copy_expert_start: + h = x_expert + if not self.config.moe_apply_probs_on_input: + h = self.mul(h, permuted_probs_splits[expert_idx].reshape(-1, 1)) + out_experts_splits.append(h) + continue h = self.matmul(x_expert, w1[expert_idx]) - h1, h2 = self.chunk(h, 2, -1) - h1 = self.activation_func(h1) - h = self.mul(h1, h2) + if self.activation_type == 'fusedswiglu': + h = self.activation_func(h, -1).reshape((-1, w2.shape[1])) + else: + x0, x1 = self.chunk(h, 2, -1) + act_out = self.activation_func(x0) + h = self.mul(act_out, x1) h = self.mul(h, permuted_probs_splits[expert_idx].reshape(-1, 1)) h = self.matmul(h, w2[expert_idx]) out_experts_splits.append(h) @@ -206,23 +224,47 @@ class GroupedMLP(nn.Cell): if self.moe_use_experts_for_loop: return self._run_experts_for_loop(w1, w2, permuted_local_hidden_states, tokens_per_expert, permuted_probs) - # Original grouped_mm implementation - tokens_per_expert = self.cumsum(tokens_per_expert, dim=0, dtype=ms.int64) - fc1_output = GroupedMatmul(split_item=3, group_type=0)( - [permuted_local_hidden_states], [w1], None, None, None, None, None, tokens_per_expert)[0] + # Split tokens into non-copy experts and copy experts to avoid extra compute + counts_list = tokens_per_expert.asnumpy().tolist() + non_copy_experts = self.num_ffn_experts + non_copy_tokens = sum(counts_list[:non_copy_experts]) + copy_tokens = sum(counts_list[non_copy_experts:]) if self.num_copy_experts > 0 else 0 - if self.gated_linear_unit: - if self.activation_type == 'fusedswiglu': - intermediate_parallel = self.activation_func(fc1_output, -1).reshape((-1, w2.shape[1])) + outputs = [] + if non_copy_experts > 0 and non_copy_tokens > 0: + # Run GroupedGEMM only for non-copy experts + non_copy_input = permuted_local_hidden_states[:non_copy_tokens] + non_copy_probs = permuted_probs[:non_copy_tokens] + tokens_per_expert_nc = tokens_per_expert[:non_copy_experts] + tokens_per_expert_nc = self.cumsum(tokens_per_expert_nc, dim=0, dtype=ms.int64) + + fc1_output = GroupedMatmul(split_item=3, group_type=0)( + [non_copy_input], [w1], None, None, None, None, None, tokens_per_expert_nc)[0] + + if self.gated_linear_unit: + if self.activation_type == 'fusedswiglu': + intermediate_parallel = self.activation_func(fc1_output, -1).reshape((-1, w2.shape[1])) + else: + x0, x1 = self.chunk(fc1_output, 2, -1) + act_out = self.activation_func(x0) + intermediate_parallel = self.mul(act_out, x1) else: - x0, x1 = self.chunk(fc1_output, 2, -1) - act_out = self.activation_func(x0) - intermediate_parallel = self.mul(act_out, x1) - else: - intermediate_parallel = self.activation_func(fc1_output) - - permuted_probs = self.cast(permuted_probs, intermediate_parallel.dtype) - intermediate_parallel = self.mul(intermediate_parallel, self.unsqueeze(permuted_probs, -1)) - fc2_output = GroupedMatmul(split_item=3, group_type=0)( - [intermediate_parallel], [w2], None, None, None, None, None, tokens_per_expert)[0] - return fc2_output + intermediate_parallel = self.activation_func(fc1_output) + + non_copy_probs = self.cast(non_copy_probs, intermediate_parallel.dtype) + intermediate_parallel = self.mul(intermediate_parallel, self.unsqueeze(non_copy_probs, -1)) + fc2_output = GroupedMatmul(split_item=3, group_type=0)( + [intermediate_parallel], [w2], None, None, None, None, None, tokens_per_expert_nc)[0] + outputs.append(fc2_output) + + if self.num_copy_experts > 0 and copy_tokens > 0: + copy_segment = permuted_local_hidden_states[non_copy_tokens: non_copy_tokens + copy_tokens] + if not self.config.moe_apply_probs_on_input: + copy_probs = self.cast(permuted_probs[non_copy_tokens: non_copy_tokens + copy_tokens], + permuted_local_hidden_states.dtype) + copy_segment = self.mul(copy_segment, copy_probs.reshape(-1, 1)) + outputs.append(copy_segment) + + if outputs: + return self.cat(outputs, dim=0) + return permuted_local_hidden_states[:0] diff --git a/mindformers/pynative/transformers/moe/moe_layer.py b/mindformers/pynative/transformers/moe/moe_layer.py index 2290d92f5c4fe36d0ed3acf98fba2e9c3a8cdaab..1429f517e0931631a091661df081f16d7ad754d8 100644 --- a/mindformers/pynative/transformers/moe/moe_layer.py +++ b/mindformers/pynative/transformers/moe/moe_layer.py @@ -20,7 +20,7 @@ from mindspore.common.parameter import Parameter from mindformers.parallel_core.transformer_config import TransformerConfig from mindformers.pynative.layers.linear import Linear from mindformers.pynative.transformers.mlp import MLPSubmodules -from .router import TopKRouter +from .router import TopKRouterWithLoadBalancing, TopKRouter from .experts import GroupedMLP from .shared_experts import SharedExpertMLP @@ -37,7 +37,12 @@ class MoELayer(nn.Cell): self.top_k = config.moe_router_topk # Router - self.router = TopKRouter(config) + # Prefer TopKRouterWithLoadBalancing when copy experts or aux loss is enabled + use_topk_router_with_load_balancing = self.config.use_topk_router_with_load_balancing + if use_topk_router_with_load_balancing: + self.router = TopKRouterWithLoadBalancing(config) + else: + self.router = TopKRouter(config) # Experts self.experts = GroupedMLP(config) @@ -94,7 +99,15 @@ class MoELayer(nn.Cell): bs, slen, dim = hidden_states.shape x_flat = self.reshape(hidden_states, (-1, dim)) - top_scores, selected_experts_indices, num_tokens_per_expert = self.router(x_flat, self.expert_bias) + if isinstance(self.router, TopKRouter): + top_scores, selected_experts_indices, num_tokens_per_expert = self.router( + x_flat, self.expert_bias + ) + aux_loss = None + else: + top_scores, selected_experts_indices, num_tokens_per_expert, aux_loss = self.router( + x_flat, training=self.training + ) self.tokens_per_expert.add_(num_tokens_per_expert) @@ -111,4 +124,4 @@ class MoELayer(nn.Cell): else: final_out = out_experts - return final_out, None + return final_out, aux_loss diff --git a/mindformers/pynative/transformers/moe/router.py b/mindformers/pynative/transformers/moe/router.py index 8c8d469c38490e9a0b33d0e13961b5310cf7026d..6ebbf4e6d1901c1e5b5da72d9747511f1b1affe7 100644 --- a/mindformers/pynative/transformers/moe/router.py +++ b/mindformers/pynative/transformers/moe/router.py @@ -220,3 +220,280 @@ class TopKRouter(nn.Cell): ) return top_scores, selected_experts_indices, num_tokens_per_expert + + +class TopKRouterWithLoadBalancing(nn.Cell): + """ + Routing mechanism with Computational Budget Control and + Load Balance Control. + + Computational Budget Control: + - PID-like bias update to maintain expected FFN activation rate. + - Bias updates are excluded for zero-computation (copy) experts. + + Load Balance Control: + - Auxiliary loss computed over expert groups. + - FFN experts are divided into D groups; copy experts form group D+1. + """ + + def __init__(self, config: TransformerConfig): + super().__init__() + self.config = config + + # Dimensions + self.hidden_size = config.hidden_size + self.num_total_experts = config.num_moe_experts + self.top_k = config.moe_router_topk + + # examine numbers of experts + self.num_copy_experts = self.config.num_copy_experts + if self.num_copy_experts < 0 or self.num_copy_experts > self.num_total_experts: + raise ValueError( + f"num_copy_experts must be in [0, num_moe_experts], but got {self.num_copy_experts}." + ) + self.num_ffn_experts = self.num_total_experts - self.num_copy_experts + if self.num_ffn_experts <= 0: + raise ValueError( + f"num_ffn_experts must be positive, got {self.num_ffn_experts}" + ) + + # Expected number of activated FFN experts + expected_ffn_k = getattr( + config, + "moe_expected_ffn_experts", + float(self.top_k * self.num_ffn_experts / self.num_total_experts), + ) + try: + self.expected_ffn_k = float(expected_ffn_k) + except (TypeError, ValueError) as exc: + raise ValueError( + f"moe_expected_ffn_experts must be a number, but got {expected_ffn_k}" + ) from exc + + # Bias update rate for computational budget control + self.bias_update_rate = self.config.moe_router_bias_update_rate + # PID gains and integral decay + self.pid_p = float(getattr(config, "moe_router_pid_p", 1.0)) + self.pid_i = float(getattr(config, "moe_router_pid_i", 0.01)) + self.pid_d = float(getattr(config, "moe_router_pid_d", 0.05)) + self.pid_i_decay = float(getattr(config, "moe_router_pid_i_decay", 0.9)) + + # Load balance control (grouped) + self.num_ffn_groups = config.moe_router_num_groups or 1 + if self.num_ffn_groups <= 0: + raise ValueError(f"moe_router_num_groups must be > 0, got {self.num_ffn_groups}") + if self.num_ffn_experts % self.num_ffn_groups != 0: + raise ValueError( + f"num_ffn_experts ({self.num_ffn_experts}) must be divisible by " + f"moe_router_num_groups ({self.num_ffn_groups})" + ) + self.aux_loss_coeff = config.moe_aux_loss_coeff + + # Routing parameters + self.score_func = config.moe_router_score_function + self.route_norm = config.norm_topk_prob + self.route_scale = ( + config.moe_router_topk_scaling_factor + if config.moe_router_topk_scaling_factor is not None + else 1.0 + ) + + # Learnable router weight + self.weight = Parameter( + init_method_normal(0.02)((self.num_total_experts, self.hidden_size)), + name="weight", + ) + + # Expert bias (updated by controller; not trained by gradients) + self.expert_bias = Parameter( + mint.zeros(self.num_total_experts, dtype=mstype.float32), + name="expert_bias", + requires_grad=False, + ) + # PID controller states + self.integral = Parameter( + mint.zeros(self.num_total_experts, dtype=mstype.float32), + name="expert_bias_integral", + requires_grad=False, + ) + self.previous_error = Parameter( + mint.zeros(self.num_total_experts, dtype=mstype.float32), + name="expert_bias_prev_error", + requires_grad=False, + ) + + # Ops + self.linear = mint.nn.functional.linear + self.sigmoid = mint.nn.functional.sigmoid + self.softmax = mint.nn.functional.softmax + self.cast = ops.cast + self.topk = mint.topk + self.gather = mint.gather + self.mul = mint.mul + self.div = mint.div + self.sum = mint.sum + self.histc = mint.histc + self.ones_like = mint.ones_like + self.reshape = mint.reshape + self.cat = mint.cat + self.zeros = mint.zeros + + def update_expert_bias(self, num_tokens_per_expert: Tensor, total_tokens: int): + """ + Bias update for Computational Budget Control. + delta_b_i = mu * (target_rate - current_rate) + Only updates FFN experts; copy experts are excluded. + """ + if not self.training or self.bias_update_rate <= 0 or total_tokens <= 0: + return + + # Current selection rate: Ti / (K * T_all) + current_rate = self.cast(num_tokens_per_expert, mstype.float32) / ( + self.top_k * total_tokens + ) + + # Target selection rate for FFN experts: Ke / (K * N) + target_rate_ffn = self.expected_ffn_k / (self.top_k * self.num_ffn_experts) + target_rates = self.ones_like(current_rate) * target_rate_ffn + + error = target_rates - current_rate + + # Mask out copy experts (no bias update) + if self.num_copy_experts > 0: + mask = self.zeros((self.num_total_experts,), dtype=mstype.float32) + mask[: self.num_ffn_experts] = 1.0 + error = self.mul(error, mask) + + # PID terms + p_term = self.mul(error, self.pid_p) + integral = self.mul(self.integral, self.pid_i_decay) + self.mul( + error, 1.0 - self.pid_i_decay + ) + i_term = self.mul(integral, self.pid_i) + d_term = self.mul(error - self.previous_error, self.pid_d) + + update_step = self.bias_update_rate * (p_term + i_term + d_term) + + ops.assign_add(self.expert_bias, update_step) + ops.assign(self.integral, integral) + ops.assign(self.previous_error, error) + + def compute_load_balancing_loss(self, router_probs: Tensor, expert_indices: Tensor) -> Tensor: + """ + Grouped load balancing loss: + L_LB = alpha * sum_{j=1}^{D+1} (f_j * P_j) + """ + if self.aux_loss_coeff <= 0: + return self.zeros((), dtype=router_probs.dtype) + + bs_slen = router_probs.shape[0] + if bs_slen == 0: + return self.zeros((), dtype=router_probs.dtype) + + # 1) f_j: frequency of selection per group + selected_flat = self.reshape(expert_indices, (-1,)) + expert_counts = self.histc( + selected_flat, + bins=self.num_total_experts, + min=0, + max=self.num_total_experts, + ) + expert_counts = self.cast(expert_counts, router_probs.dtype) + + experts_per_ffn_group = self.num_ffn_experts // self.num_ffn_groups + ffn_counts = expert_counts[: self.num_ffn_experts] + ffn_counts_reshaped = self.reshape( + ffn_counts, (self.num_ffn_groups, experts_per_ffn_group) + ) + group_counts_ffn = self.sum(ffn_counts_reshaped, dim=1) + + if self.num_copy_experts > 0: + copy_counts = expert_counts[self.num_ffn_experts :] + group_count_copy = self.sum(copy_counts).unsqueeze(0) + all_group_counts = self.cat((group_counts_ffn, group_count_copy)) + else: + all_group_counts = group_counts_ffn + + f_j = all_group_counts / (bs_slen * self.top_k) + + # 2) P_j: probability mass per group + expert_prob_sum = self.sum(router_probs, dim=0) + ffn_probs = expert_prob_sum[: self.num_ffn_experts] + ffn_probs_reshaped = self.reshape( + ffn_probs, (self.num_ffn_groups, experts_per_ffn_group) + ) + group_probs_ffn = self.sum(ffn_probs_reshaped, dim=1) + + if self.num_copy_experts > 0: + copy_probs = expert_prob_sum[self.num_ffn_experts :] + group_prob_copy = self.sum(copy_probs).unsqueeze(0) + all_group_probs = self.cat((group_probs_ffn, group_prob_copy)) + else: + all_group_probs = group_probs_ffn + + P_j = all_group_probs / bs_slen + + num_groups = self.num_ffn_groups + (1 if self.num_copy_experts > 0 else 0) + loss = self.aux_loss_coeff * num_groups * self.sum(f_j * P_j) + return loss + + def construct( + self, x: Tensor, training: bool = True + ) -> Tuple[Tensor, Tensor, Tensor, Tensor]: + """ + Returns: + top_scores: [bs*slen, K] + selected_indices: [bs*slen, K] + num_tokens_per_expert: [num_experts] + aux_loss: scalar + """ + bs_slen, _ = x.shape + router_dtype = self.config.moe_router_dtype + + # 1) Router logits + x_cast = self.cast(x, router_dtype) + weight = self.cast(self.weight, router_dtype) + logits = self.linear(x_cast, weight) + + # 2) Probabilities + if self.score_func == "sigmoid": + probs = self.sigmoid(self.cast(logits, mstype.float32)) + elif self.score_func == "softmax": + probs = self.softmax(self.cast(logits, mstype.float32), dim=1) + else: + raise NotImplementedError(f"Unknown score function {self.score_func}") + + # 3) Add bias for routing only + probs_for_routing = probs + self.expert_bias + + # 4) TopK selection + _, selected_indices = self.topk( + probs_for_routing, k=self.top_k, dim=-1, sorted=False + ) + selected_indices = self.cast(selected_indices, mstype.int64) + + # 5) Gather routing scores + top_scores = self.gather(probs, dim=1, index=selected_indices) + + # 6) Normalize and scale + if self.route_norm: + denominator = self.sum(top_scores, dim=-1, keepdim=True) + 1e-20 + top_scores = self.div(top_scores, denominator) + top_scores = self.mul(top_scores, self.route_scale) + + # 7) Statistics + num_tokens_per_expert = self.histc( + selected_indices, + bins=self.num_total_experts, + min=0, + max=self.num_total_experts, + ) + + # 8) Update bias (PID-like control) + if self.training and training: + self.update_expert_bias(num_tokens_per_expert, bs_slen) + + # 9) Load balance loss + aux_loss = self.compute_load_balancing_loss(probs, selected_indices) + + return top_scores, selected_indices, num_tokens_per_expert, aux_loss \ No newline at end of file