diff --git a/ds_pynative.yaml b/ds_pynative.yaml index 5affff04d88d50dfe9d78007e46a2f043c7eaf97..31b9c8551feb6dfa1704fc90f4da8a79006cc542 100644 --- a/ds_pynative.yaml +++ b/ds_pynative.yaml @@ -65,6 +65,19 @@ optimizer: eps: 1.e-8 weight_decay: 0.0 +# Muon optimizer configuration example (uncomment to use): +# optimizer: +# type: Muon +# learning_rate: 2.e-2 +# weight_decay: 0.1 +# matched_adamw_rms: 0.2 +# momentum: 0.95 +# nesterov: True +# ns_steps: 5 +# adamw_betas: [0.95, 0.95] +# adamw_eps: 1.e-8 +# qk_clip_threshold: 100 + # lr schedule lr_schedule: type: ConstantWarmUpLR @@ -179,6 +192,9 @@ model: v_head_dim: 192 qk_nope_head_dim: 128 qk_layernorm: True + # QK-clip scaling for Muon optimizer (tracks max attention logit per head) + # Enable this when using Muon optimizer with qk_clip_threshold + # track_max_attention_logit: True attention_dropout: 0.0 hidden_dropout: 0.0 params_dtype: "float32" diff --git a/mindformers/pynative/base_models/gpt/gpt_model.py b/mindformers/pynative/base_models/gpt/gpt_model.py index 95cefef4f502213737ed5d8483a6501145e528e6..4f515667dd0b0ac8495e5f39aaf87d942b82f2cf 100644 --- a/mindformers/pynative/base_models/gpt/gpt_model.py +++ b/mindformers/pynative/base_models/gpt/gpt_model.py @@ -33,6 +33,7 @@ from mindformers.pynative.base_models.common.embeddings.rotary_pos_embedding imp from mindformers.pynative.base_models.common.embeddings.yarn_rotary_pos_embedding import YarnRotaryEmbedding from mindformers.pynative.transformers.transformer_block import TransformerBlock, TransformerBlockSubmodules from mindformers.pynative.layers.linear import Linear +from mindformers.pynative.optimizer.muon_utils import make_muon_fns class GPTModel(nn.Cell): @@ -241,7 +242,7 @@ class GPTModel(nn.Cell): attention_mask (Tensor, optional): Attention mask tensor, used to mask padding tokens. Default is None. decoder_input (Tensor, optional): Decoder input tensor. Default is None. - labels (Tensor, optional): The label tensor, used for calculating the loss. + labels (Tensor, optional): The label tensor, used for calculating the loss. Default is None. loss_mask (Tensor, optional): Loss mask tensor, used to specify which positions are included in the loss calculation. Default is None. @@ -404,3 +405,286 @@ class GPTModel(nn.Cell): elif attention_mask is None: attention_mask = self.casual_mask(input_ids) return labels, attention_mask, loss_mask + + def get_gpt_transformer_config(self): + """Get the transformer config for GPT model. + + Returns: + TransformerConfig: The transformer configuration. + """ + return self.config + + def make_model_muon_fns(self): + """Read values from TransformersConfig and generate schema.""" + num_moe_experts = self.config.num_moe_experts + hidden_size = self.config.hidden_size + moe_ffn_hidden_size = self.config.moe_ffn_hidden_size + qk_head_dim = self.config.qk_head_dim + qk_pos_emb_head_dim = self.config.qk_pos_emb_head_dim + num_attention_heads = self.config.num_attention_heads + kv_lora_rank = self.config.kv_lora_rank + value_head_dim = self.config.v_head_dim + + # Calculate q_rank based on q_lora_rank + q_lora_rank = self.config.q_lora_rank + if q_lora_rank is None: + q_rank = num_attention_heads * (qk_head_dim + qk_pos_emb_head_dim) + else: + q_rank = q_lora_rank + + schema = [ + # experts.weight1: reshape → split into two [num_moe_experts, hidden_size, moe_ffn_hidden_size] + { + "patterns": ["*mlp.experts.weight1*"], + "kind": "reshape_concat", + "reshape": (num_moe_experts, hidden_size, 2 * moe_ffn_hidden_size), + }, + # experts.weight2: reshape → [num_moe_experts, moe_ffn_hidden_size, hidden_size] + { + "patterns": ["*mlp.experts.weight2*"], + "kind": "reshape_only", + "reshape": (num_moe_experts, moe_ffn_hidden_size, hidden_size), + }, + # linear_qkv: split into three parts along the first dimension + # Weight shape: [q_rank + kv_lora_rank + qk_pos_emb_head_dim, hidden_size] + # Split into: [q_part, kv_compressed_part, k_pe_part] + { + "patterns": ["*self_attention.linear_qkv.weight*"], + "kind": "multi_split", + "parts": (q_rank, kv_lora_rank, qk_pos_emb_head_dim), + }, + # q_up_proj (linear_qb in pynative): periodic split across heads + { + "patterns": ["*self_attention.linear_qb.weight*"], + "kind": "periodic", + "parts": (qk_head_dim, qk_pos_emb_head_dim, num_attention_heads), + }, + # kv_up_proj (linear_kvb in pynative): periodic split across heads + { + "patterns": ["*self_attention.linear_kvb.weight*"], + "kind": "periodic", + "parts": (qk_head_dim, value_head_dim, num_attention_heads), + }, + # fc1 and shared_fc1: alternating 1,1 split along rows + { + "patterns": [ + "*mlp.shared_experts.linear_fc1.weight*", + "*mlp.linear_fc1.weight*", + ], + "kind": "alt_pair_periodic", + }, + ] + + return make_muon_fns(schema) + + def get_muon_filter(self): + """Return a filter function to determine if a parameter should use Muon optimization. + + Returns: + A function that takes a parameter and returns True if it should use Muon. + """ + def muon_filter(param): + return ( + (len(param.shape) == 2 or len(param.shape) == 3) + and "word_embeddings" not in param.name + and "output_layer" not in param.name + ) + return muon_filter + + def get_tp_dims(self, params): + """Return tensor parallel dimensions for each parameter. + + Args: + params: List of parameters from the optimizer. + + Returns: + Tuple of TP dimensions for each parameter. + """ + no_tp_list = [ + "linear_q_down_proj", + "linear_kv_down_proj", + "shared_experts", + "mlp.router", + "hnorm.weight", "enorm.weight", "eh_proj.weight", + ] + + tp_dim_1_list = [ + "self_attention.linear_proj.weight", + "mlp.linear_fc2.weight" + ] + + def name_filter(param_name, full_name_list): + for full_name in full_name_list: + if full_name in param_name: + return True + return False + + tp_dims = [] + for param in params: + if name_filter(param.name, tp_dim_1_list): + tp_dims.append(1) + elif name_filter(param.name, no_tp_list): + tp_dims.append(-1) + else: + tp_dims.append(0) + return tuple(tp_dims) + + def get_op_groups_info(self, params, op): + """Return optimizer parallel group information for each parameter. + + Args: + params: List of parameters from the optimizer. + op: Optimizer parallel size. + + Returns: + Tuple of (ops, op_groups) where: + - ops: tuple of op values for each parameter + - op_groups: tuple of group names for each parameter + """ + no_op_list = [ + "self_attention.linear_q_proj.weight", + "self_attention.linear_q_up_proj.weight", + "self_attention.linear_q_down_proj.weight", + "self_attention.linear_kv_up_proj.weight", + "self_attention.linear_kv_down_proj.weight", + "eh_proj", + "max_logits_val" + ] + + def name_filter(param_name, full_name_list): + for full_name in full_name_list: + if full_name in param_name: + return True + return False + + op_list = [] + op_groups = [] + + for param in params: + if name_filter(param.name, no_op_list): + op_list.append(1) + op_groups.append("") + else: + # For pynative mode, use simplified logic + # In distributed training, this would need proper group computation + op_list.append(op if op > 1 else 1) + op_groups.append("") + + return tuple(op_list), tuple(op_groups) + + def get_param_layer_indices(self, params): + """Return layer indices for each parameter (used for QK-clip). + + Args: + params: List of parameters from the optimizer. + + Returns: + Tuple of layer indices for each parameter, where: + - layer_idx >= 0 stands for the layer_idx-th decoder layer + - layer_idx < 0 stands for the -(layer_idx+1)-th MTP layer + """ + param_layer = [] + for param in params: + name = param.name + try: + layer_idx = int(name.split(".")[2]) + except (ValueError, IndexError): + layer_idx = 0 + if name.startswith('mtp'): + layer_idx = -layer_idx - 1 + param_layer.append(layer_idx) + return tuple(param_layer) + + def apply_qk_clip_scaling(self, params, param_names, param_layer, logit_threshold, + muon_split_fn, muon_merge_fn): + """Apply QK-clip scaling to attention weight parameters. + + Args: + params: List of all parameters. + param_names: Tuple of parameter names. + param_layer: Tuple of layer indices for each parameter. + logit_threshold: Threshold for logit clipping. + muon_split_fn: Function to split parameters. + muon_merge_fn: Function to merge parameters. + + Returns: + List of (param_idx, scaled_weights) tuples to be updated. + """ + if not self.config.multi_latent_attention: + return [] + + # Check if track_max_attention_logit is enabled + if not getattr(self.config, 'track_max_attention_logit', False): + return [] + + ones = mint.ones((1,), dtype=dtype.float32) + qk_head_dim = self.config.qk_head_dim + qk_pos_emb_head_dim = self.config.qk_pos_emb_head_dim + + def get_scale_broadcast(scales, head_dim): + scale_broadcast = mint.tile( + mint.unsqueeze(scales, 1), (1, head_dim) + ).reshape(-1) + scale_broadcast = mint.unsqueeze(scale_broadcast, 1) + return scale_broadcast + + updates = [] + for idx, param_name in enumerate(param_names): + # In pynative mode, linear_qb corresponds to linear_q_up_proj, + # and linear_kvb corresponds to linear_kv_up_proj + if ( + "self_attention.linear_qb.weight" not in param_name + and "self_attention.linear_kvb.weight" not in param_name + ): + continue + + layer_idx = param_layer[idx] + param = params[idx] + + # Compute per-head scale factor + logit_threshold_f32 = self.cast(logit_threshold, dtype.float32) + if layer_idx >= 0: + logits_row = ( + self.decoder.layers[layer_idx] + .self_attention + .core_attention + .max_logits_val + .value() + ) + else: + # MTP layer + if hasattr(self, 'mtp') and self.mtp is not None: + logits_row = ( + self.mtp.layers[-(layer_idx + 1)] + .transformer_layer + .self_attention + .core_attention + .max_logits_val + .value() + ) + else: + continue + + logits_row = logits_row.reshape(-1) + mask = mint.greater_equal(logits_row, logit_threshold_f32) + safe_den = mint.where(mask, logits_row, ones) + scales = mint.where(mask, logit_threshold_f32 / safe_den, ones) + + weights = None + # In pynative mode, linear_qb corresponds to linear_q_up_proj + if "self_attention.linear_qb.weight" in param_name: + l2q_nope_proj, l2q_pe_proj = muon_split_fn(param_name, param) + l2q_nope_proj = l2q_nope_proj * get_scale_broadcast(mint.sqrt(scales), qk_head_dim) + l2q_pe_proj = l2q_pe_proj * get_scale_broadcast(scales, qk_pos_emb_head_dim) + weights = muon_merge_fn(param_name, [l2q_nope_proj, l2q_pe_proj]) + # In pynative mode, linear_kvb corresponds to linear_kv_up_proj + elif "self_attention.linear_kvb.weight" in param_name: + lkv2kv_k_nope, lkv2kv_v = muon_split_fn(param_name, param) + lkv2kv_k_nope = lkv2kv_k_nope * get_scale_broadcast(mint.sqrt(scales), qk_head_dim) + weights = muon_merge_fn(param_name, [lkv2kv_k_nope, lkv2kv_v]) + + if weights is not None: + updates.append((idx, weights)) + + return updates + diff --git a/mindformers/pynative/layers/flash_attention.py b/mindformers/pynative/layers/flash_attention.py index e3b4b500410ac72d6e441a64b8be99a4d15ba241..2af0028d928040d3b8fb356069543d04ab1bd353 100644 --- a/mindformers/pynative/layers/flash_attention.py +++ b/mindformers/pynative/layers/flash_attention.py @@ -18,9 +18,10 @@ __all__ = ['FlashAttention'] import math from typing import Union +import numpy as np import mindspore.common.dtype as mstype import mindspore as ms -from mindspore import ops, mint +from mindspore import ops, mint, Parameter from mindspore.common.tensor import Tensor from mindspore.nn.cell import Cell @@ -114,6 +115,22 @@ class FlashAttention(Cell): self.fa_out_transpose = mint.permute self.cast = ops.cast + # Track max attention logit for QK-clip scaling (used by Muon optimizer) + self.track_max_attention_logit = getattr(config, 'track_max_attention_logit', False) + + if self.track_max_attention_logit: + # Parameter to store the maximum attention logit value per head. + # Note: This is a local max within each device's partition. Cross-device + # synchronization (AllReduce-Max across DP/CP dimensions) is performed + # later in GPTModel.allreduce_max_attention_logit() to obtain the global max. + self.max_logits_val = Parameter( + Tensor(np.zeros((self.head_num,)), dtype=mstype.float32), + requires_grad=False + ) + self.matmul_qk = mint.matmul + self.reduce_max = mint.amax + self.maximum = mint.maximum + def construct(self, query: Tensor, key: Tensor, @@ -128,6 +145,10 @@ class FlashAttention(Cell): if attention_mask is not None: attention_mask = self.cast(attention_mask, ms.uint8) + # Track max attention logit if enabled + if self.track_max_attention_logit and self.training: + self._update_max_attention_logit(query, key) + if self.input_layout == "TND": output = self.flash_attention(query=query, key=key, @@ -200,3 +221,62 @@ class FlashAttention(Cell): x_merge = self.reshape(x, new_shape) x_merge = self.fa_out_transpose(x_merge, (1, 0, 2)) return x_merge + + def _update_max_attention_logit(self, query, key): + """ + Compute and update the maximum attention logit value per head. + + This method computes attention scores (Q @ K^T * scale) and tracks + the maximum value per attention head for QK-clip scaling. + + Args: + query: Query tensor with shape depending on input_layout. + key: Key tensor with shape depending on input_layout. + """ + # Convert to BNSD format for unified processing + if self.input_layout == "TND": + # TND: (T, N, D) where T = batch * seq, N = num_heads, D = head_dim + # We need to handle this carefully + t, n, d = query.shape + # For TND, we compute a simplified max logit estimation + # by sampling a subset of the sequence to reduce computation + sample_size = min(64, t) + q_sample = query[:sample_size] # (sample, N, D) + k_sample = key[:sample_size] # (sample, N, D) + # Compute attention scores: (sample, N, D) @ (sample, N, D)^T -> need per-head + # Reshape to (N, sample, D) for batch matmul + q_sample = mint.permute(q_sample, (1, 0, 2)) # (N, sample, D) + k_sample = mint.permute(k_sample, (1, 0, 2)) # (N, sample, D) + # (N, sample, D) @ (N, D, sample) -> (N, sample, sample) + scores = self.matmul_qk(q_sample, mint.permute(k_sample, (0, 2, 1))) + scores = scores * self.scalar_value + elif self.input_layout == "BNSD": + # BNSD: (B, N, S, D) + # Sample to reduce computation + seq_len = query.shape[2] + sample_size = min(64, seq_len) + q_sample = query[:, :, :sample_size, :] # (B, N, sample, D) + k_sample = key[:, :, :sample_size, :] # (B, N, sample, D) + # (B, N, sample, D) @ (B, N, D, sample) -> (B, N, sample, sample) + scores = self.matmul_qk(q_sample, mint.permute(k_sample, (0, 1, 3, 2))) + scores = scores * self.scalar_value + # Reduce over batch: max over (B, sample, sample) -> (N,) + scores = self.reduce_max(scores, dim=(0, 2, 3)) + else: + # BSH or SBH: (B, S, H) or (S, B, H) where H = N * D + # For simplicity, skip tracking for these layouts + return + + # Compute max per head + if self.input_layout == "TND": + # scores shape: (N, sample, sample) + max_logits = self.reduce_max(scores, dim=(1, 2)) # (N,) + else: + max_logits = scores # Already reduced to (N,) + + # Cast to float32 for stable comparison + max_logits = self.cast(max_logits, mstype.float32) + + # Update max_logits_val with element-wise maximum + new_max = self.maximum(self.max_logits_val, max_logits) + ops.assign(self.max_logits_val, new_max) diff --git a/mindformers/pynative/optimizer/__init__.py b/mindformers/pynative/optimizer/__init__.py index a125dde4986c6165f6b5783e79a7d0f2402d8aed..0446c1e22bf604462d5077e1c09abc9a5e7fb314 100644 --- a/mindformers/pynative/optimizer/__init__.py +++ b/mindformers/pynative/optimizer/__init__.py @@ -13,3 +13,6 @@ # limitations under the License. # ============================================================================ """optimizer modules""" +from .muon import Muon + +__all__ = ['Muon'] diff --git a/mindformers/pynative/optimizer/muon.py b/mindformers/pynative/optimizer/muon.py new file mode 100755 index 0000000000000000000000000000000000000000..f9a79136248b858132176d949f19e89ab1047253 --- /dev/null +++ b/mindformers/pynative/optimizer/muon.py @@ -0,0 +1,528 @@ +# Copyright 2025 Huawei Technologies Co., Ltd +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================ +"""Muon optimizer for pynative mode. + +This module provides the Muon optimizer implementation for pynative mode training. +Unlike the static graph version, this version: +- Does not use @jit decorator +- Uses Python for loops instead of hyper_map + MultitypeFuncGraph +- Uses direct reshape operations instead of Morph operator +""" + +from __future__ import absolute_import + +import hashlib + +import numpy as np +from mindspore.common import dtype as mstype +from mindspore.ops import functional as F, operations as P +from mindspore.nn.optim.optimizer import Optimizer +from mindspore.common.tensor import Tensor +from mindspore.common.parameter import Parameter, ParameterTuple +from mindspore.communication.management import create_group, get_rank +from mindspore.ops.auto_generate import Chunk +from mindspore import get_auto_parallel_context + +from mindformers.tools.register import MindFormerRegister, MindFormerModuleType +from mindformers.core.context import is_legacy_model +from mindformers.tools.logger import logger + + +def _perform_allgather_op(ns_inputs_item, op, tp, tp_dim, op_group, tp_group, param_name): + """Perform AllGather operations based on op and tp settings.""" + if "mlp.experts.weight" not in param_name: + # all gather op_shard + if op > 1: + ns_inputs_item = P.AllGather(group=op_group)(ns_inputs_item) + + # all gather tp_shard + if tp > 1: + if tp_dim == 0: + ns_inputs_item = P.AllGather(group=tp_group)(ns_inputs_item) + elif tp_dim == 1: + ns_inputs_item = P.AllGather(group=tp_group)(ns_inputs_item.T) + ns_inputs_item = ns_inputs_item.T + return ns_inputs_item + + +def zeropower_via_newtonschulz5_2d(x, dim_a, dim_b): + """Apply Newton-Schulz iteration for 2D tensors.""" + a, b, c = (3.4445, -4.7750, 2.0315) + + if dim_a > dim_b: + x = x.T + # Ensure spectral norm is at most 1 + x = x / (x.norm() + 1e-7) + # Perform the NS iterations + for _ in range(5): + a_mat = x @ x.T + b_mat = b * a_mat + c * a_mat @ a_mat + x = a * x + b_mat @ x + if dim_a > dim_b: + x = x.T + return x + + +def zeropower_via_newtonschulz5_3d(x, dim_a, dim_b): + """Apply Newton-Schulz iteration for 3D tensors.""" + a, b, c = (3.4445, -4.7750, 2.0315) + + if dim_a > dim_b: + x = P.Transpose()(x, (0, 2, 1)) + # Ensure spectral norm is at most 1 + x = x / P.ExpandDims()(P.ExpandDims()((x.norm(dim=(1, 2)) + 1e-7), 1), 1) + # Perform the NS iterations + for _ in range(5): + a_mat = P.BatchMatMul(transpose_b=True)(x, x) + b_mat = b * a_mat + c * P.BatchMatMul()(a_mat, a_mat) + x = a * x + P.BatchMatMul()(b_mat, x) + if dim_a > dim_b: + x = P.Transpose()(x, (0, 2, 1)) + return x + + +def _slice_tensor_to_shards(x, tp, tp_dim, op, rank_id, op_group, tp_group): + """Slice tensor to tp_shard and op_shard.""" + # slice X to tp_shard and slice X to op_shard + if tp > 1: + if tp_dim >= 0: + chunk_id = rank_id % tp + x = Chunk()(x, tp, tp_dim)[chunk_id] + + if op > 1: + if tp_dim == -1: + chunk_id = rank_id % op + else: + chunk_id = rank_id // tp % op + x = Chunk()(x, op)[chunk_id] + return x + + +def _apply_muon_update( + gradient, muon_m, momentum, use_nesterov, param, lr, weight_decay, + matched_adamw_rms, muon_split_fn, muon_merge_fn, param_name, + op, tp, tp_dim, rank_id, op_group, tp_group): + """Apply Muon optimizer update.""" + op_sqrt = P.Sqrt() + op_cast = P.Cast() + op_reshape = P.Reshape() + op_shape = P.Shape() + + m_fp32 = op_cast(muon_m, mstype.float32) + gradient_fp32 = op_cast(gradient, mstype.float32) + next_m = m_fp32 * momentum + gradient_fp32 + + if use_nesterov: + gradient_fp32 = gradient_fp32 + next_m * momentum + else: + gradient_fp32 = next_m + + ns_inputs = op_cast(gradient_fp32, mstype.bfloat16) + ns_inputs_list = muon_split_fn(param_name, ns_inputs) + x_list = [] + + dim_a, dim_b = None, None + for ns_inputs_item in ns_inputs_list: + dim_a, dim_b = op_shape(ns_inputs_item)[-2:] + + if len(op_shape(ns_inputs_item)) == 2: + ns_inputs_item = _perform_allgather_op( + ns_inputs_item, op, tp, tp_dim, op_group, tp_group, param_name) + x = zeropower_via_newtonschulz5_2d(ns_inputs_item, dim_a, dim_b) + x = _slice_tensor_to_shards(x, tp, tp_dim, op, rank_id, op_group, tp_group) + else: + x = zeropower_via_newtonschulz5_3d(ns_inputs_item, dim_a, dim_b) + + x_list.append(x) + + x_ret = muon_merge_fn(param_name, x_list) + param_fp32 = op_cast(param, mstype.float32) + param_fp32 = param_fp32 * (1 - lr * weight_decay) + + adjusted_ratio = op_sqrt(op_cast(max(dim_a, dim_b), mstype.float32)) * matched_adamw_rms + adjusted_lr = lr * adjusted_ratio + update_with_lr = adjusted_lr * x_ret + next_param = param_fp32 - op_reshape(update_with_lr, op_shape(param_fp32)) + next_param = F.depend(next_param, F.assign(param, op_cast(next_param, F.dtype(param)))) + next_param = F.depend(next_param, F.assign(muon_m, op_cast(next_m, F.dtype(muon_m)))) + return op_cast(next_param, F.dtype(param)) + + +def _apply_adamw_update(param, exp_avg, exp_avg_sq, gradient, beta1, beta2, step, eps, lr, weight_decay): + """Apply AdamW optimizer update.""" + op_mul = P.Mul() + op_pow = P.Pow() + op_sqrt = P.Sqrt() + op_cast = P.Cast() + addcmul = P.Addcmul() + + param_fp32 = op_cast(param, mstype.float32) + next_param = op_mul(param_fp32, 1 - lr * weight_decay) + gradient_fp32 = op_cast(gradient, mstype.float32) + + next_param = F.depend( + next_param, + F.assign( + exp_avg, + op_mul(exp_avg, beta1) + + op_mul(gradient_fp32, op_cast(F.tuple_to_array((1.0,)), mstype.float32) - beta1), + ), + ) + next_param = F.depend( + next_param, + F.assign( + exp_avg_sq, + addcmul( + op_mul(exp_avg_sq, beta2), + gradient_fp32, + gradient_fp32, + op_cast(F.tuple_to_array((1.0,)), mstype.float32) - beta2, + ), + ), + ) + + bias_correction1 = 1 - op_pow(op_cast(beta1, mstype.float32), step) + bias_correction2 = 1 - op_pow(op_cast(beta2, mstype.float32), step) + step_size = lr / bias_correction1 + denom = op_sqrt(exp_avg_sq / bias_correction2) + eps + return_param = next_param - op_mul(exp_avg / denom, step_size) + F.assign(param, op_cast(return_param, F.dtype(param))) + return op_cast(return_param, F.dtype(param)) + + +@MindFormerRegister.register(MindFormerModuleType.OPTIMIZER) +class Muon(Optimizer): + """ + Muon optimizer implementation for pynative mode. + + Args: + params: model parameters to optimize. + learning_rate (float): Learning rate. Default: ``2e-2``. + weight_decay (float): Weight decay factor. Default: ``0.1``. + matched_adamw_rms (float): RMS matching parameter for AdamW. Default: ``0.2``. + momentum (float): Momentum factor. Default: ``0.95``. + nesterov (bool): Whether to use Nesterov momentum. Default: ``True``. + ns_steps (int): Number of Newton-Schulz steps. Default: ``5``. + adamw_betas (tuple): Beta parameters for AdamW. Default: ``(0.95, 0.95)``. + adamw_eps (float): Epsilon for AdamW. Default: ``1e-8``. + qk_clip_threshold (float): QK clip threshold. Default: ``100``. + model: The model model. Default: ``None``. + """ + + def __init__( + self, + params, + learning_rate=2e-2, + weight_decay=0.1, + matched_adamw_rms=0.2, + momentum=0.95, + nesterov=True, + ns_steps=5, + adamw_betas=(0.95, 0.95), + adamw_eps=1e-8, + qk_clip_threshold=100, + model=None, + **kwargs, + ): + super().__init__(learning_rate, params, weight_decay) + if kwargs.get('swap', False): + raise ValueError("Muon does not support swap.") + + self._verify_model(model) + + # Initialize basic parameters + self._initialize_basic_params(adamw_betas, adamw_eps, momentum, matched_adamw_rms, nesterov) + + # Initialize model configuration + self._initialize_network_config(model) + + # Initialize parameter layers + self._initialize_param_layers(model) + + # Initialize QK-clip parameters + self.ones = Tensor([1.0], mstype.float32) + self.rank_id = get_rank() + self.rank_ids = tuple(self.rank_id for _ in self._parameters) + self.logit_threshold = Tensor([qk_clip_threshold], dtype=mstype.float32) + + # Initialize Muon momentum + self._initialize_muon_moments(model) + + # Initialize tensor parallel dimensions + self._initialize_tp_dims(model) + + # Initialize AdamW moments + self._initialize_adamw_moments(model) + + # Initialize parallel configuration + self._initialize_parallel_config(model) + + # Initialize communication groups + self._initialize_communication_groups() + + # Initialize optimizer parallel groups + self._initialize_op_groups(model) + + # Store model for QK-clip + self.model = model + self.ns_steps = ns_steps + + def _verify_model(self, model): + """Verify if the model is compatible with Muon optimizer.""" + if model is None: + raise ValueError("Model must be provided for Muon optimizer.") + + if is_legacy_model(): + raise ValueError("Muon does not support Legacy Model.") + + config = model.get_gpt_transformer_config() + + if not config.multi_latent_attention: + raise ValueError("Current Muon implementation only supports models with Multi-Latent Attention enabled.") + + def _initialize_basic_params(self, adamw_betas, adamw_eps, momentum, matched_adamw_rms, nesterov): + """Initialize basic optimizer parameters.""" + self.beta1 = Tensor(np.array([adamw_betas[0]]).astype(np.float32)) + self.beta2 = Tensor(np.array([adamw_betas[1]]).astype(np.float32)) + self.eps = Tensor(np.array([adamw_eps]).astype(np.float32)) + self.muon_momentum = Tensor(np.array([momentum]).astype(np.float32)) + self.matched_adamw_rms = Tensor(np.array([matched_adamw_rms]).astype(np.float32)) + self.use_nesterov = tuple(nesterov for _ in self._parameters) + self.param_name_tuple = tuple(p.name for p in self._parameters) + + def _initialize_network_config(self, model): + """Initialize Model configuration and split/merge functions.""" + self.muon_split_fn, self.muon_merge_fn = model.make_model_muon_fns() + self.muon_split_fns = tuple(self.muon_split_fn for _ in self._parameters) + self.muon_merge_fns = tuple(self.muon_merge_fn for _ in self._parameters) + + def _initialize_param_layers(self, model): + """Initialize parameter layer indices.""" + self.param_layer = model.get_param_layer_indices(self._parameters) + + def _initialize_muon_moments(self, model): + """Initialize Muon momentum parameters.""" + muon_filter = model.get_muon_filter() + + self.muon_m = [] + self.param_idx_in_opt = {} + for idx, param in enumerate(self._parameters): + self.param_idx_in_opt[param.name] = idx + + for param in self._parameters: + if muon_filter(param): + x1 = param.clone("zeros") + x1.name = "muon_m" + "." + x1.name + self.muon_m.append(x1) + logger.info(f"Muon apply: {param}") + else: + self.muon_m.append(Parameter(Tensor(np.array([0]).astype(np.float32)), name="muon_m." + param.name)) + self.muon_m = ParameterTuple(self.muon_m) + self.use_muon = tuple(muon_filter(param) for param in self._parameters) + + def _initialize_tp_dims(self, model): + """Initialize tensor parallel dimensions.""" + self.tp_dims = model.get_tp_dims(self._parameters) + + def _initialize_adamw_moments(self, model): + """Initialize AdamW momentum parameters.""" + muon_filter = model.get_muon_filter() + + self.moments1 = [] + self.moments2 = [] + for param in self._parameters: + if not muon_filter(param): + x1 = param.clone("zeros") + x1.name = "adam_m" + "." + x1.name + self.moments1.append(x1) + x2 = param.clone("zeros") + x2.name = "adam_v" + "." + x2.name + self.moments2.append(x2) + logger.info(f"Adam apply: {param}") + else: + self.moments1.append(Parameter(Tensor(np.array([0]).astype(np.float32)), name="adam_m." + param.name)) + self.moments2.append(Parameter(Tensor(np.array([0]).astype(np.float32)), name="adam_v." + param.name)) + self.moments1 = ParameterTuple(self.moments1) + self.moments2 = ParameterTuple(self.moments2) + + def _initialize_parallel_config(self, model): + """Initialize parallel configuration.""" + self.tp = model.get_gpt_transformer_config().tensor_model_parallel_size + self.tps = tuple(self.tp for _ in self._parameters) + self.dp = model.get_gpt_transformer_config().data_parallel_size + logger.info(f"Muon tp group size is: {self.tp}") + + if not get_auto_parallel_context('enable_parallel_optimizer'): + self.op = 1 + else: + self.op = get_auto_parallel_context('optimizer_weight_shard_size') + if self.op < 1: + raise ValueError( + "Must set parallel.parallel_optimizer_config.optimizer_weight_shard_size > 1 " + "when enable_parallel_optimizer is True.") + if self.dp < self.op: + raise ValueError('Must set parallel_config.data_parallel >= ' + 'parallel.parallel_optimizer_config.optimizer_weight_shard_size when using Muon.') + logger.info(f"Muon op group size is: {self.op}") + + def _initialize_communication_groups(self): + """Initialize communication groups for parallel training.""" + # Only create groups if parallel training is needed + if self.tp > 1 or self.op > 1: + self.tp_group = self._get_tp_group_name(self.rank_id, self.tp) + self.op_group, self.op_in_tp_group = self._get_op_group_name(self.rank_id, self.tp, self.op, self.tp_group) + self.tp_groups = tuple(self.tp_group for _ in self._parameters) + else: + self.tp_group = None + self.op_group = None + self.op_in_tp_group = None + self.tp_groups = tuple(None for _ in self._parameters) + + def _initialize_op_groups(self, model): + """Initialize optimizer parallel groups for parameters.""" + self.ops, self.op_groups = model.get_op_groups_info(self._parameters, self.op) + + def _create_communication_group(self, rank_list): + """ + Create a communication group with a hashed name. + + Args: + rank_list: List of ranks in the communication group + + Returns: + str: The created group name + """ + rank_list_str = "-".join([str(i) for i in rank_list]) + hashed = hashlib.md5(rank_list_str.encode()).hexdigest()[:48] + group_name = str(hashed) + create_group(group_name, rank_list) + return group_name + + def _get_op_group_name(self, rank_id, tp, op, tp_group): + """ + Generates a unique group name for optimizer parallel communication group. + + Returns: + tuple: The optimizer group name and optimizer-in-tensor-parallel group name + """ + dp_range = tp + op_range = tp * op + rank_start = rank_id % dp_range + rank_id // op_range * op_range + rank_end = rank_start + op_range + rank_list = list(range(rank_start, rank_end, dp_range)) + logger.info(f"Muon op group list is: {rank_list}") + op_group_name = self._create_communication_group(rank_list) + + if tp == op: + logger.info( + f"op_in_tp group will reuse tp group" + f", since tensor_parallel_size({tp}) == optimizer_parallel_size({op})." + ) + op_in_tp_group_name = tp_group + else: + logger.info(f"Muon op_in_tp group list is: {rank_list}") + op_in_tp_group_name = self._get_tp_group_name(rank_id, op) + + return op_group_name, op_in_tp_group_name + + def _get_tp_group_name(self, rank_id, tp): + """ + Generates a unique group name for tensor parallel communication group. + + Returns: + str: The tensor parallel group name + """ + rank_start = rank_id // tp * tp + rank_end = rank_id // tp * tp + tp + rank_list = list(range(rank_start, rank_end)) + logger.info(f"Muon tp group list is: {rank_list}") + tp_group_name = self._create_communication_group(rank_list) + return tp_group_name + + def construct(self, gradients): + """Construct method for optimizer. + + Args: + gradients: Gradients for optimization. + + Returns: + Updated gradients after optimization. + """ + gradients = self.flatten_gradients(gradients) + weight_decay = self.get_weight_decay() + lr = self.get_lr() + self.assignadd(self.global_step, self.global_step_increase_tensor) + + # Get current step for AdamW bias correction + step = self.global_step + + # Process each parameter using for loop (pynative mode) + optim_result = [] + for i, (param, gradient) in enumerate(zip(self._parameters, gradients)): + param_name = self.param_name_tuple[i] + + # Skip max_logits_val parameters + if "max_logits_val" in param_name: + optim_result.append(P.Cast()(gradient, F.dtype(param))) + continue + + # Skip if not in optimization filter + if not self.optim_filter[i]: + optim_result.append(gradient) + continue + + # Get learning rate and weight decay for this parameter + if self.is_group: + if self.is_group_lr: + param_lr = lr[i] + param_wd = weight_decay[i] + else: + param_lr = lr + param_wd = weight_decay[i] + else: + param_lr = lr + param_wd = weight_decay + + if self.use_muon[i]: + # Apply Muon update + result = _apply_muon_update( + gradient, self.muon_m[i], self.muon_momentum, + self.use_nesterov[i], param, param_lr, param_wd, + self.matched_adamw_rms, self.muon_split_fn, self.muon_merge_fn, + param_name, self.ops[i], self.tps[i], self.tp_dims[i], + self.rank_id, self.op_groups[i], self.tp_groups[i]) + else: + # Apply AdamW update + result = _apply_adamw_update( + param, self.moments1[i], self.moments2[i], gradient, + self.beta1, self.beta2, step, self.eps, param_lr, param_wd) + + optim_result.append(result) + + # Apply QK-clip scaling + updates = self.model.apply_qk_clip_scaling( + self._parameters, + self.param_name_tuple, + self.param_layer, + self.logit_threshold, + self.muon_split_fn, + self.muon_merge_fn, + ) + + # Apply the weight updates + for param_idx, weights in updates: + optim_result = F.depend(optim_result, F.assign(self._parameters[param_idx], weights)) + + return optim_result diff --git a/mindformers/pynative/optimizer/muon_utils.py b/mindformers/pynative/optimizer/muon_utils.py new file mode 100755 index 0000000000000000000000000000000000000000..a5f7ac43cccfe4636743ba2d8c406fc708254d26 --- /dev/null +++ b/mindformers/pynative/optimizer/muon_utils.py @@ -0,0 +1,266 @@ +# Copyright 2025 Huawei Technologies Co., Ltd +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================ +"""Muon utils for pynative mode. + +This module provides utility functions for the Muon optimizer in pynative mode. +Unlike the static graph version which uses Morph operator for dynamic shape handling, +this version uses direct ops.reshape since shapes are known at runtime in pynative mode. +""" + +import math +from fnmatch import fnmatch +from mindspore import ops + + +def block_split_reshape(tensor, block): + """ + Reshape tensor by splitting its last dimension into blocks. + + This operation takes a tensor and splits its last dimension into equal-sized blocks, + adding a new dimension for the block index. + + Args: + tensor: Input tensor. + block: Block size for splitting the last dimension. + + Returns: + Reshaped tensor with last dimension split into blocks. + """ + shape = tensor.shape + *prefix, dim = shape + new_shape = tuple(prefix) + (dim // block, block) + return ops.reshape(tensor, new_shape) + + +def tensor_reshape_to_3d(tensor, dim1, dim2): + """ + Reshape tensor to 3D with specified middle and last dimensions. + + This operation reshapes a tensor to a 3-dimensional tensor where the first dimension + is automatically calculated from the total size, and the last two dimensions are fixed. + + Args: + tensor: Input tensor. + dim1: The second dimension (middle dimension) of the output 3D tensor. + dim2: The third dimension (last dimension) of the output 3D tensor. + + Returns: + Reshaped 3D tensor. + """ + total = math.prod(tensor.shape) + new_shape = (total // (dim1 * dim2), dim1, dim2) + return ops.reshape(tensor, new_shape) + + +def prefix_dimension_reshape(tensor, *prefix): + """ + Reshape tensor with fixed prefix dimensions and calculated last dimension. + + This operation reshapes a tensor by specifying the leading (prefix) dimensions, + while the last dimension is automatically calculated from the total size. + + Args: + tensor: Input tensor. + *prefix: Variable number of prefix dimensions for the output tensor shape. + + Returns: + Reshaped tensor with specified prefix dimensions. + """ + total = math.prod(tensor.shape) + prefix_total = math.prod(prefix) + new_shape = tuple(prefix) + (total // prefix_total,) + return ops.reshape(tensor, new_shape) + + +def tensor_reshape_to_2d(tensor, dim): + """ + Reshape tensor to 2D with specified last dimension. + + This operation flattens a tensor to a 2-dimensional tensor where the last dimension + is fixed and the first dimension is automatically calculated from the total size. + + Args: + tensor: Input tensor. + dim: The second dimension (last dimension) of the output 2D tensor. + + Returns: + Reshaped 2D tensor. + """ + total = math.prod(tensor.shape) + new_shape = (total // dim, dim) + return ops.reshape(tensor, new_shape) + + +def muon_split(tensor, part_a: int, part_b: int, num_blocks: int): + """ + Split a 2D tensor into two periodic parts along its first dimension. + The split pattern repeats every (part_a + part_b) elements. + + Args: + tensor: Input tensor of shape (M, N). + part_a: Number of elements in the first part of each block. + part_b: Number of elements in the second part of each block. + num_blocks: Total number of (part_a + part_b) blocks. + + Returns: + A tuple of two tensors (first_part, second_part), + where: + - first_part contains all part_a segments of each block. + - second_part contains all part_b segments of each block. + """ + tensor = tensor.T + *prefix, _ = tensor.shape + block = part_a + part_b + t = block_split_reshape(tensor, block) + + first_part = prefix_dimension_reshape(t[..., :part_a], *prefix).T + second_part = prefix_dimension_reshape(t[..., part_a:], *prefix).T + return first_part, second_part + + +def muon_merge(tensor_a, tensor_b, part_a: int, part_b: int, num_blocks: int): + """ + Merge two tensors back into the original periodic layout + that was split by muon_split(). + + Args: + tensor_a: Tensor containing the first part of each block. + tensor_b: Tensor containing the second part of each block. + part_a: Number of elements in the first part of each block. + part_b: Number of elements in the second part of each block. + num_blocks: Total number of (part_a + part_b) blocks. + + Returns: + A single tensor of the same shape as before muon_split(). + """ + tensor_a = tensor_a.T + tensor_b = tensor_b.T + *prefix, _ = tensor_a.shape + + a = block_split_reshape(tensor_a, part_a) + b = block_split_reshape(tensor_b, part_b) + t = ops.Concat(axis=-1)([a, b]) + out = prefix_dimension_reshape(t, *prefix).T + return out + + +def _eval_tuple(spec, name, tensor): + """Evaluate spec if callable, otherwise return as-is.""" + return spec(name, tensor) if callable(spec) else spec + + +def make_muon_fns(schema): + """ + Generate two generic functions: + - split_one(param_name, tensor) -> List[tensor] + - merge_one(param_name, parts_list) -> tensor + + Dimensions in schema should be either numbers or callback functions: + - periodic: rule["parts"] = (a, b, num_blocks) or lambda(name, tensor)->(a,b,blocks) + - reshape_* : rule["reshape"] = (x, y, z) or lambda(name, tensor)->(x,y,z) + + Args: + schema: List of rules defining how to split/merge tensors. + + Returns: + Tuple of (split_fn, merge_fn) functions. + """ + + def split_fn(param_name, tensor): + """ + Input a 2D tensor, split it according to schema rules, and return several segments (List[tensor]). + """ + for rule in schema: + if not any(fnmatch(param_name, pat) for pat in rule["patterns"]): + continue + + kind = rule["kind"] + + if kind == "periodic": + part_a, part_b, num_blocks = _eval_tuple(rule["parts"], param_name, tensor) + first_part, second_part = muon_split(tensor, part_a, part_b, num_blocks) + return [first_part, second_part] + + if kind == "reshape_concat": + # e.g. experts.weight1: first reshape to [E, H, 2I], then split into two halves + _, hidden_size, total_intermediate = _eval_tuple(rule["reshape"], param_name, tensor) + half_intermediate = total_intermediate // 2 + t3 = tensor_reshape_to_3d(tensor, hidden_size, total_intermediate) + return [t3[..., :half_intermediate], t3[..., half_intermediate:]] + + if kind == "reshape_only": + # e.g. experts.weight2: just reshape to [E, I, H], no split + _, intermediate_size, hidden_size = _eval_tuple(rule["reshape"], param_name, tensor) + return [tensor_reshape_to_3d(tensor, intermediate_size, hidden_size)] + + if kind == "alt_pair_periodic": + # Alternating rows 1,1 (blocks = M//2) + num_blocks = tensor.shape[0] // 2 + a, b = muon_split(tensor, 1, 1, num_blocks) + return [a, b] + + if kind == "multi_split": + # Split tensor into multiple parts along the first dimension + # parts: list of sizes for each segment, e.g., [size1, size2, size3] + parts = _eval_tuple(rule["parts"], param_name, tensor) + result = [] + start = 0 + for size in parts: + result.append(tensor[start:start + size, :]) + start += size + return result + + # Default: no processing, return as whole block + return [tensor] + + def merge_fn(param_name, parts_list): + """ + Merge the output of split_one (List[tensor]) back to 2D according to the same rules. + """ + concat = ops.Concat(axis=-1) + + for rule in schema: + if not any(fnmatch(param_name, pat) for pat in rule["patterns"]): + continue + + kind = rule["kind"] + + if kind == "periodic": + part_a, part_b, num_blocks = _eval_tuple(rule["parts"], param_name, parts_list[0]) + # Convention: periodic always has two segments + return muon_merge(parts_list[0], parts_list[1], part_a, part_b, num_blocks) + + if kind == "reshape_concat": + _, hidden_size, total_intermediate = _eval_tuple(rule["reshape"], param_name, parts_list[0]) + cat = concat([parts_list[0], parts_list[1]]) # [..., I] + [..., I] -> [..., 2I] + return tensor_reshape_to_2d(cat, total_intermediate) + + if kind == "reshape_only": + _, _, hidden_size = _eval_tuple(rule["reshape"], param_name, parts_list[0]) + # Only one segment, directly restore to 2D + return tensor_reshape_to_2d(parts_list[0], hidden_size) + + if kind == "alt_pair_periodic": + num_blocks = parts_list[0].shape[0] # 1 row per block + return muon_merge(parts_list[0], parts_list[1], 1, 1, num_blocks) + + if kind == "multi_split": + # Merge multiple parts back by concatenating along the first dimension + return ops.Concat(axis=0)(parts_list) + + # Default: directly take the first segment + return parts_list[0] + + return split_fn, merge_fn diff --git a/mindformers/pynative/trainer/trainer.py b/mindformers/pynative/trainer/trainer.py index 10969d0e275a3e387dd2b4743ff2e262ba3154fa..274e968a91ce6c893860ee6557718b5c3f07f00b 100644 --- a/mindformers/pynative/trainer/trainer.py +++ b/mindformers/pynative/trainer/trainer.py @@ -373,8 +373,16 @@ class Trainer: grouped_lr_scheduler=None, ) - # Build optimizer using default_args to inject params and lr - default_args = {"params": grouped_params, "learning_rate": lr} + # Check optimizer type for special handling + optimizer_type = getattr(optimizer_config, "type", "AdamW") + + if optimizer_type == "Muon": + # Muon optimizer requires model for QK-clip and muon functions + default_args = {"params": grouped_params, "learning_rate": lr, "model": self.model} + else: + # Standard optimizer (AdamW, etc.) + default_args = {"params": grouped_params, "learning_rate": lr} + optimizer = build_optim(optimizer_config, default_args=default_args) return optimizer, lr