Download apply_patch.py from Radioheading/Kimi-K3-W4AFP8: direct link, hf CLI and curl.
- Browser
- Download file 9.14 kB
-
https://huggingface.co/Radioheading/Kimi-K3-W4AFP8/resolve/main/apply_patch.py
- Command line
-
hf download hf://Radioheading/Kimi-K3-W4AFP8/apply_patch.py
-
curl -L -o apply_patch.py https://huggingface.co/Radioheading/Kimi-K3-W4AFP8/resolve/main/apply_patch.py
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()) | |