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())