diff --git a/src/mcore_bridge/bridge/gpt_bridge.py b/src/mcore_bridge/bridge/gpt_bridge.py index 41caf7c..eda8814 100644 --- a/src/mcore_bridge/bridge/gpt_bridge.py +++ b/src/mcore_bridge/bridge/gpt_bridge.py @@ -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)) @@ -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: @@ -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): diff --git a/src/mcore_bridge/config/model_config.py b/src/mcore_bridge/config/model_config.py index a472003..adf5c8d 100644 --- a/src/mcore_bridge/config/model_config.py +++ b/src/mcore_bridge/config/model_config.py @@ -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 @@ -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') diff --git a/src/mcore_bridge/model/gpt_model.py b/src/mcore_bridge/model/gpt_model.py index 8ed9994..e3a1a92 100644 --- a/src/mcore_bridge/model/gpt_model.py +++ b/src/mcore_bridge/model/gpt_model.py @@ -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, ) @@ -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) diff --git a/src/mcore_bridge/model/register.py b/src/mcore_bridge/model/register.py index 847c9b0..7b9db92 100644 --- a/src/mcore_bridge/model/register.py +++ b/src/mcore_bridge/model/register.py @@ -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 @@ -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() @@ -123,12 +126,19 @@ 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: @@ -136,8 +146,15 @@ def get_transformer_layer_spec(self, vp_stage: Optional[int] = None): 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