Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
36 changes: 36 additions & 0 deletions python/sglang/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,42 @@

_redirect_third_party_caches()

# Kimi-K3 may opt into an SGLang-owned AITER tuning profile. Configure it
# before any downstream import can initialize AITER_CONFIGS.
import importlib.util as _importlib_util
import os as _os
from pathlib import Path as _Path

import torch as _torch

if (
_torch.version.hip is not None
and _os.environ.get("SGLANG_K3_AITER_M16384_PROFILE", "0").lower() in ("1", "true")
and "AITER_CONFIG_GEMM_BF16" not in _os.environ
):
_aiter_spec = _importlib_util.find_spec("aiter")
if _aiter_spec is not None and _aiter_spec.origin is not None:
_aiter_root = _Path(_aiter_spec.origin).resolve().parent
_base = _aiter_root / "configs" / "bf16_tuned_gemm.csv"
_model_configs = sorted(
(_aiter_root / "configs" / "model_configs").glob("*bf16_tuned_gemm*.csv")
)
_profile = (
_Path(__file__).resolve().parent
/ "kernels"
/ "ops"
/ "kimi_k3"
/ "configs"
/ "kimik3_m16384_profile.csv"
)
_paths = [_base, *_model_configs, _profile]
if all(_path.is_file() for _path in _paths):
_os.environ["AITER_CONFIG_GEMM_BF16"] = _os.pathsep.join(map(str, _paths))
del _importlib_util
del _os
del _Path
del _torch

if _sys.platform == "darwin" and _platform.machine() == "arm64":
from sglang._platform_stubs import install_platform_stubs as _install_platform_stubs

Expand Down
21 changes: 21 additions & 0 deletions python/sglang/kernels/ops/kimi_k3/aiter_fusion.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,21 @@
"""Shared helpers for fail-closed Kimi-K3 AITER batch-1 fusions."""

from __future__ import annotations

import torch

_FP8_MAX = 448.0


def quantize_fp8_rows(weight: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
"""Quantize a BF16/FP32 [out, in] weight with one FP32 scale per row."""
if weight.ndim != 2 or not weight.is_cuda or not weight.is_contiguous():
raise ValueError("expected a contiguous CUDA [out, in] weight")
row_amax = weight.float().abs().amax(dim=1)
scale = (row_amax / _FP8_MAX).clamp_min(torch.finfo(torch.float32).tiny)
quantized = (
(weight.float() / scale[:, None])
.clamp(-_FP8_MAX, _FP8_MAX)
.to(torch.float8_e4m3fn)
)
return quantized.contiguous(), scale.contiguous()
Original file line number Diff line number Diff line change
@@ -0,0 +1,9 @@
gfx,cu_num,M,N,K,bias,dtype,outdtype,scaleAB,bpreshuffle,libtype,solidx,splitK,us,kernelName,err_ratio,tflops,bw
gfx950,256,16384,2304,1536,False,torch.bfloat16,torch.bfloat16,False,False,hipblaslt,438972,0,116.1558,Cijk_Alik_Bljk_BBS_BH_Bias_HA_S_SAV_UserArgs_MT256x224x64_MI16x16x1_SN_LDSB0_AFC0_AFEM1_AFEM1_ASEM1_CLR0_CADS0_DTLA1_DTLB1_DTVA0_DTVB0_EPS0_FDSI0_GRPM1_GRVWA8_GRVWB8_GSU0_GSUAMB_GLS0_ISA950_IU1_K1_LDSTI0_LBSPPA1024_LBSPPB1024_LBSPPM0_LPA16_LPB16_LPM0_LRVW8_LWPMn1_MIAV0_MIWT8_7_MO40_NTn1_NTA2_NTB3_NTC7_NTD5_NTM0_NEPBS14_NLCA1_NLCB1_ONLL1_PGR2_PLR1_PKA1_SIA3_SS1_SPO0_SRVW0_SSO1_SVW8_SK3_SKFTR0_SKXCCM8_TLDS1_ULSGRO0_USL1_UIOFGRO0_USFGRO0_VSn1_VWA8_VWB1_WSGRA0_WSGRB0_WS64_WG32_8_1,0.0,998.35,1144.21
gfx950,256,16384,3072,512,False,torch.bfloat16,torch.bfloat16,False,False,hipblaslt,439111,0,58.0215,Cijk_Alik_Bljk_BBS_BH_Bias_HA_S_SAV_UserArgs_MT256x256x64_MI16x16x1_CMS_SN_LDSB0_AFC0_AFEM1_AFEM1_ASEM1_CLR1_CADS0_DTLA1_DTLB1_DTVA0_DTVB0_EPS0_FDSI0_GRPM1_GRVWA8_GRVWB8_GSU0_GSUAMB_GLS0_ISA950_IU1_K1_LDSTI0_LBSPPA1024_LBSPPB1024_LBSPPM0_LPA16_LPB16_LPM0_LRVW8_LWPMn1_MIAV0_MIWT8_8_MO40_NTn1_NTA2_NTB1_NTC4_NTD2_NTM0_NEPBS12_NLCA1_NLCB1_ONLL1_PGR2_PLR1_PKA1_SIA3_SS1_SPO0_SRVW0_SSO0_SVW8_SK3_SKFTR0_SKXCCM0_TLDS1_ULSGRO0_USL1_UIOFGRO0_USFGRO0_VSn1_VWA8_VWB8_WSGRA0_WSGRB0_WS64_WG32_8_1,0.0,888.28,2078.3
gfx950,256,16384,6144,7168,False,torch.bfloat16,torch.bfloat16,False,False,hipblaslt,439111,0,917.9176,Cijk_Alik_Bljk_BBS_BH_Bias_HA_S_SAV_UserArgs_MT256x256x64_MI16x16x1_CMS_SN_LDSB0_AFC0_AFEM1_AFEM1_ASEM1_CLR1_CADS0_DTLA1_DTLB1_DTVA0_DTVB0_EPS0_FDSI0_GRPM1_GRVWA8_GRVWB8_GSU0_GSUAMB_GLS0_ISA950_IU1_K1_LDSTI0_LBSPPA1024_LBSPPB1024_LBSPPM0_LPA16_LPB16_LPM0_LRVW8_LWPMn1_MIAV0_MIWT8_8_MO40_NTn1_NTA2_NTB1_NTC4_NTD2_NTM0_NEPBS12_NLCA1_NLCB1_ONLL1_PGR2_PLR1_PKA1_SIA3_SS1_SPO0_SRVW0_SSO0_SVW8_SK3_SKFTR0_SKXCCM0_TLDS1_ULSGRO0_USL1_UIOFGRO0_USFGRO0_VSn1_VWA8_VWB8_WSGRA0_WSGRB0_WS64_WG32_8_1,0.0,1572.16,571.17
gfx950,256,16384,7168,1536,False,torch.bfloat16,torch.bfloat16,False,False,hipblaslt,438241,0,285.7495,Cijk_Alik_Bljk_BBS_BH_Bias_HA_S_SAV_UserArgs_MT256x256x64_MI16x16x1_CMS_SN_LDSB0_AFC0_AFEM1_AFEM1_ASEM1_CLR1_CADS0_DTLA1_DTLB1_DTVA0_DTVB0_EPS0_FDSI0_GRPM1_GRVWA8_GRVWB8_GSU0_GSUAMB_GLS0_ISA950_IU1_K1_LDSTI0_LBSPPA1024_LBSPPB1024_LBSPPM0_LPA16_LPB16_LPM0_LRVW8_LWPMn1_MIAV0_MIWT8_8_MO40_NTn1_NTA3_NTB2_NTC4_NTD4_NTM0_NEPBS16_NLCA1_NLCB1_ONLL1_PGR2_PLR1_PKA1_SIA3_SS1_SPO0_SRVW0_SSO0_SVW8_SK3_SKFTR0_SKXCCM4_TLDS1_ULSGRO0_USL1_UIOFGRO0_USFGRO0_VSn1_VWA8_VWB8_WSGRA0_WSGRB0_WS64_WG32_8_1,0.0,1262.56,1075.18
gfx950,256,16384,7168,3584,False,torch.bfloat16,torch.bfloat16,False,False,hipblaslt,439112,0,575.7408,Cijk_Alik_Bljk_BBS_BH_Bias_HA_S_SAV_UserArgs_MT256x256x64_MI16x16x1_CMS_SN_LDSB0_AFC0_AFEM1_AFEM1_ASEM1_CLR1_CADS0_DTLA1_DTLB1_DTVA0_DTVB0_EPS0_FDSI0_GRPM1_GRVWA8_GRVWB8_GSU0_GSUAMB_GLS0_ISA950_IU1_K1_LDSTI0_LBSPPA1024_LBSPPB1024_LBSPPM0_LPA16_LPB16_LPM0_LRVW8_LWPMn1_MIAV0_MIWT8_8_MO40_NTn1_NTA1_NTB3_NTC5_NTD4_NTM0_NEPBS16_NLCA1_NLCB1_ONLL1_PGR2_PLR1_PKA1_SIA3_SS1_SPO1_SRVW0_SSO4_SVW8_SK3_SKFTR0_SKXCCM0_TLDS1_ULSGRO0_USL1_UIOFGRO0_USFGRO0_VSn1_VWA8_VWB8_WSGRA0_WSGRB0_WS64_WG32_8_1,0.0,1462.14,701.19
gfx950,256,16384,7168,4224,False,torch.bfloat16,torch.bfloat16,False,False,hipblaslt,439111,0,663.2511,Cijk_Alik_Bljk_BBS_BH_Bias_HA_S_SAV_UserArgs_MT256x256x64_MI16x16x1_CMS_SN_LDSB0_AFC0_AFEM1_AFEM1_ASEM1_CLR1_CADS0_DTLA1_DTLB1_DTVA0_DTVB0_EPS0_FDSI0_GRPM1_GRVWA8_GRVWB8_GSU0_GSUAMB_GLS0_ISA950_IU1_K1_LDSTI0_LBSPPA1024_LBSPPB1024_LBSPPM0_LPA16_LPB16_LPM0_LRVW8_LWPMn1_MIAV0_MIWT8_8_MO40_NTn1_NTA2_NTB1_NTC4_NTD2_NTM0_NEPBS12_NLCA1_NLCB1_ONLL1_PGR2_PLR1_PKA1_SIA3_SS1_SPO0_SRVW0_SSO0_SVW8_SK3_SKFTR0_SKXCCM0_TLDS1_ULSGRO0_USL1_UIOFGRO0_USFGRO0_VSn1_VWA8_VWB8_WSGRA0_WSGRB0_WS64_WG32_8_1,0.0,1495.87,654.12
gfx950,256,16384,8448,7168,False,torch.bfloat16,torch.bfloat16,False,False,hipblaslt,439048,0,1400.0434,Cijk_Alik_Bljk_BBS_BH_Bias_HA_S_SAV_UserArgs_MT256x256x64_MI16x16x1_SN_LDSB0_AFC0_AFEM1_AFEM1_ASEM1_CLR0_CADS0_DTLA1_DTLB1_DTVA0_DTVB0_EPS0_FDSI0_GRPM1_GRVWA8_GRVWB8_GSUAMB_GLS0_ISA950_IU1_K1_LDSTI0_LBSPPA1024_LBSPPB1024_LBSPPM0_LPA16_LPB16_LPM0_LRVW8_LWPMn1_MIAV0_MIWT8_8_MO40_NTn1_NTA0_NTB3_NTC4_NTD5_NTM0_NEPBS0_NLCA1_NLCB1_ONLL1_PGR2_PLR1_PKA1_SIA3_SS1_SPO0_SRVW0_SSO0_SVW8_SK0_SKFTR0_SKXCCM0_TLDS1_ULSGRO0_USL1_UIOFGRO0_USFGRO1_VSn1_VWA8_VWB8_WSGRA0_WSGRB0_WS64_WG32_8_1,0.0,1417.3,452.0
gfx950,256,16384,1536,7168,False,torch.bfloat16,torch.bfloat16,False,False,torch,0,0,297.4444,native,0.0,1212.92,1032.91
39 changes: 39 additions & 0 deletions python/sglang/kernels/ops/kimi_k3/flydsl/__init__.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,39 @@
"""SGLang-maintained Kimi-K3 FlyDSL specializations."""

# AITER owns the FlyDSL toolchain bootstrap and shared tensor/buffer shims.
# Import it before local kernel modules so its vendored FlyDSL path is active.
import aiter as _aiter # noqa: F401

from .kimi_k3_kda_input_group64 import (
kimi_k3_kda_input_group64,
quantize_kimi_k3_kda_input_group64,
supports_kimi_k3_kda_input_group64,
)
from .kimi_k3_mla_gate import kimi_k3_mla_gate, supports_kimi_k3_mla_gate
from .kimi_k3_moe_preroute_fp8 import (
kimi_k3_moe_tri_projection_fp8,
kimi_k3_shared_down_fp8,
supports_kimi_k3_moe_tri_projection_fp8,
supports_kimi_k3_shared_down_fp8,
supports_kimi_k3_shared_down_fp8_weight,
)
from .latent_moe_tail_fp8 import (
latent_moe_tail_fp8,
quantize_latent_moe_tail_weight,
supports_latent_moe_tail_fp8,
)

__all__ = [
"kimi_k3_kda_input_group64",
"kimi_k3_mla_gate",
"kimi_k3_moe_tri_projection_fp8",
"kimi_k3_shared_down_fp8",
"latent_moe_tail_fp8",
"quantize_kimi_k3_kda_input_group64",
"quantize_latent_moe_tail_weight",
"supports_kimi_k3_kda_input_group64",
"supports_kimi_k3_mla_gate",
"supports_kimi_k3_moe_tri_projection_fp8",
"supports_kimi_k3_shared_down_fp8",
"supports_kimi_k3_shared_down_fp8_weight",
]
Empty file.
Loading
Loading