diff --git a/.github/workflows/integration_test_8gpu_features.yaml b/.github/workflows/integration_test_8gpu_features.yaml index 422d36968c..157029a7c7 100644 --- a/.github/workflows/integration_test_8gpu_features.yaml +++ b/.github/workflows/integration_test_8gpu_features.yaml @@ -68,7 +68,7 @@ jobs: TORCH_SPEC="torch==${{ matrix.torch-version }}" fi python -m pip install --force-reinstall --pre \ - "${TORCH_SPEC}" --index-url ${{ matrix.index-url }} + "${TORCH_SPEC}" torchvision --index-url ${{ matrix.index-url }} # The torchcomms feature tests are currently disabled, so do not install # torchcomms in the main feature job. Its wheel pins torch exactly and @@ -93,6 +93,12 @@ jobs: python -m pytest tests/unit_tests/flex_shard/test_dist_muon.py \ --durations=20 -vv + if [[ "${{ matrix.gpu-arch-type }}" == "cuda" ]]; then + CUDA_VISIBLE_DEVICES=0 python -m pytest \ + tests/unit_tests/test_kimi_k3.py::TestKimiK3::test_fla_kda_kernel_matches_recurrent_reference \ + -vv + fi + sudo mkdir -p "$RUNNER_TEMP/artifacts-to-be-uploaded" sudo chown -R $(id -u):$(id -g) "$RUNNER_TEMP/artifacts-to-be-uploaded" diff --git a/scripts/checkpoint_conversion/numerical_tests_kimi_k3.py b/scripts/checkpoint_conversion/numerical_tests_kimi_k3.py new file mode 100644 index 0000000000..e89210606a --- /dev/null +++ b/scripts/checkpoint_conversion/numerical_tests_kimi_k3.py @@ -0,0 +1,501 @@ +#!/usr/bin/env python3 +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the BSD-style license found in the +# LICENSE file in the root directory of this source tree. + +"""Full text+image logit parity: TorchTitan Kimi K3 vs HuggingFace Kimi K3. + +Runs the released HuggingFace model code and TorchTitan in one process on the +same text+image prompt. Each side performs its own image preprocessing, so the +comparison covers preprocessing, vision, projector, scatter, and decoder. + +The released checkpoint is MXFP4-quantized and is not loaded. Instead, the +script reduces the HuggingFace config to TorchTitan's debug model, initializes +TorchTitan, and strictly transfers its state dict to HuggingFace. + +The script downloads the config, modeling, processor, and tokenizer assets from +a pinned HuggingFace revision without downloading the released weight shards. +The released code requires ``transformers==4.56.2`` and ``tiktoken``. + +Usage: + CUDA_VISIBLE_DEVICES=0 python -m \ + scripts.checkpoint_conversion.numerical_tests_kimi_k3 + +Add ``--force-hf-routing`` for the routing-fixed diagnostic. +""" + +import argparse +from typing import Any, cast + +import torch +import torch.nn.functional as F +from huggingface_hub import snapshot_download +from PIL import Image + +from torchtitan.hf_datasets.multimodal.utils.image import ( + process_image, + resize_to_patch_budget, + vision_to_patches, +) +from torchtitan.models.kimi_k3 import model_registry +from torchtitan.models.kimi_k3.model import KimiK3Model +from torchtitan.models.kimi_k3.state_dict_adapter import KimiK3StateDictAdapter +from transformers import AutoConfig, AutoModelForCausalLM, AutoProcessor + + +_HF_REPO_ID = "moonshotai/Kimi-K3" +_HF_REVISION = "9f62e4e9fffbd0a83ddd60e1c209d828994b3569" +_DTYPE = torch.bfloat16 +_HF_ATTN_BACKEND = "flash_attention_2" +_MEDIA_TOKEN_ID = 163605 +_PATCH_SIZE = 14 +_MERGE_SIZE = 2 +_MAX_PATCHES = 65536 +_MAX_PATCHES_PER_SIDE = 512 +_PROMPT = ( + "<|kimi_image_placeholder|>\n" "What is shown in this image? Describe it briefly." +) + + +def _reduce_hf_config(hf_config, tt_config, hf_model_path: str) -> None: + """Reduce the released HuggingFace config to the TorchTitan debug model.""" + text_config = hf_config.text_config + hf_full_attention_layers = [ + layer_idx + 1 + for layer_idx, layer in enumerate(tt_config.layers) + if layer.attention is not None + ] + hf_kda_layers = [ + layer_idx + 1 + for layer_idx, layer in enumerate(tt_config.layers) + if layer.delta_attention is not None + ] + mla = next(layer.attention for layer in tt_config.layers if layer.attention) + kda = next( + layer.delta_attention for layer in tt_config.layers if layer.delta_attention + ) + dense_ffn = next( + layer.feed_forward for layer in tt_config.layers if layer.feed_forward + ) + moe = next(layer.moe for layer in tt_config.layers if layer.moe) + + text_overrides = { + "vocab_size": tt_config.vocab_size, + "hidden_size": tt_config.dim, + "intermediate_size": dense_ffn.w1.out_features, + "num_hidden_layers": len(tt_config.layers), + "num_attention_heads": mla.n_heads, + "num_key_value_heads": mla.n_heads, + "rms_norm_eps": tt_config.norm.eps, + "q_lora_rank": mla.wq_a.out_features, + "kv_lora_rank": mla.kv_lora_rank, + "qk_nope_head_dim": mla.qk_nope_head_dim, + "qk_rope_head_dim": mla.qk_rope_head_dim, + "v_head_dim": mla.v_head_dim, + "activation_situ_beta": dense_ffn.beta, + "activation_situ_linear_beta": dense_ffn.linear_beta, + "num_experts": moe.num_experts, + "num_experts_per_token": moe.router.top_k, + "num_shared_experts": ( + moe.shared_experts.w1.out_features + // moe.routed_experts.inner_experts.hidden_dim + ), + "moe_renormalize": moe.router.route_norm, + "moe_intermediate_size": moe.routed_experts.inner_experts.hidden_dim, + "routed_expert_hidden_size": moe.routed_down.out_features, + "routed_scaling_factor": moe.router.route_scale, + "first_k_dense_replace": next( + layer_idx for layer_idx, layer in enumerate(tt_config.layers) if layer.moe + ), + "attn_res_block_size": tt_config.layers[0].attn_res_block_size, + "linear_attn_config": { + "full_attn_layers": hf_full_attention_layers, + "kda_layers": hf_kda_layers, + "head_dim": kda.head_dim, + "num_heads": kda.num_heads, + "short_conv_kernel_size": kda.conv_kernel_size, + "gate_lower_bound": kda.kernel.lower_bound, + "use_full_rank_gate": True, + }, + } + for name, value in text_overrides.items(): + setattr(text_config, name, value) + + vision = tt_config.vision_encoder + assert vision is not None + vision_overrides = { + "patch_size": vision.patch_size, + "init_pos_emb_height": vision.init_pos_emb_height, + "init_pos_emb_width": vision.init_pos_emb_width, + "init_pos_emb_time": vision.max_num_frames, + "vt_num_attention_heads": vision.block.attn.num_heads, + "vt_num_hidden_layers": vision.num_layers, + "vt_hidden_size": vision.dim, + "vt_intermediate_size": vision.block.mlp.fc1.out_features, + "merge_kernel_size": vision.merge_kernel_size, + "mm_hidden_size": vision.dim, + "qkv_hidden_size": vision.block.attn.dim, + "text_hidden_size": tt_config.dim, + "pos_emb_interpolation_mode": vision.interpolation_mode, + } + for name, value in vision_overrides.items(): + setattr(hf_config.vision_config, name, value) + + for config in (hf_config, text_config): + if hasattr(config, "quantization_config"): + delattr(config, "quantization_config") + text_config._name_or_path = hf_model_path + hf_config._name_or_path = hf_model_path + + +def _build_tt_model(tt_config, dtype: torch.dtype) -> KimiK3Model: + with torch.device("meta"): + model = tt_config.build() + model.to_empty(device="cpu") + model.to(dtype=dtype) + model.init_states(buffer_device=torch.device("cpu")) + return model.eval() + + +def _build_hf_model( + hf_model_path: str, + tt_config, + hf_state_dict: dict[str, Any], + dtype: torch.dtype, +): + hf_config = AutoConfig.from_pretrained( + hf_model_path, + trust_remote_code=True, + local_files_only=True, + ) + _reduce_hf_config(hf_config, tt_config, hf_model_path) + hf_config.text_config._attn_implementation = _HF_ATTN_BACKEND + hf_config.vision_config._attn_implementation = _HF_ATTN_BACKEND + model = AutoModelForCausalLM.from_config(hf_config, trust_remote_code=True) + model.language_model.config._attn_implementation = _HF_ATTN_BACKEND + model.to(dtype=dtype) + model.load_state_dict(hf_state_dict, strict=True) + return model.eval() + + +@torch.no_grad() +def run_hf( + hf_model_path: str, + tt_config, + hf_state_dict: dict[str, Any], + image_size: int, + dtype: torch.dtype, + device: torch.device, +) -> dict[str, Any]: + """Run HuggingFace preprocessing and the reduced HuggingFace model.""" + print(f"Loading released HuggingFace Kimi K3 code on {device} ...") + processor: Any = AutoProcessor.from_pretrained( + hf_model_path, + trust_remote_code=True, + local_files_only=True, + ) + model = _build_hf_model(hf_model_path, tt_config, hf_state_dict, dtype).to(device) + + raw_image = ( + torch.linspace(0, 255, image_size * image_size * 3) + .reshape(image_size, image_size, 3) + .to(torch.uint8) + ) + pil_image = Image.fromarray(raw_image.numpy()) + batch = processor( + medias=[{"type": "image", "image": pil_image}], # codespell:ignore medias + text=_PROMPT, + return_tensors="pt", + ) + + vision_features: dict[str, torch.Tensor] = {} + + def record_vision_features(_module, _inputs, output) -> None: + features = output[0] if isinstance(output, (list, tuple)) else output + vision_features["output"] = features.detach().float().cpu() + + model.mm_projector.register_forward_hook(record_vision_features) + + expert_indices: dict[int, torch.Tensor] = {} + for layer_idx, layer in enumerate(model.language_model.model.layers): + moe = getattr(layer, "block_sparse_moe", None) + if moe is not None: + moe.gate.register_forward_hook( + lambda _module, _inputs, output, layer_idx=layer_idx: ( + expert_indices.__setitem__( + layer_idx, + output[0].detach().cpu(), + ) + ) + ) + + inputs = { + key: value.to(device) if isinstance(value, torch.Tensor) else value + for key, value in batch.items() + } + inputs["pixel_values"] = inputs["pixel_values"].to(dtype) + output = model(**inputs, use_cache=False) + ref = { + "input_ids": batch["input_ids"].cpu(), + "raw_image": raw_image, + "last_logits": output.logits[:, -1, :].float().cpu(), + "pixel_values": batch["pixel_values"].float().cpu(), + "grid_thws": batch["grid_thws"].cpu(), + "vision_features": vision_features["output"], + "expert_indices": expert_indices, + } + del model + torch.cuda.empty_cache() + return ref + + +def _expand_image_placeholder( + input_ids: torch.Tensor, + image_token_id: int, + num_vision_tokens: int, +) -> torch.Tensor: + """Expand the single HF media placeholder for TorchTitan's scatter path.""" + if input_ids.shape[0] != 1: + raise ValueError("The Kimi K3 numerical test expects a batch size of one.") + positions = (input_ids[0] == image_token_id).nonzero().flatten() + if positions.numel() != 1: + raise ValueError(f"Expected one image placeholder, found {positions.numel()}.") + position = positions.item() + image_tokens = input_ids.new_full((1, num_vision_tokens), image_token_id) + return torch.cat( + (input_ids[:, :position], image_tokens, input_ids[:, position + 1 :]), + dim=1, + ).squeeze(0) + + +def _print_routing_comparison( + hf_expert_indices: dict[int, torch.Tensor], + tt_expert_indices: dict[int, torch.Tensor], +) -> None: + hf_layers = set(hf_expert_indices) + tt_layers = set(tt_expert_indices) + if not hf_layers: + raise ValueError("No MoE routing choices were recorded.") + if hf_layers != tt_layers: + raise ValueError( + "Routing layers differ: " + f"HF-only {sorted(hf_layers - tt_layers)}, " + f"TorchTitan-only {sorted(tt_layers - hf_layers)}." + ) + + num_matching = 0 + num_routings = 0 + for layer_idx in sorted(hf_layers): + top_k = hf_expert_indices[layer_idx].shape[-1] + hf_ids = hf_expert_indices[layer_idx].reshape(-1, top_k).sort(dim=-1).values + tt_ids = tt_expert_indices[layer_idx].reshape(-1, top_k).sort(dim=-1).values + if hf_ids.shape != tt_ids.shape: + raise ValueError( + f"Layer {layer_idx} routing shapes differ: " + f"HF {tuple(hf_ids.shape)} vs TT {tuple(tt_ids.shape)}." + ) + num_matching += int((hf_ids == tt_ids).sum().item()) + num_routings += hf_ids.numel() + match_rate = num_matching / num_routings if num_routings else 0.0 + print(f"router choices: {num_matching}/{num_routings} match " f"({match_rate:.1%})") + + +def _force_hf_routing(model, expert_indices, device) -> None: + """Use HF expert IDs with TorchTitan's independently computed scores.""" + for layer_idx, layer in model.layers.items(): + if (moe := cast(Any, layer.moe)) is None: + continue + ids = expert_indices[int(layer_idx)].to(device) + router, original = moe.router, moe.router.forward + + def forced_forward( + x_TD, + expert_bias_E=None, + _router=router, + _original=original, + _ids=ids, + ): + _, _, scores_TE = _original(x_TD, expert_bias_E) + weights = scores_TE.gather(dim=-1, index=_ids) + if _router.route_norm: + weights = weights / (weights.sum(dim=-1, keepdim=True) + 1e-20) + return weights * _router.route_scale, _ids, scores_TE + + router.forward = forced_forward + + +@torch.no_grad() +def run_tt( + model: KimiK3Model, + ref: dict[str, Any], + vision_dtype: torch.dtype, + device: torch.device, + force_hf_routing: bool, +) -> torch.Tensor: + """Run TorchTitan preprocessing and the reduced TorchTitan model.""" + print(f"Loading TorchTitan Kimi K3 (debugmodel) on {device} ...") + model.to(device) + assert model.vision_encoder is not None + + if force_hf_routing: + print("Using HF expert selections with TorchTitan router scores") + _force_hf_routing(model, ref["expert_indices"], device) + + expert_indices: dict[int, torch.Tensor] = {} + for layer_idx, layer in model.layers.items(): + if layer.moe is not None: + # pyrefly: ignore [missing-attribute] + layer.moe.router.register_forward_hook( + lambda _module, _inputs, output, layer_idx=int(layer_idx): ( + expert_indices.__setitem__( + layer_idx, + output[1].detach().cpu(), + ) + ) + ) + + image = process_image( + Image.fromarray(ref["raw_image"].numpy()), + patch_size=_PATCH_SIZE, + merge_size=_MERGE_SIZE, + resize_fn=resize_to_patch_budget, + max_patches=_MAX_PATCHES, + max_patches_per_side=_MAX_PATCHES_PER_SIDE, + image_mean=(0.5, 0.5, 0.5), + image_std=(0.5, 0.5, 0.5), + ) + if image is None: + raise ValueError("TorchTitan failed to process the numerical test image.") + patches, grid = vision_to_patches( + image, + patch_size=_PATCH_SIZE, + temporal_patch_size=1, + merge_size=_MERGE_SIZE, + patch_order="raster", + ) + pixel_values = patches.to(device=device, dtype=vision_dtype) + grid_thw = grid.unsqueeze(0).to(device) + num_vision_tokens = (grid[1] // _MERGE_SIZE) * (grid[2] // _MERGE_SIZE) + tokens = _expand_image_placeholder( + ref["input_ids"], + _MEDIA_TOKEN_ID, + int(num_vision_tokens.item()), + ).to(device) + positions = torch.arange( + tokens.shape[0], + dtype=torch.int32, + device=device, + ) + attention_masks = model.get_attention_masks(positions) + + print( + f"tokens={tuple(tokens.shape)} pixel_values={tuple(pixel_values.shape)} " + f"grid_thw={grid_thw.tolist()} vision_tokens={num_vision_tokens.item()}" + ) + + hf_pixels = ref["pixel_values"].flatten(1) + pixel_diff = (hf_pixels - patches.float()).abs() + print( + f"pixel values: max_diff={pixel_diff.max().item():.3e} " + f"num_differ={(pixel_diff > 1e-6).sum().item()}/{pixel_diff.numel()}" + ) + + tt_features = model.vision_encoder(pixel_values, grid_thw=grid_thw) + hf_features = ref["vision_features"].reshape(-1, tt_features.shape[-1]) + tt_features = tt_features.float().cpu().reshape(-1, tt_features.shape[-1]) + vision_cos = F.cosine_similarity( + hf_features.flatten(), tt_features.flatten(), dim=0 + ).item() + vision_max_diff = (hf_features - tt_features).abs().max().item() + print( + f"vision features: shape={tuple(tt_features.shape)} " + f"cos={vision_cos:.6f} max_diff={vision_max_diff:.3e}" + ) + + logits = model( + tokens, + pixel_values=pixel_values, + grid_thw=grid_thw, + special_tokens={"image_id": _MEDIA_TOKEN_ID}, + positions=positions, + attention_masks=attention_masks, + ) + _print_routing_comparison(ref["expert_indices"], expert_indices) + return logits[-1].float().cpu() + + +def compare(ref_logits: torch.Tensor, tt_logits: torch.Tensor) -> None: + """Print last-token parity metrics.""" + ref = ref_logits.squeeze() + tt = tt_logits.squeeze() + log_ref = F.log_softmax(ref, dim=-1) + log_tt = F.log_softmax(tt, dim=-1) + kl = F.kl_div(log_tt, log_ref, log_target=True, reduction="sum").item() + cosine = F.cosine_similarity(ref, tt, dim=-1).item() + max_diff = (ref - tt).abs().max().item() + top1 = (ref.argmax() == tt.argmax()).item() + top5_overlap = ( + len(set(ref.topk(5).indices.tolist()) & set(tt.topk(5).indices.tolist())) / 5 + ) + print("\nFull multimodal last-token logit parity (TorchTitan vs HuggingFace)") + print( + f"KL={kl:.4e} cos={cosine:.6f} max_diff={max_diff:.4e} " + f"top1={'Y' if top1 else 'N'} top5={top5_overlap:.0%}" + ) + + +@torch.no_grad() +def main() -> None: + parser = argparse.ArgumentParser() + parser.add_argument("--model_flavor", default="debugmodel") + parser.add_argument("--image_size", type=int, default=336) + parser.add_argument( + "--force-hf-routing", + action="store_true", + help="Use HF expert selections with TorchTitan router scores.", + ) + parser.add_argument("--seed", type=int, default=42) + args = parser.parse_args() + + if not torch.cuda.is_available(): + parser.error("Kimi K3 numerical parity requires a CUDA GPU.") + + hf_model_path = snapshot_download( + repo_id=_HF_REPO_ID, + revision=_HF_REVISION, + allow_patterns=["*.json", "*.py", "tiktoken.model"], + ) + device = torch.device("cuda") + dtype = _DTYPE + print(f"dtype={dtype} hf_attn={_HF_ATTN_BACKEND}") + + tt_config = cast(KimiK3Model.Config, model_registry(args.model_flavor).model) + torch.manual_seed(args.seed) + tt_model = _build_tt_model(tt_config, dtype) + hf_state_dict = KimiK3StateDictAdapter(tt_config, hf_assets_path=None).to_hf( + tt_model.state_dict() + ) + + ref = run_hf( + hf_model_path, + tt_config, + hf_state_dict, + args.image_size, + dtype, + device, + ) + del hf_state_dict + tt_logits = run_tt( + tt_model, + ref, + dtype, + device, + args.force_hf_routing, + ) + compare(ref["last_logits"], tt_logits) + + +if __name__ == "__main__": + main() diff --git a/tests/integration_tests/models.py b/tests/integration_tests/models.py index 974c6fb86b..74f3627d1b 100755 --- a/tests/integration_tests/models.py +++ b/tests/integration_tests/models.py @@ -129,4 +129,11 @@ def build_model_tests_list() -> list[OverrideDefinitions]: test_name="muse_glimmer_mm_fsdp+tp+sp", ngpu=4, ), + # Integration Test Case for Kimi K3 + OverrideDefinitions( + configs=[recipes.kimi_k3_debugmodel_mm_fsdp2], + test_descr="Kimi K3 multimodal FSDP", + test_name="kimi_k3_mm_fsdp", + ngpu=2, + ), ] diff --git a/tests/unit_tests/test_kimi_k3.py b/tests/unit_tests/test_kimi_k3.py new file mode 100644 index 0000000000..9d02f1e010 --- /dev/null +++ b/tests/unit_tests/test_kimi_k3.py @@ -0,0 +1,222 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the BSD-style license found in the +# LICENSE file in the root directory of this source tree. + +import unittest + +import torch +import torch.nn.functional as F +from torch.nn.attention.flex_attention import BlockMask + +from torchtitan.models.kimi_k3 import _kimi_k3_config, _vision_encoder_config +from torchtitan.models.kimi_k3.kda import KimiKDAKernel +from torchtitan.models.kimi_k3.model import KimiK3Model +from torchtitan.models.kimi_k3.state_dict_adapter import KimiK3StateDictAdapter + + +def _small_model_config() -> KimiK3Model.Config: + """Build a reduced KDA+MLA, dense+MoE, multimodal Kimi K3 config.""" + dim = 64 + return _kimi_k3_config( + dim=dim, + vocab_size=32, + num_layers=2, + full_attention_layers={1}, + attn_res_block_size=1, + num_heads=2, + q_lora_rank=32, + kv_lora_rank=32, + qk_nope_head_dim=16, + qk_rope_head_dim=16, + v_head_dim=16, + kda_head_dim=64, + conv_kernel_size=3, + dense_hidden_dim=128, + latent_dim=32, + expert_hidden_dim=32, + num_experts=2, + top_k=1, + num_shared_experts=1, + vision_encoder=_vision_encoder_config( + text_dim=dim, + dim=48, + qkv_dim=48, + hidden_dim=96, + num_layers=1, + num_heads=3, + patch_size=2, + merge_kernel_size=(2, 2), + init_pos_emb_height=2, + init_pos_emb_width=2, + max_num_frames=1, + ), + attn_backend="flex", + ) + + +def _kda_recurrent_reference( + q_BLHK: torch.Tensor, + k_BLHK: torch.Tensor, + v_BLHV: torch.Tensor, + gate_BLHK: torch.Tensor, + beta_BLH: torch.Tensor, + A_log_H: torch.Tensor, + dt_bias_HK: torch.Tensor, + *, + lower_bound: float | None, +) -> torch.Tensor: + """Explicit KDA recurrence in FP32, matching the released Kimi K3 math. + + ``lower_bound`` selects the same two gate activations FLA exposes through + ``safe_gate``: the bounded ``lower_bound * sigmoid(...)`` form when set, + and ``-exp(A_log) * softplus(...)`` when ``None``. + """ + input_dtype = q_BLHK.dtype + q_BLHK = q_BLHK.float() + k_BLHK = k_BLHK.float() + q_BLHK = q_BLHK * torch.rsqrt(q_BLHK.square().sum(dim=-1, keepdim=True) + 1e-6) + k_BLHK = k_BLHK * torch.rsqrt(k_BLHK.square().sum(dim=-1, keepdim=True) + 1e-6) + v_BLHV = v_BLHV.float() + if lower_bound is None: + log_decay_BLHK = -torch.exp(A_log_H.float()).view(1, 1, -1, 1) * F.softplus( + gate_BLHK.float() + dt_bias_HK.float() + ) + else: + log_decay_BLHK = lower_bound * torch.sigmoid( + torch.exp(A_log_H.float()).view(1, 1, -1, 1) + * (gate_BLHK.float() + dt_bias_HK.float()) + ) + decay_BLHK = torch.exp(log_decay_BLHK) + beta_BLH = torch.sigmoid(beta_BLH.float()) + + B, L, H, K = q_BLHK.shape + V = v_BLHV.shape[-1] + state_BHKV = torch.zeros(B, H, K, V, device=q_BLHK.device) + outputs_BHV = [] + for token_idx in range(L): + state_BHKV = state_BHKV * decay_BLHK[:, token_idx].unsqueeze(-1) + old_value_BHV = torch.matmul( + k_BLHK[:, token_idx].unsqueeze(-2), + state_BHKV, + ).squeeze(-2) + delta_BHV = (v_BLHV[:, token_idx] - old_value_BHV) * beta_BLH[ + :, token_idx + ].unsqueeze(-1) + state_BHKV = state_BHKV + ( + k_BLHK[:, token_idx].unsqueeze(-1) * delta_BHV.unsqueeze(-2) + ) + outputs_BHV.append( + torch.matmul( + q_BLHK[:, token_idx].unsqueeze(-2), + state_BHKV, + ).squeeze(-2) + * (K**-0.5) + ) + return torch.stack(outputs_BHV, dim=1).to(input_dtype) + + +class TestKimiK3(unittest.TestCase): + def test_flex_attention_mask(self): + config = _small_model_config() + model = config.build() + positions = torch.arange(4, dtype=torch.int32) + attention_masks = model.get_attention_masks(positions) + self.assertIsInstance(attention_masks, BlockMask) + + @unittest.skipIf(not torch.cuda.is_available(), "FLA KDA kernel requires CUDA.") + def test_fla_kda_kernel_matches_recurrent_reference(self): + torch.manual_seed(1) + head_dim = 64 + num_heads = 3 + + def parameter(*shape: int) -> torch.Tensor: + return torch.randn( + *shape, + device="cuda", + dtype=torch.bfloat16, + requires_grad=True, + ) + + for lower_bound in (-5.0, None): + with self.subTest(lower_bound=lower_bound): + A_log_H = torch.rand(num_heads, device="cuda") + A_log_H = A_log_H.uniform_(1.0, 16.0).log().requires_grad_() + actual_inputs = ( + parameter(2, 64, num_heads, head_dim), + parameter(2, 64, num_heads, head_dim), + parameter(2, 64, num_heads, head_dim), + parameter(2, 64, num_heads, head_dim), + parameter(2, 64, num_heads), + A_log_H, + parameter(num_heads, head_dim), + ) + expected_inputs = tuple( + tensor.detach().clone().requires_grad_() for tensor in actual_inputs + ) + + kernel = KimiKDAKernel.Config(lower_bound=lower_bound).build() + actual_BLHV = kernel(*actual_inputs) + expected_BLHV = _kda_recurrent_reference( + *expected_inputs, + lower_bound=lower_bound, + ) + + # The chunked kernel accumulates over chunk boundaries and uses + # reduced-precision matmuls internally, so it does not reproduce + # the sequential FP32 recurrence bit for bit. + torch.testing.assert_close( + actual_BLHV, + expected_BLHV, + atol=2e-3, + rtol=2e-3, + ) + output_grad_BLHV = torch.randn_like(actual_BLHV) + actual_grads = torch.autograd.grad( + actual_BLHV, + actual_inputs, + grad_outputs=output_grad_BLHV, + ) + expected_grads = torch.autograd.grad( + expected_BLHV, + expected_inputs, + grad_outputs=output_grad_BLHV, + ) + for actual_grad, expected_grad in zip( + actual_grads, + expected_grads, + strict=True, + ): + torch.testing.assert_close( + actual_grad, + expected_grad, + atol=2e-2, + rtol=2e-2, + ) + + def test_state_dict_round_trips_through_hf_adapter(self): + torch.manual_seed(2) + config = _small_model_config() + model = config.build() + model.init_states() + + state_dict = model.state_dict() + adapter = KimiK3StateDictAdapter(config, hf_assets_path=None) + hf_state_dict = adapter.to_hf(state_dict) + self.assertIn( + "layers.1.moe.routed_experts.inner_experts.w1_EFD", + state_dict, + ) + self.assertIn( + "language_model.model.layers.1.block_sparse_moe.experts.0.w1.weight", + hf_state_dict, + ) + roundtrip_state_dict = adapter.from_hf(hf_state_dict) + self.assertEqual(state_dict.keys(), roundtrip_state_dict.keys()) + for key, value in state_dict.items(): + torch.testing.assert_close(value, roundtrip_state_dict[key]) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/unit_tests/test_no_new_cli_options.py b/tests/unit_tests/test_no_new_cli_options.py index bbd8af5218..014d08b7ec 100644 --- a/tests/unit_tests/test_no_new_cli_options.py +++ b/tests/unit_tests/test_no_new_cli_options.py @@ -366,6 +366,7 @@ def _declared_cli_options( ("gpt_oss", "gpt_oss_debugmodel"), ("flux", "flux_debugmodel"), ("kimi_k2_7", "kimi_k2_5_debugmodel"), + ("kimi_k3", "kimi_k3_debugmodel"), ("muse_glimmer", "muse_glimmer_debugmodel_mm"), ) diff --git a/torchtitan/models/__init__.py b/torchtitan/models/__init__.py index 784b4110ca..0f22e40f1a 100644 --- a/torchtitan/models/__init__.py +++ b/torchtitan/models/__init__.py @@ -10,6 +10,7 @@ "flux", "gpt_oss", "kimi_k2_7", + "kimi_k3", "llama3", "muse_glimmer", "qwen3", diff --git a/torchtitan/models/common/decoder.py b/torchtitan/models/common/decoder.py index d6f52233f2..a208df1244 100644 --- a/torchtitan/models/common/decoder.py +++ b/torchtitan/models/common/decoder.py @@ -146,10 +146,10 @@ def update_from_config( When *config* is a ``Trainer.Config``, validates ``training.max_context_length`` against each attention layer's intrinsic - RoPE max sequence length, resizes RoPE caches, and propagates - debug flags. Non-trainer callers may pass any config-like - object with a ``ParallelismConfig`` in its ``parallelism`` - field; in that case the training/debug setup is skipped. + RoPE max context length, resizes RoPE caches when present, and + propagates debug flags. Non-trainer callers may pass any config-like + object with a ``ParallelismConfig`` in its ``parallelism`` field; in + that case the training/debug setup is skipped. """ from torchtitan.config import ParallelismConfig from torchtitan.distributed.context_parallel import validate_cp_backend @@ -216,20 +216,24 @@ def update_from_config( if isinstance(config, Trainer.Config): debug = config.debug seq_len = config.training.max_context_length - max_context_length = self.max_context_length - if seq_len > max_context_length: - raise ValueError( - f"Training sequence length {seq_len} exceeds " - f"attention RoPE maximum supported sequence " - f"length {max_context_length}." - ) + rope_cfg = getattr(attention, "rope", None) + if rope_cfg is not None: + max_context_length = self.max_context_length + if seq_len > max_context_length: + raise ValueError( + f"Training sequence length {seq_len} exceeds " + f"attention RoPE maximum supported sequence " + f"length {max_context_length}." + ) for layer_cfg in self.layers: attention_cfg = getattr(layer_cfg, "attention", None) if attention_cfg is not None: - attention_cfg.rope = dataclasses.replace( - attention_cfg.rope, max_context_length=seq_len - ) + rope_cfg = getattr(attention_cfg, "rope", None) + if rope_cfg is not None: + attention_cfg.rope = dataclasses.replace( + rope_cfg, max_context_length=seq_len + ) if hasattr(layer_cfg, "moe") and layer_cfg.moe is not None: layer_cfg.moe.router._debug_force_load_balance = ( debug.moe_force_load_balance diff --git a/torchtitan/models/common/vision_encoder.py b/torchtitan/models/common/vision_encoder.py index 4576a4bfd9..37f446a2ac 100644 --- a/torchtitan/models/common/vision_encoder.py +++ b/torchtitan/models/common/vision_encoder.py @@ -28,7 +28,7 @@ from torchtitan.models.common import Linear from torchtitan.models.common.attention import FlexAttention, local_head_split -from torchtitan.models.common.nn_modules import GELU, LayerNorm +from torchtitan.models.common.nn_modules import GELU, LayerNorm, RMSNorm from torchtitan.protocols.module import Module compiled_create_block_mask = torch.compile(create_block_mask) @@ -150,8 +150,9 @@ class VisionTransformerBlock(Module): @dataclass(kw_only=True, slots=True) class Config(Module.Config): - norm1: LayerNorm.Config - norm2: LayerNorm.Config + # MoonViT normalizes with RMSNorm; Qwen3.5 and Muse Glimmer use LayerNorm. + norm1: LayerNorm.Config | RMSNorm.Config + norm2: LayerNorm.Config | RMSNorm.Config attn: VisionAttention.Config mlp: VisionMLP.Config diff --git a/torchtitan/models/kimi_k2_7/vision_encoder.py b/torchtitan/models/kimi_k2_7/vision_encoder.py index b3b00ccbab..7041ed9182 100644 --- a/torchtitan/models/kimi_k2_7/vision_encoder.py +++ b/torchtitan/models/kimi_k2_7/vision_encoder.py @@ -27,7 +27,7 @@ from torchtitan.distributed.utils import get_spmd_backend from torchtitan.models.common import Linear -from torchtitan.models.common.nn_modules import GELU, LayerNorm +from torchtitan.models.common.nn_modules import GELU, LayerNorm, RMSNorm from torchtitan.models.common.rope import _maybe_wrap_positions, ComplexRoPE from torchtitan.models.common.vision_encoder import ( create_block_diagonal_mask, @@ -325,31 +325,26 @@ def forward(self, merged_MK: torch.Tensor) -> torch.Tensor: return self.linear_2(x) -class KimiK25VisionEncoder(Module): - """MoonViT3d vision tower + multimodal projector for Kimi K2.5.""" +class MoonViTEncoder(Module): + """MoonViT3d vision tower + multimodal projector.""" @dataclass(kw_only=True, slots=True) class Config(Module.Config): - dim: int = 1152 - num_layers: int = 27 - num_heads: int = 16 - - patch_size: int = 14 - in_channels: int = 3 - merge_kernel_size: list[int] = field(default_factory=lambda: [2, 2]) - text_hidden_size: int = 7168 + dim: int + num_layers: int + merge_kernel_size: list[int] # Learnable 2D spatial position table, shape (height, width, dim). - init_pos_emb_height: int = 64 - init_pos_emb_width: int = 64 - interpolation_mode: str = "bicubic" + init_pos_emb_height: int + init_pos_emb_width: int + interpolation_mode: str # Sub-modules. patch_embed_proj: Linear.Config rotary_pos_emb: VisionRotaryEmbedding2D.Config block: VisionTransformerBlock.Config - final_norm: LayerNorm.Config - projector: VisionProjector.Config + final_norm: LayerNorm.Config | RMSNorm.Config + projector: Module.Config def __init__(self, config: Config): super().__init__() @@ -472,3 +467,25 @@ def forward( # pyrefly: ignore [bad-argument-type] merged = _tpool_patch_merger(x, grids, self.merge_kernel_size) return self.projector(merged) + + +class KimiK25VisionEncoder(MoonViTEncoder): + """MoonViT3d vision tower + multimodal projector for Kimi K2.5.""" + + @dataclass(kw_only=True, slots=True) + class Config(MoonViTEncoder.Config): + dim: int = 1152 + num_layers: int = 27 + num_heads: int = 16 + + patch_size: int = 14 + in_channels: int = 3 + merge_kernel_size: list[int] = field(default_factory=lambda: [2, 2]) + text_hidden_size: int = 7168 + + init_pos_emb_height: int = 64 + init_pos_emb_width: int = 64 + interpolation_mode: str = "bicubic" + + final_norm: LayerNorm.Config # pyrefly: ignore [bad-override] + projector: VisionProjector.Config # pyrefly: ignore [bad-override] diff --git a/torchtitan/models/kimi_k3/README.md b/torchtitan/models/kimi_k3/README.md new file mode 100644 index 0000000000..2a0f09205e --- /dev/null +++ b/torchtitan/models/kimi_k3/README.md @@ -0,0 +1,60 @@ +# Kimi K3 + +Kimi K3 combines a hybrid Kimi Delta Attention (KDA) and Multi-head Latent +Attention (MLA) decoder with LatentMoE and a MoonViT-V2 vision encoder. + +## Prerequisites + +Install the additional dependencies: + +```bash +pip install av einops pillow torchvision flash-linear-attention +``` + +## Architecture + +Kimi K3 is built on Kimi Delta Attention (KDA) and Attention Residuals +(AttnRes), with 69 KDA layers and 24 Gated MLA layers. Stable LatentMoE selects +16 of 896 experts per token, and MoonViT-V2 provides native vision input. + +## Released Model Configuration + +The values below follow the +[official Kimi K3 model card](https://huggingface.co/moonshotai/Kimi-K3) and +describe the released model. + +| Component | Configuration | +|-----------|---------------| +| Architecture | Mixture-of-Experts (MoE) | +| Parameters | 2.8T total, 104B activated | +| Decoder | 93 layers, 1 dense layer, hidden size 7168, 96 attention heads | +| Attention | 69 KDA layers and 24 Gated MLA layers, context length 1048576 | +| LatentMoE | Dimension 3584, hidden size 3072 per expert, 896 experts, top-16 routing, 2 shared experts | +| Vocabulary | 160K | +| Activation | SiTU-GLU | +| Vision encoder | MoonViT-V2, 401M parameters | +| Quantization | MXFP4 weights and MXFP8 activations with quantization-aware training | +| Modality | Text and image | + +## Supported Parallelisms + +| Feature | Notes | +|---------|-------| +| FSDP2 / HSDP | Decoder sharded per layer; vision encoder sharded as a separate unit | + +## Numerical Parity + +The parity script reduces the released Hugging Face configuration to match +TorchTitan's local `debugmodel` configuration before initializing both models. + +End-to-end KL divergence against the Hugging Face implementation (multimodal +inputs): **6.7634e-7**, with **100% top-1 and top-5 match**. + +Vision parity: pixel preprocessing max difference **1.192e-7**; projected vision +features cosine similarity **1.000000** and max difference **2.730e-3**. + +Test scripts: + +- `scripts/checkpoint_conversion/numerical_tests_kimi_k3.py` -- Hugging Face vs. + TorchTitan comparison +- `tests/unit_tests/test_kimi_k3.py` -- KDA and FSDP2 correctness diff --git a/torchtitan/models/kimi_k3/__init__.py b/torchtitan/models/kimi_k3/__init__.py new file mode 100644 index 0000000000..e559f1469c --- /dev/null +++ b/torchtitan/models/kimi_k3/__init__.py @@ -0,0 +1,551 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the BSD-style license found in the +# LICENSE file in the root directory of this source tree. + +from collections.abc import Callable +from functools import partial + +import torch +import torch.nn as nn + +from torchtitan.components.optimizer import register_moe_load_balancing_hook +from torchtitan.models.common import Conv1d, Embedding, Linear +from torchtitan.models.common.config_utils import get_attention_config +from torchtitan.models.common.moe import RoutedExperts, TokenChoiceTopKRouter +from torchtitan.models.common.nn_modules import GELU, RMSNorm +from torchtitan.models.common.token_dispatcher import LocalTokenDispatcher +from torchtitan.models.common.vision_encoder import ( + VisionAttention, + VisionMLP, + VisionTransformerBlock, +) +from torchtitan.models.kimi_k2_7.vision_encoder import VisionRotaryEmbedding2D +from torchtitan.models.utils import validate_converter_order +from torchtitan.protocols.model import ModelConfigConverter +from torchtitan.protocols.model_spec import ModelSpec + +from .kda import KimiDeltaAttention, KimiKDAKernel, KimiRMSNormGated +from .model import KimiK3Model, KimiK3TransformerBlock, KimiMLAAttention +from .moe import KimiFeedForward, KimiGroupedExperts, KimiLatentMoE +from .parallelize import parallelize_kimi_k3 +from .state_dict_adapter import KimiK3StateDictAdapter +from .vision_encoder import KimiK3VisionEncoder, KimiK3VisionProjector + +__all__ = [ + "KIMI_K3_SPECIAL_TOKENS", + "KimiK3Model", + "KimiK3StateDictAdapter", + "KimiK3VisionEncoder", + "kimi_k3_configs", + "model_registry", + "parallelize_kimi_k3", +] + + +KIMI_K3_SPECIAL_TOKENS = { + "image_token": "<|media_pad|>", + "video_token": "<|media_pad|>", + "vision_start_token": "<|media_begin|>", + "vision_end_token": "<|media_end|>", + "pad_token": "[PAD]", +} + + +_LINEAR_INIT = { + "weight": partial(nn.init.trunc_normal_, std=0.02), + "bias": nn.init.zeros_, +} +_CONV_INIT = {"weight": partial(nn.init.trunc_normal_, std=0.02)} +_NORM_INIT = {"weight": nn.init.ones_} +_EMBEDDING_INIT = {"weight": partial(nn.init.normal_, std=1.0)} +_POS_EMBED_INIT = {"pos_embed": partial(nn.init.normal_, std=1.0)} + + +def _output_linear_init(dim: int) -> dict[str, Callable]: + scale = dim**-0.5 + return { + "weight": partial( + nn.init.trunc_normal_, + std=scale, + a=-3 * scale, + b=3 * scale, + ) + } + + +def _fan_in_linear_init(in_features: int) -> dict[str, Callable]: + return { + "weight": partial( + nn.init.trunc_normal_, + std=(2.0 / in_features) ** 0.5, + ), + "bias": nn.init.zeros_, + } + + +def _a_log_init(parameter: nn.Parameter) -> None: + with torch.no_grad(): + nn.init.uniform_(parameter, 1.0, 16.0) + parameter.log_() + + +def _linear( + in_features: int, + out_features: int, + *, + bias: bool = False, + param_init: dict[str, Callable] | None = None, +) -> Linear.Config: + return Linear.Config( + in_features=in_features, + out_features=out_features, + bias=bias, + param_init=param_init or _LINEAR_INIT, + ) + + +def _norm(dim: int, eps: float = 1e-5) -> RMSNorm.Config: + return RMSNorm.Config( + normalized_shape=dim, + eps=eps, + param_init=_NORM_INIT, + ) + + +def _feed_forward_config( + *, + dim: int, + hidden_dim: int, +) -> KimiFeedForward.Config: + return KimiFeedForward.Config( + w1=_linear(dim, hidden_dim), + w2=_linear(hidden_dim, dim), + w3=_linear(dim, hidden_dim), + beta=4.0, + linear_beta=25.0, + ) + + +def _mla_config( + *, + dim: int, + num_heads: int, + q_lora_rank: int, + kv_lora_rank: int, + qk_nope_head_dim: int, + qk_rope_head_dim: int, + v_head_dim: int, + attn_backend: str, +) -> KimiMLAAttention.Config: + inner_attention = get_attention_config(attn_backend) + + q_head_dim = qk_nope_head_dim + qk_rope_head_dim + return KimiMLAAttention.Config( + dim=dim, + n_heads=num_heads, + kv_lora_rank=kv_lora_rank, + qk_nope_head_dim=qk_nope_head_dim, + qk_rope_head_dim=qk_rope_head_dim, + v_head_dim=v_head_dim, + wq_a=_linear(dim, q_lora_rank), + q_norm=_norm(q_lora_rank), + wq_b=_linear(q_lora_rank, num_heads * q_head_dim), + wkv_a=_linear(dim, kv_lora_rank + qk_rope_head_dim), + kv_norm=_norm(kv_lora_rank), + wkv_b=_linear( + kv_lora_rank, + num_heads * (qk_nope_head_dim + v_head_dim), + ), + gate=_linear(dim, num_heads * v_head_dim), + wo=_linear(num_heads * v_head_dim, dim), + inner_attention=inner_attention, + ) + + +def _kda_config( + *, + dim: int, + num_heads: int, + head_dim: int, + conv_kernel_size: int, +) -> KimiDeltaAttention.Config: + projection_dim = num_heads * head_dim + + def conv() -> Conv1d.Config: + return Conv1d.Config( + in_channels=projection_dim, + out_channels=projection_dim, + kernel_size=conv_kernel_size, + groups=projection_dim, + bias=False, + param_init=_CONV_INIT, + ) + + return KimiDeltaAttention.Config( + dim=dim, + num_heads=num_heads, + head_dim=head_dim, + conv_kernel_size=conv_kernel_size, + q_proj=_linear(dim, projection_dim), + k_proj=_linear(dim, projection_dim), + v_proj=_linear(dim, projection_dim), + q_conv=conv(), + k_conv=conv(), + v_conv=conv(), + forget_a=_linear(dim, head_dim), + forget_b=_linear(head_dim, projection_dim), + beta=_linear(dim, num_heads), + output_gate=_linear(dim, projection_dim), + kernel=KimiKDAKernel.Config(lower_bound=-5.0), + output_norm=KimiRMSNormGated.Config( + dim=head_dim, + eps=1e-5, + param_init=_NORM_INIT, + ), + output_proj=_linear(projection_dim, dim), + param_init={ + "A_log": _a_log_init, + "dt_bias": nn.init.zeros_, + }, + ) + + +def _latent_moe_config( + *, + dim: int, + latent_dim: int, + expert_hidden_dim: int, + num_experts: int, + top_k: int, + num_shared_experts: int, +) -> KimiLatentMoE.Config: + return KimiLatentMoE.Config( + num_experts=num_experts, + router=TokenChoiceTopKRouter.Config( + num_experts=num_experts, + top_k=top_k, + gate=_linear(dim, num_experts), + score_func="sigmoid", + route_norm=True, + route_scale=1.0, + ), + routed_down=_linear(dim, latent_dim), + routed_experts=RoutedExperts.Config( + inner_experts=KimiGroupedExperts.Config( + dim=latent_dim, + hidden_dim=expert_hidden_dim, + num_experts=num_experts, + beta=4.0, + linear_beta=25.0, + param_init={ + "w1_EFD": partial(nn.init.trunc_normal_, std=0.02), + "w2_EDF": partial(nn.init.trunc_normal_, std=0.02), + "w3_EFD": partial(nn.init.trunc_normal_, std=0.02), + }, + ), + token_dispatcher=LocalTokenDispatcher.Config( + num_experts=num_experts, + top_k=top_k, + ), + ), + routed_norm=_norm(latent_dim), + routed_up=_linear(latent_dim, dim), + shared_experts=_feed_forward_config( + dim=dim, + hidden_dim=num_shared_experts * expert_hidden_dim, + ), + load_balance_coeff=1e-3, + ) + + +def _vision_encoder_config( + *, + text_dim: int, + dim: int, + qkv_dim: int, + hidden_dim: int, + num_layers: int, + num_heads: int, + patch_size: int = 14, + in_channels: int = 3, + merge_kernel_size: tuple[int, int] = (2, 2), + init_pos_emb_height: int = 16, + init_pos_emb_width: int = 16, + max_num_frames: int = 4, +) -> KimiK3VisionEncoder.Config: + patch_dim = in_channels * patch_size * patch_size + head_dim = qkv_dim // num_heads + merged_dim = dim * merge_kernel_size[0] * merge_kernel_size[1] + vision_norm = RMSNorm.Config( + normalized_shape=dim, + eps=1e-5, + param_init=_NORM_INIT, + ) + block = VisionTransformerBlock.Config( + norm1=vision_norm, + norm2=vision_norm, + attn=VisionAttention.Config( + dim=qkv_dim, + num_heads=num_heads, + wq=_linear(dim, qkv_dim), + wk=_linear(dim, qkv_dim), + wv=_linear(dim, qkv_dim), + proj=_linear(qkv_dim, dim), + ), + mlp=VisionMLP.Config( + fc1=_linear( + dim, + hidden_dim, + param_init=_fan_in_linear_init(dim), + ), + fc2=_linear( + hidden_dim, + dim, + param_init=_fan_in_linear_init(hidden_dim), + ), + act_fn=GELU.Config(approximate="tanh"), + ), + ) + return KimiK3VisionEncoder.Config( + dim=dim, + num_layers=num_layers, + patch_size=patch_size, + in_channels=in_channels, + merge_kernel_size=merge_kernel_size, + init_pos_emb_height=init_pos_emb_height, + init_pos_emb_width=init_pos_emb_width, + max_num_frames=max_num_frames, + interpolation_mode="bilinear", + patch_embed_proj=_linear(patch_dim, dim), + rotary_pos_emb=VisionRotaryEmbedding2D.Config(head_dim=head_dim), + block=block, + final_norm=vision_norm, + projector=KimiK3VisionProjector.Config( + linear_1=_linear( + merged_dim, + merged_dim, + param_init=_fan_in_linear_init(merged_dim), + ), + linear_2=_linear( + merged_dim, + text_dim, + param_init=_fan_in_linear_init(merged_dim), + ), + post_norm=RMSNorm.Config( + normalized_shape=text_dim, + eps=1e-5, + param_init=_NORM_INIT, + ), + activation=GELU.Config(), + ), + param_init=_POS_EMBED_INIT, + ) + + +def _kimi_k3_config( + *, + dim: int, + vocab_size: int, + num_layers: int, + full_attention_layers: set[int], + attn_res_block_size: int, + num_heads: int, + q_lora_rank: int, + kv_lora_rank: int, + qk_nope_head_dim: int, + qk_rope_head_dim: int, + v_head_dim: int, + kda_head_dim: int, + conv_kernel_size: int, + dense_hidden_dim: int, + latent_dim: int, + expert_hidden_dim: int, + num_experts: int, + top_k: int, + num_shared_experts: int, + vision_encoder: KimiK3VisionEncoder.Config, + attn_backend: str, +) -> KimiK3Model.Config: + """Assemble a Kimi K3 config from the released topology's free parameters. + + ``full_attention_layers`` holds zero-based layer indices. Every other layer + is KDA. Layer 0 is the single dense FFN layer (released + ``first_k_dense_replace=1``); the rest are LatentMoE. + """ + layers = [] + for layer_idx in range(num_layers): + is_full_attention = layer_idx in full_attention_layers + layers.append( + KimiK3TransformerBlock.Config( + layer_id=layer_idx, + attn_res_block_size=attn_res_block_size, + attention=( + _mla_config( + dim=dim, + num_heads=num_heads, + q_lora_rank=q_lora_rank, + kv_lora_rank=kv_lora_rank, + qk_nope_head_dim=qk_nope_head_dim, + qk_rope_head_dim=qk_rope_head_dim, + v_head_dim=v_head_dim, + attn_backend=attn_backend, + ) + if is_full_attention + else None + ), + delta_attention=( + None + if is_full_attention + else _kda_config( + dim=dim, + num_heads=num_heads, + head_dim=kda_head_dim, + conv_kernel_size=conv_kernel_size, + ) + ), + feed_forward=( + _feed_forward_config(dim=dim, hidden_dim=dense_hidden_dim) + if layer_idx == 0 + else None + ), + moe=( + None + if layer_idx == 0 + else _latent_moe_config( + dim=dim, + latent_dim=latent_dim, + expert_hidden_dim=expert_hidden_dim, + num_experts=num_experts, + top_k=top_k, + num_shared_experts=num_shared_experts, + ) + ), + attention_norm=_norm(dim), + ffn_norm=_norm(dim), + attention_res_norm=None if layer_idx == 0 else _norm(dim), + attention_res_proj=None if layer_idx == 0 else _linear(dim, 1), + ffn_res_norm=_norm(dim), + ffn_res_proj=_linear(dim, 1), + ) + ) + + return KimiK3Model.Config( + dim=dim, + vocab_size=vocab_size, + tok_embeddings=Embedding.Config( + num_embeddings=vocab_size, + embedding_dim=dim, + param_init=_EMBEDDING_INIT, + ), + layers=layers, + norm=_norm(dim), + lm_head=_linear( + dim, + vocab_size, + param_init=_output_linear_init(dim), + ), + output_res_norm=_norm(dim), + output_res_proj=_linear(dim, 1), + vision_encoder=vision_encoder, + ) + + +def _debugmodel(attn_backend: str) -> KimiK3Model.Config: + dim = 1024 + return _kimi_k3_config( + dim=dim, + vocab_size=163840, + num_layers=24, + full_attention_layers={3, 7, 11, 15, 19, 23}, + attn_res_block_size=12, + num_heads=16, + q_lora_rank=512, + kv_lora_rank=256, + qk_nope_head_dim=64, + qk_rope_head_dim=32, + v_head_dim=64, + kda_head_dim=64, + conv_kernel_size=4, + dense_hidden_dim=4096, + latent_dim=512, + expert_hidden_dim=384, + num_experts=32, + top_k=4, + num_shared_experts=2, + vision_encoder=_vision_encoder_config( + text_dim=dim, + dim=512, + qkv_dim=768, + hidden_dim=2048, + num_layers=8, + num_heads=6, + init_pos_emb_height=32, + init_pos_emb_width=32, + ), + attn_backend=attn_backend, + ) + + +def _kimi_k3(attn_backend: str) -> KimiK3Model.Config: + dim = 7168 + return _kimi_k3_config( + dim=dim, + vocab_size=163840, + num_layers=93, + full_attention_layers=set(range(3, 92, 4)) | {92}, + attn_res_block_size=12, + num_heads=96, + q_lora_rank=1536, + kv_lora_rank=512, + qk_nope_head_dim=128, + qk_rope_head_dim=64, + v_head_dim=128, + kda_head_dim=128, + conv_kernel_size=4, + dense_hidden_dim=33792, + latent_dim=3584, + expert_hidden_dim=3072, + num_experts=896, + top_k=16, + num_shared_experts=2, + vision_encoder=_vision_encoder_config( + text_dim=dim, + dim=1024, + qkv_dim=1536, + hidden_dim=4096, + num_layers=27, + num_heads=12, + init_pos_emb_height=64, + init_pos_emb_width=64, + ), + attn_backend=attn_backend, + ) + + +kimi_k3_configs = { + "debugmodel": _debugmodel, + "Kimi-K3": _kimi_k3, +} + + +def model_registry( + flavor: str, + attn_backend: str = "flex", + converters: list[ModelConfigConverter.Config] | None = None, +) -> ModelSpec: + config = kimi_k3_configs[flavor](attn_backend=attn_backend) + if converters is not None: + validate_converter_order(converters) + for converter in converters: + config = converter.build().convert(config) + return ModelSpec( + name="kimi_k3", + flavor=flavor, + model=config, + parallelize_fn=parallelize_kimi_k3, + pipelining_fn=None, + post_optimizer_build_fn=register_moe_load_balancing_hook, + state_dict_adapter=KimiK3StateDictAdapter, + ) diff --git a/torchtitan/models/kimi_k3/config_registry.py b/torchtitan/models/kimi_k3/config_registry.py new file mode 100644 index 0000000000..f1eac6b4fb --- /dev/null +++ b/torchtitan/models/kimi_k3/config_registry.py @@ -0,0 +1,96 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the BSD-style license found in the +# LICENSE file in the root directory of this source tree. + +from dataclasses import replace + +from torchtitan.components.checkpointer import CheckpointManager +from torchtitan.components.data import GrainDataLoader, SingleDatasetConfig +from torchtitan.components.loss import ChunkedLossWrapper, CrossEntropyLoss +from torchtitan.components.metrics import MetricsProcessor +from torchtitan.components.optimizer import default_adamw, LRSchedulersContainer +from torchtitan.components.tokenizer import MultiModalTokenizer +from torchtitan.config import ParallelismConfig, TrainingConfig +from torchtitan.distributed.activation_checkpoint import SelectiveAC +from torchtitan.hf_datasets.multimodal.mm_collator import MultiModalCollator +from torchtitan.hf_datasets.multimodal.mm_datasets import ( + MM_DATASETS, + MultiModalProcessor, +) +from torchtitan.hf_datasets.multimodal.utils.image import resize_to_patch_budget +from torchtitan.models.common.config_utils import decoder_vocab_size +from torchtitan.trainer import Trainer + +from . import KIMI_K3_SPECIAL_TOKENS, model_registry + + +def _kimi_k3_multimodal_dataloader( + dataset: SingleDatasetConfig, +) -> GrainDataLoader.Config: + processor = dataset.processor + if not isinstance(processor, MultiModalProcessor.Config): + raise ValueError("Kimi K3 multimodal data requires MultiModalProcessor.Config") + + processor = MultiModalProcessor.Config( + sample_processor=processor.sample_processor, + patch_size=14, + temporal_patch_size=1, + spatial_merge_size=2, + resize_fn=resize_to_patch_budget, + min_pixels=56 * 56, + max_pixels=224 * 224, + max_patches=256, + max_patches_per_side=16, + image_mean=(0.5, 0.5, 0.5), + image_std=(0.5, 0.5, 0.5), + ) + return GrainDataLoader.Config( + dataset=replace(dataset, processor=processor), + collator=MultiModalCollator.Config( + max_images_per_batch=8, + patch_size=processor.patch_size, + temporal_patch_size=processor.temporal_patch_size, + spatial_merge_size=processor.spatial_merge_size, + patch_order="raster", + build_mrope_positions=False, + ), + ) + + +def kimi_k3_debugmodel() -> Trainer.Config: + model_spec = model_registry("debugmodel") + return Trainer.Config( + loss=ChunkedLossWrapper.Config( + loss_fn=CrossEntropyLoss.Config( + global_vocab_size=decoder_vocab_size(model_spec), + ), + ), + hf_assets_path="./tests/assets/tokenizer", + tokenizer=MultiModalTokenizer.Config(**KIMI_K3_SPECIAL_TOKENS), + metrics=MetricsProcessor.Config(log_freq=1), + model_spec=model_spec, + dataloader=_kimi_k3_multimodal_dataloader(MM_DATASETS["cc12m-test"]), + optimizer=default_adamw(lr=8e-4), + lr_scheduler=LRSchedulersContainer.Config( + warmup_steps=2, + decay_ratio=0.8, + decay_type="linear", + min_lr_factor=0.0, + ), + # TODO: Kimi K3 has no spmd_types annotations yet. + parallelism=ParallelismConfig(spmd_backend="partial_dtensor"), + training=TrainingConfig( + num_tokens_per_microbatch_per_dp_rank=256, + max_context_length=256, + steps=10, + dtype="bfloat16", + disable_cuda_graphs=True, + ), + checkpoint=CheckpointManager.Config( + interval=10, + last_save_model_only=False, + ), + activation_checkpoint=SelectiveAC.Config(), + ) diff --git a/torchtitan/models/kimi_k3/kda.py b/torchtitan/models/kimi_k3/kda.py new file mode 100644 index 0000000000..f31a8eb9b1 --- /dev/null +++ b/torchtitan/models/kimi_k3/kda.py @@ -0,0 +1,175 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the BSD-style license found in the +# LICENSE file in the root directory of this source tree. + +"""Kimi Delta Attention modules for Kimi K3.""" + +from dataclasses import dataclass + +import torch +import torch.nn.functional as F +from fla.ops.kda import chunk_kda +from torch import nn + +from torchtitan.models.common import Conv1d, Linear +from torchtitan.models.common.attention import AttentionMasksType +from torchtitan.protocols.module import Module + +# Shape suffixes: +# T = packed tokens, D = model dimension, H = heads, +# K = key head dimension, V = value head dimension, C = projection channels. + + +class KimiRMSNormGated(Module): + """Per-head RMSNorm followed by a sigmoid output gate.""" + + @dataclass(kw_only=True, slots=True) + class Config(Module.Config): + dim: int + eps: float = 1e-5 + + def __init__(self, config: Config): + super().__init__() + self.eps = config.eps + self.weight = nn.Parameter(torch.empty(config.dim)) + + def forward(self, x: torch.Tensor, gate: torch.Tensor) -> torch.Tensor: + input_dtype = x.dtype + x_float = x.float() + variance = x_float.pow(2).mean(dim=-1, keepdim=True) + x_float = x_float * torch.rsqrt(variance + self.eps) + x_float = self.weight.float() * x_float + return (x_float * torch.sigmoid(gate.float())).to(input_dtype) + + +class KimiKDAKernel(Module): + """Stateless dispatch to FLA's chunked KDA kernel.""" + + @dataclass(kw_only=True, slots=True) + class Config(Module.Config): + lower_bound: float | None = -5.0 + + def __init__(self, config: Config): + super().__init__() + self.lower_bound = config.lower_bound + if self.lower_bound is not None and not (-5.0 <= self.lower_bound < 0.0): + raise ValueError("KDA lower_bound must be in the safe range [-5, 0).") + + def forward( + self, + q_BLHK: torch.Tensor, + k_BLHK: torch.Tensor, + v_BLHV: torch.Tensor, + gate_BLHK: torch.Tensor, + beta_BLH: torch.Tensor, + A_log_H: torch.Tensor, + dt_bias_HK: torch.Tensor, + ) -> torch.Tensor: + out_BLHV, _ = chunk_kda( + q_BLHK, + k_BLHK, + v_BLHV, + gate_BLHK, + beta_BLH, + A_log=A_log_H, + dt_bias=dt_bias_HK.reshape(-1), + use_qk_l2norm_in_kernel=True, + use_gate_in_kernel=True, + use_beta_sigmoid_in_kernel=True, + safe_gate=self.lower_bound is not None, + lower_bound=self.lower_bound, + ) + return out_BLHV + + +class KimiDeltaAttention(Module): + @dataclass(kw_only=True, slots=True) + class Config(Module.Config): + dim: int + num_heads: int + head_dim: int + conv_kernel_size: int + q_proj: Linear.Config + k_proj: Linear.Config + v_proj: Linear.Config + q_conv: Conv1d.Config + k_conv: Conv1d.Config + v_conv: Conv1d.Config + forget_a: Linear.Config + forget_b: Linear.Config + beta: Linear.Config + output_gate: Linear.Config + kernel: Module.Config + output_norm: KimiRMSNormGated.Config + output_proj: Linear.Config + + def __init__(self, config: Config): + super().__init__() + self.num_heads = config.num_heads + self.head_dim = config.head_dim + self.conv_kernel_size = config.conv_kernel_size + + self.q_proj = config.q_proj.build() + self.k_proj = config.k_proj.build() + self.v_proj = config.v_proj.build() + self.q_conv = config.q_conv.build() + self.k_conv = config.k_conv.build() + self.v_conv = config.v_conv.build() + self.forget_a = config.forget_a.build() + self.forget_b = config.forget_b.build() + self.beta = config.beta.build() + self.output_gate = config.output_gate.build() + self.kernel = config.kernel.build() + self.output_norm = config.output_norm.build() + self.output_proj = config.output_proj.build() + + self.A_log = nn.Parameter(torch.empty(config.num_heads)) + self.dt_bias = nn.Parameter(torch.empty(config.num_heads, config.head_dim)) + + def _causal_conv(self, x_TC: torch.Tensor, conv: Conv1d) -> torch.Tensor: + x_1CT = F.pad(x_TC.T.unsqueeze(0), (self.conv_kernel_size - 1, 0)) + return F.silu(conv(x_1CT)).squeeze(0).T + + def forward( + self, + x_TD: torch.Tensor, + attention_masks: AttentionMasksType | None = None, + positions: torch.Tensor | None = None, + ) -> torch.Tensor: + del positions + if attention_masks is not None: + raise NotImplementedError( + "Kimi K3 reference KDA does not support packed-document masks." + ) + + num_tokens = x_TD.shape[0] + q_THK = self._causal_conv(self.q_proj(x_TD), self.q_conv).view( + num_tokens, self.num_heads, self.head_dim + ) + k_THK = self._causal_conv(self.k_proj(x_TD), self.k_conv).view( + num_tokens, self.num_heads, self.head_dim + ) + v_THV = self._causal_conv(self.v_proj(x_TD), self.v_conv).view( + num_tokens, self.num_heads, self.head_dim + ) + forget_THK = self.forget_b(self.forget_a(x_TD)).view( + num_tokens, self.num_heads, self.head_dim + ) + beta_TH = self.beta(x_TD).float() + + out_THV = self.kernel( + q_THK.unsqueeze(0), + k_THK.unsqueeze(0), + v_THV.unsqueeze(0), + forget_THK.unsqueeze(0), + beta_TH.unsqueeze(0), + self.A_log, + self.dt_bias, + ).squeeze(0) + output_gate_THV = self.output_gate(x_TD).view( + num_tokens, self.num_heads, self.head_dim + ) + out_THV = self.output_norm(out_THV, output_gate_THV) + return self.output_proj(out_THV.reshape(num_tokens, -1)) diff --git a/torchtitan/models/kimi_k3/model.py b/torchtitan/models/kimi_k3/model.py new file mode 100644 index 0000000000..25f473ff21 --- /dev/null +++ b/torchtitan/models/kimi_k3/model.py @@ -0,0 +1,391 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the BSD-style license found in the +# LICENSE file in the root directory of this source tree. + +from dataclasses import dataclass, field + +import torch +from torch import nn + +from torchtitan.hf_datasets.multimodal.mm_datasets import MMSamplePackingConfig + +from torchtitan.models.common import Linear +from torchtitan.models.common.attention import ( + AttentionMasksType, + BaseAttention, + FlexAttention, +) +from torchtitan.models.common.decoder import Decoder +from torchtitan.models.common.multimodal import ( + get_vision_positions, + scatter_vision_embeds, +) +from torchtitan.models.common.nn_modules import RMSNorm +from torchtitan.models.utils import get_moe_model_nparams_and_flops +from torchtitan.protocols.module import Module + +from .kda import KimiDeltaAttention +from .moe import KimiFeedForward, KimiLatentMoE +from .vision_encoder import KimiK3VisionEncoder + +# Shape suffixes: +# T = packed tokens, D = model dimension, H = heads, +# K = key head dimension, V = value head dimension, +# N = attention-residual entries. + + +class KimiMLAAttention(BaseAttention): + """Kimi K3 multi-head latent attention. + + Unlike DeepSeek-V3 MLA, the released K3 configuration sets + ``mla_use_nope=True``: the RoPE-sized query/key slices remain part of the + projected head, but no rotary transform is applied, so this has no rope + config at all. Attention delegates to the configured inner backend. + """ + + @dataclass(kw_only=True, slots=True) + class Config(BaseAttention.Config): + dim: int + kv_lora_rank: int + qk_nope_head_dim: int + qk_rope_head_dim: int + v_head_dim: int + wq_a: Linear.Config + q_norm: RMSNorm.Config + wq_b: Linear.Config + wkv_a: Linear.Config + kv_norm: RMSNorm.Config + wkv_b: Linear.Config + gate: Linear.Config + wo: Linear.Config + inner_attention: Module.Config = field(default_factory=FlexAttention.Config) + + def __init__(self, config: Config): + super().__init__() + self.n_heads = config.n_heads + self.qk_nope_head_dim = config.qk_nope_head_dim + self.qk_rope_head_dim = config.qk_rope_head_dim + self.q_head_dim = config.qk_nope_head_dim + config.qk_rope_head_dim + self.v_head_dim = config.v_head_dim + self.kv_lora_rank = config.kv_lora_rank + self.scale = self.q_head_dim**-0.5 + + self.wq_a = config.wq_a.build() + self.q_norm = config.q_norm.build() + self.wq_b = config.wq_b.build() + self.wkv_a = config.wkv_a.build() + self.kv_norm = config.kv_norm.build() + self.wkv_b = config.wkv_b.build() + self.gate = config.gate.build() + self.wo = config.wo.build() + self.inner_attention = config.inner_attention.build() + + def forward( + self, + x_TD: torch.Tensor, + attention_masks: AttentionMasksType | None = None, + positions: torch.Tensor | None = None, + ) -> torch.Tensor: + del positions + + num_tokens = x_TD.shape[0] + q_THK = self.wq_b(self.q_norm(self.wq_a(x_TD))).view( + num_tokens, self.n_heads, self.q_head_dim + ) + + compressed_kv_TC = self.wkv_a(x_TD) + kv_latent_TC, k_rope_TK = torch.split( + compressed_kv_TC, + [self.kv_lora_rank, self.qk_rope_head_dim], + dim=-1, + ) + kv_THC = self.wkv_b(self.kv_norm(kv_latent_TC)).view( + num_tokens, + self.n_heads, + self.qk_nope_head_dim + self.v_head_dim, + ) + k_nope_THK, v_THV = torch.split( + kv_THC, + [self.qk_nope_head_dim, self.v_head_dim], + dim=-1, + ) + k_rope_THK = k_rope_TK.view(num_tokens, 1, self.qk_rope_head_dim).expand( + -1, self.n_heads, -1 + ) + k_THK = torch.cat((k_nope_THK, k_rope_THK), dim=-1) + + out_THV = self.inner_attention( + q_THK, + k_THK, + v_THV, + attention_masks=attention_masks, + scale=self.scale, + ) + out_TD = out_THV.reshape(num_tokens, self.n_heads * self.v_head_dim) + out_TD = out_TD * torch.sigmoid(self.gate(x_TD)) + return self.wo(out_TD) + + +def _apply_attention_residual( + prefix_sum_TD: torch.Tensor, + block_residual_TND: torch.Tensor, + projection: Linear, + norm: RMSNorm, +) -> torch.Tensor: + """Apply Kimi's block-level attention residual in FP32. + + TODO: Add TP Support. The current implementation assumes that the input tensors are on a single device. + """ + assert norm.eps is not None + + values_TND = torch.cat((block_residual_TND, prefix_sum_TD.unsqueeze(1)), dim=1) + values_float = values_TND.float() + variance = values_float.pow(2).mean(dim=-1, keepdim=True) + keys_TND = values_float * torch.rsqrt(variance + norm.eps) + score_weight_D = norm.weight.float() * projection.weight.squeeze(0).float() + scores_TN = (keys_TND * score_weight_D).sum(dim=-1) + probs_T1N = torch.softmax(scores_TN, dim=-1).unsqueeze(1) + output_TD = torch.matmul(probs_T1N, values_float).squeeze(1) + return output_TD.to(values_TND.dtype) + + +class KimiK3TransformerBlock(Module): + """Hybrid KDA/MLA decoder block with Kimi attention residuals.""" + + @dataclass(kw_only=True, slots=True) + class Config(Module.Config): + layer_id: int + attn_res_block_size: int + attention: KimiMLAAttention.Config | None + delta_attention: KimiDeltaAttention.Config | None + feed_forward: KimiFeedForward.Config | None + moe: KimiLatentMoE.Config | None + attention_norm: RMSNorm.Config + ffn_norm: RMSNorm.Config + attention_res_norm: RMSNorm.Config | None + attention_res_proj: Linear.Config | None + ffn_res_norm: RMSNorm.Config + ffn_res_proj: Linear.Config + + def __init__(self, config: Config): + super().__init__() + if (config.attention is None) == (config.delta_attention is None): + raise ValueError( + "Exactly one of attention or delta_attention must be configured." + ) + if (config.feed_forward is None) == (config.moe is None): + raise ValueError("Exactly one of feed_forward or moe must be configured.") + self.layer_id = config.layer_id + self.attn_res_block_size = config.attn_res_block_size + self.attention = ( + config.attention.build() if config.attention is not None else None + ) + self.delta_attention = ( + config.delta_attention.build() + if config.delta_attention is not None + else None + ) + self.feed_forward = ( + config.feed_forward.build() if config.feed_forward is not None else None + ) + self.moe = config.moe.build() if config.moe is not None else None + self.moe_enabled = self.moe is not None + self.attention_norm = config.attention_norm.build() + self.ffn_norm = config.ffn_norm.build() + self.attention_res_norm = ( + config.attention_res_norm.build() + if config.attention_res_norm is not None + else None + ) + self.attention_res_proj = ( + config.attention_res_proj.build() + if config.attention_res_proj is not None + else None + ) + self.ffn_res_norm = config.ffn_res_norm.build() + self.ffn_res_proj = config.ffn_res_proj.build() + + def forward( + self, + x_TD: torch.Tensor, + block_residual_TND: torch.Tensor, + attention_masks: AttentionMasksType | None = None, + positions: torch.Tensor | None = None, + ) -> tuple[torch.Tensor, torch.Tensor]: + prefix_sum_TD = x_TD + + if block_residual_TND.shape[1] > 0: + assert self.attention_res_proj is not None + assert self.attention_res_norm is not None + x_TD = _apply_attention_residual( + prefix_sum_TD, + block_residual_TND, + self.attention_res_proj, + self.attention_res_norm, + ) + + opens_block = self.layer_id % self.attn_res_block_size == 0 + if opens_block: + block_residual_TND = torch.cat( + ( + block_residual_TND, + prefix_sum_TD.unsqueeze(1), + ), + dim=1, + ) + + h_TD = self.attention_norm(x_TD) + if self.attention is not None: + h_TD = self.attention(h_TD, attention_masks, positions) + else: + assert self.delta_attention is not None + h_TD = self.delta_attention(h_TD, None, positions) + prefix_sum_TD = h_TD if opens_block else prefix_sum_TD + h_TD + + h_TD = _apply_attention_residual( + prefix_sum_TD, + block_residual_TND, + self.ffn_res_proj, + self.ffn_res_norm, + ) + h_TD = self.ffn_norm(h_TD) + if self.moe is not None: + h_TD = self.moe(h_TD) + else: + assert self.feed_forward is not None + h_TD = self.feed_forward(h_TD) + return prefix_sum_TD + h_TD, block_residual_TND + + +class KimiK3Model(Decoder): + @dataclass(kw_only=True, slots=True) + class Config(Decoder.Config): + layers: list[KimiK3TransformerBlock.Config] + output_res_norm: RMSNorm.Config + output_res_proj: Linear.Config + vision_encoder: KimiK3VisionEncoder.Config | None = None + + def update_from_config(self, *, config, **kwargs) -> None: + dataset = config.dataloader.dataset + # TODO: Support sample packing by resetting the Q/K/V causal-convolution + # and KDA recurrent states at document boundaries. + if isinstance(dataset, MMSamplePackingConfig): + raise ValueError("Kimi K3 does not yet support sample packing.") + Decoder.Config.update_from_config(self, config=config, **kwargs) + + def get_nparams_and_flops( + self, model: nn.Module, seq_len: int + ) -> tuple[int, int]: + attention_config = self.first_attention + if not isinstance(attention_config, KimiMLAAttention.Config): + raise ValueError( + "Kimi K3 requires at least one MLA layer for FLOP accounting." + ) + # KDA and the vision encoder have no dedicated term here, so their + # parameters only contribute the dense 6*N estimate; reported MFU is + # approximate. + return get_moe_model_nparams_and_flops( + self, + model, + attention_config.n_heads, + attention_config.qk_nope_head_dim + + attention_config.qk_rope_head_dim + + attention_config.v_head_dim, + seq_len, + ) + + def __init__(self, config: Config): + super().__init__(config) + self.output_res_norm = config.output_res_norm.build() + self.output_res_proj = config.output_res_proj.build() + self.vision_encoder = ( + config.vision_encoder.build() if config.vision_encoder is not None else None + ) + + def _prepare_multimodal_embeds( + self, + tokens: torch.Tensor, + *, + pixel_values: torch.Tensor | None, + grid_thw: torch.Tensor | None, + special_tokens: dict[str, int] | None, + ) -> torch.Tensor: + embeddings_TD = self.tok_embeddings(tokens) + if (pixel_values is None) != (grid_thw is None): + raise ValueError( + "pixel_values and grid_thw must either both be provided or " + "both be omitted." + ) + if pixel_values is None: + return embeddings_TD + assert grid_thw is not None + if self.vision_encoder is None: + raise ValueError("pixel_values were provided without a vision encoder.") + if special_tokens is None: + raise ValueError("special_tokens are required for multimodal inputs.") + + pixel_values = pixel_values.to(self.vision_encoder.patch_embed.weight.dtype) + vision_embeds = self.vision_encoder(pixel_values, grid_thw=grid_thw) + # MoonViT collapses time and merges spatially, so the text-side token + # count per item is (h/kh)*(w/kw), independent of t. + kernel_h, kernel_w = self.vision_encoder.merge_kernel_size + num_tokens_per_item = (grid_thw[:, 1] // kernel_h) * ( + grid_thw[:, 2] // kernel_w + ) + vision_positions = get_vision_positions( + tokens, + num_tokens_per_item, + special_tokens["image_id"], + ) + return scatter_vision_embeds( + embeddings_TD, + vision_embeds=vision_embeds, + vision_positions=vision_positions, + ) + + def forward( # pyrefly: ignore [bad-override] + self, + tokens: torch.Tensor, + *, + pixel_values: torch.Tensor | None = None, + grid_thw: torch.Tensor | None = None, + pixel_values_videos: torch.Tensor | None = None, + grid_thw_videos: torch.Tensor | None = None, + special_tokens: dict[str, int] | None = None, + positions: torch.Tensor | None = None, + attention_masks: AttentionMasksType | None = None, + ) -> torch.Tensor: + if pixel_values_videos is not None or grid_thw_videos is not None: + raise NotImplementedError("Kimi K3 v1 supports images but not videos.") + if self.tok_embeddings is not None: + h_TD = self._prepare_multimodal_embeds( + tokens, + pixel_values=pixel_values, + grid_thw=grid_thw, + special_tokens=special_tokens, + ) + else: + h_TD = tokens + + num_tokens, D = h_TD.shape + block_residual_TND = h_TD.new_zeros(num_tokens, 0, D) + for layer in self.layers.values(): + h_TD, block_residual_TND = layer( + h_TD, + block_residual_TND, + attention_masks, + positions, + ) + + h_TD = _apply_attention_residual( + h_TD, + block_residual_TND, + self.output_res_proj, + self.output_res_norm, + ) + h_TD = self.norm(h_TD) if self.norm is not None else h_TD + if self._skip_lm_head: + return h_TD + return self.lm_head(h_TD) if self.lm_head is not None else h_TD diff --git a/torchtitan/models/kimi_k3/moe.py b/torchtitan/models/kimi_k3/moe.py new file mode 100644 index 0000000000..990104a9d8 --- /dev/null +++ b/torchtitan/models/kimi_k3/moe.py @@ -0,0 +1,144 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the BSD-style license found in the +# LICENSE file in the root directory of this source tree. + +"""SiTU feed-forward and latent MoE modules for Kimi K3.""" + +from dataclasses import dataclass + +import torch +from torch.distributed.tensor import DTensor + +from torchtitan.models.common import Linear +from torchtitan.models.common.feed_forward import FeedForward +from torchtitan.models.common.moe import GroupedExperts, MoE +from torchtitan.models.common.nn_modules import RMSNorm + +# Shape suffixes: +# T = packed tokens, D = model dimension, E = experts, +# F = expert hidden dimension, R = routed tokens, K = selected experts per token. + + +def _situ_glu( + gate: torch.Tensor, + up: torch.Tensor, + beta: float, + linear_beta: float | None, +) -> torch.Tensor: + """Kimi's SiTU-GLU activation, evaluated in FP32.""" + input_dtype = gate.dtype + gate = gate.float() + up = up.float() + gate = beta * torch.tanh(gate / beta) * torch.sigmoid(gate) + if linear_beta is not None: + up = linear_beta * torch.tanh(up / linear_beta) + return (gate * up).to(input_dtype) + + +class KimiFeedForward(FeedForward): + """FeedForward with Kimi's SiTU activation.""" + + @dataclass(kw_only=True, slots=True) + class Config(FeedForward.Config): + beta: float = 1.0 + linear_beta: float | None = None + + def __init__(self, config: Config): + super().__init__(config) + self.beta = config.beta + self.linear_beta = config.linear_beta + + def forward(self, x: torch.Tensor) -> torch.Tensor: + return self.w2( + _situ_glu(self.w1(x), self.w3(x), self.beta, self.linear_beta), + ) + + +class KimiGroupedExperts(GroupedExperts): + """``common/moe.py::GroupedExperts`` with Kimi's SiTU activation.""" + + @dataclass(kw_only=True, slots=True) + class Config(GroupedExperts.Config): + beta: float = 1.0 + linear_beta: float | None = None + + def __init__(self, config: Config): + super().__init__(config) + self.beta = config.beta + self.linear_beta = config.linear_beta + + def forward( + self, + x_RD: torch.Tensor, + num_tokens_per_expert_E: torch.Tensor, + ) -> torch.Tensor: + if isinstance(self.w1_EFD, DTensor): + w1_EFD = self.w1_EFD.to_local() + assert isinstance(self.w2_EDF, DTensor) + w2_EDF = self.w2_EDF.to_local() + assert isinstance(self.w3_EFD, DTensor) + w3_EFD = self.w3_EFD.to_local() + else: + w1_EFD = self.w1_EFD + w2_EDF = self.w2_EDF + w3_EFD = self.w3_EFD + + offsets_E = torch.cumsum(num_tokens_per_expert_E, dim=0, dtype=torch.int32) + + gate_RF = self._grouped_mm( + A=x_RD.bfloat16(), + B_t=w1_EFD.bfloat16().transpose(-2, -1), + offs=offsets_E, + ) + up_RF = self._grouped_mm( + A=x_RD.bfloat16(), + B_t=w3_EFD.bfloat16().transpose(-2, -1), + offs=offsets_E, + ) + + h_RF = _situ_glu(gate_RF, up_RF, self.beta, self.linear_beta) + + return self._grouped_mm( + A=h_RF, + B_t=w2_EDF.bfloat16().transpose(-2, -1), + offs=offsets_E, + ).type_as(x_RD) + + +class KimiLatentMoE(MoE): + """``common/moe.py::MoE`` with Kimi's latent routed-expert path.""" + + @dataclass(kw_only=True, slots=True) + class Config(MoE.Config): + routed_down: Linear.Config + routed_norm: RMSNorm.Config + routed_up: Linear.Config + + def __init__(self, config: Config): + super().__init__(config) + self.routed_down = config.routed_down.build() + self.routed_norm = config.routed_norm.build() + self.routed_up = config.routed_up.build() + + def forward(self, x_TD: torch.Tensor) -> torch.Tensor: + weights_TK, expert_ids_TK, scores_TE = self.router(x_TD, self.expert_bias_E) + routing_map_TE = torch.zeros_like(scores_TE, dtype=torch.bool).scatter_( + -1, expert_ids_TK, True + ) + num_tokens_per_expert_E = routing_map_TE.sum(dim=0) + if self.training: + with torch.no_grad(): + self.tokens_per_expert_E.add_(num_tokens_per_expert_E) + + routed_TD = self.routed_experts( + self.routed_down(x_TD), + weights_TK, + expert_ids_TK, + num_tokens_per_expert_E, + ) + out_TD = self.routed_up(self.routed_norm(routed_TD)) + if self.shared_experts is not None: + out_TD = out_TD + self.shared_experts(x_TD) + return out_TD diff --git a/torchtitan/models/kimi_k3/parallelize.py b/torchtitan/models/kimi_k3/parallelize.py new file mode 100644 index 0000000000..8a7d604a02 --- /dev/null +++ b/torchtitan/models/kimi_k3/parallelize.py @@ -0,0 +1,97 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the BSD-style license found in the +# LICENSE file in the root directory of this source tree. + +import torch.nn as nn + +from torchtitan.config import ( + CompileConfig, + ParallelismConfig, + TORCH_DTYPE_MAP, + TrainingConfig, +) +from torchtitan.distributed import ParallelDims +from torchtitan.distributed.activation_checkpoint import ActivationCheckpointingConfig +from torchtitan.distributed.fsdp import ( + apply_fsdp_to_decoder, + apply_fsdp_to_vision_encoder, +) +from .model import KimiK3Model + + +def parallelize_kimi_k3( + model: nn.Module, + *, + parallel_dims: ParallelDims, + training: TrainingConfig, + parallelism: ParallelismConfig, + compile_config: CompileConfig, + ac_config: ActivationCheckpointingConfig, + dump_folder: str, +) -> nn.Module: + """Apply FSDP2 to the Kimi K3 decoder and vision encoder.""" + + unsupported_parallelisms = [ + name + for name, enabled in ( + ("tensor parallel", parallel_dims.tp_enabled), + ("pipeline parallel", parallel_dims.pp_enabled), + ("context parallel", parallel_dims.cp_enabled), + ("expert parallel", parallel_dims.ep_enabled), + ) + if enabled + ] + if unsupported_parallelisms: + raise NotImplementedError( + "Kimi K3 currently supports FSDP2 data parallelism " + f"only; disable {', '.join(unsupported_parallelisms)}." + ) + if parallelism.spmd_backend != "partial_dtensor": + raise NotImplementedError( + "Kimi K3 FSDP2 currently supports the partial_dtensor SPMD backend " + "only; the config registry pins it." + ) + if compile_config.enable and "model" in compile_config.components: + raise NotImplementedError("Kimi K3 does not support model compilation yet.") + + dp_mesh_names = ( + ["dp_replicate", "fsdp"] if parallel_dims.dp_replicate_enabled else ["fsdp"] + ) + dp_mesh = parallel_dims.get_mesh(dp_mesh_names) + + assert isinstance(model, KimiK3Model) + if ac_config is not None: + ac_policy = ac_config.build(dump_folder=dump_folder) + ac_policy.apply(model) + if model.vision_encoder is not None: + ac_policy.apply(model.vision_encoder) + + vision_encoder = model.vision_encoder + if vision_encoder is not None: + # TODO: An image batch on one DP rank and a text-only batch on another + # execute different FSDP collectives, deadlock, and hit a 90-second + # timeout. A general solution is needed. + apply_fsdp_to_vision_encoder( + vision_encoder, + dp_mesh, + param_dtype=TORCH_DTYPE_MAP[training.mixed_precision_param], + reduce_dtype=TORCH_DTYPE_MAP[training.mixed_precision_reduce], + reshard_after_forward_policy=parallelism.fsdp_reshard_after_forward, + pp_enabled=False, + ) + + apply_fsdp_to_decoder( + model, + dp_mesh, + param_dtype=TORCH_DTYPE_MAP[training.mixed_precision_param], + reduce_dtype=TORCH_DTYPE_MAP[training.mixed_precision_reduce], + pp_enabled=False, + cpu_offload=training.enable_cpu_offload, + reshard_after_forward_policy=parallelism.fsdp_reshard_after_forward, + ep_degree=1, + enable_symm_mem=parallelism.enable_fsdp_symm_mem, + ) + + return model diff --git a/torchtitan/models/kimi_k3/requirements.txt b/torchtitan/models/kimi_k3/requirements.txt new file mode 100644 index 0000000000..1c4403544a --- /dev/null +++ b/torchtitan/models/kimi_k3/requirements.txt @@ -0,0 +1 @@ +../../../.ci/docker/requirements-vlm.txt diff --git a/torchtitan/models/kimi_k3/state_dict_adapter.py b/torchtitan/models/kimi_k3/state_dict_adapter.py new file mode 100644 index 0000000000..ca28b3851c --- /dev/null +++ b/torchtitan/models/kimi_k3/state_dict_adapter.py @@ -0,0 +1,379 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the BSD-style license found in the +# LICENSE file in the root directory of this source tree. + +"""HuggingFace checkpoint adapter for unquantized Kimi K3 weights.""" + +import re +from typing import Any + +import torch +from torch.distributed.tensor import DTensor + +from torchtitan.models.utils import MoEStateDictAdapter + +from .model import KimiK3Model + + +_UNUSED_HF_LAYER_ZERO_ATTN_RES_KEYS = { + "language_model.model.layers.0.self_attention_res_norm.weight", + "language_model.model.layers.0.self_attention_res_proj.weight", +} + + +class KimiK3StateDictAdapter(MoEStateDictAdapter): + def __init__( + self, + model_config: KimiK3Model.Config, + hf_assets_path: str | None, + ): + super().__init__(model_config, hf_assets_path) + self.kimi_config = model_config + + self.from_hf_map = { + # Language model. + "language_model.model.embed_tokens.weight": "tok_embeddings.weight", + "language_model.model.layers.{}.input_layernorm.weight": "layers.{}.attention_norm.weight", + "language_model.model.layers.{}.post_attention_layernorm.weight": "layers.{}.ffn_norm.weight", + "language_model.model.layers.{}.self_attention_res_norm.weight": "layers.{}.attention_res_norm.weight", + "language_model.model.layers.{}.self_attention_res_proj.weight": "layers.{}.attention_res_proj.weight", + "language_model.model.layers.{}.mlp_res_norm.weight": "layers.{}.ffn_res_norm.weight", + "language_model.model.layers.{}.mlp_res_proj.weight": "layers.{}.ffn_res_proj.weight", + "language_model.model.layers.{}.mlp.gate_proj.weight": "layers.{}.feed_forward.w1.weight", + "language_model.model.layers.{}.mlp.up_proj.weight": "layers.{}.feed_forward.w3.weight", + "language_model.model.layers.{}.mlp.down_proj.weight": "layers.{}.feed_forward.w2.weight", + # MoE. + "language_model.model.layers.{}.block_sparse_moe.experts.{}.w1.weight": ( + "layers.{}.moe.routed_experts.inner_experts.w1_EFD" + ), + "language_model.model.layers.{}.block_sparse_moe.experts.{}.w2.weight": ( + "layers.{}.moe.routed_experts.inner_experts.w2_EDF" + ), + "language_model.model.layers.{}.block_sparse_moe.experts.{}.w3.weight": ( + "layers.{}.moe.routed_experts.inner_experts.w3_EFD" + ), + "language_model.model.layers.{}.block_sparse_moe.gate.weight": "layers.{}.moe.router.gate.weight", + "language_model.model.layers.{}.block_sparse_moe.gate.e_score_correction_bias": "layers.{}.moe.expert_bias_E", + "language_model.model.layers.{}.block_sparse_moe.routed_expert_down_proj.weight": "layers.{}.moe.routed_down.weight", + "language_model.model.layers.{}.block_sparse_moe.routed_expert_up_proj.weight": "layers.{}.moe.routed_up.weight", + "language_model.model.layers.{}.block_sparse_moe.routed_expert_norm.weight": "layers.{}.moe.routed_norm.weight", + "language_model.model.layers.{}.block_sparse_moe.shared_experts.gate_proj.weight": ( + "layers.{}.moe.shared_experts.w1.weight" + ), + "language_model.model.layers.{}.block_sparse_moe.shared_experts.up_proj.weight": ( + "layers.{}.moe.shared_experts.w3.weight" + ), + "language_model.model.layers.{}.block_sparse_moe.shared_experts.down_proj.weight": ( + "layers.{}.moe.shared_experts.w2.weight" + ), + "language_model.model.output_attn_res_norm.weight": "output_res_norm.weight", + "language_model.model.output_attn_res_proj.weight": "output_res_proj.weight", + "language_model.model.norm.weight": "norm.weight", + "language_model.lm_head.weight": "lm_head.weight", + # Vision encoder. + "vision_tower.patch_embed.proj.weight": "vision_encoder.patch_embed.weight", + "vision_tower.patch_embed.pos_emb.weight": "vision_encoder.pos_embed", + "vision_tower.encoder.blocks.{}.norm0.weight": "vision_encoder.layers.{}.norm1.weight", + "vision_tower.encoder.blocks.{}.norm1.weight": "vision_encoder.layers.{}.norm2.weight", + "vision_tower.encoder.blocks.{}.wo.weight": "vision_encoder.layers.{}.attn.proj.weight", + "vision_tower.encoder.blocks.{}.mlp.fc0.weight": "vision_encoder.layers.{}.mlp.linear_fc1.weight", + "vision_tower.encoder.blocks.{}.mlp.fc1.weight": "vision_encoder.layers.{}.mlp.linear_fc2.weight", + "vision_tower.encoder.final_layernorm.weight": "vision_encoder.final_norm.weight", + "mm_projector.proj.0.weight": "vision_encoder.projector.linear_1.weight", + "mm_projector.proj.2.weight": "vision_encoder.projector.linear_2.weight", + "mm_projector.post_norm.weight": "vision_encoder.projector.post_norm.weight", + } + self.mla_from_hf_map = { + "language_model.model.layers.{}.self_attn.q_a_proj.weight": "layers.{}.attention.wq_a.weight", + "language_model.model.layers.{}.self_attn.q_a_layernorm.weight": "layers.{}.attention.q_norm.weight", + "language_model.model.layers.{}.self_attn.q_b_proj.weight": "layers.{}.attention.wq_b.weight", + "language_model.model.layers.{}.self_attn.kv_a_proj_with_mqa.weight": "layers.{}.attention.wkv_a.weight", + "language_model.model.layers.{}.self_attn.kv_a_layernorm.weight": "layers.{}.attention.kv_norm.weight", + "language_model.model.layers.{}.self_attn.kv_b_proj.weight": "layers.{}.attention.wkv_b.weight", + "language_model.model.layers.{}.self_attn.g_proj.weight": "layers.{}.attention.gate.weight", + "language_model.model.layers.{}.self_attn.o_proj.weight": "layers.{}.attention.wo.weight", + } + self.kda_from_hf_map = { + "language_model.model.layers.{}.self_attn.q_proj.weight": "layers.{}.delta_attention.q_proj.weight", + "language_model.model.layers.{}.self_attn.k_proj.weight": "layers.{}.delta_attention.k_proj.weight", + "language_model.model.layers.{}.self_attn.v_proj.weight": "layers.{}.delta_attention.v_proj.weight", + "language_model.model.layers.{}.self_attn.q_conv1d.weight": "layers.{}.delta_attention.q_conv.weight", + "language_model.model.layers.{}.self_attn.k_conv1d.weight": "layers.{}.delta_attention.k_conv.weight", + "language_model.model.layers.{}.self_attn.v_conv1d.weight": "layers.{}.delta_attention.v_conv.weight", + "language_model.model.layers.{}.self_attn.f_a_proj.weight": "layers.{}.delta_attention.forget_a.weight", + "language_model.model.layers.{}.self_attn.f_b_proj.weight": "layers.{}.delta_attention.forget_b.weight", + "language_model.model.layers.{}.self_attn.b_proj.weight": "layers.{}.delta_attention.beta.weight", + "language_model.model.layers.{}.self_attn.g_proj.weight": "layers.{}.delta_attention.output_gate.weight", + "language_model.model.layers.{}.self_attn.o_norm.weight": "layers.{}.delta_attention.output_norm.weight", + "language_model.model.layers.{}.self_attn.o_proj.weight": "layers.{}.delta_attention.output_proj.weight", + "language_model.model.layers.{}.self_attn.A_log": "layers.{}.delta_attention.A_log", + "language_model.model.layers.{}.self_attn.dt_bias": "layers.{}.delta_attention.dt_bias", + } + + # The released index contains MXFP4 packed/scale FQNs, while this + # adapter exports unquantized weights. + self.fqn_to_index_mapping = None + + def _map_from_hf_layer_key( + self, + abstract_key: str, + layer_num: str, + ) -> str | None: + new_key = self.from_hf_map.get(abstract_key) + if new_key is not None: + return new_key + + layer_config = self.kimi_config.layers[int(layer_num)] + attention_map = ( + self.mla_from_hf_map + if layer_config.attention is not None + else self.kda_from_hf_map + ) + return attention_map.get(abstract_key) + + def to_hf(self, state_dict: dict[str, Any]) -> dict[str, Any]: + """Convert a TorchTitan state dict to unquantized HuggingFace format.""" + to_hf_map = { + tt_key: hf_key + for mapping in ( + self.from_hf_map, + self.mla_from_hf_map, + self.kda_from_hf_map, + ) + for hf_key, tt_key in mapping.items() + } + hf_state_dict: dict[str, Any] = {} + vision_qkv_by_layer: dict[str, dict[str, torch.Tensor]] = {} + unmapped: list[str] = [] + + for key, value in state_dict.items(): + if "moe.routed_experts.inner_experts" in key: + abstract_key = re.sub(r"(?<=\.)\d+(?=\.)", "{}", key, count=1) + layer_num_match = re.search(r"layers\.(\d+)\.", key) + assert layer_num_match is not None + layer_num = layer_num_match.group(1) + hf_abstract_key = to_hf_map.get(abstract_key) + if hf_abstract_key is None: + unmapped.append(key) + continue + + if isinstance(value, DTensor): + self.grouped_expert_weight_placements[ + abstract_key + ] = value.placements + self.grouped_expert_weight_shape[abstract_key] = value.shape + self.grouped_expert_weight_mesh[abstract_key] = value.device_mesh + hf_state_dict.update( + self._get_local_experts_weights( + hf_abstract_key, + abstract_key, + layer_num, + value, + ) + ) + else: + moe_config = self.kimi_config.layers[int(layer_num)].moe + assert moe_config is not None + split_values = self._split_experts_weights( + value, + moe_config.num_experts, + ) + for expert_num, expert_weight in enumerate(split_values): + hf_state_dict[ + hf_abstract_key.format(layer_num, expert_num) + ] = expert_weight.squeeze(0) + continue + + vision_qkv_match = re.fullmatch( + r"vision_encoder\.layers\.(\d+)\.attn\.w(q|k|v)\.weight", + key, + ) + if vision_qkv_match is not None: + layer_num, projection = vision_qkv_match.groups() + vision_qkv_by_layer.setdefault(layer_num, {})[projection] = value + continue + + layer_num_match = re.search(r"(?<=\.)\d+(?=\.)", key) + if layer_num_match is not None: + layer_num = layer_num_match.group(0) + abstract_key = re.sub( + r"(?<=\.)\d+(?=\.)", + "{}", + key, + count=1, + ) + hf_abstract_key = to_hf_map.get(abstract_key) + if hf_abstract_key is None: + unmapped.append(key) + continue + if abstract_key == "layers.{}.delta_attention.dt_bias": + value = value.reshape(-1) + hf_state_dict[hf_abstract_key.format(layer_num)] = value + continue + + hf_key = to_hf_map.get(key) + if hf_key is None: + unmapped.append(key) + continue + if key == "vision_encoder.patch_embed.weight": + vision_config = self.kimi_config.vision_encoder + if vision_config is None: + raise ValueError( + "Vision state was provided for a text-only config." + ) + value = value.reshape( + value.shape[0], + vision_config.in_channels, + vision_config.patch_size, + vision_config.patch_size, + ) + hf_state_dict[hf_key] = value + + for layer_num, qkv in vision_qkv_by_layer.items(): + missing = {"q", "k", "v"} - qkv.keys() + if missing: + raise ValueError( + f"Vision layer {layer_num} is missing QKV parts: {sorted(missing)}." + ) + hf_state_dict[ + f"vision_tower.encoder.blocks.{layer_num}.wqkv.weight" + ] = torch.cat((qkv["q"], qkv["k"], qkv["v"]), dim=0) + + # The released HF model contain these unused layer-0 attn res parameters. + # TT omits them, so synthesize deterministic, placeholders to preserve strict HF state-dict loading. + if self.kimi_config.layers[0].attention_res_norm is None: + norm_template_key = ( + "language_model.model.layers.1.self_attention_res_norm.weight" + ) + proj_template_key = ( + "language_model.model.layers.1.self_attention_res_proj.weight" + ) + hf_state_dict[ + "language_model.model.layers.0.self_attention_res_norm.weight" + ] = torch.ones_like(hf_state_dict[norm_template_key]) + hf_state_dict[ + "language_model.model.layers.0.self_attention_res_proj.weight" + ] = torch.zeros_like(hf_state_dict[proj_template_key]) + + if unmapped: + raise ValueError( + "KimiK3StateDictAdapter found TorchTitan keys without a " + f"mapping: {unmapped}." + ) + return hf_state_dict + + def from_hf(self, hf_state_dict: dict[str, Any]) -> dict[str, Any]: + """Convert an unquantized HuggingFace state dict to TorchTitan.""" + state_dict: dict[str, Any] = {} + expert_weights_by_layer: dict[str, dict[str, dict[int, torch.Tensor]]] = {} + unmapped: list[str] = [] + + for key, value in hf_state_dict.items(): + if key in _UNUSED_HF_LAYER_ZERO_ATTN_RES_KEYS: + continue + if key.endswith("rotary_emb.inv_freq"): + continue + + new_key = self.from_hf_map.get(key) + if new_key is not None: + if key == "vision_tower.patch_embed.proj.weight": + value = value.reshape(value.shape[0], -1) + state_dict[new_key] = value + continue + + if "block_sparse_moe.experts" in key: + abstract_key = re.sub( + r"(?<=\.)\d+(?=\.)", + "{}", + key, + count=2, + ) + indices = re.findall(r"(?<=\.)\d+(?=\.)", key) + if len(indices) != 2: + unmapped.append(key) + continue + layer_num, expert_num = indices + titan_abstract_key = self.from_hf_map.get(abstract_key) + if titan_abstract_key is None: + unmapped.append(key) + continue + new_key = titan_abstract_key.format(layer_num) + + experts = expert_weights_by_layer.setdefault(layer_num, {}).setdefault( + titan_abstract_key, {} + ) + experts[int(expert_num)] = value + + if titan_abstract_key in self.local_experts_indices: + stacked_value = self._concatenate_expert_weights_dtensor( + expert_weights_by_layer, + titan_abstract_key, + layer_num, + ) + else: + moe_config = self.kimi_config.layers[int(layer_num)].moe + assert moe_config is not None + stacked_value = self._concatenate_expert_weights( + expert_weights_by_layer, + titan_abstract_key, + layer_num, + moe_config.num_experts, + ) + if stacked_value is not None: + state_dict[new_key] = stacked_value + continue + + layer_num_match = re.search(r"(?<=\.)\d+(?=\.)", key) + if layer_num_match is not None: + layer_num = layer_num_match.group(0) + abstract_key = re.sub( + r"(?<=\.)\d+(?=\.)", + "{}", + key, + count=1, + ) + + if abstract_key == "vision_tower.encoder.blocks.{}.wqkv.weight": + q, k, v = torch.chunk(value, 3, dim=0) + base = f"vision_encoder.layers.{layer_num}.attn" + state_dict[f"{base}.wq.weight"] = q + state_dict[f"{base}.wk.weight"] = k + state_dict[f"{base}.wv.weight"] = v + continue + + new_abstract_key = ( + self._map_from_hf_layer_key(abstract_key, layer_num) + if key.startswith("language_model.model.layers.") + else self.from_hf_map.get(abstract_key) + ) + if new_abstract_key is None: + unmapped.append(key) + continue + if new_abstract_key == "layers.{}.delta_attention.dt_bias": + delta_config = self.kimi_config.layers[ + int(layer_num) + ].delta_attention + if delta_config is None: + raise ValueError(f"HF key '{key}' targets a non-KDA layer.") + value = value.reshape( + delta_config.num_heads, + delta_config.head_dim, + ) + state_dict[new_abstract_key.format(layer_num)] = value + continue + + unmapped.append(key) + + if unmapped: + raise ValueError( + "KimiK3StateDictAdapter found HuggingFace keys without a " + f"mapping: {unmapped}." + ) + if expert_weights_by_layer: + raise ValueError( + "KimiK3StateDictAdapter received an incomplete set of " + f"routed-expert weights: {expert_weights_by_layer.keys()}." + ) + return state_dict diff --git a/torchtitan/models/kimi_k3/vision_encoder.py b/torchtitan/models/kimi_k3/vision_encoder.py new file mode 100644 index 0000000000..6c480c8815 --- /dev/null +++ b/torchtitan/models/kimi_k3/vision_encoder.py @@ -0,0 +1,56 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the BSD-style license found in the +# LICENSE file in the root directory of this source tree. + +"""MoonViT3d vision encoder used by Kimi K3. + +Shape suffixes: +- M = total merged tokens +- F = merged feature dimension +- O = projected text dimension +""" + +from dataclasses import dataclass, field + +import torch + +from torchtitan.models.common import Linear +from torchtitan.models.common.nn_modules import GELU, RMSNorm +from torchtitan.models.kimi_k2_7.vision_encoder import MoonViTEncoder +from torchtitan.protocols.module import Module + + +class KimiK3VisionProjector(Module): + """PatchMergerMLPV2 projector from merged vision features to text width.""" + + @dataclass(kw_only=True, slots=True) + class Config(Module.Config): + linear_1: Linear.Config + linear_2: Linear.Config + post_norm: RMSNorm.Config + activation: GELU.Config = field(default_factory=GELU.Config) + + def __init__(self, config: Config): + super().__init__() + self.linear_1 = config.linear_1.build() + self.linear_2 = config.linear_2.build() + self.post_norm = config.post_norm.build() + self.activation = config.activation.build() + + def forward(self, merged_MF: torch.Tensor) -> torch.Tensor: + projected_MO = self.linear_2(self.activation(self.linear_1(merged_MF))) + return self.post_norm(projected_MO) + + +class KimiK3VisionEncoder(MoonViTEncoder): + @dataclass(kw_only=True, slots=True) + class Config(MoonViTEncoder.Config): + patch_size: int + in_channels: int + merge_kernel_size: tuple[int, int] # pyrefly: ignore [bad-override] + max_num_frames: int + + final_norm: RMSNorm.Config # pyrefly: ignore [bad-override] + projector: KimiK3VisionProjector.Config # pyrefly: ignore [bad-override] diff --git a/torchtitan_recipes/tests/models.py b/torchtitan_recipes/tests/models.py index 59a61e8b60..87cf8d3e78 100644 --- a/torchtitan_recipes/tests/models.py +++ b/torchtitan_recipes/tests/models.py @@ -215,6 +215,14 @@ def kimi_k2_5_debugmodel_muon_fsdp2_pp2_ep2() -> Trainer.Config: return config +def kimi_k3_debugmodel_mm_fsdp2() -> Trainer.Config: + from torchtitan.models.kimi_k3.config_registry import kimi_k3_debugmodel + + config = kimi_k3_debugmodel() + config.parallelism.data_parallel_shard_degree = 2 + return config + + def muse_glimmer_debugmodel_mm_fsdp2_tp2() -> Trainer.Config: from torchtitan.models.muse_glimmer.config_registry import ( muse_glimmer_debugmodel_mm,