mirror of
https://github.com/radixark/miles.git
synced 2026-10-02 07:14:53 +08:00
Kimi-K3 full-model and LoRA RL support (#1825)
Co-authored-by: Zhichenzzz <zczeng@uw.edu>
This commit is contained in:
co-authored by
Zhichenzzz
parent
b3c8f9c8f2
commit
8cdf3794d9
+45
-31
@@ -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/`.
|
||||
|
||||
@@ -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)
|
||||
|
||||
+25
-27
@@ -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
|
||||
|
||||
@@ -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]:
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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()
|
||||
@@ -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",
|
||||
]
|
||||
|
||||
@@ -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)
|
||||
@@ -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"]
|
||||
@@ -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
|
||||
@@ -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
|
||||
]
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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 "
|
||||
)
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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}")
|
||||
@@ -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
|
||||
@@ -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"
|
||||
@@ -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"
|
||||
@@ -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,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
|
||||
|
||||
@@ -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()
|
||||
Reference in New Issue
Block a user