Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions src/mcore_bridge/config/parser.py
Original file line number Diff line number Diff line change
Expand Up @@ -213,6 +213,8 @@ def hf_to_mcore_config(hf_config: PretrainedConfig) -> Dict[str, Any]:
res.setdefault('linear_attention_freq', 4)
elif llm_model_type == 'minimax_m2':
res['add_qkv_bias'] = False
elif llm_model_type == 'olmoe':
res['qk_layernorm'] = True
elif hf_model_type == 'llama4':
qk_layernorm = res.pop('qk_layernorm', False)
if qk_layernorm:
Expand Down
249 changes: 61 additions & 188 deletions src/mcore_bridge/model/gpts/olmoe.py
Original file line number Diff line number Diff line change
@@ -1,222 +1,95 @@
import torch
import torch.distributed as dist
from copy import deepcopy
from megatron.core.extensions.transformer_engine import SplitAlongDim, TENorm
from megatron.core.models.gpt.gpt_layer_specs import get_gpt_layer_with_transformer_engine_spec
from megatron.core.transformer.attention import SelfAttention as SelfAttentionBase
from megatron.core.transformer.attention import SelfAttentionSubmodules
# Copyright (c) ModelScope Contributors. All rights reserved.
from megatron.core import mpu
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, SelfAttentionSubmodules
from megatron.core.transformer.identity_op import IdentityOp
from megatron.core.transformer.spec_utils import build_module
from megatron.core.transformer.transformer_block import TransformerBlockSubmodules, get_num_layers_to_build
from megatron.core.transformer.transformer_layer import get_transformer_layer_offset
from typing import Optional

from mcore_bridge.bridge import GPTBridge
from mcore_bridge.config import ModelConfig
from mcore_bridge.tuners import LoraParallelLinear

from ..constant import ModelType
from ..register import ModelLoader, ModelMeta, register_model


class OLMoESelfAttention(SelfAttentionBase):

def __init__(self, config: ModelConfig, submodules: SelfAttentionSubmodules, *args, **kwargs):
super().__init__(config, submodules, *args, **kwargs)
self.q_layernorm = build_module(
class OLMoESelfAttention(SelfAttention):

def __init__(
self,
config: ModelConfig,
submodules: SelfAttentionSubmodules,
*args,
**kwargs,
):
q_layernorm = submodules.q_layernorm
k_layernorm = submodules.k_layernorm
submodules.q_layernorm = IdentityOp
submodules.k_layernorm = IdentityOp
try:
super().__init__(config, submodules, *args, **kwargs)
finally:
submodules.q_layernorm = q_layernorm
submodules.k_layernorm = k_layernorm
# OLMoE applies q/k RMSNorm over the full Q/K channels (not per-head).
self.q_norm = build_module(
submodules.q_layernorm,
hidden_size=self.hidden_size_per_attention_head * self.num_attention_heads_per_partition,
hidden_size=self.hidden_size_per_attention_head * config.num_attention_heads,
config=self.config,
eps=self.config.layernorm_epsilon,
)
self.k_layernorm = build_module(
self.k_norm = build_module(
submodules.k_layernorm,
hidden_size=self.hidden_size_per_attention_head * self.num_query_groups_per_partition,
hidden_size=self.hidden_size_per_attention_head * config.num_query_groups,
config=self.config,
eps=self.config.layernorm_epsilon,
)

def get_query_key_value_tensors(self, hidden_states, key_value_states=None, *args, **kwargs):
"""
Derives `query`, `key` and `value` tensors from `hidden_states`.
"""
# Attention heads [sq, b, h] --> [sq, b, ng * (np/ng + 2) * hn)]
mixed_qkv, _ = self.linear_qkv(hidden_states)

# [sq, b, ng * (np/ng + 2) * hn] -> [sq, b, np * hn], [sq, b, ng * hn], [sq, b, ng * hn]
split_arg_list = [
self.hidden_size_per_attention_head * self.num_attention_heads_per_partition,
self.hidden_size_per_attention_head * self.num_query_groups_per_partition,
self.hidden_size_per_attention_head * self.num_query_groups_per_partition
]

if SplitAlongDim is not None:
(query, key, value) = SplitAlongDim(mixed_qkv, 2, split_arg_list)
else:
(query, key, value) = torch.split(mixed_qkv, split_arg_list, 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)
query = query.reshape(query.size(0), query.size(1), -1, self.hidden_size_per_attention_head)
key = key.reshape(key.size(0), key.size(1), -1, self.hidden_size_per_attention_head)
value = value.reshape(value.size(0), value.size(1), -1, self.hidden_size_per_attention_head)

if self.config.test_mode:
self.run_realtime_tests()

# These norms operate on the gathered (TP-replicated) tensor, so they must
# NOT participate in the cross-TP SP/QK-layernorm grad AllReduce,
# otherwise the gradients will be amplified by tp_size.
for norm in (self.q_norm, self.k_norm):
for p in norm.parameters():
p.sequence_parallel = False
p.average_gradients_across_tp_domain = True
Comment on lines +51 to +54

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

medium

If submodules.q_layernorm or submodules.k_layernorm is None (for example, if QK layernorm is disabled or not configured in a custom setup), self.q_norm or self.k_norm will be initialized as None. Accessing .parameters() on a None object will raise an AttributeError. Adding a defensive check to ensure the norm is not None before iterating over its parameters prevents potential runtime crashes.

Suggested change
for norm in (self.q_norm, self.k_norm):
for p in norm.parameters():
p.sequence_parallel = False
p.average_gradients_across_tp_domain = True
for norm in (self.q_norm, self.k_norm):
if norm is not None:
for p in norm.parameters():
p.sequence_parallel = False
p.average_gradients_across_tp_domain = True


def get_query_key_value_tensors(self, *_args, **kwargs):
enable_tp = mpu.get_tensor_model_parallel_world_size() > 1
query, key, value = super().get_query_key_value_tensors(*_args, **kwargs)
query = query.reshape(*query.shape[:-2], -1)
key = key.reshape(*key.shape[:-2], -1)
if enable_tp:
query = gather_from_tensor_model_parallel_region(query)
key = gather_from_tensor_model_parallel_region(key)
query = self.q_norm(query)
key = self.k_norm(key)
Comment on lines +64 to +65

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

medium

Similarly, if self.q_norm or self.k_norm is None, calling them directly will raise a TypeError. Adding defensive None checks before applying the norms ensures robustness.

Suggested change
query = self.q_norm(query)
key = self.k_norm(key)
if self.q_norm is not None:
query = self.q_norm(query)
if self.k_norm is not None:
key = self.k_norm(key)

if enable_tp:
query = scatter_to_tensor_model_parallel_region(query)
key = scatter_to_tensor_model_parallel_region(key)
query = query.view(*query.shape[:2], -1, self.hidden_size_per_attention_head)
key = key.view(*key.shape[:2], -1, self.hidden_size_per_attention_head)
return query, key, value


def get_olmoe_decoder_block_spec(
config: ModelConfig,
vp_stage: Optional[int] = None,
) -> TransformerBlockSubmodules:
"""GPT block spec."""
layer_norm_impl = TENorm
moe_layer_spec = get_gpt_layer_with_transformer_engine_spec(
num_experts=config.num_moe_experts,
moe_grouped_gemm=config.moe_grouped_gemm,
qk_layernorm=True,
multi_latent_attention=False,
use_kitchen=config.use_kitchen,
)
layer_specs = []
for _ in range(config.num_layers):
layer_spec = deepcopy(moe_layer_spec)
layer_spec.submodules.self_attention.module = OLMoESelfAttention
layer_specs.append(layer_spec)

num_layers_to_build = get_num_layers_to_build(config, vp_stage=vp_stage)

if config.pipeline_model_parallel_layout is not None:
from megatron.core.transformer.enums import LayerType
local_layer_specs = [
layer_specs[layer_id] for layer_id in config.pipeline_model_parallel_layout.get_layer_id_list(
layer_type=LayerType.decoder, vp_stage=vp_stage)
]
else:
offset = get_transformer_layer_offset(config, vp_stage=vp_stage)
local_layer_specs = layer_specs[offset:offset + num_layers_to_build]

# Block spec.
block_spec = TransformerBlockSubmodules(layer_specs=local_layer_specs, layer_norm=layer_norm_impl)

return block_spec


class OLMoEBridge(GPTBridge):

def _set_attn_state(self, mg_attn, hf_state_dict, hf_prefix: str, layer_idx: int, to_mcore: bool):
if to_mcore:
hf_state_dict = self._remove_prefix(hf_state_dict, hf_prefix)
else:
hf_state_dict = {}
config = self.config
if to_mcore:
if isinstance(mg_attn.linear_qkv, LoraParallelLinear):
lora_A = hf_state_dict['q_proj.lora_A.weight'].load()
assert (lora_A == hf_state_dict['k_proj.lora_A.weight'].load()).all() and (
lora_A == hf_state_dict['v_proj.lora_A.weight'].load()
).all(), 'Need to ensure QKV\'s lora_A are consistent'
lora_B = torch.cat([
hf_state_dict['q_proj.lora_B.weight'].load(),
hf_state_dict['k_proj.lora_B.weight'].load(),
hf_state_dict['v_proj.lora_B.weight'].load(),
],
dim=0)
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:
linear_qkv_weight = torch.cat([
hf_state_dict['q_proj.weight'].load(),
hf_state_dict['k_proj.weight'].load(),
hf_state_dict['v_proj.weight'].load(),
],
dim=0)
qkv_scale_inv = None
if 'q_proj.weight_scale_inv' in hf_state_dict:
qkv_scale_inv = torch.cat([
hf_state_dict['q_proj.weight_scale_inv'].load(),
hf_state_dict['k_proj.weight_scale_inv'].load(),
hf_state_dict['v_proj.weight_scale_inv'].load(),
],
dim=0)
self._set_weight(
mg_attn.linear_qkv.weight, linear_qkv_weight, 'linear_qkv.weight', hf_scale_inv=qkv_scale_inv)
else:
q_dim = self.config.kv_channels * self.config.num_attention_heads // self.config.num_query_groups
kv_dim = self.config.kv_channels
q_block = q_dim // self.fp8_block_size
kv_block = kv_dim // self.fp8_block_size
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({'q_proj', 'k_proj', 'v_proj'})
for key in ['q_proj', 'k_proj', 'v_proj']:
hf_state_dict[f'{key}.lora_A.weight'] = lora_A.clone()
hf_state_dict['q_proj.lora_B.weight'] = lora_B[:q_dim, :].clone()
hf_state_dict['k_proj.lora_B.weight'] = lora_B[q_dim:-kv_dim, :].clone()
hf_state_dict['v_proj.lora_B.weight'] = lora_B[-kv_dim:, :].clone()
elif not self._peft_format:
mg_attn_weight, scale_inv = self._get_weight(
None if mg_attn is None else mg_attn.linear_qkv.weight.data, 'linear_qkv.weight')
if mg_attn_weight is not None:
hf_state_dict['q_proj.weight'] = mg_attn_weight[:q_dim, :].clone()
hf_state_dict['k_proj.weight'] = mg_attn_weight[q_dim:-kv_dim, :].clone()
hf_state_dict['v_proj.weight'] = mg_attn_weight[-kv_dim:, :].clone()
if scale_inv is not None:
hf_state_dict['q_proj.weight_scale_inv'] = scale_inv[:q_block, :].clone()
hf_state_dict['k_proj.weight_scale_inv'] = scale_inv[q_block:-kv_block, :].clone()
hf_state_dict['v_proj.weight_scale_inv'] = scale_inv[-kv_block:, :].clone()
del mg_attn_weight
self._set_state_dict(mg_attn, 'linear_proj.weight', hf_state_dict, 'o_proj.weight', to_mcore)
if config.add_qkv_bias and not self._peft_format:
if to_mcore:
linear_qkv_bias = torch.cat([
hf_state_dict['q_proj.bias'].load(),
hf_state_dict['k_proj.bias'].load(),
hf_state_dict['v_proj.bias'].load(),
],
dim=0)
self._set_weight(mg_attn.linear_qkv.bias, linear_qkv_bias, 'linear_qkv.bias')
else:
mg_attn_bias, _ = self._get_weight(None if mg_attn is None else mg_attn.linear_qkv.bias.data,
'linear_qkv.bias')
if mg_attn_bias is not None:
hf_state_dict['q_proj.bias'] = mg_attn_bias[:q_dim].clone()
hf_state_dict['k_proj.bias'] = mg_attn_bias[q_dim:-kv_dim].clone()
hf_state_dict['v_proj.bias'] = mg_attn_bias[-kv_dim:].clone()
self._set_state_dict(mg_attn, 'q_layernorm.weight', hf_state_dict, 'q_norm.weight', to_mcore)
self._set_state_dict(mg_attn, 'k_layernorm.weight', hf_state_dict, 'k_norm.weight', to_mcore)
if to_mcore:
hf_state_dict = {}
else:
hf_state_dict = self._add_prefix(hf_state_dict, hf_prefix)
return hf_state_dict
def _set_qk_layernorm(self, mg_attn, hf_state_dict, to_mcore, **kwargs):
self._set_state_dict(mg_attn, 'q_norm.weight', hf_state_dict, 'q_norm.weight', to_mcore)
self._set_state_dict(mg_attn, 'k_norm.weight', hf_state_dict, 'k_norm.weight', to_mcore)


class OlMoELoader(ModelLoader):
class OLMoELoader(ModelLoader):

def get_transformer_layer_spec(self, vp_stage: Optional[int] = None):
return get_olmoe_decoder_block_spec(self.config, vp_stage)
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 = OLMoESelfAttention
return transformer_layer_spec


register_model(ModelMeta(
ModelType.olmoe,
['olmoe'],
bridge_cls=OLMoEBridge,
loader=OlMoELoader,
loader=OLMoELoader,
))
Loading