Kimi-K3-W4AFP8 / apply_patch.py
Radioheading's picture
Add files using upload-large-folder tool
66a097f verified
Raw History Blame Contribute Delete
9.14 kB
#!/usr/bin/env python3
"""Teach SGLang how to load this checkpoint. Run once, then serve.
python3 apply_patch.py # apply
python3 apply_patch.py --check # report status only
python3 apply_patch.py --revert # undo
What it does — two files, both additive:
1. writes sglang/srt/layers/quantization/kimi_k3_attnfp8.py (new file)
2. adds one import + one dict entry to
sglang/srt/layers/quantization/__init__.py
It does not modify `W4AFp8Config`, so other w4afp8 checkpoints are unaffected.
─────────────────────────────────────────────────────────────────────────────
Why a patch is needed at all
This checkpoint mixes three schemes:
MoE experts INT4 group-128 + FP8 activations
4 attention projs FP8 E4M3, per-tensor
every other linear bf16
Stock `W4AFp8Config` has one behaviour for linears — `Fp8LinearMethod(self)`
with `weight_block_size = [128, 128]` — and it cannot be steered from
config.json, because `from_config()` hardcodes the block size and never passes
`ignored_layers`. `b_proj` (output size 6) then fails:
ValueError: Weight output_partition_size = 6 is not divisible by block_n = 128
Upstream tracking: sgl-project/sglang#16643, #22806, #30598.
─────────────────────────────────────────────────────────────────────────────
!! Five attention projections must stay bf16
`kimi_k3.py` reads `.weight` directly on `kv_b_proj`, `q_b_proj`, `f_b_proj`,
`f_a_proj` and `b_proj`, bypassing `quant_method`. The FP8 path stores weights
transposed, so a direct reader gets a flipped layout and crashes:
RuntimeError: mat1 and mat2 shapes cannot be multiplied (384x128 and 768x128)
They are bf16 in the checkpoint and this patch keeps them unquantized.
"""
from __future__ import annotations
import argparse
import pathlib
import sys
MARKER = "_KIMI_K3_ATTNFP8"
MODULE_NAME = "kimi_k3_attnfp8.py"
IMPORT_LINE = (
f"from sglang.srt.layers.quantization.kimi_k3_attnfp8 import ( # {MARKER}\n"
" KimiK3AttnFp8Config,\n"
")\n"
)
REGISTRY_ANCHOR = 'BASE_QUANTIZATION_METHODS: Dict[str, Type[QuantizationConfig]] = {\n'
REGISTRY_LINE = f' "w4afp8_attnfp8": KimiK3AttnFp8Config, # {MARKER}\n'
MODULE_SRC = '''"""Quantization config for Kimi-K3-W4AFP8 with FP8 attention projections.
Installed by the checkpoint's apply_patch.py. See that file for why this exists.
"""
from __future__ import annotations
import logging
import torch
from sglang.srt.layers.linear import LinearBase
from sglang.srt.layers.quantization.fp8 import Fp8Config, Fp8LinearMethod
from sglang.srt.layers.quantization.unquant import UnquantizedLinearMethod
from sglang.srt.layers.quantization.w4afp8 import W4AFp8Config
logger = logging.getLogger(__name__)
# Attention projections stored as FP8 in the checkpoint (runtime module names).
FP8_LINEARS = {
"fused_qkvg_proj", # KDA layers, fused q/k/v/g (attn_tp == tp)
"qkv_proj", # same layers when dp-attention splits the fusion
"g_proj", # MLA gate, plus KDA g when unfused
"o_proj", # attention output projection
}
# Kept bf16 — these have raw `.weight` readers in kimi_k3.py. Do not add them.
EXCLUDED = ("kv_b_proj", "q_b_proj", "f_b_proj", "f_a_proj", "b_proj")
def _patch_kda_precompile_dtype() -> None:
"""kimi_k3.py uses o_proj.weight.dtype as a stand-in for the activation
dtype when precompiling the KDA kernel. With o_proj in FP8 that stand-in
lies and the wrong kernel is compiled — silently wrong, not a crash.
"""
try:
from sglang.kernels.ops.attention.fla import kda
from sglang.srt.models import kimi_k3 as model_mod
except Exception as exc: # noqa: BLE001
logger.warning("kimi_k3_attnfp8: KDA dtype patch skipped (%s)", exc)
return
orig = kda.precompile_k3_recompute_w_u_kernel
if getattr(orig, "_kimi_k3_dtype_fixed", False):
return
def wrapped(*, num_heads, dtype, device):
if dtype in (torch.float8_e4m3fn, torch.float8_e5m2):
dtype = torch.bfloat16
return orig(num_heads=num_heads, dtype=dtype, device=device)
wrapped._kimi_k3_dtype_fixed = True
kda.precompile_k3_recompute_w_u_kernel = wrapped
if hasattr(model_mod, "precompile_k3_recompute_w_u_kernel"):
model_mod.precompile_k3_recompute_w_u_kernel = wrapped
logger.info("kimi_k3_attnfp8: KDA precompile dtype patch applied")
class KimiK3AttnFp8Config(W4AFp8Config):
"""MoE is delegated to W4AFp8; linear routing is decided here."""
@classmethod
def get_name(cls) -> str:
return "w4afp8_attnfp8"
@classmethod
def from_config(cls, config):
self = cls(
is_checkpoint_fp8_serialized=True,
is_checkpoint_w4afp8_serialized=True,
linear_activation_scheme="dynamic",
moe_activation_scheme="static",
group_size=int(config.get("group_size", 128)),
)
# weight_block_size=None matters: the parent's [128, 128] block quant
# kills b_proj, and the checkpoint's attention scales are per-tensor.
self._attn_cfg = Fp8Config(
is_checkpoint_fp8_serialized=True,
activation_scheme="dynamic",
weight_block_size=None,
)
_patch_kda_precompile_dtype()
logger.info("kimi_k3_attnfp8: config ready (FP8 linears: %s)",
", ".join(sorted(FP8_LINEARS)))
return self
def get_quant_method(self, layer, prefix: str):
if isinstance(layer, LinearBase):
name = prefix.rsplit(".", 1)[-1] if prefix else ""
if name in FP8_LINEARS:
return Fp8LinearMethod(self._attn_cfg)
# Everything else stays bf16. Sending these to Fp8LinearMethod(self)
# — what stock does — routes them into the [128,128] block-quant
# path and kills b_proj (output size 6).
return UnquantizedLinearMethod()
return super().get_quant_method(layer, prefix)
'''
def sglang_quant_dir() -> pathlib.Path:
try:
import sglang
except ImportError:
sys.exit("sglang is not importable. Install SGLang first, then re-run.")
d = pathlib.Path(sglang.__file__).resolve().parent / "srt" / "layers" / "quantization"
if not (d / "__init__.py").exists():
sys.exit(f"unexpected SGLang layout: {d} has no __init__.py")
return d
def status(d: pathlib.Path) -> tuple[bool, bool]:
return (d / MODULE_NAME).exists(), MARKER in (d / "__init__.py").read_text()
def main() -> int:
ap = argparse.ArgumentParser(description=__doc__,
formatter_class=argparse.RawDescriptionHelpFormatter)
ap.add_argument("--check", action="store_true")
ap.add_argument("--revert", action="store_true")
a = ap.parse_args()
d = sglang_quant_dir()
init = d / "__init__.py"
has_mod, has_reg = status(d)
print(f" sglang quantization dir : {d}")
print(f" module installed : {has_mod}")
print(f" registry entry : {has_reg}")
if a.check:
ok = has_mod and has_reg
print(f" → {'ready' if ok else 'NOT applied'}")
return 0 if ok else 1
if a.revert:
(d / MODULE_NAME).unlink(missing_ok=True)
src = init.read_text()
kept = [l for l in src.splitlines(keepends=True) if MARKER not in l]
# the import spans 3 lines; drop its continuation lines too
text = "".join(kept).replace(" KimiK3AttnFp8Config,\n)\n", "", 1)
init.write_text(text)
print(" → reverted")
return 0
if has_mod and has_reg:
print(" → already applied, nothing to do")
return 0
(d / MODULE_NAME).write_text(MODULE_SRC)
src = init.read_text()
if MARKER not in src:
if REGISTRY_ANCHOR not in src:
sys.exit(
"could not find BASE_QUANTIZATION_METHODS in "
f"{init}.\nThis SGLang version is not supported by this patch; "
"the checkpoint README lists the verified version."
)
src = src.replace(REGISTRY_ANCHOR, REGISTRY_ANCHOR + REGISTRY_LINE, 1)
# import must come after the other quantization imports; append near top
# of the registry block's preceding import section is fragile, so put it
# immediately before the registry definition.
src = src.replace(REGISTRY_ANCHOR, IMPORT_LINE + "\n" + REGISTRY_ANCHOR, 1)
init.write_text(src)
import py_compile
py_compile.compile(str(init), doraise=True)
py_compile.compile(str(d / MODULE_NAME), doraise=True)
print(" → applied and compiled OK")
print(" serve with: --trust-remote-code (config declares w4afp8_attnfp8)")
return 0
if __name__ == "__main__":
raise SystemExit(main())