#!/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())