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#

KimiK3ShortConvolution

TP-sharded KDA short convolution with FP32 state.

KimiK3Attention

Select KDA or Kimi’s no-RoPE MLA according to the global layer number.

KimiK3MoELayer

MCore latent MoE with Kimi’s post-combine RMS normalization.

KimiK3TransformerLayer

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',
) None#
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.ShortConvolution

TP-sharded KDA short convolution with FP32 state.

Initialization

sharded_state_dict(
prefix: str = '',
sharded_offsets: tuple = (),
metadata: dict | None = None,
) megatron.core.dist_checkpointing.mapping.ShardedStateDict#

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.MegatronModule

Select KDA or Kimi’s no-RoPE MLA according to the global layer number.

Initialization

_duplicated_linear(
input_size: int,
output_size: int,
) megatron.core.extensions.transformer_engine.TELinear#
_column_linear(
input_size: int,
output_size: int,
) megatron.core.extensions.transformer_engine.TEColumnParallelLinear#
_row_linear(
input_size: int,
output_size: int,
) megatron.core.extensions.transformer_engine.TERowParallelLinear#
_init_kda(config) None#
_init_mla(config) None#
sharded_state_dict(
prefix: str = '',
sharded_offsets: tuple = (),
metadata: dict | None = None,
) megatron.core.dist_checkpointing.mapping.ShardedStateDict#

Return attention state with explicit KDA TP sharding.

_forward_kda(
hidden_states: torch.Tensor,
packed_seq_params: megatron.core.packed_seq_params.PackedSeqParams | None,
) torch.Tensor#
_forward_mla(
hidden_states: torch.Tensor,
packed_seq_params: megatron.core.packed_seq_params.PackedSeqParams | None,
) torch.Tensor#
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,
) tuple[torch.Tensor, None]#

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.MoELayer

MCore latent MoE with Kimi’s post-combine RMS normalization.

Initialization

postprocess(
output: torch.Tensor,
shared_expert_output: torch.Tensor | None,
) torch.Tensor#

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.TransformerLayer

Transformer layer implementing Kimi K3’s AttnRes residual bank.

Initialization

static _add_bias(
output_with_bias: tuple[torch.Tensor, torch.Tensor | None],
) torch.Tensor#
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,
) tuple[torch.Tensor, torch.Tensor]#

Apply attention, MLP/MoE, and AttnRes state updates.