Kimi-K3 full-model and LoRA RL support (#1825)

Co-authored-by: Zhichenzzz <zczeng@uw.edu>
This commit is contained in:
Yueming Yuan
2026-09-22 18:50:50 -07:00
committed by GitHub
co-authored by Zhichenzzz
parent b3c8f9c8f2
commit 8cdf3794d9
49 changed files with 3971 additions and 203 deletions
+45 -31
View File
@@ -29,49 +29,64 @@ as HF-named chunks over CUDA IPC.
## 2. Supported Variants
| Variant | Layers | Purpose | GPUs |
| `--model-name` | Layers | Purpose | GPUs |
|---|---|---|---|
| `full` | full stack | the real model | 64 (16 × 4), validated |
| `4layer` | 4 | smoke test, default | single node |
| `Kimi-K3` | 93 | the release | 64 (16 × 4), validated |
| `Kimi-K3-4layer-64experts` | 4 (1 dense + 3 MoE), 64 routed experts | smoke test and CI, default | one node |
| `Kimi-K3-4layer` | 4 (1 dense + 3 MoE), all 896 experts | rollout TP16 and EP layouts of the release | one node; two for rollout TP16 |
`--model-variant` selects between them and sets the matching checkpoint paths and
`megatron_model_type`.
The name sets the checkpoint paths under `--model-dir` and the `megatron_model_type`;
`--train-mode lora|full` picks the recipe.
Architecture, from `scripts/models/kimi-k3.sh`: hidden 7168, FFN 33792, 96 attention heads,
Architecture, from `scripts/models/kimi-k3.py`: hidden 7168, FFN 33792, 96 attention heads,
`kv_channels=256`, MLA with `q_lora_rank=1536` / `kv_lora_rank=512` /
`qk_head_dim=128` / `qk_pos_emb_head_dim=64` / `v_head_dim=128`, 896 experts at
`moe_ffn_hidden_size=3072`, shared expert 6144, vocab 163840, no position embedding.
## 3. Environment Setup
Use the `docker.io/radixark/miles:kimi-k3` image, which pins miles, SGLang (the
[`sglang-miles-k3`](https://github.com/sgl-project/sglang/tree/sglang-miles-k3) branch) and
flashinfer `0.6.15.post1` at the validated versions. On Hopper set
`SGLANG_K3_ATTN_RES_MODE=jit`.
Use the `radixark/miles:dev` image with the Megatron and SGLang changes from
[radixark/Megatron-LM#94](https://github.com/radixark/Megatron-LM/pull/94) and
[sgl-project/sglang#37704](https://github.com/sgl-project/sglang/pull/37704) until they merge.
The only external asset is the Kimi-K3 MXFP4 HF checkpoint. Everything else derives in-repo.
The only external asset is the native MXFP4 checkpoint; the BF16 dequantization and the
`torch_dist` conversion derive from it. `scripts/run_kimi_k3.py` names the three by model:
`{model_dir}/{model_name}`, `{model_dir}/{model_name}-bf16` and
`{model_dir}/{model_name}-bf16_torch_dist`, each overridable with `--hf-checkpoint`,
`--bf16-checkpoint` and `--ref-load`.
### 3.1 Data
### 3.1 Four-layer prune (one node)
```bash
python scripts/run_kimi_k3_lora.py prepare-data --task dapo-math --data-dir <datasets>
python scripts/run_kimi_k3.py prepare-download --model-name Kimi-K3-4layer-64experts --task gsm8k
python scripts/run_kimi_k3.py prepare-bf16 --model-name Kimi-K3-4layer-64experts
python scripts/run_kimi_k3.py prepare-torch-dist --model-name Kimi-K3-4layer-64experts
```
### 3.2 MXFP4 to BF16
`Pinaster/Kimi-K3-4layer` is the first dense layer plus three MoE layers of the release;
`Pinaster/Kimi-K3-4layer-64experts` keeps only the first 64 routed experts of each MoE layer
(the router is sliced to match) so full-parameter training fits one node's host memory. The
`run-ci-model-scripts` recipes train the 64-expert prune; the 896-expert prune is for layouts
that depend on the release's expert count.
### 3.2 Full model
Download the release into `{model_dir}/Kimi-K3` and dequantize it shard by shard across the
nodes (`--shard-rank/--num-shards`, then `--finalize-only` once):
```bash
python tools/convert_mxfp4_to_bf16.py --model-dir <native-mxfp4> --save-dir <bf16-hf>
python tools/convert_mxfp4_to_bf16.py --model-dir <native-mxfp4> --save-dir <bf16-hf> --device cuda \
--shard-rank $RANK --num-shards $NUM_NODES
```
### 3.3 BF16 to `torch_dist`
Unlike the bridge-mode recipes, K3 needs an offline conversion. Run it on 32 ranks; the
output re-shards at load, so the conversion layout does not have to match the training one:
The `torch_dist` conversion runs on 32 ranks; the output re-shards at load, so the conversion
layout does not have to match the training one:
```bash
source scripts/models/kimi-k3.sh # defines MODEL_ARGS
MODEL_ARGS_LINE="$(python3 miles/utils/external_utils/model_args_utils.py kimi-k3)" || exit 1
read -ra MODEL_ARGS <<< "${MODEL_ARGS_LINE}"
torchrun --nnodes=8 --nproc-per-node=4 ... \
tools/convert_hf_to_torch_dist.py "${MODEL_ARGS[@]}" \
tools/convert_hf_to_torch_dist.py ${MODEL_ARGS[@]} \
--hf-checkpoint <bf16-hf> --save <torch-dist-dcp> \
--bf16 --tensor-model-parallel-size 32 --sequence-parallel \
--pipeline-model-parallel-size 1 --context-parallel-size 1 \
@@ -79,17 +94,14 @@ torchrun --nnodes=8 --nproc-per-node=4 ... \
--megatron-to-hf-mode raw
```
Training then takes the **MXFP4** directory as `--hf-checkpoint` and the converted
`torch_dist` as `--ref-load`.
## 4. Launch
Validated on **16 nodes × 4 GPUs**. One container per node; bring up a ray cluster across
them, `export MILES_SCRIPT_EXTERNAL_RAY=1`, then:
`--train-mode lora` (default) or `full`. Validated on **16 nodes × 4 GPUs**: one container per
node, a ray cluster across them, `export MILES_SCRIPT_EXTERNAL_RAY=1`, then:
```bash
python scripts/run_kimi_k3_lora.py train \
--mode normal --model-variant full --task dapo-math --reward-model deepscaler \
python scripts/run_kimi_k3.py train \
--mode normal --model-name Kimi-K3 --train-mode lora --task dapo-math --reward-model deepscaler \
--num-nodes 16 --num-gpus-per-node 4 \
--pipeline-parallel-size 8 --context-parallel-size 2 \
--rollout-tp-size 16 --rollout-max-concurrency 8 \
@@ -104,9 +116,12 @@ python scripts/run_kimi_k3_lora.py train \
```
`--rollout-max-concurrency 8` is passed explicitly: the field default is 64, and the
validated runs pin 8.
validated runs pin 8. Off the validated 64 GPUs the full model needs `--tp-size-override`
(and `--ep-size-override`).
For a single-node smoke test, drop to the default `--model-variant 4layer`.
For a single-node smoke test use the default `--model-name Kimi-K3-4layer`: the trainer layout
follows the GPU count (TP8/EP8 on 8 GPUs) and the rollout TP/EP default to it; `--rollout-tp-size 16
--rollout-ep-size 1` on two nodes reproduces the TP16 Marlin layout of the full recipe.
## 5. Recipe Configuration
@@ -164,7 +179,6 @@ weight sync, the adapter export is leaking and the run will die later rather tha
- The image preloads a small `shm_unlink` shim through `/etc/ld.so.preload`. It tolerates a
benign PyTorch CUDA-IPC unlink race that otherwise aborts colocated weight sync at scale.
- On Hopper, set `SGLANG_K3_ATTN_RES_MODE=jit`.
- Weight conversion for K3 lives in
`miles/backends/megatron_utils/megatron_to_hf/kimi_k3.py`, and the model itself in
`miles_plugins/models/kimi_k3/`.
+59 -36
View File
@@ -423,6 +423,9 @@ def save_lora_checkpoint(
checkpoint resume without name/weight conversion. Each TP/PP rank saves its
own shard with original parameter names.
Raw-mode (native) adapters skip the HF PEFT export: they have no bridge, and Kimi K3's
896-expert adapter could not be materialized on every rank anyway.
When ``optimizer`` is provided, resume state (iteration + LR scheduler) is also
saved per-rank. ``--no-save-optim`` drops the optimizer entry from that file and
nothing else, so a resumed run still restarts at the right step and LR.
@@ -433,15 +436,14 @@ def save_lora_checkpoint(
"""
import json
from megatron.bridge import AutoBridge
from miles.utils import megatron_bridge_utils
save_path = Path(save_dir)
# raw-mode adapters have no bridge to export HF PEFT through
native = args.megatron_to_hf_mode == "raw"
parallel_state = get_parallel_state()
is_dp_cp_rank_0 = parallel_state.effective_dp.rank == 0 and parallel_state.cp.rank == 0
tp_rank = parallel_state.tp.rank
pp_rank = parallel_state.pp.rank
global_rank = dist.get_rank() if dist.is_initialized() else 0
save_path.mkdir(parents=True, exist_ok=True)
if dist.is_initialized():
@@ -453,53 +455,74 @@ def save_lora_checkpoint(
if _is_adapter_param_name(name):
adapter_state[name] = param.data.cpu()
global_rank = dist.get_rank() if dist.is_initialized() else 0
native_path = save_path / f"adapter_megatron_rank{global_rank}.pt"
torch.save(adapter_state, native_path)
logger.info(f"Saved {len(adapter_state)} adapter tensors (native) to {native_path}")
# ---- HF PEFT format (uses bridge for correct name/weight conversion) ----
# Bridge export is collective: all TP ranks participate in the all-gather,
# so every rank must call export_adapter_weights.
try:
bridge = AutoBridge.from_hf_pretrained(args.hf_checkpoint, trust_remote_code=True)
lora_state_dict: dict[str, torch.Tensor] = {}
with megatron_bridge_utils.patch_megatron_model(model):
for hf_name, weight, *_ in bridge.export_adapter_weights(
model,
cpu=True,
show_progress=False,
):
lora_state_dict[hf_name] = weight
if is_dp_cp_rank_0 and tp_rank == 0 and pp_rank == 0:
torch.save(lora_state_dict, save_path / "adapter_model.bin")
target_modules_hf = (
convert_target_modules_to_hf(list(args.target_modules))
if args.target_modules
else ["q_proj", "k_proj", "v_proj", "o_proj", "gate_proj", "up_proj", "down_proj"]
)
if native:
if global_rank == 0:
config = {
"peft_type": "LORA",
"r": args.lora_rank,
"lora_alpha": args.lora_alpha,
"target_modules": target_modules_hf,
"target_modules": convert_target_modules_to_hf(list(args.target_modules)),
"lora_dropout": args.lora_dropout,
"bias": "none",
"task_type": "CAUSAL_LM",
"experts_shared_outer_loras": bool(args.experts_shared_outer_loras),
"format": "megatron_rank_sharded",
}
with open(save_path / "adapter_config.json", "w") as f:
json.dump(config, f, indent=2)
os.sync()
logger.info(f"Saved HF PEFT adapter to {save_path} with {len(lora_state_dict)} tensors")
except Exception as hf_export_err:
logger.warning(
f"HF PEFT adapter export skipped ({hf_export_err}); the per-rank native "
f"shards + training state are sufficient for training resume."
)
logger.info(f"Saved rank-sharded adapter config to {save_path}")
else:
# ---- HF PEFT format (uses bridge for correct name/weight conversion) ----
# Bridge export is collective: all TP ranks participate in the all-gather,
# so every rank must call export_adapter_weights.
try:
from megatron.bridge import AutoBridge
from miles.utils import megatron_bridge_utils
bridge = AutoBridge.from_hf_pretrained(args.hf_checkpoint, trust_remote_code=True)
lora_state_dict: dict[str, torch.Tensor] = {}
with megatron_bridge_utils.patch_megatron_model(model):
for hf_name, weight, *_ in bridge.export_adapter_weights(
model,
cpu=True,
show_progress=False,
):
lora_state_dict[hf_name] = weight
if is_dp_cp_rank_0 and tp_rank == 0 and pp_rank == 0:
torch.save(lora_state_dict, save_path / "adapter_model.bin")
target_modules_hf = (
convert_target_modules_to_hf(list(args.target_modules))
if args.target_modules
else ["q_proj", "k_proj", "v_proj", "o_proj", "gate_proj", "up_proj", "down_proj"]
)
config = {
"peft_type": "LORA",
"r": args.lora_rank,
"lora_alpha": args.lora_alpha,
"target_modules": target_modules_hf,
"lora_dropout": args.lora_dropout,
"bias": "none",
"task_type": "CAUSAL_LM",
}
with open(save_path / "adapter_config.json", "w") as f:
json.dump(config, f, indent=2)
os.sync()
logger.info(f"Saved HF PEFT adapter to {save_path} with {len(lora_state_dict)} tensors")
except Exception as hf_export_err:
logger.warning(
f"HF PEFT adapter export skipped ({hf_export_err}); the per-rank native "
f"shards + training state are sufficient for training resume."
)
# ---- Training state (iteration + scheduler, and the optimizer unless opted out) ----
if optimizer is not None:
@@ -3,6 +3,7 @@ from .deepseekv4 import convert_deepseekv4_to_hf
from .glm4 import convert_glm4_to_hf
from .glm4moe import convert_glm4moe_to_hf
from .inkling import convert_inkling_to_hf
from .kimi_k3 import convert_kimi_k3_to_hf
from .kimi_vl import convert_kimi_k25_to_hf, convert_kimivl_to_hf
from .llama import convert_llama_to_hf
from .mimo import convert_mimo_to_hf
@@ -21,12 +22,12 @@ def postprocess_hf_param(args, megatron_param_name, hf_param_name, param):
# TODO optimize code details
def convert_to_hf(args, model_name, name, param, quantization_config=None):
def convert_to_hf(args, model_name, name, param, quantization_config=None, packed_weight_basenames=None):
param = remove_padding(name, param, args.vocab_size)
converted_named_tensors = _convert_to_hf_core(args, model_name, name, param)
return quantize_params(args, name, converted_named_tensors, quantization_config)
return quantize_params(args, name, converted_named_tensors, quantization_config, packed_weight_basenames)
# TODO optimize code details
@@ -59,6 +60,8 @@ def _convert_to_hf_core(args, model_name, name, param):
converted_named_tensors = convert_llama_to_hf(args, name, param)
elif "mimo" in model_name:
converted_named_tensors = convert_mimo_to_hf(args, name, param)
elif "kimi_k3" in model_name:
converted_named_tensors = convert_kimi_k3_to_hf(args, name, param)
elif "kimivl" in model_name:
converted_named_tensors = convert_kimivl_to_hf(args, name, param)
elif "kimi_k25" in model_name:
@@ -0,0 +1,97 @@
import re
_PREFIX = "language_model.model"
_DIRECT_NAMES = {
"module.module.embedding.word_embeddings.weight": f"{_PREFIX}.embed_tokens.weight",
"module.module.decoder.final_layernorm.weight": f"{_PREFIX}.norm.weight",
"module.module.output_layer.weight": "language_model.lm_head.weight",
}
_OUTPUT_HEAD_NAMES = {
"output_attn_res_norm.weight": f"{_PREFIX}.output_attn_res_norm.weight",
"output_attn_res_proj.weight": f"{_PREFIX}.output_attn_res_proj.weight",
}
_LAYER_NAMES = {
"input_layernorm.weight": "input_layernorm.weight",
"pre_mlp_layernorm.weight": "post_attention_layernorm.weight",
"self_attention_res_norm.weight": "self_attention_res_norm.weight",
"self_attention_res_proj.weight": "self_attention_res_proj.weight",
"mlp_res_norm.weight": "mlp_res_norm.weight",
"mlp_res_proj.weight": "mlp_res_proj.weight",
"self_attention.q_proj.weight": "self_attn.q_proj.weight",
"self_attention.k_proj.weight": "self_attn.k_proj.weight",
"self_attention.v_proj.weight": "self_attn.v_proj.weight",
"self_attention.q_conv1d.weight": "self_attn.q_conv1d.weight",
"self_attention.k_conv1d.weight": "self_attn.k_conv1d.weight",
"self_attention.v_conv1d.weight": "self_attn.v_conv1d.weight",
"self_attention.A_log": "self_attn.A_log",
"self_attention.dt_bias": "self_attn.dt_bias",
"self_attention.f_a_proj.weight": "self_attn.f_a_proj.weight",
"self_attention.f_b_proj.weight": "self_attn.f_b_proj.weight",
"self_attention.b_proj.weight": "self_attn.b_proj.weight",
"self_attention.g_proj.weight": "self_attn.g_proj.weight",
"self_attention.o_norm.weight": "self_attn.o_norm.weight",
"self_attention.o_proj.weight": "self_attn.o_proj.weight",
"self_attention.q_a_proj.weight": "self_attn.q_a_proj.weight",
"self_attention.q_a_layernorm.weight": "self_attn.q_a_layernorm.weight",
"self_attention.q_b_proj.weight": "self_attn.q_b_proj.weight",
"self_attention.kv_a_proj_with_mqa.weight": "self_attn.kv_a_proj_with_mqa.weight",
"self_attention.kv_a_layernorm.weight": "self_attn.kv_a_layernorm.weight",
"self_attention.kv_b_proj.weight": "self_attn.kv_b_proj.weight",
"mlp.linear_fc2.weight": "mlp.down_proj.weight",
"mlp.router.weight": "block_sparse_moe.gate.weight",
"mlp.router.expert_bias": "block_sparse_moe.gate.e_score_correction_bias",
"mlp.fc1_latent_proj.weight": "block_sparse_moe.routed_expert_down_proj.weight",
"mlp.routed_expert_norm.weight": "block_sparse_moe.routed_expert_norm.weight",
"mlp.fc2_latent_proj.weight": "block_sparse_moe.routed_expert_up_proj.weight",
"mlp.shared_experts.linear_fc2.weight": "block_sparse_moe.shared_experts.down_proj.weight",
}
def convert_kimi_k3_to_hf(args, name, param):
del args
if name in _DIRECT_NAMES:
return [(_DIRECT_NAMES[name], param)]
match = re.fullmatch(r"module\.module\.decoder\.layers\.(\d+)\.(.+)", name)
if match is None:
raise ValueError(f"Unknown Kimi K3 parameter name: {name}")
layer_idx, rest = match.groups()
layer_prefix = f"{_PREFIX}.layers.{layer_idx}"
expert_match = re.fullmatch(r"mlp\.experts\.(linear_fc[12])\.weight(\d+)", rest)
if expert_match is not None:
projection, expert_idx = expert_match.groups()
expert_prefix = f"{layer_prefix}.block_sparse_moe.experts.{expert_idx}"
if projection == "linear_fc1":
gate, up = param.chunk(2, dim=0)
return [
(f"{expert_prefix}.w1.weight", gate),
(f"{expert_prefix}.w3.weight", up),
]
return [(f"{expert_prefix}.w2.weight", param)]
if rest == "mlp.linear_fc1.weight":
gate, up = param.chunk(2, dim=0)
return [
(f"{layer_prefix}.mlp.gate_proj.weight", gate),
(f"{layer_prefix}.mlp.up_proj.weight", up),
]
if rest == "mlp.shared_experts.linear_fc1.weight":
gate, up = param.chunk(2, dim=0)
shared_prefix = f"{layer_prefix}.block_sparse_moe.shared_experts"
return [
(f"{shared_prefix}.gate_proj.weight", gate),
(f"{shared_prefix}.up_proj.weight", up),
]
if rest in _OUTPUT_HEAD_NAMES:
return [(_OUTPUT_HEAD_NAMES[rest], param)]
if rest not in _LAYER_NAMES:
raise ValueError(f"Unknown Kimi K3 layer parameter name: {name}")
return [(f"{layer_prefix}.{_LAYER_NAMES[rest]}", param)]
@@ -14,7 +14,7 @@ __all__ = [
]
def quantize_params(args, megatron_name, converted_named_params, quantization_config):
def quantize_params(args, megatron_name, converted_named_params, quantization_config, packed_weight_basenames=None):
if quantization_config is None:
return converted_named_params
elif quantization_config["quant_method"] == "fp8":
@@ -24,5 +24,7 @@ def quantize_params(args, megatron_name, converted_named_params, quantization_co
elif quantization_config.get("quant_algo") == "NVFP4" or quantization_config["quant_method"] == "nvfp4":
return quantize_params_nvfp4(args, megatron_name, converted_named_params, quantization_config)
elif quantization_config["quant_method"] == "compressed-tensors":
# only int4 at the moment.
return quantize_params_compressed_tensors(converted_named_params, quantization_config)
assert (
packed_weight_basenames is not None
), "compressed-tensors quantization needs the checkpoint's packed weight names"
return quantize_params_compressed_tensors(converted_named_params, quantization_config, packed_weight_basenames)
@@ -5,6 +5,8 @@ import re
import torch
import torch.nn as nn
from miles.utils.mxfp4 import quantize_mxfp4
try:
import fake_int4_quant_cuda
except ImportError:
@@ -263,43 +265,39 @@ def pack_layer(weight, group_size, sym=True):
return packed_weight, scale, packed_zp
def quantize_params_compressed_tensors(converted_named_params, quantization_config):
def quantize_params_compressed_tensors(converted_named_params, quantization_config, packed_weight_basenames):
quant_format = quantization_config["format"]
w_cfg = quantization_config["config_groups"]["group_0"]["weights"]
group_size = w_cfg["group_size"]
is_symmetric = w_cfg["symmetric"]
ignore_rules = quantization_config.get("ignore", [])
# Base names of params the checkpoint actually stores packed (see
# HfWeightIteratorBridge). The published ignore list of multimodal
# checkpoints (e.g. Kimi-K2.5 VL) only covers LLM submodules, so relying
# on it alone would wrongly quantize the vision tower / projector.
quantized_basenames = quantization_config.get("_miles_quantized_basenames")
is_mxfp4 = quant_format == "mxfp4-pack-quantized"
if is_mxfp4:
assert w_cfg["type"] == "float"
assert w_cfg["num_bits"] == 4
assert w_cfg["scale_dtype"] == "torch.uint8"
assert is_symmetric
results = []
for name, param in converted_named_params:
if quantized_basenames is not None:
should_quantize = name.endswith(".weight") and name.removesuffix(".weight") in quantized_basenames
else:
is_ignored = any(
(r.startswith("re:") and re.match(r[3:], name)) or r == name or name.startswith(r)
for r in ignore_rules
)
should_quantize = not is_ignored and name.endswith(".weight") and param.dim() >= 2
if not should_quantize:
if not (name.endswith(".weight") and name.removesuffix(".weight") in packed_weight_basenames):
results.append((name, param))
continue
qw, s, zp = pack_layer(param, group_size, is_symmetric)
qweight_name = name.replace(".weight", ".weight_packed")
scale_name = name.replace(".weight", ".weight_scale")
weight_shape = torch.tensor(param.shape, dtype=torch.int32, device="cuda")
weight_shape_name = name.replace(".weight", ".weight_shape")
if zp is not None:
zp_name = name.replace(".weight", ".weight_zero_point")
results.append((zp_name, zp))
results.append((qweight_name, qw))
results.append((scale_name, s))
results.append((weight_shape_name, weight_shape))
if is_mxfp4:
qw, s = quantize_mxfp4(param, group_size)
results.append((qweight_name, qw))
results.append((scale_name, s))
else:
qw, s, zp = pack_layer(param, group_size, is_symmetric)
weight_shape = torch.tensor(param.shape, dtype=torch.int32, device="cuda")
weight_shape_name = name.replace(".weight", ".weight_shape")
if zp is not None:
zp_name = name.replace(".weight", ".weight_zero_point")
results.append((zp_name, zp))
results.append((qweight_name, qw))
results.append((scale_name, s))
results.append((weight_shape_name, weight_shape))
return results
+15 -7
View File
@@ -63,6 +63,7 @@ def _has_loadable_ckpt(load_dir: str | None) -> bool:
return bool(load_dir) and Path(load_dir).is_dir() and any(Path(load_dir).iterdir())
from .fp32_param_utils import enforce_marked_param_dtypes
from .lora.bridge import _ensure_model_list, _setup_lora_model_via_bridge # noqa: F401
@@ -150,16 +151,23 @@ def setup_model_and_optimizer(
is_lora_enabled(args) and role == "actor" and args.megatron_to_hf_mode == "bridge"
):
model = _setup_lora_model_via_bridge(args)
enforce_marked_param_dtypes(model)
else:
provider_func = get_model_provider_func(args, role)
if (
is_lora_enabled(args)
and role == "actor"
and "inkling" in (getattr(args, "custom_model_provider_path", None) or "")
):
from miles_plugins.models.inkling.lora import wrap_model_provider_with_inkling_lora
if is_lora_enabled(args) and role == "actor":
if "inkling" in (getattr(args, "custom_model_provider_path", None) or ""):
from miles_plugins.models.inkling.lora import wrap_model_provider_with_inkling_lora
provider_func = wrap_model_provider_with_inkling_lora(provider_func, args)
provider_func = wrap_model_provider_with_inkling_lora(provider_func, args)
# TODO: will rewrite in native lora refactor
elif "kimi_k3" in (args.model_name or "").lower():
from miles_plugins.models.kimi_k3.lora import wrap_model_provider_with_kimi_k3_lora
from .lora.utils import patch_param_grad_buffer_for_colocate_mode_lora
provider_func = wrap_model_provider_with_kimi_k3_lora(provider_func, args)
if args.offload_train:
patch_param_grad_buffer_for_colocate_mode_lora()
model = get_model(provider_func, ModelType.encoder_or_decoder)
if args.debug_disable_optimizer:
@@ -1,8 +1,10 @@
"""Megatron implementations' shared base and factory for the backend-neutral
HF weight iterator API."""
import json
import logging
import math
import os
from abc import abstractmethod
from argparse import Namespace
from collections.abc import Sequence
@@ -28,6 +30,12 @@ class MegatronHfWeightIteratorBase(HfWeightIteratorBase):
def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
self.packed_weight_basenames = (
get_packed_weight_basenames(self.args.hf_checkpoint)
if self.quantization_config is not None
and self.quantization_config["quant_method"] == "compressed-tensors"
else None
)
trainer_has_mtp = bool(unwrap_model(self.model)[0].config.mtp_num_layers)
if self.args.sglang_speculative_algorithm and not trainer_has_mtp:
self.weight_update_selector = "target"
@@ -47,9 +55,7 @@ class MegatronHfWeightIteratorBase(HfWeightIteratorBase):
return
if not named_tensors:
raise RuntimeError(
f"LoRA weight sync failed: the adapter export produced zero tensors"
f"{f' for adapter {adapter!r}' if adapter is not None else ''}. "
"This usually means the Megatron-Bridge or SGLang version is incompatible."
f"LoRA weight sync failed: the adapter export produced zero tensors{f' for adapter {adapter!r}' if adapter is not None else ''}. This usually means the Megatron-Bridge or SGLang version is incompatible."
)
if not any(is_lora_weight_name(name) for name, _tensor in named_tensors):
raise RuntimeError("LoRA weight sync failed: the adapter export contains no lora_A/lora_B names.")
@@ -87,6 +93,15 @@ def get_hf_weight_iterator(
)
def get_packed_weight_basenames(hf_checkpoint: str) -> set[str]:
"""Base names the checkpoint stores as compressed-tensors `weight_packed`; the quantizer
re-quantizes exactly these, since the published `ignore` list is written for loaders and
leaves out BF16 weights such as routers, residual projections and the vision tower."""
with open(os.path.join(hf_checkpoint, "model.safetensors.index.json")) as index_file:
names = json.load(index_file)["weight_map"]
return {n.removesuffix(".weight_packed") for n in names if n.endswith(".weight_packed")}
def _gather_pp_full_adapter(
hf_named_tensors: Sequence[tuple[str, torch.Tensor]],
) -> list[tuple[str, torch.Tensor]]:
@@ -1,8 +1,6 @@
import dataclasses
import inspect
import itertools
import json
import os
from miles.backends.megatron_utils.update_weight.hf_weight_iterator import (
MegatronHfWeightIteratorBase,
@@ -24,21 +22,6 @@ class HfWeightIteratorBridge(MegatronHfWeightIteratorBase):
self._bridge = AutoBridge.from_hf_pretrained(self.args.hf_checkpoint, trust_remote_code=True)
if (
self.quantization_config is not None
and self.quantization_config.get("quant_method") == "compressed-tensors"
):
quantized_basenames = _load_quantized_param_basenames(self.args.hf_checkpoint)
if quantized_basenames is not None:
# Quantize exactly the params the checkpoint stores packed; the
# published ignore list of multimodal checkpoints (e.g.
# Kimi-K2.5 VL) omits vision_tower/mm_projector, so it cannot
# be trusted as the sole quantization criterion.
self.quantization_config = {
**self.quantization_config,
"_miles_quantized_basenames": quantized_basenames,
}
def _iter_hf_param_units(self, weights, *, materialize):
renamed_megatron_local_weights = {strip_param_name_prefix(k): v for k, v in weights.items()}
with megatron_bridge_utils.patch_megatron_model(self.model):
@@ -129,23 +112,17 @@ class HfWeightIteratorBridge(MegatronHfWeightIteratorBase):
# A tensor with no Megatron source (HF-only passthrough) is not a trainable weight: pass it through.
qmegatron_name = f"module.module.{megatron_param_name}"
for q_hf_name, q_weight in quantize_params(
self.args, qmegatron_name, [(hf_name, weight)], self.quantization_config
self.args,
qmegatron_name,
[(hf_name, weight)],
self.quantization_config,
self.packed_weight_basenames,
):
yield q_hf_name, q_weight, megatron_param_names
else:
yield hf_name, weight, megatron_param_names
def _load_quantized_param_basenames(hf_checkpoint):
"""Base names of params stored packed (`<base>.weight_packed`) in the checkpoint, or None if unknown."""
index_path = os.path.join(hf_checkpoint, "model.safetensors.index.json")
if not os.path.exists(index_path):
return None
with open(index_path) as f:
names = json.load(f)["weight_map"]
return {n.removesuffix(".weight_packed") for n in names if n.endswith(".weight_packed")}
def _process_conversion_tasks(vanilla_conversion_tasks, new_weight_dict):
def _handle_one(task):
if task is None:
@@ -62,13 +62,24 @@ class HfWeightIteratorDirect(MegatronHfWeightIteratorBase):
def _export_pp_local_lora(self, adapter):
assert adapter is None, "multi-LoRA export requires --megatron-to-hf-mode bridge"
from miles_plugins.models.inkling.lora import export_inkling_lora_hf_named
# TODO: will rewrite in native lora refactor
if "kimi_k3" in self.model_name.lower():
from miles_plugins.models.kimi_k3.lora import export_kimi_k3_lora_hf_chunks
return export_inkling_lora_hf_named(self.model)
return [named_tensor for chunk in export_kimi_k3_lora_hf_chunks(self.model) for named_tensor in chunk]
if "inkling" in (self.args.custom_model_provider_path or ""):
from miles_plugins.models.inkling.lora import export_inkling_lora_hf_named
return export_inkling_lora_hf_named(self.model)
raise NotImplementedError(f"Raw LoRA export is not implemented for model {self.model_name!r}")
def _convert_to_hf_param_units(self, named_params: Sequence[tuple[str, torch.Tensor]]):
for name, param in named_params:
yield list(convert_to_hf(self.args, self.model_name, name, param, self.quantization_config))
yield list(
convert_to_hf(
self.args, self.model_name, name, param, self.quantization_config, self.packed_weight_basenames
)
)
def _load_or_allocate_params(param_infos: Sequence[ParamInfo], megatron_local_weights) -> list[torch.Tensor]:
+5 -8
View File
@@ -36,14 +36,11 @@ def _grade_boxed_solution(model_solution, label):
def get_deepscaler_rule_based_reward(response, label):
if "</think>" in response:
model_solution = response.split("</think>")[-1]
elif "###Response" in response:
model_solution = response.split("###Response")[1]
else:
return 0
return _grade_boxed_solution(model_solution, label)
# markers that end the reasoning segment; Kimi K-series uses the tagged <|open|>response<|sep|>
for closer in ("<|open|>response<|sep|>", "</think>", "###Response"):
if closer in response:
return _grade_boxed_solution(response.split(closer)[-1], label)
return 0
def get_gemma_math_reward(response, label):
+4 -1
View File
@@ -1827,7 +1827,10 @@ def get_miles_extra_args_provider(add_custom_arguments=None):
help=(
"LoRA + colocate: keep SGLang-side CPU mirror of base weights "
"and skip per-step base sync. Trades host RAM for faster "
"onload/offload. Ignored unless --colocate and LoRA are both on."
"onload/offload. Ignored unless --colocate and LoRA are both on. "
"Also needs 'weight' in --offload-rollout-level: SGLang populates "
"the mirror during release_weights_occupation, so with the weights "
"never released the mirror is never built and the flag does nothing."
),
)
parser.add_argument(
+59
View File
@@ -0,0 +1,59 @@
"""MXFP4 (E2M1 elements + E8M0 block scales) pack/unpack.
Torch-only, so checkpoint tooling can use it without importing Megatron.
"""
import torch
_E2M1_VALUES = (0.0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0)
def quantize_mxfp4(weight, group_size):
assert weight.shape[-1] % group_size == 0
assert weight.shape[-1] % 2 == 0
blocks = weight.reshape(-1, group_size)
amax = blocks.abs().amax(dim=-1, keepdim=True).float()
scale_exp = torch.ceil(torch.log2(amax / 6.0)).clamp_(-127, 127)
normalized = blocks.float() * torch.exp2(-scale_exp)
magnitude = torch.zeros_like(normalized, dtype=torch.uint8)
for bound in (0.25, 0.75, 1.25, 1.75, 2.5, 3.5, 5.0):
magnitude.add_(normalized.abs() > bound)
encoded = magnitude | (torch.signbit(normalized).to(torch.uint8) << 3)
packed = encoded[:, 0::2] | (encoded[:, 1::2] << 4)
packed = packed.reshape(*weight.shape[:-1], weight.shape[-1] // 2).contiguous()
scale = (scale_exp + 127).to(torch.uint8)
scale = scale.reshape(*weight.shape[:-1], weight.shape[-1] // group_size).contiguous()
return packed, scale
def dequantize_mxfp4(
weight_packed: torch.Tensor,
weight_scale: torch.Tensor,
group_size: int,
) -> torch.Tensor:
assert weight_packed.dtype == torch.uint8
assert weight_scale.dtype == torch.uint8
assert group_size > 0
unpacked = torch.empty(
*weight_packed.shape[:-1],
weight_packed.shape[-1] * 2,
dtype=torch.uint8,
device=weight_packed.device,
)
unpacked[..., 0::2] = weight_packed & 0x0F
unpacked[..., 1::2] = (weight_packed >> 4) & 0x0F
signs = 1.0 - 2.0 * ((unpacked & 0b1000) >> 3).float()
magnitudes = unpacked & 0b0111
values = torch.tensor(_E2M1_VALUES, dtype=torch.float32, device=weight_packed.device)
dequantized = signs * values[magnitudes.long()]
assert dequantized.numel() % group_size == 0
assert weight_scale.numel() == dequantized.numel() // group_size
dequantized = dequantized.reshape(-1, group_size)
scales = torch.exp2(weight_scale.float().reshape(-1, 1) - 127.0)
return (dequantized * scales).reshape(unpacked.shape).to(torch.bfloat16).contiguous()
+2
View File
@@ -5,6 +5,7 @@ from .glm4moe import GLM4MoEBridge
from .glm4moe_lite import GLM4MoELiteBridge
from .inkling import InklingBridge
from .joyai_llm_flash import JoyAILLMFlashBridge
from .kimi_k3 import KimiK3Bridge
from .mimo import MimoBridge
from .qwen3_5 import Qwen3_5Bridge
from .qwen3_next import Qwen3NextBridge
@@ -20,4 +21,5 @@ __all__ = [
"DeepseekV4Bridge",
"JoyAILLMFlashBridge",
"InklingBridge",
"KimiK3Bridge",
]
+217
View File
@@ -0,0 +1,217 @@
import torch
from megatron.core.transformer import MLATransformerConfig
from megatron.core.transformer.enums import AttnBackend
from mbridge.core import register_model
from mbridge.core.safetensor_io import SafeTensorIO
from mbridge.models import DeepseekV3Bridge
from miles_plugins.models.kimi_k3.model import build_kimi_k3_spec
from miles_plugins.models.kimi_k3.ops import situ_and_mul
@register_model("kimi_k3")
class KimiK3Bridge(DeepseekV3Bridge):
TransformerConfigClass = MLATransformerConfig
_CONFIG_MAPPING = {
"num_layers": "num_hidden_layers",
"hidden_size": "hidden_size",
"num_attention_heads": "num_attention_heads",
"num_query_groups": "num_key_value_heads",
"ffn_hidden_size": "intermediate_size",
"attention_dropout": ("attention_dropout", 0.0),
"layernorm_epsilon": "rms_norm_eps",
"hidden_dropout": ("hidden_dropout", 0.0),
"kv_channels": "head_dim",
}
_DIRECT_MAPPING = {
"embedding.word_embeddings.weight": "language_model.model.embed_tokens.weight",
"decoder.final_layernorm.weight": "language_model.model.norm.weight",
"output_layer.weight": "language_model.lm_head.weight",
}
_ATTENTION_MAPPING = {
"input_layernorm.weight": ["language_model.model.layers.{layer_number}.input_layernorm.weight"],
"self_attention.q_proj.weight": ["language_model.model.layers.{layer_number}.self_attn.q_proj.weight"],
"self_attention.k_proj.weight": ["language_model.model.layers.{layer_number}.self_attn.k_proj.weight"],
"self_attention.v_proj.weight": ["language_model.model.layers.{layer_number}.self_attn.v_proj.weight"],
"self_attention.q_conv1d.weight": ["language_model.model.layers.{layer_number}.self_attn.q_conv1d.weight"],
"self_attention.k_conv1d.weight": ["language_model.model.layers.{layer_number}.self_attn.k_conv1d.weight"],
"self_attention.v_conv1d.weight": ["language_model.model.layers.{layer_number}.self_attn.v_conv1d.weight"],
"self_attention.A_log": ["language_model.model.layers.{layer_number}.self_attn.A_log"],
"self_attention.dt_bias": ["language_model.model.layers.{layer_number}.self_attn.dt_bias"],
"self_attention.f_a_proj.weight": ["language_model.model.layers.{layer_number}.self_attn.f_a_proj.weight"],
"self_attention.f_b_proj.weight": ["language_model.model.layers.{layer_number}.self_attn.f_b_proj.weight"],
"self_attention.b_proj.weight": ["language_model.model.layers.{layer_number}.self_attn.b_proj.weight"],
"self_attention.g_proj.weight": ["language_model.model.layers.{layer_number}.self_attn.g_proj.weight"],
"self_attention.o_norm.weight": ["language_model.model.layers.{layer_number}.self_attn.o_norm.weight"],
"self_attention.o_proj.weight": ["language_model.model.layers.{layer_number}.self_attn.o_proj.weight"],
"self_attention.q_a_proj.weight": ["language_model.model.layers.{layer_number}.self_attn.q_a_proj.weight"],
"self_attention.q_a_layernorm.weight": [
"language_model.model.layers.{layer_number}.self_attn.q_a_layernorm.weight"
],
"self_attention.q_b_proj.weight": ["language_model.model.layers.{layer_number}.self_attn.q_b_proj.weight"],
"self_attention.kv_a_proj_with_mqa.weight": [
"language_model.model.layers.{layer_number}.self_attn.kv_a_proj_with_mqa.weight"
],
"self_attention.kv_a_layernorm.weight": [
"language_model.model.layers.{layer_number}.self_attn.kv_a_layernorm.weight"
],
"self_attention.kv_b_proj.weight": ["language_model.model.layers.{layer_number}.self_attn.kv_b_proj.weight"],
"self_attention_res_norm.weight": [
"language_model.model.layers.{layer_number}.self_attention_res_norm.weight"
],
"self_attention_res_proj.weight": [
"language_model.model.layers.{layer_number}.self_attention_res_proj.weight"
],
}
_MLP_MAPPING = {
"pre_mlp_layernorm.weight": ["language_model.model.layers.{layer_number}.post_attention_layernorm.weight"],
"mlp_res_norm.weight": ["language_model.model.layers.{layer_number}.mlp_res_norm.weight"],
"mlp_res_proj.weight": ["language_model.model.layers.{layer_number}.mlp_res_proj.weight"],
"mlp.linear_fc1.weight": [
"language_model.model.layers.{layer_number}.mlp.gate_proj.weight",
"language_model.model.layers.{layer_number}.mlp.up_proj.weight",
],
"mlp.linear_fc2.weight": ["language_model.model.layers.{layer_number}.mlp.down_proj.weight"],
"mlp.router.weight": ["language_model.model.layers.{layer_number}.block_sparse_moe.gate.weight"],
"mlp.router.expert_bias": [
"language_model.model.layers.{layer_number}.block_sparse_moe.gate.e_score_correction_bias"
],
"mlp.fc1_latent_proj.weight": [
"language_model.model.layers.{layer_number}.block_sparse_moe.routed_expert_down_proj.weight"
],
"mlp.routed_expert_norm.weight": [
"language_model.model.layers.{layer_number}.block_sparse_moe.routed_expert_norm.weight"
],
"mlp.fc2_latent_proj.weight": [
"language_model.model.layers.{layer_number}.block_sparse_moe.routed_expert_up_proj.weight"
],
"mlp.shared_experts.linear_fc1.weight": [
"language_model.model.layers.{layer_number}.block_sparse_moe.shared_experts.gate_proj.weight",
"language_model.model.layers.{layer_number}.block_sparse_moe.shared_experts.up_proj.weight",
],
"mlp.shared_experts.linear_fc2.weight": [
"language_model.model.layers.{layer_number}.block_sparse_moe.shared_experts.down_proj.weight"
],
"mlp.experts.linear_fc1.weight": [
"language_model.model.layers.{layer_number}.block_sparse_moe.experts.{expert_id}.w1.weight",
"language_model.model.layers.{layer_number}.block_sparse_moe.experts.{expert_id}.w3.weight",
],
"mlp.experts.linear_fc2.weight": [
"language_model.model.layers.{layer_number}.block_sparse_moe.experts.{expert_id}.w2.weight"
],
}
_OTHER_MAPPING = {
"output_attn_res_norm.weight": ["language_model.model.output_attn_res_norm.weight"],
"output_attn_res_proj.weight": ["language_model.model.output_attn_res_proj.weight"],
}
@property
def text_config(self):
return self.hf_config.text_config
def _build_config(self):
hf_config = self.text_config
moe_layer_freq = [0] * hf_config.num_hidden_layers
for layer_idx in range(hf_config.first_k_dense_replace, hf_config.num_hidden_layers):
if layer_idx % hf_config.moe_layer_freq == 0:
moe_layer_freq[layer_idx] = 1
config = self._build_base_config(
text_config_key="text_config",
attention_backend=AttnBackend.auto,
multi_latent_attention=True,
qk_layernorm=True,
q_lora_rank=hf_config.q_lora_rank,
kv_lora_rank=hf_config.kv_lora_rank,
qk_head_dim=hf_config.qk_nope_head_dim,
qk_pos_emb_head_dim=hf_config.qk_rope_head_dim,
v_head_dim=hf_config.v_head_dim,
gated_activation_func=situ_and_mul,
bias_activation_fusion=False,
bias_dropout_fusion=False,
use_te_activation_func=False,
persist_layer_norm=True,
moe_ffn_hidden_size=hf_config.moe_intermediate_size,
moe_latent_size=hf_config.routed_expert_hidden_size,
moe_latent_use_norm=hf_config.latent_moe_use_norm,
num_moe_experts=hf_config.num_experts,
moe_router_topk=hf_config.num_experts_per_token,
moe_router_score_function=hf_config.moe_router_activation_func,
moe_router_pre_softmax=True,
moe_router_topk_scaling_factor=hf_config.routed_scaling_factor,
moe_router_enable_expert_bias=True,
freeze_e_score_correction_bias=True,
moe_router_bias_update_rate=0.0,
moe_router_dtype="fp32",
moe_router_load_balancing_type="none",
moe_aux_loss_coeff=0.0,
moe_grouped_gemm=True,
moe_shared_expert_intermediate_size=(hf_config.moe_intermediate_size * hf_config.num_shared_experts),
moe_shared_expert_overlap=False,
moe_layer_freq=moe_layer_freq,
disable_bf16_reduced_precision_matmul=True,
)
config.kimi_kda_layers = tuple(hf_config.linear_attn_config["kda_layers"])
config.kimi_linear_num_heads = hf_config.linear_attn_config["num_heads"]
config.kimi_linear_head_dim = hf_config.linear_attn_config["head_dim"]
config.kimi_linear_conv_kernel_size = hf_config.linear_attn_config["short_conv_kernel_size"]
config.kimi_kda_gate_lower_bound = hf_config.linear_attn_config["gate_lower_bound"]
config.kimi_attn_res_block_size = hf_config.attn_res_block_size
return config
def _get_gptmodel_args(self) -> dict:
return {
"vocab_size": self.text_config.vocab_size,
"max_sequence_length": self.text_config.max_position_embeddings,
"position_embedding_type": "none",
}
def _get_transformer_layer_spec(self, vp_stage=None):
self.has_vp_stage = True
return build_kimi_k3_spec(self.config, vp_stage=vp_stage)
def _get_safetensor_io(self, weights_path: str):
safetensor_io = SafeTensorIO(self._get_actual_hf_path(weights_path))
assert not any(name.endswith(".weight_packed") for name in safetensor_io.index), (
"Kimi K3 Megatron loading requires the experts to be converted from MXFP4 "
"to standard BF16 safetensors first"
)
return safetensor_io
def _weight_name_mapping_mcore_to_hf(self, mcore_weights_name: str) -> list[str]:
assert "_extra_state" not in mcore_weights_name
if mcore_weights_name in self._DIRECT_MAPPING:
return [self._DIRECT_MAPPING[mcore_weights_name]]
if "self_attention" in mcore_weights_name or "input_layernorm" in mcore_weights_name:
return self._weight_name_mapping_attention(mcore_weights_name)
if "mlp" in mcore_weights_name:
return self._weight_name_mapping_mlp(mcore_weights_name)
return self._weight_name_mapping_other(mcore_weights_name)
def _weight_to_mcore_format(
self,
mcore_weights_name: str,
hf_weights: list[torch.Tensor],
) -> torch.Tensor:
if mcore_weights_name.endswith(
(
"self_attention.q_conv1d.weight",
"self_attention.k_conv1d.weight",
"self_attention.v_conv1d.weight",
)
):
assert len(hf_weights) == 1
return hf_weights[0].float().contiguous()
if mcore_weights_name.endswith("self_attention.A_log"):
assert len(hf_weights) == 1
return hf_weights[0][: self.config.kimi_linear_num_heads].float().contiguous()
if mcore_weights_name.endswith("self_attention.dt_bias"):
assert len(hf_weights) == 1
return hf_weights[0].float().contiguous()
return super()._weight_to_mcore_format(mcore_weights_name, hf_weights)
+7
View File
@@ -0,0 +1,7 @@
def get_kimi_k3_spec(*args, **kwargs):
from .model import get_kimi_k3_spec as get_spec
return get_spec(*args, **kwargs)
__all__ = ["get_kimi_k3_spec"]
+572
View File
@@ -0,0 +1,572 @@
import copy
import torch
import torch.nn as nn
from einops import rearrange
from fla.modules import FusedRMSNormGated, ShortConvolution
from megatron.core.dist_checkpointing.mapping import ShardedStateDict
from megatron.core.extensions.transformer_engine import (
TEColumnParallelLinear,
TEDotProductAttention,
TELinear,
TERowParallelLinear,
)
from megatron.core.inference.contexts import BaseInferenceContext
from megatron.core.packed_seq_params import PackedSeqParams
from megatron.core.process_groups_config import ProcessGroupCollection
from megatron.core.tensor_parallel.layers import set_tensor_model_parallel_attributes
from megatron.core.tensor_parallel.mappings import (
copy_to_tensor_model_parallel_region,
gather_from_sequence_parallel_region,
scatter_to_sequence_parallel_region,
)
from megatron.core.transformer.enums import AttnMaskType
from megatron.core.transformer.module import MegatronModule
from megatron.core.transformer.transformer_layer import TransformerLayer, get_transformer_layer_offset
from megatron.core.transformer.utils import ensure_metadata_has_dp_cp_group, make_sharded_tensors_for_checkpoint
from miles.backends.megatron_utils.fp32_param_utils import mark_param_dtype
from miles_plugins.models.cp_utils import build_gdn_cp_context, packed_shard_to_zigzag, zigzag_to_packed_shard
from miles_plugins.models.kimi_k3.ops import KimiRMSNorm, attn_res_aggregate, kda
from miles_plugins.models.kimi_k3.pipeline import bank_num_rows, pack_stage_boundary, unpack_stage_boundary
def _mark_tp_replicated(module: nn.Module) -> None:
for parameter in module.parameters():
parameter.sum_gradients_across_tp_domain = True
def _linear(module: nn.Module, inputs: torch.Tensor) -> torch.Tensor:
output, bias = module(inputs)
assert bias is None
return output
class KimiK3ShortConvolution(ShortConvolution):
def __init__(self, *args, tp_group, **kwargs) -> None:
super().__init__(*args, **kwargs)
self.tp_group = tp_group
mark_param_dtype(self.weight, torch.float32)
set_tensor_model_parallel_attributes(self.weight, True, 0, 1)
def sharded_state_dict(
self,
prefix: str = "",
sharded_offsets: tuple = (),
metadata: dict | None = None,
) -> ShardedStateDict:
metadata = ensure_metadata_has_dp_cp_group(metadata)
return make_sharded_tensors_for_checkpoint(
self.state_dict(prefix="", keep_vars=True),
prefix,
{"weight": 0},
sharded_offsets,
tp_group=self.tp_group,
dp_cp_group=metadata["dp_cp_group"],
)
class KimiK3Attention(MegatronModule):
def __init__(
self,
config,
layer_number: int,
cp_comm_type: str | None = None,
pg_collection=None,
name: str | None = None,
) -> None:
super().__init__(config=config)
del name # build_module forwards the module path; K3 constructs its submodules directly
self.cp_comm_type = cp_comm_type
if pg_collection is None:
pg_collection = ProcessGroupCollection.use_mpu_process_groups(required_pgs=["tp", "cp"])
else:
assert hasattr(pg_collection, "tp") and hasattr(pg_collection, "cp")
self.pg_collection = pg_collection
self.tp_group = pg_collection.tp
self.tp_size = self.tp_group.size()
self.cp_group = pg_collection.cp
self.cp_size = self.cp_group.size()
self.sequence_parallel = config.sequence_parallel
self.linear_config = copy.copy(config)
self.linear_config.sequence_parallel = False
self.layer_idx = layer_number - 1
self.is_kda = layer_number in config.kimi_kda_layers
if self.is_kda:
self._init_kda(config)
else:
self._init_mla(config)
def _duplicated_linear(self, input_size: int, output_size: int) -> TELinear:
return TELinear(
input_size,
output_size,
config=self.linear_config,
init_method=self.config.init_method,
bias=False,
skip_bias_add=False,
skip_weight_param_allocation=False,
parallel_mode="duplicated",
)
def _column_linear(self, input_size: int, output_size: int) -> TEColumnParallelLinear:
return TEColumnParallelLinear(
input_size,
output_size,
config=self.linear_config,
init_method=self.config.init_method,
bias=False,
gather_output=False,
skip_bias_add=False,
is_expert=False,
tp_group=self.tp_group,
)
def _row_linear(self, input_size: int, output_size: int) -> TERowParallelLinear:
return TERowParallelLinear(
input_size,
output_size,
config=self.linear_config,
init_method=self.config.init_method,
bias=False,
input_is_parallel=True,
skip_bias_add=False,
is_expert=False,
tp_group=self.tp_group,
)
def _init_kda(self, config) -> None:
hidden_size = config.hidden_size
device = torch.cuda.current_device()
dtype = config.params_dtype
self.num_heads = config.kimi_linear_num_heads
assert self.num_heads % self.tp_size == 0
self.local_num_heads = self.num_heads // self.tp_size
self.head_dim = config.kimi_linear_head_dim
# build_gdn_cp_context reads this off the module.
self.conv_kernel_size = config.kimi_linear_conv_kernel_size
self.projection_size = self.num_heads * self.head_dim
self.local_projection_size = self.local_num_heads * self.head_dim
self.q_proj = self._column_linear(hidden_size, self.projection_size)
self.k_proj = self._column_linear(hidden_size, self.projection_size)
self.v_proj = self._column_linear(hidden_size, self.projection_size)
self.q_conv1d = KimiK3ShortConvolution(
hidden_size=self.local_projection_size,
kernel_size=config.kimi_linear_conv_kernel_size,
activation="silu",
device=device,
dtype=dtype,
tp_group=self.tp_group,
)
self.k_conv1d = KimiK3ShortConvolution(
hidden_size=self.local_projection_size,
kernel_size=config.kimi_linear_conv_kernel_size,
activation="silu",
device=device,
dtype=dtype,
tp_group=self.tp_group,
)
self.v_conv1d = KimiK3ShortConvolution(
hidden_size=self.local_projection_size,
kernel_size=config.kimi_linear_conv_kernel_size,
activation="silu",
device=device,
dtype=dtype,
tp_group=self.tp_group,
)
self.f_a_proj = self._duplicated_linear(hidden_size, self.head_dim)
self.f_b_proj = self._column_linear(self.head_dim, self.projection_size)
self.b_proj = self._column_linear(hidden_size, self.num_heads)
self.g_proj = self._column_linear(hidden_size, self.projection_size)
self.A_log = nn.Parameter(torch.empty(self.local_num_heads, dtype=torch.float32, device=device))
self.dt_bias = nn.Parameter(torch.empty(self.local_projection_size, dtype=torch.float32, device=device))
mark_param_dtype(self.A_log, torch.float32)
mark_param_dtype(self.dt_bias, torch.float32)
set_tensor_model_parallel_attributes(self.A_log, True, 0, 1)
set_tensor_model_parallel_attributes(self.dt_bias, True, 0, 1)
self.o_norm = FusedRMSNormGated(
self.head_dim,
eps=config.layernorm_epsilon,
activation="sigmoid",
device=device,
dtype=dtype,
)
_mark_tp_replicated(self.o_norm)
self.o_proj = self._row_linear(self.projection_size, hidden_size)
self.gate_lower_bound = config.kimi_kda_gate_lower_bound
def _init_mla(self, config) -> None:
hidden_size = config.hidden_size
device = torch.cuda.current_device()
dtype = config.params_dtype
self.num_heads = config.num_attention_heads
assert self.num_heads % self.tp_size == 0
self.local_num_heads = self.num_heads // self.tp_size
self.q_lora_rank = config.q_lora_rank
self.kv_lora_rank = config.kv_lora_rank
self.qk_nope_head_dim = config.qk_head_dim
self.qk_extra_head_dim = config.qk_pos_emb_head_dim
self.v_head_dim = config.v_head_dim
self.q_head_dim = self.qk_nope_head_dim + self.qk_extra_head_dim
self.q_a_proj = self._duplicated_linear(hidden_size, self.q_lora_rank)
self.q_a_layernorm = KimiRMSNorm(self.q_lora_rank, config.layernorm_epsilon, device=device, dtype=dtype)
self.q_b_proj = self._column_linear(self.q_lora_rank, self.num_heads * self.q_head_dim)
self.kv_a_proj_with_mqa = self._duplicated_linear(hidden_size, self.kv_lora_rank + self.qk_extra_head_dim)
self.kv_a_layernorm = KimiRMSNorm(self.kv_lora_rank, config.layernorm_epsilon, device=device, dtype=dtype)
self.kv_b_proj = self._column_linear(
self.kv_lora_rank,
self.num_heads * (self.qk_nope_head_dim + self.v_head_dim),
)
self.g_proj = self._column_linear(hidden_size, self.num_heads * self.v_head_dim)
self.o_proj = self._row_linear(self.num_heads * self.v_head_dim, hidden_size)
# Megatron's MLA core carries CP and varlen; only its module assumes the rotated 64-wide field K3 lacks
self.core_attention = TEDotProductAttention(
config=self.config,
layer_number=self.layer_idx + 1,
attn_mask_type=AttnMaskType.causal,
attention_type="self",
softmax_scale=self.q_head_dim**-0.5,
k_channels=self.q_head_dim,
v_channels=self.v_head_dim,
cp_comm_type=self.cp_comm_type,
pg_collection=self.pg_collection,
)
def sharded_state_dict(
self,
prefix: str = "",
sharded_offsets: tuple = (),
metadata: dict | None = None,
) -> ShardedStateDict:
sharded_state_dict = super().sharded_state_dict(prefix, sharded_offsets, metadata)
if not self.is_kda:
return sharded_state_dict
metadata = ensure_metadata_has_dp_cp_group(metadata)
sharded_state_dict.update(
make_sharded_tensors_for_checkpoint(
{"A_log": self.A_log, "dt_bias": self.dt_bias},
prefix,
{"A_log": 0, "dt_bias": 0},
sharded_offsets,
tp_group=self.tp_group,
dp_cp_group=metadata["dp_cp_group"],
)
)
return sharded_state_dict
def _cp_global_cu_seqlens(
self,
hidden_states: torch.Tensor,
packed_seq_params: PackedSeqParams | None,
) -> torch.Tensor:
"""Global packed-sequence boundaries; under BSHD the whole sequence is one segment."""
if packed_seq_params is not None and packed_seq_params.cu_seqlens_q is not None:
return packed_seq_params.cu_seqlens_q
total = hidden_states.shape[0] * self.cp_size
return torch.tensor([0, total], dtype=torch.int32, device=hidden_states.device)
def _forward_kda(
self,
hidden_states: torch.Tensor,
packed_seq_params: PackedSeqParams | None,
) -> torch.Tensor:
core = self._kda_core(hidden_states, packed_seq_params)
return _linear(self.o_proj, core.to(hidden_states.dtype)).transpose(0, 1)
def _kda_core(
self,
hidden_states: torch.Tensor,
packed_seq_params: PackedSeqParams | None,
) -> torch.Tensor:
"""KDA delta-rule core: projections + conv + recurrence + gated norm,
returning the flattened pre-o_proj output."""
cp_context = (
build_gdn_cp_context(
self, self._cp_global_cu_seqlens(hidden_states, packed_seq_params), hidden_states.device
)
if self.cp_size > 1
else None
)
x = hidden_states.transpose(0, 1)
cu_seqlens = packed_seq_params.cu_seqlens_q if packed_seq_params is not None else None
if cp_context is not None:
# the context carries the rank-local boundaries; the global ones only located this shard
cu_seqlens = cp_context.cu_seqlens
# packed input with neither channel would silently leak recurrent state across samples
assert (
packed_seq_params is None or cu_seqlens is not None or cp_context is not None
), "packed (THD) input reached the KDA core without sequence boundaries"
conv_kwargs = {"output_final_state": False, "cu_seqlens": cu_seqlens, "cp_context": cp_context}
q, _ = self.q_conv1d(x=_linear(self.q_proj, x), **conv_kwargs)
k, _ = self.k_conv1d(x=_linear(self.k_proj, x), **conv_kwargs)
v, _ = self.v_conv1d(x=_linear(self.v_proj, x), **conv_kwargs)
q = rearrange(q, "b s (h d) -> b s h d", h=self.local_num_heads)
k = rearrange(k, "b s (h d) -> b s h d", h=self.local_num_heads)
v = rearrange(v, "b s (h d) -> b s h d", h=self.local_num_heads)
forget_gate = rearrange(
_linear(self.f_b_proj, _linear(self.f_a_proj, x)), "b s (h d) -> b s h d", h=self.local_num_heads
)
beta = _linear(self.b_proj, x).float().sigmoid()
output = kda(
q,
k,
v,
forget_gate,
beta,
self.A_log,
self.dt_bias,
self.gate_lower_bound,
cu_seqlens=cu_seqlens,
cp_context=cp_context,
)
gate = rearrange(_linear(self.g_proj, x), "b s (h d) -> b s h d", h=self.local_num_heads)
output = self.o_norm(output.reshape(-1, self.head_dim), gate.reshape(-1, self.head_dim))
return output.view(*gate.shape).flatten(-2)
def _forward_mla(
self,
hidden_states: torch.Tensor,
packed_seq_params: PackedSeqParams | None,
) -> torch.Tensor:
# stay in TE's sbhd [s, b, ...] layout so q/k/v need no transpose; the KDA kernels want [b, s, ...]
x = hidden_states
query = _linear(
self.q_b_proj,
self.q_a_layernorm(_linear(self.q_a_proj, x)),
)
query = query.view(*query.shape[:-1], self.local_num_heads, self.q_head_dim)
compressed_kv = _linear(self.kv_a_proj_with_mqa, x)
kv_latent, key_extra = torch.split(
compressed_kv,
[self.kv_lora_rank, self.qk_extra_head_dim],
dim=-1,
)
key_value = _linear(self.kv_b_proj, self.kv_a_layernorm(kv_latent))
key_value = key_value.view(
*key_value.shape[:-1],
self.local_num_heads,
self.qk_nope_head_dim + self.v_head_dim,
)
key_nope, value = torch.split(
key_value,
[self.qk_nope_head_dim, self.v_head_dim],
dim=-1,
)
# TE classifies the qkv layout from strides, so the sliced value needs canonical ones too
value = value.contiguous()
key_extra = copy_to_tensor_model_parallel_region(key_extra, group=self.tp_group)
key_extra = key_extra.unsqueeze(-2).expand(*key_nope.shape[:-1], -1)
key = torch.cat((key_nope, key_extra), dim=-1)
# thd packing wants 3D [t, h, d]; mirror Megatron's Attention.forward squeeze/reshape
is_thd = packed_seq_params is not None and packed_seq_params.qkv_format == "thd"
if is_thd:
query = query.squeeze(1)
key = key.squeeze(1)
value = value.squeeze(1)
# TE returns the head dims already fused, so there is no flatten here.
output = self.core_attention(
query,
key,
value,
None,
packed_seq_params=packed_seq_params,
attn_mask_type=AttnMaskType.causal,
)
if is_thd:
output = output.reshape(output.size(0), 1, -1)
output = output * torch.sigmoid(_linear(self.g_proj, x))
return _linear(self.o_proj, output)
def forward(
self,
hidden_states: torch.Tensor,
attention_mask: torch.Tensor | None = None,
key_value_states: torch.Tensor | None = None,
inference_context: BaseInferenceContext | None = None,
rotary_pos_emb: torch.Tensor | None = None,
rotary_pos_cos: torch.Tensor | None = None,
rotary_pos_sin: torch.Tensor | None = None,
rotary_pos_cos_sin: torch.Tensor | None = None,
attention_bias: torch.Tensor | None = None,
packed_seq_params: PackedSeqParams | None = None,
sequence_len_offset: int | None = None,
**kwargs,
) -> tuple[torch.Tensor, None]:
del (
attention_mask,
key_value_states,
inference_context,
rotary_pos_emb,
rotary_pos_cos,
rotary_pos_sin,
rotary_pos_cos_sin,
attention_bias,
sequence_len_offset,
kwargs,
)
if self.sequence_parallel:
hidden_states = gather_from_sequence_parallel_region(
hidden_states,
tensor_parallel_output_grad=False,
group=self.tp_group,
)
# CP tokens sit in ring attention's zigzag order; fla's CP kernels want a contiguous rank-local chunk
relayout = self.is_kda and self.cp_size > 1
if relayout:
cp_cu_seqlens = self._cp_global_cu_seqlens(hidden_states, packed_seq_params)
hidden_states = zigzag_to_packed_shard(
hidden_states, cp_cu_seqlens, self.cp_group, self.cp_group.rank(), self.cp_size
)
output = (
self._forward_kda(hidden_states, packed_seq_params)
if self.is_kda
else self._forward_mla(hidden_states, packed_seq_params)
)
if relayout:
output = packed_shard_to_zigzag(output, cp_cu_seqlens, self.cp_group, self.cp_group.rank(), self.cp_size)
if self.sequence_parallel:
output = scatter_to_sequence_parallel_region(output, group=self.tp_group)
return output, None
class KimiK3TransformerLayer(TransformerLayer):
def __init__(self, *args, **kwargs) -> None:
super().__init__(*args, **kwargs)
assert self.config.hidden_dropout == 0.0, "Kimi K3 requires hidden dropout 0"
hidden_size = self.config.hidden_size
eps = self.config.layernorm_epsilon
device = torch.cuda.current_device()
dtype = self.config.params_dtype
self.attn_res_block_size = self.config.kimi_attn_res_block_size
self.self_attention_res_norm = KimiRMSNorm(hidden_size, eps, device=device, dtype=dtype)
self.self_attention_res_proj = nn.Linear(hidden_size, 1, bias=False, device=device, dtype=dtype)
self.mlp_res_norm = KimiRMSNorm(hidden_size, eps, device=device, dtype=dtype)
self.mlp_res_proj = nn.Linear(hidden_size, 1, bias=False, device=device, dtype=dtype)
_mark_tp_replicated(self.self_attention_res_norm)
_mark_tp_replicated(self.self_attention_res_proj)
_mark_tp_replicated(self.mlp_res_norm)
_mark_tp_replicated(self.mlp_res_proj)
if self.layer_number == self.config.num_layers:
self.output_attn_res_norm = KimiRMSNorm(hidden_size, eps, device=device, dtype=dtype)
self.output_attn_res_proj = nn.Linear(hidden_size, 1, bias=False, device=device, dtype=dtype)
_mark_tp_replicated(self.output_attn_res_norm)
_mark_tp_replicated(self.output_attn_res_proj)
# stage entry/exit layers from the per-rank offsets; VPP is rejected, so vp_stage is None
pp_size = self.config.pipeline_model_parallel_size
stage_starts = {get_transformer_layer_offset(self.config, None, r) for r in range(pp_size)}
layer_idx = self.layer_number - 1
self.is_stage_entry = layer_idx in stage_starts and layer_idx > 0
self.is_stage_exit = (layer_idx + 1) in stage_starts and layer_idx + 1 < self.config.num_layers
@staticmethod
def _add_bias(output_with_bias: tuple[torch.Tensor, torch.Tensor | None]) -> torch.Tensor:
output, bias = output_with_bias
return output if bias is None else output + bias
def forward(
self,
hidden_states: torch.Tensor,
attention_mask: torch.Tensor | None = None,
context: torch.Tensor | None = None,
context_mask: torch.Tensor | None = None,
rotary_pos_emb: torch.Tensor | None = None,
rotary_pos_cos: torch.Tensor | None = None,
rotary_pos_sin: torch.Tensor | None = None,
rotary_pos_cos_sin: torch.Tensor | None = None,
attention_bias: torch.Tensor | None = None,
inference_context: BaseInferenceContext | None = None,
packed_seq_params: PackedSeqParams | None = None,
sequence_len_offset: torch.Tensor | None = None,
padding_mask: torch.Tensor | None = None,
input_ids: torch.Tensor | None = None,
**kwargs,
) -> tuple[torch.Tensor, torch.Tensor]:
del context_mask, kwargs
layer_idx = self.layer_number - 1
if context is not None:
prefix_sum = hidden_states
block_residual = context
elif layer_idx == 0:
prefix_sum = hidden_states
block_residual = hidden_states.new_empty(*hidden_states.shape[:-1], 0, hidden_states.shape[-1])
else:
assert self.is_stage_entry, "Attention-residual snapshot bank is missing"
prefix_sum, block_residual = unpack_stage_boundary(
hidden_states,
self.config.hidden_size,
bank_num_rows(layer_idx, self.attn_res_block_size),
)
if block_residual.shape[-2] > 0:
attention_input = attn_res_aggregate(
prefix_sum,
block_residual,
self.self_attention_res_proj,
self.self_attention_res_norm,
self.input_layernorm,
)
else:
attention_input = self.input_layernorm(prefix_sum)
is_block_write_layer = layer_idx % self.attn_res_block_size == 0
if is_block_write_layer:
block_residual = torch.cat((block_residual, prefix_sum.unsqueeze(-2)), dim=-2)
attention_output = self._add_bias(
self.self_attention(
attention_input,
attention_mask=attention_mask,
inference_context=inference_context,
rotary_pos_emb=rotary_pos_emb,
rotary_pos_cos=rotary_pos_cos,
rotary_pos_sin=rotary_pos_sin,
rotary_pos_cos_sin=rotary_pos_cos_sin,
attention_bias=attention_bias,
packed_seq_params=packed_seq_params,
sequence_len_offset=sequence_len_offset,
)
)
prefix_sum = attention_output if is_block_write_layer else prefix_sum + attention_output
mlp_input = attn_res_aggregate(
prefix_sum,
block_residual,
self.mlp_res_proj,
self.mlp_res_norm,
self.pre_mlp_layernorm,
)
mlp_kwargs = {"padding_mask": padding_mask}
if self.is_moe_layer:
mlp_kwargs["input_ids"] = input_ids
mlp_output = self._add_bias(self.mlp(mlp_input, **mlp_kwargs))
prefix_sum = prefix_sum + mlp_output
if self.layer_number == self.config.num_layers:
prefix_sum = attn_res_aggregate(
prefix_sum,
block_residual,
self.output_attn_res_proj,
self.output_attn_res_norm,
nn.Identity(),
)
if self.is_stage_exit:
return pack_stage_boundary(prefix_sum, block_residual), block_residual
return prefix_sum, block_residual
+684
View File
@@ -0,0 +1,684 @@
from __future__ import annotations
import logging
from collections import Counter
from collections.abc import Callable
import torch
import torch.nn as nn
import torch.nn.functional as F
logger = logging.getLogger(__name__)
# TODO: will rewrite in native lora refactor: the adapter class, parameter registration, gather batch and
# export collation below duplicate miles_plugins/models/inkling/lora.py
_SUPPORTED_TARGET_SUFFIXES = {
"self_attention.o_proj",
"self_attention.q_a_proj",
"self_attention.kv_a_proj_with_mqa",
"mlp.linear_fc1",
"mlp.linear_fc2",
"mlp.experts.linear_fc1",
"mlp.experts.linear_fc2",
}
# the routed-expert down-proj may be omitted: its EP-shared w2_lora_B dominates adapter growth (#1559)
_OPTIONAL_TARGET_SUFFIXES = {"mlp.experts.linear_fc2"}
class KimiK3LoRAAdapter(nn.Module):
def __init__(self, kind: str, hf_prefix: str) -> None:
super().__init__()
self.kind = kind
self.hf_prefix = hf_prefix
self.include_fc2 = True # experts only: _apply_expert_lora overrides per target set
def sharded_state_dict(self, prefix="", sharded_offsets=(), metadata=None):
del prefix, sharded_offsets, metadata
return {}
def _new_param(
ref_weight: torch.Tensor,
shape: tuple[int, ...],
*,
init: str,
grad_sum_group: str | None = None,
expert: bool = False,
) -> nn.Parameter:
tensor = torch.empty(*shape, dtype=ref_weight.dtype, device=ref_weight.device)
if init == "zero":
tensor.zero_()
elif init == "xavier":
if tensor.ndim == 2:
nn.init.xavier_uniform_(tensor)
else:
for expert_tensor in tensor:
nn.init.xavier_uniform_(expert_tensor)
else:
raise ValueError(f"Unsupported Kimi K3 LoRA init method: {init}")
param = nn.Parameter(tensor)
param.tensor_model_parallel = False
param.partition_dim = -1
param.partition_stride = 1
if expert:
param.allreduce = False
if grad_sum_group == "tp":
# Megatron sums TP-replicated partial grads itself in finalize_model_grads
param.sum_gradients_across_tp_domain = True
elif grad_sum_group is not None:
# Megatron never reduces over EP (experts are assumed partitioned); reduce_marked_lora_grads does
assert grad_sum_group == "ep", f"unsupported LoRA gradient sum group {grad_sum_group!r}"
param._lora_grad_sum_group = grad_sum_group
return param
def _register_param(
adapter: KimiK3LoRAAdapter,
name: str,
ref_weight: torch.Tensor,
shape: tuple[int, ...],
*,
init: str,
grad_sum_group: str | None = None,
expert: bool = False,
) -> None:
adapter.register_parameter(
name,
_new_param(
ref_weight,
shape,
init=init,
grad_sum_group=grad_sum_group,
expert=expert,
),
)
def _dropout(inputs: torch.Tensor, probability: float, training: bool) -> torch.Tensor:
if probability and training:
return F.dropout(inputs, p=probability, training=True)
return inputs
def _grouped_linear(inputs: torch.Tensor, weights: torch.Tensor, tokens_per_expert: list[int]) -> torch.Tensor:
if inputs.is_cuda:
offsets = torch.as_tensor(tokens_per_expert, device=inputs.device, dtype=torch.int32).cumsum(
0, dtype=torch.int32
)
return F.grouped_mm(inputs, weights.transpose(1, 2), offs=offsets)
segments = torch.split(inputs, tokens_per_expert, dim=0)
return torch.cat([F.linear(segment, weights[idx]) for idx, segment in enumerate(segments)], dim=0)
def _validate_targets(args) -> None:
targets = set(args.target_modules)
suffixes = {target.split("decoder.layers.*.", 1)[-1] for target in targets}
unsupported = suffixes - _SUPPORTED_TARGET_SUFFIXES
missing = _SUPPORTED_TARGET_SUFFIXES - suffixes - _OPTIONAL_TARGET_SUFFIXES
if unsupported or missing:
raise NotImplementedError(
"Kimi K3 native LoRA currently requires the verified target set; "
f"unsupported={sorted(unsupported)}, missing={sorted(missing)}"
)
if not args.experts_shared_outer_loras:
raise NotImplementedError("Kimi K3 native LoRA currently requires --experts-shared-outer-loras")
def _enable_full_recompute_input_grads(model) -> None:
if model.config.recompute_granularity != "full" or not model.pre_process:
return
def enable_grad(module, _inputs, output):
if module.training:
if not isinstance(output, torch.Tensor):
raise TypeError(f"Kimi K3 embedding returned {type(output)}, expected torch.Tensor")
output.requires_grad_(True)
return output
model.embedding.register_forward_hook(enable_grad)
def _apply_attention_lora(attention, args, layer_idx: int, scale: float, dropout: float) -> None:
from megatron.core.tensor_parallel.mappings import reduce_from_tensor_model_parallel_region
rank = int(args.lora_rank)
hidden_size = attention.config.hidden_size
adapter = KimiK3LoRAAdapter(
"kda_attention" if attention.is_kda else "mla_attention",
f"language_model.model.layers.{layer_idx}.self_attn.",
)
_register_param(
adapter,
"o_lora_A",
attention.o_proj.weight,
(rank, attention.o_proj.weight.shape[1]),
init="xavier",
)
_register_param(
adapter,
"o_lora_B",
attention.o_proj.weight,
(hidden_size, rank),
init="zero",
)
o_proj = attention.o_proj
original_o_proj = o_proj.forward
def o_proj_forward(inputs, *forward_args, **forward_kwargs):
output, bias = original_o_proj(inputs, *forward_args, **forward_kwargs)
local = F.linear(_dropout(inputs, dropout, o_proj.training), adapter.o_lora_A)
reduced = reduce_from_tensor_model_parallel_region(local, group=attention.tp_group)
delta = F.linear(reduced, adapter.o_lora_B)
return torch.add(output, delta, alpha=scale), bias
o_proj.forward = o_proj_forward
if not attention.is_kda:
for module_name, output_size in (
("q_a_proj", attention.q_lora_rank),
("kv_a_proj_with_mqa", attention.kv_lora_rank + attention.qk_extra_head_dim),
):
module = getattr(attention, module_name)
prefix = "q_a" if module_name == "q_a_proj" else "kv_a"
_register_param(
adapter,
f"{prefix}_lora_A",
module.weight,
(rank, hidden_size),
init="xavier",
)
_register_param(
adapter,
f"{prefix}_lora_B",
module.weight,
(output_size, rank),
init="zero",
)
original_forward = module.forward
lora_a = getattr(adapter, f"{prefix}_lora_A")
lora_b = getattr(adapter, f"{prefix}_lora_B")
def duplicated_forward(
inputs,
*forward_args,
_module=module,
_original=original_forward,
_a=lora_a,
_b=lora_b,
**forward_kwargs,
):
output, bias = _original(inputs, *forward_args, **forward_kwargs)
delta = F.linear(F.linear(_dropout(inputs, dropout, _module.training), _a), _b)
# the TE column backward reduces the latent dgrad over TP; KimiK3Attention reduces the key-extra slice
return torch.add(output, delta, alpha=scale), bias
module.forward = duplicated_forward
attention.lora_adapter = adapter
def _apply_dense_mlp_lora(
mlp,
args,
layer_idx: int,
scale: float,
dropout: float,
*,
adapter_kind: str = "dense_mlp",
hf_prefix: str | None = None,
) -> None:
from megatron.core.tensor_parallel.mappings import (
gather_from_sequence_parallel_region,
reduce_from_tensor_model_parallel_region,
reduce_scatter_to_sequence_parallel_region,
)
rank = int(args.lora_rank)
sequence_parallel = bool(mlp.config.sequence_parallel)
tp_group = mlp.tp_group
fc1 = mlp.linear_fc1
fc2 = mlp.linear_fc2
adapter = KimiK3LoRAAdapter(
adapter_kind,
hf_prefix or f"language_model.model.layers.{layer_idx}.mlp.",
)
_register_param(
adapter,
"fc1_lora_A",
fc1.weight,
(rank, mlp.config.hidden_size),
init="xavier",
grad_sum_group="tp",
)
_register_param(
adapter,
"fc1_lora_B",
fc1.weight,
(fc1.weight.shape[0], rank),
init="zero",
)
_register_param(
adapter,
"fc2_lora_A",
fc2.weight,
(rank, fc2.weight.shape[1]),
init="xavier",
)
_register_param(
adapter,
"fc2_lora_B",
fc2.weight,
(mlp.config.hidden_size, rank),
init="zero",
grad_sum_group="tp" if sequence_parallel else None,
)
original_fc1 = fc1.forward
def fc1_forward(inputs, *forward_args, **forward_kwargs):
output, bias = original_fc1(inputs, *forward_args, **forward_kwargs)
adapter_inputs = gather_from_sequence_parallel_region(inputs, group=tp_group) if sequence_parallel else inputs
delta = F.linear(
F.linear(_dropout(adapter_inputs, dropout, fc1.training), adapter.fc1_lora_A),
adapter.fc1_lora_B,
)
return torch.add(output, delta, alpha=scale), bias
fc1.forward = fc1_forward
original_fc2 = fc2.forward
def fc2_forward(inputs, *forward_args, **forward_kwargs):
output, bias = original_fc2(inputs, *forward_args, **forward_kwargs)
local = F.linear(_dropout(inputs, dropout, fc2.training), adapter.fc2_lora_A)
reduced = (
reduce_scatter_to_sequence_parallel_region(local, group=tp_group)
if sequence_parallel
else reduce_from_tensor_model_parallel_region(local, group=tp_group)
)
delta = F.linear(reduced, adapter.fc2_lora_B)
return torch.add(output, delta, alpha=scale), bias
fc2.forward = fc2_forward
mlp.lora_adapter = adapter
def _apply_expert_lora(
moe,
args,
layer_idx: int,
scale: float,
dropout: float,
*,
include_fc2: bool,
) -> None:
experts = moe.experts
if (moe.config.expert_tensor_parallel_size or 1) != 1:
raise NotImplementedError("Kimi K3 native expert LoRA currently requires ETP=1")
rank = int(args.lora_rank)
num_local_experts = experts.num_local_experts
latent_size = moe.config.moe_latent_size
intermediate_size = moe.config.moe_ffn_hidden_size
ref_fc1 = experts.linear_fc1.weight0
ref_fc2 = experts.linear_fc2.weight0
adapter = KimiK3LoRAAdapter(
"experts",
f"language_model.model.layers.{layer_idx}.block_sparse_moe.experts.",
)
adapter.include_fc2 = include_fc2
_register_param(
adapter,
"w1_lora_A",
ref_fc1,
(rank, latent_size),
init="xavier",
expert=True,
grad_sum_group="ep",
)
_register_param(
adapter,
"w3_lora_A",
ref_fc1,
(rank, latent_size),
init="xavier",
expert=True,
grad_sum_group="ep",
)
_register_param(
adapter,
"w1_lora_B",
ref_fc1,
(num_local_experts, intermediate_size, rank),
init="zero",
expert=True,
)
_register_param(
adapter,
"w3_lora_B",
ref_fc1,
(num_local_experts, intermediate_size, rank),
init="zero",
expert=True,
)
if include_fc2:
_register_param(
adapter,
"w2_lora_A",
ref_fc2,
(num_local_experts, rank, intermediate_size),
init="xavier",
expert=True,
)
_register_param(
adapter,
"w2_lora_B",
ref_fc2,
(latent_size, rank),
init="zero",
grad_sum_group="ep",
expert=True,
)
fc1 = experts.linear_fc1
original_fc1 = fc1.forward
def expert_fc1_forward(inputs, tokens_per_expert, *forward_args, **forward_kwargs):
output, bias = original_fc1(inputs, tokens_per_expert, *forward_args, **forward_kwargs)
adapter_inputs = _dropout(inputs, dropout, fc1.training)
shared = F.linear(
adapter_inputs,
torch.cat((adapter.w1_lora_A, adapter.w3_lora_A), dim=0),
)
w1_shared, w3_shared = shared.chunk(2, dim=-1)
w1_delta = _grouped_linear(w1_shared.contiguous(), adapter.w1_lora_B, tokens_per_expert)
w3_delta = _grouped_linear(w3_shared.contiguous(), adapter.w3_lora_B, tokens_per_expert)
delta = torch.cat((w1_delta, w3_delta), dim=-1)
return torch.add(output, delta, alpha=scale), bias
fc1.forward = expert_fc1_forward
if include_fc2:
fc2 = experts.linear_fc2
original_fc2 = fc2.forward
def expert_fc2_forward(inputs, tokens_per_expert, *forward_args, **forward_kwargs):
output, bias = original_fc2(inputs, tokens_per_expert, *forward_args, **forward_kwargs)
inner = _grouped_linear(
_dropout(inputs, dropout, fc2.training),
adapter.w2_lora_A,
tokens_per_expert,
)
delta = F.linear(inner, adapter.w2_lora_B)
return torch.add(output, delta, alpha=scale), bias
fc2.forward = expert_fc2_forward
experts.lora_adapter = adapter
def apply_kimi_k3_lora(model, args):
from megatron.core.transformer.mlp import MLP
from megatron.core.transformer.moe.moe_layer import MoELayer
from .layers import KimiK3Attention
_validate_targets(args)
rank = int(args.lora_rank)
if rank <= 0:
raise ValueError("apply_kimi_k3_lora requires --lora-rank > 0")
scale = float(args.lora_alpha) / rank
dropout = float(args.lora_dropout or 0.0)
for parameter in model.parameters():
parameter.requires_grad = False
_enable_full_recompute_input_grads(model)
for layer in model.decoder.layers:
layer_idx = layer.layer_number - 1
if not isinstance(layer.self_attention, KimiK3Attention):
raise TypeError(f"Kimi K3 layer {layer_idx} has unexpected attention type {type(layer.self_attention)}")
_apply_attention_lora(layer.self_attention, args, layer_idx, scale, dropout)
if isinstance(layer.mlp, MLP):
_apply_dense_mlp_lora(layer.mlp, args, layer_idx, scale, dropout)
elif isinstance(layer.mlp, MoELayer):
_apply_expert_lora(
layer.mlp,
args,
layer_idx,
scale,
dropout,
include_fc2=any(target.endswith("mlp.experts.linear_fc2") for target in args.target_modules),
)
if layer.mlp.shared_experts is None:
raise RuntimeError(f"Kimi K3 MoE layer {layer_idx} is missing shared experts")
_apply_dense_mlp_lora(
layer.mlp.shared_experts,
args,
layer_idx,
scale,
dropout,
adapter_kind="shared_experts",
hf_prefix=(f"language_model.model.layers.{layer_idx}.block_sparse_moe.shared_experts."),
)
else:
raise TypeError(f"Kimi K3 layer {layer_idx} has unexpected MLP type {type(layer.mlp)}")
trainable = sum(parameter.numel() for parameter in model.parameters() if parameter.requires_grad)
total = sum(parameter.numel() for parameter in model.parameters())
logger.info(
"Kimi K3 native LoRA applied: rank=%d alpha=%s trainable=%d total=%d ratio=%.6f%%",
rank,
args.lora_alpha,
trainable,
total,
100.0 * trainable / total,
)
return model
def wrap_model_provider_with_kimi_k3_lora(provider_func, args):
def wrapped(*provider_args, **provider_kwargs):
return apply_kimi_k3_lora(provider_func(*provider_args, **provider_kwargs), args)
return wrapped
class _GatherBatch:
class _Token:
def __init__(self, batch: _GatherBatch, kind: str, index: int) -> None:
self.batch = batch
self.kind = kind
self.index = index
def get(self) -> torch.Tensor:
return self.batch.resolved[self.kind][self.index]
def __init__(self) -> None:
self.requests: dict[str, list[tuple[torch.Tensor, int]]] = {"tp": [], "ep": []}
self.resolved: dict[str, list[torch.Tensor]] = {"tp": [], "ep": []}
def add(self, kind: str, local: torch.Tensor, dim: int) -> _Token:
self.requests[kind].append((local, dim))
return self._Token(self, kind, len(self.requests[kind]) - 1)
def flush(self) -> int:
from megatron.core import parallel_state
groups = {
"tp": (
parallel_state.get_tensor_model_parallel_group,
parallel_state.get_tensor_model_parallel_world_size(),
),
"ep": (
parallel_state.get_expert_model_parallel_group,
parallel_state.get_expert_model_parallel_world_size(),
),
}
calls = 0
for kind, requests in self.requests.items():
if not requests:
continue
get_group, world_size = groups[kind]
if world_size == 1:
self.resolved[kind] = [local for local, _dim in requests]
continue
group = get_group()
dtypes = {local.dtype for local, _dim in requests}
if len(dtypes) != 1:
raise TypeError(f"Kimi K3 LoRA {kind} gather has mixed dtypes: {dtypes}")
flat_parts = [local.detach().contiguous().view(-1) for local, _dim in requests]
sizes = [part.numel() for part in flat_parts]
flat_local = torch.cat(flat_parts)
gathered = flat_local.new_empty(world_size * flat_local.numel())
torch.distributed.all_gather_into_tensor(gathered, flat_local, group=group)
per_rank = gathered.view(world_size, flat_local.numel())
offset = 0
resolved = []
for (local, dim), size in zip(requests, sizes, strict=True):
partitions = [per_rank[rank, offset : offset + size].view(local.shape) for rank in range(world_size)]
resolved.append(torch.cat(partitions, dim=dim))
offset += size
self.resolved[kind] = resolved
calls += 1
return calls
def _unwrap_model_chunks(model_chunks):
for chunk in model_chunks:
while hasattr(chunk, "module"):
chunk = chunk.module
yield chunk
def _validate_adapter_layout(models, adapters: list[KimiK3LoRAAdapter]) -> None:
expected = []
for model in models:
for layer in model.decoder.layers:
layer_idx = layer.layer_number - 1
attention_kind = "kda_attention" if layer.self_attention.is_kda else "mla_attention"
expected.append(
(
attention_kind,
f"language_model.model.layers.{layer_idx}.self_attn.",
)
)
if hasattr(layer.mlp, "experts"):
expected.extend(
(
(
"experts",
f"language_model.model.layers.{layer_idx}.block_sparse_moe.experts.",
),
(
"shared_experts",
f"language_model.model.layers.{layer_idx}.block_sparse_moe.shared_experts.",
),
)
)
else:
expected.append(
(
"dense_mlp",
f"language_model.model.layers.{layer_idx}.mlp.",
)
)
expected_counts = Counter(expected)
actual_counts = Counter((adapter.kind, adapter.hf_prefix) for adapter in adapters)
if actual_counts != expected_counts:
missing = list((expected_counts - actual_counts).elements())
unexpected = list((actual_counts - expected_counts).elements())
raise RuntimeError(
"Kimi K3 LoRA adapter layout is incomplete: "
f"expected={sum(expected_counts.values())}, actual={sum(actual_counts.values())}, "
f"missing={missing[:5]}, unexpected={unexpected[:5]}"
)
def _export_attention(adapter: KimiK3LoRAAdapter):
batch = _GatherBatch()
plans: list[tuple[str, torch.Tensor | Callable[[], torch.Tensor]]] = []
prefix = adapter.hf_prefix
if adapter.kind == "mla_attention":
for hf_name, parameter_a, parameter_b in (
("q_a_proj", adapter.q_a_lora_A, adapter.q_a_lora_B),
("kv_a_proj_with_mqa", adapter.kv_a_lora_A, adapter.kv_a_lora_B),
):
plans.append((f"{prefix}{hf_name}.lora_A.weight", parameter_a))
plans.append((f"{prefix}{hf_name}.lora_B.weight", parameter_b))
o_a = batch.add("tp", adapter.o_lora_A, 1)
plans.append((f"{prefix}o_proj.lora_A.weight", o_a.get))
plans.append((f"{prefix}o_proj.lora_B.weight", adapter.o_lora_B))
batch.flush()
return plans
def _export_dense_mlp(adapter: KimiK3LoRAAdapter):
batch = _GatherBatch()
prefix = adapter.hf_prefix
fc1_a = adapter.fc1_lora_A
gate_b_local, up_b_local = adapter.fc1_lora_B.chunk(2, dim=0)
gate_b = batch.add("tp", gate_b_local, 0)
up_b = batch.add("tp", up_b_local, 0)
down_a = batch.add("tp", adapter.fc2_lora_A, 1)
fc2_b = adapter.fc2_lora_B
batch.flush()
return [
(f"{prefix}gate_proj.lora_A.weight", fc1_a),
(f"{prefix}gate_proj.lora_B.weight", gate_b.get),
(f"{prefix}up_proj.lora_A.weight", fc1_a),
(f"{prefix}up_proj.lora_B.weight", up_b.get),
(f"{prefix}down_proj.lora_A.weight", down_a.get),
(f"{prefix}down_proj.lora_B.weight", fc2_b),
]
def _export_experts(adapter: KimiK3LoRAAdapter):
batch = _GatherBatch()
prefix = adapter.hf_prefix
w1_a = adapter.w1_lora_A
w3_a = adapter.w3_lora_A
w1_b = batch.add("ep", adapter.w1_lora_B, 0)
w3_b = batch.add("ep", adapter.w3_lora_B, 0)
if adapter.include_fc2:
w2_a = batch.add("ep", adapter.w2_lora_A, 0)
w2_b = adapter.w2_lora_B
batch.flush()
plans = [
(f"{prefix}w1.lora_A.weight", w1_a.unsqueeze(0)),
(f"{prefix}w1.lora_B.weight", w1_b.get),
(f"{prefix}w3.lora_A.weight", w3_a.unsqueeze(0)),
(f"{prefix}w3.lora_B.weight", w3_b.get),
]
if adapter.include_fc2:
plans.append((f"{prefix}w2.lora_A.weight", w2_a.get))
plans.append((f"{prefix}w2.lora_B.weight", w2_b.unsqueeze(0)))
return plans
def export_kimi_k3_lora_hf_chunks(model_chunks):
models = list(_unwrap_model_chunks(model_chunks))
adapters: list[KimiK3LoRAAdapter] = []
for model in models:
adapters.extend(module for module in model.modules() if isinstance(module, KimiK3LoRAAdapter))
if not adapters:
raise RuntimeError("Kimi K3 native LoRA export found no adapters")
_validate_adapter_layout(models, adapters)
for adapter in adapters:
if adapter.kind in ("kda_attention", "mla_attention"):
plans = _export_attention(adapter)
elif adapter.kind in ("dense_mlp", "shared_experts"):
plans = _export_dense_mlp(adapter)
elif adapter.kind == "experts":
plans = _export_experts(adapter)
else:
raise ValueError(f"Unknown Kimi K3 LoRA adapter kind: {adapter.kind}")
yield [
(name, (value() if callable(value) else value).detach().to(torch.bfloat16).contiguous())
for name, value in plans
]
+66
View File
@@ -0,0 +1,66 @@
import copy
from megatron.core.extensions.transformer_engine import TEColumnParallelLinear, TENorm
from megatron.core.models.gpt.gpt_layer_specs import get_gpt_decoder_block_spec
from megatron.core.transformer.spec_utils import ModuleSpec
from megatron.core.transformer.transformer_layer import get_transformer_layer_offset
from .layers import KimiK3Attention, KimiK3TransformerLayer
from .ops import situ_and_mul
KIMI_K3_KDA_LAYERS = tuple(
layer_number for layer_number in range(1, 94) if layer_number % 4 != 0 and layer_number != 93
)
def configure_kimi_k3(config) -> None:
assert config.num_layers <= 93
assert config.num_attention_heads == 96
assert config.moe_latent_size == 3584
config.kimi_kda_layers = tuple(
layer_number for layer_number in KIMI_K3_KDA_LAYERS if layer_number <= config.num_layers
)
config.kimi_linear_num_heads = 96
config.kimi_linear_head_dim = 128
config.kimi_linear_conv_kernel_size = 4
config.kimi_kda_gate_lower_bound = -5.0
config.kimi_attn_res_block_size = 12
config.gated_activation_func = situ_and_mul
config.moe_latent_use_norm = True
config.bias_activation_fusion = False
config.use_te_activation_func = False
def build_kimi_k3_spec(config, vp_stage=None):
assert config.virtual_pipeline_model_parallel_size is None, "Kimi K3 does not support VPP yet"
block_spec = get_gpt_decoder_block_spec(
config,
use_transformer_engine=True,
vp_stage=vp_stage,
)
# layer_specs holds only this stage's layers; moe_layer_freq is indexed globally
layer_offset = get_transformer_layer_offset(config, vp_stage)
layer_specs = []
for layer_spec in block_spec.layer_specs:
layer_spec = copy.deepcopy(layer_spec)
layer_spec.module = KimiK3TransformerLayer
layer_spec.submodules.self_attention = ModuleSpec(module=KimiK3Attention)
layer_spec.submodules.input_layernorm = TENorm
layer_spec.submodules.pre_mlp_layernorm = TENorm
if not config.moe_layer_freq[layer_offset + len(layer_specs)]:
# The dense MLP spec is partial(MLP.as_mlp_submodule, submodules=...).
layer_spec.submodules.mlp.keywords["submodules"].linear_fc1 = TEColumnParallelLinear
layer_specs.append(layer_spec)
block_spec.layer_specs = layer_specs
return block_spec
def get_kimi_k3_spec(args, config, vp_stage):
del args
configure_kimi_k3(config)
return build_kimi_k3_spec(config, vp_stage=vp_stage)
+86
View File
@@ -0,0 +1,86 @@
import torch
import torch.nn as nn
def kda(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
g: torch.Tensor,
beta: torch.Tensor,
A_log: torch.Tensor,
dt_bias: torch.Tensor,
lower_bound: float,
*,
cu_seqlens: torch.Tensor | None = None,
cp_context=None,
) -> torch.Tensor:
"""fla delta-rule core; boundaries travel as ``cu_seqlens`` without CP or ``cp_context`` under CP, never both."""
from fla.ops.kda import chunk_kda
boundaries = {"cp_context": cp_context} if cp_context is not None else {"cu_seqlens": cu_seqlens}
output, _ = chunk_kda(
q=q,
k=k,
v=v,
g=g,
beta=beta,
A_log=A_log,
dt_bias=dt_bias,
initial_state=None,
output_final_state=False,
use_qk_l2norm_in_kernel=True,
use_gate_in_kernel=True,
safe_gate=True,
lower_bound=lower_bound,
transpose_state_layout=True,
**boundaries,
)
return output
def situ_and_mul(
x: torch.Tensor,
beta: float = 4.0,
linear_beta: float = 25.0,
) -> torch.Tensor:
gate, linear = torch.chunk(x.float(), 2, dim=-1)
gate = beta * torch.tanh(gate / beta) * torch.sigmoid(gate)
linear = linear_beta * torch.tanh(linear / linear_beta)
return (gate * linear).to(x.dtype)
class KimiRMSNorm(nn.Module):
def __init__(
self,
hidden_size: int,
eps: float,
device: torch.device | int | None = None,
dtype: torch.dtype | None = None,
) -> None:
super().__init__()
self.weight = nn.Parameter(torch.ones(hidden_size, device=device, dtype=dtype))
self.eps = eps
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
input_dtype = hidden_states.dtype
normalized = hidden_states.float()
normalized = normalized * torch.rsqrt(normalized.square().mean(dim=-1, keepdim=True) + self.eps)
return self.weight * normalized.to(input_dtype)
def attn_res_aggregate(
prefix_sum: torch.Tensor,
block_residual: torch.Tensor,
score_proj: nn.Linear,
score_norm: KimiRMSNorm,
output_norm: nn.Module,
) -> torch.Tensor:
rows = torch.cat((block_residual, prefix_sum.unsqueeze(-2)), dim=-2)
rows_float = rows.float()
normalized = rows_float * torch.rsqrt(rows_float.square().mean(dim=-1, keepdim=True) + score_norm.eps)
score_weight = score_norm.weight.float() * score_proj.weight.squeeze(0).float()
scores = (normalized * score_weight).sum(dim=-1)
probabilities = torch.softmax(scores, dim=-1)
mixed = (probabilities.unsqueeze(-1) * rows_float).sum(dim=-2).to(rows.dtype)
return output_norm(mixed)
+28
View File
@@ -0,0 +1,28 @@
"""Stage-boundary packing for pipeline-parallel Kimi K3.
Megatron's p2p carries one hidden-states tensor, so the stage-exit layer packs
`[prefix_sum, snapshot bank]` along the hidden dim and the stage-entry layer unpacks it.
"""
import torch
def bank_num_rows(layer_idx: int, block_size: int) -> int:
"""Snapshot rows present before global layer ``layer_idx`` executes."""
assert layer_idx > 0
return (layer_idx + block_size - 1) // block_size
def pack_stage_boundary(prefix_sum: torch.Tensor, block_residual: torch.Tensor) -> torch.Tensor:
assert block_residual.shape[-2] > 0, "stage boundary before the first snapshot write"
return torch.cat((prefix_sum, block_residual.flatten(-2)), dim=-1)
def unpack_stage_boundary(packed: torch.Tensor, hidden_size: int, num_rows: int) -> tuple[torch.Tensor, torch.Tensor]:
expected = (1 + num_rows) * hidden_size
assert (
packed.shape[-1] == expected
), f"stage-boundary payload width {packed.shape[-1]} != (1 + {num_rows}) * {hidden_size}"
prefix_sum = packed[..., :hidden_size].contiguous()
block_residual = packed[..., hidden_size:].unflatten(-1, (num_rows, hidden_size)).contiguous()
return prefix_sum, block_residual
@@ -0,0 +1,5 @@
from model_args_utils import load_sibling_model_args
def model_args() -> str:
return load_sibling_model_args(__file__, "kimi-k3", nlayers=4, num_experts=64)
+5
View File
@@ -0,0 +1,5 @@
from model_args_utils import load_sibling_model_args
def model_args() -> str:
return load_sibling_model_args(__file__, "kimi-k3", nlayers=4)
+53
View File
@@ -0,0 +1,53 @@
import os
from model_args_utils import moe_layer_freq
def model_args(nlayers: int | None = None, num_experts: int = 896) -> str:
nlayers = nlayers if nlayers is not None else int(os.environ.get("MODEL_ARGS_NUM_LAYERS") or 93)
return (
"--spec miles_plugins.models.kimi_k3 get_kimi_k3_spec "
"--disable-bias-linear "
f"--num-layers {nlayers} "
"--hidden-size 7168 "
"--ffn-hidden-size 33792 "
"--num-attention-heads 96 "
"--num-query-groups 96 "
"--kv-channels 256 "
"--normalization RMSNorm "
"--position-embedding-type none "
"--rotary-base 10000 "
"--norm-epsilon 1e-5 "
"--hidden-dropout 0 "
"--attention-dropout 0 "
"--disable-bf16-reduced-precision-matmul "
"--swiglu "
"--untie-embeddings-and-output-weights "
"--vocab-size 163840 "
"--make-vocab-size-divisible-by 128 "
"--multi-latent-attention "
"--q-lora-rank 1536 "
"--kv-lora-rank 512 "
"--qk-head-dim 128 "
"--qk-pos-emb-head-dim 64 "
"--v-head-dim 128 "
"--qk-layernorm "
"--attention-softmax-in-fp32 "
f"--num-experts {num_experts} "
f"--moe-layer-freq {moe_layer_freq(nlayers=nlayers, first_k_dense_replace=1)} "
"--moe-ffn-hidden-size 3072 "
"--moe-latent-size 3584 "
"--moe-shared-expert-intermediate-size 6144 "
"--moe-router-topk 16 "
"--moe-router-score-function sigmoid "
"--moe-router-pre-softmax "
"--moe-router-enable-expert-bias "
"--moe-router-load-balancing-type none "
"--moe-token-dispatcher-type alltoall "
"--moe-router-bias-update-rate 0 "
"--moe-aux-loss-coeff 0 "
"--moe-router-topk-scaling-factor 1.0 "
"--moe-router-dtype fp32 "
"--moe-grouped-gemm "
"--moe-permute-fusion "
)
+551
View File
@@ -0,0 +1,551 @@
"""Kimi K3 RL launcher: LoRA or full-parameter training on the native MXFP4 checkpoint.
python scripts/run_kimi_k3.py prepare-download --model-name Kimi-K3-4layer-64experts
python scripts/run_kimi_k3.py prepare-bf16 --model-name Kimi-K3-4layer-64experts
python scripts/run_kimi_k3.py prepare-torch-dist --model-name Kimi-K3-4layer-64experts
python scripts/run_kimi_k3.py train --model-name Kimi-K3-4layer-64experts --train-mode lora
"""
import os
from dataclasses import dataclass
from pathlib import Path
from typing import Literal
import typer
import miles.utils.external_utils.command_utils as U
app = typer.Typer()
_DEFAULT_MODEL_ORG = {
"Kimi-K3": "moonshotai",
# prunes of the native MXFP4 checkpoint: the first dense layer plus three MoE layers, with all 896 routed
# experts or the first 64 of them
"Kimi-K3-4layer": "Pinaster",
"Kimi-K3-4layer-64experts": "Pinaster",
}
_MEGATRON_MODEL_TYPE = {
"Kimi-K3": "kimi-k3",
"Kimi-K3-4layer": "kimi-k3-4layer",
"Kimi-K3-4layer-64experts": "kimi-k3-4layer-64experts",
}
_NUM_LAYERS = {"Kimi-K3": 93, "Kimi-K3-4layer": 4, "Kimi-K3-4layer-64experts": 4}
_NUM_EXPERTS = {"Kimi-K3": 896, "Kimi-K3-4layer": 896, "Kimi-K3-4layer-64experts": 64}
_NUM_ATTENTION_HEADS = 96
_VALIDATED_FULL_MODEL_GPUS = 64
_LAYERS = "decoder.layers.*"
_DEFAULT_TARGET_MODULES = ",".join(
[
f"{_LAYERS}.self_attention.o_proj",
f"{_LAYERS}.self_attention.q_a_proj",
f"{_LAYERS}.self_attention.kv_a_proj_with_mqa",
f"{_LAYERS}.mlp.linear_fc1",
f"{_LAYERS}.mlp.linear_fc2",
f"{_LAYERS}.mlp.experts.linear_fc1",
f"{_LAYERS}.mlp.experts.linear_fc2",
]
)
@dataclass
class ScriptArgs(U.ExecuteTrainConfig):
mode: Literal["normal", "debug_minimal"] = "debug_minimal"
run_id: str = U.create_run_id()
model_org: str = ""
model_name: Literal["Kimi-K3", "Kimi-K3-4layer", "Kimi-K3-4layer-64experts"] = "Kimi-K3-4layer-64experts"
train_mode: Literal["lora", "full"] = "lora"
task: Literal["gsm8k", "dapo-math"] = "gsm8k"
hardware: Literal["auto", "H100", "H200", "B200", "B300", "GB200", "GB300"] = "auto"
num_gpus_per_node: int | None = None
# native MXFP4 checkpoint (rollout), its BF16 dequantization and the torch_dist conversion (trainer);
# None derives each from model_dir/model_name
hf_checkpoint: str | None = None
bf16_checkpoint: str | None = None
ref_load: str | None = None
data_dir: str = "/root/datasets"
model_dir: str = "/root/models"
save_dir: str = "/personal/checkpoints"
megatron_path: str = "/root/Megatron-LM"
sglang_path: str = "/root/sglang/python"
pipeline_parallel_size: int = 1
context_parallel_size: int = 1
# for layouts the PP-only derivation cannot express, e.g. TP4/CP2/PP4/EP8 (DP=2)
tp_size_override: int | None = None
ep_size_override: int | None = None
rollout_tp_size: int | None = None
rollout_ep_size: int | None = None
rollout_max_concurrency: int = 64
lora_rank: int = 16
lora_alpha: int = 32
lora_dropout: float = 0.0
target_modules: str = _DEFAULT_TARGET_MODULES
experts_shared_outer_loras: bool = True
reward_model: Literal["deterministic_random", "deepscaler", "math"] | None = None
num_rollout: int | None = None
rollout_batch_size: int | None = None
n_samples_per_prompt: int | None = None
rollout_max_response_len: int | None = None
sglang_max_total_tokens: int | None = None
global_batch_size: int | None = None
eval_interval: int | None = None
lr: float | None = None
max_tokens_per_gpu: int | None = None
distributed_timeout_minutes: int = 10
save_debug_rollout_data: str | None = None
enable_wandb: bool = False
skip_saving: bool = False
check_weight_update_equal: bool = False
check_lora_weight_equal: bool = False
update_weight_buffer_size: int | None = None
extra_args: str = ""
def __post_init__(self):
self.hardware = U.resolve_hardware(self)
self.num_gpus_per_node = self.num_gpus_per_node or U.NUM_GPUS_OF_HARDWARE[self.hardware]
if not self.model_org:
self.model_org = _DEFAULT_MODEL_ORG[self.model_name]
if self.hf_checkpoint is None:
self.hf_checkpoint = f"{self.model_dir}/{self.model_name}"
if self.bf16_checkpoint is None:
self.bf16_checkpoint = f"{self.model_dir}/{self.bf16_name}"
if self.ref_load is None:
self.ref_load = f"{self.model_dir}/{self.bf16_name}_torch_dist"
if self.lr is None:
self.lr = 1e-5 if self.train_mode == "lora" else 1e-6
if self.rollout_tp_size is None:
self.rollout_tp_size = min(8, self.num_gpus)
if self.rollout_ep_size is None:
# the Marlin LoRA MoE runner serves the experts replicated across the TP group
self.rollout_ep_size = 1
if self.train_mode == "lora" and self.lora_rank <= 0:
raise ValueError(f"LoRA rank must be positive, got {self.lora_rank}")
if self.sglang_max_total_tokens is not None and self.sglang_max_total_tokens <= 0:
raise ValueError("SGLang max total tokens must be positive")
if self.distributed_timeout_minutes <= 0:
raise ValueError("Distributed timeout must be positive")
if self.num_gpus % self.rollout_tp_size != 0 or _NUM_ATTENTION_HEADS % self.rollout_tp_size != 0:
raise ValueError(
f"rollout_tp_size must divide {self.num_gpus} GPUs and {_NUM_ATTENTION_HEADS} attention heads, "
f"got {self.rollout_tp_size}"
)
if self.rollout_tp_size % self.rollout_ep_size != 0 or self.num_experts % self.rollout_ep_size != 0:
raise ValueError(
f"rollout_ep_size must divide rollout_tp_size and {self.num_experts} experts, got {self.rollout_ep_size}"
)
if self.is_4layer:
if self.pipeline_parallel_size != 1 or self.context_parallel_size != 1:
raise NotImplementedError("Pipeline and context parallelism are only wired for the full model")
return
if self.num_gpus != _VALIDATED_FULL_MODEL_GPUS and self.tp_size_override is None:
raise ValueError(
f"The full-model layout is derived for {_VALIDATED_FULL_MODEL_GPUS} GPUs; pass --tp-size-override "
f"(and --ep-size-override) for {self.num_gpus} GPUs"
)
if self.pipeline_parallel_size > 1:
first, last = self.pipeline_layer_split
if not 1 <= last <= first:
raise ValueError(
f"Cannot split {self.num_layers} layers over pipeline_parallel_size={self.pipeline_parallel_size}: "
f"first={first}, last={last}"
)
model_parallel = self.tensor_parallel_size * self.context_parallel_size * self.pipeline_parallel_size
if self.num_gpus % model_parallel != 0:
raise ValueError(
f"TP{self.tensor_parallel_size}*CP{self.context_parallel_size}*PP{self.pipeline_parallel_size}"
f"={model_parallel} must divide the {self.num_gpus} training GPUs"
)
non_pp_ranks = self.tensor_parallel_size * self.context_parallel_size * (self.num_gpus // model_parallel)
if non_pp_ranks % self.expert_parallel_size != 0:
raise ValueError(
f"expert_parallel_size={self.expert_parallel_size} must divide the DP*CP*TP ranks "
f"of one pipeline stage ({non_pp_ranks})"
)
@property
def is_4layer(self) -> bool:
return self.model_name != "Kimi-K3"
@property
def num_layers(self) -> int:
return _NUM_LAYERS[self.model_name]
@property
def num_experts(self) -> int:
return _NUM_EXPERTS[self.model_name]
@property
def num_gpus(self) -> int:
return self.num_nodes * self.num_gpus_per_node
@property
def bf16_name(self) -> str:
return f"{self.model_name}-bf16"
@property
def megatron_model_type(self) -> str:
return _MEGATRON_MODEL_TYPE[self.model_name]
@property
def tensor_parallel_size(self) -> int:
if self.tp_size_override is not None:
return self.tp_size_override
if self.is_4layer:
return min(8, self.num_gpus)
if self.pipeline_parallel_size == 1:
return 32
return self.num_gpus // (self.pipeline_parallel_size * self.context_parallel_size)
@property
def expert_parallel_size(self) -> int:
if self.ep_size_override is not None:
return self.ep_size_override
if self.is_4layer:
return self.tensor_parallel_size
# EP is bounded by the DP*CP*TP ranks inside one pipeline stage.
model_parallel = self.tensor_parallel_size * self.context_parallel_size * self.pipeline_parallel_size
data_parallel = self.num_gpus // model_parallel
return self.tensor_parallel_size * self.context_parallel_size * data_parallel
@property
def pipeline_layer_split(self) -> tuple[int, int]:
"""(first, last) stage layer counts; middle stages each take `first` layers."""
pp = self.pipeline_parallel_size
first = -(-self.num_layers // pp)
last = self.num_layers - first * (pp - 1)
return first, last
def _download_dataset(args: ScriptArgs) -> None:
if args.task == "gsm8k":
U.hf_download_dataset("zhuzilin/gsm8k", data_dir=args.data_dir)
else:
U.hf_download_dataset("zhuzilin/dapo-math-17k", data_dir=args.data_dir)
U.hf_download_dataset("zhuzilin/aime-2024", data_dir=args.data_dir)
@app.command()
@U.dataclass_cli
def prepare_data(args: ScriptArgs) -> None:
_download_dataset(args)
def _prepare_download(args: ScriptArgs) -> None:
"""Native MXFP4 checkpoint + task dataset. Idempotent: hf skips existing blobs."""
U.exec_command_cpu(f"mkdir -p {args.model_dir} {args.data_dir}")
if args.hf_checkpoint == f"{args.model_dir}/{args.model_name}":
U.exec_command_cpu(f"hf download {args.model_org}/{args.model_name} --local-dir {args.hf_checkpoint}")
_download_dataset(args)
@app.command()
@U.dataclass_cli
def prepare_download(args: ScriptArgs) -> None:
_prepare_download(args)
def _prepare_bf16(args: ScriptArgs) -> None:
"""Dequantize the MXFP4 experts; Megatron loads BF16. One node, GPU."""
U.exec_command_gpu(
f"python {U.repo_base_dir}/tools/convert_mxfp4_to_bf16.py "
f"--model-dir {args.hf_checkpoint} --save-dir {args.bf16_checkpoint} --device cuda"
)
@app.command()
@U.dataclass_cli
def prepare_bf16(args: ScriptArgs) -> None:
_prepare_bf16(args)
def _prepare_torch_dist(args: ScriptArgs) -> None:
"""BF16 HF -> torch_dist in the training layout. The output re-shards at load."""
if not args.is_4layer:
raise NotImplementedError(
"The full model converts on 32 ranks; run tools/convert_hf_to_torch_dist.py as documented in "
"docs/models/kimi/kimi-k3.md"
)
# TP>1 needs CUDA_DEVICE_MAX_CONNECTIONS=1, which the converter does not set; EP alone shards the
# experts that dominate the 4-layer prune, and the torch_dist output re-shards at load
U.convert_checkpoint(
model_name=args.bf16_name,
megatron_model_type=args.megatron_model_type,
num_gpus_per_node=args.num_gpus_per_node,
extra_args=(
"--bf16 --tensor-model-parallel-size 1 "
"--pipeline-model-parallel-size 1 --context-parallel-size 1 "
f"--expert-model-parallel-size {args.num_gpus_per_node} --expert-tensor-parallel-size 1 "
"--megatron-to-hf-mode raw "
),
dir_dst=args.model_dir,
hf_checkpoint=args.bf16_checkpoint,
megatron_path=args.megatron_path,
)
@app.command()
@U.dataclass_cli
def prepare_torch_dist(args: ScriptArgs) -> None:
_prepare_torch_dist(args)
def _train(args: ScriptArgs) -> None:
is_debug = args.mode == "debug_minimal"
is_lora = args.train_mode == "lora"
if args.task == "gsm8k":
dataset = Path(args.data_dir) / "gsm8k" / "train.parquet"
input_key = "messages"
else:
dataset = Path(args.data_dir) / "dapo-math-17k" / "dapo-math-17k.jsonl"
input_key = "prompt"
ckpt_args = (
f"--hf-checkpoint {args.hf_checkpoint} "
f"--ref-load {args.ref_load} "
"--megatron-to-hf-mode raw "
"--model-name kimi_k3 "
)
lora_args = ""
if is_lora:
lora_args = (
f"--lora-rank {args.lora_rank} "
f"--lora-alpha {args.lora_alpha} "
f"--lora-dropout {args.lora_dropout} "
f'--target-modules "{args.target_modules}" '
"--no-gradient-accumulation-fusion "
# host mirror of the rollout base, so releasing it does not re-ship the base every step
"--lora-base-cpu-backup "
)
if args.experts_shared_outer_loras:
lora_args += "--experts-shared-outer-loras "
if args.check_lora_weight_equal:
lora_args += "--check-lora-weight-equal "
reward_model = args.reward_model or (
"math" if args.task == "gsm8k" else ("deterministic_random" if is_debug else "deepscaler")
)
num_rollout = args.num_rollout if args.num_rollout is not None else (2 if is_debug else 3000)
rollout_batch_size = args.rollout_batch_size if args.rollout_batch_size is not None else (8 if is_debug else 32)
n_samples_per_prompt = (
args.n_samples_per_prompt if args.n_samples_per_prompt is not None else (2 if is_debug else 8)
)
rollout_max_response_len = (
args.rollout_max_response_len
if args.rollout_max_response_len is not None
else (256 if args.task == "gsm8k" else (32 if is_debug else 16384))
)
global_batch_size = args.global_batch_size if args.global_batch_size is not None else (16 if is_debug else 256)
rollout_args = (
f"--prompt-data {dataset} "
f"--input-key {input_key} "
"--label-key label "
"--apply-chat-template "
"--rollout-shuffle "
"--balance-data "
f"--rm-type {reward_model} "
f"--num-rollout {num_rollout} "
f"--rollout-batch-size {rollout_batch_size} "
f"--n-samples-per-prompt {n_samples_per_prompt} "
f"--rollout-max-response-len {rollout_max_response_len} "
"--rollout-temperature 1 "
f"--global-batch-size {global_batch_size} "
"--use-dynamic-global-batch-size "
)
if args.eval_interval is not None:
# eval on the training task's held-out set, keyed by the eval dataset's own column
if args.task == "gsm8k":
eval_name, eval_rel, eval_input_key = "gsm8k", "gsm8k/test.parquet", "messages"
else:
eval_name, eval_rel, eval_input_key = "aime", "aime-2024/aime-2024.jsonl", "prompt"
rollout_args += (
f"--eval-interval {args.eval_interval} "
f"--eval-prompt-data {eval_name} {Path(args.data_dir) / eval_rel} "
f"--eval-input-key {eval_input_key} "
"--n-samples-per-eval-prompt 1 "
"--eval-temperature 0 "
f"--eval-max-response-len {rollout_max_response_len} "
)
if args.save_debug_rollout_data is not None:
rollout_args += f"--save-debug-rollout-data {args.save_debug_rollout_data} "
pp_split_args = ""
if args.pipeline_parallel_size > 1:
first, last = args.pipeline_layer_split
if first != last:
pp_split_args = f"--decoder-first-pipeline-num-layers {first} --decoder-last-pipeline-num-layers {last} "
max_tokens_per_gpu = (
args.max_tokens_per_gpu if args.max_tokens_per_gpu is not None else (512 if args.is_4layer else 8192)
)
perf_args = (
f"--tensor-model-parallel-size {args.tensor_parallel_size} "
"--sequence-parallel "
f"--pipeline-model-parallel-size {args.pipeline_parallel_size} "
f"{pp_split_args}"
f"--context-parallel-size {args.context_parallel_size} "
f"--expert-model-parallel-size {args.expert_parallel_size} "
"--expert-tensor-parallel-size 1 "
"--recompute-granularity full "
"--recompute-method uniform "
"--recompute-num-layers 1 "
"--use-dynamic-batch-size "
f"--max-tokens-per-gpu {max_tokens_per_gpu} "
"--log-probs-chunk-size 512 "
f"--distributed-timeout-minutes {args.distributed_timeout_minutes} "
)
optimizer_args = (
"--optimizer sgd --sgd-momentum 0 --lr 1e-6 --lr-decay-style constant "
"--weight-decay 0 --use-distributed-optimizer "
if is_debug
else (
f"--optimizer adam --lr {args.lr} --lr-decay-style constant "
"--weight-decay 0.1 --adam-beta1 0.9 --adam-beta2 0.98 "
"--optimizer-cpu-offload --optimizer-offload-fraction 0.8 "
"--overlap-cpu-optimizer-d2h-h2d "
"--use-precision-aware-optimizer --use-distributed-optimizer "
)
)
grpo_args = (
"--advantage-estimator grpo "
"--kl-loss-coef 0.0 "
"--kl-loss-type low_var_kl "
"--entropy-coef 0.0 "
"--eps-clip 0.2 "
"--eps-clip-high 0.28 "
)
update_weight_buffer_size = args.update_weight_buffer_size
if update_weight_buffer_size is None:
update_weight_buffer_size = 2 * 1024**3 if args.is_4layer or not is_lora else 256 * 1024**2
# the radix extra-buffer strategy needs five KDA cache slots per running request
sglang_request_capacity = 16 if args.is_4layer else args.rollout_max_concurrency
sglang_mamba_capacity = 16 if args.is_4layer else 5 * args.rollout_max_concurrency
graph_bs = " ".join(str(b) for b in (1, 2, 4, 8, 16, 32) if b <= max(1, args.rollout_max_concurrency))
sglang_args = (
f"--rollout-num-gpus-per-engine {args.rollout_tp_size} "
f"--sglang-tp-size {args.rollout_tp_size} "
f"--sglang-ep-size {args.rollout_ep_size} "
"--sglang-server-concurrency 16 "
f"--sglang-max-running-requests {sglang_request_capacity} "
f"--sglang-max-mamba-cache-size {sglang_mamba_capacity} "
"--use-miles-router "
)
if is_lora:
sglang_args += (
"--sglang-lora-backend triton " "--sglang-lora-strict-loading " f"--sglang-max-lora-rank {args.lora_rank} "
)
if args.rollout_tp_size > 8:
# above TP8 the Marlin MoE intermediate is tile-padded, which the virtual-experts LoRA kernel rejects
sglang_args += "--no-sglang-lora-use-virtual-experts "
if not args.is_4layer:
# the adapter is re-streamed every step; a host copy per TP rank (~45 GB) is never read
sglang_args += "--sglang-lora-no-cpu-backup "
# Marlin is the one MXFP4 MoE runner with a LoRA path (Mxfp4MoEMethod has no triton quant info)
sglang_args += "--sglang-moe-runner-backend marlin "
if args.is_4layer:
sglang_args += (
"--sglang-cuda-graph-bs-decode 1 2 4 8 16 "
"--sglang-mem-fraction-static 0.7 "
"--sglang-disable-shared-experts-fusion "
)
else:
sglang_args += (
"--sglang-decode-attention-backend trtllm_mla "
"--sglang-mamba-radix-cache-strategy extra_buffer "
f"--sglang-cuda-graph-bs-decode {graph_bs} "
"--sglang-cuda-graph-backend-prefill disabled "
)
if args.sglang_max_total_tokens is not None:
sglang_args += f"--sglang-max-total-tokens {args.sglang_max_total_tokens} "
if is_debug:
sglang_args += "--sglang-context-length 8192 "
misc_args = (
"--attention-dropout 0.0 "
"--hidden-dropout 0.0 "
"--accumulate-allreduce-grads-in-fp32 "
"--colocate "
"--offload-train "
f"--update-weight-buffer-size {update_weight_buffer_size} "
f"--train-memory-margin-bytes {(2 if is_debug else 4) * 1024**3} "
f"--actor-num-nodes {args.num_nodes} "
f"--actor-num-gpus-per-node {args.num_gpus_per_node} "
f"--num-gpus-per-node {args.num_gpus_per_node} "
)
if args.is_4layer:
misc_args += "--no-check-for-nan-in-loss-and-grad "
# the trainer has no vision tower, and the MXFP4 experts round-trip through BF16 on every sync
misc_args += "--check-weight-update-skip-list vision_tower. mm_projector. --check-weight-update-allow-quant-error "
if args.check_weight_update_equal:
misc_args += "--check-weight-update-equal "
if not args.skip_saving:
misc_args += f"--save {args.save_dir}/{args.run_id} --save-interval 50 "
if args.enable_wandb:
wandb_args = (
"--use-wandb "
"--wandb-project miles-run_kimi_k3 "
f"--wandb-group {args.run_id} "
"--disable-wandb-random-suffix "
)
else:
wandb_args = U.get_default_wandb_args(__file__, run_id=args.run_id)
train_args = (
f"{ckpt_args} "
f"{lora_args} "
f"{rollout_args} "
f"{optimizer_args} "
f"{grpo_args} "
f"{wandb_args} "
f"{perf_args} "
f"{sglang_args} "
f"{misc_args} "
f"{args.extra_args} "
)
extra_env_vars = {
"NCCL_TIMEOUT": "3600",
"PYTHONPATH": os.pathsep.join((str(Path(__file__).resolve().parents[1]), args.sglang_path)),
"SGLANG_JIT_ROUTE_RADIX": "1",
# sglang's membind pins the whole TMS host backup to one NUMA node, which cannot hold it
"SGLANG_NUMA_BIND_V2": os.environ.get("SGLANG_NUMA_BIND_V2", "0"),
}
if is_lora:
# the LoRA wrapper bypasses the o_proj.forward patch the K3 all-reduce fusion relies on
extra_env_vars["SGLANG_K3_AR_FUSION"] = "0"
# Ray's runtime_env replaces the actor env; without the JIT cache dirs the KDA kernels recompile every run
for cache_var in ("TRITON_CACHE_DIR", "TORCHINDUCTOR_CACHE_DIR"):
cache_dir = os.environ.get(cache_var)
if cache_dir:
extra_env_vars[cache_var] = cache_dir
U.execute_train(
train_args=train_args,
config=args,
num_gpus_per_node=args.num_gpus_per_node,
megatron_model_type=args.megatron_model_type,
extra_env_vars=extra_env_vars,
megatron_path=args.megatron_path,
)
@app.command()
@U.dataclass_cli
def train(args: ScriptArgs) -> None:
_train(args)
if __name__ == "__main__":
app()
@@ -0,0 +1,60 @@
import os
from scripts.run_kimi_k3 import ScriptArgs, _prepare_bf16, _prepare_download, _prepare_torch_dist, _train
from tests.ci.ci_register import register_cuda_ci
from tests.ci.metric_history import register_ci_gate
register_cuda_ci(
est_time=3000,
suite="stage-c-8-gpu-h200",
labels=["megatron", "model-scripts"],
hardware=["hopper", "blackwell"],
)
register_ci_gate(metric_key="train/grad_norm")
register_ci_gate(metric_key="train/ppo_kl")
register_ci_gate(metric_key="train/train_rollout_logprob_abs_diff")
register_ci_gate(metric_key="train/train_rollout_kl")
register_ci_gate(metric_key="rollout/raw_reward")
def _args() -> ScriptArgs:
return ScriptArgs(
model_name="Kimi-K3-4layer-64experts",
train_mode="full",
mode="normal",
task="gsm8k",
# the pruned model scores 0 on gsm8k, which zeroes every advantage; a fixed pseudo-random
# reward keeps the weights moving so the MXFP4 re-quantization carries real deltas
reward_model="deterministic_random",
hardware="H200",
num_nodes=1,
num_gpus_per_node=8,
num_rollout=2,
rollout_batch_size=8,
n_samples_per_prompt=8,
global_batch_size=64,
rollout_max_response_len=256,
rollout_max_concurrency=16,
check_weight_update_equal=True,
skip_saving=True,
extra_args="--ci-test --ci-disable-logprobs-checker ",
)
def prepare(args: ScriptArgs):
_prepare_download(args)
_prepare_bf16(args)
_prepare_torch_dist(args)
def execute(args: ScriptArgs):
_train(args)
if __name__ == "__main__":
args = _args()
prepare(args)
for proxy_var in ("http_proxy", "https_proxy", "HTTP_PROXY", "HTTPS_PROXY"):
os.environ.pop(proxy_var, None)
execute(args)
@@ -0,0 +1,62 @@
import os
from scripts.run_kimi_k3 import ScriptArgs, _prepare_bf16, _prepare_download, _prepare_torch_dist, _train
from tests.ci.ci_register import register_cuda_ci
from tests.ci.metric_history import register_ci_gate
register_cuda_ci(
est_time=2400,
suite="stage-c-8-gpu-h200",
labels=["megatron", "model-scripts", "lora"],
hardware=["hopper", "blackwell"],
)
register_ci_gate(metric_key="train/grad_norm")
register_ci_gate(metric_key="train/ppo_kl")
register_ci_gate(metric_key="train/train_rollout_logprob_abs_diff")
register_ci_gate(metric_key="train/train_rollout_kl")
register_ci_gate(metric_key="rollout/raw_reward")
def _args() -> ScriptArgs:
return ScriptArgs(
model_name="Kimi-K3-4layer-64experts",
train_mode="lora",
mode="normal",
task="gsm8k",
# the pruned model scores 0 on gsm8k, which zeroes every advantage; a fixed pseudo-random
# reward keeps the adapter moving so the weight sync carries real deltas
reward_model="deterministic_random",
hardware="H200",
num_nodes=1,
num_gpus_per_node=8,
lora_rank=32,
lora_alpha=64,
num_rollout=2,
rollout_batch_size=8,
n_samples_per_prompt=8,
global_batch_size=64,
rollout_max_response_len=256,
rollout_max_concurrency=16,
check_lora_weight_equal=True,
skip_saving=True,
extra_args="--ci-test --ci-disable-logprobs-checker ",
)
def prepare(args: ScriptArgs):
_prepare_download(args)
_prepare_bf16(args)
_prepare_torch_dist(args)
def execute(args: ScriptArgs):
_train(args)
if __name__ == "__main__":
args = _args()
prepare(args)
for proxy_var in ("http_proxy", "https_proxy", "HTTP_PROXY", "HTTPS_PROXY"):
os.environ.pop(proxy_var, None)
execute(args)
+59 -53
View File
@@ -1,80 +1,86 @@
"""Tests for the prefix-matching ignore rule added to quantize_params_compressed_tensors.
The change adds `name.startswith(r)` to the ignore matching logic, so rules like
"model.layers.0.self_attn" now ignore all weights under that prefix.
"""
"""`quantize_params_compressed_tensors` re-quantizes exactly the weights the checkpoint stores packed."""
from tests.ci.ci_register import register_cuda_ci
# The quantizer hardcodes `device="cuda"` throughout; this test drives it with
# real CUDA tensors to exercise the ignore-rule name-matching path. Fast enough
# for the GPU fast suite; only needs 1 GPU.
# The quantizer hardcodes `device="cuda"` throughout, so it runs on a GPU worker.
register_cuda_ci(est_time=60, suite="stage-b-2-gpu-h200", labels=["precision"], hardware=["hopper", "blackwell"])
import pytest
import torch
from miles.backends.megatron_utils.megatron_to_hf.processors.quantizer_compressed_tensors import (
quantize_params_compressed_tensors,
)
from miles.utils.mxfp4 import dequantize_mxfp4, quantize_mxfp4
CONFIG = {
"format": "int-quantized",
"config_groups": {"group_0": {"weights": {"group_size": 128, "symmetric": True}}},
"ignore": [], # overridden per test
}
def _quantize_names(name, ignore_rules):
"""Run quantization on a single 2D weight and return output names."""
config = {**CONFIG, "ignore": ignore_rules}
results = quantize_params_compressed_tensors([(name, torch.randn(256, 256, device="cuda"))], config)
return [r[0] for r in results]
def test_only_the_packed_basenames_are_quantized():
"""A 2-D BF16 weight the checkpoint keeps unpacked (router, residual projection) must pass through."""
params = [
("model.layers.0.experts.0.w1.weight", torch.randn(256, 256, device="cuda")),
("model.layers.0.gate.weight", torch.randn(256, 256, device="cuda")),
("model.layers.0.self_attn.q_proj.weight", torch.randn(256, 256, device="cuda")),
]
result_names = [
name for name, _ in quantize_params_compressed_tensors(params, CONFIG, {"model.layers.0.experts.0.w1"})
]
assert "model.layers.0.experts.0.w1.weight_packed" in result_names
assert "model.layers.0.experts.0.w1.weight" not in result_names
assert "model.layers.0.gate.weight" in result_names
assert "model.layers.0.self_attn.q_proj.weight" in result_names
def _is_ignored(name, ignore_rules):
"""Check if a weight name is ignored (returned as-is, not quantized)."""
names = _quantize_names(name, ignore_rules)
return name in names and f"{name}_packed" not in names
def test_mxfp4_quantize_inverts_the_checkpoint_dequantizer():
"""``quantize_mxfp4`` must be the exact inverse of the dequantizer used to
read the K3 checkpoint: any value that came out of an MXFP4 checkpoint has
to re-encode to the same bits, or every weight sync degrades the rollout
weights a little further. The magnitude thresholds, sign-bit position,
nibble order and exponent bias all have to agree.
Also pins the mxfp4 branch of ``quantize_params_compressed_tensors``, which
emits only weight_packed/weight_scale -- the int-quantized branch's extra
weight_shape/weight_zero_point tensors would be rejected by the receiver.
"""
group_size = 32
packed = torch.randint(0, 256, (16, 64), dtype=torch.uint8, device="cuda")
scale = torch.randint(96, 144, (16, 4), dtype=torch.uint8, device="cuda")
weight = dequantize_mxfp4(packed, scale, group_size)
class TestIgnoreRulePrefixMatching:
"""Tests for the new prefix-matching ignore rule (name.startswith(r))."""
actual_packed, actual_scale = quantize_mxfp4(weight, group_size)
actual = dequantize_mxfp4(actual_packed, actual_scale, group_size)
@pytest.mark.parametrize(
"rule,name,expected_ignored",
[
# exact match (pre-existing)
("model.layer.weight", "model.layer.weight", True),
# regex match (pre-existing)
("re:.*embed.*", "model.embed_tokens.weight", True),
# prefix match (NEW)
("model.layers.0.self_attn", "model.layers.0.self_attn.q_proj.weight", True),
("model.layers.", "model.layers.5.attn.weight", True),
("model.embed", "model.embed_tokens.weight", True),
# non-matching
("model.layers.1", "model.layers.0.attn.weight", False),
("other.prefix", "model.layers.0.attn.weight", False),
],
torch.testing.assert_close(actual, weight, rtol=0, atol=0)
config = {
"format": "mxfp4-pack-quantized",
"config_groups": {
"group_0": {
"weights": {
"group_size": group_size,
"symmetric": True,
"type": "float",
"num_bits": 4,
"scale_dtype": "torch.uint8",
}
}
},
}
results = quantize_params_compressed_tensors(
[("model.layers.0.experts.0.w1.weight", torch.randn(64, 64, device="cuda"))],
config,
{"model.layers.0.experts.0.w1"},
)
def test_ignore_rule_matching(self, rule, name, expected_ignored):
assert _is_ignored(name, [rule]) == expected_ignored
def test_prefix_selectively_ignores(self):
"""Prefix rule ignores matching params while others get quantized."""
config = {**CONFIG, "ignore": ["model.layers.0.self_attn"]}
params = [
("model.layers.0.self_attn.q_proj.weight", torch.randn(256, 256, device="cuda")),
("model.layers.0.self_attn.k_proj.weight", torch.randn(256, 256, device="cuda")),
("model.layers.0.mlp.gate_proj.weight", torch.randn(256, 256, device="cuda")),
]
result_names = [r[0] for r in quantize_params_compressed_tensors(params, config)]
# self_attn params ignored (passed through)
assert "model.layers.0.self_attn.q_proj.weight" in result_names
assert "model.layers.0.self_attn.k_proj.weight" in result_names
# mlp param quantized
assert "model.layers.0.mlp.gate_proj.weight_packed" in result_names
assert [(name, tensor.dtype, tuple(tensor.shape)) for name, tensor in results] == [
("model.layers.0.experts.0.w1.weight_packed", torch.uint8, (64, 32)),
("model.layers.0.experts.0.w1.weight_scale", torch.uint8, (64, 2)),
]
if __name__ == "__main__":
@@ -0,0 +1,33 @@
from types import SimpleNamespace
import torch
from miles.backends.megatron_utils.megatron_to_hf.kimi_k3 import convert_kimi_k3_to_hf
def _convert(name, shape=(4, 2)):
param = torch.arange(torch.tensor(shape).prod()).reshape(shape)
return convert_kimi_k3_to_hf(SimpleNamespace(), name, param)
def test_fused_fc1_splits_gate_before_up():
"""Megatron fuses gate and up into one linear_fc1; HF wants them separate.
Which half is which is pure convention -- swapping them produces a model
that runs and generates plausible-looking garbage, on both the dense MLP
and the per-expert path.
"""
dense = _convert("module.module.decoder.layers.0.mlp.linear_fc1.weight")
assert [name for name, _param in dense] == [
"language_model.model.layers.0.mlp.gate_proj.weight",
"language_model.model.layers.0.mlp.up_proj.weight",
]
assert dense[0][1].tolist() == [[0, 1], [2, 3]]
assert dense[1][1].tolist() == [[4, 5], [6, 7]]
expert = _convert("module.module.decoder.layers.2.mlp.experts.linear_fc1.weight17")
assert [name for name, _param in expert] == [
"language_model.model.layers.2.block_sparse_moe.experts.17.w1.weight",
"language_model.model.layers.2.block_sparse_moe.experts.17.w3.weight",
]
assert expert[0][1].tolist() == [[0, 1], [2, 3]]
assert expert[1][1].tolist() == [[4, 5], [6, 7]]
@@ -68,6 +68,7 @@ class TestSetupModelAndOptimizerLoraBranch:
lora_rank=lora_rank,
lora_adapter_path=None,
megatron_to_hf_mode=mode,
model_name=None,
moe_use_upcycling=False,
debug_disable_optimizer=False,
stream_optimizer_state_to_disk=False,
@@ -437,6 +437,8 @@ class TestSaveLoraCheckpointTrainingState:
model = [SimpleNamespace(named_parameters=lambda: [("layers.0.self_attention.lora_A.weight", adapter)])]
args = Namespace(
hf_checkpoint="/nonexistent",
# the bridge export path; raw mode writes the rank-sharded config instead
megatron_to_hf_mode="bridge",
target_modules=None,
lora_rank=8,
lora_alpha=16,
@@ -0,0 +1,16 @@
import torch
from miles.utils.mxfp4 import dequantize_mxfp4
def test_dequantize_mxfp4_decodes_nibbles_and_e8m0_scales() -> None:
"""Golden values pin the wire convention a pack/unpack round trip cannot catch."""
packed = torch.tensor([[0x10, 0x32, 0x54, 0x76]], dtype=torch.uint8)
scales = torch.tensor([127, 128], dtype=torch.uint8)
expected = torch.tensor(
[[0.0, 0.5, 1.0, 1.5, 4.0, 6.0, 8.0, 12.0]],
dtype=torch.bfloat16,
)
actual = dequantize_mxfp4(packed, scales, group_size=4)
torch.testing.assert_close(actual, expected, rtol=0, atol=0)
+121
View File
@@ -0,0 +1,121 @@
"""Boundary and aliasing contracts of the Kimi K3 KDA delta-rule core.
Both cases guard failure modes that are silent in training: the run keeps
learning, only worse, so nothing short of a numerical check catches them.
"""
import pytest
import torch
from miles_plugins.models.kimi_k3.ops import kda
# KDA kernel construction, validated on H200: q/k/v/g scaled bf16, beta uniform on
# [0, 1] because it is the delta rule's step size and falls outside the operator's
# domain when negative or above 1. (The "KDA goes non-finite with random weights"
# result is model-level -- random *projections* produce out-of-distribution q/k/v/g
# that compound across layers -- and does not apply to the kernel in isolation.)
_KDA_HEADS = 4
_KDA_HEAD_DIM = 128
_KDA_LOWER_BOUND = -5.0
def _kda_inputs(seq_len: int, seed: int) -> dict[str, torch.Tensor]:
torch.manual_seed(seed)
def activation() -> torch.Tensor:
return torch.randn(1, seq_len, _KDA_HEADS, _KDA_HEAD_DIM, device="cuda", dtype=torch.bfloat16) * 0.5
return {
"q": activation(),
"k": activation(),
"v": activation(),
"g": activation(),
"beta": torch.rand(1, seq_len, _KDA_HEADS, device="cuda", dtype=torch.float32),
"A_log": torch.randn(_KDA_HEADS, device="cuda", dtype=torch.float32),
"dt_bias": torch.randn(_KDA_HEADS * _KDA_HEAD_DIM, device="cuda", dtype=torch.float32),
}
def _run_kda(inputs: dict[str, torch.Tensor], cu_seqlens: torch.Tensor | None) -> torch.Tensor:
return kda(
inputs["q"],
inputs["k"],
inputs["v"],
inputs["g"],
inputs["beta"],
inputs["A_log"],
inputs["dt_bias"],
_KDA_LOWER_BOUND,
cu_seqlens=cu_seqlens,
)
def _require_kda() -> None:
if not torch.cuda.is_available():
pytest.skip("CUDA device required to run the KDA kernels")
pytest.importorskip("fla.ops.kda")
def _relative_l2(actual: torch.Tensor, expected: torch.Tensor) -> float:
return ((actual - expected).float().norm() / expected.float().norm()).item()
def test_kda_applies_packed_sequence_boundaries() -> None:
"""A packed batch must reproduce, per sequence, what that sequence produces alone.
``kda`` routes boundaries through one of two mutually exclusive channels
(``cu_seqlens`` off CP, ``cp_context`` under CP). A caller that selects the CP
channel and drops ``cu_seqlens`` -- the shape the CP-only helper had before the
two kernel paths were merged -- leaves the recurrence with no boundaries at all,
so the whole packed microbatch becomes one sequence and every sample inherits its
predecessors' state. Nothing raises.
Thresholds are relative because the assertion must not be satisfiable by the
output being small: these activations sit around 1e-3, so any absolute tolerance
worth the name swallows the entire signal. Measured on H200: the per-sequence
equality is exact, and dropping the boundaries moves the second sequence by 0.40
while leaving the first at exactly 0 -- the first sequence has no predecessor to
inherit from, which is what makes this leakage rather than noise.
"""
_require_kda()
first_len, second_len = 96, 160
total = first_len + second_len
packed = _kda_inputs(total, seed=0)
packed_output = _run_kda(packed, torch.tensor([0, first_len, total], dtype=torch.int32, device="cuda"))
for offset, length in ((0, first_len), (first_len, second_len)):
alone = {
name: tensor[:, offset : offset + length] if tensor.dim() > 1 else tensor
for name, tensor in packed.items()
}
alone_output = _run_kda(alone, torch.tensor([0, length], dtype=torch.int32, device="cuda"))
torch.testing.assert_close(packed_output[:, offset : offset + length], alone_output, rtol=0, atol=0)
unbounded_output = _run_kda(packed, None)
assert _relative_l2(unbounded_output[:, first_len:], packed_output[:, first_len:]) > 1e-1, (
"dropping cu_seqlens barely moved the second sequence, so boundaries are not reaching the kernel "
"and the equality above passed vacuously"
)
torch.testing.assert_close(unbounded_output[:, :first_len], packed_output[:, :first_len], rtol=0, atol=0)
def test_kda_does_not_mutate_its_inputs() -> None:
"""The forward must leave q/k/v/g/beta untouched.
SGLang's vendored ``chunk_kda`` overwrote its ``v`` buffer with a WY-representation
intermediate. Any caller that reads an input back afterwards -- a manual
``save_for_backward`` re-derivation being the case that bit us -- then
differentiates at a point the forward never evaluated, which put the q/k/g/beta
gradients ~110x low and near-orthogonal. Exact comparison: aliasing is a
mutation, not a rounding.
"""
_require_kda()
seq_len = 256
inputs = _kda_inputs(seq_len, seed=1)
before = {name: tensor.clone() for name, tensor in inputs.items()}
_run_kda(inputs, torch.tensor([0, seq_len], dtype=torch.int32, device="cuda"))
for name, tensor in inputs.items():
torch.testing.assert_close(tensor, before[name], rtol=0, atol=0, msg=f"KDA forward mutated input {name}")
+191
View File
@@ -0,0 +1,191 @@
from types import SimpleNamespace
import pytest
import torch
import torch.nn as nn
from torch.utils.checkpoint import checkpoint
from miles_plugins.models.kimi_k3.lora import (
KimiK3LoRAAdapter,
_enable_full_recompute_input_grads,
_grouped_linear,
export_kimi_k3_lora_hf_chunks,
)
def _parameter(*shape):
return nn.Parameter(torch.arange(torch.tensor(shape).prod()).reshape(shape).float())
def _mla_attention_adapter():
adapter = KimiK3LoRAAdapter("mla_attention", "language_model.model.layers.3.self_attn.")
adapter.register_parameter("q_a_lora_A", _parameter(2, 8))
adapter.register_parameter("q_a_lora_B", _parameter(4, 2))
adapter.register_parameter("kv_a_lora_A", _parameter(2, 8))
adapter.register_parameter("kv_a_lora_B", _parameter(6, 2))
adapter.register_parameter("o_lora_A", _parameter(2, 3))
adapter.register_parameter("o_lora_B", _parameter(8, 2))
return adapter
def _expert_adapter():
adapter = KimiK3LoRAAdapter(
"experts",
"language_model.model.layers.4.block_sparse_moe.experts.",
)
adapter.register_parameter("w1_lora_A", _parameter(2, 8))
adapter.register_parameter("w3_lora_A", _parameter(2, 8))
adapter.register_parameter("w1_lora_B", _parameter(3, 5, 2))
adapter.register_parameter("w3_lora_B", _parameter(3, 5, 2))
adapter.register_parameter("w2_lora_A", _parameter(3, 2, 5))
adapter.register_parameter("w2_lora_B", _parameter(8, 2))
return adapter
def _shared_expert_adapter():
adapter = KimiK3LoRAAdapter(
"shared_experts",
"language_model.model.layers.4.block_sparse_moe.shared_experts.",
)
adapter.register_parameter("fc1_lora_A", _parameter(2, 8))
adapter.register_parameter("fc1_lora_B", _parameter(10, 2))
adapter.register_parameter("fc2_lora_A", _parameter(2, 5))
adapter.register_parameter("fc2_lora_B", _parameter(8, 2))
return adapter
def _dense_adapter():
adapter = KimiK3LoRAAdapter(
"dense_mlp",
"language_model.model.layers.3.mlp.",
)
adapter.register_parameter("fc1_lora_A", _parameter(2, 8))
adapter.register_parameter("fc1_lora_B", _parameter(10, 2))
adapter.register_parameter("fc2_lora_A", _parameter(2, 5))
adapter.register_parameter("fc2_lora_B", _parameter(8, 2))
return adapter
def _kda_attention_adapter():
adapter = KimiK3LoRAAdapter(
"kda_attention",
"language_model.model.layers.4.self_attn.",
)
adapter.register_parameter("o_lora_A", _parameter(2, 3))
adapter.register_parameter("o_lora_B", _parameter(8, 2))
return adapter
def _model_with_adapters(*, include_shared_experts=True):
mla_attention = _mla_attention_adapter()
dense = _dense_adapter()
kda_attention = _kda_attention_adapter()
experts = _expert_adapter()
shared_experts = _shared_expert_adapter()
adapters = [mla_attention, dense, kda_attention, experts]
if include_shared_experts:
adapters.append(shared_experts)
model = nn.Module()
model.adapters = nn.ModuleList(adapters)
model.decoder = SimpleNamespace(
layers=[
SimpleNamespace(
layer_number=4,
self_attention=SimpleNamespace(is_kda=False),
mlp=SimpleNamespace(),
),
SimpleNamespace(
layer_number=5,
self_attention=SimpleNamespace(is_kda=True),
mlp=SimpleNamespace(
experts=SimpleNamespace(),
shared_experts=SimpleNamespace(),
),
),
]
)
return model
def test_native_export_is_chunked_by_adapter(monkeypatch):
"""A wrong expert dim is a shape mismatch in SGLang's LoRA pool; a wrong HF name is silently dropped."""
from megatron.core import parallel_state
monkeypatch.setattr(parallel_state, "get_tensor_model_parallel_world_size", lambda: 1)
monkeypatch.setattr(parallel_state, "get_expert_model_parallel_world_size", lambda: 1)
model = _model_with_adapters()
chunks = list(export_kimi_k3_lora_hf_chunks([model]))
assert len(chunks) == 5
attention = dict(chunks[0])
assert attention["language_model.model.layers.3.self_attn.q_a_proj.lora_A.weight"].shape == (2, 8)
assert attention["language_model.model.layers.3.self_attn.kv_a_proj_with_mqa.lora_B.weight"].shape == (
6,
2,
)
assert attention["language_model.model.layers.3.self_attn.o_proj.lora_A.weight"].shape == (2, 3)
experts = dict(chunks[3])
prefix = "language_model.model.layers.4.block_sparse_moe.experts."
assert experts[f"{prefix}w1.lora_A.weight"].shape == (1, 2, 8)
assert experts[f"{prefix}w1.lora_B.weight"].shape == (3, 5, 2)
assert experts[f"{prefix}w2.lora_A.weight"].shape == (3, 2, 5)
assert experts[f"{prefix}w2.lora_B.weight"].shape == (1, 8, 2)
shared_experts = dict(chunks[4])
prefix = "language_model.model.layers.4.block_sparse_moe.shared_experts."
assert shared_experts[f"{prefix}gate_proj.lora_A.weight"].shape == (2, 8)
assert shared_experts[f"{prefix}gate_proj.lora_B.weight"].shape == (5, 2)
assert shared_experts[f"{prefix}up_proj.lora_A.weight"].shape == (2, 8)
assert shared_experts[f"{prefix}up_proj.lora_B.weight"].shape == (5, 2)
assert shared_experts[f"{prefix}down_proj.lora_A.weight"].shape == (2, 5)
assert shared_experts[f"{prefix}down_proj.lora_B.weight"].shape == (8, 2)
def test_native_export_rejects_missing_shared_expert_adapter(monkeypatch):
"""A layer without its adapter must fail the export, not ship a partial adapter SGLang accepts."""
from megatron.core import parallel_state
monkeypatch.setattr(parallel_state, "get_tensor_model_parallel_world_size", lambda: 1)
monkeypatch.setattr(parallel_state, "get_expert_model_parallel_world_size", lambda: 1)
model = _model_with_adapters(include_shared_experts=False)
with pytest.raises(RuntimeError, match="adapter layout is incomplete"):
list(export_kimi_k3_lora_hf_chunks([model]))
def test_grouped_linear_uses_expert_token_boundaries():
"""Wrong boundaries apply expert i's adapter to expert j's tokens without any error."""
inputs = torch.tensor([[1.0, 2.0], [3.0, 4.0], [5.0, 6.0]])
weights = torch.tensor([[[1.0, 0.0]], [[0.0, 1.0]]])
output = _grouped_linear(inputs, weights, [1, 2])
torch.testing.assert_close(output, torch.tensor([[1.0], [4.0], [6.0]]))
def test_full_recompute_keeps_native_lora_in_autograd_graph():
"""Under full recompute with a frozen base the segment input carries no grad, so every LoRA
gradient is zero; the embedding hook must fix that in training and stay out of eval."""
model = nn.Module()
model.embedding = nn.Embedding.from_pretrained(torch.ones(8, 4), freeze=True)
model.adapter = nn.Parameter(torch.ones(4, 4))
model.config = SimpleNamespace(recompute_granularity="full")
model.pre_process = True
model.embedding.requires_grad_(False)
_enable_full_recompute_input_grads(model)
model.train()
hidden_states = model.embedding(torch.tensor([[1, 2]]))
output = checkpoint(lambda inputs: inputs @ model.adapter, hidden_states, use_reentrant=True)
output.sum().backward()
assert hidden_states.requires_grad
assert model.embedding.weight.grad is None
torch.testing.assert_close(model.adapter.grad, torch.full_like(model.adapter, 2.0))
model.eval()
assert not model.embedding(torch.tensor([[1, 2]])).requires_grad
+27
View File
@@ -0,0 +1,27 @@
from types import SimpleNamespace
import pytest
import torch
pytest.importorskip("mbridge")
from miles_plugins.mbridge.kimi_k3 import KimiK3Bridge # noqa: E402
@pytest.mark.parametrize("name", ("q_conv1d.weight", "k_conv1d.weight", "v_conv1d.weight", "A_log", "dt_bias"))
def test_kda_state_stays_fp32_against_the_bridge_dtype(name):
"""Losing the fp32 override silently downcasts the KDA state during conversion."""
bridge = object.__new__(KimiK3Bridge)
bridge.dtype = torch.bfloat16
bridge.config = SimpleNamespace(kimi_linear_num_heads=1)
weight = torch.tensor([0.1234567], dtype=torch.float32)
converted = bridge._weight_to_mcore_format(f"decoder.layers.0.self_attention.{name}", [weight])
assert converted.dtype == torch.float32
torch.testing.assert_close(converted, weight, rtol=0, atol=0)
bridge.make_vocab_size_divisible_by = None
assert (
bridge._weight_to_mcore_format("decoder.layers.0.self_attention.q_proj.weight", [weight]).dtype
== torch.bfloat16
)
@@ -0,0 +1,43 @@
import pytest
import torch
from miles_plugins.models.kimi_k3.pipeline import bank_num_rows, pack_stage_boundary, unpack_stage_boundary
def test_pack_unpack_is_exact_in_both_directions() -> None:
"""A transposed unflatten mixes bank rows into prefix_sum with no shape error."""
torch.manual_seed(0)
prefix_sum = torch.randn(5, 2, 16, dtype=torch.bfloat16)
block_residual = torch.randn(5, 2, 3, 16, dtype=torch.bfloat16)
packed = pack_stage_boundary(prefix_sum, block_residual)
assert packed.shape == (5, 2, 4 * 16)
prefix_out, bank_out = unpack_stage_boundary(packed, 16, 3)
torch.testing.assert_close(prefix_out, prefix_sum, rtol=0, atol=0)
torch.testing.assert_close(bank_out, block_residual, rtol=0, atol=0)
grad_prefix_sum = torch.randn(4, 1, 8, requires_grad=True)
grad_block_residual = torch.randn(4, 1, 2, 8, requires_grad=True)
prefix_out, bank_out = unpack_stage_boundary(pack_stage_boundary(grad_prefix_sum, grad_block_residual), 8, 2)
grad_prefix = torch.randn_like(prefix_out)
grad_bank = torch.randn_like(bank_out)
torch.autograd.backward([prefix_out, bank_out], [grad_prefix, grad_bank])
torch.testing.assert_close(grad_prefix_sum.grad, grad_prefix, rtol=0, atol=0)
torch.testing.assert_close(grad_block_residual.grad, grad_bank, rtol=0, atol=0)
with pytest.raises(AssertionError, match="stage-boundary payload width"):
unpack_stage_boundary(torch.zeros(2, 1, 3 * 16), 16, 3)
def test_bank_num_rows_matches_write_schedule() -> None:
"""One row off between the write schedule and the receiver reinterprets the payload."""
block_size = 12
for layer_idx in range(1, 93):
rows_written_before = sum(1 for w in range(layer_idx) if w % block_size == 0)
assert bank_num_rows(layer_idx, block_size) == rows_written_before
for last_layer_of_stage in range(92):
rows_after_exit = last_layer_of_stage // block_size + 1
assert bank_num_rows(last_layer_of_stage + 1, block_size) == rows_after_exit
@@ -0,0 +1,76 @@
from __future__ import annotations
import pytest
from scripts import run_kimi_k3
def _full(**kwargs):
return run_kimi_k3.ScriptArgs(model_name="Kimi-K3", hardware="GB300", num_nodes=16, num_gpus_per_node=4, **kwargs)
def _four_layer(**kwargs):
kwargs.setdefault("num_nodes", 1)
return run_kimi_k3.ScriptArgs(model_name="Kimi-K3-4layer", hardware="H200", **kwargs)
# (pipeline, context, expected tp, expected ep) for every layout run in production.
@pytest.mark.parametrize(
"pipeline_parallel_size,context_parallel_size,expected_tp,expected_ep",
[(1, 1, 32, 64), (8, 1, 8, 8), (8, 2, 4, 8)],
)
def test_full_model_parallel_derivation(pipeline_parallel_size, context_parallel_size, expected_tp, expected_ep):
args = _full(pipeline_parallel_size=pipeline_parallel_size, context_parallel_size=context_parallel_size)
assert args.tensor_parallel_size == expected_tp
assert args.expert_parallel_size == expected_ep
def test_derived_ep_saturates_the_bound_post_init_validates():
"""EP must fill the non-PP ranks of one stage; TP alone under-uses it whenever CP > 1."""
args = _full(pipeline_parallel_size=8, context_parallel_size=2)
model_parallel = args.tensor_parallel_size * args.context_parallel_size * args.pipeline_parallel_size
data_parallel = 64 // model_parallel
assert args.expert_parallel_size == args.tensor_parallel_size * args.context_parallel_size * data_parallel
def test_ep_override_wins_over_derivation():
args = _full(pipeline_parallel_size=8, context_parallel_size=2, ep_size_override=4)
assert args.expert_parallel_size == 4
def test_tp_override_wins_and_feeds_the_ep_derivation():
args = _full(pipeline_parallel_size=4, context_parallel_size=1, tp_size_override=8)
assert args.tensor_parallel_size == 8
# DP is 2 here, so the stage still holds 16 non-PP ranks.
assert args.expert_parallel_size == 16
def test_full_model_off_the_validated_gpu_count_needs_an_explicit_layout():
"""The 64-GPU derivation would silently produce a layout nobody has run."""
with pytest.raises(ValueError, match="tp-size-override"):
run_kimi_k3.ScriptArgs(model_name="Kimi-K3", hardware="H200", num_nodes=4, num_gpus_per_node=8)
@pytest.mark.parametrize("num_gpus_per_node,expected_tp", [(8, 8), (4, 4)], ids=["8-gpu-node", "4-gpu-node"])
def test_four_layer_layout_follows_the_gpu_count(num_gpus_per_node, expected_tp):
"""The rollout serves the experts replicated (EP1): Marlin is the only MXFP4 MoE runner with a LoRA path."""
args = _four_layer(num_gpus_per_node=num_gpus_per_node)
assert args.tensor_parallel_size == expected_tp
assert args.expert_parallel_size == expected_tp
assert args.rollout_tp_size == expected_tp
assert args.rollout_ep_size == 1
def test_four_layer_rollout_tp_can_span_nodes():
"""TP16 pads the Marlin MoE intermediate, the layout the padding fix is only decidable on."""
args = _four_layer(num_nodes=2, num_gpus_per_node=8, rollout_tp_size=16, rollout_ep_size=1)
assert args.tensor_parallel_size == 8
assert args.rollout_tp_size == 16
assert args.rollout_ep_size == 1
def test_checkpoint_paths_derive_from_the_model_name():
args = _four_layer(model_dir="/m")
assert args.hf_checkpoint == "/m/Kimi-K3-4layer"
assert args.bf16_checkpoint == "/m/Kimi-K3-4layer-bf16"
assert args.ref_load == "/m/Kimi-K3-4layer-bf16_torch_dist"
+6
View File
@@ -21,6 +21,12 @@ class TestAsyncRm:
"rm_type,response,label,expected",
[
("math", r"\boxed{42}", "42", 1),
(
"math",
r"<|open|>think<|sep|>work<|close|>think<|open|>response<|sep|>Answer: \boxed{42}<|close|>response",
"42",
1,
),
("math", r"\boxed{wrong}", "42", 0),
("f1", "hello world", "hello world", 1.0),
("dapo", "Answer: 42", "42", {"score": 1.0}),
@@ -59,6 +59,8 @@ _SCRIPTS_WHOSE_DEFAULTS_ARE_UNSUPPORTED: dict[str, Callable[[Path], dict[str, ob
"scripts/run_glm5_744b_a40b.py": lambda sandbox: _glm_checkpoint(sandbox, "GLM-5", 78),
"scripts/run_glm5_2_744b_a40b.py": lambda sandbox: _glm_checkpoint(sandbox, "GLM-5.2", 78),
"scripts/run_inkling.py": lambda sandbox: {"model_name": "Inkling-4layer"},
# convert_checkpoint locks a file under model_dir, which must be writable while recording
"scripts/run_kimi_k3.py": lambda sandbox: {"model_dir": str(sandbox / "models")},
"scripts/run_nemotron_3_nano_4b_fsdp.py": _nemotron_checkpoint,
"scripts/run_nemotron_3_ultra_550b_a55b.py": lambda sandbox: {
"model_name": "NVIDIA-Nemotron-3-Ultra-550B-A55B-BF16-4layer"
@@ -0,0 +1,5 @@
### 0
python <REPO_ROOT>/tools/convert_mxfp4_to_bf16.py
--model-dir <SANDBOX>/models/Kimi-K3-4layer-64experts
--save-dir <SANDBOX>/models/Kimi-K3-4layer-64experts-bf16
--device cuda
@@ -0,0 +1,4 @@
### 0
hf download
--repo-type dataset zhuzilin/gsm8k
--local-dir /root/datasets/gsm8k
@@ -0,0 +1,11 @@
### 0
mkdir -p <SANDBOX>/models /root/datasets
### 1
hf download Pinaster/Kimi-K3-4layer-64experts
--local-dir <SANDBOX>/models/Kimi-K3-4layer-64experts
### 2
hf download
--repo-type dataset zhuzilin/gsm8k
--local-dir /root/datasets/gsm8k
@@ -0,0 +1,56 @@
### 0
PYTHONPATH=<REPO_ROOT>:/root/Megatron-LM:/frozen/pythonpath torchrun
--nproc-per-node 8 <REPO_ROOT>/tools/convert_hf_to_torch_dist.py
--spec miles_plugins.models.kimi_k3 get_kimi_k3_spec
--disable-bias-linear
--num-layers 4
--hidden-size 7168
--ffn-hidden-size 33792
--num-attention-heads 96
--num-query-groups 96
--kv-channels 256
--normalization RMSNorm
--position-embedding-type none
--rotary-base 10000
--norm-epsilon 1e-5
--hidden-dropout 0
--attention-dropout 0
--disable-bf16-reduced-precision-matmul
--swiglu
--untie-embeddings-and-output-weights
--vocab-size 163840
--make-vocab-size-divisible-by 128
--multi-latent-attention
--q-lora-rank 1536
--kv-lora-rank 512
--qk-head-dim 128
--qk-pos-emb-head-dim 64
--v-head-dim 128
--qk-layernorm
--attention-softmax-in-fp32
--num-experts 64
--moe-layer-freq '[0,1,1,1]'
--moe-ffn-hidden-size 3072
--moe-latent-size 3584
--moe-shared-expert-intermediate-size 6144
--moe-router-topk 16
--moe-router-score-function sigmoid
--moe-router-pre-softmax
--moe-router-enable-expert-bias
--moe-router-load-balancing-type none
--moe-token-dispatcher-type alltoall
--moe-router-bias-update-rate 0
--moe-aux-loss-coeff 0
--moe-router-topk-scaling-factor 1.0
--moe-router-dtype fp32
--moe-grouped-gemm
--moe-permute-fusion
--hf-checkpoint <SANDBOX>/models/Kimi-K3-4layer-64experts-bf16
--save <SANDBOX>/models/Kimi-K3-4layer-64experts-bf16_torch_dist
--bf16
--tensor-model-parallel-size 1
--pipeline-model-parallel-size 1
--context-parallel-size 1
--expert-model-parallel-size 8
--expert-tensor-parallel-size 1
--megatron-to-hf-mode raw
@@ -0,0 +1,148 @@
### 0
pkill -9 sglang; sleep 3; ray stop
--force; pkill -9 ray; pkill -9 miles; sleep 3; pkill -9 ray; pkill -9 miles; pkill -9 redis; true;
### 1
export PYTHONUNBUFFERED=1 && ray start
--head
--node-ip-address 127.0.0.1
--num-gpus 8
--disable-usage-stats
### 2
nvidia-smi topo -m 2>/dev/null | grep -o 'NV[0-9][0-9]*' | wc -l
### 3
export no_proxy=127.0.0.1 && export PYTHONUNBUFFERED=1 && ray job submit
--address="http://127.0.0.1:8265"
--runtime-env-json='{"env_vars": {"PYTHONUNBUFFERED": "1", "CUDA_DEVICE_MAX_CONNECTIONS": "1", "NCCL_NVLS_ENABLE": "0", "no_proxy": "127.0.0.1,127.0.0.1", "MASTER_ADDR": "127.0.0.1", "NCCL_TIMEOUT": "3600", "PYTHONPATH": "<REPO_ROOT>:/root/Megatron-LM:/root/sglang/python:/frozen/pythonpath", "SGLANG_JIT_ROUTE_RADIX": "1", "SGLANG_NUMA_BIND_V2": "0", "SGLANG_K3_AR_FUSION": "0"}}'
-- python3 <REPO_ROOT>/train.py
--spec miles_plugins.models.kimi_k3 get_kimi_k3_spec
--disable-bias-linear
--num-layers 4
--hidden-size 7168
--ffn-hidden-size 33792
--num-attention-heads 96
--num-query-groups 96
--kv-channels 256
--normalization RMSNorm
--position-embedding-type none
--rotary-base 10000
--norm-epsilon 1e-5
--hidden-dropout 0
--attention-dropout 0
--disable-bf16-reduced-precision-matmul
--swiglu
--untie-embeddings-and-output-weights
--vocab-size 163840
--make-vocab-size-divisible-by 128
--multi-latent-attention
--q-lora-rank 1536
--kv-lora-rank 512
--qk-head-dim 128
--qk-pos-emb-head-dim 64
--v-head-dim 128
--qk-layernorm
--attention-softmax-in-fp32
--num-experts 64
--moe-layer-freq '[0,1,1,1]'
--moe-ffn-hidden-size 3072
--moe-latent-size 3584
--moe-shared-expert-intermediate-size 6144
--moe-router-topk 16
--moe-router-score-function sigmoid
--moe-router-pre-softmax
--moe-router-enable-expert-bias
--moe-router-load-balancing-type none
--moe-token-dispatcher-type alltoall
--moe-router-bias-update-rate 0
--moe-aux-loss-coeff 0
--moe-router-topk-scaling-factor 1.0
--moe-router-dtype fp32
--moe-grouped-gemm
--moe-permute-fusion
--hf-checkpoint <SANDBOX>/models/Kimi-K3-4layer-64experts
--ref-load <SANDBOX>/models/Kimi-K3-4layer-64experts-bf16_torch_dist
--megatron-to-hf-mode raw
--model-name kimi_k3
--lora-rank 16
--lora-alpha 32
--lora-dropout 0.0
--target-modules "decoder.layers.*.self_attention.o_proj,decoder.layers.*.self_attention.q_a_proj,decoder.layers.*.self_attention.kv_a_proj_with_mqa,decoder.layers.*.mlp.linear_fc1,decoder.layers.*.mlp.linear_fc2,decoder.layers.*.mlp.experts.linear_fc1,decoder.layers.*.mlp.experts.linear_fc2"
--no-gradient-accumulation-fusion
--lora-base-cpu-backup
--experts-shared-outer-loras
--prompt-data /root/datasets/gsm8k/train.parquet
--input-key messages
--label-key label
--apply-chat-template
--rollout-shuffle
--balance-data
--rm-type math
--num-rollout 2
--rollout-batch-size 8
--n-samples-per-prompt 2
--rollout-max-response-len 256
--rollout-temperature 1
--global-batch-size 16
--use-dynamic-global-batch-size
--optimizer sgd
--sgd-momentum 0
--lr 1e-6
--lr-decay-style constant
--weight-decay 0
--use-distributed-optimizer
--advantage-estimator grpo
--kl-loss-coef 0.0
--kl-loss-type low_var_kl
--entropy-coef 0.0
--eps-clip 0.2
--eps-clip-high 0.28
--use-wandb
--wandb-project miles-run_kimi_k3
--wandb-group 260101-000000-000
--wandb-key 'frozen-wandb-api-key'
--disable-wandb-random-suffix
--tensor-model-parallel-size 8
--sequence-parallel
--pipeline-model-parallel-size 1
--context-parallel-size 1
--expert-model-parallel-size 8
--expert-tensor-parallel-size 1
--recompute-granularity full
--recompute-method uniform
--recompute-num-layers 1
--use-dynamic-batch-size
--max-tokens-per-gpu 512
--log-probs-chunk-size 512
--distributed-timeout-minutes 10
--rollout-num-gpus-per-engine 8
--sglang-tp-size 8
--sglang-ep-size 1
--sglang-server-concurrency 16
--sglang-max-running-requests 16
--sglang-max-mamba-cache-size 16
--use-miles-router
--sglang-lora-backend triton
--sglang-lora-strict-loading
--sglang-max-lora-rank 16
--sglang-moe-runner-backend marlin
--sglang-cuda-graph-bs-decode 1 2 4 8 16
--sglang-mem-fraction-static 0.7
--sglang-disable-shared-experts-fusion
--sglang-context-length 8192
--attention-dropout 0.0
--hidden-dropout 0.0
--accumulate-allreduce-grads-in-fp32
--colocate
--offload-train
--update-weight-buffer-size 2147483648
--train-memory-margin-bytes 2147483648
--actor-num-nodes 1
--actor-num-gpus-per-node 8
--num-gpus-per-node 8
--no-check-for-nan-in-loss-and-grad
--check-weight-update-skip-list vision_tower. mm_projector.
--check-weight-update-allow-quant-error
--save /personal/checkpoints/260101-000000-000
--save-interval 50
@@ -0,0 +1,78 @@
"--spec"
"miles_plugins.models.kimi_k3"
"get_kimi_k3_spec"
"--disable-bias-linear"
"--num-layers"
"4"
"--hidden-size"
"7168"
"--ffn-hidden-size"
"33792"
"--num-attention-heads"
"96"
"--num-query-groups"
"96"
"--kv-channels"
"256"
"--normalization"
"RMSNorm"
"--position-embedding-type"
"none"
"--rotary-base"
"10000"
"--norm-epsilon"
"1e-5"
"--hidden-dropout"
"0"
"--attention-dropout"
"0"
"--disable-bf16-reduced-precision-matmul"
"--swiglu"
"--untie-embeddings-and-output-weights"
"--vocab-size"
"163840"
"--make-vocab-size-divisible-by"
"128"
"--multi-latent-attention"
"--q-lora-rank"
"1536"
"--kv-lora-rank"
"512"
"--qk-head-dim"
"128"
"--qk-pos-emb-head-dim"
"64"
"--v-head-dim"
"128"
"--qk-layernorm"
"--attention-softmax-in-fp32"
"--num-experts"
"64"
"--moe-layer-freq"
"[0,1,1,1]"
"--moe-ffn-hidden-size"
"3072"
"--moe-latent-size"
"3584"
"--moe-shared-expert-intermediate-size"
"6144"
"--moe-router-topk"
"16"
"--moe-router-score-function"
"sigmoid"
"--moe-router-pre-softmax"
"--moe-router-enable-expert-bias"
"--moe-router-load-balancing-type"
"none"
"--moe-token-dispatcher-type"
"alltoall"
"--moe-router-bias-update-rate"
"0"
"--moe-aux-loss-coeff"
"0"
"--moe-router-topk-scaling-factor"
"1.0"
"--moe-router-dtype"
"fp32"
"--moe-grouped-gemm"
"--moe-permute-fusion"
@@ -0,0 +1,78 @@
"--spec"
"miles_plugins.models.kimi_k3"
"get_kimi_k3_spec"
"--disable-bias-linear"
"--num-layers"
"4"
"--hidden-size"
"7168"
"--ffn-hidden-size"
"33792"
"--num-attention-heads"
"96"
"--num-query-groups"
"96"
"--kv-channels"
"256"
"--normalization"
"RMSNorm"
"--position-embedding-type"
"none"
"--rotary-base"
"10000"
"--norm-epsilon"
"1e-5"
"--hidden-dropout"
"0"
"--attention-dropout"
"0"
"--disable-bf16-reduced-precision-matmul"
"--swiglu"
"--untie-embeddings-and-output-weights"
"--vocab-size"
"163840"
"--make-vocab-size-divisible-by"
"128"
"--multi-latent-attention"
"--q-lora-rank"
"1536"
"--kv-lora-rank"
"512"
"--qk-head-dim"
"128"
"--qk-pos-emb-head-dim"
"64"
"--v-head-dim"
"128"
"--qk-layernorm"
"--attention-softmax-in-fp32"
"--num-experts"
"896"
"--moe-layer-freq"
"[0,1,1,1]"
"--moe-ffn-hidden-size"
"3072"
"--moe-latent-size"
"3584"
"--moe-shared-expert-intermediate-size"
"6144"
"--moe-router-topk"
"16"
"--moe-router-score-function"
"sigmoid"
"--moe-router-pre-softmax"
"--moe-router-enable-expert-bias"
"--moe-router-load-balancing-type"
"none"
"--moe-token-dispatcher-type"
"alltoall"
"--moe-router-bias-update-rate"
"0"
"--moe-aux-loss-coeff"
"0"
"--moe-router-topk-scaling-factor"
"1.0"
"--moe-router-dtype"
"fp32"
"--moe-grouped-gemm"
"--moe-permute-fusion"
+78
View File
@@ -0,0 +1,78 @@
"--spec"
"miles_plugins.models.kimi_k3"
"get_kimi_k3_spec"
"--disable-bias-linear"
"--num-layers"
"93"
"--hidden-size"
"7168"
"--ffn-hidden-size"
"33792"
"--num-attention-heads"
"96"
"--num-query-groups"
"96"
"--kv-channels"
"256"
"--normalization"
"RMSNorm"
"--position-embedding-type"
"none"
"--rotary-base"
"10000"
"--norm-epsilon"
"1e-5"
"--hidden-dropout"
"0"
"--attention-dropout"
"0"
"--disable-bf16-reduced-precision-matmul"
"--swiglu"
"--untie-embeddings-and-output-weights"
"--vocab-size"
"163840"
"--make-vocab-size-divisible-by"
"128"
"--multi-latent-attention"
"--q-lora-rank"
"1536"
"--kv-lora-rank"
"512"
"--qk-head-dim"
"128"
"--qk-pos-emb-head-dim"
"64"
"--v-head-dim"
"128"
"--qk-layernorm"
"--attention-softmax-in-fp32"
"--num-experts"
"896"
"--moe-layer-freq"
"[0,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1]"
"--moe-ffn-hidden-size"
"3072"
"--moe-latent-size"
"3584"
"--moe-shared-expert-intermediate-size"
"6144"
"--moe-router-topk"
"16"
"--moe-router-score-function"
"sigmoid"
"--moe-router-pre-softmax"
"--moe-router-enable-expert-bias"
"--moe-router-load-balancing-type"
"none"
"--moe-token-dispatcher-type"
"alltoall"
"--moe-router-bias-update-rate"
"0"
"--moe-aux-loss-coeff"
"0"
"--moe-router-topk-scaling-factor"
"1.0"
"--moe-router-dtype"
"fp32"
"--moe-grouped-gemm"
"--moe-permute-fusion"
+12 -1
View File
@@ -12,6 +12,7 @@ from megatron.training.training import get_model
import miles_plugins.mbridge # noqa: F401
from mbridge import AutoBridge
from miles.backends.megatron_utils.arguments import set_default_megatron_args
from miles.backends.megatron_utils.fp32_param_utils import enforce_marked_param_dtypes
from miles.backends.megatron_utils.initialize import init
from miles.backends.megatron_utils.model_provider import get_model_provider_func
from miles.utils.logging_utils import configure_logger_raw
@@ -69,7 +70,16 @@ def get_args():
def ceildiv(a, b):
return -(a // -b)
if args.pipeline_model_parallel_size == 1 and world_size > 1 and not os.environ.get("CONVERT_KEEP_PP1"):
auto_pipeline_parallel = (
args.pipeline_model_parallel_size == 1
and args.tensor_model_parallel_size == 1
and args.context_parallel_size == 1
and args.expert_model_parallel_size == 1
and args.expert_tensor_parallel_size == 1
and world_size > 1
and not os.environ.get("CONVERT_KEEP_PP1")
)
if auto_pipeline_parallel:
pp_size = world_size
while True:
args.pipeline_model_parallel_size = pp_size
@@ -117,6 +127,7 @@ def main():
args = get_args()
init(args)
model = get_model(get_model_provider_func(args), ModelType.encoder_or_decoder, wrap_with_ddp=False)
enforce_marked_param_dtypes(model)
# Load model
hf_model_path = args.hf_checkpoint
+170
View File
@@ -0,0 +1,170 @@
"""Convert an ``mxfp4-pack-quantized`` compressed-tensors checkpoint to BF16.
Every ``<name>.weight_packed`` / ``<name>.weight_scale`` pair becomes a dequantized
``<name>.weight`` and ``quantization_config`` is dropped from the copied config; shard the
file list over processes with ``--shard-rank/--num-shards`` and run ``--finalize-only`` once.
"""
import argparse
import json
import math
import os
import shutil
from pathlib import Path
from safetensors.torch import safe_open, save_file
from miles.utils.mxfp4 import dequantize_mxfp4
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description="Convert MXFP4-packed compressed-tensors weights to BF16.")
parser.add_argument("--model-dir", type=Path, required=True)
parser.add_argument("--save-dir", type=Path, required=True)
parser.add_argument("--device", choices=("cpu", "cuda"), default="cpu")
parser.add_argument("--files", nargs="+")
parser.add_argument("--shard-rank", type=int)
parser.add_argument("--num-shards", type=int)
parser.add_argument("--finalize-only", action="store_true")
parser.add_argument("--overwrite", action="store_true")
return parser.parse_args()
def quantization_section(config: dict) -> dict:
"""Multimodal checkpoints nest ``quantization_config`` under ``text_config``."""
return config.get("text_config", config)
def load_group_size(model_dir: Path) -> int:
with (model_dir / "config.json").open() as file:
config = json.load(file)
quantization_config = quantization_section(config)["quantization_config"]
assert quantization_config["format"] == "mxfp4-pack-quantized", quantization_config["format"]
weights_config = quantization_config["config_groups"]["group_0"]["weights"]
assert weights_config["type"] == "float"
assert weights_config["num_bits"] == 4
assert weights_config["scale_dtype"] == "torch.uint8"
return int(weights_config["group_size"])
def copy_metadata(model_dir: Path, output_dir: Path) -> None:
output_dir.mkdir(parents=True, exist_ok=True)
for source in model_dir.iterdir():
if not source.is_file() or source.suffix == ".safetensors" or source.name.endswith(".index.json"):
continue
shutil.copy2(source, output_dir / source.name)
config_path = output_dir / "config.json"
with config_path.open() as file:
config = json.load(file)
del quantization_section(config)["quantization_config"]
with config_path.open("w") as file:
json.dump(config, file, indent=2)
def convert_file(
source_path: Path,
output_path: Path,
group_size: int,
device: str,
) -> None:
tensors = {}
with safe_open(source_path, framework="pt", device="cpu") as reader:
keys = set(reader.keys())
packed_keys = sorted(name for name in keys if name.endswith(".weight_packed"))
scale_keys = {name for name in keys if name.endswith(".weight_scale")}
for packed_name in packed_keys:
prefix = packed_name.removesuffix(".weight_packed")
scale_name = f"{prefix}.weight_scale"
assert scale_name in scale_keys, f"Missing scale for {packed_name}"
scale_keys.remove(scale_name)
weight_packed = reader.get_tensor(packed_name).to(device)
weight_scale = reader.get_tensor(scale_name).to(device)
weight = dequantize_mxfp4(weight_packed, weight_scale, group_size).cpu()
output_name = f"{prefix}.weight"
tensors[output_name] = weight
assert not scale_keys, f"Orphan MXFP4 scales in {source_path.name}: {sorted(scale_keys)}"
quantized_keys = set(packed_keys) | {
f"{name.removesuffix('.weight_packed')}.weight_scale" for name in packed_keys
}
for name in sorted(keys - quantized_keys):
tensor = reader.get_tensor(name)
tensors[name] = tensor
output_path.parent.mkdir(parents=True, exist_ok=True)
temporary_path = output_path.with_suffix(f"{output_path.suffix}.tmp")
save_file(tensors, temporary_path)
os.replace(temporary_path, output_path)
def build_index(output_dir: Path) -> None:
dtype_sizes = {
"BF16": 2,
"F16": 2,
"F32": 4,
"F64": 8,
"I32": 4,
"I64": 8,
"U8": 1,
}
weight_map = {}
total_size = 0
for path in sorted(output_dir.glob("*.safetensors")):
with safe_open(path, framework="pt", device="cpu") as reader:
for name in reader.keys():
assert name not in weight_map, f"Duplicate tensor: {name}"
weight_map[name] = path.name
tensor_slice = reader.get_slice(name)
shape = tensor_slice.get_shape()
dtype = tensor_slice.get_dtype()
assert dtype in dtype_sizes, f"Unsupported safetensors dtype: {dtype}"
total_size += math.prod(shape) * dtype_sizes[dtype]
with (output_dir / "model.safetensors.index.json").open("w") as file:
json.dump({"metadata": {"total_size": total_size}, "weight_map": weight_map}, file, indent=2)
def main() -> None:
args = parse_args()
assert args.model_dir.is_dir(), args.model_dir
assert args.save_dir != args.model_dir
assert (args.shard_rank is None) == (args.num_shards is None)
if args.num_shards is not None:
assert args.num_shards > 0
assert 0 <= args.shard_rank < args.num_shards
assert args.files is None, "--files cannot be combined with sharded conversion"
group_size = load_group_size(args.model_dir)
args.save_dir.mkdir(parents=True, exist_ok=True)
if args.finalize_only:
copy_metadata(args.model_dir, args.save_dir)
build_index(args.save_dir)
return
if args.shard_rank in (None, 0):
copy_metadata(args.model_dir, args.save_dir)
filenames = args.files or [path.name for path in sorted(args.model_dir.glob("*.safetensors"))]
if args.num_shards is not None:
filenames = filenames[args.shard_rank :: args.num_shards]
assert filenames, "No safetensors files found"
for filename in filenames:
source_path = args.model_dir / filename
output_path = args.save_dir / filename
assert source_path.is_file(), source_path
if output_path.exists() and not args.overwrite:
print(f"Skipping existing {output_path}", flush=True)
continue
print(f"Converting {source_path} -> {output_path}", flush=True)
convert_file(source_path, output_path, group_size, args.device)
if args.num_shards is None:
build_index(args.save_dir)
if __name__ == "__main__":
main()