diff --git a/README.md b/README.md index 5d0dfe0..329941e 100644 --- a/README.md +++ b/README.md @@ -131,7 +131,7 @@ The following is the list of models supported by MCore-Bridge: | GLM | glm4, glm4_moe, glm4_moe_lite
glm_moe_dsa | | MiniMax | minimax_m2 | | Kimi | kimi_k2, kimi_k25 | -| Bailing | bailing_moe | +| Bailing | bailing_moe, bailing_hybrid | | InternLM | internlm3 | | Llama | llama | | GPT-OSS | gpt_oss | diff --git a/README_zh.md b/README_zh.md index a5b6d17..60f9798 100644 --- a/README_zh.md +++ b/README_zh.md @@ -127,7 +127,7 @@ uv pip install -e . --torch-backend=auto | GLM | glm4, glm4_moe, glm4_moe_lite
glm_moe_dsa | | MiniMax | minimax_m2 | | Kimi | kimi_k2, kimi_k25 | -| Bailing | bailing_moe | +| Bailing | bailing_moe, bailing_hybrid | | InternLM | internlm3 | | Llama | llama | | GPT-OSS | gpt_oss | diff --git a/src/mcore_bridge/bridge/gpt_bridge.py b/src/mcore_bridge/bridge/gpt_bridge.py index 9be4917..f3111ff 100644 --- a/src/mcore_bridge/bridge/gpt_bridge.py +++ b/src/mcore_bridge/bridge/gpt_bridge.py @@ -31,6 +31,7 @@ class GPTBridge: hf_mtp_prefix = 'model.layers' hf_embed_key = 'model.embed_tokens.weight' hf_final_layernorm_key = 'model.norm.weight' + hf_mtp_final_layernorm_key = 'shared_head.norm.weight' hf_lm_head_key = 'lm_head.weight' hf_score_key = 'score.weight' hf_state_dict_mapping = {} @@ -542,6 +543,8 @@ def _reduce_tensor_pp_group(self, tensor, to_mcore, dtype=torch.bool, op=dist.Re return tensor def _set_qkv(self, mg_attn, hf_state_dict, to_mcore: bool, **kwargs): + # qkv: split along dim=0: [H*{qkv*a}, b] + # linear_fc1: split along dim=1, [2, x, y] config = self.config num_query_groups = kwargs.get('num_query_groups') if num_query_groups is None: @@ -1562,7 +1565,7 @@ def _set_mla_attn_state( hf_state_dict = self._remove_prefix(hf_state_dict, hf_prefix) else: hf_state_dict = {} - self._set_state_dict(mg_attn, 'linear_proj.weight', hf_state_dict, 'o_proj.weight', to_mcore) + self._set_state_dict(mg_attn, 'linear_proj.weight', hf_state_dict, f'{self.hf_o_proj_key}.weight', to_mcore) if self.config.q_lora_rank is None: self._set_state_dict(mg_attn, 'linear_q_proj.weight', hf_state_dict, 'q_proj.weight', to_mcore) else: @@ -1814,7 +1817,8 @@ def _convert_mtp_extra(self, mtp_layer, hf_state_dict, to_mcore, origin_hf_state for key in ['enorm.weight', 'hnorm.weight', 'eh_proj.weight']: self._set_state_dict(mtp_layer, key, hf_state_dict, key, to_mcore) self._fp8_skip_modules.update({'eh_proj'}) - self._set_state_dict(mtp_layer, 'final_layernorm.weight', hf_state_dict, 'shared_head.norm.weight', to_mcore) + self._set_state_dict(mtp_layer, 'final_layernorm.weight', hf_state_dict, self.hf_mtp_final_layernorm_key, + to_mcore) def _convert_mtp_layer(self, lm_model, hf_state_dict, hf_prefix: str, layer_idx: int, to_mcore: bool): mtp_layer = lm_model.mtp.layers[layer_idx] if hasattr(lm_model, 'mtp') else None diff --git a/src/mcore_bridge/config/parser.py b/src/mcore_bridge/config/parser.py index 99962ee..badaac3 100644 --- a/src/mcore_bridge/config/parser.py +++ b/src/mcore_bridge/config/parser.py @@ -233,6 +233,11 @@ def hf_to_mcore_config(hf_config: PretrainedConfig) -> Dict[str, Any]: res['moe_layer_freq'] = f"[{','.join(moe_layer_freq)}]" elif hf_model_type == 'glm4v': res['rotary_interleaved'] = True + elif llm_model_type == 'bailing_hybrid': + res['qk_layernorm'] = True + res['add_qkv_bias'] = False + res['moe_router_score_function'] = 'sigmoid' + res['moe_router_load_balancing_type'] = 'seq_aux_loss' if 'partial_rotary_factor' not in res and 'partial_rotary_factor' in rope_scaling: res['partial_rotary_factor'] = rope_scaling['partial_rotary_factor'] diff --git a/src/mcore_bridge/model/constant.py b/src/mcore_bridge/model/constant.py index 9708f6a..30bf289 100644 --- a/src/mcore_bridge/model/constant.py +++ b/src/mcore_bridge/model/constant.py @@ -9,6 +9,7 @@ class LLMModelType: minimax_m2 = 'minimax_m2' hy_v3 = 'hy_v3' bailing_moe = 'bailing_moe' + bailing_hybrid = 'bailing_hybrid' deepseek_v4 = 'deepseek_v4' qwen3_emb = 'qwen3_emb' diff --git a/src/mcore_bridge/model/gpts/__init__.py b/src/mcore_bridge/model/gpts/__init__.py index 6eb44db..52025b8 100644 --- a/src/mcore_bridge/model/gpts/__init__.py +++ b/src/mcore_bridge/model/gpts/__init__.py @@ -1,2 +1,2 @@ # Copyright (c) ModelScope Contributors. All rights reserved. -from . import bailing_moe, deepseek_v4, glm4, hunyuan, llm, minimax_m2, olmoe, qwen3_emb, qwen3_next +from . import bailing_hybrid, bailing_moe, deepseek_v4, glm4, hunyuan, llm, minimax_m2, olmoe, qwen3_emb, qwen3_next diff --git a/src/mcore_bridge/model/gpts/bailing_hybrid.py b/src/mcore_bridge/model/gpts/bailing_hybrid.py new file mode 100644 index 0000000..25cdae2 --- /dev/null +++ b/src/mcore_bridge/model/gpts/bailing_hybrid.py @@ -0,0 +1,258 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +import math +import torch +from contextlib import contextmanager +from megatron.core import parallel_state +from megatron.core.extensions.transformer_engine import TEColumnParallelLinear, TELinear +from megatron.core.models.common.embeddings.rope_utils import apply_rotary_pos_emb +from megatron.core.models.common.embeddings.yarn_rotary_pos_embedding import _yarn_get_concentration_factor_from_config +from megatron.core.tensor_parallel.mappings import (gather_from_tensor_model_parallel_region, + scatter_to_tensor_model_parallel_region) +from megatron.core.transformer.attention import SelfAttention +from megatron.core.transformer.transformer_config import TransformerConfig +from megatron.core.utils import nvtx_range_pop, nvtx_range_push +from torch import Tensor, nn +from typing import Optional, Tuple + +from ..constant import ModelType +from ..register import ModelLoader, ModelMeta, register_model +from .bailing_moe import BailingMoeBridge + +try: + from fla.ops.simple_gla.fused_recurrent import fused_recurrent_simple_gla +except ImportError: + fused_recurrent_simple_gla = None + + +class BailingHybridBridge(BailingMoeBridge): + additional_dim0_keys = {'g_proj'} + + def _set_layer_attn(self, mg_layer, hf_state_dict, layer_idx: int, to_mcore: bool): + layer_type = self.config.hf_config.attention_layer_type[layer_idx] + mg_attn = None if mg_layer is None else mg_layer.self_attention + if layer_type == 'attention': + hf_state_dict.update( + self._set_mla_attn_state(mg_attn, hf_state_dict, f'{self.hf_attn_prefix}.', layer_idx, to_mcore)) + + elif layer_type == 'linear_attention': + hf_state_dict.update( + self._set_attn_state(mg_attn, hf_state_dict, f'{self.hf_attn_prefix}.', layer_idx, to_mcore)) + for key in ['g_proj', 'g_norm']: + self._set_state_dict(mg_layer, f'self_attention.{key}.weight', hf_state_dict, f'attention.{key}.weight', + to_mcore) + self._set_state_dict(mg_layer, 'input_layernorm.weight', hf_state_dict, self.hf_input_layernorm_key, to_mcore) + return hf_state_dict + + +class BailingMoeV2_5GroupRMSNorm(nn.Module): + + def __init__(self, config, hidden_size, group_norm_size, eps=1e-6): + super().__init__() + self.config = config + assert hidden_size % group_norm_size == 0, 'hidden_size must be divisible by group_norm_size' + self.hidden_size = hidden_size + self.group_norm_size = group_norm_size + self.variance_epsilon = eps + self.weight = nn.Parameter(torch.ones(hidden_size)) + + def forward(self, hidden_states): + input_dtype = hidden_states.dtype + input_shape = hidden_states.size() + group_input_shape = input_shape[:-1] + (self.group_norm_size, input_shape[-1] // self.group_norm_size) + hidden_states = hidden_states.view(group_input_shape) + hidden_states = hidden_states.to(torch.float32) + variance = hidden_states.pow(2).mean(-1, keepdim=True) + hidden_states = hidden_states * torch.rsqrt(variance + self.variance_epsilon) + return self.weight * hidden_states.to(input_dtype).view(input_shape) + + +class LinearAttention(SelfAttention): + + def __init__(self, config: TransformerConfig, *args, **kwargs): + if fused_recurrent_simple_gla is None: + raise ImportError('flash-linear-attention is required but not installed. ' + 'Please install it via: ' + "`pip install -U 'flash-linear-attention' --no-build-isolation`") + super().__init__(config, *args, **kwargs) + self.g_proj = TEColumnParallelLinear( + input_size=config.hidden_size, + output_size=self.query_projection_size, + bias=False, + skip_bias_add=False, + init_method=config.init_method, + skip_weight_param_allocation=False, + gather_output=False, + is_expert=False, + config=config, + ) + self.g_norm = BailingMoeV2_5GroupRMSNorm( + config, + self.query_projection_size, + group_norm_size=config.hf_config.group_norm_size, + eps=config.layernorm_epsilon) + self.g_norm.weight.average_gradients_across_tp_domain = True # No need to set `sequence_parallel`. + # https://github.com/sgl-project/sglang/blob/8e0ed75f2d5417015329095dc9a1626df2895acf/python/sglang/srt/layers/attention/linear/lightning_backend.py#L144C12-L149 # noqa + slope = -self.build_slope_tensor(config.num_attention_heads) * (1 - (self.layer_number - 1) / + (config.num_layers - 1) + 1e-5) + # Slice slope to current TP rank: each rank only owns `num_attention_heads_per_partition` heads. + tp_rank = parallel_state.get_tensor_model_parallel_rank() + heads_per_partition = self.num_attention_heads_per_partition + slope = slope[tp_rank * heads_per_partition:(tp_rank + 1) * heads_per_partition].contiguous() + self.register_buffer('slope', slope, persistent=False) + + @staticmethod + def build_slope_tensor(n_attention_heads: int): + """ + Build a tensor of slopes for Lightning Attention-2 as described in the paper: + "Lightning Attention-2: A Free Lunch for Handling Unlimited Sequence Lengths in Large Language Models" + (https://arxiv.org/abs/2401.04658) + This function computes the slope values that control the decay rate of attention scores + based on the number of attention heads. The slopes are designed to have specific + mathematical properties that work optimally when the number of heads is a power of 2. + For non-power-of-2 head counts, a workaround is implemented to maintain similar properties. + Args: + n_attention_heads (int): Number of attention heads in the model + Returns: + torch.Tensor: A tensor of shape [n_attention_heads] containing the computed slopes + Note: + Code copied from: https://github.com/OpenNLPLab/lightning-attention/blob/d15c38529bbd5c2c82b44ddda3cac885825aa873/lightning_attn/utils/utils.py#L6 # noqa + """ + + def get_slopes(n): + + def get_slopes_power_of_2(n): + start = 2**(-(2**-(math.log2(n) - 3))) + ratio = start + return [start * ratio**i for i in range(n)] + + if math.log2(n).is_integer(): + return get_slopes_power_of_2( + n) # In the paper, we only train models that have 2^a heads for some a. This function has + else: # some good properties that only occur when the input is a power of 2. To maintain that even + closest_power_of_2 = 2**math.floor( + math.log2(n)) # when the number of heads is not a power of 2, we use this workaround. + return (get_slopes_power_of_2(closest_power_of_2) + + get_slopes(2 * closest_power_of_2)[0::2][:n - closest_power_of_2]) + + slopes = torch.tensor(get_slopes(n_attention_heads), dtype=torch.float) + return slopes + + @contextmanager + def _patch_attention_scaling(self): + multi_latent_attention = self.config.multi_latent_attention + self.config.multi_latent_attention = False + try: + yield + finally: + self.config.multi_latent_attention = multi_latent_attention + + def _apply_rotary(self, query, key, rotary_pos_emb, cu_seqlens=None): + if cu_seqlens is not None: + query = query.squeeze(1) + key = key.squeeze(1) + nvtx_range_push(suffix='rotary_pos_emb') + q_pos_emb, k_pos_emb = rotary_pos_emb + + if q_pos_emb is not None: + # TODO VIJAY: simplify + query = apply_rotary_pos_emb( + query, + q_pos_emb, + config=self.config, + cu_seqlens=cu_seqlens, + mscale=_yarn_get_concentration_factor_from_config(self.config), + cp_group=self.pg_collection.cp, + ) + if k_pos_emb is not None: + key = apply_rotary_pos_emb( + key, + k_pos_emb, + config=self.config, + cu_seqlens=cu_seqlens, + mscale=_yarn_get_concentration_factor_from_config(self.config), + cp_group=self.pg_collection.cp, + ) + nvtx_range_pop(suffix='rotary_pos_emb') + if cu_seqlens is not None: + query = query.unsqueeze(1) + key = key.unsqueeze(1) + return query, key + + def _forward_core_attention(self, query, key, value, attention_mask, cu_seqlens=None): + nvtx_range_push(suffix='core_attention') + query = query.transpose(0, 1) + core_attn_out, _ = fused_recurrent_simple_gla( + q=query, + k=key.transpose(0, 1), + v=value.transpose(0, 1), + g=self.slope[None, None, :].expand(*query.shape[:2], self.num_attention_heads_per_partition), + initial_state=None, + output_final_state=False, + cu_seqlens=cu_seqlens, + ) + nvtx_range_pop(suffix='core_attention') + core_attn_out = core_attn_out.view(*core_attn_out.shape[:2], -1) + return core_attn_out.transpose(0, 1) + + def forward(self, hidden_states: Tensor, attention_mask: Tensor, **kwargs) -> Tuple[Tensor, Tensor]: + rotary_pos_emb = kwargs.get('rotary_pos_emb') + packed_seq_params = kwargs.get('packed_seq_params') + query, key, value = self.get_query_key_value_tensors(hidden_states) + if isinstance(rotary_pos_emb, torch.Tensor): + rotary_pos_emb = (rotary_pos_emb, ) * 2 + if packed_seq_params is not None and packed_seq_params.qkv_format == 'thd': + if packed_seq_params.cu_seqlens_q_padded is not None: + cu_seqlens_q = packed_seq_params.cu_seqlens_q_padded + else: + cu_seqlens_q = packed_seq_params.cu_seqlens_q + else: + cu_seqlens_q = None + with self._patch_attention_scaling(): + query, key = self._apply_rotary(query, key, rotary_pos_emb, cu_seqlens_q) + core_attn_out = self._forward_core_attention(query, key, value, attention_mask, cu_seqlens_q) + enable_tp = self.config.tensor_model_parallel_size > 1 + if enable_tp: + core_attn_out = gather_from_tensor_model_parallel_region(core_attn_out) + core_attn_out = self.g_norm(core_attn_out) + if enable_tp: + core_attn_out = scatter_to_tensor_model_parallel_region(core_attn_out) + g_proj = self.g_proj(hidden_states)[0] + core_attn_out = core_attn_out * torch.sigmoid_(g_proj) + nvtx_range_push(suffix='linear_proj') + output, bias = self.linear_proj(core_attn_out) + nvtx_range_pop(suffix='linear_proj') + return output, bias + + +class BailingHybridLoader(ModelLoader): + + def get_transformer_layer_spec(self, vp_stage: Optional[int] = None): + hf_config = self.config.hf_config + num_layers = hf_config.num_hidden_layers + group_size = hf_config.layer_group_size + tail_start = num_layers // group_size * group_size + hf_config.attention_layer_type = [ + 'attention' if (layer_idx + 1) % group_size == 0 or layer_idx >= tail_start else 'linear_attention' + for layer_idx in range(num_layers) + ] + layer_specs = super().get_transformer_layer_spec(vp_stage=vp_stage) + multi_latent_attention = self.config.multi_latent_attention + self.config.multi_latent_attention = False + linear_layer_specs = super().get_transformer_layer_spec(vp_stage=vp_stage) + self.config.multi_latent_attention = multi_latent_attention + for i, layer_spec in enumerate(layer_specs.layer_specs): + if hf_config.attention_layer_type[i] == 'linear_attention': + linear_spec = linear_layer_specs.layer_specs[i].submodules.self_attention + linear_spec.module = LinearAttention + linear_spec.submodules.linear_qkv = TEColumnParallelLinear + layer_spec.submodules.self_attention = linear_spec + return layer_specs + + +register_model( + ModelMeta( + ModelType.bailing_hybrid, + ['bailing_hybrid'], + bridge_cls=BailingHybridBridge, + loader=BailingHybridLoader, + )) diff --git a/src/mcore_bridge/model/gpts/bailing_moe.py b/src/mcore_bridge/model/gpts/bailing_moe.py index 7e2a70d..5ee79b7 100644 --- a/src/mcore_bridge/model/gpts/bailing_moe.py +++ b/src/mcore_bridge/model/gpts/bailing_moe.py @@ -1,66 +1,12 @@ # Copyright (c) ModelScope Contributors. All rights reserved. import torch -from megatron.core.transformer.attention import SelfAttention -from torch import Tensor -from typing import Optional +import torch.distributed as dist from mcore_bridge.bridge import GPTBridge +from mcore_bridge.tuners import LoraParallelLinear from ..constant import ModelType -from ..register import ModelLoader, ModelMeta, register_model - - -class BailingMoeSelfAttention(SelfAttention): - - def get_query_key_value_tensors( - self, - hidden_states: Tensor, - key_value_states: Optional[Tensor] = None, - *args, - **kwargs, - ): - """Override to handle BailingMoE's non-interleaved QKV weight layout. - - BailingMoE stores weights as [Q_all | K_all | V_all] (split by head count), - not Megatron's interleaved [q1 q2 k1 v1 | q3 q4 k2 v2 | ...]. - """ - # [sq, b, h] --> [sq, b, (num_heads + 2 * num_kv_heads) * head_dim] - mixed_qkv, _ = self.linear_qkv(hidden_states) - - # [sq, b, (num_heads + 2 * num_kv_heads) * head_dim] - # --> [sq, b, num_heads + 2 * num_kv_heads, head_dim] - new_tensor_shape = mixed_qkv.size()[:-1] + ( - self.num_attention_heads_per_partition + 2 * self.num_query_groups_per_partition, - self.hidden_size_per_attention_head, - ) - mixed_qkv = mixed_qkv.view(*new_tensor_shape) - - # Split by head count: [sq, b, num_heads, hn], [sq, b, num_kv_heads, hn], [sq, b, num_kv_heads, hn] - query, key, value = torch.split( - mixed_qkv, - [ - self.num_attention_heads_per_partition, self.num_query_groups_per_partition, - self.num_query_groups_per_partition - ], - dim=2, - ) - - if self.q_layernorm is not None: - query = self.q_layernorm(query) - - if self.k_layernorm is not None: - key = self.k_layernorm(key) - - return query, key, value - - -class BailingMoeLoader(ModelLoader): - - def get_transformer_layer_spec(self, vp_stage: Optional[int] = None): - transformer_layer_spec = super().get_transformer_layer_spec(vp_stage) - for layer_spec in transformer_layer_spec.layer_specs: - layer_spec.submodules.self_attention.module = BailingMoeSelfAttention - return transformer_layer_spec +from ..register import ModelMeta, register_model class BailingMoeBridge(GPTBridge): @@ -70,17 +16,92 @@ class BailingMoeBridge(GPTBridge): hf_k_norm_key = 'key_layernorm.weight' hf_expert_bias_key = 'gate.expert_bias' hf_o_proj_key = 'dense' + hf_mtp_final_layernorm_key = 'final_layernorm.weight' def _set_qkv(self, mg_attn, hf_state_dict, to_mcore: bool, **kwargs): - self._set_state_dict(mg_attn, 'linear_qkv.weight', hf_state_dict, 'query_key_value.weight', to_mcore) + config = self.config + num_heads = config.num_attention_heads + num_query_groups = config.num_query_groups if config.num_query_groups is not None else num_heads + assert num_heads % num_query_groups == 0, ( + f'num_attention_heads ({num_heads}) must be divisible by num_query_groups ({num_query_groups})') + q_per_group = num_heads // num_query_groups + head_dim = config.kv_channels + hidden_size = config.hidden_size + hidden_size_block = hidden_size // self.fp8_block_size + + def hf_to_mg(w, per_head_rows, last_dim): + # HF [Q_all (N*r) | K_all (G*r) | V_all (G*r)] -> MG grouped interleaved (G,(qpg+2)*r) + total_q = num_heads * per_head_rows + total_kv = num_query_groups * per_head_rows + q = w[:total_q].reshape(num_query_groups, q_per_group, per_head_rows, last_dim) + k = w[total_q:total_q + total_kv].reshape(num_query_groups, 1, per_head_rows, last_dim) + v = w[total_q + total_kv:].reshape(num_query_groups, 1, per_head_rows, last_dim) + return torch.cat([q, k, v], dim=1).reshape(-1, last_dim) + + def mg_to_hf(w, per_head_rows, last_dim): + # MG grouped interleaved -> HF [Q_all | K_all | V_all] + w = w.reshape(num_query_groups, q_per_group + 2, per_head_rows, last_dim) + q = w[:, :q_per_group, :, :].reshape(-1, last_dim) + k = w[:, q_per_group:q_per_group + 1, :, :].reshape(-1, last_dim) + v = w[:, q_per_group + 1:, :, :].reshape(-1, last_dim) + return torch.cat([q, k, v], dim=0) + + if to_mcore: + if isinstance(mg_attn.linear_qkv, LoraParallelLinear): + # LoRA on fused QKV: lora_A is shared (input side), lora_B needs same row layout transform. + lora_A = hf_state_dict['query_key_value.lora_A.weight'].load() + lora_B = hf_state_dict['query_key_value.lora_B.weight'].load() + lora_B = hf_to_mg(lora_B, head_dim, lora_B.shape[-1]) + self._set_weight(mg_attn.linear_qkv.lora_A[self._adapter_name].weight, lora_A, + 'linear_qkv.lora_A.weight') + self._set_weight(mg_attn.linear_qkv.lora_B[self._adapter_name].weight, lora_B, + 'linear_qkv.lora_B.weight') + elif not self._peft_format: + qkv = hf_state_dict['query_key_value.weight'].load() + qkv = hf_to_mg(qkv, head_dim, hidden_size) + qkv_scale_inv = None + if 'query_key_value.weight_scale_inv' in hf_state_dict: + assert head_dim % self.fp8_block_size == 0, ( + f'head_dim ({head_dim}) must be divisible by fp8_block_size ({self.fp8_block_size})') + head_dim_block = head_dim // self.fp8_block_size + qkv_scale_inv = hf_state_dict['query_key_value.weight_scale_inv'].load() + qkv_scale_inv = hf_to_mg(qkv_scale_inv, head_dim_block, hidden_size_block) + self._set_weight(mg_attn.linear_qkv.weight, qkv, 'linear_qkv.weight', hf_scale_inv=qkv_scale_inv) + else: + is_lora = False if mg_attn is None else (isinstance(mg_attn.linear_qkv, LoraParallelLinear) + and self._peft_format) + is_lora = torch.tensor([is_lora], dtype=torch.bool, device='cuda') + if self.pp_size > 1: + dist.all_reduce(is_lora, group=self.pp_group) + if is_lora: + lora_A, _ = self._get_weight( + None if mg_attn is None else mg_attn.linear_qkv.lora_A[self._adapter_name].weight.data, + f'linear_qkv.lora_A.{self._adapter_name}.weight') + lora_B, _ = self._get_weight( + None if mg_attn is None else mg_attn.linear_qkv.lora_B[self._adapter_name].weight.data, + f'linear_qkv.lora_B.{self._adapter_name}.weight') + if lora_A is not None: + self._peft_target_modules.update({'query_key_value'}) + hf_state_dict['query_key_value.lora_A.weight'] = lora_A.clone() + hf_state_dict['query_key_value.lora_B.weight'] = mg_to_hf(lora_B, head_dim, lora_B.shape[-1]) + elif not self._peft_format: + mg_w, scale_inv = self._get_weight(None if mg_attn is None else mg_attn.linear_qkv.weight.data, + 'linear_qkv.weight') + if mg_w is not None: + hf_state_dict['query_key_value.weight'] = mg_to_hf(mg_w, head_dim, hidden_size) + if scale_inv is not None: + assert head_dim % self.fp8_block_size == 0, ( + f'head_dim ({head_dim}) must be divisible by fp8_block_size ({self.fp8_block_size})') + head_dim_block = head_dim // self.fp8_block_size + hf_state_dict['query_key_value.weight_scale_inv'] = mg_to_hf(scale_inv, head_dim_block, + hidden_size_block) + del mg_w assert not self.config.add_bias_linear return hf_state_dict -register_model( - ModelMeta( - ModelType.bailing_moe, - ['bailing_moe'], - bridge_cls=BailingMoeBridge, - loader=BailingMoeLoader, - )) +register_model(ModelMeta( + ModelType.bailing_moe, + ['bailing_moe'], + bridge_cls=BailingMoeBridge, +)) diff --git a/src/mcore_bridge/model/register.py b/src/mcore_bridge/model/register.py index eb7166a..be81845 100644 --- a/src/mcore_bridge/model/register.py +++ b/src/mcore_bridge/model/register.py @@ -1,6 +1,7 @@ # Copyright (c) ModelScope Contributors. All rights reserved. import megatron.core from contextlib import contextmanager +from copy import deepcopy from dataclasses import dataclass from functools import partial from megatron.core import mpu @@ -104,13 +105,28 @@ def _replace_spec_dsa(self, layer_spec): dsa_spec.submodules.linear_kv_up_proj = linear_q_up_proj layer_spec.submodules.self_attention = dsa_spec + @contextmanager + def _patch_experimental_attention_variant(self): + experimental_attention_variant = self.config.experimental_attention_variant + self.config.experimental_attention_variant = None + try: + yield + finally: + self.config.experimental_attention_variant = experimental_attention_variant + + def _deepcopy_layer_spec(self, transformer_layer_spec): + for i, layer_spec in enumerate(transformer_layer_spec.layer_specs): + transformer_layer_spec.layer_specs[i] = deepcopy(layer_spec) + def get_transformer_layer_spec(self, vp_stage: Optional[int] = None): - transformer_layer_spec = get_gpt_decoder_block_spec( - self.config, - use_transformer_engine=True, - normalization=self.config.normalization, - qk_l2_norm=self.config.qk_l2_norm, - vp_stage=vp_stage) + with self._patch_experimental_attention_variant(): + transformer_layer_spec = get_gpt_decoder_block_spec( + self.config, + use_transformer_engine=True, + normalization=self.config.normalization, + qk_l2_norm=self.config.qk_l2_norm, + vp_stage=vp_stage) + self._deepcopy_layer_spec(transformer_layer_spec) if self.config.experimental_attention_variant == 'dsa': for layer_spec in transformer_layer_spec.layer_specs: self._replace_spec_dsa(layer_spec) diff --git a/src/mcore_bridge/patcher.py b/src/mcore_bridge/patcher.py index 4d00cd3..0649fa7 100644 --- a/src/mcore_bridge/patcher.py +++ b/src/mcore_bridge/patcher.py @@ -195,28 +195,9 @@ def forward(self, position_ids, mrope_section: List[int], mrope_interleaved: boo MultimodalRotaryEmbedding.forward = forward _origin_apply_rotary_pos_emb_thd = rope_utils._apply_rotary_pos_emb_thd - def _apply_rotary_pos_emb_thd( - t: torch.Tensor, - cu_seqlens: torch.Tensor, - freqs: torch.Tensor, - rotary_interleaved: bool = False, - multi_latent_attention: bool = False, - mscale: float = 1.0, - cp_group: torch.distributed.ProcessGroup = None, - **kwargs, - ) -> torch.Tensor: - """A baseline implementation of applying RoPE for `thd` format. - - Args: - t (Tensor): Input tensor T is of shape [t, h, d] - cu_seqlens(Tensor): Cumulative sum of sequence lengths in a batch for `t`, - with shape [b + 1] and dtype torch.int32. - freqs (Tensor): Rotary Positional embedding tensor freq is of shape [max_s, 1, 1, d] - cp_group (torch.distributed.ProcessGroup): The context parallel group - - Returns: - Tensor: Shape [t, h, d]. The input tensor after applying RoPE. - """ + def _apply_rotary_pos_emb_thd(t: torch.Tensor, cu_seqlens: torch.Tensor, freqs: torch.Tensor, *args, + **kwargs) -> torch.Tensor: + cp_group = kwargs.pop('cp_group', None) if cp_group is not None: cp_size = cp_group.size() else: @@ -225,25 +206,9 @@ def _apply_rotary_pos_emb_thd( use_batched_rope = (freqs.dim() >= 1 and freqs.shape[0] == cu_seqlens_for_batched[-1]).item() if not use_batched_rope: logger.warning_once('Using non-batched RoPE, which may affect performance.') - return _origin_apply_rotary_pos_emb_thd( - t, - cu_seqlens, - freqs, - rotary_interleaved=rotary_interleaved, - multi_latent_attention=multi_latent_attention, - mscale=mscale, - cp_group=cp_group, - **kwargs, - ) + return _origin_apply_rotary_pos_emb_thd(t, cu_seqlens, freqs, *args, **kwargs) - return rope_utils._apply_rotary_pos_emb_bshd( - t.unsqueeze(1), - freqs, - rotary_interleaved=rotary_interleaved, - multi_latent_attention=multi_latent_attention, - mscale=mscale, - **kwargs, - ).squeeze(1) + return rope_utils._apply_rotary_pos_emb_bshd(t.unsqueeze(1), freqs, *args, **kwargs).squeeze(1) rope_utils._apply_rotary_pos_emb_thd = _apply_rotary_pos_emb_thd