bridge.models.kimi.kimi_k3_layers#
KDA, no-RoPE MLA, latent-MoE, and AttnRes layers for Kimi K3.
The initial implementation was adapted from the Apache-2.0 Miles Kimi K3 backend: https://github.com/radixark/miles/tree/dc62a0bd4b7af1c59ee2084852eb18b5585ec082/miles_plugins/models/kimi_k3
Module Contents#
Classes#
TP-sharded KDA short convolution with FP32 state. |
|
Select KDA or Kimi’s no-RoPE MLA according to the global layer number. |
|
MCore latent MoE with Kimi’s post-combine RMS normalization. |
|
Transformer layer implementing Kimi K3’s AttnRes residual bank. |
Functions#
API#
- bridge.models.kimi.kimi_k3_layers._mark_tp_replicated(
- module: torch.nn.Module,
- *,
- reduction: str = 'average',
- bridge.models.kimi.kimi_k3_layers._linear(module: torch.nn.Module, inputs: torch.Tensor) torch.Tensor#
- class bridge.models.kimi.kimi_k3_layers.KimiK3ShortConvolution(*args, tp_group, **kwargs)#
Bases:
fla.modules.ShortConvolutionTP-sharded KDA short convolution with FP32 state.
Initialization
- sharded_state_dict(
- prefix: str = '',
- sharded_offsets: tuple = (),
- metadata: dict | None = None,
Return the TP-sharded convolution state.
- class bridge.models.kimi.kimi_k3_layers.KimiK3Attention(
- config,
- layer_number: int,
- cp_comm_type: str | None = None,
- pg_collection=None,
- pp_layer_offset: int | None = None,
- name: str | None = None,
Bases:
megatron.core.transformer.module.MegatronModuleSelect KDA or Kimi’s no-RoPE MLA according to the global layer number.
Initialization
- _duplicated_linear(
- input_size: int,
- output_size: int,
- _column_linear(
- input_size: int,
- output_size: int,
- _row_linear(
- input_size: int,
- output_size: int,
- _init_kda(config) None#
- _init_mla(config) None#
- sharded_state_dict(
- prefix: str = '',
- sharded_offsets: tuple = (),
- metadata: dict | None = None,
Return attention state with explicit KDA TP sharding.
- _forward_kda(
- hidden_states: torch.Tensor,
- packed_seq_params: megatron.core.packed_seq_params.PackedSeqParams | None,
- _forward_mla(
- hidden_states: torch.Tensor,
- packed_seq_params: megatron.core.packed_seq_params.PackedSeqParams | None,
- forward(
- hidden_states: torch.Tensor,
- attention_mask: torch.Tensor | None = None,
- key_value_states: torch.Tensor | None = None,
- inference_context: megatron.core.inference.contexts.BaseInferenceContext | None = None,
- rotary_pos_emb: torch.Tensor | None = None,
- rotary_pos_cos: torch.Tensor | None = None,
- rotary_pos_sin: torch.Tensor | None = None,
- rotary_pos_cos_sin: torch.Tensor | None = None,
- attention_bias: torch.Tensor | None = None,
- packed_seq_params: megatron.core.packed_seq_params.PackedSeqParams | None = None,
- sequence_len_offset: int | None = None,
- **kwargs,
Run the attention implementation selected for this layer.
- class bridge.models.kimi.kimi_k3_layers.KimiK3MoELayer(*args, **kwargs)#
Bases:
megatron.core.transformer.moe.moe_layer.MoELayerMCore latent MoE with Kimi’s post-combine RMS normalization.
Initialization
- postprocess(
- output: torch.Tensor,
- shared_expert_output: torch.Tensor | None,
Normalize the combined routed-expert output before projecting up.
- class bridge.models.kimi.kimi_k3_layers.KimiK3TransformerLayer(*args, **kwargs)#
Bases:
megatron.core.transformer.transformer_layer.TransformerLayerTransformer layer implementing Kimi K3’s AttnRes residual bank.
Initialization
- static _add_bias(
- output_with_bias: tuple[torch.Tensor, torch.Tensor | None],
- forward(
- hidden_states: torch.Tensor,
- attention_mask: torch.Tensor | None = None,
- context: torch.Tensor | None = None,
- context_mask: torch.Tensor | None = None,
- rotary_pos_emb: torch.Tensor | None = None,
- rotary_pos_cos: torch.Tensor | None = None,
- rotary_pos_sin: torch.Tensor | None = None,
- rotary_pos_cos_sin: torch.Tensor | None = None,
- attention_bias: torch.Tensor | None = None,
- inference_context: megatron.core.inference.contexts.BaseInferenceContext | None = None,
- packed_seq_params: megatron.core.packed_seq_params.PackedSeqParams | None = None,
- sequence_len_offset: torch.Tensor | None = None,
- padding_mask: torch.Tensor | None = None,
- input_ids: torch.Tensor | None = None,
- **kwargs,
Apply attention, MLP/MoE, and AttnRes state updates.