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