Skip to content
Draft
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
13 changes: 9 additions & 4 deletions src/mcore_bridge/bridge/gpt_bridge.py
Original file line number Diff line number Diff line change
Expand Up @@ -51,6 +51,10 @@ class GPTBridge:
additional_dim1_keys = set()
_support_hf_grouped_lora = True

@property
def use_transformer_engine(self):
return self.config.transformer_impl == 'transformer_engine'

def __init__(self, config: ModelConfig):
self.config = config
self._disable_tqdm = False
Expand Down Expand Up @@ -1638,8 +1642,9 @@ def _set_layer_attn(self, mg_layer, hf_state_dict, layer_idx: int, to_mcore: boo
else:
hf_state_dict.update(
self._set_attn_state(mg_attn, hf_state_dict, f'{self.hf_attn_prefix}.', layer_idx, to_mcore))
self._set_state_dict(mg_layer, 'self_attention.linear_qkv.layer_norm_weight', hf_state_dict,
self.hf_input_layernorm_key, to_mcore)
mg_key = ('self_attention.linear_qkv.layer_norm_weight'
if self.use_transformer_engine else 'input_layernorm.weight')
self._set_state_dict(mg_layer, mg_key, hf_state_dict, self.hf_input_layernorm_key, to_mcore)
return hf_state_dict

def _set_layer_mlp(self, mg_layer, hf_state_dict, layer_idx: int, to_mcore: bool, is_mtp: bool = False):
Expand All @@ -1658,8 +1663,8 @@ def _set_layer_mlp(self, mg_layer, hf_state_dict, layer_idx: int, to_mcore: bool
else:
hf_state_dict.update(
self._set_mlp_state(mg_mlp, hf_state_dict, f'{self.hf_mlp_prefix}.', layer_idx, to_mcore))
self._set_state_dict(mg_layer, 'mlp.linear_fc1.layer_norm_weight', hf_state_dict,
self.hf_post_attention_layernorm_key, to_mcore)
mg_key = 'mlp.linear_fc1.layer_norm_weight' if self.use_transformer_engine else 'pre_mlp_layernorm.weight'
self._set_state_dict(mg_layer, mg_key, hf_state_dict, self.hf_post_attention_layernorm_key, to_mcore)
return hf_state_dict

def _set_hyper_connection(self, mg_layer, hf_state_dict, layer_idx, to_mcore):
Expand Down
14 changes: 10 additions & 4 deletions src/mcore_bridge/model/mm_gpts/gemma4.py
Original file line number Diff line number Diff line change
Expand Up @@ -468,8 +468,8 @@ def _set_router(self, mg_mlp, hf_state_dict, to_mcore, **kwargs):
def _set_layer_mlp(self, mg_layer, hf_state_dict, layer_idx: int, to_mcore: bool, is_mtp: bool = False):
mg_mlp = None if mg_layer is None else mg_layer.mlp
hf_state_dict.update(self._set_mlp_state(mg_mlp, hf_state_dict, f'{self.hf_mlp_prefix}.', layer_idx, to_mcore))
self._set_state_dict(mg_layer, 'mlp.linear_fc1.layer_norm_weight', hf_state_dict,
'pre_feedforward_layernorm.weight', to_mcore)
mg_key = 'mlp.linear_fc1.layer_norm_weight' if self.use_transformer_engine else 'pre_mlp_layernorm.weight'
self._set_state_dict(mg_layer, mg_key, hf_state_dict, 'pre_feedforward_layernorm.weight', to_mcore)
if self.text_config.enable_moe_block:
mg_experts = None if mg_layer is None else mg_layer.experts_mlp
hf_state_dict.update(self._set_moe_state(mg_experts, hf_state_dict, '', layer_idx, to_mcore, is_mtp=is_mtp))
Expand Down Expand Up @@ -840,14 +840,20 @@ def get_transformer_layer_spec(self, vp_stage: Optional[int] = None):
num_moe_experts = self.config.num_moe_experts
self.config.num_moe_experts = None
layer_specs = get_gpt_decoder_block_spec(
self.config, use_transformer_engine=True, normalization=self.config.normalization, vp_stage=vp_stage)
self.config,
use_transformer_engine=self.use_transformer_engine,
normalization=self.config.normalization,
vp_stage=vp_stage)
for layer_spec in layer_specs.layer_specs:
layer_spec.submodules.self_attention.module = Gemma4SelfAttention
self._set_mlp_spec(layer_spec.submodules, Gemma4MLP)
if num_moe_experts is not None:
self.config.num_moe_experts = num_moe_experts
moe_layer_specs = get_gpt_decoder_block_spec(
self.config, use_transformer_engine=True, normalization=self.config.normalization, vp_stage=vp_stage)
self.config,
use_transformer_engine=self.use_transformer_engine,
normalization=self.config.normalization,
vp_stage=vp_stage)
for layer_spec, moe_layer_spec in zip(layer_specs.layer_specs, moe_layer_specs.layer_specs):
layer_spec.submodules.experts_mlp = moe_layer_spec.submodules.mlp
self._set_mlp_spec(layer_spec.submodules, Gemma4MoELayer, mlp_key='experts_mlp')
Expand Down
17 changes: 15 additions & 2 deletions src/mcore_bridge/model/register.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,11 +9,13 @@
from megatron.core.extensions.transformer_engine import TEGroupedLinear, TELayerNormColumnParallelLinear, TELinear
from megatron.core.models.gpt import gpt_model
from megatron.core.models.gpt.gpt_layer_specs import get_gpt_decoder_block_spec, get_gpt_mtp_block_spec
from megatron.core.transformer.dot_product_attention import DotProductAttention
from megatron.core.transformer.moe.router import TopKRouter as McoreTopKRouter
from megatron.core.transformer.multi_latent_attention import MLASelfAttention as McoreMLASelfAttention
from megatron.core.transformer.transformer_layer import TransformerLayer as McoreTransformerLayer
from packaging import version
from torch import nn
from transformers.utils import is_torch_npu_available
from typing import TYPE_CHECKING, List, Optional, Type, Union

from mcore_bridge.bridge import GPTBridge
Expand Down Expand Up @@ -82,6 +84,10 @@ def __init__(self, config: ModelConfig):
if self.model_cls is None:
self.model_cls = MultimodalGPTModel if config.is_multimodal else GPTModel

@property
def use_transformer_engine(self):
return self.config.transformer_impl == 'transformer_engine'

def _set_mlp_spec(self, layer_submodules, mlp_module, mlp_key='mlp'):
mlp_spec = getattr(layer_submodules, mlp_key)
if isinstance(mlp_spec, partial):
Expand Down Expand Up @@ -121,23 +127,30 @@ 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 _replace_unfused_attention(self, transformer_layer_spec):
if not is_torch_npu_available() or self.config.attention_backend.name != 'unfused':
return
for layer_spec in transformer_layer_spec.layer_specs:
layer_spec.submodules.self_attention.submodules.core_attention = DotProductAttention

def get_transformer_layer_spec(self, vp_stage: Optional[int] = None):
with self._patch_experimental_attention_variant():
transformer_layer_spec = get_gpt_decoder_block_spec(
self.config,
use_transformer_engine=True,
use_transformer_engine=self.use_transformer_engine,
normalization=self.config.normalization,
qk_l2_norm=self.config.qk_l2_norm,
vp_stage=vp_stage)
self._deepcopy_layer_spec(transformer_layer_spec)
self._replace_unfused_attention(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)
return transformer_layer_spec

def get_mtp_block_spec(self, transformer_layer_spec, vp_stage: Optional[int] = None):
mtp_block_spec = get_gpt_mtp_block_spec(
self.config, transformer_layer_spec, use_transformer_engine=True, vp_stage=vp_stage)
self.config, transformer_layer_spec, use_transformer_engine=self.use_transformer_engine, vp_stage=vp_stage)
if mtp_block_spec is not None:
for layer_spec in mtp_block_spec.layer_specs:
layer_spec.module = MultiTokenPredictionLayer
Expand Down
63 changes: 43 additions & 20 deletions src/mcore_bridge/tuners/lora.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,7 @@
TERowParallelGroupedLinear, TERowParallelLinear)
from megatron.core.models.common.embeddings.language_model_embedding import LanguageModelEmbedding
from megatron.core.parallel_state import get_expert_tensor_parallel_world_size, get_tensor_model_parallel_world_size
from megatron.core.tensor_parallel.layers import ColumnParallelLinear, RowParallelLinear
from megatron.core.tensor_parallel.random import get_cuda_rng_tracker, get_expert_parallel_rng_tracker_name
from megatron.core.transformer.mlp import apply_swiglu_sharded_factory
from megatron.core.transformer.module import MegatronModule
Expand Down Expand Up @@ -117,6 +118,8 @@ def __init__(
raise ValueError(f'{self.__class__.__name__} does not support DoRA yet, please set it to False')

self.is_parallel_a = isinstance(base_layer, (TERowParallelLinear, TERowParallelGroupedLinear))
self.is_local_linear = isinstance(base_layer, (ColumnParallelLinear, RowParallelLinear))
self.is_parallel_a = self.is_parallel_a or isinstance(base_layer, RowParallelLinear)
self.is_grouped = isinstance(base_layer, TEGroupedLinear)
self.fan_in_fan_out = fan_in_fan_out
self._active_adapter = adapter_name
Expand Down Expand Up @@ -180,7 +183,7 @@ def update_layer(self, adapter_name, r, *, lora_alpha, **kwargs):
lora_a = _build_local_te_linear(router_shape[1], r, lora_bias, **kwargs)
lora_b = _build_local_te_linear(r, router_shape[0], lora_bias, **kwargs)
elif self.is_parallel_a:
in_features = self.in_features * self.tp_size
in_features = self.in_features if self.is_local_linear else self.in_features * self.tp_size
if self.is_grouped:
if is_torch_npu_available():
lora_a = NpuGroupedLoraLinear(
Expand Down Expand Up @@ -214,15 +217,25 @@ def update_layer(self, adapter_name, r, *, lora_alpha, **kwargs):
**kwargs,
)
else:
lora_a = TERowParallelLinear(
input_size=in_features,
output_size=r,
bias=False,
input_is_parallel=True,
**kwargs,
)
lora_b = _build_local_te_linear(r, self.out_features, lora_bias, **kwargs)
lora_a.parallel_mode = self.base_layer.parallel_mode # fix moe_shared_expert_overlap
if self.is_local_linear:
lora_a = RowParallelLinear(
input_size=in_features,
output_size=r,
bias=False,
input_is_parallel=True,
**kwargs,
)
lora_b = nn.Linear(r, self.out_features, bias=lora_bias)
else:
lora_a = TERowParallelLinear(
input_size=in_features,
output_size=r,
bias=False,
input_is_parallel=True,
**kwargs,
)
lora_b = _build_local_te_linear(r, self.out_features, lora_bias, **kwargs)
lora_a.parallel_mode = self.base_layer.parallel_mode # fix moe_shared_expert_overlap
else:
if is_torch_npu_available():
out_features = self.out_features
Expand Down Expand Up @@ -260,15 +273,25 @@ def update_layer(self, adapter_name, r, *, lora_alpha, **kwargs):
**kwargs,
)
else:
lora_a = _build_local_te_linear(self.in_features, r, lora_bias, **kwargs)
lora_b = TEColumnParallelLinear(
input_size=r,
output_size=out_features,
bias=lora_bias,
gather_output=False,
**kwargs,
)
lora_b.parallel_mode = self.base_layer.parallel_mode # fix moe_shared_expert_overlap
if self.is_local_linear:
lora_a = nn.Linear(self.in_features, r, bias=lora_bias)
lora_b = ColumnParallelLinear(
input_size=r,
output_size=out_features,
bias=lora_bias,
gather_output=False,
**kwargs,
)
else:
lora_a = _build_local_te_linear(self.in_features, r, lora_bias, **kwargs)
lora_b = TEColumnParallelLinear(
input_size=r,
output_size=out_features,
bias=lora_bias,
gather_output=False,
**kwargs,
)
lora_b.parallel_mode = self.base_layer.parallel_mode # fix moe_shared_expert_overlap
for lora in [lora_a, lora_b]:
# When parallel_mode is set to None by moe_shared_expert_overlap,
# disable UB comm overlap; the corresponding collectives are driven
Expand Down Expand Up @@ -412,7 +435,7 @@ def forward(self, x: torch.Tensor, *args: Any, **kwargs: Any):
f'Got base_layer type: {type(self.base_layer)}. ')
else:
(result, x), bias = self.base_layer(x, *args, **kwargs)
elif isinstance(self.base_layer, (TELinear, TEGroupedLinear)):
elif isinstance(self.base_layer, (TELinear, TEGroupedLinear, ColumnParallelLinear, RowParallelLinear)):
result, bias = self.base_layer(x, *args, **kwargs)
elif isinstance(self.base_layer, TopKRouter):
with self._patch_router_gating():
Expand Down
4 changes: 3 additions & 1 deletion src/mcore_bridge/tuners/patcher.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
# Copyright (c) ModelScope Contributors. All rights reserved.
from megatron.core.extensions.transformer_engine import TEGroupedLinear, TELayerNormColumnParallelLinear, TELinear
from megatron.core.tensor_parallel.layers import ColumnParallelLinear, RowParallelLinear
from megatron.core.transformer.module import MegatronModule
from megatron.core.transformer.moe.router import TopKRouter
from peft import LoraModel
Expand Down Expand Up @@ -27,7 +28,8 @@ def dispatch_megatron(
else:
target_base_layer = target

linear_cls = (TELayerNormColumnParallelLinear, TELinear, TEGroupedLinear, TopKRouter)
linear_cls = (TELayerNormColumnParallelLinear, TELinear, TEGroupedLinear, ColumnParallelLinear, RowParallelLinear,
TopKRouter)
if isinstance(target_base_layer, linear_cls):
new_module = LoraParallelLinear(base_layer=target, adapter_name=adapter_name, **kwargs)

Expand Down
96 changes: 96 additions & 0 deletions tests/test_model_register.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,96 @@
import unittest
from types import SimpleNamespace
from unittest.mock import patch

from mcore_bridge.bridge.gpt_bridge import GPTBridge
from mcore_bridge.model.register import DotProductAttention, ModelLoader


class TestModelLoader(unittest.TestCase):

@staticmethod
def _get_loader(transformer_impl):
loader = ModelLoader.__new__(ModelLoader)
loader.config = SimpleNamespace(
transformer_impl=transformer_impl,
experimental_attention_variant=None,
normalization='RMSNorm',
qk_l2_norm=False,
attention_backend=SimpleNamespace(name='flash'),
)
return loader

def test_transformer_impl_controls_decoder_layer_spec(self):
for transformer_impl, expected in [('local', False), ('transformer_engine', True)]:
with self.subTest(transformer_impl=transformer_impl):
loader = self._get_loader(transformer_impl)
with patch('mcore_bridge.model.register.get_gpt_decoder_block_spec') as get_spec:
get_spec.return_value = SimpleNamespace(layer_specs=[])
loader.get_transformer_layer_spec()
self.assertEqual(get_spec.call_args.kwargs['use_transformer_engine'], expected)

@patch('mcore_bridge.model.register.is_torch_npu_available', return_value=True)
def test_unfused_attention_uses_mcore_dot_product_on_npu(self, _):
loader = self._get_loader('transformer_engine')
loader.config.attention_backend.name = 'unfused'
core_attention = object()
transformer_layer_spec = SimpleNamespace(layer_specs=[
SimpleNamespace(
submodules=SimpleNamespace(
self_attention=SimpleNamespace(submodules=SimpleNamespace(core_attention=core_attention))))
])

loader._replace_unfused_attention(transformer_layer_spec)

self.assertIs(transformer_layer_spec.layer_specs[0].submodules.self_attention.submodules.core_attention,
DotProductAttention)

def test_transformer_impl_controls_mtp_layer_spec(self):
for transformer_impl, expected in [('local', False), ('transformer_engine', True)]:
with self.subTest(transformer_impl=transformer_impl):
loader = self._get_loader(transformer_impl)
with patch('mcore_bridge.model.register.get_gpt_mtp_block_spec') as get_spec:
get_spec.return_value = None
loader.get_mtp_block_spec(SimpleNamespace())
self.assertEqual(get_spec.call_args.kwargs['use_transformer_engine'], expected)


class TestGPTBridge(unittest.TestCase):

@staticmethod
def _get_bridge(transformer_impl):
bridge = GPTBridge.__new__(GPTBridge)
bridge.config = SimpleNamespace(transformer_impl=transformer_impl, multi_latent_attention=False)
return bridge

def test_transformer_impl_controls_attention_layernorm_mapping(self):
expected_keys = {
'local': 'input_layernorm.weight',
'transformer_engine': 'self_attention.linear_qkv.layer_norm_weight',
}
layer = SimpleNamespace(self_attention=SimpleNamespace())
for transformer_impl, expected_key in expected_keys.items():
with self.subTest(transformer_impl=transformer_impl):
bridge = self._get_bridge(transformer_impl)
with patch.object(bridge, '_set_attn_state', return_value={}), \
patch.object(bridge, '_set_state_dict') as set_state_dict:
bridge._set_layer_attn(layer, {}, 0, True)
self.assertEqual(set_state_dict.call_args.args[1], expected_key)

def test_transformer_impl_controls_mlp_layernorm_mapping(self):
expected_keys = {
'local': 'pre_mlp_layernorm.weight',
'transformer_engine': 'mlp.linear_fc1.layer_norm_weight',
}
layer = SimpleNamespace(mlp=SimpleNamespace())
for transformer_impl, expected_key in expected_keys.items():
with self.subTest(transformer_impl=transformer_impl):
bridge = self._get_bridge(transformer_impl)
with patch.object(bridge, '_set_mlp_state', return_value={}), \
patch.object(bridge, '_set_state_dict') as set_state_dict:
bridge._set_layer_mlp(layer, {}, 0, True)
self.assertEqual(set_state_dict.call_args.args[1], expected_key)


if __name__ == '__main__':
unittest.main()
22 changes: 22 additions & 0 deletions tests/test_tuners.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,22 @@
import unittest
from unittest.mock import patch

from megatron.core.tensor_parallel.layers import ColumnParallelLinear, RowParallelLinear

from mcore_bridge.tuners.patcher import dispatch_megatron


class TestMegatronLoraDispatch(unittest.TestCase):

def test_dispatches_local_parallel_linears(self):
for linear_cls in (ColumnParallelLinear, RowParallelLinear):
with self.subTest(linear_cls=linear_cls.__name__):
target = linear_cls.__new__(linear_cls)
with patch('mcore_bridge.tuners.patcher.LoraParallelLinear', return_value='adapter') as lora_cls:
result = dispatch_megatron(target, 'default')
self.assertEqual(result, 'adapter')
self.assertIs(lora_cls.call_args.kwargs['base_layer'], target)


if __name__ == '__main__':
unittest.main()