Compare commits

...
10 Commits
Author SHA1 Message Date
Kevin Mi 8ac19cc19f [AMD][Kimi-K3] Fix deferred KDA gate projection and update DCP cookbook (#39066)
PR Test (XPU) / finish (push) Blocked by required conditions
PR Test (Arm64) / check-changes (push) Successful in 10s
PR Test (NPU) / set-image-config (push) Successful in 1s
PR Test (NPU) / Recommend tests from coverage (push) Skipped
PR Test (NPU) / check-changes (push) Successful in 12s
PR Test (sgl-router) / gate (push) Successful in 8s
PR Test (Xeon) / check-changes (push) Successful in 9s
PR Test (XPU) / check-changes (push) Successful in 16s
pr-test-arm64.yml / pr-gate (push) Successful in 3s
PR Test (Arm64) / pr-gate (push) Successful in 3s
PR Test (Arm64) / build-test (push) Waiting to run
pr-test-npu.yml / pr-gate (push) Successful in 2s
PR Test (NPU) / pr-gate (push) Successful in 2s
PR Test (sgl-router) / tier-1 — lint (push) Failing after 33s
PR Test (sgl-router) / tier-2 — build + test (push) Skipped
PR Test (sgl-router) / tier-3 — docker (placeholder) (push) Skipped
PR Test (sgl-router) / tier-3 — k8s integration (push) Skipped
PR Test (sgl-router) / tier-3 — e2e (push) Skipped
pr-test-xpu.yml / pr-gate (push) Successful in 2s
PR Test (XPU) / pr-gate (push) Successful in 2s
pr-test-xeon.yml / pr-gate (push) Successful in 3s
PR Test (XPU) / stage-a-test-1-gpu-xpu (push) Waiting to run
PR Test (XPU) / multimodal-gen-test-1-gpu-xpu (push) Waiting to run
PR Test (Xeon) / pr-gate (push) Successful in 3s
PR Test (sgl-router) / finish (push) Successful in 1s
PR Test (Xeon) / build-test (gnr, gnr, xeon-gnr, stage-a-tp-test-cpu-intel) (push) Waiting to run
PR Test (Xeon) / build-test (spr1, 0, 3, spr, xeon-spr, stage-a-test-cpu-intel,stage-b-test-cpu-intel) (push) Waiting to run
PR Test (Xeon) / build-test (spr2, 1, 3, spr, xeon-spr, stage-a-test-cpu-intel,stage-b-test-cpu-intel) (push) Waiting to run
PR Test (Xeon) / build-test (spr3, 2, 3, spr, xeon-spr, stage-a-test-cpu-intel,stage-b-test-cpu-intel) (push) Waiting to run
Lint / lint (push) Failing after 2m51s
PR Test (NPU) / base-a-test-1-npu-a2 (push) Canceled after 0s
PR Test (NPU) / base-b-test-1-npu-a3 (push) Canceled after 0s
PR Test (NPU) / base-b-test-2-npu-a3 (push) Canceled after 0s
PR Test (NPU) / base-b-test-4-npu-a3 (push) Canceled after 0s
PR Test (NPU) / base-b-test-8-npu-a3 (push) Canceled after 0s
PR Test (NPU) / base-b-test-16-npu-a3 (push) Canceled after 0s
PR Test (NPU) / multimodal-gen-test-1-npu-a3 (0) (push) Canceled after 0s
PR Test (NPU) / multimodal-gen-test-1-npu-a3 (1) (push) Canceled after 0s
PR Test (NPU) / multimodal-gen-test-4-npu-a3 (0) (push) Canceled after 0s
PR Test (NPU) / multimodal-gen-test-4-npu-a3 (1) (push) Canceled after 0s
PR Test (NPU) / base-c-test-acc-2-npu-a3 (push) Canceled after 0s
PR Test (NPU) / base-c-test-acc-16-npu-a3 (push) Canceled after 0s
PR Test (NPU) / base-c-test-perf-2-npu-a3 (push) Canceled after 0s
PR Test (NPU) / base-c-test-perf-16-npu-a3 (push) Canceled after 0s
PR Test (NPU) / Analyze failure report (push) Canceled after 0s
PR Test (NPU) / setup-covstub (push) Canceled after 0s
PR Test (NPU) / pr-test-npu-finish (push) Canceled after 0s
pr-test-npu.yml / run (${{ fromJson(inputs.partitions).arr }}) (push) Canceled after 0s
2026-09-22 06:35:12 +00:00
Mohammad Miadh AngkadandMohammad Angkad 4c81cd1b09 [KDA] Fix missing beta sigmoid in PTX prefill (#40685)
Co-authored-by: Mohammad Angkad <mohammad.angkad@radixark.ai>
2026-09-21 23:24:38 -07:00
Khoa Pham bc22e1de9e [DSpark] Fix draft CUDA graph stream explosion (#40658) 2026-09-21 23:05:45 -07:00
Brayden ZhongandXinyuan Tong a0781f2714 [Docs] GLM-5.3/5.3-Flash cookbooks: enable reasoning/tool-call parsers by default via auto (#40497)
Co-authored-by: Xinyuan Tong <xinyuantong.cs@gmail.com>
2026-09-22 13:54:08 +08:00
Guangda LiuandGuangda Liu 04c0913434 [HiSparse] Add MHA hisparse support for MiniMax M3 (#31446)
Co-authored-by: Guangda Liu <bingps@users.noreply.github.com>
2026-09-22 13:28:03 +08:00
095e45100b [AMD] [GLM-5.3-Flash Day 0] Route mHC through AITER on gfx950 (#38545)
Co-authored-by: Raiden-Makoto <Raiden-Makoto@users.noreply.github.com>
Co-authored-by: Thomas Wang <thomawan@amd.com>
Co-authored-by: Kevin Mi <45493463+kevin-mii@users.noreply.github.com>
Co-authored-by: Kevin Mi <mikevin920@yahoo.com>
Co-authored-by: Claude Fable 5.1 <noreply@anthropic.com>
2026-09-21 22:24:29 -07:00
kangwangamd 264da63319 [AMD] Update ROCm AITER pin to acf8fdf9 (#39965) 2026-09-21 22:01:17 -07:00
Piotr Mazurek b01961e295 [LFM2-VL] Add DSpark speculative decoding (#40651) 2026-09-21 21:49:51 -07:00
YAMY 9b59fc5db5 [ModelOpt][PP] Keep BF16 shared experts out of the NVFP4 fusion so TP1 pipeline stages can load (#40628) 2026-09-21 21:45:58 -07:00
2032f3a071 [Router] Abort the engine when a client disconnects mid-request (#39461)
Co-authored-by: Kangyan Zhou <kangyan.zhou@radixark.ai>
Co-authored-by: Claude Opus 5 (1M context) <noreply@anthropic.com>
Co-authored-by: Shangming Cai <csmthu@gmail.com>
Co-authored-by: Kan Wu <wukanustc@gmail.com>
2026-09-22 12:35:56 +08:00
59 changed files with 2217 additions and 301 deletions
+8 -8
View File
@@ -61,7 +61,7 @@ ENV BUILD_TRITON="0"
ENV BUILD_LLVM="0" ENV BUILD_LLVM="0"
ENV BUILD_AITER_ALL="1" ENV BUILD_AITER_ALL="1"
ENV BUILD_MOONCAKE="1" ENV BUILD_MOONCAKE="1"
ENV AITER_COMMIT_DEFAULT="4ad99832823dde2315b361cbd3b54b1c5c12acd5" ENV AITER_COMMIT_DEFAULT="acf8fdf9307431ece8ee275971c41cb3d1a7020b"
# =============================== # ===============================
# Base image 942 with rocm720 and args # Base image 942 with rocm720 and args
@@ -71,7 +71,7 @@ ENV BUILD_TRITON="1"
ENV BUILD_LLVM="0" ENV BUILD_LLVM="0"
ENV BUILD_AITER_ALL="1" ENV BUILD_AITER_ALL="1"
ENV BUILD_MOONCAKE="1" ENV BUILD_MOONCAKE="1"
ENV AITER_COMMIT_DEFAULT="4ad99832823dde2315b361cbd3b54b1c5c12acd5" ENV AITER_COMMIT_DEFAULT="acf8fdf9307431ece8ee275971c41cb3d1a7020b"
ENV TRITON_COMMIT_DEFAULT="42270451990532c67e69d753fbd026f28fcc4840" ENV TRITON_COMMIT_DEFAULT="42270451990532c67e69d753fbd026f28fcc4840"
# =============================== # ===============================
@@ -82,7 +82,7 @@ ENV BUILD_TRITON="1"
ENV BUILD_LLVM="0" ENV BUILD_LLVM="0"
ENV BUILD_AITER_ALL="1" ENV BUILD_AITER_ALL="1"
ENV BUILD_MOONCAKE="1" ENV BUILD_MOONCAKE="1"
ENV AITER_COMMIT_DEFAULT="4ad99832823dde2315b361cbd3b54b1c5c12acd5" ENV AITER_COMMIT_DEFAULT="acf8fdf9307431ece8ee275971c41cb3d1a7020b"
# Pin the ROCm torch stack for every pip invocation in this flavor. The file is # Pin the ROCm torch stack for every pip invocation in this flavor. The file is
# filled in after the torch 2.11 upgrade below; it must already exist (empty is # filled in after the torch 2.11 upgrade below; it must already exist (empty is
# valid) because pip reads PIP_CONSTRAINT from the first pip call onwards. # valid) because pip reads PIP_CONSTRAINT from the first pip call onwards.
@@ -106,7 +106,7 @@ ENV BUILD_TRITON="0"
ENV BUILD_LLVM="0" ENV BUILD_LLVM="0"
ENV BUILD_AITER_ALL="1" ENV BUILD_AITER_ALL="1"
ENV BUILD_MOONCAKE="1" ENV BUILD_MOONCAKE="1"
ENV AITER_COMMIT_DEFAULT="4ad99832823dde2315b361cbd3b54b1c5c12acd5" ENV AITER_COMMIT_DEFAULT="acf8fdf9307431ece8ee275971c41cb3d1a7020b"
# =============================== # ===============================
# Base image 950 with rocm720 and args # Base image 950 with rocm720 and args
@@ -116,7 +116,7 @@ ENV BUILD_TRITON="1"
ENV BUILD_LLVM="0" ENV BUILD_LLVM="0"
ENV BUILD_AITER_ALL="1" ENV BUILD_AITER_ALL="1"
ENV BUILD_MOONCAKE="1" ENV BUILD_MOONCAKE="1"
ENV AITER_COMMIT_DEFAULT="4ad99832823dde2315b361cbd3b54b1c5c12acd5" ENV AITER_COMMIT_DEFAULT="acf8fdf9307431ece8ee275971c41cb3d1a7020b"
ENV TRITON_COMMIT_DEFAULT="42270451990532c67e69d753fbd026f28fcc4840" ENV TRITON_COMMIT_DEFAULT="42270451990532c67e69d753fbd026f28fcc4840"
# =============================== # ===============================
@@ -127,7 +127,7 @@ ENV BUILD_TRITON="1"
ENV BUILD_LLVM="0" ENV BUILD_LLVM="0"
ENV BUILD_AITER_ALL="1" ENV BUILD_AITER_ALL="1"
ENV BUILD_MOONCAKE="1" ENV BUILD_MOONCAKE="1"
ENV AITER_COMMIT_DEFAULT="4ad99832823dde2315b361cbd3b54b1c5c12acd5" ENV AITER_COMMIT_DEFAULT="acf8fdf9307431ece8ee275971c41cb3d1a7020b"
# Pin the ROCm torch stack for every pip invocation in this flavor. The file is # Pin the ROCm torch stack for every pip invocation in this flavor. The file is
# filled in after the torch 2.11 upgrade below; it must already exist (empty is # filled in after the torch 2.11 upgrade below; it must already exist (empty is
# valid) because pip reads PIP_CONSTRAINT from the first pip call onwards. # valid) because pip reads PIP_CONSTRAINT from the first pip call onwards.
@@ -286,7 +286,7 @@ ENV BUILD_TRITON="0"
ENV BUILD_LLVM="0" ENV BUILD_LLVM="0"
ENV BUILD_AITER_ALL="1" ENV BUILD_AITER_ALL="1"
ENV BUILD_MOONCAKE="1" ENV BUILD_MOONCAKE="1"
ENV AITER_COMMIT_DEFAULT="4ad99832823dde2315b361cbd3b54b1c5c12acd5" ENV AITER_COMMIT_DEFAULT="acf8fdf9307431ece8ee275971c41cb3d1a7020b"
# Same reasoning as the rocm724 stages: keep pip from resolving the image's # Same reasoning as the rocm724 stages: keep pip from resolving the image's
# ROCm torch away to a PyPI CUDA build. Populated after the stack is in place. # ROCm torch away to a PyPI CUDA build. Populated after the stack is in place.
ENV PIP_CONSTRAINT="/etc/sglang/constraints/torch-rocm.txt" ENV PIP_CONSTRAINT="/etc/sglang/constraints/torch-rocm.txt"
@@ -300,7 +300,7 @@ ENV BUILD_TRITON="0"
ENV BUILD_LLVM="0" ENV BUILD_LLVM="0"
ENV BUILD_AITER_ALL="1" ENV BUILD_AITER_ALL="1"
ENV BUILD_MOONCAKE="1" ENV BUILD_MOONCAKE="1"
ENV AITER_COMMIT_DEFAULT="4ad99832823dde2315b361cbd3b54b1c5c12acd5" ENV AITER_COMMIT_DEFAULT="acf8fdf9307431ece8ee275971c41cb3d1a7020b"
ENV PIP_CONSTRAINT="/etc/sglang/constraints/torch-rocm.txt" ENV PIP_CONSTRAINT="/etc/sglang/constraints/torch-rocm.txt"
RUN mkdir -p /etc/sglang/constraints && : > /etc/sglang/constraints/torch-rocm.txt RUN mkdir -p /etc/sglang/constraints && : > /etc/sglang/constraints/torch-rocm.txt
@@ -10,10 +10,10 @@ tag: NEW
<Accordion title="Install SGLang"> <Accordion title="Install SGLang">
Use an SGLang build that includes GLM-5.3-Flash support. Use an SGLang build that includes GLM-5.3-Flash support (v0.5.20 or later).
```bash Command ```bash Command
docker pull lmsysorg/sglang:glm-5.3-flash docker pull lmsysorg/sglang:latest
``` ```
The deployment panel can render a complete `docker run` command for the selected hardware and options. See [Install SGLang with Docker](/docs/get-started/install#method-3-using-docker) for host setup. The deployment panel can render a complete `docker run` command for the selected hardware and options. See [Install SGLang with Docker](/docs/get-started/install#method-3-using-docker) for host setup.
@@ -104,7 +104,7 @@ The **Speculative** card in the Playground changes the algorithm without leaving
- **EAGLE / MTP 5-1-6** is exactly what Low Latency serves, so a Low Latency base starts on this chip. Pick it from a High Throughput base to keep that recipe's other settings and add the MTP head. - **EAGLE / MTP 5-1-6** is exactly what Low Latency serves, so a Low Latency base starts on this chip. Pick it from a High Throughput base to keep that recipe's other settings and add the MTP head.
- **Off (greedy)** strips the whole `--speculative-*` family, which is what High Throughput already starts from. - **Off (greedy)** strips the whole `--speculative-*` family, which is what High Throughput already starts from.
- **DFlash2** swaps the in-checkpoint MTP head for the trained block-diffusion draft in [`incoai/GLM-5.3-Flash-DFlash2`](https://huggingface.co/incoai/GLM-5.3-Flash-DFlash2). The draft proposes a whole block per step and the target verifies it in one forward pass, so output quality stays the target's. Its block size comes from the draft checkpoint, and the draft runs on `fa4` rather than the target's DSA backends. It needs a build that carries the GLM-5.3-Flash hidden-state capture from [PR #36708](https://github.com/sgl-project/sglang/pull/36708), which is merged into the [PR #36507](https://github.com/sgl-project/sglang/pull/36507) support branch (`xinyuan/glm-5.3-flash-support`) rather than into `main`, so the image pinned above is not enough on its own — pull that branch at its current head, or add #36708's commit on top of an older checkout. The draft repository is also access-gated: request access on its model page, then download it alongside the target before serving. This combination is not yet measured on the cookbook hardware, so treat it as a starting point. - **DFlash2** swaps the in-checkpoint MTP head for the trained block-diffusion draft in [`incoai/GLM-5.3-Flash-DFlash2`](https://huggingface.co/incoai/GLM-5.3-Flash-DFlash2). The draft proposes a whole block per step and the target verifies it in one forward pass, so output quality stays the target's. Its block size comes from the draft checkpoint, and the draft runs on `fa4` rather than the target's DSA backends. The hidden-state capture it needs ([PR #36708](https://github.com/sgl-project/sglang/pull/36708)) shipped with the GLM-5.3-Flash support in v0.5.20, so the image pinned above is enough. The draft repository is access-gated: request access on its model page, then download it alongside the target before serving. This combination is not yet measured on the cookbook hardware, so treat it as a starting point.
Neither algorithm runs with DP-Attention; the card disables the affected chips and names the reason. Neither algorithm runs with DP-Attention; the card disables the affected chips and names the reason.
@@ -134,13 +134,13 @@ The default multimodal feature transport is automatic, and on a single CUDA node
### 3.1 Reasoning ### 3.1 Reasoning
Thinking is enabled by the checkpoint's generation configuration, and generated commands enable `--reasoning-parser glm45` by default. The OpenAI-compatible API then places thinking in `message.reasoning_content` and the final answer in `message.content`. You can disable **Reasoning Parser** in the Playground when an integration needs the raw response format. Thinking is enabled by the checkpoint's generation configuration, and generated commands enable `--reasoning-parser auto` (which resolves to `glm45` for GLM-5.3-Flash) by default. The OpenAI-compatible API then places thinking in `message.reasoning_content` and the final answer in `message.content`. You can disable **Reasoning Parser** in the Playground when an integration needs the raw response format.
To disable thinking for a request, pass `chat_template_kwargs: {"thinking": false}` in the request body. To disable thinking for a request, pass `chat_template_kwargs: {"thinking": false}` in the request body.
### 3.2 Tool calling ### 3.2 Tool calling
Generated commands enable `--tool-call-parser glm47` by default, so structured calls are returned in `message.tool_calls`. You can disable **Tool Call Parser** in the Playground when tool calling is not needed. On follow-up turns, read both `reasoning_content` and `content` because a thinking model can use either field around tool execution. Generated commands enable `--tool-call-parser auto` (which resolves to `glm47` for GLM-5.3-Flash) by default, so structured calls are returned in `message.tool_calls`. You can disable **Tool Call Parser** in the Playground when tool calling is not needed. On follow-up turns, read both `reasoning_content` and `content` because a thinking model can use either field around tool execution.
### 3.3 Multimodal serving ### 3.3 Multimodal serving
+4 -4
View File
@@ -103,7 +103,7 @@ import { Playground } from "/src/snippets/_playground.jsx";
- **DeepSeek Sparse Attention (DSA).** GLM-5.3 uses the `glm_moe_dsa` architecture; SGLang auto-selects the DSA attention backends (`flashmla_sparse` prefill, `fa3` decode, `sgl-kernel` indexer topk). No attention-backend flag is needed on the supported hardware. SGLang also auto-selects the KV-cache dtype for DSA models — `fp8_e4m3` on Blackwell (B200/GB300/B300, which then routes DSA through the TensorRT-LLM backend) and `bf16` on Hopper (H200) — so no `--kv-cache-dtype` flag is required. On Hopper, pairing `--kv-cache-dtype fp8_e4m3` with `--dsa-prefill-backend flashmla_sparse_q8 --dsa-decode-backend flashmla_kv` selects the native FP8 sparse prefill kernel (computes directly on the fp8 KV cache with no fp8→bf16 dequantization round-trip; GLM-5.3's 64 query heads match the kernel's native tile) — see the [DeepSeek-V3.2 page](../DeepSeek/DeepSeek-V3_2) for kernel details; the optional `SGLANG_ENABLE_DSA_Q8KV8_*` performance env vars are documented in `python/sglang/srt/environ.py`. - **DeepSeek Sparse Attention (DSA).** GLM-5.3 uses the `glm_moe_dsa` architecture; SGLang auto-selects the DSA attention backends (`flashmla_sparse` prefill, `fa3` decode, `sgl-kernel` indexer topk). No attention-backend flag is needed on the supported hardware. SGLang also auto-selects the KV-cache dtype for DSA models — `fp8_e4m3` on Blackwell (B200/GB300/B300, which then routes DSA through the TensorRT-LLM backend) and `bf16` on Hopper (H200) — so no `--kv-cache-dtype` flag is required. On Hopper, pairing `--kv-cache-dtype fp8_e4m3` with `--dsa-prefill-backend flashmla_sparse_q8 --dsa-decode-backend flashmla_kv` selects the native FP8 sparse prefill kernel (computes directly on the fp8 KV cache with no fp8→bf16 dequantization round-trip; GLM-5.3's 64 query heads match the kernel's native tile) — see the [DeepSeek-V3.2 page](../DeepSeek/DeepSeek-V3_2) for kernel details; the optional `SGLANG_ENABLE_DSA_Q8KV8_*` performance env vars are documented in `python/sglang/srt/environ.py`.
- **MTP / speculative decoding.** The checkpoint ships one nextn layer. Enable EAGLE MTP for lower latency (`--speculative-algorithm EAGLE --speculative-num-steps 5 --speculative-eagle-topk 1 --speculative-num-draft-tokens 6` for low-latency; `1-1-2` for balanced). The config's `index_share_for_mtp_iteration` reuses the DSA indexer's topk across draft steps (effective only at `--speculative-eagle-topk 1`). Watch the server's reported **accept length** and adjust `--speculative-num-steps` / `--speculative-num-draft-tokens`: lower the draft length when rejected draft tokens create excess verification work. - **MTP / speculative decoding.** The checkpoint ships one nextn layer. Enable EAGLE MTP for lower latency (`--speculative-algorithm EAGLE --speculative-num-steps 5 --speculative-eagle-topk 1 --speculative-num-draft-tokens 6` for low-latency; `1-1-2` for balanced). The config's `index_share_for_mtp_iteration` reuses the DSA indexer's topk across draft steps (effective only at `--speculative-eagle-topk 1`). Watch the server's reported **accept length** and adjust `--speculative-num-steps` / `--speculative-num-draft-tokens`: lower the draft length when rejected draft tokens create excess verification work.
- **DFlash2 (block-diffusion draft).** The **Speculative** card in the [Playground above](#playground) also offers **DFlash2**, which replaces the in-checkpoint MTP layer with the separately trained block-diffusion drafter [`incoai/GLM-5.3-DFlash2`](https://huggingface.co/incoai/GLM-5.3-DFlash2). It proposes a whole block per step and the target verifies the block in one forward pass, so output quality stays the target's. The block size — 8, i.e. 7 draft tokens per verification step — comes from the draft checkpoint's own `dflash_config`, so no `--speculative-num-draft-tokens` is passed; the draft is a small dense model and runs on `fa4` instead of the target's DSA backends. Two prerequisites: the DFlash2 drafter ([PR #35371](https://github.com/sgl-project/sglang/pull/35371)) merged **after v0.5.18**, so install SGLang from `main` (or use a nightly image) rather than the release this page pins; and DFLASH runs on **CUDA/NPU only** and rejects **DP-Attention**, so turn DP-Attention off in the **Attention** card before selecting it on a high-throughput base. The draft repository is public but licensed CC BY-NC-ND 4.0 for research and evaluation. - **DFlash2 (block-diffusion draft).** The **Speculative** card in the [Playground above](#playground) also offers **DFlash2**, which replaces the in-checkpoint MTP layer with the separately trained block-diffusion drafter [`incoai/GLM-5.3-DFlash2`](https://huggingface.co/incoai/GLM-5.3-DFlash2). It proposes a whole block per step and the target verifies the block in one forward pass, so output quality stays the target's. The block size — 8, i.e. 7 draft tokens per verification step — comes from the draft checkpoint's own `dflash_config`, so no `--speculative-num-draft-tokens` is passed; the draft is a small dense model and runs on `fa4` instead of the target's DSA backends. Note that DFLASH runs on **CUDA/NPU only** and rejects **DP-Attention**, so turn DP-Attention off in the **Attention** card before selecting it on a high-throughput base. The draft repository is public but licensed CC BY-NC-ND 4.0 for research and evaluation.
- **Memory.** The FP8 weights are large (MoE total, not active params). Start around `--mem-fraction-static 0.8` on H200 (TP8) and tune up; raise it for the 4-GPU GB300 single-node layout (TP4). - **Memory.** The FP8 weights are large (MoE total, not active params). Start around `--mem-fraction-static 0.8` on H200 (TP8) and tune up; raise it for the 4-GPU GB300 single-node layout (TP4).
- **DP-Attention + DeepEP** for the balanced/high-throughput strategies spreads attention across data-parallel ranks and routes MoE through DeepEP. - **DP-Attention + DeepEP** for the balanced/high-throughput strategies spreads attention across data-parallel ranks and routes MoE through DeepEP.
- **BF16 weights need more GPUs.** The full-precision build (`zai-org/GLM-5.3-BF16`, ~1.5 TB) does not fit a single 8×H200 / 8×B200 / 4×GB300 node. It fits single-node on **8×B300** (TP8, ~2.1 TB HBM); on the smaller GPUs it needs a **multi-node** layout (e.g. 2×8×H200 or 2×8×B200 at TP16, 2×4×GB300 at TP8). FP8 is the recommended deployment. Use the same DSA / MTP / chunked-prefill guidance as FP8. - **BF16 weights need more GPUs.** The full-precision build (`zai-org/GLM-5.3-BF16`, ~1.5 TB) does not fit a single 8×H200 / 8×B200 / 4×GB300 node. It fits single-node on **8×B300** (TP8, ~2.1 TB HBM); on the smaller GPUs it needs a **multi-node** layout (e.g. 2×8×H200 or 2×8×B200 at TP16, 2×4×GB300 at TP8). FP8 is the recommended deployment. Use the same DSA / MTP / chunked-prefill guidance as FP8.
@@ -117,7 +117,7 @@ import { Playground } from "/src/snippets/_playground.jsx";
### 3.1 Reasoning ### 3.1 Reasoning
GLM-5.3 is a reasoning model. Enable the `glm45` reasoning parser (toggle **Reasoning Parser** in the **Parsers** card of the [Playground above](#playground)) to separate thinking from the final answer — thinking lands in `message.reasoning_content`, the answer in `message.content`. The chat template defaults `clear_thinking` to `false`; for multi-turn chat, pass `chat_template_kwargs: {"clear_thinking": True}` so previous reasoning is cleared before the next response. GLM-5.3 is a reasoning model, and generated commands enable `--reasoning-parser auto` (which resolves to `glm45` for GLM-5.3) by default so thinking is separated from the final answer — thinking lands in `message.reasoning_content`, the answer in `message.content`. Without the parser the server returns the thinking and the answer as one `content` string with a stray `</think>` between them, because the chat template opens `<think>` in the generation prompt. You can disable **Reasoning Parser** in the **Parsers** card of the [Playground above](#playground) when an integration needs that raw format. The chat template defaults `clear_thinking` to `false`; for multi-turn chat, pass `chat_template_kwargs: {"clear_thinking": True}` so previous reasoning is cleared before the next response.
**Reasoning effort.** Pass `chat_template_kwargs: {"reasoning_effort": ...}` to select `low`, `high`, or `max`. If you omit it or pass another value, the template uses `max`. **Reasoning effort.** Pass `chat_template_kwargs: {"reasoning_effort": ...}` to select `low`, `high`, or `max`. If you omit it or pass another value, the template uses `max`.
@@ -164,7 +164,7 @@ Here is how you can calculate it:
### 3.2 Tool Calling ### 3.2 Tool Calling
Enable the `glm47` tool-call parser (toggle **Tool Call Parser** in the **Parsers** card of the [Playground above](#playground)) to surface structured tool calls via `message.tool_calls`. GLM-5.3 emits the newer `<tool_call>…<arg_key>…<arg_value>…` format, so it needs the **`glm47`** parser — the older `glm45` parser does not parse it (the call would be left as raw text in `content`). On thinking mode the turn also fills `reasoning_content`, so print both fields. Generated commands enable `--tool-call-parser auto` by default, so structured calls are returned in `message.tool_calls` with `finish_reason: "tool_calls"`. `auto` resolves to **`glm47`** for GLM-5.3: the model emits the newer `<tool_call>…<arg_key>…<arg_value>…` format, which the older `glm45` parser does not parse (the call would be left as raw text in `content`). Running with no tool-call parser fails the same way, and `finish_reason` stays `"stop"`, so an agent loop never sees the call. You can disable **Tool Call Parser** in the **Parsers** card of the [Playground above](#playground) when tool calling is not needed. On thinking mode the turn also fills `reasoning_content`, so print both fields.
<Accordion title="Tool Calling Example (Python)"> <Accordion title="Tool Calling Example (Python)">
@@ -218,7 +218,7 @@ For long-context, prefix-heavy workloads, enable hierarchical KV caching to spil
### 3.4 Claude Code Integration ### 3.4 Claude Code Integration
GLM-5.3's strong reasoning + tool-calling makes it a good backend for [Claude Code](https://code.claude.com/docs/en/overview), Anthropic's agentic CLI. SGLang exposes the Anthropic-compatible `/v1/messages` endpoint on every server, so Claude Code can talk to a GLM-5.3 server with only environment variables — no code change. Launch the server with `--reasoning-parser glm45 --tool-call-parser glm47` (any recipe from the Deployment panel above works), then: GLM-5.3's strong reasoning + tool-calling makes it a good backend for [Claude Code](https://code.claude.com/docs/en/overview), Anthropic's agentic CLI. SGLang exposes the Anthropic-compatible `/v1/messages` endpoint on every server, so Claude Code can talk to a GLM-5.3 server with only environment variables — no code change. Launch the server with `--reasoning-parser auto --tool-call-parser auto` (any recipe from the Deployment panel above works), then:
```bash Command ```bash Command
export ANTHROPIC_BASE_URL="http://127.0.0.1:30000" export ANTHROPIC_BASE_URL="http://127.0.0.1:30000"
@@ -129,6 +129,7 @@ The NVIDIA Blackwell recipes are validated single-node: **B200 at `--tp 8`** and
- **Memory**: `--mem-fraction-static` reserves GPU memory for weights + KV pool; the rest is prefill **activation headroom**. The value scales with *free* memory per GPU (card capacity minus per-GPU weight), so it tracks the card more than the TP degree: **`0.65` on B200** (180 GB — less headroom once weights are resident) and **`0.75` on the larger-memory B300 / GB300** (`0.80` on AMD). Lower TP packs more weight per GPU, so a tighter config needs a *lower* value — B200 needs `0.65` even at `--tp 4`. Raising it past the validated value is fine only for low-concurrency single-stream serving; it OOMs under high concurrency or long context. - **Memory**: `--mem-fraction-static` reserves GPU memory for weights + KV pool; the rest is prefill **activation headroom**. The value scales with *free* memory per GPU (card capacity minus per-GPU weight), so it tracks the card more than the TP degree: **`0.65` on B200** (180 GB — less headroom once weights are resident) and **`0.75` on the larger-memory B300 / GB300** (`0.80` on AMD). Lower TP packs more weight per GPU, so a tighter config needs a *lower* value — B200 needs `0.65` even at `--tp 4`. Raising it past the validated value is fine only for low-concurrency single-stream serving; it OOMs under high concurrency or long context.
- **Long context (32K+)**: keep `--mem-fraction-static` at the platform default and raise `--chunked-prefill-size` to `16384`. Decode TPOT stays roughly flat in context length thanks to sparse attention; 1K128K prompts are validated. - **Long context (32K+)**: keep `--mem-fraction-static` at the platform default and raise `--chunked-prefill-size` to `16384`. Decode TPOT stays roughly flat in context length thanks to sparse attention; 1K128K prompts are validated.
- **HiSparse for decode capacity**: on NVIDIA CUDA, HiSparse keeps the three dense layers on GPU, moves the 57 sparse-layer K/V caches to pinned host memory, and feeds selected block IDs directly to the swap-in kernel. For the released four-KV-head model, use `--tp 4` or greater, `--disable-radix-cache`, and `device_buffer_size >= 2048`. Enable it with `--enable-hisparse --hisparse-config='{"device_buffer_size":4096,"host_to_device_ratio":2}'` on the Triton launch command.
- **Scaling TP**: B200 is documented at `--tp 8`; B300 / GB200 / GB300 at `--tp 4` (the single-node cross-family common denominator). On an 8-GPU B300 host you can also raise to `--tp 8` for more throughput / KV headroom. - **Scaling TP**: B200 is documented at `--tp 8`; B300 / GB200 / GB300 at `--tp 4` (the single-node cross-family common denominator). On an 8-GPU B300 host you can also raise to `--tp 8` for more throughput / KV headroom.
- **Expert parallelism**: to trade latency for throughput add `--ep` (see [Expert Parallelism Deployment](../../../docs/advanced_features/expert_parallelism)). On AMD, set `--ep` equal to `--tp`. Shared-experts fusion is automatically disabled when EP > 1; on AMD standard EP the server also disables `--enable-aiter-allreduce-fusion` automatically to preserve accuracy. - **Expert parallelism**: to trade latency for throughput add `--ep` (see [Expert Parallelism Deployment](../../../docs/advanced_features/expert_parallelism)). On AMD, set `--ep` equal to `--tp`. Shared-experts fusion is automatically disabled when EP > 1; on AMD standard EP the server also disables `--enable-aiter-allreduce-fusion` automatically to preserve accuracy.
- `--trust-remote-code` is required to load the MiniMax config / processor classes. - `--trust-remote-code` is required to load the MiniMax config / processor classes.
@@ -30,7 +30,7 @@ Then run the **Python** output of the command panel below in that environment.
```bash Command ```bash Command
docker pull lmsysorg/sglang:latest # NVIDIA (CUDA) docker pull lmsysorg/sglang:latest # NVIDIA (CUDA)
docker pull lmsysorg/sglang-rocm:v0.5.19-rocm720-mi35x-20260910 # AMD MI350X / MI355X (ROCm) docker pull lmsysorg/sglang-rocm:v0.5.19-rocm720-mi35x-20260916 # AMD MI350X / MI355X (ROCm)
``` ```
For how to launch the image, see [Install → Method 3: Using Docker](../../../docs/get-started/install#method-3-using-docker). Substitute the inner `sglang serve ...` with what the command generator below produces. For how to launch the image, see [Install → Method 3: Using Docker](../../../docs/get-started/install#method-3-using-docker). Substitute the inner `sglang serve ...` with what the command generator below produces.
@@ -65,7 +65,7 @@ Pick your hardware, then the deployment shape and operating point. Node count fo
**Strategy** — the operating point within that shape: **Strategy** — the operating point within that shape:
- **Low-Latency** — no DCP, so the MLA KV stays TP-replicated. For chat. B200 splits its two nodes into PP2 × TP8; every other platform is flat TP. - **Low-Latency** — no DCP, so the MLA KV stays TP-replicated. For chat. B200 splits its two nodes into PP2 × TP8; every other platform is flat TP.
- **Balanced** — the accuracy-preserving default: PP2 × DCPEP8 on B200 (the two pipeline stages and DCP8 split KV and KDA state), TP16/DCP16 on GB200, TP8/DCP8 on B300/GB300, TP8 ROCm/AITER on MI35x. - **Balanced** — the accuracy-preserving default: PP2 × DCPEP8 on B200 (the two pipeline stages and DCP8 split KV and KDA state), TP16/DCP16 on GB200, TP8/DCP8 on B300/GB300, TP8/DCP8 ROCm/AITER on MI35x.
- **High-Throughput** — the large-scale lane: pick a **Cluster Size** and **Large-Scale Preset** in the Playground ([details](#large-scale-presets)). The cell itself is Balanced, except on H100 (plus `extra_buffer_lazy`) and H200 (widens to 4×8 TP32/EP32 at `--mem-fraction-static 0.90`). - **High-Throughput** — the large-scale lane: pick a **Cluster Size** and **Large-Scale Preset** in the Playground ([details](#large-scale-presets)). The cell itself is Balanced, except on H100 (plus `extra_buffer_lazy`) and H200 (widens to 4×8 TP32/EP32 at `--mem-fraction-static 0.90`).
`Long-Context` appears only under the `Prefill` PD mode; for long-context unified serving on B200, start from High-Throughput and raise `--context-length`. `Long-Context` appears only under the `Prefill` PD mode; for long-context unified serving on B200, start from High-Throughput and raise `--context-length`.
@@ -93,6 +93,22 @@ import { KimiK3MambaRatioCalculator } from "/src/snippets/_kimi_k3_mamba_ratio_c
NVFP4 NOSPEC / NVFP4 DSPARK), which is why no point past concurrency 64 is published for Balanced. NVFP4 NOSPEC / NVFP4 DSPARK), which is why no point past concurrency 64 is published for Balanced.
</Note> </Note>
### AMD AITER with DCP8
The MI350X/MI355X unified Balanced recipe uses TP8/DCP8 with AITER prefill and
decode attention. DCP shards the target MLA KV cache; RadixArk DSPARK's draft KV
remains replicated. The pinned `v0.5.19-rocm720-mi35x-20260916` image records
SGLang revision `e7f7447333`, which includes
[AITER DCP support (#34432)](https://github.com/sgl-project/sglang/pull/34432) and
the [DCP KV-free fix (#38941)](https://github.com/sgl-project/sglang/pull/38941).
No source overlay is required for DCP.
Keep `SGLANG_K3_KDA_FUSED_BACKEND` unset with this image. The separate fused-KDA
opt-in requires the [deferred-gate fix (#39066)](https://github.com/sgl-project/sglang/pull/39066),
which is not included in this image. This updated recipe remains **Final
Verification In Progress**; the recorded speed numbers use their original
configurations and do not validate the new image or DCP8 recipe.
### Mamba ratio calculator ### Mamba ratio calculator
<KimiK3MambaRatioCalculator /> <KimiK3MambaRatioCalculator />
@@ -180,11 +196,11 @@ Speculation: DSPARK holds block size + 1 (= 8) intermediate states per request
| GB200 4×4 | TP16/DCP16 | MNNVL auto-detected | | GB200 4×4 | TP16/DCP16 | MNNVL auto-detected |
| H200 2×8 (4×8 on Unified High-Throughput) | TP16/EP16 + symm-mem, Marlin + FlashMLA; High-Throughput widens to TP32/EP32 over 4 nodes at mem-frac 0.90 with `extra_buffer_lazy` | same block on every node; export the cross-node NIC (`GLOO_SOCKET_IFNAME` / `NCCL_SOCKET_IFNAME`, `SGLANG_HOST_IP`); keep `NCCL_MNNVL_ENABLE=1 NCCL_CUMEM_ENABLE=1` | | H200 2×8 (4×8 on Unified High-Throughput) | TP16/EP16 + symm-mem, Marlin + FlashMLA; High-Throughput widens to TP32/EP32 over 4 nodes at mem-frac 0.90 with `extra_buffer_lazy` | same block on every node; export the cross-node NIC (`GLOO_SOCKET_IFNAME` / `NCCL_SOCKET_IFNAME`, `SGLANG_HOST_IP`); keep `NCCL_MNNVL_ENABLE=1 NCCL_CUMEM_ENABLE=1` |
| H100 4×8 | TP32/EP32, Marlin + FlashMLA | SM90a build of the K3 image; pin NCCL/Gloo to the same NIC on all nodes; least post-weight headroom (80 GB) | | H100 4×8 | TP32/EP32, Marlin + FlashMLA | SM90a build of the K3 image; pin NCCL/Gloo to the same NIC on all nodes; least post-weight headroom (80 GB) |
| MI350X/MI355X 1×8 | TP8 ROCm/AITER | AITER A8W4 FlyDSL MoE, Triton attention (`SGLANG_MLA_DECODE_TUNE=1` for gfx950 MLA decode geometry), graph bs up to 256, fp8 kvcache; DSPARK supported. Activation-quant and fused-KDA-decode knobs: [AMD ROCm/AITER environment](#amd-env) | | MI350X/MI355X 1×8 | TP8/DCP8 ROCm/AITER (Unified Balanced) | AITER A8W4 FlyDSL MoE, AITER prefill/decode attention with sharded target MLA KV, graph bs up to 256, fp8 kvcache; DSPARK supported. Activation-quant and fused-KDA-decode knobs: [AMD ROCm/AITER environment](#amd-env) |
| Ascend A3 Series 4×8 (32 cards / 64 dies) | TP64/DP4 + DeepEP | PD-mixed `Unified` only; DSPARK baked in; pin `GLOO`/`HCCL_SOCKET_IFNAME` on every node | | Ascend A3 Series 4×8 (32 cards / 64 dies) | TP64/DP4 + DeepEP | PD-mixed `Unified` only; DSPARK baked in; pin `GLOO`/`HCCL_SOCKET_IFNAME` on every node |
| Ascend 950PR/DT Series 4×8 | TP32/dp1 + DeepEP | PD-mixed `Unified` only; DSPARK baked in; shared experts / dense MLP shard over attention-TP (`--shared-experts-tp-size 4`); radix cache off; pin `GLOO`/`HCCL_SOCKET_IFNAME` on every node | | Ascend 950PR/DT Series 4×8 | TP32/dp1 + DeepEP | PD-mixed `Unified` only; DSPARK baked in; shared experts / dense MLP shard over attention-TP (`--shared-experts-tp-size 4`); radix cache off; pin `GLOO`/`HCCL_SOCKET_IFNAME` on every node |
**DCP notes** — the DCP cells are Balanced and High-Throughput on every Blackwell platform, in both the `Unified` and `Decode` roles: **Blackwell DCP notes** — the DCP cells are Balanced and High-Throughput on every Blackwell platform, in both the `Unified` and `Decode` roles:
- DCP is the only axis that shards the TP-replicated MLA KV; Low-Latency skips it. - DCP is the only axis that shards the TP-replicated MLA KV; Low-Latency skips it.
- Leave `--dcp-comm-backend` unset (fabric-resolved: `fi_a2a` on GB200/GB300, `a2a` on B200/B300). - Leave `--dcp-comm-backend` unset (fabric-resolved: `fi_a2a` on GB200/GB300, `a2a` on B200/B300).
+10 -1
View File
@@ -6,7 +6,7 @@ metatags:
HiSparse reduces per-request GPU memory consumption during the decode phase by maintaining only a small "hot" KV buffer on GPU while keeping complete KV data in CPU pinned memory. Combined with PD disaggregation, it enables significantly higher decode concurrency. HiSparse reduces per-request GPU memory consumption during the decode phase by maintaining only a small "hot" KV buffer on GPU while keeping complete KV data in CPU pinned memory. Combined with PD disaggregation, it enables significantly higher decode concurrency.
> **Prerequisites**: HiSparse works with models that use **DeepSeek Sparse Attention (DSA)** architectures (e.g., DeepSeek-V3.2, GLM-5.1) and **DeepSeek V4**. These models natively select a subset of tokens for attention, making it possible to keep only the top-k KV on GPU while storing the full KV in host memory — without accuracy loss. Additionally, HiSparse currently requires **PD disaggregation mode** and is enabled on the **decode instance** only. > **Prerequisites**: HiSparse works with models that use **DeepSeek Sparse Attention (DSA)** architectures (e.g., DeepSeek-V3.2, GLM-5.1), **DeepSeek V4**, and **MiniMax M3**. These models natively select a subset of tokens for attention, making it possible to keep only the top-k KV on GPU while storing the full KV in host memory — without accuracy loss. Additionally, HiSparse currently requires **PD disaggregation mode** and is enabled on the **decode instance** only.
## Why HiSparse? ## Why HiSparse?
@@ -165,6 +165,15 @@ python3 -m sglang.launch_server \
> **Note**: For DSA models, `--kv-cache-dtype` defaults to `auto`, which resolves to `fp8_e4m3` on SM100+ (Blackwell) and `bfloat16` on older architectures. The DSA decode backend is automatically selected based on KV dtype (`bfloat16` → `flashmla_sparse`, `fp8_e4m3` → `flashmla_kv`), except for GLM DSA models on SM120/SM121 with `fp8_e4m3`, which use `flashinfer_sparse_mla`. DSA backend flags apply only to DSA models; DeepSeek V4 uses its own `dsv4` attention backend. > **Note**: For DSA models, `--kv-cache-dtype` defaults to `auto`, which resolves to `fp8_e4m3` on SM100+ (Blackwell) and `bfloat16` on older architectures. The DSA decode backend is automatically selected based on KV dtype (`bfloat16` → `flashmla_sparse`, `fp8_e4m3` → `flashmla_kv`), except for GLM DSA models on SM120/SM121 with `fp8_e4m3`, which use `flashinfer_sparse_mla`. DSA backend flags apply only to DSA models; DeepSeek V4 uses its own `dsv4` attention backend.
### MiniMax M3
Dense-layer K/V and index K stay on GPU; sparse-layer K/V use host memory plus a GPU working set.
- Use TP 4 or greater, with the same TP size and PP 1 on both PD instances.
- Use `--attention-backend triton`, `--mm-attention-backend triton_attn`, `--disable-prefill-cuda-graph`, and `--disable-radix-cache`.
- Set `device_buffer_size` to at least 2048 in `--hisparse-config`; `top_k` does not override the model's selection width.
- PD retraction backup is not supported. Use `--num-reserved-decode-tokens` to reserve capacity for the expected output length.
### Benchmark ### Benchmark
```bash Command ```bash Command
@@ -478,8 +478,8 @@ export const config = {
gb200: "lmsysorg/sglang:kimi-k3", gb200: "lmsysorg/sglang:kimi-k3",
// 20260903 or newer: the AITER SiTU A4W4/A8W4 layout fix (sgl-project/sglang#33838, // 20260903 or newer: the AITER SiTU A4W4/A8W4 layout fix (sgl-project/sglang#33838,
// merged Sep 3) and the fused gfx950 KDA decode boundary (#34198) first ship here. // merged Sep 3) and the fused gfx950 KDA decode boundary (#34198) first ship here.
mi350x: "lmsysorg/sglang-rocm:v0.5.19-rocm720-mi35x-20260910", mi350x: "lmsysorg/sglang-rocm:v0.5.19-rocm720-mi35x-20260916",
mi355x: "lmsysorg/sglang-rocm:v0.5.19-rocm720-mi35x-20260910", mi355x: "lmsysorg/sglang-rocm:v0.5.19-rocm720-mi35x-20260916",
// NVFP4 needs a build with sgl-project/sglang#35077; the purpose-built dev // NVFP4 needs a build with sgl-project/sglang#35077; the purpose-built dev
// image is cut from that PR's head (CUDA 13). // image is cut from that PR's head (CUDA 13).
"b300|nvfp4": "lmsysorg/sglang:dev-dev-kimi-k3-nvfp4", "b300|nvfp4": "lmsysorg/sglang:dev-dev-kimi-k3-nvfp4",
@@ -1217,7 +1217,7 @@ export const config = {
], ],
}, },
{ {
// MI350X and MI355X use the same single-node TP8 ROCm/AITER profile. // MI350X and MI355X use the same single-node TP8/DCP8 ROCm/AITER profile.
match: { hw: "mi350x", pdMode: "unified", strategy: "balanced" }, match: { hw: "mi350x", pdMode: "unified", strategy: "balanced" },
nnodes: 1, nnodes: 1,
verified: false, verified: false,
@@ -1232,7 +1232,10 @@ export const config = {
"--model-path {{MODEL_NAME}}", "--model-path {{MODEL_NAME}}",
"--trust-remote-code", "--trust-remote-code",
"--tp-size 8", "--tp-size 8",
"--attention-backend triton", "--dcp-size 8",
"--dcp-comm-backend a2a",
"--prefill-attention-backend aiter",
"--decode-attention-backend aiter",
"--kv-cache-dtype fp8_e4m3", "--kv-cache-dtype fp8_e4m3",
"--dtype bfloat16", "--dtype bfloat16",
"--mem-fraction-static 0.85", "--mem-fraction-static 0.85",
@@ -1259,7 +1262,10 @@ export const config = {
"--model-path {{MODEL_NAME}}", "--model-path {{MODEL_NAME}}",
"--trust-remote-code", "--trust-remote-code",
"--tp-size 8", "--tp-size 8",
"--attention-backend triton", "--dcp-size 8",
"--dcp-comm-backend a2a",
"--prefill-attention-backend aiter",
"--decode-attention-backend aiter",
"--kv-cache-dtype fp8_e4m3", "--kv-cache-dtype fp8_e4m3",
"--dtype bfloat16", "--dtype bfloat16",
"--mem-fraction-static 0.85", "--mem-fraction-static 0.85",
@@ -119,6 +119,32 @@ export const config = {
}, },
], ],
}, },
// Parser flags live in one overlay dim so every generated command gets
// them without per-cell duplication; the Parsers card toggles derive
// on/off from the composed flags. `auto` needs the GLM-5.3 template
// detection (v0.5.20+).
{
id: "parsers",
title: "Parsers",
default: "auto",
options: [
{
id: "auto",
label: "Auto (glm45 + glm47)",
stripPrefixes: ["--reasoning-parser", "--tool-call-parser"],
flags: [
"--reasoning-parser auto",
"--tool-call-parser auto",
],
},
{
id: "off",
label: "Off",
stripPrefixes: ["--reasoning-parser", "--tool-call-parser"],
flags: [],
},
],
},
], ],
modelNames: { modelNames: {
@@ -174,15 +200,16 @@ sgl-eval run gsm8k \\
["aime2026_pct", "AIME 2026", "%"], ["aime2026_pct", "AIME 2026", "%"],
], ],
// Support is not in a public sglang release yet, so the nightly images do // v0.5.20 (= latest) carries GLM-5.3-Flash support (#36507) and the GLM-5.3
// not work; every NVIDIA lane uses the purpose-built CUDA 13 image. // template parser detection (#38297) that `--*-parser auto` needs; the old
// glm-5.3-flash dev image (2026-09-03) predates #38297 and misdetects.
dockerImages: { dockerImages: {
gb300: "lmsysorg/sglang:glm-5.3-flash", gb300: "lmsysorg/sglang:latest",
h100: "lmsysorg/sglang:glm-5.3-flash", h100: "lmsysorg/sglang:latest",
h200: "lmsysorg/sglang:glm-5.3-flash", h200: "lmsysorg/sglang:latest",
b200: "lmsysorg/sglang:glm-5.3-flash", b200: "lmsysorg/sglang:latest",
b300: "lmsysorg/sglang:glm-5.3-flash", b300: "lmsysorg/sglang:latest",
gb200: "lmsysorg/sglang:glm-5.3-flash", gb200: "lmsysorg/sglang:latest",
}, },
github: { github: {
@@ -259,8 +286,8 @@ sgl-eval run gsm8k \\
parsers: { parsers: {
items: [ items: [
{ id: "reasoning", label: "Reasoning Parser", flag: "--reasoning-parser glm45" }, { id: "reasoning", label: "Reasoning Parser", flag: "--reasoning-parser auto" },
{ id: "toolCall", label: "Tool Call Parser", flag: "--tool-call-parser glm47" }, { id: "toolCall", label: "Tool Call Parser", flag: "--tool-call-parser auto" },
], ],
}, },
@@ -305,12 +332,7 @@ sgl-eval run gsm8k \\
"--speculative-draft-model-path incoai/GLM-5.3-Flash-DFlash2", "--speculative-draft-model-path incoai/GLM-5.3-Flash-DFlash2",
"--speculative-draft-attention-backend fa4", "--speculative-draft-attention-backend fa4",
], ],
// DFLASH needs this model's hidden-state capture, which landed on the note: "⚠️ The draft checkpoint incoai/GLM-5.3-Flash-DFlash2 is access-gated: request access on its Hugging Face page, then download it alongside the target before serving.",
// GLM-5.3-Flash support branch (PR #36708 into #36507's
// xinyuan/glm-5.3-flash-support), not on main — so it postdates the
// image the Install accordion pins. Drop this note once #36507 merges
// and a published image carries it.
note: "⚠️ Needs the GLM-5.3-Flash hidden-state capture from PR #36708. It is merged into the PR #36507 support branch (xinyuan/glm-5.3-flash-support), not into main, so pull that branch at its current head — or add #36708's commit on top of an older checkout — before serving. The lmsysorg/sglang:glm-5.3-flash image alone is not enough.",
disable: [ disable: [
{ {
when: { dpAttnOn: [true] }, when: { dpAttnOn: [true] },
@@ -346,8 +368,6 @@ sgl-eval run gsm8k \\
"--speculative-num-steps 5", "--speculative-num-steps 5",
"--speculative-eagle-topk 1", "--speculative-eagle-topk 1",
"--speculative-num-draft-tokens 6", "--speculative-num-draft-tokens 6",
"--reasoning-parser glm45",
"--tool-call-parser glm47",
"--host {{HOST_IP}}", "--host {{HOST_IP}}",
"--port {{PORT}}", "--port {{PORT}}",
], ],
@@ -371,8 +391,6 @@ sgl-eval run gsm8k \\
"--dsa-decode-backend trtllm", "--dsa-decode-backend trtllm",
"--kv-cache-dtype fp8_e4m3", "--kv-cache-dtype fp8_e4m3",
"--moe-runner-backend flashinfer_trtllm", "--moe-runner-backend flashinfer_trtllm",
"--reasoning-parser glm45",
"--tool-call-parser glm47",
"--host {{HOST_IP}}", "--host {{HOST_IP}}",
"--port {{PORT}}", "--port {{PORT}}",
], ],
@@ -406,8 +424,6 @@ sgl-eval run gsm8k \\
"--speculative-num-steps 5", "--speculative-num-steps 5",
"--speculative-eagle-topk 1", "--speculative-eagle-topk 1",
"--speculative-num-draft-tokens 6", "--speculative-num-draft-tokens 6",
"--reasoning-parser glm45",
"--tool-call-parser glm47",
"--mem-fraction-static 0.85", "--mem-fraction-static 0.85",
"--cuda-graph-max-bs-decode 32", "--cuda-graph-max-bs-decode 32",
"--host {{HOST_IP}}", "--host {{HOST_IP}}",
@@ -434,8 +450,6 @@ sgl-eval run gsm8k \\
"--dsa-decode-backend tilelang", "--dsa-decode-backend tilelang",
"--kv-cache-dtype bfloat16", "--kv-cache-dtype bfloat16",
"--moe-runner-backend flashinfer_cutlass", "--moe-runner-backend flashinfer_cutlass",
"--reasoning-parser glm45",
"--tool-call-parser glm47",
"--mem-fraction-static 0.85", "--mem-fraction-static 0.85",
"--host {{HOST_IP}}", "--host {{HOST_IP}}",
"--port {{PORT}}", "--port {{PORT}}",
@@ -461,8 +475,6 @@ sgl-eval run gsm8k \\
"--speculative-num-steps 5", "--speculative-num-steps 5",
"--speculative-eagle-topk 1", "--speculative-eagle-topk 1",
"--speculative-num-draft-tokens 6", "--speculative-num-draft-tokens 6",
"--reasoning-parser glm45",
"--tool-call-parser glm47",
"--mem-fraction-static 0.85", "--mem-fraction-static 0.85",
"--cuda-graph-max-bs-decode 32", "--cuda-graph-max-bs-decode 32",
"--host {{HOST_IP}}", "--host {{HOST_IP}}",
@@ -482,8 +494,6 @@ sgl-eval run gsm8k \\
"--dsa-decode-backend tilelang", "--dsa-decode-backend tilelang",
"--kv-cache-dtype bfloat16", "--kv-cache-dtype bfloat16",
"--moe-runner-backend flashinfer_cutlass", "--moe-runner-backend flashinfer_cutlass",
"--reasoning-parser glm45",
"--tool-call-parser glm47",
"--mem-fraction-static 0.85", "--mem-fraction-static 0.85",
"--host {{HOST_IP}}", "--host {{HOST_IP}}",
"--port {{PORT}}", "--port {{PORT}}",
@@ -506,8 +516,6 @@ sgl-eval run gsm8k \\
"--speculative-num-steps 5", "--speculative-num-steps 5",
"--speculative-eagle-topk 1", "--speculative-eagle-topk 1",
"--speculative-num-draft-tokens 6", "--speculative-num-draft-tokens 6",
"--reasoning-parser glm45",
"--tool-call-parser glm47",
"--mem-fraction-static 0.85", "--mem-fraction-static 0.85",
"--cuda-graph-max-bs-decode 32", "--cuda-graph-max-bs-decode 32",
"--host {{HOST_IP}}", "--host {{HOST_IP}}",
@@ -527,8 +535,6 @@ sgl-eval run gsm8k \\
"--dsa-decode-backend tilelang", "--dsa-decode-backend tilelang",
"--kv-cache-dtype bfloat16", "--kv-cache-dtype bfloat16",
"--moe-runner-backend flashinfer_cutlass", "--moe-runner-backend flashinfer_cutlass",
"--reasoning-parser glm45",
"--tool-call-parser glm47",
"--mem-fraction-static 0.85", "--mem-fraction-static 0.85",
"--host {{HOST_IP}}", "--host {{HOST_IP}}",
"--port {{PORT}}", "--port {{PORT}}",
@@ -551,8 +557,6 @@ sgl-eval run gsm8k \\
"--speculative-num-steps 5", "--speculative-num-steps 5",
"--speculative-eagle-topk 1", "--speculative-eagle-topk 1",
"--speculative-num-draft-tokens 6", "--speculative-num-draft-tokens 6",
"--reasoning-parser glm45",
"--tool-call-parser glm47",
"--mem-fraction-static 0.85", "--mem-fraction-static 0.85",
"--cuda-graph-max-bs-decode 32", "--cuda-graph-max-bs-decode 32",
"--host {{HOST_IP}}", "--host {{HOST_IP}}",
@@ -572,8 +576,6 @@ sgl-eval run gsm8k \\
"--dsa-decode-backend tilelang", "--dsa-decode-backend tilelang",
"--kv-cache-dtype bfloat16", "--kv-cache-dtype bfloat16",
"--moe-runner-backend flashinfer_cutlass", "--moe-runner-backend flashinfer_cutlass",
"--reasoning-parser glm45",
"--tool-call-parser glm47",
"--mem-fraction-static 0.85", "--mem-fraction-static 0.85",
"--host {{HOST_IP}}", "--host {{HOST_IP}}",
"--port {{PORT}}", "--port {{PORT}}",
@@ -603,8 +605,6 @@ sgl-eval run gsm8k \\
"--speculative-num-steps 5", "--speculative-num-steps 5",
"--speculative-eagle-topk 1", "--speculative-eagle-topk 1",
"--speculative-num-draft-tokens 6", "--speculative-num-draft-tokens 6",
"--reasoning-parser glm45",
"--tool-call-parser glm47",
"--host {{HOST_IP}}", "--host {{HOST_IP}}",
"--port {{PORT}}", "--port {{PORT}}",
], ],
@@ -626,8 +626,6 @@ sgl-eval run gsm8k \\
"--dsa-decode-backend tilelang", "--dsa-decode-backend tilelang",
"--kv-cache-dtype bfloat16", "--kv-cache-dtype bfloat16",
"--moe-runner-backend deep_gemm", "--moe-runner-backend deep_gemm",
"--reasoning-parser glm45",
"--tool-call-parser glm47",
"--host {{HOST_IP}}", "--host {{HOST_IP}}",
"--port {{PORT}}", "--port {{PORT}}",
], ],
@@ -655,8 +653,6 @@ sgl-eval run gsm8k \\
"--speculative-num-steps 5", "--speculative-num-steps 5",
"--speculative-eagle-topk 1", "--speculative-eagle-topk 1",
"--speculative-num-draft-tokens 6", "--speculative-num-draft-tokens 6",
"--reasoning-parser glm45",
"--tool-call-parser glm47",
"--host {{HOST_IP}}", "--host {{HOST_IP}}",
"--port {{PORT}}", "--port {{PORT}}",
], ],
@@ -677,8 +673,6 @@ sgl-eval run gsm8k \\
"--dsa-decode-backend tilelang", "--dsa-decode-backend tilelang",
"--kv-cache-dtype bfloat16", "--kv-cache-dtype bfloat16",
"--moe-runner-backend deep_gemm", "--moe-runner-backend deep_gemm",
"--reasoning-parser glm45",
"--tool-call-parser glm47",
"--host {{HOST_IP}}", "--host {{HOST_IP}}",
"--port {{PORT}}", "--port {{PORT}}",
], ],
@@ -701,8 +695,6 @@ sgl-eval run gsm8k \\
"--speculative-num-steps 5", "--speculative-num-steps 5",
"--speculative-eagle-topk 1", "--speculative-eagle-topk 1",
"--speculative-num-draft-tokens 6", "--speculative-num-draft-tokens 6",
"--reasoning-parser glm45",
"--tool-call-parser glm47",
"--host {{HOST_IP}}", "--host {{HOST_IP}}",
"--port {{PORT}}", "--port {{PORT}}",
], ],
@@ -722,8 +714,6 @@ sgl-eval run gsm8k \\
"--dsa-decode-backend trtllm", "--dsa-decode-backend trtllm",
"--kv-cache-dtype fp8_e4m3", "--kv-cache-dtype fp8_e4m3",
"--moe-runner-backend flashinfer_trtllm", "--moe-runner-backend flashinfer_trtllm",
"--reasoning-parser glm45",
"--tool-call-parser glm47",
"--host {{HOST_IP}}", "--host {{HOST_IP}}",
"--port {{PORT}}", "--port {{PORT}}",
], ],
@@ -746,8 +736,6 @@ sgl-eval run gsm8k \\
"--speculative-num-steps 5", "--speculative-num-steps 5",
"--speculative-eagle-topk 1", "--speculative-eagle-topk 1",
"--speculative-num-draft-tokens 6", "--speculative-num-draft-tokens 6",
"--reasoning-parser glm45",
"--tool-call-parser glm47",
"--host {{HOST_IP}}", "--host {{HOST_IP}}",
"--port {{PORT}}", "--port {{PORT}}",
], ],
@@ -767,8 +755,6 @@ sgl-eval run gsm8k \\
"--dsa-decode-backend trtllm", "--dsa-decode-backend trtllm",
"--kv-cache-dtype fp8_e4m3", "--kv-cache-dtype fp8_e4m3",
"--moe-runner-backend flashinfer_trtllm", "--moe-runner-backend flashinfer_trtllm",
"--reasoning-parser glm45",
"--tool-call-parser glm47",
"--host {{HOST_IP}}", "--host {{HOST_IP}}",
"--port {{PORT}}", "--port {{PORT}}",
], ],
@@ -789,8 +775,6 @@ sgl-eval run gsm8k \\
"--speculative-num-steps 5", "--speculative-num-steps 5",
"--speculative-eagle-topk 1", "--speculative-eagle-topk 1",
"--speculative-num-draft-tokens 6", "--speculative-num-draft-tokens 6",
"--reasoning-parser glm45",
"--tool-call-parser glm47",
"--host {{HOST_IP}}", "--host {{HOST_IP}}",
"--port {{PORT}}", "--port {{PORT}}",
], ],
@@ -807,8 +791,6 @@ sgl-eval run gsm8k \\
"--dsa-decode-backend trtllm", "--dsa-decode-backend trtllm",
"--kv-cache-dtype fp8_e4m3", "--kv-cache-dtype fp8_e4m3",
"--moe-runner-backend flashinfer_trtllm", "--moe-runner-backend flashinfer_trtllm",
"--reasoning-parser glm45",
"--tool-call-parser glm47",
"--host {{HOST_IP}}", "--host {{HOST_IP}}",
"--port {{PORT}}", "--port {{PORT}}",
], ],
+36 -9
View File
@@ -96,15 +96,45 @@ sgl-eval run aime25 \\
b200: "lmsysorg/sglang:latest", b200: "lmsysorg/sglang:latest",
gb300: "lmsysorg/sglang:latest", gb300: "lmsysorg/sglang:latest",
b300: "lmsysorg/sglang:latest", b300: "lmsysorg/sglang:latest",
mi355x: "lmsysorg/sglang-rocm:v0.5.13.post1-rocm720-mi35x-20260618", // >= v0.5.20 so `--*-parser auto` detects GLM-5.3 (#38297); the rocm700
mi325x: "lmsysorg/sglang-rocm:v0.5.13.post1-rocm700-mi30x-20260616", // line stopped at v0.5.19, so mi30x moves to the rocm720 build.
mi300x: "lmsysorg/sglang-rocm:v0.5.13.post1-rocm700-mi30x-20260616", mi355x: "lmsysorg/sglang-rocm:v0.5.20-rocm720-mi35x-20260920",
mi325x: "lmsysorg/sglang-rocm:v0.5.20-rocm720-mi30x-20260920",
mi300x: "lmsysorg/sglang-rocm:v0.5.20-rocm720-mi30x-20260920",
}, },
github: { github: {
cookbookModel: "zai-org/glm-5.3", cookbookModel: "zai-org/glm-5.3",
}, },
// Parser flags live in one overlay dim so every generated command gets them
// without per-cell duplication; the Parsers card toggles derive on/off from
// the composed flags. `auto` needs the GLM-5.3 template detection (v0.5.20+).
overlayDims: [
{
id: "parsers",
title: "Parsers",
default: "auto",
options: [
{
id: "auto",
label: "Auto (glm45 + glm47)",
stripPrefixes: ["--reasoning-parser", "--tool-call-parser"],
flags: [
"--reasoning-parser auto",
"--tool-call-parser auto",
],
},
{
id: "off",
label: "Off",
stripPrefixes: ["--reasoning-parser", "--tool-call-parser"],
flags: [],
},
],
},
],
playgroundFeatures: { playgroundFeatures: {
// ----- Card 1: "Attention Parallelism" ----- // ----- Card 1: "Attention Parallelism" -----
@@ -162,8 +192,8 @@ sgl-eval run aime25 \\
// ----- Card 3: "Parsers" ----- // ----- Card 3: "Parsers" -----
parsers: { parsers: {
items: [ items: [
{ id: "reasoning", label: "Reasoning Parser", flag: "--reasoning-parser glm45" }, { id: "reasoning", label: "Reasoning Parser", flag: "--reasoning-parser auto" },
{ id: "toolCall", label: "Tool Call Parser", flag: "--tool-call-parser glm47" }, { id: "toolCall", label: "Tool Call Parser", flag: "--tool-call-parser auto" },
], ],
}, },
@@ -194,10 +224,7 @@ sgl-eval run aime25 \\
flags: ["--speculative-algorithm DFLASH", flags: ["--speculative-algorithm DFLASH",
"--speculative-draft-model-path incoai/GLM-5.3-DFlash2", "--speculative-draft-model-path incoai/GLM-5.3-DFlash2",
"--speculative-draft-attention-backend fa4"], "--speculative-draft-attention-backend fa4"],
// The DFlash2 drafter (PR #35371) merged after v0.5.18, so neither the note: "⚠️ The draft is a separate checkpoint: fetch incoai/GLM-5.3-DFlash2 alongside the target. It is public but licensed CC BY-NC-ND 4.0 for research and evaluation.",
// release wheel nor the lmsysorg/sglang:latest image this page pins
// carries it. Drop this note once a release ships it.
note: "⚠️ Needs a nightly image: the DFlash2 drafter (PR #35371) is not in the release wheel nor the lmsysorg/sglang:latest image this page pins — install SGLang from main or use a lmsysorg/sglang:dev image. The draft is a separate checkpoint, so fetch incoai/GLM-5.3-DFlash2 alongside the target; it is public but licensed CC BY-NC-ND 4.0 for research and evaluation.",
disable: [ disable: [
{ when: { dpAttnOn: [true] }, { when: { dpAttnOn: [true] },
reason: "DFLASH speculative decoding does not support DP-Attention — the server rejects the combination at startup. Turn DP-Attention off in the Attention card above (the high-throughput recipes enable it)." }, reason: "DFLASH speculative decoding does not support DP-Attention — the server rejects the combination at startup. Turn DP-Attention off in the Attention card above (the high-throughput recipes enable it)." },
@@ -0,0 +1,47 @@
// SPDX-FileCopyrightText: Copyright (c) 2026 The SGLang Authors
// SPDX-License-Identifier: Apache-2.0
use axum::http::{header::AUTHORIZATION, HeaderMap};
use reqwest::{Client, RequestBuilder, Url};
use std::time::Duration;
/// Cancels unfinished engine work without delaying request cleanup.
pub(super) struct AbortOnDrop(Option<RequestBuilder>);
impl AbortOnDrop {
// Only router-minted IDs are safe: the engine aborts by prefix.
pub(super) fn new(
client: &Client,
worker: &Url,
headers: &HeaderMap,
rid: Option<&str>,
) -> Self {
Self(rid.filter(|rid| !rid.is_empty()).map(|rid| {
let mut request = client
.post(worker.join("/abort_request").expect("validated worker URL"))
.json(&serde_json::json!({"rid": rid, "abort_all": false}))
.timeout(Duration::from_secs(5));
if let Some(auth) = headers.get(AUTHORIZATION) {
request = request.header(AUTHORIZATION, auth);
}
request
}))
}
pub(super) fn disarm(&mut self) {
self.0 = None;
}
}
impl Drop for AbortOnDrop {
fn drop(&mut self) {
let Some(request) = self.0.take() else { return };
if let Ok(runtime) = tokio::runtime::Handle::try_current() {
runtime.spawn(async move {
if let Err(error) = request.send().await.and_then(|r| r.error_for_status()) {
tracing::warn!(%error, "engine abort failed");
}
});
}
}
}
+24
View File
@@ -3,8 +3,11 @@
//! HTTP proxy — forwards requests to the upstream SGLang worker. //! HTTP proxy — forwards requests to the upstream SGLang worker.
mod abort;
pub mod sse; pub mod sse;
use abort::AbortOnDrop;
use crate::health::circuit_breaker::CircuitBreaker; use crate::health::circuit_breaker::CircuitBreaker;
use crate::server::error::ApiError; use crate::server::error::ApiError;
use crate::server::header_utils::should_forward_request_header; use crate::server::header_utils::should_forward_request_header;
@@ -203,6 +206,7 @@ impl Proxy {
/// path concatenation (no double-slash) and pass a typed URL to the /// path concatenation (no double-slash) and pass a typed URL to the
/// split error variants (`UpstreamUnreachable` / `UpstreamTimeout` / /// split error variants (`UpstreamUnreachable` / `UpstreamTimeout` /
/// `UpstreamStatus`). /// `UpstreamStatus`).
#[allow(clippy::too_many_arguments)]
pub async fn forward_json_to( pub async fn forward_json_to(
&self, &self,
worker_url: &str, worker_url: &str,
@@ -211,6 +215,7 @@ impl Proxy {
path: &str, path: &str,
headers: &HeaderMap, headers: &HeaderMap,
body: Bytes, body: Bytes,
abort_rid: Option<&str>,
) -> Result<Response<Body>, ApiError> { ) -> Result<Response<Body>, ApiError> {
let permit = breaker.acquire().ok_or_else(|| ApiError::BreakerOpen { let permit = breaker.acquire().ok_or_else(|| ApiError::BreakerOpen {
worker: worker_url.to_string(), worker: worker_url.to_string(),
@@ -228,6 +233,8 @@ impl Proxy {
req = req req = req
.header("content-type", "application/json") .header("content-type", "application/json")
.timeout(self.request_timeout); .timeout(self.request_timeout);
let mut abort =
AbortOnDrop::new(self.client_for(protocol), &worker_url, headers, abort_rid);
let resp = req.send().await.map_err(|e| { let resp = req.send().await.map_err(|e| {
breaker.record_failure(); breaker.record_failure();
Self::classify_reqwest_error_for(worker_url.clone(), e, path) Self::classify_reqwest_error_for(worker_url.clone(), e, path)
@@ -257,6 +264,7 @@ impl Proxy {
return Err(ApiError::UpstreamStatus { status }); return Err(ApiError::UpstreamStatus { status });
} }
}; };
abort.disarm();
match breaker_outcome(status) { match breaker_outcome(status) {
BreakerOutcome::Failure => breaker.record_failure(), BreakerOutcome::Failure => breaker.record_failure(),
BreakerOutcome::Success => breaker.record_success(), BreakerOutcome::Success => breaker.record_success(),
@@ -301,6 +309,7 @@ impl Proxy {
path: &str, path: &str,
headers: &HeaderMap, headers: &HeaderMap,
body: Bytes, body: Bytes,
abort_rid: Option<&str>,
stream_guards: Option<Box<dyn Send + 'static>>, stream_guards: Option<Box<dyn Send + 'static>>,
on_first_byte: Option<Box<dyn FnOnce() + Send + 'static>>, on_first_byte: Option<Box<dyn FnOnce() + Send + 'static>>,
on_stream_end: Option<Box<dyn FnOnce(sse::StreamEnd) + Send + 'static>>, on_stream_end: Option<Box<dyn FnOnce(sse::StreamEnd) + Send + 'static>>,
@@ -322,11 +331,16 @@ impl Proxy {
req = req req = req
.header("content-type", "application/json") .header("content-type", "application/json")
.header("accept", "text/event-stream"); .header("accept", "text/event-stream");
let mut abort =
AbortOnDrop::new(self.client_for(protocol), &worker_url, headers, abort_rid);
let resp = req.send().await.map_err(|e| { let resp = req.send().await.map_err(|e| {
breaker.record_failure(); breaker.record_failure();
Self::classify_reqwest_error_for(worker_url.clone(), e, path) Self::classify_reqwest_error_for(worker_url.clone(), e, path)
})?; })?;
let status = resp.status(); let status = resp.status();
if !status.is_success() {
abort.disarm();
}
let upstream_ct = resp let upstream_ct = resp
.headers() .headers()
.get(reqwest::header::CONTENT_TYPE) .get(reqwest::header::CONTENT_TYPE)
@@ -366,6 +380,9 @@ impl Proxy {
BreakerOutcome::Success => { BreakerOutcome::Success => {
let breaker_for_hook = Arc::clone(breaker); let breaker_for_hook = Arc::clone(breaker);
Some(Box::new(move |end| { Some(Box::new(move |end| {
if end.reason == sse::StreamEndReason::Completed {
abort.disarm();
}
match stream_breaker_outcome(end) { match stream_breaker_outcome(end) {
BreakerOutcome::Success => breaker_for_hook.record_success(), BreakerOutcome::Success => breaker_for_hook.record_success(),
BreakerOutcome::Failure => breaker_for_hook.record_failure(), BreakerOutcome::Failure => breaker_for_hook.record_failure(),
@@ -465,6 +482,7 @@ mod tests {
"/chat", "/chat",
&headers, &headers,
Bytes::new(), Bytes::new(),
None,
) )
.now_or_never() .now_or_never()
.is_none()); .is_none());
@@ -481,6 +499,7 @@ mod tests {
None, None,
None, None,
None, None,
None,
) )
.now_or_never() .now_or_never()
.is_none()); .is_none());
@@ -583,6 +602,7 @@ mod tests {
None, None,
None, None,
None, None,
None,
expiration, expiration,
) )
.await .await
@@ -694,6 +714,7 @@ mod tests {
"/v1/chat/completions", "/v1/chat/completions",
&headers, &headers,
Bytes::from_static(b"{}"), Bytes::from_static(b"{}"),
None,
) )
.await .await
.expect("dispatch should reach the worker (breaker must stay closed)"); .expect("dispatch should reach the worker (breaker must stay closed)");
@@ -737,6 +758,7 @@ mod tests {
"/v1/chat/completions", "/v1/chat/completions",
&headers, &headers,
Bytes::from_static(b"{}"), Bytes::from_static(b"{}"),
None,
) )
.await; .await;
} }
@@ -780,6 +802,7 @@ mod tests {
"/v1/chat/completions", "/v1/chat/completions",
&headers, &headers,
Bytes::from_static(b"{}"), Bytes::from_static(b"{}"),
None,
) )
.await .await
.expect("the half-open probe must be admitted and reach the worker"); .expect("the half-open probe must be admitted and reach the worker");
@@ -819,6 +842,7 @@ mod tests {
None, None,
None, None,
None, None,
None,
) )
.await .await
.expect("streaming dispatch should reach the worker"); .expect("streaming dispatch should reach the worker");
@@ -150,6 +150,9 @@ async fn access_log_and_record(
.unwrap_or_else(|| outcome_from_status(status.as_u16())) .unwrap_or_else(|| outcome_from_status(status.as_u16()))
.as_str(), .as_str(),
worker = log_ctx.map(|c| c.worker_url.as_str()).unwrap_or(""), worker = log_ctx.map(|c| c.worker_url.as_str()).unwrap_or(""),
engine_rid = log_ctx
.and_then(|c| c.engine_rid.as_deref())
.unwrap_or(""),
model = log_ctx.map(|c| c.model_id.as_str()).unwrap_or(""), model = log_ctx.map(|c| c.model_id.as_str()).unwrap_or(""),
stream = log_ctx.is_some_and(|c| c.streaming), stream = log_ctx.is_some_and(|c| c.streaming),
latency_ms, latency_ms,
@@ -446,6 +449,7 @@ mod tests {
model_id: "tiny".into(), model_id: "tiny".into(),
streaming: false, streaming: false,
outcome: RequestOutcome::Cancelled, outcome: RequestOutcome::Cancelled,
engine_rid: Some("1f0c2b7a4e9d4f3ab6c5d8e7f0a1b2c3".into()),
}); });
resp resp
}), }),
@@ -468,6 +472,11 @@ mod tests {
logs.contains("worker=\"http://worker-a:30000\"") && logs.contains("model=\"tiny\""), logs.contains("worker=\"http://worker-a:30000\"") && logs.contains("model=\"tiny\""),
"a routed request must be logged with its worker and model; captured:\n{logs}", "a routed request must be logged with its worker and model; captured:\n{logs}",
); );
assert!(
logs.contains("engine_rid=\"1f0c2b7a4e9d4f3ab6c5d8e7f0a1b2c3\"")
&& logs.contains("request_id="),
"the minted rid must be logged beside the caller's request id; captured:\n{logs}",
);
// The handler's outcome must win over the status-derived fallback — // The handler's outcome must win over the status-derived fallback —
// otherwise the log and `worker_requests_total` can disagree about a // otherwise the log and `worker_requests_total` can disagree about a
// request the handler classified itself (here, a cancellation served // request the handler classified itself (here, a cancellation served
@@ -223,6 +223,8 @@ pub struct RequestLogContext {
/// line and `worker_requests_total` cannot disagree — the middleware can /// line and `worker_requests_total` cannot disagree — the middleware can
/// only see the status, which cannot express a router-side cancellation. /// only see the status, which cannot express a router-side cancellation.
pub outcome: RequestOutcome, pub outcome: RequestOutcome,
/// Router-minted engine ID, logged beside the caller's correlation ID.
pub engine_rid: Option<String>,
} }
/// Final outcome of a 2xx SSE stream. /// Final outcome of a 2xx SSE stream.
@@ -81,7 +81,12 @@ pub(super) async fn forward_chat_request(
}; };
(decode, bootstrap) (decode, bootstrap)
}); });
let body = request.into_outgoing_body(ctx, pd.as_ref().map(|(_, bootstrap)| bootstrap))?; let engine_rid = request.engine_rid(pd.is_some());
let body = request.into_outgoing_body(
ctx,
pd.as_ref().map(|(_, bootstrap)| bootstrap),
engine_rid.as_deref(),
)?;
let prefill_load_guards = (worker_load_guard, active_request_guard); let prefill_load_guards = (worker_load_guard, active_request_guard);
// In PD mode, prefill runs independently and decode supplies the client response. // In PD mode, prefill runs independently and decode supplies the client response.
@@ -112,6 +117,7 @@ pub(super) async fn forward_chat_request(
&response_worker, &response_worker,
&headers, &headers,
body, body,
engine_rid.as_deref(),
response_load_guards, response_load_guards,
&metrics, &metrics,
expiration_token.clone(), expiration_token.clone(),
@@ -124,7 +130,7 @@ pub(super) async fn forward_chat_request(
model: metrics.model.clone(), model: metrics.model.clone(),
}), }),
}; };
let log_context = metrics.record_dispatch_result(&result); let log_context = metrics.record_dispatch_result(&result, engine_rid);
// Materialize dispatch errors here so the access log retains the selected worker. // Materialize dispatch errors here so the access log retains the selected worker.
let mut response = match result { let mut response = match result {
Ok(mut response) => { Ok(mut response) => {
@@ -171,6 +177,7 @@ fn spawn_prefill_request(
CHAT_PATH, CHAT_PATH,
&headers, &headers,
body, body,
None,
) )
.await .await
{ {
@@ -188,11 +195,13 @@ fn spawn_prefill_request(
}); });
} }
#[allow(clippy::too_many_arguments)]
async fn forward_to_response_worker( async fn forward_to_response_worker(
ctx: &AppContext, ctx: &AppContext,
worker: &Worker, worker: &Worker,
headers: &HeaderMap, headers: &HeaderMap,
body: Bytes, body: Bytes,
engine_rid: Option<&str>,
load_guards: LoadGuards, load_guards: LoadGuards,
metrics: &DispatchMetrics, metrics: &DispatchMetrics,
expiration: CancellationToken, expiration: CancellationToken,
@@ -209,6 +218,7 @@ async fn forward_to_response_worker(
CHAT_PATH, CHAT_PATH,
headers, headers,
body, body,
engine_rid,
Some(stream_guards), Some(stream_guards),
Some(metrics.first_byte_callback()), Some(metrics.first_byte_callback()),
Some(metrics.stream_end_callback(worker.url.clone())), Some(metrics.stream_end_callback(worker.url.clone())),
@@ -226,6 +236,7 @@ async fn forward_to_response_worker(
CHAT_PATH, CHAT_PATH,
headers, headers,
body, body,
engine_rid,
) )
.await .await
} }
@@ -295,6 +306,7 @@ impl DispatchMetrics {
fn record_dispatch_result( fn record_dispatch_result(
&self, &self,
result: &Result<Response<Body>, ApiError>, result: &Result<Response<Body>, ApiError>,
engine_rid: Option<String>,
) -> RequestLogContext { ) -> RequestLogContext {
let http_status = match result { let http_status = match result {
Ok(response) => response.status().as_u16(), Ok(response) => response.status().as_u16(),
@@ -326,6 +338,7 @@ impl DispatchMetrics {
model_id: self.model.clone(), model_id: self.model.clone(),
streaming: self.streaming, streaming: self.streaming,
outcome, outcome,
engine_rid,
} }
} }
} }
@@ -26,6 +26,8 @@ pub(super) struct PreparedChatRequest {
pub(super) tokens: Option<RequestTokens>, pub(super) tokens: Option<RequestTokens>,
/// Token count for routing/load accounting; estimated from body size when unavailable. /// Token count for routing/load accounting; estimated from body size when unavailable.
pub(super) input_token_count: usize, pub(super) input_token_count: usize,
caller_set_rid: bool,
fans_out: bool,
can_forward_input_ids: bool, can_forward_input_ids: bool,
parsed_body: Option<Value>, parsed_body: Option<Value>,
sampling_defaults: Vec<(SamplingField, Number)>, sampling_defaults: Vec<(SamplingField, Number)>,
@@ -69,16 +71,27 @@ impl PreparedChatRequest {
body, body,
tokens, tokens,
input_token_count, input_token_count,
caller_set_rid: fields.caller_set_rid,
fans_out: requests_multiple_samples(&fields, &sampling_defaults),
can_forward_input_ids, can_forward_input_ids,
parsed_body, parsed_body,
sampling_defaults, sampling_defaults,
}) })
} }
pub(super) fn engine_rid(&self, pd_mode: bool) -> Option<String> {
// Caller IDs are unsafe for prefix aborts; fan-out regenerates IDs; PD must finish KV transfer.
if self.caller_set_rid || self.fans_out || pd_mode {
return None;
}
Some(uuid::Uuid::new_v4().simple().to_string())
}
pub(super) fn into_outgoing_body( pub(super) fn into_outgoing_body(
self, self,
ctx: &AppContext, ctx: &AppContext,
bootstrap: Option<&BootstrapFields>, bootstrap: Option<&BootstrapFields>,
engine_rid: Option<&str>,
) -> Result<Bytes, ApiError> { ) -> Result<Bytes, ApiError> {
// Routing tokens can replace engine tokenization only for supported chat templates. // Routing tokens can replace engine tokenization only for supported chat templates.
let input_ids = match (self.tokens.as_ref(), self.parsed_body.as_ref()) { let input_ids = match (self.tokens.as_ref(), self.parsed_body.as_ref()) {
@@ -104,6 +117,7 @@ impl PreparedChatRequest {
input_ids, input_ids,
bootstrap, bootstrap,
&self.sampling_defaults, &self.sampling_defaults,
engine_rid,
) )
} }
} }
@@ -116,6 +130,8 @@ pub(super) struct RoutingFields {
max_tokens: Option<u64>, max_tokens: Option<u64>,
max_completion_tokens: Option<u64>, max_completion_tokens: Option<u64>,
sampling: [SamplingValue; SamplingField::ALL.len()], sampling: [SamplingValue; SamplingField::ALL.len()],
// Preserve both string and list IDs without retaining their contents.
caller_set_rid: bool,
} }
/// Null is absent; unrepresentable values are rejected only under a sampling contract. /// Null is absent; unrepresentable values are rejected only under a sampling contract.
@@ -248,6 +264,7 @@ impl RoutingKey {
enum RequestKey { enum RequestKey {
Routing(RoutingKey), Routing(RoutingKey),
Sampling(SamplingField), Sampling(SamplingField),
Rid,
Other, Other,
} }
@@ -267,6 +284,7 @@ impl<'de> Deserialize<'de> for RequestKey {
"model" => RequestKey::Routing(RoutingKey::Model), "model" => RequestKey::Routing(RoutingKey::Model),
"max_tokens" => RequestKey::Routing(RoutingKey::MaxTokens), "max_tokens" => RequestKey::Routing(RoutingKey::MaxTokens),
"max_completion_tokens" => RequestKey::Routing(RoutingKey::MaxCompletionTokens), "max_completion_tokens" => RequestKey::Routing(RoutingKey::MaxCompletionTokens),
"rid" => RequestKey::Rid,
other => match SamplingField::from_wire_name(other) { other => match SamplingField::from_wire_name(other) {
Some(field) => RequestKey::Sampling(field), Some(field) => RequestKey::Sampling(field),
None => RequestKey::Other, None => RequestKey::Other,
@@ -324,6 +342,9 @@ impl<'de> serde::de::Visitor<'de> for RoutingFieldsVisitor {
SamplingValue::Unusable SamplingValue::Unusable
}; };
} }
RequestKey::Rid => {
fields.caller_set_rid = map.next_value::<Option<IgnoredAny>>()?.is_some();
}
RequestKey::Other => { RequestKey::Other => {
// Validate unrelated JSON without retaining its contents. // Validate unrelated JSON without retaining its contents.
map.next_value::<IgnoredAny>()?; map.next_value::<IgnoredAny>()?;
@@ -374,6 +395,22 @@ fn should_tokenize_request(
can_forward_input_ids || policy_needs_request_tokens || bucket_routing_enabled can_forward_input_ids || policy_needs_request_tokens || bucket_routing_enabled
} }
// Use the effective n, including injected defaults; unreadable values opt out.
fn requests_multiple_samples(
fields: &RoutingFields,
sampling_defaults: &[(SamplingField, Number)],
) -> bool {
match fields.sampling_field(SamplingField::N) {
SamplingValue::Number(n) => n > 1.0,
SamplingValue::Unusable => true,
SamplingValue::Absent => sampling_defaults
.iter()
.find(|(field, _)| *field == SamplingField::N)
.and_then(|(_, value)| value.as_f64())
.is_some_and(|n| n > 1.0),
}
}
fn estimate_prefill_tokens(body: &Bytes) -> usize { fn estimate_prefill_tokens(body: &Bytes) -> usize {
// Never 0: a zero-load entry is invisible to the cache-aware imbalance fast path. // Never 0: a zero-load entry is invisible to the cache-aware imbalance fast path.
(body.len() / BYTES_PER_TOKEN_ESTIMATE).max(1) (body.len() / BYTES_PER_TOKEN_ESTIMATE).max(1)
@@ -391,9 +428,10 @@ pub(super) struct BootstrapFields {
} }
/// Append before the closing brace so injected values win over explicit nulls. /// Append before the closing brace so injected values win over explicit nulls.
fn append_sampling_defaults( fn append_top_level_fields(
body: &Bytes, body: &Bytes,
sampling_defaults: &[(SamplingField, Number)], sampling_defaults: &[(SamplingField, Number)],
rid: Option<&str>,
) -> Option<Bytes> { ) -> Option<Bytes> {
use std::io::Write as _; use std::io::Write as _;
@@ -405,13 +443,23 @@ fn append_sampling_defaults(
let has_members = body[open + 1..close] let has_members = body[open + 1..close]
.iter() .iter()
.any(|b| !b.is_ascii_whitespace()); .any(|b| !b.is_ascii_whitespace());
let mut output = Vec::with_capacity(body.len() + 24 * sampling_defaults.len() + 1); let rid_budget = rid.map_or(0, |rid| rid.len() + ",\"rid\":\"\"".len());
let mut output = Vec::with_capacity(body.len() + 24 * sampling_defaults.len() + rid_budget + 1);
output.extend_from_slice(&body[..close]); output.extend_from_slice(&body[..close]);
for (i, (field, value)) in sampling_defaults.iter().enumerate() { let mut wrote_any = has_members;
if has_members || i > 0 { for (field, value) in sampling_defaults {
if wrote_any {
output.push(b','); output.push(b',');
} }
write!(output, "\"{}\":{}", field.wire_name(), value).ok()?; write!(output, "\"{}\":{}", field.wire_name(), value).ok()?;
wrote_any = true;
}
if let Some(rid) = rid {
if wrote_any {
output.push(b',');
}
output.extend_from_slice(b"\"rid\":");
serde_json::to_writer(&mut output, rid).ok()?;
} }
output.extend_from_slice(&body[close..]); output.extend_from_slice(&body[close..]);
Some(Bytes::from(output)) Some(Bytes::from(output))
@@ -424,14 +472,15 @@ fn build_outgoing_body(
input_ids: Option<&[u32]>, input_ids: Option<&[u32]>,
bootstrap: Option<&BootstrapFields>, bootstrap: Option<&BootstrapFields>,
sampling_defaults: &[(SamplingField, Number)], sampling_defaults: &[(SamplingField, Number)],
rid: Option<&str>,
) -> Result<Bytes, ApiError> { ) -> Result<Bytes, ApiError> {
let sampling_only = input_ids.is_none() && bootstrap.is_none(); let needs_parse = input_ids.is_some() || bootstrap.is_some();
if sampling_only && sampling_defaults.is_empty() { if !needs_parse && sampling_defaults.is_empty() && rid.is_none() {
// Cloning Bytes shares the original allocation when no injection is needed. // Cloning Bytes shares the original allocation when no injection is needed.
return Ok(body.clone()); return Ok(body.clone());
} }
if sampling_only { if !needs_parse {
if let Some(spliced) = append_sampling_defaults(body, sampling_defaults) { if let Some(spliced) = append_top_level_fields(body, sampling_defaults, rid) {
return Ok(spliced); return Ok(spliced);
} }
} }
@@ -445,6 +494,9 @@ fn build_outgoing_body(
return Err(invalid_request()); return Err(invalid_request());
} }
}; };
if let Some(rid) = rid {
body_fields.insert("rid".into(), Value::String(rid.to_owned()));
}
for (field, default) in sampling_defaults { for (field, default) in sampling_defaults {
body_fields.insert(field.wire_name().into(), default.clone().into()); body_fields.insert(field.wire_name().into(), default.clone().into());
} }
@@ -739,7 +791,8 @@ mod tests {
.unwrap() .unwrap()
.extend(fields.as_object().unwrap().clone()); .extend(fields.as_object().unwrap().clone());
for value in [None, Some(original.clone())] { for value in [None, Some(original.clone())] {
let out = build_outgoing_body(&body, value, ids, bootstrap.as_ref(), &[]).unwrap(); let out =
build_outgoing_body(&body, value, ids, bootstrap.as_ref(), &[], None).unwrap();
assert_eq!(serde_json::from_slice::<Value>(&out).unwrap(), expected); assert_eq!(serde_json::from_slice::<Value>(&out).unwrap(), expected);
} }
} }
@@ -750,7 +803,7 @@ mod tests {
for raw in [r#"{"model":"x"}"#, r#"{"model":"x","messages":[]}"#] { for raw in [r#"{"model":"x"}"#, r#"{"model":"x","messages":[]}"#] {
let body = Bytes::copy_from_slice(raw.as_bytes()); let body = Bytes::copy_from_slice(raw.as_bytes());
for value in [None, Some(serde_json::from_slice(&body).unwrap())] { for value in [None, Some(serde_json::from_slice(&body).unwrap())] {
let out = build_outgoing_body(&body, value, None, None, &[]).unwrap(); let out = build_outgoing_body(&body, value, None, None, &[], None).unwrap();
assert_eq!(out, body); assert_eq!(out, body);
assert_eq!(out.as_ptr(), body.as_ptr()); assert_eq!(out.as_ptr(), body.as_ptr());
} }
@@ -1108,7 +1161,7 @@ mod tests {
&metrics(), &metrics(),
) )
.unwrap(); .unwrap();
let out = build_outgoing_body(&body, None, Some(&[1, 2, 3]), None, &inject).unwrap(); let out = build_outgoing_body(&body, None, Some(&[1, 2, 3]), None, &inject, None).unwrap();
assert_eq!( assert_eq!(
serde_json::from_slice::<Value>(&out).unwrap(), serde_json::from_slice::<Value>(&out).unwrap(),
json!({ json!({
@@ -1240,7 +1293,7 @@ mod tests {
let inject = resolve_sampling_defaults(&config, &fields_of(raw), &metrics()).unwrap(); let inject = resolve_sampling_defaults(&config, &fields_of(raw), &metrics()).unwrap();
assert_eq!(inject.len(), 2); assert_eq!(inject.len(), 2);
for value in [None, Some(serde_json::from_slice(&body).unwrap())] { for value in [None, Some(serde_json::from_slice(&body).unwrap())] {
let out = build_outgoing_body(&body, value, None, None, &inject).unwrap(); let out = build_outgoing_body(&body, value, None, None, &inject, None).unwrap();
assert_eq!(std::str::from_utf8(&out).unwrap(), expected, "{raw}"); assert_eq!(std::str::from_utf8(&out).unwrap(), expected, "{raw}");
let parsed: Value = serde_json::from_slice(&out).unwrap(); let parsed: Value = serde_json::from_slice(&out).unwrap();
assert_eq!(parsed["temperature"], json!(1.0)); assert_eq!(parsed["temperature"], json!(1.0));
@@ -1464,4 +1517,35 @@ mod tests {
"a value at the cap is still read" "a value at the cap is still read"
); );
} }
#[test]
fn abort_opt_outs_follow_caller_rid_and_effective_sample_count() {
for raw in [r#"{"rid":"abc"}"#, r#"{"rid":["a","b"]}"#] {
assert!(fields_of(raw).caller_set_rid);
}
assert!(!fields_of(r#"{"rid":null}"#).caller_set_rid);
for (raw, fan_out) in [
(r#"{}"#, false),
(r#"{"n":1}"#, false),
(r#"{"n":2}"#, true),
(r#"{"n":"3"}"#, true),
(r#"{"n":[2]}"#, true),
] {
assert_eq!(
requests_multiple_samples(&fields_of(raw), &[]),
fan_out,
"{raw}"
);
}
for (config, fan_out) in [(r#"{"n":1}"#, false), (r#"{"n":4}"#, true)] {
let fields = fields_of("{}");
let defaults = resolve_sampling_defaults(
&overrides_of(ConflictPolicy::Reject, config),
&fields,
&metrics(),
)
.unwrap();
assert_eq!(requests_multiple_samples(&fields, &defaults), fan_out);
}
}
} }
@@ -198,7 +198,18 @@ async fn caller_input_ids_are_used_for_routing_and_preserved() {
send(Arc::clone(&ctx), request.clone()).await, send(Arc::clone(&ctx), request.clone()).await,
StatusCode::OK StatusCode::OK
); );
assert_eq!(captured(&mock), request, "body must be forwarded untouched"); let mut forwarded = captured(&mock);
let rid = forwarded
.as_object_mut()
.expect("a forwarded chat body is an object")
.remove("rid");
assert!(
rid.as_ref()
.and_then(Value::as_str)
.is_some_and(crate::common::is_engine_shaped_rid),
"plain mode must mint an abort rid; got {rid:?}",
);
assert_eq!(forwarded, request, "body must be forwarded untouched");
} }
// Bypasses are not rendering failures. // Bypasses are not rendering failures.
assert!(!ctx assert!(!ctx
@@ -17,10 +17,14 @@ use sgl_router::workers::{WireProtocol, Worker, WorkerRegistry};
use axum::body::Body; use axum::body::Body;
use axum::http::{Request, StatusCode}; use axum::http::{Request, StatusCode};
use http_body_util::BodyExt; use http_body_util::BodyExt;
use sgl_router::state::load_monitor::router_inflight_load::{
spawn_janitor, JanitorHandle, RouterInflightLoadRegistry,
};
use std::sync::Arc; use std::sync::Arc;
use std::time::Duration; use std::time::Duration;
use tower::ServiceExt; use tower::ServiceExt;
mod cancellation;
mod reorg; mod reorg;
const TEST_TIMEOUT: Duration = Duration::from_secs(5); const TEST_TIMEOUT: Duration = Duration::from_secs(5);
@@ -74,6 +78,36 @@ fn build_ctx_with_worker(url: &str) -> Arc<AppContext> {
Arc::new(AppContext::new(cfg, tokenizers, proxy, registry, policies)) Arc::new(AppContext::new(cfg, tokenizers, proxy, registry, policies))
} }
/// Expire requests after 50ms; keep the janitor handle alive during the test.
fn build_ctx_with_janitor(url: &str) -> (Arc<AppContext>, JanitorHandle) {
let cfg = config_for(url);
let registry = Arc::new(WorkerRegistry::default());
let _ = registry.add(WorkerSpec {
id: WorkerId("w1".into()),
url: url.to_string(),
mode: WorkerMode::Plain,
model_ids: vec![ModelId("tiny".into())],
bootstrap_port: None,
});
let policies = Arc::new(build_policy_registry(&cfg).unwrap());
let tokenizers = Arc::new(TokenizerRegistry::load_from_config(&cfg).unwrap());
let proxy = Arc::new(Proxy::new(TEST_TIMEOUT).unwrap());
let router_inflight_load = RouterInflightLoadRegistry::new(
Arc::new(sgl_router::state::load_monitor::router_inflight_load::SystemTimeClock),
Duration::from_millis(50),
);
let janitor = spawn_janitor(Arc::clone(&router_inflight_load), Duration::from_millis(20));
let ctx = Arc::new(AppContext::with_router_inflight_load(
cfg,
tokenizers,
proxy,
registry,
policies,
router_inflight_load,
));
(ctx, janitor)
}
#[tokio::test] #[tokio::test]
async fn non_streaming_returns_200() { async fn non_streaming_returns_200() {
let worker = crate::common::mock_worker::MockWorker::start(vec![]).await; let worker = crate::common::mock_worker::MockWorker::start(vec![]).await;
@@ -1033,6 +1067,7 @@ async fn forward_json_to_records_failure_on_body_drop() {
"/v1/chat/completions", "/v1/chat/completions",
&headers, &headers,
body, body,
None,
) )
.await; .await;
assert!(res.is_err(), "body drop should surface as ApiError"); assert!(res.is_err(), "body drop should surface as ApiError");
@@ -1090,6 +1125,7 @@ async fn forward_json_to_records_success_only_after_body_completes() {
"/v1/chat/completions", "/v1/chat/completions",
&headers, &headers,
bytes::Bytes::from_static(b"{}"), bytes::Bytes::from_static(b"{}"),
None,
) )
.await; .await;
assert!(res.is_ok(), "clean OK call must succeed: {res:?}"); assert!(res.is_ok(), "clean OK call must succeed: {res:?}");
@@ -1145,6 +1181,7 @@ async fn forward_streaming_to_records_failure_on_mid_stream_drop() {
None, None,
None, None,
None, None,
None,
) )
.await; .await;
@@ -1248,6 +1285,7 @@ async fn forward_json_to_records_failure_on_5xx() {
"/v1/chat/completions", "/v1/chat/completions",
&headers, &headers,
body, body,
None,
) )
.await; .await;
@@ -1282,6 +1320,7 @@ async fn forward_json_to_rejects_when_breaker_open() {
"/v1/chat/completions", "/v1/chat/completions",
&headers, &headers,
body, body,
None,
) )
.await; .await;
@@ -1320,6 +1359,7 @@ async fn forward_json_to_malformed_url_returns_worker_misconfigured_and_trips_br
"/v1/chat/completions", "/v1/chat/completions",
&headers, &headers,
body, body,
None,
) )
.await; .await;
@@ -1607,58 +1647,17 @@ async fn streaming_active_load_drops_on_client_disconnect() {
/// returns; cancellation fires; handler returns 504. /// returns; cancellation fires; handler returns 504.
#[tokio::test] #[tokio::test]
async fn janitor_expiry_returns_504_stale_request_expired() { async fn janitor_expiry_returns_504_stale_request_expired() {
use sgl_router::state::load_monitor::router_inflight_load::{ // Upstream that takes 2s to respond — longer than the helper's 50ms
spawn_janitor, RouterInflightLoadRegistry, // stale_request_timeout, so the janitor sweeps before it answers.
};
// Upstream that takes 2s to respond — longer than our 50ms
// stale_request_timeout.
let worker = let worker =
crate::common::mock_worker::MockWorker::start_hanging(Duration::from_secs(2)).await; crate::common::mock_worker::MockWorker::start_hanging(Duration::from_secs(2)).await;
let (ctx, _janitor) = build_ctx_with_janitor(&worker.url);
let cfg = config_for(&worker.url);
let registry = Arc::new(WorkerRegistry::default());
let _ = registry.add(WorkerSpec {
id: WorkerId("w1".into()),
url: worker.url.clone(),
mode: WorkerMode::Plain,
model_ids: vec![ModelId("tiny".into())],
bootstrap_port: None,
});
let policies = Arc::new(build_policy_registry(&cfg).unwrap());
let tokenizers = Arc::new(TokenizerRegistry::load_from_config(&cfg).unwrap());
let proxy = Arc::new(Proxy::new(TEST_TIMEOUT).unwrap());
// Aggressive 50ms timeout: the janitor will sweep on the next
// tick (every 20ms) and fire the cancellation token before the
// upstream returns.
let router_inflight_load = RouterInflightLoadRegistry::new(
Arc::new(sgl_router::state::load_monitor::router_inflight_load::SystemTimeClock),
Duration::from_millis(50),
);
let _janitor = spawn_janitor(Arc::clone(&router_inflight_load), Duration::from_millis(20));
let ctx = Arc::new(AppContext::with_router_inflight_load(
cfg,
tokenizers,
proxy,
registry,
policies,
router_inflight_load,
));
let app = build_router(ctx); let app = build_router(ctx);
let req = Request::builder() let res = app
.method("POST") .oneshot(cancellation::request(serde_json::json!({})))
.uri("/v1/chat/completions") .await
.header("content-type", "application/json")
.body(Body::from(
serde_json::to_vec(&serde_json::json!({
"model": "tiny",
"messages": [{"role": "user", "content": "hi"}],
"stream": false,
}))
.unwrap(),
))
.unwrap(); .unwrap();
let res = app.oneshot(req).await.unwrap();
assert_eq!( assert_eq!(
res.status(), res.status(),
StatusCode::GATEWAY_TIMEOUT, StatusCode::GATEWAY_TIMEOUT,
@@ -1677,6 +1676,53 @@ async fn janitor_expiry_returns_504_stale_request_expired() {
body_str.contains("\"code\":\"stale_request_expired\""), body_str.contains("\"code\":\"stale_request_expired\""),
"504 body must encode the same code in the JSON envelope: {body_str}", "504 body must encode the same code in the JSON envelope: {body_str}",
); );
assert_engine_abort(&worker).await;
}
#[tokio::test]
async fn janitor_expiry_aborts_before_headers_and_mid_stream() {
use crate::common::mock_worker::MockWorker;
for before_headers in [true, false] {
let worker = if before_headers {
MockWorker::start_hanging(Duration::from_secs(2)).await
} else {
MockWorker::start_slow_stream(vec!["data: a\n\n"], Duration::from_secs(2)).await
};
let (ctx, _janitor) = build_ctx_with_janitor(&worker.url);
let response = build_router(ctx)
.oneshot(cancellation::request(serde_json::json!({"stream":true})))
.await
.unwrap();
assert_eq!(
response.status(),
if before_headers {
StatusCode::GATEWAY_TIMEOUT
} else {
StatusCode::OK
}
);
let result = response.into_body().collect().await;
assert_eq!(result.is_ok(), before_headers);
assert_engine_abort(&worker).await;
}
}
async fn assert_engine_abort(worker: &crate::common::mock_worker::MockWorker) {
tokio::time::timeout(TEST_TIMEOUT, async {
while worker.abort_log.lock().unwrap().is_empty() {
tokio::time::sleep(Duration::from_millis(10)).await;
}
})
.await
.unwrap();
let forwarded: serde_json::Value =
serde_json::from_slice(worker.captured.lock().unwrap().last_body.as_ref().unwrap())
.unwrap();
assert_eq!(
*worker.abort_log.lock().unwrap(),
vec![serde_json::json!({"rid":forwarded["rid"], "abort_all":false})]
);
} }
/// Task A: a non-streaming request that errors out (upstream /// Task A: a non-streaming request that errors out (upstream
@@ -0,0 +1,205 @@
// SPDX-FileCopyrightText: Copyright (c) 2026 The SGLang Authors
// SPDX-License-Identifier: Apache-2.0
use super::*;
use axum::{extract::State, http::HeaderMap, routing::post, Json, Router};
use serde_json::{json, Value};
use tokio::sync::mpsc;
type Event = (&'static str, Value);
type Events = mpsc::UnboundedSender<Event>;
async fn chat(State(events): State<Events>, Json(body): Json<Value>) -> (StatusCode, Body) {
events.send(("chat", body.clone())).unwrap();
if body["before_headers"] == true || (body["hold"] == true && body["stream"] != true) {
std::future::pending::<()>().await;
}
let response = if body["hold"] == true {
Body::from_stream(futures::stream::pending::<
Result<bytes::Bytes, std::io::Error>,
>())
} else if body["stream"] == true {
Body::from("data: [DONE]\n\n")
} else {
Body::from("{}")
};
(
StatusCode::from_u16(body["status"].as_u64().unwrap_or(200) as u16).unwrap(),
response,
)
}
async fn abort(
State(events): State<Events>,
headers: HeaderMap,
Json(body): Json<Value>,
) -> StatusCode {
assert_eq!(headers["authorization"], "Bearer test");
events.send(("abort", body)).unwrap();
StatusCode::INTERNAL_SERVER_ERROR // Abort failures must not affect the worker's breaker.
}
struct Harness {
ctx: Arc<AppContext>,
events: mpsc::UnboundedReceiver<Event>,
server: tokio::task::JoinHandle<()>,
}
impl Harness {
async fn new() -> Self {
let (events, rx) = mpsc::unbounded_channel();
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let ctx = build_ctx_with_worker(&format!("http://{}", listener.local_addr().unwrap()));
let server = tokio::spawn(async move {
axum::serve(
listener,
Router::new()
.route("/v1/chat/completions", post(chat))
.route("/abort_request", post(abort))
.with_state(events),
)
.await
.unwrap();
});
Self {
ctx,
events: rx,
server,
}
}
async fn event(&mut self, expected: &str) -> Value {
let (kind, body) = tokio::time::timeout(TEST_TIMEOUT, self.events.recv())
.await
.unwrap()
.unwrap();
assert_eq!(kind, expected);
body
}
async fn quiet(&mut self) {
assert!(
tokio::time::timeout(Duration::from_millis(50), self.events.recv())
.await
.is_err()
);
}
}
impl Drop for Harness {
fn drop(&mut self) {
self.server.abort();
}
}
pub(super) fn request(mut body: Value) -> Request<Body> {
body["model"] = json!("tiny");
Request::post("/v1/chat/completions")
.header("content-type", "application/json")
.header("authorization", "Bearer test")
.header("x-request-id", "reused-gateway-id")
.body(Body::from(body.to_string()))
.unwrap()
}
#[tokio::test]
async fn only_unfinished_requests_abort() {
let mut h = Harness::new().await;
for (stream, hold, before_headers) in [
(false, false, false),
(true, false, false),
(false, true, false),
(true, true, true),
(true, true, false),
] {
let task = tokio::spawn(build_router(h.ctx.clone()).oneshot(request(json!({
"stream": stream, "hold": hold, "before_headers": before_headers,
}))));
let forwarded = h.event("chat").await;
assert!(crate::common::is_engine_shaped_rid(
forwarded["rid"].as_str().unwrap()
));
let worker = h.ctx.registry.get(&WorkerId("w1".into())).unwrap();
if hold && (!stream || before_headers) {
worker.breaker.record_failure();
worker.breaker.record_failure();
task.abort();
assert!(task.await.unwrap_err().is_cancelled());
} else {
let response = tokio::time::timeout(TEST_TIMEOUT, task)
.await
.unwrap()
.unwrap()
.unwrap();
if hold {
drop(response); // Silent upstream: cancellation must not wait for a token.
} else {
response.into_body().collect().await.unwrap();
}
}
if hold {
assert_eq!(
h.event("abort").await,
json!({"rid": forwarded["rid"], "abort_all": false})
);
}
h.quiet().await;
assert_eq!(worker.breaker.snapshot().state_code, 0);
worker.breaker.record_success();
}
}
#[tokio::test]
async fn caller_ids_fan_out_and_rejected_streams_do_not_abort() {
let mut h = Harness::new().await;
for fields in [
json!({"rid":"a"}),
json!({"rid":["a","b"]}),
json!({"n":2}),
json!({"status":400}),
json!({"status":429}),
json!({"status":500}),
json!({"status":503}),
] {
let mut body = json!({"stream":true, "hold":true});
body.as_object_mut()
.unwrap()
.extend(fields.as_object().unwrap().clone());
let response = build_router(h.ctx.clone())
.oneshot(request(body))
.await
.unwrap();
let forwarded = h.event("chat").await;
if fields.get("status").is_none() {
assert_eq!(forwarded.get("rid"), fields.get("rid"));
}
drop(response);
h.quiet().await;
}
}
#[tokio::test]
async fn concurrent_requests_with_the_same_header_get_distinct_abort_ids() {
let mut h = Harness::new().await;
let mut tasks = Vec::new();
for _ in 0..2 {
tasks.push(tokio::spawn(
build_router(h.ctx.clone()).oneshot(request(json!({"hold":true}))),
));
}
let mut rids = Vec::new();
for _ in 0..2 {
rids.push(h.event("chat").await["rid"].clone());
}
assert_ne!(rids[0], rids[1]);
for task in tasks {
task.abort();
assert!(task.await.unwrap_err().is_cancelled());
}
for _ in 0..2 {
let aborted = h.event("abort").await;
let index = rids.iter().position(|rid| rid == &aborted["rid"]).unwrap();
rids.remove(index);
}
h.quiet().await;
}
@@ -230,7 +230,19 @@ async fn length_selects_plain_bucket_before_engine_selection() {
let response = app.clone().oneshot(request(body("hi"))).await.unwrap(); let response = app.clone().oneshot(request(body("hi"))).await.unwrap();
assert_eq!(response.status(), StatusCode::OK); assert_eq!(response.status(), StatusCode::OK);
let _ = response.into_body().collect().await.unwrap(); let _ = response.into_body().collect().await.unwrap();
assert!(short_worker.captured.lock().unwrap().last_body.is_some()); let forwarded: serde_json::Value = serde_json::from_slice(
short_worker
.captured
.lock()
.unwrap()
.last_body
.as_ref()
.unwrap(),
)
.unwrap();
assert!(crate::common::is_engine_shaped_rid(
forwarded["rid"].as_str().unwrap()
));
assert!(long_worker.captured.lock().unwrap().last_body.is_none()); assert!(long_worker.captured.lock().unwrap().last_body.is_none());
let response = app let response = app
@@ -293,6 +305,7 @@ async fn pd_picks_both_groups_from_selected_bucket_and_shares_bootstrap() {
let d: serde_json::Value = let d: serde_json::Value =
serde_json::from_slice(decode.captured.lock().unwrap().last_body.as_ref().unwrap()) serde_json::from_slice(decode.captured.lock().unwrap().last_body.as_ref().unwrap())
.unwrap(); .unwrap();
assert!(p.get("rid").is_none() && d.get("rid").is_none());
assert!(p["bootstrap_room"].is_number()); assert!(p["bootstrap_room"].is_number());
assert_eq!(p["bootstrap_room"], d["bootstrap_room"]); assert_eq!(p["bootstrap_room"], d["bootstrap_room"]);
let calls = policy.calls.lock().unwrap(); let calls = policy.calls.lock().unwrap();
@@ -39,9 +39,22 @@ pub struct MockWorker {
// Used in header_forwarding_test; not every test file reads captured headers. // Used in header_forwarding_test; not every test file reads captured headers.
#[allow(dead_code)] #[allow(dead_code)]
pub captured: Arc<Mutex<CapturedHeaders>>, pub captured: Arc<Mutex<CapturedHeaders>>,
#[allow(dead_code)]
pub abort_log: Arc<Mutex<Vec<Value>>>,
_shutdown: oneshot::Sender<()>, _shutdown: oneshot::Sender<()>,
} }
#[allow(dead_code)] // shared across all axum variants
fn abort_request_route<S>(log: Arc<Mutex<Vec<Value>>>) -> axum::routing::MethodRouter<S>
where
S: Clone + Send + Sync + 'static,
{
post(move |Json(body): Json<Value>| async move {
log.lock().unwrap().push(body);
StatusCode::OK
})
}
impl MockWorker { impl MockWorker {
/// Bind to a random port on 127.0.0.1 and start serving. /// Bind to a random port on 127.0.0.1 and start serving.
/// ///
@@ -50,6 +63,7 @@ impl MockWorker {
#[allow(dead_code)] // Only used by some test files. #[allow(dead_code)] // Only used by some test files.
pub async fn start(stream_chunks: Vec<&'static str>) -> Self { pub async fn start(stream_chunks: Vec<&'static str>) -> Self {
let captured = Arc::new(Mutex::new(CapturedHeaders::default())); let captured = Arc::new(Mutex::new(CapturedHeaders::default()));
let abort_log: Arc<Mutex<Vec<Value>>> = Arc::new(Mutex::new(Vec::new()));
let state = MockWorkerState { let state = MockWorkerState {
captured: captured.clone(), captured: captured.clone(),
stream_chunks: Arc::new(stream_chunks), stream_chunks: Arc::new(stream_chunks),
@@ -60,6 +74,7 @@ impl MockWorker {
let app = axum::Router::new() let app = axum::Router::new()
.route("/v1/chat/completions", post(chat)) .route("/v1/chat/completions", post(chat))
.route("/server_info", get(serve_tiny_server_info)) .route("/server_info", get(serve_tiny_server_info))
.route("/abort_request", abort_request_route(abort_log.clone()))
.with_state(state); .with_state(state);
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
@@ -77,6 +92,7 @@ impl MockWorker {
Self { Self {
url, url,
captured, captured,
abort_log,
_shutdown: tx, _shutdown: tx,
} }
} }
@@ -88,6 +104,7 @@ impl MockWorker {
#[allow(dead_code)] #[allow(dead_code)]
pub async fn start_hanging(delay: Duration) -> Self { pub async fn start_hanging(delay: Duration) -> Self {
let captured = Arc::new(Mutex::new(CapturedHeaders::default())); let captured = Arc::new(Mutex::new(CapturedHeaders::default()));
let abort_log: Arc<Mutex<Vec<Value>>> = Arc::new(Mutex::new(Vec::new()));
#[derive(Clone)] #[derive(Clone)]
struct HangState { struct HangState {
@@ -127,6 +144,7 @@ impl MockWorker {
let app = axum::Router::new() let app = axum::Router::new()
.route("/v1/chat/completions", post(hang_handler)) .route("/v1/chat/completions", post(hang_handler))
.route("/server_info", get(serve_tiny_server_info)) .route("/server_info", get(serve_tiny_server_info))
.route("/abort_request", abort_request_route(abort_log.clone()))
.with_state(state); .with_state(state);
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
@@ -144,6 +162,7 @@ impl MockWorker {
Self { Self {
url, url,
captured, captured,
abort_log,
_shutdown: tx, _shutdown: tx,
} }
} }
@@ -154,6 +173,7 @@ impl MockWorker {
#[allow(dead_code)] #[allow(dead_code)]
pub async fn start_slow_stream(chunks: Vec<&'static str>, delay: Duration) -> Self { pub async fn start_slow_stream(chunks: Vec<&'static str>, delay: Duration) -> Self {
let captured = Arc::new(Mutex::new(CapturedHeaders::default())); let captured = Arc::new(Mutex::new(CapturedHeaders::default()));
let abort_log: Arc<Mutex<Vec<Value>>> = Arc::new(Mutex::new(Vec::new()));
#[derive(Clone)] #[derive(Clone)]
struct SlowState { struct SlowState {
@@ -207,6 +227,7 @@ impl MockWorker {
let app = axum::Router::new() let app = axum::Router::new()
.route("/v1/chat/completions", post(slow_chat)) .route("/v1/chat/completions", post(slow_chat))
.route("/server_info", get(serve_tiny_server_info)) .route("/server_info", get(serve_tiny_server_info))
.route("/abort_request", abort_request_route(abort_log.clone()))
.with_state(state); .with_state(state);
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
@@ -224,6 +245,7 @@ impl MockWorker {
Self { Self {
url, url,
captured, captured,
abort_log,
_shutdown: tx, _shutdown: tx,
} }
} }
@@ -248,6 +270,7 @@ impl MockWorker {
partial_body_bytes: &'static [u8], partial_body_bytes: &'static [u8],
) -> Self { ) -> Self {
let captured = Arc::new(Mutex::new(CapturedHeaders::default())); let captured = Arc::new(Mutex::new(CapturedHeaders::default()));
let abort_log: Arc<Mutex<Vec<Value>>> = Arc::new(Mutex::new(Vec::new()));
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr: SocketAddr = listener.local_addr().unwrap(); let addr: SocketAddr = listener.local_addr().unwrap();
let url = format!("http://{addr}"); let url = format!("http://{addr}");
@@ -309,6 +332,7 @@ impl MockWorker {
Self { Self {
url, url,
captured, captured,
abort_log,
_shutdown: tx, _shutdown: tx,
} }
} }
@@ -319,6 +343,7 @@ impl MockWorker {
#[allow(dead_code)] #[allow(dead_code)]
pub async fn start_returning_error(status: StatusCode, body: Value) -> Self { pub async fn start_returning_error(status: StatusCode, body: Value) -> Self {
let captured = Arc::new(Mutex::new(CapturedHeaders::default())); let captured = Arc::new(Mutex::new(CapturedHeaders::default()));
let abort_log: Arc<Mutex<Vec<Value>>> = Arc::new(Mutex::new(Vec::new()));
let body_arc = Arc::new(body.to_string()); let body_arc = Arc::new(body.to_string());
#[derive(Clone)] #[derive(Clone)]
@@ -360,6 +385,7 @@ impl MockWorker {
let app = axum::Router::new() let app = axum::Router::new()
.route("/v1/chat/completions", post(error_handler)) .route("/v1/chat/completions", post(error_handler))
.route("/server_info", get(serve_tiny_server_info)) .route("/server_info", get(serve_tiny_server_info))
.route("/abort_request", abort_request_route(abort_log.clone()))
.with_state(state); .with_state(state);
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
@@ -377,6 +403,7 @@ impl MockWorker {
Self { Self {
url, url,
captured, captured,
abort_log,
_shutdown: tx, _shutdown: tx,
} }
} }
@@ -6,3 +6,11 @@
pub mod cache_aware_fixture; pub mod cache_aware_fixture;
pub mod mock_worker; pub mod mock_worker;
pub mod streaming; pub mod streaming;
#[allow(dead_code)] // not every test file inspects forwarded rids
pub fn is_engine_shaped_rid(rid: &str) -> bool {
rid.len() == 32
&& rid
.bytes()
.all(|b| b.is_ascii_hexdigit() && !b.is_ascii_uppercase())
}
@@ -72,6 +72,7 @@ async fn h2c_client_reaches_http2_only_worker() {
"/v1/chat/completions", "/v1/chat/completions",
&axum::http::HeaderMap::new(), &axum::http::HeaderMap::new(),
Bytes::from_static(b"{}"), Bytes::from_static(b"{}"),
None,
) )
.await .await
.expect("h2c client must reach an HTTP/2-only worker"); .expect("h2c client must reach an HTTP/2-only worker");
@@ -96,6 +97,7 @@ async fn http1_client_cannot_reach_http2_only_worker() {
"/v1/chat/completions", "/v1/chat/completions",
&axum::http::HeaderMap::new(), &axum::http::HeaderMap::new(),
Bytes::from_static(b"{}"), Bytes::from_static(b"{}"),
None,
) )
.await; .await;
assert!( assert!(
@@ -164,6 +166,7 @@ async fn h2c_client_streams_sse_from_http2_only_worker() {
&axum::http::HeaderMap::new(), &axum::http::HeaderMap::new(),
Bytes::from_static(b"{}"), Bytes::from_static(b"{}"),
None, None,
None,
Some(Box::new(move || { Some(Box::new(move || {
flag.store(true, std::sync::atomic::Ordering::SeqCst); flag.store(true, std::sync::atomic::Ordering::SeqCst);
})), })),
@@ -376,3 +376,61 @@ async fn pd_mode_prefill_5xx_does_not_poison_decode_response() {
let pv = parse_body(&prefill_body); let pv = parse_body(&prefill_body);
assert_eq!(bootstrap_port(&pv), Some(8997)); assert_eq!(bootstrap_port(&pv), Some(8997));
} }
fn streaming_chat_request() -> Request<Body> {
Request::builder()
.method("POST")
.uri("/v1/chat/completions")
.header("content-type", "application/json")
.body(Body::from(
serde_json::to_vec(&json!({
"model": "tiny",
"messages": [{"role": "user", "content": "hi"}],
"stream": true,
}))
.unwrap(),
))
.unwrap()
}
#[tokio::test]
async fn pd_mode_disconnect_does_not_abort_either_worker() {
let prefill = crate::common::mock_worker::MockWorker::start(vec![]).await;
let decode = crate::common::mock_worker::MockWorker::start_slow_stream(
vec!["data: a\n\n", "data: b\n\n", "data: c\n\n"],
Duration::from_millis(50),
)
.await;
let ctx = build_ctx(vec![
WorkerSpec {
id: WorkerId("p1".into()),
url: prefill.url.clone(),
mode: WorkerMode::Prefill,
model_ids: vec![ModelId("tiny".into())],
bootstrap_port: Some(8997),
},
WorkerSpec {
id: WorkerId("d1".into()),
url: decode.url.clone(),
mode: WorkerMode::Decode,
model_ids: vec![ModelId("tiny".into())],
bootstrap_port: None,
},
]);
let app = build_router(ctx);
let res = app.oneshot(streaming_chat_request()).await.unwrap();
assert_eq!(res.status(), StatusCode::OK);
use futures::StreamExt;
let mut data_stream = res.into_body().into_data_stream();
assert!(data_stream.next().await.is_some());
drop(data_stream);
tokio::time::sleep(Duration::from_millis(100)).await;
for worker in [&prefill, &decode] {
assert!(worker.abort_log.lock().unwrap().is_empty());
let body = await_captured_body(worker, Duration::from_secs(2), "PD worker").await;
assert!(parse_body(&body).get("rid").is_none());
}
}
@@ -106,9 +106,23 @@ fn without_forwarding(mut cfg: Config, policy: PolicyKind) -> Config {
cfg cfg
} }
fn without_minted_rid(mut body: Value) -> Value {
let rid = body
.as_object_mut()
.expect("a forwarded chat body is an object")
.remove("rid");
assert!(
rid.as_ref()
.and_then(Value::as_str)
.is_some_and(crate::common::is_engine_shaped_rid),
"plain mode must mint an abort rid; got {rid:?}",
);
body
}
async fn assert_forwarded_unchanged(ctx: &Arc<AppContext>, mock: &MockWorker, request: &Value) { async fn assert_forwarded_unchanged(ctx: &Arc<AppContext>, mock: &MockWorker, request: &Value) {
assert_eq!(send(Arc::clone(ctx), request.clone()).await, StatusCode::OK); assert_eq!(send(Arc::clone(ctx), request.clone()).await, StatusCode::OK);
assert_eq!(captured(mock), *request); assert_eq!(without_minted_rid(captured(mock)), *request);
assert!(!ctx assert!(!ctx
.metrics .metrics
.render() .render()
@@ -428,6 +442,6 @@ async fn kimi_ids_forward_with_engine_rendering_fallback() {
} else { } else {
assert!(ids.is_none()); assert!(ids.is_none());
} }
assert_eq!(captured(&mock), request); assert_eq!(without_minted_rid(captured(&mock)), request);
} }
} }
@@ -303,13 +303,26 @@ template <int NUM_TOP_K, int HOT_BUFFER_SIZE>
struct SmemLayout { struct SmemLayout {
static constexpr int HASH_SIZE = NUM_TOP_K * 2; static constexpr int HASH_SIZE = NUM_TOP_K * 2;
static constexpr int NUM_BUFFER_CHUNKS = (HOT_BUFFER_SIZE + WARP_SIZE - 1) / WARP_SIZE; static constexpr int NUM_BUFFER_CHUNKS = (HOT_BUFFER_SIZE + WARP_SIZE - 1) / WARP_SIZE;
// int32_t region: top_k_tokens + chunk_offset + evict_chunk_offset + hash_keys + total_hits + newest_hit // int32_t region: top_k_tokens + chunk offsets + hash keys + hit counters
static constexpr int TOTAL_INT32 = NUM_TOP_K + (NUM_BUFFER_CHUNKS + 1) + (NUM_BUFFER_CHUNKS + 1) + HASH_SIZE + 2; static constexpr int TOTAL_INT32 = NUM_TOP_K + (NUM_BUFFER_CHUNKS + 1) + (NUM_BUFFER_CHUNKS + 1) + HASH_SIZE + 2;
// int16_t region: lru_slots_out + hash_vals // int16_t region: lru_slots_out + hash_vals
static constexpr int TOTAL_INT16 = HOT_BUFFER_SIZE + HASH_SIZE; static constexpr int TOTAL_INT16 = HOT_BUFFER_SIZE + HASH_SIZE;
static constexpr size_t BYTES = TOTAL_INT32 * sizeof(int32_t) + TOTAL_INT16 * sizeof(int16_t); static constexpr size_t BYTES = TOTAL_INT32 * sizeof(int32_t) + TOTAL_INT16 * sizeof(int16_t);
}; };
template <int SPARSE_BLOCK_SIZE, bool TopKIsBlocks>
__device__ __forceinline__ int32_t resolve_selected_token(const int32_t* top_k, int32_t token_index) {
if constexpr (TopKIsBlocks) {
const int32_t block_index = top_k[token_index / SPARSE_BLOCK_SIZE];
if (block_index < 0) {
return -1;
}
return block_index * SPARSE_BLOCK_SIZE + token_index % SPARSE_BLOCK_SIZE;
} else {
return top_k[token_index];
}
}
// Each block processes one request // Each block processes one request
// req_pool_indices and seq_lens can each be int32_t or int64_t // req_pool_indices and seq_lens can each be int32_t or int64_t
// Layout: [HOT_BUFFER_SIZE slots for LRU] + [page_size slots for newest token] // Layout: [HOT_BUFFER_SIZE slots for LRU] + [page_size slots for newest token]
@@ -319,23 +332,28 @@ struct SmemLayout {
// false -> generic byte-stride: device + host both linear, stride = item_size_bytes // false -> generic byte-stride: device + host both linear, stride = item_size_bytes
// true -> DSv4 page-padded device + page-padded host (kvcacheio.cuh constants) // true -> DSv4 page-padded device + page-padded host (kvcacheio.cuh constants)
// //
// TopKIsBlocks makes the kernel consume block ids directly. It resolves token
// positions in registers and writes the flattened token-slot table expected by
// sparse attention without materializing an intermediate token-index tensor.
// RecordMissPlan records this step's miss plan (miss_src/dst = host/device loc // RecordMissPlan records this step's miss plan (miss_src/dst = host/device loc
// per miss, miss_count per request) for shared-index skip layers to replay via // per miss, miss_count per request) for shared-index skip layers to replay via
// copy_cache_planned_kernel. SkipIO elides only the KV byte movement (timing // copy_cache_planned_kernel. SkipIO elides only the KV byte movement (timing
// probe; output is garbage). Both are compile-time flags so the production // probe; output is garbage). These are compile-time flags, so inactive paths
// (false, false) instantiation stays byte-identical. // are removed from each specialization.
template < template <
int BLOCK_SIZE, int BLOCK_SIZE,
int NUM_TOP_K, int NUM_TOP_K,
int HOT_BUFFER_SIZE, int HOT_BUFFER_SIZE,
bool IsMLA, bool IsMLA,
bool IsDsv4Layout, bool IsDsv4Layout,
int SPARSE_BLOCK_SIZE,
bool TopKIsBlocks,
bool RecordMissPlan, bool RecordMissPlan,
bool SkipIO, bool SkipIO,
typename SeqLensT, typename SeqLensT,
typename ReqPoolIndicesT> typename ReqPoolIndicesT>
__global__ void load_cache_to_device_buffer_kernel( __global__ void load_cache_to_device_buffer_kernel(
const int32_t* __restrict__ top_k_tokens, const int32_t* __restrict__ top_k,
int32_t* __restrict__ device_buffer_tokens, int32_t* __restrict__ device_buffer_tokens,
const int64_t* __restrict__ host_cache_locs, const int64_t* __restrict__ host_cache_locs,
const int32_t* __restrict__ device_buffer_locs, const int32_t* __restrict__ device_buffer_locs,
@@ -351,7 +369,7 @@ __global__ void load_cache_to_device_buffer_kernel(
int64_t buffer_stride_0, int64_t buffer_stride_0,
int64_t host_stride, int64_t host_stride,
int64_t lru_slot_stride_0, int64_t lru_slot_stride_0,
int64_t top_k_tokens_stride, int64_t top_k_stride,
int64_t top_k_device_locs_stride, int64_t top_k_device_locs_stride,
int64_t page_size, int64_t page_size,
int64_t item_size_bytes, int64_t item_size_bytes,
@@ -360,9 +378,12 @@ __global__ void load_cache_to_device_buffer_kernel(
int32_t* __restrict__ miss_count_out, int32_t* __restrict__ miss_count_out,
int64_t plan_stride) { int64_t plan_stride) {
static_assert(!IsDsv4Layout || IsMLA, "DSv4 page-padded layout is K-only (MLA)."); static_assert(!IsDsv4Layout || IsMLA, "DSv4 page-padded layout is K-only (MLA).");
// todo hisparse: support page wise sparsity static_assert(SPARSE_BLOCK_SIZE > 0, "SPARSE_BLOCK_SIZE must be positive.");
// Cache residency and LRU replacement remain token-granular even when the
// sparse-attention selection arrives as block ids.
constexpr int NUM_TOP_K_TOKENS = NUM_TOP_K * (TopKIsBlocks ? SPARSE_BLOCK_SIZE : 1);
constexpr int NUM_WARPS = BLOCK_SIZE / WARP_SIZE; constexpr int NUM_WARPS = BLOCK_SIZE / WARP_SIZE;
constexpr int NUM_TOKEN_CHUNKS = (NUM_TOP_K + WARP_SIZE - 1) / WARP_SIZE; constexpr int NUM_TOKEN_CHUNKS = (NUM_TOP_K_TOKENS + WARP_SIZE - 1) / WARP_SIZE;
constexpr int NUM_BUFFER_CHUNKS = (HOT_BUFFER_SIZE + WARP_SIZE - 1) / WARP_SIZE; constexpr int NUM_BUFFER_CHUNKS = (HOT_BUFFER_SIZE + WARP_SIZE - 1) / WARP_SIZE;
const int bid = blockIdx.x; const int bid = blockIdx.x;
@@ -372,7 +393,7 @@ __global__ void load_cache_to_device_buffer_kernel(
// CUDA graph pads the batch to a captured size. Keep padded output rows // CUDA graph pads the batch to a captured size. Keep padded output rows
// invalid without a separate fill kernel. // invalid without a separate fill kernel.
if (bid >= num_real_reqs[0]) { if (bid >= num_real_reqs[0]) {
for (int i = tid; i < NUM_TOP_K; i += BLOCK_SIZE) { for (int i = tid; i < NUM_TOP_K_TOKENS; i += BLOCK_SIZE) {
req_top_k_device_locs[i] = -1; req_top_k_device_locs[i] = -1;
} }
return; return;
@@ -386,7 +407,7 @@ __global__ void load_cache_to_device_buffer_kernel(
const int64_t seq_len = seq_lens[bid]; const int64_t seq_len = seq_lens[bid];
// Calculate offsets for this request // Calculate offsets for this request
const int32_t* req_top_k_tokens = top_k_tokens + bid * top_k_tokens_stride; const int32_t* req_top_k = top_k + bid * top_k_stride;
const int64_t buffer_offset = rid * buffer_stride_0; const int64_t buffer_offset = rid * buffer_stride_0;
int32_t* req_device_buffer_tokens = device_buffer_tokens + buffer_offset; int32_t* req_device_buffer_tokens = device_buffer_tokens + buffer_offset;
@@ -396,14 +417,16 @@ __global__ void load_cache_to_device_buffer_kernel(
// Fast path: short sequences have all tokens in the device buffer in order. // Fast path: short sequences have all tokens in the device buffer in order.
if (seq_len <= HOT_BUFFER_SIZE) { if (seq_len <= HOT_BUFFER_SIZE) {
const int count = (seq_len < NUM_TOP_K) ? static_cast<int>(seq_len) : NUM_TOP_K; const int count = (seq_len < NUM_TOP_K_TOKENS) ? static_cast<int>(seq_len) : NUM_TOP_K_TOKENS;
for (int i = tid; i < NUM_TOP_K; i += BLOCK_SIZE) { for (int i = tid; i < NUM_TOP_K_TOKENS; i += BLOCK_SIZE) {
int32_t device_loc = -1; int32_t device_loc = -1;
if (i < count) { const int32_t token_pos = resolve_selected_token<SPARSE_BLOCK_SIZE, TopKIsBlocks>(req_top_k, i);
int32_t token_pos = req_top_k_tokens[i]; if constexpr (TopKIsBlocks) {
if (token_pos >= 0) { if (token_pos >= 0 && token_pos < seq_len) {
device_loc = req_device_buffer_locs[token_pos]; device_loc = req_device_buffer_locs[token_pos];
} }
} else if (i < count && token_pos >= 0) {
device_loc = req_device_buffer_locs[token_pos];
} }
req_top_k_device_locs[i] = device_loc; req_top_k_device_locs[i] = device_loc;
} }
@@ -418,21 +441,21 @@ __global__ void load_cache_to_device_buffer_kernel(
// Dynamic shared memory layout: int32_t arrays first, then int16_t arrays. // Dynamic shared memory layout: int32_t arrays first, then int16_t arrays.
extern __shared__ char smem_raw[]; extern __shared__ char smem_raw[];
using Layout = SmemLayout<NUM_TOP_K, HOT_BUFFER_SIZE>; using Layout = SmemLayout<NUM_TOP_K_TOKENS, HOT_BUFFER_SIZE>;
constexpr int HASH_SIZE = Layout::HASH_SIZE; constexpr int HASH_SIZE = Layout::HASH_SIZE;
int32_t* smem_i32 = reinterpret_cast<int32_t*>(smem_raw); int32_t* smem_i32 = reinterpret_cast<int32_t*>(smem_raw);
// Top-k token positions; reused as miss-token scratch in the copy phase // Top-k token positions; reused as miss-token scratch in the copy phase
int32_t* s_top_k_tokens = smem_i32; int32_t* s_top_k_tokens = smem_i32;
// Prefix-sum offsets for hit counting and miss counting // Prefix-sum offsets for hit counting and miss counting
int32_t* s_chunk_offset = s_top_k_tokens + NUM_TOP_K; int32_t* s_chunk_offset = s_top_k_tokens + NUM_TOP_K_TOKENS;
// Prefix-sum offsets for evictable counting // Prefix-sum offsets for evictable counting
int32_t* s_evict_chunk_offset = s_chunk_offset + (NUM_BUFFER_CHUNKS + 1); int32_t* s_evict_chunk_offset = s_chunk_offset + (NUM_BUFFER_CHUNKS + 1);
// Open-addressing hash table: top-k token_id -> top-k index (keys) // Open-addressing hash table: top-k token_id -> top-k index (keys)
int32_t* s_hash_keys = s_evict_chunk_offset + (NUM_BUFFER_CHUNKS + 1); int32_t* s_hash_keys = s_evict_chunk_offset + (NUM_BUFFER_CHUNKS + 1);
// Scalar counters // Scalar counters
int32_t& s_total_hits = s_hash_keys[HASH_SIZE]; int32_t& s_total_hits = s_hash_keys[HASH_SIZE];
int32_t& s_newest_hit = s_hash_keys[HASH_SIZE + 1]; int32_t& s_total_misses = s_hash_keys[HASH_SIZE + 1];
int16_t* smem_i16 = reinterpret_cast<int16_t*>(smem_i32 + Layout::TOTAL_INT32); int16_t* smem_i16 = reinterpret_cast<int16_t*>(smem_i32 + Layout::TOTAL_INT32);
// Compacted slot ordering: [hits fwd-> ... <-evictables bwd] // Compacted slot ordering: [hits fwd-> ... <-evictables bwd]
@@ -443,7 +466,7 @@ __global__ void load_cache_to_device_buffer_kernel(
// Initialize shared memory: counters, hash table, prefix-sum offsets. // Initialize shared memory: counters, hash table, prefix-sum offsets.
if (tid == 0) { if (tid == 0) {
s_total_hits = 0; s_total_hits = 0;
s_newest_hit = 0; s_total_misses = 0;
} }
for (int i = tid; i < HASH_SIZE; i += BLOCK_SIZE) { for (int i = tid; i < HASH_SIZE; i += BLOCK_SIZE) {
s_hash_keys[i] = HASH_EMPTY; s_hash_keys[i] = HASH_EMPTY;
@@ -458,14 +481,20 @@ __global__ void load_cache_to_device_buffer_kernel(
const int32_t newest_token = seq_len - 1; const int32_t newest_token = seq_len - 1;
// Insert top-k tokens into shared-memory hash table. // Insert top-k tokens into shared-memory hash table.
for (int i = tid; i < NUM_TOP_K; i += BLOCK_SIZE) { for (int i = tid; i < NUM_TOP_K_TOKENS; i += BLOCK_SIZE) {
int32_t token_idx = req_top_k_tokens[i]; const int32_t token_idx = resolve_selected_token<SPARSE_BLOCK_SIZE, TopKIsBlocks>(req_top_k, i);
if constexpr (TopKIsBlocks) {
if (token_idx < 0 || token_idx >= seq_len) {
s_top_k_tokens[i] = TOKEN_HIT;
req_top_k_device_locs[i] = -1;
continue;
}
}
if (token_idx == newest_token) { if (token_idx == newest_token) {
// If topk includes the latest token, bind its canonical occurrence to newest_slot (at HOT_BUFFER_SIZE) and mark // If topk includes the latest token, bind its canonical occurrence to newest_slot (at HOT_BUFFER_SIZE) and mark
// it as a hit. newest_slot is at the first position of the extra page, excluded from LRU tracking. // it as a hit. newest_slot is at the first position of the extra page, excluded from LRU tracking.
s_top_k_tokens[i] = TOKEN_HIT; s_top_k_tokens[i] = TOKEN_HIT;
req_top_k_device_locs[i] = req_device_buffer_locs[newest_slot]; req_top_k_device_locs[i] = req_device_buffer_locs[newest_slot];
s_newest_hit = 1;
} else { } else {
int slot = hash_slot(token_idx, HASH_SIZE); int slot = hash_slot(token_idx, HASH_SIZE);
while (true) { while (true) {
@@ -580,7 +609,7 @@ __global__ void load_cache_to_device_buffer_kernel(
const int chunk_token_start = chunk_idx * WARP_SIZE; const int chunk_token_start = chunk_idx * WARP_SIZE;
const int my_token_idx = chunk_token_start + lane_id; const int my_token_idx = chunk_token_start + lane_id;
const bool has_valid_token = has_valid_chunk && (my_token_idx < NUM_TOP_K); const bool has_valid_token = has_valid_chunk && (my_token_idx < NUM_TOP_K_TOKENS);
int32_t my_token = 0; int32_t my_token = 0;
bool is_miss = false; bool is_miss = false;
@@ -611,6 +640,9 @@ __global__ void load_cache_to_device_buffer_kernel(
#else #else
total_misses = warp_inclusive_scan(s_chunk_offset, lane_id, chunk_idx + 1, NUM_TOKEN_CHUNKS + 1, total_misses); total_misses = warp_inclusive_scan(s_chunk_offset, lane_id, chunk_idx + 1, NUM_TOKEN_CHUNKS + 1, total_misses);
#endif #endif
if (tid == 0) {
s_total_misses = total_misses;
}
} }
__syncthreads(); __syncthreads();
@@ -632,7 +664,7 @@ __global__ void load_cache_to_device_buffer_kernel(
} }
__syncthreads(); __syncthreads();
total_misses = NUM_TOP_K - s_total_hits - s_newest_hit; total_misses = s_total_misses;
if constexpr (RecordMissPlan) { if constexpr (RecordMissPlan) {
if (tid == 0) { if (tid == 0) {
miss_count_out[bid] = total_misses; miss_count_out[bid] = total_misses;
@@ -695,10 +727,12 @@ template <
int HOT_BUFFER_SIZE, int HOT_BUFFER_SIZE,
bool IsMLA, bool IsMLA,
bool IsDsv4Layout, bool IsDsv4Layout,
int SPARSE_BLOCK_SIZE,
bool TopKIsBlocks,
bool RecordMissPlan, bool RecordMissPlan,
bool SkipIO> bool SkipIO>
void load_cache_to_device_buffer( void load_cache_to_device_buffer(
tvm::ffi::TensorView top_k_tokens, tvm::ffi::TensorView top_k,
tvm::ffi::TensorView device_buffer_tokens, tvm::ffi::TensorView device_buffer_tokens,
tvm::ffi::TensorView host_cache_locs, tvm::ffi::TensorView host_cache_locs,
tvm::ffi::TensorView device_buffer_locs, tvm::ffi::TensorView device_buffer_locs,
@@ -718,7 +752,8 @@ void load_cache_to_device_buffer(
tvm::ffi::TensorView miss_count_out) { tvm::ffi::TensorView miss_count_out) {
using namespace host; using namespace host;
const int64_t bs = top_k_tokens.shape()[0]; constexpr int NUM_TOP_K_TOKENS = NUM_TOP_K * (TopKIsBlocks ? SPARSE_BLOCK_SIZE : 1);
const int64_t bs = top_k.shape()[0];
const int64_t host_stride = host_cache_locs.shape()[1]; const int64_t host_stride = host_cache_locs.shape()[1];
// Miss-plan side outputs; 0-dim sentinels when RecordMissPlan is false. // Miss-plan side outputs; 0-dim sentinels when RecordMissPlan is false.
int64_t* const miss_src_ptr = RecordMissPlan ? static_cast<int64_t*>(miss_src_out.data_ptr()) : nullptr; int64_t* const miss_src_ptr = RecordMissPlan ? static_cast<int64_t*>(miss_src_out.data_ptr()) : nullptr;
@@ -730,9 +765,9 @@ void load_cache_to_device_buffer(
} }
const int64_t buffer_stride_0 = device_buffer_tokens.strides()[0]; const int64_t buffer_stride_0 = device_buffer_tokens.strides()[0];
const int64_t lru_slot_stride_0 = lru_slots.strides()[0]; const int64_t lru_slot_stride_0 = lru_slots.strides()[0];
const int64_t top_k_tokens_stride = top_k_tokens.strides()[0]; const int64_t top_k_stride = top_k.strides()[0];
const int64_t top_k_device_locs_stride = top_k_device_locs.strides()[0]; const int64_t top_k_device_locs_stride = top_k_device_locs.strides()[0];
const auto kernel_device = top_k_tokens.device(); const auto kernel_device = top_k.device();
const auto device = LaunchKernel::resolve_device(kernel_device); const auto device = LaunchKernel::resolve_device(kernel_device);
const void* const host_cache_k_ptr = runtime::get_device_accessible_ptr(host_cache_k); const void* const host_cache_k_ptr = runtime::get_device_accessible_ptr(host_cache_k);
const void* const host_cache_v_ptr = const void* const host_cache_v_ptr =
@@ -741,7 +776,7 @@ void load_cache_to_device_buffer(
// Generic lambda: int32/int64 kernel variants are compiled for both // Generic lambda: int32/int64 kernel variants are compiled for both
// seq_lens and req_pool_indices; the correct combo is selected at runtime. // seq_lens and req_pool_indices; the correct combo is selected at runtime.
auto launch = [&](auto kernel_fn, const auto* seq_lens_ptr, const auto* req_pool_indices_ptr) { auto launch = [&](auto kernel_fn, const auto* seq_lens_ptr, const auto* req_pool_indices_ptr) {
constexpr size_t smem_bytes = SmemLayout<NUM_TOP_K, HOT_BUFFER_SIZE>::BYTES; constexpr size_t smem_bytes = SmemLayout<NUM_TOP_K_TOKENS, HOT_BUFFER_SIZE>::BYTES;
#ifndef USE_ROCM #ifndef USE_ROCM
if constexpr (smem_bytes > 48u * 1024u) { if constexpr (smem_bytes > 48u * 1024u) {
cudaFuncSetAttribute(kernel_fn, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_bytes); cudaFuncSetAttribute(kernel_fn, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_bytes);
@@ -749,7 +784,7 @@ void load_cache_to_device_buffer(
#endif #endif
LaunchKernel(bs, BLOCK_SIZE, device, smem_bytes)( LaunchKernel(bs, BLOCK_SIZE, device, smem_bytes)(
kernel_fn, kernel_fn,
static_cast<const int32_t*>(top_k_tokens.data_ptr()), static_cast<const int32_t*>(top_k.data_ptr()),
static_cast<int32_t*>(device_buffer_tokens.data_ptr()), static_cast<int32_t*>(device_buffer_tokens.data_ptr()),
static_cast<const int64_t*>(host_cache_locs.data_ptr()), static_cast<const int64_t*>(host_cache_locs.data_ptr()),
static_cast<const int32_t*>(device_buffer_locs.data_ptr()), static_cast<const int32_t*>(device_buffer_locs.data_ptr()),
@@ -765,7 +800,7 @@ void load_cache_to_device_buffer(
buffer_stride_0, buffer_stride_0,
host_stride, host_stride,
lru_slot_stride_0, lru_slot_stride_0,
top_k_tokens_stride, top_k_stride,
top_k_device_locs_stride, top_k_device_locs_stride,
page_size, page_size,
item_size_bytes, item_size_bytes,
@@ -788,6 +823,8 @@ void load_cache_to_device_buffer(
HOT_BUFFER_SIZE, HOT_BUFFER_SIZE,
IsMLA, IsMLA,
IsDsv4Layout, IsDsv4Layout,
SPARSE_BLOCK_SIZE,
TopKIsBlocks,
RecordMissPlan, RecordMissPlan,
SkipIO, SkipIO,
int64_t, int64_t,
@@ -802,6 +839,8 @@ void load_cache_to_device_buffer(
HOT_BUFFER_SIZE, HOT_BUFFER_SIZE,
IsMLA, IsMLA,
IsDsv4Layout, IsDsv4Layout,
SPARSE_BLOCK_SIZE,
TopKIsBlocks,
RecordMissPlan, RecordMissPlan,
SkipIO, SkipIO,
int64_t, int64_t,
@@ -816,6 +855,8 @@ void load_cache_to_device_buffer(
HOT_BUFFER_SIZE, HOT_BUFFER_SIZE,
IsMLA, IsMLA,
IsDsv4Layout, IsDsv4Layout,
SPARSE_BLOCK_SIZE,
TopKIsBlocks,
RecordMissPlan, RecordMissPlan,
SkipIO, SkipIO,
int32_t, int32_t,
@@ -830,6 +871,8 @@ void load_cache_to_device_buffer(
HOT_BUFFER_SIZE, HOT_BUFFER_SIZE,
IsMLA, IsMLA,
IsDsv4Layout, IsDsv4Layout,
SPARSE_BLOCK_SIZE,
TopKIsBlocks,
RecordMissPlan, RecordMissPlan,
SkipIO, SkipIO,
int32_t, int32_t,
@@ -23,6 +23,7 @@ from ..common.utils import (
"BLOCK_SIZE_T": lambda args: triton.next_power_of_2(args["max_topk"]), "BLOCK_SIZE_T": lambda args: triton.next_power_of_2(args["max_topk"]),
"HAS_SINK": lambda args: args["sink_ptr"] is not None, "HAS_SINK": lambda args: args["sink_ptr"] is not None,
"BATCH_SIZE_BUCKET": lambda args: triton.next_power_of_2(args["batch_size"]), "BATCH_SIZE_BUCKET": lambda args: triton.next_power_of_2(args["batch_size"]),
"HAS_HISPARSE_SLOTS": lambda args: args["hisparse_slots_ptr"] is not None,
} }
) )
@triton.autotune( @triton.autotune(
@@ -43,6 +44,7 @@ def _gqa_share_sparse_decode_kernel(
idx_ptr, # topk index: qh x b x topk idx_ptr, # topk index: qh x b x topk
o_ptr, # O partial: c x b x qh x d o_ptr, # O partial: c x b x qh x d
lse_ptr, # lse partial: c x b x qh lse_ptr, # lse partial: c x b x qh
hisparse_slots_ptr, # pre-resolved device slots: kh x b x (topk * block)
seq_lens, seq_lens,
slot_ids, slot_ids,
# shape # shape
@@ -52,6 +54,8 @@ def _gqa_share_sparse_decode_kernel(
head_dim, head_dim,
max_topk, max_topk,
max_kv_len, max_kv_len,
hisparse_slots_stride_h,
hisparse_slots_stride_b,
# sm_scale # sm_scale
sm_scale, sm_scale,
# per-tensor KV dequant scales (1.0 when the cache is unit-scaled) # per-tensor KV dequant scales (1.0 when the cache is unit-scaled)
@@ -89,6 +93,7 @@ def _gqa_share_sparse_decode_kernel(
NUM_TOPK_CHUNKS: tl.constexpr, NUM_TOPK_CHUNKS: tl.constexpr,
HAS_SINK: tl.constexpr, HAS_SINK: tl.constexpr,
IS_FP8: tl.constexpr, IS_FP8: tl.constexpr,
HAS_HISPARSE_SLOTS: tl.constexpr,
): ):
# decode program ids: split-K over the topk dimension to give every SM # decode program ids: split-K over the topk dimension to give every SM
# something to do at small batch. pid(0) folds (batch, chunk) together so # something to do at small batch. pid(0) folds (batch, chunk) together so
@@ -161,18 +166,30 @@ def _gqa_share_sparse_decode_kernel(
# only iterate over this chunk's topk slice. the load must respect the # only iterate over this chunk's topk slice. the load must respect the
# per-chunk start offset. # per-chunk start offset.
cur_idx_ptr = idx_base + chunk_start_topk * stride_ti_t cur_idx_ptr = idx_base + chunk_start_topk * stride_ti_t
hisparse_topk_counter = chunk_start_topk
for _ in tl.range(chunk_start_topk, chunk_end_topk): for _ in tl.range(chunk_start_topk, chunk_end_topk):
# load index # load index
c = tl.load(cur_idx_ptr).to(tl.int32) * BLOCK_SIZE_N c = tl.load(cur_idx_ptr).to(tl.int32) * BLOCK_SIZE_N
cur_idx_ptr = cur_idx_ptr + stride_ti_t cur_idx_ptr = cur_idx_ptr + stride_ti_t
# resolve slots for this block via req_to_token
pos = c + off_n pos = c + off_n
pos_mask = pos < seq_len pos_mask = pos < seq_len
slots = tl.load( if HAS_HISPARSE_SLOTS:
req_to_token_ptr + sid * stride_r2t_b + pos, slots = tl.load(
mask=pos_mask, hisparse_slots_ptr
other=0, + pid_kh * hisparse_slots_stride_h
).to(tl.int64) + pid_b * hisparse_slots_stride_b
+ hisparse_topk_counter * BLOCK_SIZE_N
+ off_n,
mask=off_n < BLOCK_SIZE_N,
other=0,
).to(tl.int64)
hisparse_topk_counter = hisparse_topk_counter + 1
else:
slots = tl.load(
req_to_token_ptr + sid * stride_r2t_b + pos,
mask=pos_mask,
other=0,
).to(tl.int64)
slots = (slots + max_slots) % max_slots # safety against negative slots = (slots + max_slots) % max_slots # safety against negative
# load K as (head_dim, BLOCK_SIZE_N) via indirect addressing # load K as (head_dim, BLOCK_SIZE_N) via indirect addressing
k_off = ( k_off = (
@@ -321,6 +338,7 @@ def flash_decode_with_gqa_share_sparse(
q_scale: Optional[float] = None, q_scale: Optional[float] = None,
k_scale: Optional[float] = None, k_scale: Optional[float] = None,
v_scale: Optional[float] = None, v_scale: Optional[float] = None,
hisparse_slots: Optional[torch.Tensor] = None,
) -> torch.Tensor: ) -> torch.Tensor:
triton.set_allocator(robust_allocator) triton.set_allocator(robust_allocator)
is_fp8 = check_sparse_kv_fp8(q, k_cache, v_cache, label="decode") is_fp8 = check_sparse_kv_fp8(q, k_cache, v_cache, label="decode")
@@ -384,6 +402,7 @@ def flash_decode_with_gqa_share_sparse(
topk_idx, topk_idx,
o_partial, o_partial,
lse_partial, lse_partial,
hisparse_slots,
seq_lens, seq_lens,
slot_ids, slot_ids,
max_slots, max_slots,
@@ -392,6 +411,8 @@ def flash_decode_with_gqa_share_sparse(
head_dim, head_dim,
max_topk, max_topk,
max_kv_len, max_kv_len,
hisparse_slots.stride(0) if hisparse_slots is not None else 0,
hisparse_slots.stride(1) if hisparse_slots is not None else 0,
sm_scale, sm_scale,
k_scale, k_scale,
v_scale, v_scale,
@@ -28,6 +28,7 @@ from ..common.utils import (
"BLOCK_SIZE_T": lambda args: triton.next_power_of_2(args["max_topk"]), "BLOCK_SIZE_T": lambda args: triton.next_power_of_2(args["max_topk"]),
"BLOCK_SIZE_QH": lambda args: args["BLOCK_SIZE_Q"] * args["BLOCK_SIZE_H"], "BLOCK_SIZE_QH": lambda args: args["BLOCK_SIZE_Q"] * args["BLOCK_SIZE_H"],
"HAS_SINK": lambda args: args["sink_ptr"] is not None, "HAS_SINK": lambda args: args["sink_ptr"] is not None,
"HAS_LOC_MAPPING": lambda args: args["loc_mapping_ptr"] is not None,
} }
) )
@triton.autotune( @triton.autotune(
@@ -55,6 +56,7 @@ def _gqa_share_sparse_fwd_kernel(
t_ptr, # topk_idx: kh x n x k t_ptr, # topk_idx: kh x n x k
o_ptr, # O: n x h x d o_ptr, # O: n x h x d
req_to_token_ptr, # req_to_token: max_reqs x max_kv_len req_to_token_ptr, # req_to_token: max_reqs x max_kv_len
loc_mapping_ptr, # logical slot to HiSparse device slot
# seqlens # seqlens
cu_seqlens_q, cu_seqlens_q,
cu_seqblocks_q, cu_seqblocks_q,
@@ -106,6 +108,7 @@ def _gqa_share_sparse_fwd_kernel(
HAS_SINK: tl.constexpr, HAS_SINK: tl.constexpr,
USE_TMA: tl.constexpr, USE_TMA: tl.constexpr,
IS_FP8: tl.constexpr, IS_FP8: tl.constexpr,
HAS_LOC_MAPPING: tl.constexpr,
): ):
sm_scale_log2e = sm_scale * 1.4426950409 sm_scale_log2e = sm_scale * 1.4426950409
# get batch id and head id # get batch id and head id
@@ -199,6 +202,12 @@ def _gqa_share_sparse_fwd_kernel(
mask=pos_mask, mask=pos_mask,
other=0, other=0,
).to(tl.int64) ).to(tl.int64)
if HAS_LOC_MAPPING:
slots = tl.load(
loc_mapping_ptr + slots,
mask=pos_mask,
other=0,
).to(tl.int64)
slots = (slots + max_slots) % max_slots # safety against negative slots = (slots + max_slots) % max_slots # safety against negative
# k shape: [BLOCK_SIZE_KD, BLOCK_SIZE_K] (transposed for tl.dot) # k shape: [BLOCK_SIZE_KD, BLOCK_SIZE_K] (transposed for tl.dot)
k = tl.load( k = tl.load(
@@ -289,6 +298,7 @@ def flash_prefill_with_gqa_share_sparse(
q_scale: Optional[float] = None, q_scale: Optional[float] = None,
k_scale: Optional[float] = None, k_scale: Optional[float] = None,
v_scale: Optional[float] = None, v_scale: Optional[float] = None,
loc_mapping: Optional[torch.Tensor] = None,
) -> torch.Tensor: ) -> torch.Tensor:
triton.set_allocator(robust_allocator) triton.set_allocator(robust_allocator)
is_fp8 = check_sparse_kv_fp8(q, k_cache, v_cache, label="prefill") is_fp8 = check_sparse_kv_fp8(q, k_cache, v_cache, label="prefill")
@@ -340,6 +350,7 @@ def flash_prefill_with_gqa_share_sparse(
topk_idx, topk_idx,
o, o,
req_to_token, req_to_token,
loc_mapping,
cu_seqlens, cu_seqlens,
cu_seqblocks_q, cu_seqblocks_q,
seq_lens, seq_lens,
+84 -1
View File
@@ -174,6 +174,8 @@ def _jit_sparse_module(
hot_buffer_size: int, hot_buffer_size: int,
is_mla: bool = False, is_mla: bool = False,
is_dsv4_layout: bool = False, is_dsv4_layout: bool = False,
top_k_block_size: int = 1,
top_k_is_blocks: bool = False,
record_miss_plan: bool = False, record_miss_plan: bool = False,
skip_io: bool = False, skip_io: bool = False,
) -> Module: ) -> Module:
@@ -185,6 +187,8 @@ def _jit_sparse_module(
hot_buffer_size, hot_buffer_size,
is_mla, is_mla,
is_dsv4_layout, is_dsv4_layout,
top_k_block_size,
top_k_is_blocks,
record_miss_plan, record_miss_plan,
skip_io, skip_io,
) )
@@ -195,6 +199,8 @@ def _jit_sparse_module(
hot_buffer_size, hot_buffer_size,
is_mla, is_mla,
is_dsv4_layout, is_dsv4_layout,
top_k_block_size,
top_k_is_blocks,
record_miss_plan, record_miss_plan,
skip_io, skip_io,
) )
@@ -308,7 +314,7 @@ def _load_cache_to_device_buffer_mla(
skip_io=skip_io, skip_io=skip_io,
) )
empty = torch.empty(0) empty = torch.empty(0, device=top_k_tokens.device)
if num_real_reqs is None: if num_real_reqs is None:
num_real_reqs = torch.tensor( num_real_reqs = torch.tensor(
@@ -399,6 +405,83 @@ def load_cache_to_device_buffer_mla(
) )
def load_blocks_to_device_buffer_mha(
top_k_blocks: torch.Tensor,
device_buffer_tokens: torch.Tensor,
host_cache_locs: torch.Tensor,
device_buffer_locs: torch.Tensor,
host_cache_k: torch.Tensor,
host_cache_v: torch.Tensor,
device_buffer_k: torch.Tensor,
device_buffer_v: torch.Tensor,
top_k_device_locs: torch.Tensor,
req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor,
lru_slots: torch.Tensor,
item_size_bytes: int,
hot_buffer_size: int,
sparse_block_size: int,
page_size: int = 1,
block_size: int = 256,
num_real_reqs: torch.Tensor | None = None,
skip_io: bool = False,
) -> None:
"""Swap block-selected MHA K/V into the HiSparse device pool."""
num_top_k_blocks = top_k_blocks.size(1)
num_top_k_tokens = num_top_k_blocks * sparse_block_size
assert hot_buffer_size >= num_top_k_tokens, (
f"hot_buffer_size ({hot_buffer_size}) must be >= selected tokens "
f"({num_top_k_tokens})"
)
assert top_k_device_locs.size(1) >= num_top_k_tokens
k_stride = host_cache_k.stride(0) * host_cache_k.element_size()
v_stride = host_cache_v.stride(0) * host_cache_v.element_size()
assert k_stride == v_stride == item_size_bytes, (
"K/V token strides must equal item_size_bytes: "
f"k_stride={k_stride}, v_stride={v_stride}, "
f"item_size_bytes={item_size_bytes}"
)
module = _jit_sparse_module(
item_size_bytes,
block_size,
num_top_k_blocks,
hot_buffer_size,
is_mla=False,
is_dsv4_layout=False,
top_k_block_size=sparse_block_size,
top_k_is_blocks=True,
record_miss_plan=False,
skip_io=skip_io,
)
empty = torch.empty(0, device=top_k_blocks.device)
if num_real_reqs is None:
num_real_reqs = torch.tensor(
[top_k_blocks.size(0)], dtype=torch.int32, device=top_k_blocks.device
)
module.load_cache_to_device_buffer(
top_k_blocks,
device_buffer_tokens,
host_cache_locs,
device_buffer_locs,
host_cache_k,
host_cache_v,
device_buffer_k,
device_buffer_v,
top_k_device_locs,
req_pool_indices,
seq_lens,
lru_slots,
num_real_reqs,
page_size,
item_size_bytes,
empty,
empty,
empty,
)
def copy_cache_planned_mla( def copy_cache_planned_mla(
*, *,
miss_src: torch.Tensor, miss_src: torch.Tensor,
+139 -4
View File
@@ -18,10 +18,99 @@ from sglang.srt.layers.attention.dsa.utils import is_dsa_prefill_cp_interleave
from sglang.srt.layers.dp_attention import is_allocation_symmetric from sglang.srt.layers.dp_attention import is_allocation_symmetric
from sglang.srt.layers.utils.common import strict_contiguous from sglang.srt.layers.utils.common import strict_contiguous
from sglang.srt.runtime_context import get_parallel, get_platform from sglang.srt.runtime_context import get_parallel, get_platform
from sglang.srt.utils import is_gfx95_supported, is_hip
from sglang.srt.utils.common import is_gfx1250_supported from sglang.srt.utils.common import is_gfx1250_supported
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
_AITER_MHC_RUNTIME_DISABLED = False
_AITER_MHC_ACTIVE_LOGGED = False
def _use_aiter_mhc() -> bool:
return (
not _AITER_MHC_RUNTIME_DISABLED
and is_gfx95_supported()
and envs.SGLANG_USE_AITER.get()
)
def _try_aiter_mhc_pre(
residual: torch.Tensor,
fn: torch.Tensor,
hc_scale: torch.Tensor,
hc_base: torch.Tensor,
rms_eps: float,
hc_pre_eps: float,
hc_sinkhorn_eps: float,
hc_post_mult_value: float,
sinkhorn_repeat: int,
norm_weight: torch.Tensor | None,
norm_eps: float | None,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor] | None:
global _AITER_MHC_RUNTIME_DISABLED, _AITER_MHC_ACTIVE_LOGGED
try:
from aiter.ops.mhc import mhc_pre as aiter_mhc_pre
except Exception as err:
logger.warning("AITER mHC pre is unavailable, falling back: %s", err)
_AITER_MHC_RUNTIME_DISABLED = True
return None
kwargs = {}
if norm_weight is not None:
kwargs["norm_weight"] = norm_weight
kwargs["norm_eps"] = norm_eps if norm_eps is not None else rms_eps
try:
result = aiter_mhc_pre(
residual,
fn,
hc_scale,
hc_base,
rms_eps,
hc_pre_eps,
hc_sinkhorn_eps,
hc_post_mult_value,
sinkhorn_repeat,
**kwargs,
)
except Exception as err:
logger.warning("AITER mHC pre failed, disabling fast path: %s", err)
_AITER_MHC_RUNTIME_DISABLED = True
return None
if not _AITER_MHC_ACTIVE_LOGGED:
logger.info("Using AITER gfx950 mHC pre/post kernels")
_AITER_MHC_ACTIVE_LOGGED = True
return result
def _try_aiter_mhc_post(
x: torch.Tensor,
residual: torch.Tensor,
post_layer_mix: torch.Tensor,
comb_res_mix: torch.Tensor,
) -> torch.Tensor | None:
global _AITER_MHC_RUNTIME_DISABLED
try:
from aiter.ops.mhc import mhc_post as aiter_mhc_post
except Exception as err:
logger.warning("AITER mHC post is unavailable, falling back: %s", err)
_AITER_MHC_RUNTIME_DISABLED = True
return None
out = torch.empty_like(residual)
try:
aiter_mhc_post(out, x, residual, post_layer_mix, comb_res_mix)
except Exception as err:
logger.warning("AITER mHC post failed, disabling fast path: %s", err)
_AITER_MHC_RUNTIME_DISABLED = True
return None
return out
# This module is imported during model-registry discovery. Do not import the real # This module is imported during model-registry discovery. Do not import the real
# TileLang package here: it loads native CUDA stubs. The proxy below lets # TileLang package here: it loads native CUDA stubs. The proxy below lets
# module-level @tilelang.jit declarations parse, then imports and applies real # module-level @tilelang.jit declarations parse, then imports and applies real
@@ -119,6 +208,24 @@ pass_configs = {
tilelang.PassConfigKey.TL_DISABLE_TMA_LOWER: True, tilelang.PassConfigKey.TL_DISABLE_TMA_LOWER: True,
} }
def _use_deep_gemm_hc_prenorm() -> bool:
if is_hip() or not envs.SGLANG_OPT_DEEPGEMM_HC_PRENORM.get():
return False
from sglang.srt.layers.deep_gemm_wrapper.configurer import ENABLE_JIT_DEEPGEMM
return ENABLE_JIT_DEEPGEMM
def _use_tilelang_mhc_pre() -> bool:
return envs.SGLANG_OPT_USE_TILELANG_MHC_PRE.get() and not is_hip()
def _use_tilelang_mhc_post() -> bool:
return envs.SGLANG_OPT_USE_TILELANG_MHC_POST.get() and not is_hip()
FP8 = "float8_e4m3" FP8 = "float8_e4m3"
BF16 = "bfloat16" BF16 = "bfloat16"
FP32 = "float32" FP32 = "float32"
@@ -1041,7 +1148,7 @@ def mhc_pre(
num_tokens, hidden_size, dtype=torch.bfloat16, device=residual.device num_tokens, hidden_size, dtype=torch.bfloat16, device=residual.device
) )
if envs.SGLANG_OPT_DEEPGEMM_HC_PRENORM.get(): if _use_deep_gemm_hc_prenorm():
n_splits = _compute_num_split_for_mhc_pre(num_tokens, hc_hidden_size) n_splits = _compute_num_split_for_mhc_pre(num_tokens, hc_hidden_size)
gemm_out_mul = torch.empty( gemm_out_mul = torch.empty(
@@ -1653,7 +1760,7 @@ def mhc_fused_post_pre(
hidden_size, hidden_size,
) )
if envs.SGLANG_OPT_DEEPGEMM_HC_PRENORM.get(): if _use_deep_gemm_hc_prenorm():
import deep_gemm import deep_gemm
deep_gemm.tf32_hc_prenorm_gemm( deep_gemm.tf32_hc_prenorm_gemm(
@@ -1847,7 +1954,25 @@ def _mhc_pre_dispatch(
norm_eps: float | None = None, norm_eps: float | None = None,
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, bool]: ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, bool]:
assert residual.dim() == 3, f"residual must be (s, n, h); got {residual.shape}" assert residual.dim() == 3, f"residual must be (s, n, h); got {residual.shape}"
if not envs.SGLANG_OPT_USE_TILELANG_MHC_PRE.get(): if _use_aiter_mhc():
result = _try_aiter_mhc_pre(
residual=residual,
fn=fn,
hc_scale=hc_scale,
hc_base=hc_base,
rms_eps=rms_eps,
hc_pre_eps=hc_pre_eps,
hc_sinkhorn_eps=hc_sinkhorn_eps,
hc_post_mult_value=hc_post_mult_value,
sinkhorn_repeat=sinkhorn_repeat,
norm_weight=norm_weight,
norm_eps=norm_eps,
)
if result is not None:
post_mix, comb_mix, layer_input = result
return post_mix, comb_mix, layer_input, norm_weight is not None
if not _use_tilelang_mhc_pre():
post_mix, comb_mix, layer_input = _mhc_pre_torch( post_mix, comb_mix, layer_input = _mhc_pre_torch(
residual=residual, residual=residual,
fn=fn, fn=fn,
@@ -1886,7 +2011,17 @@ def _mhc_post_dispatch(
) -> torch.Tensor: ) -> torch.Tensor:
assert x.dim() == 2 and residual.dim() == 3 assert x.dim() == 2 and residual.dim() == 3
assert post_layer_mix.dim() == 3 and comb_res_mix.dim() == 3 assert post_layer_mix.dim() == 3 and comb_res_mix.dim() == 3
if not envs.SGLANG_OPT_USE_TILELANG_MHC_POST.get(): if _use_aiter_mhc():
result = _try_aiter_mhc_post(
x=x,
residual=residual,
post_layer_mix=post_layer_mix,
comb_res_mix=comb_res_mix,
)
if result is not None:
return result
if not _use_tilelang_mhc_post():
return _mhc_post_torch(x, residual, post_layer_mix, comb_res_mix) return _mhc_post_torch(x, residual, post_layer_mix, comb_res_mix)
return mhc_post(x, residual, post_layer_mix, comb_res_mix) return mhc_post(x, residual, post_layer_mix, comb_res_mix)
@@ -86,14 +86,16 @@ def validate_hisparse(server_args: ServerArgs) -> None:
from sglang.srt.configs.model_config import ( from sglang.srt.configs.model_config import (
is_deepseek_dsa, is_deepseek_dsa,
is_deepseek_v4, is_deepseek_v4,
is_minimax_sparse,
) )
hf_config = model_config_of(server_args).hf_config hf_config = model_config_of(server_args).hf_config
is_v4_hisparse = is_deepseek_v4(hf_config) is_v4_hisparse = is_deepseek_v4(hf_config)
is_m3_hisparse = is_minimax_sparse(hf_config)
is_hip = get_platform().is_hip is_hip = get_platform().is_hip
assert is_deepseek_dsa(hf_config) or is_v4_hisparse, ( assert is_deepseek_dsa(hf_config) or is_v4_hisparse or is_m3_hisparse, (
"--enable-hisparse is only supported for DSA (DeepSeek Sparse Attention) " "--enable-hisparse is only supported for DSA (DeepSeek Sparse Attention) "
"models (e.g., DeepSeek V3.2, GLM-5) and DeepSeek V4 now. " "models (e.g., DeepSeek V3.2, GLM-5), DeepSeek V4, and MiniMax M3 now. "
) )
assert cfg.disable_radix_cache, ( assert cfg.disable_radix_cache, (
@@ -121,6 +123,10 @@ def validate_hisparse(server_args: ServerArgs) -> None:
) )
return return
# MiniMax M3 uses its own Triton sparse kernels.
if is_m3_hisparse:
return
if resolved_view(server_args).kv_cache_dtype not in ( if resolved_view(server_args).kv_cache_dtype not in (
"bfloat16", "bfloat16",
"auto", "auto",
@@ -24,6 +24,7 @@ class StateType(str, enum.Enum):
# only the live subrange of that row for the current open pool. # only the live subrange of that row for the current open pool.
DSA_TAIL = "dsa_tail" DSA_TAIL = "dsa_tail"
MINIMAX_INDEX_K = "minimax_index_k" MINIMAX_INDEX_K = "minimax_index_k"
MINIMAX_DENSE_KV = "minimax_dense_kv"
# DeepSeek-V4 unified_kv SWA ring: addressed per-row by ring slot # DeepSeek-V4 unified_kv SWA ring: addressed per-row by ring slot
# (req_pool_idx * ring_stride + pos % ring_stride), needs its own component. # (req_pool_idx * ring_stride + pos % ring_stride), needs its own component.
SWA_RING = "swa_ring" SWA_RING = "swa_ring"
+4 -2
View File
@@ -62,6 +62,7 @@ from sglang.srt.disaggregation.utils import (
build_staging_slot_metadata, build_staging_slot_metadata,
get_dsa_tail_state_indices, get_dsa_tail_state_indices,
get_kv_class, get_kv_class,
get_kv_transfer_buf_infos,
get_qsa_pending_state_indices, get_qsa_pending_state_indices,
is_mla_backend, is_mla_backend,
is_unadmitted_reject, is_unadmitted_reject,
@@ -575,8 +576,8 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
if self.scheduler.enable_hisparse if self.scheduler.enable_hisparse
else self.token_to_kv_pool else self.token_to_kv_pool
) )
kv_data_ptrs, kv_data_lens, kv_item_lens = ( kv_data_ptrs, kv_data_lens, kv_item_lens = get_kv_transfer_buf_infos(
transfer_kv_pool.get_contiguous_buf_infos() transfer_kv_pool
) )
kv_data_mem_kinds = ( kv_data_mem_kinds = (
["DRAM"] * len(kv_data_ptrs) ["DRAM"] * len(kv_data_ptrs)
@@ -1579,6 +1580,7 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
StateType.DSA: _full_kv_pages_payload, StateType.DSA: _full_kv_pages_payload,
StateType.DSA_TAIL: _dsa_tail_payload, StateType.DSA_TAIL: _dsa_tail_payload,
StateType.MINIMAX_INDEX_K: _full_kv_pages_payload, StateType.MINIMAX_INDEX_K: _full_kv_pages_payload,
StateType.MINIMAX_DENSE_KV: _full_kv_pages_payload,
StateType.SWA_RING: _swa_ring_payload, StateType.SWA_RING: _swa_ring_payload,
StateType.DSV4_REQUEST_STATE: _request_state_payload, StateType.DSV4_REQUEST_STATE: _request_state_payload,
StateType.BLOCK_SCALE: _full_kv_pages_payload, StateType.BLOCK_SCALE: _full_kv_pages_payload,
@@ -1398,7 +1398,9 @@ class MooncakeKVManager(StagingManagerMixin, CommonKVManager):
if src_layer_ids or dst_layer_ids: if src_layer_ids or dst_layer_ids:
# Draft buffers break the flat [K block, V block] layout, so pair by # Draft buffers break the flat [K block, V block] layout, so pair by
# layer ID instead of the half-split used by get_mha_kv_ptrs_with_pp. # layer ID instead of the half-split used by get_mha_kv_ptrs_with_pp.
if any(l != src_kv_item_len for l in self.kv_args.kv_item_lens): if any(
item_len != src_kv_item_len for item_len in self.kv_args.kv_item_lens
):
logger.error( logger.error(
f"[{mooncake_session_id}] head-sliced transfer assumes one item " f"[{mooncake_session_id}] head-sliced transfer assumes one item "
f"length for every KV entry, got {set(self.kv_args.kv_item_lens)}" f"length for every KV entry, got {set(self.kv_args.kv_item_lens)}"
@@ -1852,12 +1854,11 @@ class MooncakeKVManager(StagingManagerMixin, CommonKVManager):
) )
or rc or rc
) )
elif st == StateType.MINIMAX_INDEX_K: elif st in (StateType.MINIMAX_INDEX_K, StateType.MINIMAX_DENSE_KV):
# Equal-TP / PP=1 only. Sub-pools are compacted sparse-layer # Compacted layer lists require equal TP and PP=1 on both peers.
# lists, so PP>1 mis-slices and heterogeneous TP is unsupported.
if self.pp_size is not None and self.pp_size > 1: if self.pp_size is not None and self.pp_size > 1:
raise RuntimeError( raise RuntimeError(
"PD disagg: PP>1 not supported for MiniMax sparse index yet." "PD disagg: PP>1 not supported for MiniMax state yet."
) )
if ( if (
target_rank_registration_info is not None target_rank_registration_info is not None
@@ -1866,11 +1867,17 @@ class MooncakeKVManager(StagingManagerMixin, CommonKVManager):
): ):
raise RuntimeError( raise RuntimeError(
"PD disagg: heterogeneous TP not supported for MiniMax " "PD disagg: heterogeneous TP not supported for MiniMax "
"sparse index yet." "state yet."
) )
src_indices = list(indices) src_indices = list(indices)
dst_indices_local = list(dst_indices) dst_indices_local = list(dst_indices)
if len(src_indices) > len(dst_indices_local): if st == StateType.MINIMAX_DENSE_KV:
if len(src_indices) != len(dst_indices_local):
raise RuntimeError(
f"{st.value} state index length mismatch: "
f"prefill={len(src_indices)}, dst={len(dst_indices_local)}"
)
elif len(src_indices) > len(dst_indices_local):
src_indices = src_indices[: len(dst_indices_local)] src_indices = src_indices[: len(dst_indices_local)]
elif len(src_indices) < len(dst_indices_local): elif len(src_indices) < len(dst_indices_local):
dst_indices_local = dst_indices_local[: len(src_indices)] dst_indices_local = dst_indices_local[: len(src_indices)]
@@ -1282,6 +1282,7 @@ class MoriKVManager(CommonKVManager):
"swa_ring", "swa_ring",
"c128_state", "c128_state",
"minimax_index_k", "minimax_index_k",
"minimax_dense_kv",
): ):
statuses.extend( statuses.extend(
self._send_swa_dsa_state( self._send_swa_dsa_state(
@@ -1409,7 +1410,12 @@ class MoriKVManager(CommonKVManager):
f"PD state transfer does not support TP-mismatched non-MLA SWA models " f"PD state transfer does not support TP-mismatched non-MLA SWA models "
f"(prefill_tp_size={self.attn_tp_size}, decode_tp_size={peer_info.decode_tp_size})" f"(prefill_tp_size={self.attn_tp_size}, decode_tp_size={peer_info.decode_tp_size})"
) )
if state_type in ("qsa_pending", "qsa_compressed", "minimax_index_k"): if state_type in (
"qsa_pending",
"qsa_compressed",
"minimax_index_k",
"minimax_dense_kv",
):
if self.pp_size is not None and self.pp_size > 1: if self.pp_size is not None and self.pp_size > 1:
# MORI registration does not exchange state_layer_ids. Compact # MORI registration does not exchange state_layer_ids. Compact
# sparse-state lists therefore cannot be paired safely across # sparse-state lists therefore cannot be paired safely across
@@ -1445,6 +1451,7 @@ class MoriKVManager(CommonKVManager):
"qsa_compressed", "qsa_compressed",
"swa_ring", "swa_ring",
"c128_state", "c128_state",
"minimax_dense_kv",
): ):
raise RuntimeError( raise RuntimeError(
f"{state_type.upper()} state index length mismatch: " f"{state_type.upper()} state index length mismatch: "
@@ -2696,17 +2696,16 @@ class NixlKVManager(StagingManagerMixin, CommonKVManager):
dst_layer_ids=dst_lids, dst_layer_ids=dst_lids,
dst_item_lens=dst_lens, dst_item_lens=dst_lens,
) )
elif st == StateType.MINIMAX_INDEX_K: elif st in (StateType.MINIMAX_INDEX_K, StateType.MINIMAX_DENSE_KV):
# Equal-TP / PP=1 only. Sub-pools are compacted sparse-layer # Compacted layer lists require equal TP and PP=1 on both peers.
# lists, so PP>1 mis-slices and heterogeneous TP is unsupported.
if self.pp_size is not None and self.pp_size > 1: if self.pp_size is not None and self.pp_size > 1:
raise RuntimeError( raise RuntimeError(
"PD disagg: PP>1 not supported for MiniMax sparse index yet." "PD disagg: PP>1 not supported for MiniMax state yet."
) )
if self.attn_tp_size != decode_tp_size: if self.attn_tp_size != decode_tp_size:
raise RuntimeError( raise RuntimeError(
"PD disagg: heterogeneous TP not supported for MiniMax " "PD disagg: heterogeneous TP not supported for MiniMax "
"sparse index yet." "state yet."
) )
if len(src_indices) != len(dst_indices): if len(src_indices) != len(dst_indices):
raise RuntimeError( raise RuntimeError(
+4 -2
View File
@@ -53,6 +53,7 @@ from sglang.srt.disaggregation.utils import (
build_staging_slot_metadata, build_staging_slot_metadata,
get_dsa_tail_state_indices, get_dsa_tail_state_indices,
get_kv_class, get_kv_class,
get_kv_transfer_buf_infos,
get_qsa_pending_state_indices, get_qsa_pending_state_indices,
is_aborted, is_aborted,
is_mla_backend, is_mla_backend,
@@ -256,8 +257,8 @@ class PrefillBootstrapQueue:
hf_text_config=self.scheduler.model_config.hf_text_config, hf_text_config=self.scheduler.model_config.hf_text_config,
) )
) )
kv_data_ptrs, kv_data_lens, kv_item_lens = ( kv_data_ptrs, kv_data_lens, kv_item_lens = get_kv_transfer_buf_infos(
self.token_to_kv_pool.get_contiguous_buf_infos() self.token_to_kv_pool
) )
kv_args.prefill_end_layer = ( kv_args.prefill_end_layer = (
kv_args.prefill_start_layer + len(kv_data_ptrs) kv_args.prefill_start_layer + len(kv_data_ptrs)
@@ -1424,6 +1425,7 @@ class SchedulerDisaggregationPrefillMixin:
StateType.DSA: _full_kv_pages_payload, StateType.DSA: _full_kv_pages_payload,
StateType.DSA_TAIL: _dsa_tail_payload, StateType.DSA_TAIL: _dsa_tail_payload,
StateType.MINIMAX_INDEX_K: _full_kv_pages_payload, StateType.MINIMAX_INDEX_K: _full_kv_pages_payload,
StateType.MINIMAX_DENSE_KV: _full_kv_pages_payload,
StateType.SWA_RING: _swa_ring_payload, StateType.SWA_RING: _swa_ring_payload,
StateType.DSV4_REQUEST_STATE: _request_state_payload, StateType.DSV4_REQUEST_STATE: _request_state_payload,
StateType.BLOCK_SCALE: _full_kv_pages_payload, StateType.BLOCK_SCALE: _full_kv_pages_payload,
+13
View File
@@ -1310,6 +1310,14 @@ def build_dsa_tail_transfer_blocks(
return transfer_blocks return transfer_blocks
def get_kv_transfer_buf_infos(pool):
from sglang.srt.mem_cache.memory_pool import MiniMaxSparseKVPool
if isinstance(pool, MiniMaxSparseKVPool):
return pool.get_sparse_kv_buf_infos()
return pool.get_contiguous_buf_infos()
def setup_state_kv_args( def setup_state_kv_args(
kv_args: KVArgs, kv_args: KVArgs,
token_to_kv_pool, token_to_kv_pool,
@@ -1375,6 +1383,11 @@ def setup_state_kv_args(
if token_to_kv_pool.index_k_pool is not None: if token_to_kv_pool.index_k_pool is not None:
dp, dl, il = token_to_kv_pool.get_index_k_state_buf_infos() dp, dl, il = token_to_kv_pool.get_index_k_state_buf_infos()
append_state_component(kv_args, StateType.MINIMAX_INDEX_K, dp, dl, il) append_state_component(kv_args, StateType.MINIMAX_INDEX_K, dp, dl, il)
append_state_component(
kv_args,
StateType.MINIMAX_DENSE_KV,
*token_to_kv_pool.get_dense_kv_state_buf_infos(),
)
elif hasattr(token_to_kv_pool, "get_state_buf_infos"): elif hasattr(token_to_kv_pool, "get_state_buf_infos"):
data_ptrs, data_lens, item_lens = token_to_kv_pool.get_state_buf_infos() data_ptrs, data_lens, item_lens = token_to_kv_pool.get_state_buf_infos()
@@ -1518,9 +1518,7 @@ class GroupCoordinator:
if self.world_size == 1: if self.world_size == 1:
return input_ return input_
# Always use pynccl to avoid capturing hip graph failure on torch if self.pynccl_comm is not None and not self.pynccl_comm.disabled:
# version smaller than or equal to 2.11
if is_hip() and self.pynccl_comm is not None and not self.pynccl_comm.disabled:
self.pynccl_comm.broadcast(input_, src=src) self.pynccl_comm.broadcast(input_, src=src)
else: else:
# Broadcast. # Broadcast.
@@ -19,7 +19,7 @@ Correctness-sensitive cases stay on Triton:
Single-sequence token counts that are not a multiple of the kernel's 64-token Single-sequence token counts that are not a multiple of the kernel's 64-token
chunk are padded up to a bucket (1k/2k/4k/8k/16k/32k) in a persistent staging chunk are padded up to a bucket (1k/2k/4k/8k/16k/32k) in a persistent staging
buffer, which bounds the resident workspace set. Pad rows are state-neutral: buffer, which bounds the resident workspace set. Pad rows are state-neutral:
k/v/beta zero => no rank-1 update; raw gate -1000 => transformed decay of k/v zero => no rank-1 update, even with beta sigmoid; raw gate -1000 => decay of
exactly 1. Multi-sequence batches go through the kernel's own varlen grid exactly 1. Multi-sequence batches go through the kernel's own varlen grid
(real cu_seqlens, no padding), so their shapes are whatever the scheduler (real cu_seqlens, no padding), so their shapes are whatever the scheduler
produces and each distinct shape can retain another workspace. produces and each distinct shape can retain another workspace.
@@ -327,6 +327,7 @@ class PtxKDAKernel(LinearAttnKernelBase):
dt_bias=self._flat_param(dt_bias), dt_bias=self._flat_param(dt_bias),
return_intermediate_states=return_intermediate_states, return_intermediate_states=return_intermediate_states,
use_qk_l2norm_in_kernel=True, use_qk_l2norm_in_kernel=True,
use_beta_sigmoid_in_kernel=kwargs.get("beta_is_raw", False),
) )
out, final_state, h = result[0], result[1], result[10] out, final_state, h = result[0], result[1], result[10]
ssm_states.index_copy_(0, slot, final_state.to(ssm_states.dtype)) ssm_states.index_copy_(0, slot, final_state.to(ssm_states.dtype))
@@ -116,6 +116,7 @@ class MiniMaxSparseAttnBackend(AttentionBackend):
assert isinstance(runner.token_to_kv_pool, MiniMaxSparseKVPool) assert isinstance(runner.token_to_kv_pool, MiniMaxSparseKVPool)
self.is_npu = is_npu() self.is_npu = is_npu()
self.kv_pool = runner.token_to_kv_pool self.kv_pool = runner.token_to_kv_pool
self.hisparse_coordinator = runner.hisparse_coordinator
self.token_to_kv_pool = runner.token_to_kv_pool # alias for TboAttnBackend self.token_to_kv_pool = runner.token_to_kv_pool # alias for TboAttnBackend
self.req_to_token_pool = runner.req_to_token_pool # pool obj for TboAttnBackend self.req_to_token_pool = runner.req_to_token_pool # pool obj for TboAttnBackend
self.req_to_token = runner.req_to_token_pool.req_to_token self.req_to_token = runner.req_to_token_pool.req_to_token
@@ -176,6 +177,18 @@ class MiniMaxSparseAttnBackend(AttentionBackend):
local_tokens + self.block_size_k - 1 local_tokens + self.block_size_k - 1
) // self.block_size_k + 1 ) // self.block_size_k + 1
self.topk_blocks = sparse_cfg["sparse_topk_blocks"] self.topk_blocks = sparse_cfg["sparse_topk_blocks"]
if self.hisparse_coordinator is not None:
selected_tokens = self.topk_blocks * self.block_size_k
assert selected_tokens <= self.hisparse_coordinator.device_buffer_size, (
f"MiniMax M3 selects {selected_tokens} sparse-attention tokens, "
"but the HiSparse device buffer holds only "
f"{self.hisparse_coordinator.device_buffer_size}."
)
self._loc_mapping = (
self.kv_pool.main_pool.full_to_hisparse_device_index_mapping
)
else:
self._loc_mapping = None
# MSA (fmha_sm100) is SM100-only; fall back to the Triton sparse path when # MSA (fmha_sm100) is SM100-only; fall back to the Triton sparse path when
# the kernel is unavailable or its constraints don't hold. # the kernel is unavailable or its constraints don't hold.
@@ -209,6 +222,7 @@ class MiniMaxSparseAttnBackend(AttentionBackend):
) )
self.use_msa = ( self.use_msa = (
not envs.SGLANG_DISABLE_MSA.get() not envs.SGLANG_DISABLE_MSA.get()
and self.hisparse_coordinator is None
and msa_available() and msa_available()
and self.block_size_k == 128 and self.block_size_k == 128
and self.kv_pool.page_size == self.block_size_k and self.kv_pool.page_size == self.block_size_k
@@ -245,6 +259,7 @@ class MiniMaxSparseAttnBackend(AttentionBackend):
self.page_size = self.kv_pool.page_size self.page_size = self.kv_pool.page_size
self.use_dense_sparse_decode = ( self.use_dense_sparse_decode = (
(not self.is_npu) (not self.is_npu)
and self.hisparse_coordinator is None
and envs.SGLANG_OPT_USE_MINIMAX_DENSE_SPARSE_DECODE.get() and envs.SGLANG_OPT_USE_MINIMAX_DENSE_SPARSE_DECODE.get()
and self.block_size_k % self.page_size == 0 and self.block_size_k % self.page_size == 0
# _dense_sparse_main_decode calls trtllm decode with a bf16 q and # _dense_sparse_main_decode calls trtllm decode with a bf16 q and
@@ -326,6 +341,7 @@ class MiniMaxSparseAttnBackend(AttentionBackend):
f"msa_owns_decode={self._msa_owns_decode}, " f"msa_owns_decode={self._msa_owns_decode}, "
f"decode_cuda_graph={_decode_cuda_graph}, " f"decode_cuda_graph={_decode_cuda_graph}, "
f"fp8_attn_gemm={self.fp8_attn_gemm}, " f"fp8_attn_gemm={self.fp8_attn_gemm}, "
f"hisparse={'enabled' if self._loc_mapping is not None else 'disabled'}, "
f"npu_native_attn={'on' if (self._native_sparse_ok and _native_attn_enabled()) else 'off'}, " f"npu_native_attn={'on' if (self._native_sparse_ok and _native_attn_enabled()) else 'off'}, "
f"disable_value_layers={sorted(self.disable_value_layer_ids)})" f"disable_value_layers={sorted(self.disable_value_layer_ids)})"
) )
@@ -336,6 +352,22 @@ class MiniMaxSparseAttnBackend(AttentionBackend):
"take minutes; compiles serialize across TP ranks)." "take minutes; compiles serialize across TP ranks)."
) )
def _hisparse_swap_in_blocks(
self,
forward_batch: ForwardBatch,
topk_idx: torch.Tensor,
layer_id: int,
) -> torch.Tensor:
assert topk_idx.size(0) == 1
top_k_device_locs = self.hisparse_coordinator.swap_in_selected_blocks(
req_pool_indices=forward_batch.req_pool_indices,
seq_lens=forward_batch.seq_lens,
top_k_blocks=topk_idx[0],
layer_id=layer_id,
sparse_block_size=self.block_size_k,
)
return top_k_device_locs.unsqueeze(0)
@staticmethod @staticmethod
def _choose_decode_score_max_chunks(batch_size: int) -> int: def _choose_decode_score_max_chunks(batch_size: int) -> int:
"""Score chunk count per graph bucket. """Score chunk count per graph bucket.
@@ -1549,6 +1581,7 @@ class MiniMaxSparseAttnBackend(AttentionBackend):
idx_v_scale=layer.idx_v_scale_float, idx_v_scale=layer.idx_v_scale_float,
cached_topk_idx=cached_topk_idx, cached_topk_idx=cached_topk_idx,
return_topk_idx=want_topk, return_topk_idx=want_topk,
loc_mapping=self._loc_mapping,
) )
if want_topk: if want_topk:
idx_o, o, reduced_topk_idx = result idx_o, o, reduced_topk_idx = result
@@ -1702,6 +1735,16 @@ class MiniMaxSparseAttnBackend(AttentionBackend):
else: else:
_cached_topk = _topk_buf _cached_topk = _topk_buf
hisparse_swap_in_fn = None
if self.hisparse_coordinator is not None:
def hisparse_swap_in_fn(topk_idx):
return self._hisparse_swap_in_blocks(
forward_batch=forward_batch,
topk_idx=topk_idx,
layer_id=layer.layer_id,
)
idx_o, o = minimax_sparse_decode( idx_o, o = minimax_sparse_decode(
q, q,
None, None,
@@ -1735,6 +1778,7 @@ class MiniMaxSparseAttnBackend(AttentionBackend):
idx_v_scale=layer.idx_v_scale_float, idx_v_scale=layer.idx_v_scale_float,
cached_topk_idx=_cached_topk, cached_topk_idx=_cached_topk,
topk_out=_topk_buf if _want_topk else None, topk_out=_topk_buf if _want_topk else None,
hisparse_swap_in_fn=hisparse_swap_in_fn,
) )
return ( return (
None if idx_o is None else idx_o.reshape(q.shape[0], -1).contiguous(), None if idx_o is None else idx_o.reshape(q.shape[0], -1).contiguous(),
@@ -75,6 +75,7 @@ def minimax_sparse_prefill(
idx_v_scale: Optional[float] = None, idx_v_scale: Optional[float] = None,
cached_topk_idx: Optional[torch.Tensor] = None, cached_topk_idx: Optional[torch.Tensor] = None,
return_topk_idx: bool = False, return_topk_idx: bool = False,
loc_mapping: Optional[torch.Tensor] = None,
): ):
"""Run MiniMax-M3 sparse prefill. """Run MiniMax-M3 sparse prefill.
@@ -146,7 +147,7 @@ def minimax_sparse_prefill(
# Step 3: Sparse attention using topk index (main head). The MSA path only # Step 3: Sparse attention using topk index (main head). The MSA path only
# replaces this step; the indexer above is unchanged. MSA has no attn-sink # replaces this step; the indexer above is unchanged. MSA has no attn-sink
# input, so keep the Triton path when sink is present. # input, so keep the Triton path when sink is present.
if use_msa and sink is None: if use_msa and sink is None and loc_mapping is None:
from .msa import MSAUnavailableError, msa_sparse_prefill_main from .msa import MSAUnavailableError, msa_sparse_prefill_main
try: try:
@@ -188,6 +189,7 @@ def minimax_sparse_prefill(
q_scale=q_scale, q_scale=q_scale,
k_scale=k_scale, k_scale=k_scale,
v_scale=v_scale, v_scale=v_scale,
loc_mapping=loc_mapping,
) )
else: else:
o = flash_prefill_with_gqa_share_sparse( o = flash_prefill_with_gqa_share_sparse(
@@ -210,6 +212,7 @@ def minimax_sparse_prefill(
q_scale=q_scale, q_scale=q_scale,
k_scale=k_scale, k_scale=k_scale,
v_scale=v_scale, v_scale=v_scale,
loc_mapping=loc_mapping,
) )
if return_topk_idx: if return_topk_idx:
return idx_o, o, reduced_topk_idx return idx_o, o, reduced_topk_idx
@@ -255,6 +258,7 @@ def minimax_sparse_decode(
idx_v_scale: Optional[float] = None, idx_v_scale: Optional[float] = None,
cached_topk_idx: Optional[torch.Tensor] = None, cached_topk_idx: Optional[torch.Tensor] = None,
topk_out: Optional[torch.Tensor] = None, topk_out: Optional[torch.Tensor] = None,
hisparse_swap_in_fn: Optional[Callable] = None,
) -> Tuple[torch.Tensor, torch.Tensor]: ) -> Tuple[torch.Tensor, torch.Tensor]:
# Index top-k sharing for DECODE. A group's source layer passes ``topk_out`` # Index top-k sharing for DECODE. A group's source layer passes ``topk_out``
# (a persistent buffer) and publishes its reduced top-k there; the group's # (a persistent buffer) and publishes its reduced top-k there; the group's
@@ -319,9 +323,12 @@ def minimax_sparse_decode(
f"reduced top-k shape {tuple(topk_idx.shape)}" f"reduced top-k shape {tuple(topk_idx.shape)}"
) )
topk_out.copy_(topk_idx) topk_out.copy_(topk_idx)
hisparse_slots = (
hisparse_swap_in_fn(topk_idx) if hisparse_swap_in_fn is not None else None
)
# Step 3: Sparse attention using topk index (main head). The MSA path # Step 3: Sparse attention using topk index (main head). The MSA path
# only replaces this step; keep the Triton path when sink is present. # only replaces this step; keep the Triton path when sink is present.
if use_msa and sink is None: if use_msa and sink is None and hisparse_slots is None:
from .msa import MSAUnavailableError, msa_sparse_decode_main from .msa import MSAUnavailableError, msa_sparse_decode_main
try: try:
@@ -357,6 +364,7 @@ def minimax_sparse_decode(
q_scale=q_scale, q_scale=q_scale,
k_scale=k_scale, k_scale=k_scale,
v_scale=v_scale, v_scale=v_scale,
hisparse_slots=hisparse_slots,
) )
else: else:
o = flash_decode_with_gqa_share_sparse( o = flash_decode_with_gqa_share_sparse(
@@ -373,5 +381,6 @@ def minimax_sparse_decode(
q_scale=q_scale, q_scale=q_scale,
k_scale=k_scale, k_scale=k_scale,
v_scale=v_scale, v_scale=v_scale,
hisparse_slots=hisparse_slots,
) )
return idx_o, o return idx_o, o
@@ -1497,6 +1497,15 @@ class ModelOptFp4Config(ModelOptQuantConfig):
def get_min_capability(cls) -> int: def get_min_capability(cls) -> int:
return 80 return 80
def can_fuse_shared_expert(self) -> bool:
# A shared-expert body kept BF16 via exclude_modules cannot share the packed
# FP4 FusedMoE buffers. The shared_expert_gate is a separate linear (kept
# BF16 by e.g. Qwen3-Next NVFP4 checkpoints) and must not veto fusion.
return not any(
"shared_expert" in name and "shared_expert_gate" not in name
for name in self.exclude_modules
)
@staticmethod @staticmethod
def common_group_size(cfg: dict) -> int: def common_group_size(cfg: dict) -> int:
"""Return the unique group_size across the config; raise if missing/mismatched.""" """Return the unique group_size across the config; raise if missing/mismatched."""
@@ -19,9 +19,16 @@ if is_xpu():
"copy_cache_planned_mla has no AOT sgl_kernel implementation." "copy_cache_planned_mla has no AOT sgl_kernel implementation."
) )
def load_blocks_to_device_buffer_mha(*args, **kwargs):
raise RuntimeError(
"MiniMax M3 HiSparse block swap-in is unsupported on XPU: "
"load_blocks_to_device_buffer_mha has no AOT sgl_kernel implementation."
)
else: else:
from sglang.kernels.ops.kvcache.hisparse import ( from sglang.kernels.ops.kvcache.hisparse import (
copy_cache_planned_mla, copy_cache_planned_mla,
load_blocks_to_device_buffer_mha,
load_cache_to_device_buffer_dsv4_mla, load_cache_to_device_buffer_dsv4_mla,
load_cache_to_device_buffer_mla, load_cache_to_device_buffer_mla,
) )
@@ -36,8 +43,9 @@ from sglang.srt.mem_cache.allocator.hisparse import (
from sglang.srt.mem_cache.hisparse_memory_pool import ( from sglang.srt.mem_cache.hisparse_memory_pool import (
HiSparseDSATokenToKVPool, HiSparseDSATokenToKVPool,
) )
from sglang.srt.mem_cache.memory_pool import ReqToTokenPool from sglang.srt.mem_cache.memory_pool import MiniMaxSparseKVPool, ReqToTokenPool
from sglang.srt.mem_cache.memory_pool_host import DeepSeekV4PagedHostPool from sglang.srt.mem_cache.memory_pool_host import DeepSeekV4PagedHostPool
from sglang.srt.mem_cache.pool_host.mha import HiSparseMHATokenToKVPoolHost
from sglang.srt.mem_cache.pool_host.mla import MLATokenToKVPoolHost from sglang.srt.mem_cache.pool_host.mla import MLATokenToKVPoolHost
device_module = get_device_module() device_module = get_device_module()
@@ -157,9 +165,11 @@ class HiSparseCoordinator:
) )
self.compress_ratio = self.token_to_kv_pool_allocator.compress_ratio self.compress_ratio = self.token_to_kv_pool_allocator.compress_ratio
kvcache = self.token_to_kv_pool_allocator.get_kvcache()
self.is_dsv4_hisparse = isinstance( self.is_dsv4_hisparse = isinstance(
self.token_to_kv_pool_allocator, DeepSeekV4HiSparseTokenToKVPoolAllocator self.token_to_kv_pool_allocator, DeepSeekV4HiSparseTokenToKVPoolAllocator
) )
self.is_m3_hisparse = isinstance(kvcache, MiniMaxSparseKVPool)
if self.is_dsv4_hisparse: if self.is_dsv4_hisparse:
self.mem_pool_device = self.token_to_kv_pool_allocator.hisparse_kvcache self.mem_pool_device = self.token_to_kv_pool_allocator.hisparse_kvcache
page_size = self.mem_pool_device.page_size page_size = self.mem_pool_device.page_size
@@ -184,18 +194,30 @@ class HiSparseCoordinator:
assert isinstance( assert isinstance(
self.token_to_kv_pool_allocator, HiSparseTokenToKVPoolAllocator self.token_to_kv_pool_allocator, HiSparseTokenToKVPoolAllocator
) )
self.mem_pool_device: HiSparseDSATokenToKVPool = ( if self.is_m3_hisparse:
self.token_to_kv_pool_allocator.get_kvcache() self.mem_pool_device = kvcache.main_pool
) assert self.mem_pool_device.head_num == 1, (
self.mem_pool_host = MLATokenToKVPoolHost( "MiniMax M3 HiSparse requires one KV head per TP rank, "
device_pool=self.mem_pool_device, f"got {self.mem_pool_device.head_num}. Increase the "
host_to_device_ratio=host_to_device_ratio, "tensor-parallel size."
host_size=0, )
page_size=self.mem_pool_device.page_size, self.mem_pool_host = HiSparseMHATokenToKVPoolHost(
layout="layer_first", device_pool=self.mem_pool_device,
override_kv_cache_dim=self.mem_pool_device.kv_cache_dim, host_to_device_ratio=host_to_device_ratio,
) page_size=self.mem_pool_device.page_size,
self.item_size_bytes = self.mem_pool_host.token_stride_size )
self.item_size_bytes = self.mem_pool_device.bytes_per_token_k
else:
self.mem_pool_device: HiSparseDSATokenToKVPool = kvcache
self.mem_pool_host = MLATokenToKVPoolHost(
device_pool=self.mem_pool_device,
host_to_device_ratio=host_to_device_ratio,
host_size=0,
page_size=self.mem_pool_device.page_size,
layout="layer_first",
override_kv_cache_dim=self.mem_pool_device.kv_cache_dim,
)
self.item_size_bytes = self.mem_pool_host.token_stride_size
self.page_size = self.mem_pool_device.page_size self.page_size = self.mem_pool_device.page_size
max_num_req_slots = req_to_token_pool.req_to_token.shape[0] max_num_req_slots = req_to_token_pool.req_to_token.shape[0]
@@ -263,9 +285,17 @@ class HiSparseCoordinator:
self.device_buffer_size, dtype=torch.int32, device=device self.device_buffer_size, dtype=torch.int32, device=device
) )
# Pre-allocated output buffer for swap_in_selected_pages (CUDA-graph safe) # Pre-allocated output buffer for swap-in (CUDA-graph safe). MiniMax
# selects blocks, so its flattened token-slot output can occupy any
# prefix up to the full device working-set size.
swap_output_width = (
self.device_buffer_size if self.is_m3_hisparse else self.top_k
)
self.top_k_device_locs_buffer = torch.full( self.top_k_device_locs_buffer = torch.full(
(max_num_req_slots, self.top_k), -1, dtype=torch.int32, device=device (max_num_req_slots, swap_output_width),
-1,
dtype=torch.int32,
device=device,
) )
self.raw_indices_buffer = torch.full( self.raw_indices_buffer = torch.full(
(max_num_req_slots, self.top_k), -1, dtype=torch.int32, device=device (max_num_req_slots, self.top_k), -1, dtype=torch.int32, device=device
@@ -457,7 +487,12 @@ class HiSparseCoordinator:
host_indices = self.req_to_host_pool[req.kv.req_pool_idx, :n] host_indices = self.req_to_host_pool[req.kv.req_pool_idx, :n]
device_locs = self.req_to_device_buffer[req.kv.req_pool_idx, :n] device_locs = self.req_to_device_buffer[req.kv.req_pool_idx, :n]
for layer_id in range(self.mem_pool_device.layer_num): layer_ids = (
range(self.mem_pool_device.start_layer, self.mem_pool_device.end_layer)
if self.is_m3_hisparse
else range(self.mem_pool_device.layer_num)
)
for layer_id in layer_ids:
self.mem_pool_host.load_to_device_per_layer( self.mem_pool_host.load_to_device_per_layer(
self.mem_pool_device, self.mem_pool_device,
host_indices, host_indices,
@@ -642,13 +677,9 @@ class HiSparseCoordinator:
compressed_locs = self.token_to_kv_pool_allocator.get_last_loc_compressed( compressed_locs = self.token_to_kv_pool_allocator.get_last_loc_compressed(
out_cache_loc out_cache_loc
) )
# ROCm: the decode remap creates a temporary hisparse device slot per # Page-size-one allocation creates a temporary slot before remapping
# new token (via the page_size==1 allocator path). Free the stale # the new token into the request's reserved device-buffer slot.
# slot before pointing the mapping at the reserved device-buffer slot, if _is_hip or self.mem_pool_device.page_size == 1:
# otherwise the temporary slots leak and corrupt later swap-in lookups.
# CUDA keeps the original behavior: the swap-in kernel consumes only
# top_k_device_locs, so stale mapping entries are harmless there.
if _is_hip:
previous_locs = self.mem_pool_device._translate_loc_to_hisparse_device( previous_locs = self.mem_pool_device._translate_loc_to_hisparse_device(
compressed_locs compressed_locs
) )
@@ -976,8 +1007,7 @@ class HiSparseCoordinator:
miss plan into self._miss_{src,dst,count} for the skip layers to replay. miss plan into self._miss_{src,dst,count} for the skip layers to replay.
""" """
num_reqs = req_pool_indices.size(0) num_reqs = req_pool_indices.size(0)
top_k_indices = self.top_k_device_locs_buffer[:num_reqs] top_k_indices = self.top_k_device_locs_buffer[:num_reqs, : self.top_k]
swap_in_fn = ( swap_in_fn = (
load_cache_to_device_buffer_dsv4_mla load_cache_to_device_buffer_dsv4_mla
if self.is_dsv4_hisparse if self.is_dsv4_hisparse
@@ -1015,6 +1045,46 @@ class HiSparseCoordinator:
) )
return top_k_indices return top_k_indices
def swap_in_selected_blocks(
self,
req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor,
top_k_blocks: torch.Tensor,
layer_id: int,
sparse_block_size: int,
) -> torch.Tensor:
assert self.is_m3_hisparse
num_reqs = req_pool_indices.size(0)
num_selected_tokens = top_k_blocks.size(1) * sparse_block_size
assert num_selected_tokens <= self.device_buffer_size, (
f"MiniMax M3 selected {num_selected_tokens} tokens, but the "
f"HiSparse device buffer holds only {self.device_buffer_size}."
)
top_k_indices = self.top_k_device_locs_buffer[:num_reqs, :num_selected_tokens]
host_layer = layer_id - self.mem_pool_device.start_layer
load_blocks_to_device_buffer_mha(
top_k_blocks=top_k_blocks,
device_buffer_tokens=self.req_device_buffer_tokens[host_layer],
host_cache_locs=self.req_to_host_pool,
device_buffer_locs=self.req_device_buffer_token_locs[host_layer],
host_cache_k=self.mem_pool_host.k_buffer[host_layer],
host_cache_v=self.mem_pool_host.v_buffer[host_layer],
device_buffer_k=self.mem_pool_device.get_key_buffer(layer_id),
device_buffer_v=self.mem_pool_device.get_value_buffer(layer_id),
top_k_device_locs=top_k_indices,
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
lru_slots=self.lru_slots[host_layer],
item_size_bytes=self.item_size_bytes,
hot_buffer_size=self.device_buffer_size,
sparse_block_size=sparse_block_size,
page_size=1,
block_size=self.swap_in_block_size,
num_real_reqs=self.num_real_reqs,
skip_io=self.skip_io,
)
return top_k_indices
def _run_copy_only_kernel(self, num_reqs: int, skip_layer: int) -> None: def _run_copy_only_kernel(self, num_reqs: int, skip_layer: int) -> None:
"""Replay the anchor's recorded miss plan into a skip layer's buffers """Replay the anchor's recorded miss plan into a skip layer's buffers
(IO-only; the anchor's slot table stays valid -- lockstep layout).""" (IO-only; the anchor's slot table stays valid -- lockstep layout)."""
@@ -1045,7 +1115,10 @@ class HiSparseCoordinator:
""" """
if not self.enable_prefetch: if not self.enable_prefetch:
return self._run_swap_in_kernel( return self._run_swap_in_kernel(
req_pool_indices, compressed_seq_lens, top_k_result, layer_id req_pool_indices,
compressed_seq_lens,
top_k_result,
layer_id,
) )
num_reqs = req_pool_indices.size(0) num_reqs = req_pool_indices.size(0)
@@ -1054,7 +1127,7 @@ class HiSparseCoordinator:
# applies (shared index + lockstep buffers). # applies (shared index + lockstep buffers).
slot = self._prefetch_slot[layer_id] slot = self._prefetch_slot[layer_id]
self._prefetch_events[slot].wait(device_module.current_stream()) self._prefetch_events[slot].wait(device_module.current_stream())
return self.top_k_device_locs_buffer[:num_reqs] return self.top_k_device_locs_buffer[:num_reqs, : self.top_k]
# Anchor: swap in synchronously (recording the plan), then prefetch the # Anchor: swap in synchronously (recording the plan), then prefetch the
# skip layers' copies on the side stream. # skip layers' copies on the side stream.
@@ -1,4 +1,7 @@
from __future__ import annotations
import weakref import weakref
from typing import TYPE_CHECKING
import torch import torch
@@ -11,6 +14,9 @@ from sglang.srt.mem_cache.deepseek_v4_memory_pool import (
from sglang.srt.mem_cache.hisparse_memory_pool import HiSparseDSATokenToKVPool from sglang.srt.mem_cache.hisparse_memory_pool import HiSparseDSATokenToKVPool
from sglang.srt.utils.common import get_num_new_pages from sglang.srt.utils.common import get_num_new_pages
if TYPE_CHECKING:
from sglang.srt.mem_cache.memory_pool import MiniMaxSparseKVPool
class HiSparseTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator): class HiSparseTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
def __init__( def __init__(
@@ -19,7 +25,7 @@ class HiSparseTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
page_size: int, page_size: int,
dtype: torch.dtype, dtype: torch.dtype,
device: torch.device, device: torch.device,
kvcache: HiSparseDSATokenToKVPool, kvcache: HiSparseDSATokenToKVPool | MiniMaxSparseKVPool,
need_sort: bool, need_sort: bool,
host_to_device_ratio: int = 2, host_to_device_ratio: int = 2,
): ):
@@ -9,7 +9,7 @@ from sglang.kernels.ops.kvcache.hisparse_slot_mapping import (
translate_padded_hisparse_locations, translate_padded_hisparse_locations,
) )
from sglang.srt.layers.radix_attention import RadixAttention from sglang.srt.layers.radix_attention import RadixAttention
from sglang.srt.mem_cache.memory_pool import DSATokenToKVPool from sglang.srt.mem_cache.memory_pool import DSATokenToKVPool, MHATokenToKVPool
from sglang.srt.utils import is_cuda, is_hip, is_xpu from sglang.srt.utils import is_cuda, is_hip, is_xpu
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -147,3 +147,111 @@ class HiSparseDSATokenToKVPool(DSATokenToKVPool):
self, kv_cache_cpu, indices, mamba_indices=None, req_pool_index=None self, kv_cache_cpu, indices, mamba_indices=None, req_pool_index=None
): ):
raise NotImplementedError("HiSparseDevicePool does not support load_cpu_copy") raise NotImplementedError("HiSparseDevicePool does not support load_cpu_copy")
class HiSparseMHAMainPool(MHATokenToKVPool):
"""MHA KV pool with HiSparse logical-to-device mapping.
Used by MiniMax M3 HiSparse. The index pools (index_kv_pool, index_k_pool)
stay fully resident on the device and do not use this mapping.
"""
def __init__(
self,
size: int,
page_size: int,
dtype: torch.dtype,
head_num: int,
head_dim: int,
layer_num: int,
device: str,
enable_memory_saver: bool,
start_layer: Optional[int] = None,
end_layer: Optional[int] = None,
):
super().__init__(
size=size,
page_size=page_size,
dtype=dtype,
head_num=head_num,
head_dim=head_dim,
layer_num=layer_num,
device=device,
enable_memory_saver=enable_memory_saver,
start_layer=start_layer,
end_layer=end_layer,
)
self.full_to_hisparse_device_index_mapping: Optional[torch.Tensor] = None
self.bytes_per_token_k = head_num * head_dim * self.store_dtype.itemsize
self.bytes_per_token_v = head_num * self.v_head_dim * self.store_dtype.itemsize
def register_mapping(
self, full_to_hisparse_device_index_mapping: torch.Tensor
) -> None:
self.full_to_hisparse_device_index_mapping = (
full_to_hisparse_device_index_mapping
)
def translate_loc_to_hisparse_device(self, indices: torch.Tensor) -> torch.Tensor:
assert self.full_to_hisparse_device_index_mapping is not None
return self.full_to_hisparse_device_index_mapping[indices]
def _translate_loc_to_hisparse_device(self, indices: torch.Tensor) -> torch.Tensor:
assert self.full_to_hisparse_device_index_mapping is not None
return self.full_to_hisparse_device_index_mapping[indices]
def translate_loc_from_full_to_hisparse_device(
self, full_indices: torch.Tensor
) -> torch.Tensor:
assert self.full_to_hisparse_device_index_mapping is not None
return self.full_to_hisparse_device_index_mapping[full_indices]
def translate_loc_from_full_to_compressed(
self, full_indices: torch.Tensor
) -> torch.Tensor:
return full_indices
def set_kv_buffer(
self,
layer: RadixAttention,
loc,
cache_k: torch.Tensor,
cache_v: torch.Tensor,
*args,
**kwargs,
):
from sglang.srt.mem_cache.memory_pool import unwrap_write_loc
raw_loc, _, _ = unwrap_write_loc(loc)
translated = self.translate_loc_to_hisparse_device(raw_loc)
super().set_kv_buffer(layer, translated, cache_k, cache_v, *args, **kwargs)
def transfer_values_on_device(
self,
dst_indices: torch.Tensor,
src_indices: torch.Tensor,
) -> None:
transfer_kv_all_layer_mla(
src_layers=self.k_data_ptrs,
dst_layers=self.k_data_ptrs,
src_indices=src_indices,
dst_indices=dst_indices,
item_size=self.bytes_per_token_k,
num_layers=self.layer_num,
)
transfer_kv_all_layer_mla(
src_layers=self.v_data_ptrs,
dst_layers=self.v_data_ptrs,
src_indices=src_indices,
dst_indices=dst_indices,
item_size=self.bytes_per_token_v,
num_layers=self.layer_num,
)
def get_cpu_copy(self, indices, mamba_indices=None, req_pool_index=None):
raise NotImplementedError("HiSparseMHAMainPool does not support get_cpu_copy")
def load_cpu_copy(
self, kv_cache_cpu, indices, mamba_indices=None, req_pool_index=None
):
raise NotImplementedError("HiSparseMHAMainPool does not support load_cpu_copy")
@@ -1836,6 +1836,14 @@ class KVCacheConfigurator:
disable_value_sparse_layer_ids = get_minimax_sparse_disable_value_layer_ids( disable_value_sparse_layer_ids = get_minimax_sparse_disable_value_layer_ids(
sparse_cfg sparse_cfg
) )
enable_hisparse = get_memory().enable_hisparse
hisparse_kwargs = {}
if enable_hisparse:
from sglang.srt.mem_cache.sparsity import parse_hisparse_config
hisparse_kwargs["host_to_device_ratio"] = (
parse_hisparse_config().host_to_device_ratio
)
token_to_kv_pool = MiniMaxSparseKVPool( token_to_kv_pool = MiniMaxSparseKVPool(
size=max_total_num_tokens, size=max_total_num_tokens,
page_size=self.pool_page_size, page_size=self.pool_page_size,
@@ -1861,6 +1869,8 @@ class KVCacheConfigurator:
enable_memory_saver=get_exec().features.enable_memory_saver, enable_memory_saver=get_exec().features.enable_memory_saver,
start_layer=self.layer_info.start_layer, start_layer=self.layer_info.start_layer,
end_layer=self.layer_info.end_layer, end_layer=self.layer_info.end_layer,
enable_hisparse=enable_hisparse,
**hisparse_kwargs,
) )
return token_to_kv_pool return token_to_kv_pool
+134 -24
View File
@@ -5448,6 +5448,8 @@ class MiniMaxSparseKVPool(KVCache):
main_pool_cls=MHATokenToKVPool, main_pool_cls=MHATokenToKVPool,
index_kv_pool_cls=MHATokenToKVPool, index_kv_pool_cls=MHATokenToKVPool,
index_k_pool_cls=MHATokenToKOnlyPool, index_k_pool_cls=MHATokenToKOnlyPool,
enable_hisparse: bool = False,
host_to_device_ratio: int = 2,
): ):
# Do not call super().__init__() — delegate to sub-pools instead. # Do not call super().__init__() — delegate to sub-pools instead.
self.size = size self.size = size
@@ -5466,6 +5468,7 @@ class MiniMaxSparseKVPool(KVCache):
] ]
index_dtype = index_dtype if index_dtype is not None else dtype index_dtype = index_dtype if index_dtype is not None else dtype
index_pool_size = size * host_to_device_ratio if enable_hisparse else size
# Split sparse layers by V policy: kv_sparse (index_kv_pool holds K+V) vs # Split sparse layers by V policy: kv_sparse (index_kv_pool holds K+V) vs
# k_only_sparse (index_k_pool holds only K; V is never read). # k_only_sparse (index_k_pool holds only K; V is never read).
@@ -5489,22 +5492,59 @@ class MiniMaxSparseKVPool(KVCache):
gid: i for i, gid in enumerate(local_k_only_sparse_layer_ids) gid: i for i, gid in enumerate(local_k_only_sparse_layer_ids)
} }
self.main_pool = main_pool_cls( self._dense_layer_ids = set(local_dense_layer_ids)
size=size, main_layer_num = len(local_dense_layer_ids) + len(local_sparse_layer_ids)
page_size=page_size, if enable_hisparse:
dtype=dtype, from sglang.srt.mem_cache.hisparse_memory_pool import (
head_num=head_num, HiSparseMHAMainPool,
head_dim=head_dim, )
layer_num=len(local_dense_layer_ids) + len(local_sparse_layer_ids),
device=device, self.dense_pool = (
enable_memory_saver=enable_memory_saver, main_pool_cls(
start_layer=start_layer, size=index_pool_size,
end_layer=end_layer, page_size=page_size,
) dtype=dtype,
head_num=head_num,
head_dim=head_dim,
layer_num=len(local_dense_layer_ids),
device=device,
enable_memory_saver=enable_memory_saver,
start_layer=start_layer,
end_layer=start_layer + len(local_dense_layer_ids),
)
if local_dense_layer_ids
else None
)
self.main_pool = HiSparseMHAMainPool(
size=size,
page_size=page_size,
dtype=dtype,
head_num=head_num,
head_dim=head_dim,
layer_num=len(local_sparse_layer_ids),
device=device,
enable_memory_saver=enable_memory_saver,
start_layer=local_sparse_layer_ids[0],
end_layer=end_layer,
)
else:
self.dense_pool = None
self.main_pool = main_pool_cls(
size=size,
page_size=page_size,
dtype=dtype,
head_num=head_num,
head_dim=head_dim,
layer_num=main_layer_num,
device=device,
enable_memory_saver=enable_memory_saver,
start_layer=start_layer,
end_layer=end_layer,
)
self.index_kv_pool: Optional[MHATokenToKVPool] = ( self.index_kv_pool: Optional[MHATokenToKVPool] = (
index_kv_pool_cls( index_kv_pool_cls(
size=size, size=index_pool_size,
page_size=page_size, page_size=page_size,
dtype=index_dtype, dtype=index_dtype,
head_num=1, head_num=1,
@@ -5519,7 +5559,7 @@ class MiniMaxSparseKVPool(KVCache):
self.index_k_pool: Optional[MHATokenToKOnlyPool] = ( self.index_k_pool: Optional[MHATokenToKOnlyPool] = (
index_k_pool_cls( index_k_pool_cls(
size=size, size=index_pool_size,
page_size=page_size, page_size=page_size,
dtype=index_dtype, dtype=index_dtype,
head_num=1, head_num=1,
@@ -5533,19 +5573,58 @@ class MiniMaxSparseKVPool(KVCache):
) )
self.mem_usage = self.main_pool.mem_usage self.mem_usage = self.main_pool.mem_usage
if self.dense_pool is not None:
self.mem_usage += self.dense_pool.mem_usage
if self.index_kv_pool is not None: if self.index_kv_pool is not None:
self.mem_usage += self.index_kv_pool.mem_usage self.mem_usage += self.index_kv_pool.mem_usage
if self.index_k_pool is not None: if self.index_k_pool is not None:
self.mem_usage += self.index_k_pool.mem_usage self.mem_usage += self.index_k_pool.mem_usage
# HiCacheController reads these from the top-level KV pool wrapper. # HiCacheController reads these from the top-level KV pool wrapper.
self.layer_num = self.main_pool.layer_num self.layer_num = main_layer_num
self.start_layer = self.main_pool.start_layer self.start_layer = start_layer
self.end_layer = self.main_pool.end_layer self.end_layer = end_layer
# PD disaggregation reads these directly (no fallback) off the wrapper. # PD disaggregation reads these directly (no fallback) off the wrapper.
self.head_num = self.main_pool.head_num self.head_num = self.main_pool.head_num
self.head_dim = self.main_pool.head_dim self.head_dim = self.main_pool.head_dim
self.v_head_dim = self.main_pool.v_head_dim
self.store_dtype = self.main_pool.store_dtype
self.layer_transfer_counter = None self.layer_transfer_counter = None
self._enable_hisparse = enable_hisparse
def register_mapping(self, mapping: torch.Tensor) -> None:
assert self._enable_hisparse
self.main_pool.register_mapping(mapping)
def _translate_loc_to_hisparse_device(self, indices: torch.Tensor) -> torch.Tensor:
assert self._enable_hisparse
return self.main_pool._translate_loc_to_hisparse_device(indices)
def translate_loc_to_hisparse_device(self, indices: torch.Tensor) -> torch.Tensor:
assert self._enable_hisparse
return self.main_pool.translate_loc_to_hisparse_device(indices)
def translate_loc_from_full_to_hisparse_device(
self, indices: torch.Tensor
) -> torch.Tensor:
assert self._enable_hisparse
return self.main_pool.translate_loc_from_full_to_hisparse_device(indices)
def translate_loc_from_full_to_compressed(
self, indices: torch.Tensor
) -> torch.Tensor:
assert self._enable_hisparse
return self.main_pool.translate_loc_from_full_to_compressed(indices)
@property
def bytes_per_token_k(self) -> int:
assert self._enable_hisparse
return self.main_pool.bytes_per_token_k
@property
def full_to_hisparse_device_index_mapping(self):
assert self._enable_hisparse
return self.main_pool.full_to_hisparse_device_index_mapping
def register_layer_transfer_counter( def register_layer_transfer_counter(
self, layer_transfer_counter: LayerDoneCounter self, layer_transfer_counter: LayerDoneCounter
@@ -5561,17 +5640,22 @@ class MiniMaxSparseKVPool(KVCache):
if self.layer_transfer_counter is not None: if self.layer_transfer_counter is not None:
self.layer_transfer_counter.wait_until(layer_id - self.start_layer) self.layer_transfer_counter.wait_until(layer_id - self.start_layer)
def _pool_for(self, layer_id: int) -> MHATokenToKVPool:
if self.dense_pool is not None and layer_id in self._dense_layer_ids:
return self.dense_pool
return self.main_pool
def get_key_buffer(self, layer_id: int) -> torch.Tensor: def get_key_buffer(self, layer_id: int) -> torch.Tensor:
self._wait_for_layer(layer_id) self._wait_for_layer(layer_id)
return self.main_pool.get_key_buffer(layer_id) return self._pool_for(layer_id).get_key_buffer(layer_id)
def get_value_buffer(self, layer_id: int) -> torch.Tensor: def get_value_buffer(self, layer_id: int) -> torch.Tensor:
self._wait_for_layer(layer_id) self._wait_for_layer(layer_id)
return self.main_pool.get_value_buffer(layer_id) return self._pool_for(layer_id).get_value_buffer(layer_id)
def get_kv_buffer(self, layer_id: int) -> Tuple[torch.Tensor, torch.Tensor]: def get_kv_buffer(self, layer_id: int) -> Tuple[torch.Tensor, torch.Tensor]:
self._wait_for_layer(layer_id) self._wait_for_layer(layer_id)
return self.main_pool.get_kv_buffer(layer_id) return self._pool_for(layer_id).get_kv_buffer(layer_id)
def get_index_kv_buffer(self, layer_id: int) -> Tuple[torch.Tensor, torch.Tensor]: def get_index_kv_buffer(self, layer_id: int) -> Tuple[torch.Tensor, torch.Tensor]:
self._wait_for_layer(layer_id) self._wait_for_layer(layer_id)
@@ -5613,7 +5697,7 @@ class MiniMaxSparseKVPool(KVCache):
Scale semantics follow MHATokenToKVPool: None means unit scale; Scale semantics follow MHATokenToKVPool: None means unit scale;
a non-None scale is applied with an in-place div_ before the fp8 cast. a non-None scale is applied with an in-place div_ before the fp8 cast.
""" """
self.main_pool.set_kv_buffer( self._pool_for(layer.layer_id).set_kv_buffer(
layer, layer,
loc, loc,
cache_k, cache_k,
@@ -5711,8 +5795,10 @@ class MiniMaxSparseKVPool(KVCache):
disable_value = cache_idx_v is None disable_value = cache_idx_v is None
index_pool = self.index_k_pool if disable_value else self.index_kv_pool index_pool = self.index_k_pool if disable_value else self.index_kv_pool
if index_pool is not None and self._can_fuse_kv_index_store( if (
index_pool, cache_k, cache_idx_k index_pool is not None
and not self._enable_hisparse
and self._can_fuse_kv_index_store(index_pool, cache_k, cache_idx_k)
): ):
from sglang.kernels.ops.kvcache.minimax_store_kv_index import store_kv_index from sglang.kernels.ops.kvcache.minimax_store_kv_index import store_kv_index
@@ -5759,7 +5845,12 @@ class MiniMaxSparseKVPool(KVCache):
) )
def get_kv_size_bytes(self): def get_kv_size_bytes(self):
sub_pools = [self.main_pool, self.index_kv_pool, self.index_k_pool] sub_pools = [
self.main_pool,
self.dense_pool,
self.index_kv_pool,
self.index_k_pool,
]
sizes = [p.get_kv_size_bytes() for p in sub_pools if p is not None] sizes = [p.get_kv_size_bytes() for p in sub_pools if p is not None]
return sum(k for k, _ in sizes), sum(v for _, v in sizes) return sum(k for k, _ in sizes), sum(v for _, v in sizes)
@@ -5767,6 +5858,25 @@ class MiniMaxSparseKVPool(KVCache):
# Main K/V only; index buffers ride the state-buffer channel. # Main K/V only; index buffers ride the state-buffer channel.
return self.main_pool.get_contiguous_buf_infos() return self.main_pool.get_contiguous_buf_infos()
def get_sparse_kv_buf_infos(self):
return self._get_layer_kv_buf_infos(
layer_ids=sorted(self.sparse_layer_id_mapping)
)
def get_dense_kv_state_buf_infos(self):
# Dense KV uses logical device slots, independently of sparse host slots.
return self._get_layer_kv_buf_infos(layer_ids=sorted(self._dense_layer_ids))
def _get_layer_kv_buf_infos(self, *, layer_ids):
buffers = [self.get_key_buffer(layer_id) for layer_id in layer_ids] + [
self.get_value_buffer(layer_id) for layer_id in layer_ids
]
return (
[buffer.data_ptr() for buffer in buffers],
[buffer.nbytes for buffer in buffers],
[buffer[0].nbytes * self.page_size for buffer in buffers],
)
def get_index_k_state_buf_infos(self): def get_index_k_state_buf_infos(self):
# Per-page item_len (MHATokenToKVPool convention); index rows share the # Per-page item_len (MHATokenToKVPool convention); index rows share the
# main-KV `loc`, so the transfer reuses the same page-ids. # main-KV `loc`, so the transfer reuses the same page-ids.
@@ -43,6 +43,7 @@ from sglang.srt.mem_cache.pool_host.common import (
get_allocator_from_storage, get_allocator_from_storage,
make_kernel_ptr_table, make_kernel_ptr_table,
) )
from sglang.srt.mem_cache.pool_host.hisparse import HiSparseHostPoolMixin
from sglang.srt.utils import is_cuda, is_hip, is_mps, is_npu, is_xpu from sglang.srt.utils import is_cuda, is_hip, is_mps, is_npu, is_xpu
_is_cuda = is_cuda() _is_cuda = is_cuda()
@@ -1098,6 +1099,76 @@ class MHATokenToKOnlyPoolHost(HostKVCache):
return ptr_list, element_size_list return ptr_list, element_size_list
class HiSparseMHATokenToKVPoolHost(HiSparseHostPoolMixin, MHATokenToKVPoolHost):
"""Layer-first MHA host pool with page-granular HiSparse allocation."""
def __init__(
self,
device_pool: MHATokenToKVPool,
host_to_device_ratio: float,
page_size: int,
):
super().__init__(
device_pool=device_pool,
host_to_device_ratio=host_to_device_ratio,
host_size=0,
page_size=page_size,
layout="layer_first",
)
def get_contiguous_buf_infos(self):
buffers = self.k_data_refs + self.v_data_refs
return (
[buffer.data_ptr() for buffer in buffers],
[buffer.nbytes for buffer in buffers],
[buffer[0].nbytes * self.page_size for buffer in buffers],
)
def load_to_device_per_layer(
self,
device_pool,
host_indices,
device_indices,
layer_id,
io_backend,
*,
is_draft: bool = False,
):
if io_backend != "kernel" or is_draft:
raise ValueError(
"MiniMax M3 HiSparse host transfers require the kernel backend."
)
host_layer = layer_id - device_pool.start_layer
transfer_kv_per_layer(
src_k=self.k_buffer[host_layer],
dst_k=device_pool.get_key_buffer(layer_id),
src_v=self.v_buffer[host_layer],
dst_v=device_pool.get_value_buffer(layer_id),
src_indices=host_indices,
dst_indices=device_indices,
item_size=self.token_stride_size,
)
def backup_from_device_all_layer(
self, device_pool, host_indices, device_indices, io_backend
):
if io_backend != "kernel":
raise ValueError(
"MiniMax M3 HiSparse host transfers require the kernel backend."
)
for layer_id in range(device_pool.start_layer, device_pool.end_layer):
host_layer = layer_id - device_pool.start_layer
transfer_kv_per_layer(
src_k=device_pool.get_key_buffer(layer_id),
dst_k=self.k_buffer[host_layer],
src_v=device_pool.get_value_buffer(layer_id),
dst_v=self.v_buffer[host_layer],
src_indices=device_indices,
dst_indices=host_indices,
item_size=self.token_stride_size,
)
class AsymmetricMHATokenToKVPoolHost(MHATokenToKVPoolHost): class AsymmetricMHATokenToKVPoolHost(MHATokenToKVPoolHost):
"""Host KV pool for MHA models whose K and V have different head dims """Host KV pool for MHA models whose K and V have different head dims
(``head_dim != v_head_dim``), e.g. MiMo-V2. (``head_dim != v_head_dim``), e.g. MiMo-V2.
+2 -1
View File
@@ -2029,7 +2029,8 @@ class KimiK3DeltaAttention(nn.Module):
qkv, g_proj_states, f_a, beta, _pad = torch.split( qkv, g_proj_states, f_a, beta, _pad = torch.split(
fused_states, self._qkvgbfa_sizes, dim=-1 fused_states, self._qkvgbfa_sizes, dim=-1
) )
forget_gate = gemm(f_a, self._bfa_f_b_w) # Fused KDA decode consumes f_a and applies f_b itself.
forget_gate = f_a if defer_f_b else gemm(f_a, self._bfa_f_b_w)
return qkv, beta, forget_gate, g_proj_states return qkv, beta, forget_gate, g_proj_states
if ( if (
+8
View File
@@ -180,6 +180,14 @@ class Lfm2VlForConditionalGeneration(nn.Module):
def get_input_embeddings(self) -> nn.Embedding: def get_input_embeddings(self) -> nn.Embedding:
return self.language_model.model.embed_tokens return self.language_model.model.embed_tokens
@property
def lm_head(self):
return self.language_model.lm_head
def set_dflash_layers_to_capture(self, layer_ids: List[int]):
# Lfm2ForCausalLM applies the HF-layer-k -> "before layer k+1" shift.
self.language_model.set_dflash_layers_to_capture(layer_ids)
def get_image_feature(self, items: List[MultimodalDataItem]) -> torch.Tensor: def get_image_feature(self, items: List[MultimodalDataItem]) -> torch.Tensor:
"""Process images through vision tower and projector. """Process images through vision tower and projector.
@@ -1,4 +1,5 @@
import unittest import unittest
from unittest.mock import patch
import torch import torch
import torch.nn.functional as F import torch.nn.functional as F
@@ -10,6 +11,8 @@ from sglang.kernels.ops.attention.linear.kda_nvidia_prefill import (
from sglang.kernels.ops.attention.linear.kda_ptx_prefill import ( from sglang.kernels.ops.attention.linear.kda_ptx_prefill import (
chunk_kda_fwd as ptx_chunk_kda_fwd, chunk_kda_fwd as ptx_chunk_kda_fwd,
) )
from sglang.srt.layers.attention.linear.kernels.kda_ptx import PtxKDAKernel
from sglang.srt.layers.attention.linear.kernels.kda_triton import TritonKDAKernel
from sglang.test.ci.ci_register import register_cuda_ci from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.test_utils import CustomTestCase from sglang.test.test_utils import CustomTestCase
@@ -79,6 +82,45 @@ def _reference(q, k, v, gate, beta, a_log, dt_bias, state, fused_qk_norm):
class TestKdaPrefill(CustomTestCase): class TestKdaPrefill(CustomTestCase):
@torch.inference_mode()
def test_ptx_padded_raw_beta(self):
"""Raw beta must match Triton, including final state after neutral padding."""
if not torch.cuda.is_available() or torch.cuda.get_device_capability() != (
10,
3,
):
self.skipTest("PTX KDA prefill requires GB300")
q, k, v, gate, beta, a_log, dt_bias, state = _inputs(2, seq_len=1025)
state.fill_(0.1)
actual_state = state.clone()
inputs = dict(
q=q,
k=k,
v=v,
g=gate,
beta=beta,
cache_indices=torch.zeros(1, device="cuda", dtype=torch.int32),
query_start_loc=torch.tensor([0, 1025], device="cuda", dtype=torch.int32),
A_log=a_log,
dt_bias=dt_bias,
lower_bound=-5.0,
beta_is_raw=True,
extend_seq_lens_cpu=[1025],
)
kernel = PtxKDAKernel()
with patch.object(
kernel._triton,
"extend",
side_effect=AssertionError("PTX unexpectedly fell back to Triton"),
):
actual = kernel.extend(**inputs, ssm_states=actual_state)
# Triton may mutate inputs, so run the reference last.
expected = TritonKDAKernel().extend(**inputs, ssm_states=state)
torch.testing.assert_close(
actual.float(), expected.float(), rtol=2e-2, atol=3e-2
)
torch.testing.assert_close(actual_state, state, rtol=2e-2, atol=3e-2)
@torch.inference_mode() @torch.inference_mode()
def test_nvidia_prefill(self): def test_nvidia_prefill(self):
if not torch.cuda.is_available() or torch.cuda.get_device_capability()[0] != 10: if not torch.cuda.is_available() or torch.cuda.get_device_capability()[0] != 10:
@@ -28,6 +28,7 @@ if is_xpu():
) )
else: else:
from sglang.kernels.ops.kvcache.hisparse import ( from sglang.kernels.ops.kvcache.hisparse import (
load_blocks_to_device_buffer_mha,
load_cache_to_device_buffer_dsv4_mla, load_cache_to_device_buffer_dsv4_mla,
load_cache_to_device_buffer_mla, load_cache_to_device_buffer_mla,
transfer_cache_dsv4_mla, transfer_cache_dsv4_mla,
@@ -368,6 +369,84 @@ def test_load_cache_to_device_buffer_hits_newest_and_updates_lru() -> None:
) )
@pytest.mark.skipif(is_xpu(), reason="MiniMax MHA block swap-in has no XPU kernel.")
def test_load_blocks_to_device_buffer_mha_handles_partial_newest_block() -> None:
"""A partial newest block must not consume slots for its invalid tail."""
sparse_block_size = 4
hot_buffer_size = 8
host_k = _host_cache()
host_v = _host_cache()
host_v.add_(1000)
device_k = torch.full(
(DEVICE_CACHE_SIZE, 1, KV_DIM), -1, dtype=DTYPE, device=DEVICE
)
device_v = torch.full_like(device_k, -1)
device_buffer_locs = torch.arange(
hot_buffer_size + 1, dtype=torch.int32, device=DEVICE
).view(1, -1)
device_buffer_tokens = torch.tensor(
[[0, 1, 2, 3, -1, -1, -1, -1, -1]],
dtype=torch.int32,
device=DEVICE,
)
for slot, token in enumerate([0, 1, 2, 3]):
device_k[device_buffer_locs[0, slot]].copy_(host_k[token], non_blocking=True)
device_v[device_buffer_locs[0, slot]].copy_(host_v[token], non_blocking=True)
device_k[device_buffer_locs[0, hot_buffer_size]].copy_(
host_k[10], non_blocking=True
)
device_v[device_buffer_locs[0, hot_buffer_size]].copy_(
host_v[10], non_blocking=True
)
top_k_blocks = torch.tensor([[0, 2]], dtype=torch.int32, device=DEVICE)
out = torch.full(
(1, top_k_blocks.size(1) * sparse_block_size),
-1,
dtype=torch.int32,
device=DEVICE,
)
lru_slots = torch.arange(hot_buffer_size, dtype=torch.int16, device=DEVICE).view(
1, -1
)
load_blocks_to_device_buffer_mha(
top_k_blocks=top_k_blocks,
device_buffer_tokens=device_buffer_tokens,
host_cache_locs=torch.arange(
HOST_CACHE_SIZE, dtype=torch.int64, device=DEVICE
).view(1, -1),
device_buffer_locs=device_buffer_locs,
host_cache_k=host_k,
host_cache_v=host_v,
device_buffer_k=device_k,
device_buffer_v=device_v,
top_k_device_locs=out,
req_pool_indices=torch.tensor([0], dtype=torch.int64, device=DEVICE),
seq_lens=torch.tensor([11], dtype=torch.int32, device=DEVICE),
lru_slots=lru_slots,
item_size_bytes=ITEM_SIZE_BYTES,
hot_buffer_size=hot_buffer_size,
sparse_block_size=sparse_block_size,
num_real_reqs=torch.tensor([1], dtype=torch.int32, device=DEVICE),
)
get_device_module().synchronize()
assert torch.equal(
out.cpu(), torch.tensor([[0, 1, 2, 3, 4, 5, 8, -1]], dtype=torch.int32)
)
assert torch.equal(device_k[4].cpu(), host_k[8])
assert torch.equal(device_v[4].cpu(), host_v[8])
assert torch.equal(device_k[5].cpu(), host_k[9])
assert torch.equal(device_v[5].cpu(), host_v[9])
assert torch.equal(
device_buffer_tokens.cpu(),
torch.tensor([[0, 1, 2, 3, 8, 9, -1, -1, -1]], dtype=torch.int32),
)
assert torch.equal(
lru_slots.cpu(), torch.tensor([[6, 7, 4, 5, 0, 1, 2, 3]], dtype=torch.int16)
)
def test_load_cache_to_device_buffer_miss_uses_updated_lru_slot() -> None: def test_load_cache_to_device_buffer_miss_uses_updated_lru_slot() -> None:
state = _long_case() state = _long_case()
@@ -0,0 +1,213 @@
"""The AITER mHC route on gfx950: gate, fallback latch, and kernel numerics vs the Torch oracle."""
import sys
import types
import unittest
from unittest.mock import patch
import torch
from sglang.kernels.ops.layernorm import mhc
from sglang.srt.environ import envs
from sglang.srt.utils import is_gfx95_supported, is_hip
from sglang.test.ci.ci_register import register_amd_ci
from sglang.test.test_utils import CustomTestCase
register_amd_ci(est_time=120, suite="stage-b-test-1-gpu-small-amd-mi35x")
@unittest.skipUnless(
torch.cuda.is_available() and is_hip() and is_gfx95_supported(),
"requires one gfx950 GPU",
)
class TestAiterMHCGLM53Flash(CustomTestCase):
hidden_size = 4096
hc_mult = 4
rms_eps = 1e-6
hc_eps = 1e-6
def setUp(self):
mhc._AITER_MHC_RUNTIME_DISABLED = False
def _inputs(self, tokens: int, seed: int = 0):
torch.manual_seed(seed)
device = torch.device("cuda")
mix_size = 2 * self.hc_mult + self.hc_mult**2
residual = (
torch.randn(
tokens,
self.hc_mult,
self.hidden_size,
device=device,
dtype=torch.bfloat16,
)
* 0.1
)
fn = (
torch.randn(
mix_size,
self.hc_mult * self.hidden_size,
device=device,
dtype=torch.float32,
)
* 0.01
)
scale = torch.tensor([0.5, 0.25, 0.25], device=device, dtype=torch.float32)
base = torch.zeros(mix_size, device=device, dtype=torch.float32)
return residual, fn, scale, base
def _rmsnorm(self, x, weight):
return (
x.float()
* torch.rsqrt(x.float().square().mean(dim=-1, keepdim=True) + self.rms_eps)
* weight.float()
).to(x.dtype)
def test_gate_selects_aiter_on_gfx950(self):
"""A gate that resolves False on real hardware silently serves the Torch path."""
with envs.SGLANG_USE_AITER.override(True):
self.assertTrue(mhc._use_aiter_mhc())
def test_hip_without_aiter_stays_on_torch_and_never_loads_tilelang(self):
"""The TileLang/DeepGEMM flags default on; only the HIP gate keeps them off this device."""
residual, fn, scale, base = self._inputs(8)
x = residual.reshape(8, self.hc_mult * self.hidden_size)
_, _, layer_ref = mhc._mhc_pre_torch(
residual, fn, scale, base, self.rms_eps, self.hc_eps, self.hc_eps, 2.0, 4
)
with (
envs.SGLANG_USE_AITER.override(False),
envs.SGLANG_OPT_USE_TILELANG_MHC_PRE.override(True),
envs.SGLANG_OPT_USE_TILELANG_MHC_POST.override(True),
envs.SGLANG_OPT_DEEPGEMM_HC_PRENORM.override(True),
patch.object(
mhc, "_load_tilelang", side_effect=AssertionError("TileLang imported")
),
):
self.assertFalse(mhc._use_aiter_mhc())
self.assertFalse(mhc._use_tilelang_mhc_pre())
self.assertFalse(mhc._use_tilelang_mhc_post())
self.assertFalse(mhc._use_deep_gemm_hc_prenorm())
layer_input, h_res, h_post, norm_fused = mhc.hc_pre(
x, fn, scale, base, self.hc_mult, self.rms_eps, self.hc_eps, 4
)
out = mhc.hc_post(layer_input, x, h_post, h_res, self.hc_mult)
self.assertFalse(norm_fused)
torch.testing.assert_close(layer_input, layer_ref)
self.assertTrue(torch.isfinite(out).all())
def test_aiter_import_and_runtime_failures_latch_to_torch(self):
"""A missing symbol or a raising kernel must disable the route once, not fail the request."""
residual, fn, scale, base = self._inputs(8)
x = residual.reshape(8, self.hc_mult * self.hidden_size)
modules = {
"aiter": types.ModuleType("aiter"),
"aiter.ops": types.ModuleType("aiter.ops"),
"aiter.ops.mhc": types.ModuleType("aiter.ops.mhc"),
}
with patch.dict(sys.modules, modules):
result = mhc._try_aiter_mhc_pre(
residual,
fn,
scale,
base,
self.rms_eps,
self.hc_eps,
self.hc_eps,
2.0,
4,
None,
None,
)
self.assertIsNone(result)
self.assertTrue(mhc._AITER_MHC_RUNTIME_DISABLED)
mhc._AITER_MHC_RUNTIME_DISABLED = False
def fail_post(*_args, **_kwargs):
raise RuntimeError("synthetic failure")
failing = types.ModuleType("aiter.ops.mhc")
failing.mhc_post = fail_post
modules["aiter.ops.mhc"] = failing
with envs.SGLANG_USE_AITER.override(False):
layer_input, h_res, h_post, _ = mhc.hc_pre(
x, fn, scale, base, self.hc_mult, self.rms_eps, self.hc_eps, 4
)
with (
patch.dict(sys.modules, modules),
envs.SGLANG_USE_AITER.override(True),
):
out = mhc.hc_post(layer_input, x, h_post, h_res, self.hc_mult)
self.assertTrue(mhc._AITER_MHC_RUNTIME_DISABLED)
self.assertTrue(torch.isfinite(out).all())
def test_aiter_pre_post_match_torch_oracle(self):
"""A positional or kwarg mixup in the AITER call shows up only against the real kernel."""
norm_weight = torch.linspace(
0.75, 1.25, self.hidden_size, device="cuda", dtype=torch.bfloat16
)
for tokens in (1, 8, 17, 32, 64, 128):
for sinkhorn_iters in (2, 20):
for fused_norm in (False, True):
with self.subTest(
tokens=tokens, sinkhorn_iters=sinkhorn_iters, norm=fused_norm
):
residual, fn, scale, base = self._inputs(tokens)
post_ref, comb_ref, layer_ref = mhc._mhc_pre_torch(
residual,
fn,
scale,
base,
self.rms_eps,
self.hc_eps,
self.hc_eps,
2.0,
sinkhorn_iters,
)
result = mhc._try_aiter_mhc_pre(
residual,
fn,
scale,
base,
self.rms_eps,
self.hc_eps,
self.hc_eps,
2.0,
sinkhorn_iters,
norm_weight if fused_norm else None,
self.rms_eps if fused_norm else None,
)
self.assertIsNotNone(result, "AITER mHC pre fell back")
post_out, comb_out, layer_out = result
if fused_norm:
layer_ref = self._rmsnorm(layer_ref, norm_weight)
torch.cuda.synchronize()
torch.testing.assert_close(
post_out, post_ref, atol=2e-3, rtol=2e-3
)
torch.testing.assert_close(
comb_out, comb_ref, atol=2e-3, rtol=2e-3
)
torch.testing.assert_close(
layer_out, layer_ref, atol=2e-2, rtol=2e-2
)
x = (layer_ref.float() * 0.75).to(layer_ref.dtype)
post_ref_out = mhc._mhc_post_torch(
x, residual, post_ref, comb_ref
)
post_out_actual = mhc._try_aiter_mhc_post(
x, residual, post_out, comb_out
)
self.assertIsNotNone(
post_out_actual, "AITER mHC post fell back"
)
torch.testing.assert_close(
post_out_actual, post_ref_out, atol=2e-2, rtol=2e-2
)
if __name__ == "__main__":
unittest.main()
@@ -53,14 +53,19 @@ def _make_kv_pool(start_layer: int = 0) -> MiniMaxSparseKVPool:
class TestMiniMaxSparseDisaggStateKvArgs(unittest.TestCase): class TestMiniMaxSparseDisaggStateKvArgs(unittest.TestCase):
def test_setup_state_kv_args_single_minimax_component(self): def test_setup_state_kv_args_minimax_components(self):
pool = _make_k_only_pool() pool = _make_k_only_pool()
kv_args = KVArgs() kv_args = KVArgs()
setup_state_kv_args(kv_args, pool) setup_state_kv_args(kv_args, pool)
self.assertEqual(kv_args.state_types, [StateType.MINIMAX_INDEX_K]) self.assertEqual(
self.assertEqual(len(kv_args.state_data_ptrs), 1) kv_args.state_types,
[StateType.MINIMAX_INDEX_K, StateType.MINIMAX_DENSE_KV],
)
self.assertEqual(len(kv_args.state_data_ptrs), 2)
self.assertEqual(len(kv_args.state_data_ptrs[0]), pool.index_k_pool.layer_num) self.assertEqual(len(kv_args.state_data_ptrs[0]), pool.index_k_pool.layer_num)
self.assertEqual(len(kv_args.state_item_lens[0]), pool.index_k_pool.layer_num) self.assertEqual(len(kv_args.state_item_lens[0]), pool.index_k_pool.layer_num)
self.assertEqual(len(kv_args.state_data_ptrs[1]), 6)
self.assertEqual(len(kv_args.state_item_lens[1]), 6)
def test_index_kv_pool_raises(self): def test_index_kv_pool_raises(self):
pool = _make_kv_pool() pool = _make_kv_pool()
@@ -1,4 +1,5 @@
import concurrent.futures import concurrent.futures
import ctypes
import unittest import unittest
from threading import Event from threading import Event
from types import SimpleNamespace from types import SimpleNamespace
@@ -6,6 +7,7 @@ from unittest.mock import MagicMock, call, patch
import numpy as np import numpy as np
from sglang.srt.disaggregation.base.conn import StateType
from sglang.srt.disaggregation.mooncake.conn import MooncakeKVManager from sglang.srt.disaggregation.mooncake.conn import MooncakeKVManager
from sglang.test.ci.ci_register import register_cpu_ci from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase from sglang.test.test_utils import CustomTestCase
@@ -132,6 +134,61 @@ class TestMooncakeTransferBatching(unittest.TestCase):
) )
class TestMiniMaxStateTransfer(CustomTestCase):
def test_index_truncates_but_dense_rejects_mismatched_page_lists(self):
"""Legacy index transfers copy the common prefix; incomplete dense KV must fail."""
def copy_bytes(session, sources, destinations, lengths):
for src, dst, length in zip(sources, destinations, lengths, strict=True):
ctypes.memmove(dst, src, length)
return 0
for state in (StateType.MINIMAX_INDEX_K, StateType.MINIMAX_DENSE_KV):
for src_pages, dst_pages in (([1], [0]), ([1, 2], [0]), ([1], [0, 2])):
with self.subTest(state=state, src=src_pages, dst=dst_pages):
src = np.arange(3, dtype=np.int32)
dst = np.full(3, -1, dtype=np.int32)
manager = MooncakeKVManager.__new__(MooncakeKVManager)
manager.kv_args = SimpleNamespace(
state_types=[state],
state_data_ptrs=[[src.ctypes.data]],
state_item_lens=[[src.itemsize]],
state_dim_per_tensor=[[]],
state_layer_ids=[[]],
)
manager.engine = SimpleNamespace(batch_transfer_sync=copy_bytes)
manager.pp_size = manager.attn_tp_size = 1
manager.is_mla_backend = manager.is_hybrid_mla_backend = False
manager.enable_custom_mem_pool = False
manager.max_transfer_batch_indices = 0
peer = SimpleNamespace(
dst_state_data_ptrs=[[dst.ctypes.data]],
dst_state_item_lens=[[dst.itemsize]],
dst_state_dim_per_tensor=[[]],
dst_state_layer_ids=[[]],
dst_attn_tp_size=1,
)
kwargs = dict(
req=SimpleNamespace(
mooncake_session_id="cpu", dst_state_indices=[dst_pages]
),
prefill_state_indices=[src_pages],
executor=None,
target_rank_registration_info=peer,
)
if state == StateType.MINIMAX_DENSE_KV and len(src_pages) != len(
dst_pages
):
with self.assertRaisesRegex(
RuntimeError, "state index length mismatch"
):
manager.maybe_send_extra(**kwargs)
np.testing.assert_array_equal(dst, [-1, -1, -1])
else:
self.assertEqual(manager.maybe_send_extra(**kwargs), 0)
np.testing.assert_array_equal(dst, [1, -1, -1])
class TestDcpDraftHeadTransfer(unittest.TestCase): class TestDcpDraftHeadTransfer(unittest.TestCase):
def test_transfers_draft_heads_to_logical_destination_rows(self): def test_transfers_draft_heads_to_logical_destination_rows(self):
for src_tp, dst_tp in ((4, 8), (8, 4), (8, 8), (4, 32), (32, 4)): for src_tp, dst_tp in ((4, 8), (8, 4), (8, 8), (4, 32), (32, 4)):
@@ -95,8 +95,10 @@ class TestPtxKDATrackRouting(CustomTestCase):
kernel = self._make_kernel() kernel = self._make_kernel()
kernel._triton = _RejectTriton() kernel._triton = _RejectTriton()
h = torch.zeros(3, 2, 128, 128, dtype=torch.float32) h = torch.zeros(3, 2, 128, 128, dtype=torch.float32)
beta_flags = []
def fake_fwd(*args, **kwargs): def fake_fwd(*args, **kwargs):
beta_flags.append(kwargs["use_beta_sigmoid_in_kernel"])
return [ return [
args[2].clone(), # out == v args[2].clone(), # out == v
kwargs["initial_state"].clone(), # final_state kwargs["initial_state"].clone(), # final_state
@@ -107,24 +109,16 @@ class TestPtxKDATrackRouting(CustomTestCase):
kernel._fwd = fake_fwd kernel._fwd = fake_fwd
x = self._inputs() x = self._inputs()
out, h_out = kernel.extend( for beta_kwargs in ({"beta_is_raw": True}, {}, {"beta_is_raw": False}):
x["q"], out, h_out = kernel.extend(
x["k"], **x,
x["v"], **beta_kwargs,
x["g"], return_intermediate_states=True,
x["beta"], track_ssm_h_src=torch.empty(0, dtype=torch.long),
ssm_states=x["ssm_states"], )
cache_indices=x["cache_indices"], self.assertEqual(tuple(out.shape), (1, 164, 2, 128))
query_start_loc=x["query_start_loc"], self.assertIs(h_out, h)
A_log=x["A_log"], self.assertEqual(beta_flags, [True, False, False])
dt_bias=x["dt_bias"],
return_intermediate_states=True,
track_ssm_h_src=torch.empty(0, dtype=torch.long),
extend_seq_lens_cpu=x["extend_seq_lens_cpu"],
)
self.assertEqual(tuple(out.shape), (1, 164, 2, 128))
self.assertIs(h_out, h)
if __name__ == "__main__": if __name__ == "__main__":
@@ -132,6 +132,28 @@ class TestModelOptNvfp4(CustomTestCase):
use_per_token_activation=True, use_per_token_activation=True,
) )
def test_shared_expert_fusion_requires_matching_fp4_precision(self):
quantized_shared = ModelOptFp4Config(
is_checkpoint_nvfp4_serialized=True,
group_size=16,
)
bf16_shared = ModelOptFp4Config(
is_checkpoint_nvfp4_serialized=True,
group_size=16,
exclude_modules=["model.layers.*.mlp.shared_experts*"],
)
gate_only_bf16 = ModelOptFp4Config(
is_checkpoint_nvfp4_serialized=True,
group_size=16,
exclude_modules=["model.layers.*.mlp.shared_expert_gate"],
)
self.assertTrue(quantized_shared.can_fuse_shared_expert())
self.assertFalse(bf16_shared.can_fuse_shared_expert())
# Only the gate is BF16 (Qwen3-Next NVFP4): the FP4 body still fuses.
self.assertTrue(gate_only_bf16.can_fuse_shared_expert())
if __name__ == "__main__": if __name__ == "__main__":
unittest.main() unittest.main()
@@ -1,7 +1,7 @@
import unittest import unittest
from concurrent.futures import ThreadPoolExecutor from concurrent.futures import ThreadPoolExecutor
from types import SimpleNamespace from types import SimpleNamespace
from unittest.mock import MagicMock from unittest.mock import MagicMock, patch
import numpy as np import numpy as np
import torch import torch
@@ -18,6 +18,71 @@ from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=10, suite="base-a-test-cpu") register_cpu_ci(est_time=10, suite="base-a-test-cpu")
class TestHiSparseDecodeRemap(CustomTestCase):
def test_page_size_one_reclaims_temporary_device_slot(self):
"""Decode remapping must reclaim its temporary slot without freeing the live slot."""
from sglang.srt.managers.hisparse_coordinator import HiSparseCoordinator
from sglang.srt.mem_cache.allocator.hisparse import (
HiSparseTokenToKVPoolAllocator,
)
from sglang.srt.mem_cache.memory_pool import MiniMaxSparseKVPool
pool = MiniMaxSparseKVPool(
size=8,
page_size=1,
dtype=torch.float32,
head_num=1,
head_dim=8,
idx_head_dim=16,
dense_layer_ids=[0],
sparse_layer_ids=[1],
disable_value_sparse_layer_ids=[1],
device="cpu",
start_layer=0,
end_layer=2,
enable_hisparse=True,
)
allocator = HiSparseTokenToKVPoolAllocator(
size=pool.size,
page_size=1,
dtype=pool.dtype,
device="cpu",
kvcache=pool,
need_sort=False,
)
coordinator = HiSparseCoordinator.__new__(HiSparseCoordinator)
coordinator.is_dsv4_hisparse = False
coordinator.mem_pool_device = pool.main_pool
coordinator.token_to_kv_pool_allocator = allocator
coordinator.device_buffer_size = 2
coordinator.req_to_device_buffer = allocator.hisparse_attn_allocator.alloc(
3
).reshape(1, 3)
coordinator.req_device_buffer_size = torch.tensor([3])
coordinator.req_device_buffer_token_locs = torch.zeros(
(1, 1, 3), dtype=torch.int32
)
coordinator._skip_first_backup = [True]
out_loc = allocator.alloc(1)
with patch("sglang.srt.managers.hisparse_coordinator._is_hip", False):
for _ in range(2):
coordinator._skip_first_backup[0] = True
coordinator.map_last_loc_to_buffer(
seq_lens=torch.tensor([3]),
out_cache_loc=out_loc,
req_pool_indices=torch.tensor([0]),
seq_lens_cpu=torch.tensor([3]),
req_pool_indices_cpu=torch.tensor([0]),
)
self.assertEqual(
allocator.hisparse_attn_allocator.available_size(), pool.size - 3
)
torch.testing.assert_close(
allocator.full_to_hisparse_device_index_mapping[out_loc],
coordinator.req_to_device_buffer[:, 2],
)
class TestDeepSeekV4HiSparseAllocator(CustomTestCase): class TestDeepSeekV4HiSparseAllocator(CustomTestCase):
def setUp(self): def setUp(self):
# The code under test reads its config from the bags. # The code under test reads its config from the bags.
@@ -1,14 +1,23 @@
import unittest import unittest
from unittest.mock import patch
import torch import torch
from sglang.srt.disaggregation.utils import get_kv_transfer_buf_infos
from sglang.srt.mem_cache.memory_pool import MiniMaxSparseKVPool from sglang.srt.mem_cache.memory_pool import MiniMaxSparseKVPool
from sglang.srt.mem_cache.pool_host.mha import (
HiSparseMHATokenToKVPoolHost,
MHATokenToKVPoolHost,
)
from sglang.test.ci.ci_register import register_cpu_ci from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=10, suite="base-a-test-cpu") register_cpu_ci(est_time=10, suite="base-a-test-cpu")
def _make_k_only_pool(start_layer: int = 0) -> MiniMaxSparseKVPool: def _make_k_only_pool(
start_layer: int = 0, *, enable_hisparse: bool = False
) -> MiniMaxSparseKVPool:
"""Mirror the released MiniMax-M3 config shape: all sparse layers K-only.""" """Mirror the released MiniMax-M3 config shape: all sparse layers K-only."""
dense_layer_ids = [start_layer, start_layer + 1, start_layer + 2] dense_layer_ids = [start_layer, start_layer + 1, start_layer + 2]
sparse_layer_ids = [start_layer + 3 + i for i in range(4)] sparse_layer_ids = [start_layer + 3 + i for i in range(4)]
@@ -26,10 +35,11 @@ def _make_k_only_pool(start_layer: int = 0) -> MiniMaxSparseKVPool:
device="cpu", device="cpu",
start_layer=start_layer, start_layer=start_layer,
end_layer=end_layer, end_layer=end_layer,
enable_hisparse=enable_hisparse,
) )
class TestMiniMaxSparsePoolPD(unittest.TestCase): class TestMiniMaxSparsePoolPD(CustomTestCase):
def test_contiguous_buf_infos_main_only(self): def test_contiguous_buf_infos_main_only(self):
pool = _make_k_only_pool() pool = _make_k_only_pool()
ptrs, lens, item_lens = pool.get_contiguous_buf_infos() ptrs, lens, item_lens = pool.get_contiguous_buf_infos()
@@ -53,6 +63,52 @@ class TestMiniMaxSparsePoolPD(unittest.TestCase):
self.assertEqual(lens[i], buf.nbytes) self.assertEqual(lens[i], buf.nbytes)
self.assertEqual(item_lens[i], buf[0].nbytes * pool.page_size) self.assertEqual(item_lens[i], buf[0].nbytes * pool.page_size)
def test_hisparse_host_registration(self):
"""PD startup must register every sparse host K/V buffer with page strides."""
pool = _make_k_only_pool(enable_hisparse=True)
host = HiSparseMHATokenToKVPoolHost.__new__(HiSparseMHATokenToKVPoolHost)
with patch(
"sglang.srt.mem_cache.pool_host.base.host_memory_budget_bytes",
return_value=1 << 30,
):
MHATokenToKVPoolHost.__init__(
host,
device_pool=pool.main_pool,
host_to_device_ratio=2,
host_size=0,
page_size=pool.page_size,
layout="layer_first",
pin_memory=False,
)
ptrs, lens, item_lens = get_kv_transfer_buf_infos(host)
buffers = list(host.k_buffer.unbind()) + list(host.v_buffer.unbind())
self.assertEqual(len(buffers), 8)
self.assertEqual(ptrs, [buffer.data_ptr() for buffer in buffers])
self.assertEqual(lens, [buffer.nbytes for buffer in buffers])
self.assertEqual(
item_lens, [buffer[0].nbytes * pool.page_size for buffer in buffers]
)
def test_pd_registration_separates_dense_and_sparse_layers(self):
"""Both PD peers must keep dense device KV out of the sparse transfer list."""
for hisparse in (False, True):
pool = _make_k_only_pool(enable_hisparse=hisparse)
for layers, infos in (
(range(3, 7), get_kv_transfer_buf_infos(pool)),
(range(3), pool.get_dense_kv_state_buf_infos()),
):
with self.subTest(hisparse=hisparse, layers=layers):
buffers = [pool.get_key_buffer(i) for i in layers] + [
pool.get_value_buffer(i) for i in layers
]
ptrs, lens, item_lens = infos
self.assertEqual(ptrs, [buffer.data_ptr() for buffer in buffers])
self.assertEqual(lens, [buffer.nbytes for buffer in buffers])
self.assertEqual(
item_lens,
[buffer[0].nbytes * pool.page_size for buffer in buffers],
)
if __name__ == "__main__": if __name__ == "__main__":
unittest.main() unittest.main()