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
20 changes: 20 additions & 0 deletions src/megatron/bridge/models/kimi/__init__.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,20 @@
# Copyright (c) 2025, NVIDIA CORPORATION. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.

from megatron.bridge.models.kimi.kimi_provider import KimiK2Provider


__all__ = [
"KimiK2Provider",
]
129 changes: 129 additions & 0 deletions src/megatron/bridge/models/kimi/kimi_provider.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,129 @@
# Copyright (c) 2025, NVIDIA CORPORATION. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
from dataclasses import dataclass, field
from functools import partial
from typing import TYPE_CHECKING, Callable, List, Optional, Union

import torch
import torch.nn.functional as F
from megatron.core.models.gpt.gpt_layer_specs import get_gpt_decoder_block_spec

from megatron.bridge.models.gpt_provider import GPTModelProvider
from megatron.bridge.models.transformer_config import MLATransformerConfig


try:
import transformer_engine # type: ignore # noqa: F401

HAVE_TE = True
except (ImportError, ModuleNotFoundError):
HAVE_TE = False

if TYPE_CHECKING:
from megatron.core.transformer import ModuleSpec

if HAVE_TE:
from megatron.core.utils import is_te_min_version


@dataclass
class KimiK2Provider(MLATransformerConfig, GPTModelProvider):
"""
https://moonshotai.github.io/Kimi-K2/
"""

transformer_layer_spec: Union["ModuleSpec", Callable[["GPTModelProvider"], "ModuleSpec"]] = partial(
get_gpt_decoder_block_spec, use_transformer_engine=HAVE_TE
)

# Model
num_layers: int = 61
hidden_size: int = 7168
ffn_hidden_size: int = 18432
num_moe_experts: int = 384
moe_ffn_hidden_size: int = 2048
moe_shared_expert_intermediate_size: int = 2048 # 2048 * 1 shared expert
moe_layer_freq: Union[int, List[int]] = field(default_factory=lambda: [0] + [1] * 60) # first layer are dense
normalization: str = "RMSNorm"
activation_func: Callable = F.silu
gated_linear_unit: bool = True # swiglu
position_embedding_type: str = "rope"
add_bias_linear: bool = False
share_embeddings_and_output_weights: bool = False
num_attention_heads: int = 64
kv_channels: int = 64
max_position_embeddings: int = 4096
seq_length: int = 4096
rotary_base: float = 50000.0
make_vocab_size_divisible_by: int = 1280
mtp_num_layers: Optional[int] = None
mtp_loss_scaling_factor: Optional[float] = None

# Regularization
attention_dropout: float = 0.0
hidden_dropout: float = 0.0
qk_layernorm: bool = True

# MoE
moe_router_topk: int = 8
moe_router_num_groups: int = 1
moe_router_group_topk: int = 1
moe_router_topk_scaling_factor: float = 2.827
moe_aux_loss_coeff: float = 1e-3
moe_router_score_function: str = "sigmoid"
moe_router_enable_expert_bias: bool = True
moe_router_bias_update_rate: float = 1e-3
moe_grouped_gemm: bool = True
moe_router_pre_softmax: bool = True
moe_token_dispatcher_type: str = "alltoall"
moe_router_load_balancing_type: str = "seq_aux_loss"
moe_shared_expert_overlap: bool = True
moe_router_dtype: Optional[str] = "fp32"

# MLA
multi_latent_attention: bool = True
q_lora_rank: int = 1536
kv_lora_rank: int = 512
qk_head_dim: int = 128
qk_pos_emb_head_dim: int = 64
v_head_dim: int = 128
rotary_scaling_factor: float = 32
beta_fast: float = 1.0
beta_slow: float = 1.0
mscale: float = 1.0
mscale_all_dim: float = 1.0

# Miscellaneous
init_method_std: float = 0.006
layernorm_epsilon: float = 1e-6
bf16: bool = True
params_dtype: torch.dtype = torch.bfloat16
async_tensor_model_parallel_allreduce: bool = True
attention_softmax_in_fp32: bool = True
persist_layer_norm: bool = True
num_layers_in_first_pipeline_stage: Optional[int] = None
num_layers_in_last_pipeline_stage: Optional[int] = None
account_for_embedding_in_pipeline_split: bool = False
account_for_loss_in_pipeline_split: bool = False
vocab_size: int = 163840

# fusions
apply_rope_fusion: bool = False
bias_activation_fusion: bool = True
bias_dropout_fusion: bool = True
masked_softmax_fusion: bool = True
gradient_accumulation_fusion: bool = True
cross_entropy_loss_fusion: bool = True
cross_entropy_fusion_impl: str = "te"
moe_permute_fusion: bool = is_te_min_version("2.1.0") if HAVE_TE else False
7 changes: 7 additions & 0 deletions src/megatron/bridge/models/kimi_vl/__init__.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,7 @@
from megatron.bridge.models.kimi_vl.modeling_kimi_k25_vl import KimiK25VLModel
from megatron.bridge.models.kimi_vl.kimi_k25_vl_bridge import KimiK25VLBridge

__all__ = [
"KimiK25VLModel",
"KimiK25VLBridge",
]
97 changes: 97 additions & 0 deletions src/megatron/bridge/models/kimi_vl/kimi_k25_vl_bridge.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,97 @@
import math

import torch
from transformers import Gemma3ForConditionalGeneration

from megatron.bridge.models.conversion.mapping_registry import MegatronMappingRegistry
from megatron.bridge.models.conversion.model_bridge import MegatronModelBridge

from megatron.bridge.models.conversion.param_mapping import (
AutoMapping,
GatedMLPMapping,
QKVMapping,
ReplicatedMapping,
)
from megatron.bridge.models.hf_pretrained.vlm import PreTrainedVLM
from megatron.bridge.models.kimi_vl.kimi_k25_vl_provider import KimiK25VLModelProvider
from megatron.bridge.models.kimi_vl.modeling_kimi_k25_vl import KimiK25VLModel
from megatron.bridge.models.deepseek.common import get_common_configs, get_common_mapping_list

@MegatronModelBridge.register_bridge(source="KimiK25ForConditionalGeneration", target=KimiK25VLModel)
class KimiK25VLBridge(MegatronModelBridge):
"""
Megatron Bridge for Kimi K2.5 VL.
"""

def provider_bridge(self, hf_pretrained: PreTrainedVLM) -> KimiK25VLModelProvider:
hf_config = hf_pretrained.config
text_config = hf_config.text_config
vision_config = hf_config.vision_config

# get_common_configs expects TextConfig
hf_pretrained.config = text_config
configs = get_common_configs(hf_pretrained)

configs["make_vocab_size_divisible_by"] = 1280
configs["moe_router_score_function"] = "sigmoid"
configs["moe_router_enable_expert_bias"] = True
# aux_loss_alpha is not set in all DSv3 HF configs
if hasattr(hf_config, "aux_loss_alpha"):
configs["moe_aux_loss_coeff"] = hf_config.aux_loss_alpha

provider = KimiK25VLModelProvider(
# Text configuration
**configs,
# Vision configuration
vision_config=vision_config,
# VL-specific token IDs
bos_token_id=text_config.bos_token_id,
eos_token_id=text_config.eos_token_id,
media_placeholder_token_id=hf_config.media_placeholder_token_id,
# Precision configuration
fp16=(self.dtype_from_hf(hf_config, default=torch.float32) == torch.float16),
bf16=(self.dtype_from_hf(hf_config, default=torch.float32) == torch.bfloat16),
params_dtype=self.dtype_from_hf(hf_config, default=torch.float32),
# misc
hf_model_path=hf_pretrained._model_name_or_path,
)

return provider

def mapping_registry(self) -> MegatronMappingRegistry:
# Return MegatronMappingRegistry containing parameter mappings from Megatron to HF format
# First create simple 1:1 parameter mappings using a dictionary for readability
mapping_list = get_common_mapping_list()
param_mappings = {
# expert bias
"decoder.layers.*.mlp.router.expert_bias": "model.layers.*.mlp.gate.e_score_correction_bias",
}

for megatron_param, hf_param in param_mappings.items():
mapping_list.append(AutoMapping(megatron_param=megatron_param, hf_param=hf_param))

for mapping in mapping_list:
# in HF Kimi K2.5 VL models, language component is prefixed with "language_model.model" instead of "model"
if isinstance(mapping, AutoMapping):
mapping.hf_param = "language_model." + mapping.hf_param
mapping.megatron_param = "language_model." + mapping.megatron_param
elif isinstance(mapping, GatedMLPMapping):
mapping.megatron_param = mapping.megatron_param.replace("decoder", "language_model.decoder")
mapping.hf_param["gate"] = mapping.hf_param["gate"].replace("model", "language_model.model")
mapping.hf_param["up"] = mapping.hf_param["up"].replace("model", "language_model.model")


# Add Vision and MM Projector mappings
mapping_list.extend(
[
ReplicatedMapping(
megatron_param="vision_tower.**",
hf_param="vision_tower.**",
),
ReplicatedMapping(
megatron_param="mm_projector.**",
hf_param="mm_projector.**",
),
]
)
return MegatronMappingRegistry(*mapping_list)
88 changes: 88 additions & 0 deletions src/megatron/bridge/models/kimi_vl/kimi_k25_vl_provider.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,88 @@
# Copyright (c) 2025, NVIDIA CORPORATION. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.


from dataclasses import dataclass
from typing import Optional, Any

from megatron.core.models.gpt import GPTModel

from megatron.bridge.models.kimi.kimi_provider import KimiK2Provider
from megatron.bridge.models.kimi_vl.modeling_kimi_k25_vl import KimiK25VLModel
from megatron.bridge.models.kimi_vl.modelling_kimi_vl.transfomer_config import DeepseekV3Config
from transformers.dynamic_module_utils import get_class_from_dynamic_module



@dataclass
class KimiK25VLModelProvider(KimiK2Provider):
""" """

hf_model_path: Optional[str] = None
vision_config: Any = None

hf_text_config: Optional[DeepseekV3Config] = None
pretrained_model_name: str = "moonshotai/Kimi-K2.5"

bos_token_id: int = 163584
eos_token_id: int = 163585
media_placeholder_token_id: int = 163605
freeze_language_model: bool = False
# Whether to freeze vision encoder weights
freeze_vision_model: bool = True
# Whether to freeze vision-to-language projection weights
freeze_vision_projection: bool = True
scatter_embedding_sequence_parallel: bool = False

variable_seq_lengths: bool = True
moe_token_dispatcher_type: str = "alltoall"

def finalize(self) -> None:
if self.tensor_model_parallel_size > 1:
self.sequence_parallel = True

super().finalize()

def provide(self, pre_process=None, post_process=None, vp_stage=None):
model = KimiK25VLModel(
self,
pre_process=pre_process,
post_process=post_process,
vp_stage=vp_stage,
)

# Apply freeze options if any are enabled for fine-tuning
if self.freeze_language_model or self.freeze_vision_model or self.freeze_vision_projection:
model.freeze(
freeze_language_model=self.freeze_language_model,
freeze_vision_model=self.freeze_vision_model,
freeze_vision_projection=self.freeze_vision_projection,
)

return model

def provide_language_model(self, pre_process=None, post_process=None, vp_stage=None) -> GPTModel:
"""
Provide just the language model component without vision.

Args:
pre_process: Whether this is the first stage in pipeline parallelism
post_process: Whether this is the last stage in pipeline parallelism
vp_stage: Virtual pipeline stage number

Returns:
GPTModel instance (language model only)
"""
# Use parent class to create standard language model
return super().provide(pre_process=pre_process, post_process=post_process, vp_stage=vp_stage)
Loading