File size: 9,136 Bytes
66a097f | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 | #!/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())
|