[Fix] Patch Kimi image for SM120 FlashInfer MXFP4
This commit is contained in:
parent
b50de8fe99
commit
e3974e2352
@ -1,5 +1,7 @@
|
|||||||
# sskj — 多平台大模型推理性能基准测试项目
|
# sskj — 多平台大模型推理性能基准测试项目
|
||||||
|
|
||||||
|
**更新(2026-08-18 13:03:33 CST)**:Phase 5 候选镜像改为针对 Kimi 基础镜像原生 SGLang 源码应用最小兼容补丁,仅补齐 SM120 FlashInfer MXFP4 backend 分发、权重布局与 SiTU 激活映射,避免整体替换新版 Python 包造成 `sglang-kernel` 版本错配。
|
||||||
|
|
||||||
**更新(2026-08-18 12:24:19 CST)**:新增 Kimi-K3 / 601-604 / SGLang TP32×EP32 的真实 Prefill MoE backend 验收实验;固定比较 Marlin 与 FlashInfer MXFP4 在 16K→1、C=8/16、Chunk=8K/16K 下的 TTFT、E2E 与 Input TPS,并提供四节点一致候选镜像构建、原始证据和自动汇总,详见 `experiments/pro6000/kimi3_pro6000_sglang_tp32ep32_moe_backend_prefill/README.md`。
|
**更新(2026-08-18 12:24:19 CST)**:新增 Kimi-K3 / 601-604 / SGLang TP32×EP32 的真实 Prefill MoE backend 验收实验;固定比较 Marlin 与 FlashInfer MXFP4 在 16K→1、C=8/16、Chunk=8K/16K 下的 TTFT、E2E 与 Input TPS,并提供四节点一致候选镜像构建、原始证据和自动汇总,详见 `experiments/pro6000/kimi3_pro6000_sglang_tp32ep32_moe_backend_prefill/README.md`。
|
||||||
|
|
||||||
**更新(2026-08-17 11:56:47 CST)**:完成 Kimi-K3 / RTX PRO 6000D / SGLang FlashInfer MXFP4 MoE Phase 2 correctness harness;SwiGLU 控制矩阵 8/8 通过,覆盖 edge routing、EP global ID、NaN/Inf、确定性、CUDA Graph 和 Kimi 真尺寸,SiTU 目标矩阵 8/8 稳定复现同一 C++ 缺口,详见 `experiments/pro6000/kimi3_pro6000_sglang_sm120_flashinfer_mxfp4/README.md`。
|
**更新(2026-08-17 11:56:47 CST)**:完成 Kimi-K3 / RTX PRO 6000D / SGLang FlashInfer MXFP4 MoE Phase 2 correctness harness;SwiGLU 控制矩阵 8/8 通过,覆盖 edge routing、EP global ID、NaN/Inf、确定性、CUDA Graph 和 Kimi 真尺寸,SiTU 目标矩阵 8/8 稳定复现同一 C++ 缺口,详见 `experiments/pro6000/kimi3_pro6000_sglang_sm120_flashinfer_mxfp4/README.md`。
|
||||||
|
|||||||
@ -18,9 +18,9 @@ RUN python3 -m pip install --no-deps --force-reinstall \
|
|||||||
# SM120 AOT artifacts that passed the Phase 3/4 correctness gate.
|
# SM120 AOT artifacts that passed the Phase 3/4 correctness gate.
|
||||||
COPY aot/ /usr/local/lib/python3.12/dist-packages/flashinfer/data/aot/
|
COPY aot/ /usr/local/lib/python3.12/dist-packages/flashinfer/data/aot/
|
||||||
|
|
||||||
# Apply only the two production source changes from SGLang commit fb929bb. The
|
# Apply only the SM120 backend-selection, weight-layout, and SiTU activation
|
||||||
# Kimi base image intentionally pins sglang-kernel 0.4.5, so replacing the
|
# changes against this Kimi image's exact SGLang source. Replacing the newer
|
||||||
# entire newer Python tree would create a false package-version dependency.
|
# Python tree would break the image's intentional sglang-kernel 0.4.5 pin.
|
||||||
COPY sglang.patch /tmp/sglang.patch
|
COPY sglang.patch /tmp/sglang.patch
|
||||||
RUN patch --dry-run --batch --forward -p1 -d /sgl-workspace/sglang \
|
RUN patch --dry-run --batch --forward -p1 -d /sgl-workspace/sglang \
|
||||||
< /tmp/sglang.patch && \
|
< /tmp/sglang.patch && \
|
||||||
@ -35,5 +35,5 @@ RUN python3 -c "from pathlib import Path; p=Path('/sgl-workspace/sglang/python/s
|
|||||||
python3 -c "from flashinfer.fused_moe import core; assert core.CUTLASS_FUSED_MOE_SUPPORTS_SITU is True"
|
python3 -c "from flashinfer.fused_moe import core; assert core.CUTLASS_FUSED_MOE_SUPPORTS_SITU is True"
|
||||||
|
|
||||||
LABEL ai.meta-stone.purpose="Kimi-K3 SM120 FlashInfer MXFP4 Phase 5 acceptance" \
|
LABEL ai.meta-stone.purpose="Kimi-K3 SM120 FlashInfer MXFP4 Phase 5 acceptance" \
|
||||||
ai.meta-stone.sglang.patch="fb929bb" \
|
ai.meta-stone.sglang.patch="kimi-image-sm120-flashinfer-mxfp4" \
|
||||||
ai.meta-stone.flashinfer.patch="7d1a4ee"
|
ai.meta-stone.flashinfer.patch="7d1a4ee"
|
||||||
|
|||||||
@ -15,8 +15,6 @@ BASE_IMAGE="${BASE_IMAGE:-lmsysorg/sglang:kimi-k3-fiv617situ-warm}"
|
|||||||
CANDIDATE_IMAGE="${CANDIDATE_IMAGE:-local/sglang:kimi-k3-sm120-flashinfer-mxfp4-phase5}"
|
CANDIDATE_IMAGE="${CANDIDATE_IMAGE:-local/sglang:kimi-k3-sm120-flashinfer-mxfp4-phase5}"
|
||||||
FLASHINFER_TARGET="${FLASHINFER_TARGET:-/tmp/hzy-fi-situ-phase3-v3}"
|
FLASHINFER_TARGET="${FLASHINFER_TARGET:-/tmp/hzy-fi-situ-phase3-v3}"
|
||||||
FLASHINFER_WHEEL="${FLASHINFER_WHEEL:-/data/hzy/artifacts/flashinfer-kimi-sm120-situ-phase3-v3/flashinfer_python-0.6.18-py3-none-any.whl}"
|
FLASHINFER_WHEEL="${FLASHINFER_WHEEL:-/data/hzy/artifacts/flashinfer-kimi-sm120-situ-phase3-v3/flashinfer_python-0.6.18-py3-none-any.whl}"
|
||||||
SGLANG_REPO="${SGLANG_REPO:-/data/hzy/src/sglang-kimi-sm120}"
|
|
||||||
SGLANG_PATCH_COMMIT="${SGLANG_PATCH_COMMIT:-fb929bb}"
|
|
||||||
|
|
||||||
BENCH_DATASET_NAME="${BENCH_DATASET_NAME:-random-ids}"
|
BENCH_DATASET_NAME="${BENCH_DATASET_NAME:-random-ids}"
|
||||||
BENCH_MODULE="${BENCH_MODULE:-sglang.benchmark.serving}"
|
BENCH_MODULE="${BENCH_MODULE:-sglang.benchmark.serving}"
|
||||||
|
|||||||
@ -0,0 +1,303 @@
|
|||||||
|
--- a/python/sglang/srt/layers/moe/moe_runner/flashinfer_cutlass.py
|
||||||
|
+++ b/python/sglang/srt/layers/moe/moe_runner/flashinfer_cutlass.py
|
||||||
|
@@ -114,17 +114,21 @@
|
||||||
|
from sglang.srt.layers.moe.moe_runner.flashinfer_trtllm import get_activation_type
|
||||||
|
|
||||||
|
_, ActivationType = _flashinfer_cutlass_fused_moe()
|
||||||
|
- activation = ActivationType(
|
||||||
|
- get_activation_type(
|
||||||
|
- runner_config.activation,
|
||||||
|
- is_gated=runner_config.is_gated,
|
||||||
|
+ if runner_config.activation == "situ" and runner_config.is_gated:
|
||||||
|
+ activation = ActivationType.Situ
|
||||||
|
+ else:
|
||||||
|
+ activation = ActivationType(
|
||||||
|
+ get_activation_type(
|
||||||
|
+ runner_config.activation,
|
||||||
|
+ is_gated=runner_config.is_gated,
|
||||||
|
+ )
|
||||||
|
)
|
||||||
|
- )
|
||||||
|
supported = {
|
||||||
|
ActivationType.Swiglu,
|
||||||
|
ActivationType.Geglu,
|
||||||
|
ActivationType.Relu2,
|
||||||
|
ActivationType.Identity,
|
||||||
|
+ ActivationType.Situ,
|
||||||
|
}
|
||||||
|
assert activation in supported, (
|
||||||
|
f"Activation {runner_config.activation!r} "
|
||||||
|
@@ -306,7 +310,7 @@
|
||||||
|
quant_info, FlashInferCutlassMxfp4MoeQuantInfo
|
||||||
|
), f"Unexpected quant_info type for flashinfer_mxfp4: {type(quant_info)}"
|
||||||
|
|
||||||
|
- flashinfer_cutlass_fused_moe, ActivationType = _flashinfer_cutlass_fused_moe()
|
||||||
|
+ flashinfer_cutlass_fused_moe, _ = _flashinfer_cutlass_fused_moe()
|
||||||
|
|
||||||
|
x = dispatch_output.hidden_states
|
||||||
|
topk_output = dispatch_output.topk_output
|
||||||
|
@@ -384,7 +388,7 @@
|
||||||
|
ep_rank=quant_info.moe_ep_rank,
|
||||||
|
use_w4_group_scaling=not use_mxfp8_act_scaling,
|
||||||
|
use_mxfp8_act_scaling=use_mxfp8_act_scaling,
|
||||||
|
- activation_type=ActivationType.Swiglu,
|
||||||
|
+ activation_type=_activation_type(runner_config),
|
||||||
|
tune_max_num_tokens=next_power_of_2(x.shape[0]),
|
||||||
|
output=out,
|
||||||
|
)
|
||||||
|
--- a/python/sglang/srt/layers/quantization/mxfp4.py
|
||||||
|
+++ b/python/sglang/srt/layers/quantization/mxfp4.py
|
||||||
|
@@ -340,8 +340,7 @@
|
||||||
|
self.flashinfer_mxfp4_moe_precision = (
|
||||||
|
get_server_args().flashinfer_mxfp4_moe_precision
|
||||||
|
)
|
||||||
|
- # When `flashinfer_mxfp4` is enabled, dispatch to one of two FlashInfer
|
||||||
|
- # entry points depending on the GPU:
|
||||||
|
+ # When `flashinfer_mxfp4` is enabled, dispatch by GPU architecture:
|
||||||
|
# - SM100 (Blackwell) -> trtllm_fp4_block_scale_moe (existing)
|
||||||
|
# - SM90 (Hopper) -> cutlass_fused_moe(use_w4_group_scaling=True)
|
||||||
|
# (FlashInfer PR #3084, post-0.6.10)
|
||||||
|
@@ -349,6 +348,8 @@
|
||||||
|
if self.use_flashinfer:
|
||||||
|
if is_sm100_supported():
|
||||||
|
self._fi_kernel = "trtllm_sm100"
|
||||||
|
+ elif is_sm120_supported():
|
||||||
|
+ self._fi_kernel = "cutlass_sm120"
|
||||||
|
elif is_sm90_supported():
|
||||||
|
if not _FI_HAS_SM90_CUTLASS_MXFP4:
|
||||||
|
raise RuntimeError(
|
||||||
|
@@ -360,7 +361,8 @@
|
||||||
|
self._fi_kernel = "cutlass_sm90"
|
||||||
|
else:
|
||||||
|
raise NotImplementedError(
|
||||||
|
- "moe_runner_backend=flashinfer_mxfp4 requires SM90 or SM100."
|
||||||
|
+ "moe_runner_backend=flashinfer_mxfp4 requires SM90, SM100, "
|
||||||
|
+ "or SM120."
|
||||||
|
)
|
||||||
|
|
||||||
|
def create_weights(
|
||||||
|
@@ -412,7 +414,7 @@
|
||||||
|
intermediate_size_per_partition_after_pad = round_up(
|
||||||
|
intermediate_size_per_partition, triton_kernels_padding_alignment
|
||||||
|
)
|
||||||
|
- elif self._fi_kernel == "cutlass_sm90":
|
||||||
|
+ elif self._fi_kernel in ("cutlass_sm90", "cutlass_sm120"):
|
||||||
|
# cutlass mixed-input GEMM contraction dim K must be % 128 == 0
|
||||||
|
# (interleave factor for MXFP4 group_size=32 is 4). The kernel
|
||||||
|
# also expects ``fc1_expert_weights`` in halved ``[up; gate]``
|
||||||
|
@@ -609,6 +611,9 @@
|
||||||
|
if self._fi_kernel == "cutlass_sm90":
|
||||||
|
self._process_weights_for_sm90_cutlass(layer)
|
||||||
|
return
|
||||||
|
+ if self._fi_kernel == "cutlass_sm120":
|
||||||
|
+ self._process_weights_for_sm120_cutlass(layer)
|
||||||
|
+ return
|
||||||
|
if self.use_flashinfer:
|
||||||
|
# Per-expert buffers are local (create_weights uses num_local_experts);
|
||||||
|
# the global self.num_experts here breaks EP>1. Mirrors the SM90 path.
|
||||||
|
@@ -1130,7 +1135,135 @@
|
||||||
|
layer.w2_weight_bias = Parameter(w2_bias_padded, requires_grad=False)
|
||||||
|
|
||||||
|
torch.cuda.empty_cache()
|
||||||
|
+
|
||||||
|
+ def _process_weights_for_sm120_cutlass(self, layer):
|
||||||
|
+ """Prepare Kimi/GPT-OSS MXFP4 experts for FlashInfer CUTLASS SM120."""
|
||||||
|
+ from flashinfer import block_scale_interleave
|
||||||
|
+
|
||||||
|
+ sf_block_size = 32
|
||||||
|
+ n_un = layer.w13_weight.shape[1] // 2
|
||||||
|
+ k_un = layer.w13_weight.shape[2] * 2
|
||||||
|
+ n_pad = self._padded_intermediate
|
||||||
|
+ k_pad = self._padded_hidden
|
||||||
|
+ num_experts = layer.num_local_experts
|
||||||
|
+ device = layer.w13_weight.device
|
||||||
|
+ gate_up_interleaved = getattr(
|
||||||
|
+ layer.moe_runner_config, "gate_up_interleaved", True
|
||||||
|
+ )
|
||||||
|
+
|
||||||
|
+ def split_gate_up(unpadded):
|
||||||
|
+ if gate_up_interleaved:
|
||||||
|
+ return unpadded[:, 0::2, :], unpadded[:, 1::2, :]
|
||||||
|
+ return unpadded[:, :n_un, :], unpadded[:, n_un : 2 * n_un, :]
|
||||||
|
+
|
||||||
|
+ def stack_up_gate(unpadded, last_pad, last_un):
|
||||||
|
+ gate_rows, up_rows = split_gate_up(unpadded)
|
||||||
|
+ out = torch.zeros(
|
||||||
|
+ num_experts,
|
||||||
|
+ 2 * n_pad,
|
||||||
|
+ last_pad,
|
||||||
|
+ dtype=unpadded.dtype,
|
||||||
|
+ device=device,
|
||||||
|
+ )
|
||||||
|
+ out[:, :n_un, :last_un] = up_rows
|
||||||
|
+ out[:, n_pad : n_pad + n_un, :last_un] = gate_rows
|
||||||
|
+ return out
|
||||||
|
+
|
||||||
|
+ w13_padded = stack_up_gate(
|
||||||
|
+ layer.w13_weight.data, k_pad // 2, k_un // 2
|
||||||
|
+ )
|
||||||
|
+ w13_scale_padded = stack_up_gate(
|
||||||
|
+ layer.w13_weight_scale.data,
|
||||||
|
+ k_pad // sf_block_size,
|
||||||
|
+ k_un // sf_block_size,
|
||||||
|
+ )
|
||||||
|
|
||||||
|
+ bias_dtype = layer.w13_weight_bias.dtype
|
||||||
|
+ w13_bias_padded = torch.zeros(
|
||||||
|
+ num_experts, 2 * n_pad, dtype=bias_dtype, device=device
|
||||||
|
+ )
|
||||||
|
+ if gate_up_interleaved:
|
||||||
|
+ gate_bias = layer.w13_weight_bias.data[:, 0::2]
|
||||||
|
+ up_bias = layer.w13_weight_bias.data[:, 1::2]
|
||||||
|
+ else:
|
||||||
|
+ gate_bias = layer.w13_weight_bias.data[:, :n_un]
|
||||||
|
+ up_bias = layer.w13_weight_bias.data[:, n_un : 2 * n_un]
|
||||||
|
+ w13_bias_padded[:, :n_un] = up_bias
|
||||||
|
+ w13_bias_padded[:, n_pad : n_pad + n_un] = gate_bias
|
||||||
|
+
|
||||||
|
+ def pad_w2(unpadded, last_pad, last_un):
|
||||||
|
+ out = torch.zeros(
|
||||||
|
+ num_experts,
|
||||||
|
+ k_pad,
|
||||||
|
+ last_pad,
|
||||||
|
+ dtype=unpadded.dtype,
|
||||||
|
+ device=device,
|
||||||
|
+ )
|
||||||
|
+ out[:, :k_un, :last_un] = unpadded[:, :k_un, :]
|
||||||
|
+ return out
|
||||||
|
+
|
||||||
|
+ w2_padded = pad_w2(layer.w2_weight.data, n_pad // 2, n_un // 2)
|
||||||
|
+ w2_scale_padded = pad_w2(
|
||||||
|
+ layer.w2_weight_scale.data,
|
||||||
|
+ n_pad // sf_block_size,
|
||||||
|
+ n_un // sf_block_size,
|
||||||
|
+ )
|
||||||
|
+ w2_bias_padded = torch.zeros(
|
||||||
|
+ num_experts, k_pad, dtype=bias_dtype, device=device
|
||||||
|
+ )
|
||||||
|
+ w2_bias_padded[:, :k_un] = layer.w2_weight_bias.data
|
||||||
|
+
|
||||||
|
+ layer.w13_weight = Parameter(w13_padded, requires_grad=False)
|
||||||
|
+ layer.w2_weight = Parameter(w2_padded, requires_grad=False)
|
||||||
|
+ layer.w13_weight_scale = Parameter(
|
||||||
|
+ block_scale_interleave(w13_scale_padded).reshape_as(w13_scale_padded),
|
||||||
|
+ requires_grad=False,
|
||||||
|
+ )
|
||||||
|
+ layer.w2_weight_scale = Parameter(
|
||||||
|
+ block_scale_interleave(w2_scale_padded).reshape_as(w2_scale_padded),
|
||||||
|
+ requires_grad=False,
|
||||||
|
+ )
|
||||||
|
+ layer.w13_weight_bias = Parameter(w13_bias_padded, requires_grad=False)
|
||||||
|
+ layer.w2_weight_bias = Parameter(w2_bias_padded, requires_grad=False)
|
||||||
|
+
|
||||||
|
+ activation = layer.moe_runner_config.activation
|
||||||
|
+ alpha = layer.moe_runner_config.gemm1_alpha
|
||||||
|
+ beta = layer.moe_runner_config.gemm1_beta
|
||||||
|
+ limit = layer.moe_runner_config.gemm1_clamp_limit
|
||||||
|
+ if activation == "situ":
|
||||||
|
+ alpha = 4.0 if alpha is None else alpha
|
||||||
|
+ beta = 25.0 if limit is None else limit
|
||||||
|
+ limit = None
|
||||||
|
+ else:
|
||||||
|
+ alpha = 1.702 if alpha is None else alpha
|
||||||
|
+ beta = 1.0 if beta is None else beta
|
||||||
|
+ limit = 7.0 if limit is None else limit
|
||||||
|
+
|
||||||
|
+ layer.swiglu_alpha = Parameter(
|
||||||
|
+ torch.full((num_experts,), alpha, dtype=torch.float32, device=device),
|
||||||
|
+ requires_grad=False,
|
||||||
|
+ )
|
||||||
|
+ layer.swiglu_beta = Parameter(
|
||||||
|
+ torch.full((num_experts,), beta, dtype=torch.float32, device=device),
|
||||||
|
+ requires_grad=False,
|
||||||
|
+ )
|
||||||
|
+ layer.swiglu_limit = (
|
||||||
|
+ None
|
||||||
|
+ if limit is None
|
||||||
|
+ else Parameter(
|
||||||
|
+ torch.full(
|
||||||
|
+ (num_experts,), limit, dtype=torch.float32, device=device
|
||||||
|
+ ),
|
||||||
|
+ requires_grad=False,
|
||||||
|
+ )
|
||||||
|
+ )
|
||||||
|
+ layer.mxfp4_weight_global_scale = Parameter(
|
||||||
|
+ torch.ones(num_experts, dtype=torch.float32, device=device),
|
||||||
|
+ requires_grad=False,
|
||||||
|
+ )
|
||||||
|
+ layer._mxfp4_backend = "flashinfer_cutlass_sm120"
|
||||||
|
+ torch.cuda.empty_cache()
|
||||||
|
+
|
||||||
|
def create_moe_runner(
|
||||||
|
self, layer: torch.nn.Module, moe_runner_config: MoeRunnerConfig
|
||||||
|
):
|
||||||
|
@@ -1154,10 +1287,23 @@
|
||||||
|
or moe_runner_backend.is_deep_gemm()
|
||||||
|
):
|
||||||
|
self.runner = MoeRunner(moe_runner_backend, moe_runner_config)
|
||||||
|
- elif (
|
||||||
|
- moe_runner_backend.is_flashinfer_mxfp4()
|
||||||
|
- and self._fi_kernel == "cutlass_sm90"
|
||||||
|
+ elif moe_runner_backend.is_flashinfer_mxfp4() and self._fi_kernel in (
|
||||||
|
+ "cutlass_sm90",
|
||||||
|
+ "cutlass_sm120",
|
||||||
|
):
|
||||||
|
+ if (
|
||||||
|
+ self._fi_kernel == "cutlass_sm120"
|
||||||
|
+ and moe_runner_config.activation == "situ"
|
||||||
|
+ ):
|
||||||
|
+ from flashinfer.fused_moe import core as flashinfer_moe_core
|
||||||
|
+
|
||||||
|
+ if not getattr(
|
||||||
|
+ flashinfer_moe_core, "CUTLASS_FUSED_MOE_SUPPORTS_SITU", False
|
||||||
|
+ ):
|
||||||
|
+ raise RuntimeError(
|
||||||
|
+ "Kimi-K3 FlashInfer MXFP4 MoE on SM120 requires a "
|
||||||
|
+ "FlashInfer build with CUTLASS SiTU support."
|
||||||
|
+ )
|
||||||
|
# Register the fused func at runner construction so the FusedOpPool
|
||||||
|
# lookup at `MoeRunner.__init__` finds it.
|
||||||
|
import sglang.srt.layers.moe.moe_runner.flashinfer_cutlass # noqa: F401
|
||||||
|
@@ -1215,8 +1361,33 @@
|
||||||
|
quant_info = build_marlin_moe_quant_info(layer)
|
||||||
|
return self.runner.run(
|
||||||
|
dispatch_output._replace(hidden_states=x_padded), quant_info
|
||||||
|
+ )
|
||||||
|
+
|
||||||
|
+ def _apply_sm120_cutlass(self, layer, dispatch_output):
|
||||||
|
+ """SM120 MXFP8 x MXFP4 MoE via FlashInfer CUTLASS."""
|
||||||
|
+ from sglang.srt.layers.moe.moe_runner.flashinfer_cutlass import (
|
||||||
|
+ FlashInferCutlassMxfp4MoeQuantInfo,
|
||||||
|
)
|
||||||
|
|
||||||
|
+ quant_info = FlashInferCutlassMxfp4MoeQuantInfo(
|
||||||
|
+ w13_weight=layer.w13_weight,
|
||||||
|
+ w2_weight=layer.w2_weight,
|
||||||
|
+ w13_weight_scale=layer.w13_weight_scale,
|
||||||
|
+ w2_weight_scale=layer.w2_weight_scale,
|
||||||
|
+ mxfp4_weight_global_scale=layer.mxfp4_weight_global_scale,
|
||||||
|
+ w13_bias=layer.w13_weight_bias if self.with_bias else None,
|
||||||
|
+ w2_bias=layer.w2_weight_bias if self.with_bias else None,
|
||||||
|
+ swiglu_alpha=layer.swiglu_alpha,
|
||||||
|
+ swiglu_beta=layer.swiglu_beta,
|
||||||
|
+ swiglu_limit=layer.swiglu_limit,
|
||||||
|
+ moe_tp_size=layer.moe_tp_size,
|
||||||
|
+ moe_tp_rank=layer.moe_tp_rank,
|
||||||
|
+ moe_ep_size=layer.moe_ep_size,
|
||||||
|
+ moe_ep_rank=layer.moe_ep_rank,
|
||||||
|
+ padded_hidden=self._padded_hidden,
|
||||||
|
+ )
|
||||||
|
+ return self.runner.run(dispatch_output, quant_info)
|
||||||
|
+
|
||||||
|
def apply(
|
||||||
|
self,
|
||||||
|
layer: torch.nn.Module,
|
||||||
|
@@ -1290,6 +1461,8 @@
|
||||||
|
|
||||||
|
if self._fi_kernel == "cutlass_sm90":
|
||||||
|
return self._apply_sm90_cutlass(layer, dispatch_output)
|
||||||
|
+ if self._fi_kernel == "cutlass_sm120":
|
||||||
|
+ return self._apply_sm120_cutlass(layer, dispatch_output)
|
||||||
|
if self.use_flashinfer:
|
||||||
|
# When bf16 mode is enabled, we don't need to quantize the input,
|
||||||
|
# TRT-LLM automatically handles quantization in the kernel implementation and pipelines it with GEMM operations,
|
||||||
@ -79,7 +79,8 @@ sudo_host() {
|
|||||||
check_inputs() {
|
check_inputs() {
|
||||||
local path
|
local path
|
||||||
for path in "$MODEL_PATH" "$FLASHINFER_TARGET/flashinfer" \
|
for path in "$MODEL_PATH" "$FLASHINFER_TARGET/flashinfer" \
|
||||||
"$FLASHINFER_WHEEL" "$SGLANG_REPO/.git"; do
|
"$FLASHINFER_WHEEL" \
|
||||||
|
"$SCRIPT_DIR/patches/sglang-kimi-image-sm120-flashinfer-mxfp4.patch"; do
|
||||||
[[ -e "$path" ]] || { echo "ERROR: required path missing: $path" >&2; exit 2; }
|
[[ -e "$path" ]] || { echo "ERROR: required path missing: $path" >&2; exit 2; }
|
||||||
done
|
done
|
||||||
[[ -d "$FLASHINFER_TARGET/flashinfer/data/aot/fused_moe_120" ]] || {
|
[[ -d "$FLASHINFER_TARGET/flashinfer/data/aot/fused_moe_120" ]] || {
|
||||||
@ -93,16 +94,10 @@ prepare_build_context() {
|
|||||||
rm -rf "$BUILD_CONTEXT"
|
rm -rf "$BUILD_CONTEXT"
|
||||||
mkdir -p "$BUILD_CONTEXT/aot"
|
mkdir -p "$BUILD_CONTEXT/aot"
|
||||||
cp "$SCRIPT_DIR/Dockerfile" "$BUILD_CONTEXT/Dockerfile"
|
cp "$SCRIPT_DIR/Dockerfile" "$BUILD_CONTEXT/Dockerfile"
|
||||||
|
cp "$SCRIPT_DIR/patches/sglang-kimi-image-sm120-flashinfer-mxfp4.patch" \
|
||||||
|
"$BUILD_CONTEXT/sglang.patch"
|
||||||
cp "$FLASHINFER_WHEEL" "$BUILD_CONTEXT/flashinfer_python-0.6.18-py3-none-any.whl"
|
cp "$FLASHINFER_WHEEL" "$BUILD_CONTEXT/flashinfer_python-0.6.18-py3-none-any.whl"
|
||||||
cp -a "$FLASHINFER_TARGET/flashinfer/data/aot/." "$BUILD_CONTEXT/aot/"
|
cp -a "$FLASHINFER_TARGET/flashinfer/data/aot/." "$BUILD_CONTEXT/aot/"
|
||||||
git -C "$SGLANG_REPO" show --format= --binary "$SGLANG_PATCH_COMMIT" -- \
|
|
||||||
python/sglang/srt/layers/moe/moe_runner/flashinfer_cutlass.py \
|
|
||||||
python/sglang/srt/layers/quantization/mxfp4.py \
|
|
||||||
>"$BUILD_CONTEXT/sglang.patch"
|
|
||||||
[[ -s "$BUILD_CONTEXT/sglang.patch" ]] || {
|
|
||||||
echo "ERROR: empty SGLang patch from ${SGLANG_PATCH_COMMIT}" >&2
|
|
||||||
exit 2
|
|
||||||
}
|
|
||||||
find "$BUILD_CONTEXT" -type d -name __pycache__ -prune -exec rm -rf {} +
|
find "$BUILD_CONTEXT" -type d -name __pycache__ -prune -exec rm -rf {} +
|
||||||
find "$BUILD_CONTEXT" -type f -name '*.pyc' -delete
|
find "$BUILD_CONTEXT" -type f -name '*.pyc' -delete
|
||||||
du -sh "$BUILD_CONTEXT" | tee "$RESULT_ROOT/build/context_size.txt"
|
du -sh "$BUILD_CONTEXT" | tee "$RESULT_ROOT/build/context_size.txt"
|
||||||
|
|||||||
Loading…
x
Reference in New Issue
Block a user