Skip to content
Open
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
10 changes: 6 additions & 4 deletions src/mcore_bridge/bridge/gpt_bridge.py
Original file line number Diff line number Diff line change
Expand Up @@ -1630,6 +1630,8 @@ def _set_mla_attn_state(

def _set_layer_attn(self, mg_layer, hf_state_dict, layer_idx: int, to_mcore: bool):
mg_attn = None if mg_layer is None else mg_layer.self_attention
attn_ln_name = ('input_layernorm.weight'
if self.config.use_accuracy_compatible else 'self_attention.linear_qkv.layer_norm_weight')
if self.config.multi_latent_attention:
hf_state_dict.update(
self._set_mla_attn_state(mg_attn, hf_state_dict, f'{self.hf_attn_prefix}.', layer_idx, to_mcore))
Expand All @@ -1638,13 +1640,14 @@ 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)
self._set_state_dict(mg_layer, attn_ln_name, 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):
mg_mlp = None if mg_layer is None else mg_layer.mlp
is_moe = True if hasattr(mg_mlp, 'experts') else False
mlp_ln_name = ('pre_mlp_layernorm.weight'
if self.config.use_accuracy_compatible else 'mlp.linear_fc1.layer_norm_weight')
if not to_mcore:
is_moe = torch.tensor([is_moe], dtype=torch.bool, device='cuda')
if self.pp_size > 1:
Expand All @@ -1658,8 +1661,7 @@ 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)
self._set_state_dict(mg_layer, mlp_ln_name, 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
4 changes: 4 additions & 0 deletions src/mcore_bridge/config/model_config.py
Original file line number Diff line number Diff line change
Expand Up @@ -114,6 +114,8 @@ def tuple_type(x):

@dataclass
class ModelConfig(TransformerConfig):
# Use native Megatron modules and accuracy-compatible kernels for alignment.
use_accuracy_compatible: bool = False
mcore_model_type: Optional[str] = None # Inferred from hf_model_type by default
hf_model_type: Optional[str] = None
llm_model_type: Optional[str] = None
Expand Down Expand Up @@ -291,6 +293,8 @@ def __post_init__(self):
from mcore_bridge.model import get_mcore_model_type, get_model_meta
self._augment_mindspeed_defaults()
self._format_config()
if self.use_accuracy_compatible:
self.persist_layer_norm = False
if self.experimental_attention_variant is not None:
require_version('megatron-core>=0.16.0.dev',
'experimental attention variant requires megatron-core>=0.16.0')
Expand Down
8 changes: 6 additions & 2 deletions src/mcore_bridge/model/gpt_model.py
Original file line number Diff line number Diff line change
Expand Up @@ -88,6 +88,9 @@ def __init__(
share_embeddings_and_output_weights=not config.untie_embeddings_and_output_weights,
position_embedding_type=config.position_embedding_type,
rotary_base=config.rotary_base,
**({
'rotary_percent': config.partial_rotary_factor or 1.0
} if config.use_accuracy_compatible else {}),
mtp_block_spec=mtp_block_spec,
vp_stage=vp_stage,
)
Expand Down Expand Up @@ -205,8 +208,9 @@ def _preprocess(
def _set_inv_freq(self):
if getattr(self, 'rotary_pos_emb', None) is None:
return
new_inv_freq, self.config.attention_scaling = get_rope_inv_freq(self.config)
self.rotary_pos_emb.inv_freq = new_inv_freq.to(self.rotary_pos_emb.inv_freq.device)
if not self.config.use_accuracy_compatible:
new_inv_freq, self.config.attention_scaling = get_rope_inv_freq(self.config)
self.rotary_pos_emb.inv_freq = new_inv_freq.to(self.rotary_pos_emb.inv_freq.device)

def _get_rotary_pos_emb(self, decoder_input, position_ids, packed_seq_params, inference_context=None):
# Rotary positional embeddings (embedding is None for PP intermediate devices)
Expand Down
33 changes: 25 additions & 8 deletions src/mcore_bridge/model/register.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,7 @@
from megatron.core import mpu
from megatron.core.enums import ModelType
from megatron.core.extensions.transformer_engine import TEGroupedLinear, TELayerNormColumnParallelLinear, TELinear
from megatron.core.models.gpt import gpt_layer_specs as _gpt_layer_specs
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.moe.router import TopKRouter as McoreTopKRouter
Expand All @@ -27,6 +28,8 @@
from .gpt_model import GPTModel
from .mm_gpt_model import MultimodalGPTModel

from megatron.core.transformer.torch_norm import WrappedTorchNorm

MODEL_MAPPING = {}
logger = get_logger()

Expand Down Expand Up @@ -123,21 +126,35 @@ def _deepcopy_layer_spec(self, transformer_layer_spec):

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,
normalization=self.config.normalization,
qk_l2_norm=self.config.qk_l2_norm,
vp_stage=vp_stage)
use_transformer_engine = not self.config.use_accuracy_compatible
original_ln_impl = _gpt_layer_specs.LNImpl
if not use_transformer_engine:
_gpt_layer_specs.LNImpl = WrappedTorchNorm
try:
transformer_layer_spec = get_gpt_decoder_block_spec(
self.config,
use_transformer_engine=use_transformer_engine,
normalization=self.config.normalization,
qk_l2_norm=self.config.qk_l2_norm,
vp_stage=vp_stage)
finally:
_gpt_layer_specs.LNImpl = original_ln_impl
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)
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)
use_transformer_engine = not self.config.use_accuracy_compatible
original_ln_impl = _gpt_layer_specs.LNImpl
if not use_transformer_engine:
_gpt_layer_specs.LNImpl = WrappedTorchNorm
try:
mtp_block_spec = get_gpt_mtp_block_spec(
self.config, transformer_layer_spec, use_transformer_engine=use_transformer_engine, vp_stage=vp_stage)
finally:
_gpt_layer_specs.LNImpl = original_ln_impl
if mtp_block_spec is not None:
for layer_spec in mtp_block_spec.layer_specs:
layer_spec.module = MultiTokenPredictionLayer
Expand Down