[Model] Complete dots.note.omni support with native encoders, video preprocessing, and MTP decoding (#33829)
Co-authored-by: miraclezqc <dysania@pku.edu.cn>
This commit is contained in:
co-authored by
miraclezqc
parent
c35683fda0
commit
af39ad9349
@@ -57,6 +57,14 @@ import { config } from "/src/snippets/configs/rednote/dots3-note.jsx";
|
|||||||
|
|
||||||
<Deployment config={config} />
|
<Deployment config={config} />
|
||||||
|
|
||||||
|
## Playground
|
||||||
|
|
||||||
|
The Playground is where you experiment with **SGLang features beyond the verified matrix**. The Deploy panel above only emits combinations signed off on this page; the Playground lets you turn on additional knobs on top of whichever cell the Deploy panel is currently showing.
|
||||||
|
|
||||||
|
import { Playground } from "/src/snippets/_playground.jsx";
|
||||||
|
|
||||||
|
<Playground config={config} />
|
||||||
|
|
||||||
## 1. Model Introduction
|
## 1. Model Introduction
|
||||||
|
|
||||||
dots3.note is RedNote's native multimodal omni model, built on the dots3 language model. It accepts text, image, audio, and native video input.
|
dots3.note is RedNote's native multimodal omni model, built on the dots3 language model. It accepts text, image, audio, and native video input.
|
||||||
@@ -66,7 +74,12 @@ dots3.note is RedNote's native multimodal omni model, built on the dots3 languag
|
|||||||
- **Hybrid attention** — dots3 combines MLA with full-attention and sliding-window layers of different geometry, attention gates, and optional DSA indexing on full-attention layers.
|
- **Hybrid attention** — dots3 combines MLA with full-attention and sliding-window layers of different geometry, attention gates, and optional DSA indexing on full-attention layers.
|
||||||
- **MTP speculative decoding** — a full-sharing MTP/NextN architecture exposes one recursively shared, SWA-shaped MTP layer and shares the target LM head.
|
- **MTP speculative decoding** — a full-sharing MTP/NextN architecture exposes one recursively shared, SWA-shaped MTP layer and shares the target LM head.
|
||||||
|
|
||||||
**Resources:** [Hugging Face](https://huggingface.co/dots-studio/dots3-note-prev) · [SGLang PR #33829](https://github.com/sgl-project/sglang/pull/33829)
|
**Available checkpoints:**
|
||||||
|
|
||||||
|
- **BF16**: [dots-studio/dots3-note-prev](https://huggingface.co/dots-studio/dots3-note-prev)
|
||||||
|
- **FP8**: [dots-studio/dots3-note-prev-fp8](https://huggingface.co/dots-studio/dots3-note-prev-fp8)
|
||||||
|
|
||||||
|
**Resources:** [Hugging Face (BF16)](https://huggingface.co/dots-studio/dots3-note-prev) · [Hugging Face (FP8)](https://huggingface.co/dots-studio/dots3-note-prev-fp8) · [SGLang PR #33829](https://github.com/sgl-project/sglang/pull/33829)
|
||||||
|
|
||||||
## 2. Configuration Tips
|
## 2. Configuration Tips
|
||||||
|
|
||||||
@@ -112,11 +125,13 @@ response = client.chat.completions.create(
|
|||||||
}
|
}
|
||||||
],
|
],
|
||||||
extra_body={
|
extra_body={
|
||||||
|
"video_config": {
|
||||||
"seq": 131072,
|
"seq": 131072,
|
||||||
"audio_cap": 0.5,
|
"audio_cap": 0.5,
|
||||||
"audio_sr": 16000,
|
"audio_sr": 16000,
|
||||||
"k_mode": "eval_ek",
|
"k_mode": "eval_ek",
|
||||||
},
|
},
|
||||||
|
},
|
||||||
)
|
)
|
||||||
|
|
||||||
print(response.choices[0].message.content)
|
print(response.choices[0].message.content)
|
||||||
@@ -132,28 +147,93 @@ Pending update...
|
|||||||
|
|
||||||
</Accordion>
|
</Accordion>
|
||||||
|
|
||||||
Per-request video preprocessing controls (all optional, passed via `extra_body`):
|
Per-request video preprocessing controls are grouped under `video_config` in
|
||||||
|
`extra_body`:
|
||||||
|
|
||||||
| Field | Default | Purpose |
|
<table style={{width: "100%", borderCollapse: "collapse", tableLayout: "fixed"}}>
|
||||||
|-------|---------|---------|
|
<colgroup>
|
||||||
| `seq` | `131072` | Total sequence budget used by the video flattener. |
|
<col style={{width: "22%"}} />
|
||||||
| `audio_cap` | `1.0` | Maximum fraction of the input budget assigned to audio; `0` disables audio processing. |
|
<col style={{width: "18%"}} />
|
||||||
| `audio_sr` | `16000` | Audio sample rate. |
|
<col style={{width: "60%"}} />
|
||||||
| `k_mode` | `eval_ek` | Deterministic evaluation/sampling mode of the flattener. |
|
</colgroup>
|
||||||
|
<thead>
|
||||||
|
<tr style={{borderBottom: "2px solid #d55816"}}>
|
||||||
|
<th style={{textAlign: "left", padding: "10px 12px", fontWeight: 700, backgroundColor: "rgba(255,255,255,0.02)"}}>Field</th>
|
||||||
|
<th style={{textAlign: "left", padding: "10px 12px", fontWeight: 700, backgroundColor: "rgba(255,255,255,0.05)"}}>Default</th>
|
||||||
|
<th style={{textAlign: "left", padding: "10px 12px", fontWeight: 700, backgroundColor: "rgba(255,255,255,0.02)"}}>Purpose</th>
|
||||||
|
</tr>
|
||||||
|
</thead>
|
||||||
|
<tbody>
|
||||||
|
<tr>
|
||||||
|
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}><code>seq</code></td>
|
||||||
|
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}><code>131072</code></td>
|
||||||
|
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}>Total sequence budget used by the video flattener.</td>
|
||||||
|
</tr>
|
||||||
|
<tr>
|
||||||
|
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}><code>audio_cap</code></td>
|
||||||
|
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}><code>1.0</code></td>
|
||||||
|
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}>Maximum fraction of the input budget assigned to audio; <code>0</code> disables audio processing.</td>
|
||||||
|
</tr>
|
||||||
|
<tr>
|
||||||
|
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}><code>audio_sr</code></td>
|
||||||
|
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}><code>16000</code></td>
|
||||||
|
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}>Audio sample rate.</td>
|
||||||
|
</tr>
|
||||||
|
<tr>
|
||||||
|
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}><code>k_mode</code></td>
|
||||||
|
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}><code>eval_ek</code></td>
|
||||||
|
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}>Deterministic evaluation/sampling mode of the flattener.</td>
|
||||||
|
</tr>
|
||||||
|
</tbody>
|
||||||
|
</table>
|
||||||
|
|
||||||
These controls are request-scoped so that evaluation jobs with different context budgets can share one server. The flattener reserves room for `max_new_tokens` inside the budget and falls back to visual-only processing if audio would exceed the configured token budget.
|
These controls are request-scoped so that evaluation jobs with different context budgets can share one server. For example: `extra_body={"video_config": {"seq": 131072, "audio_cap": 0.5}}`. The flattener reserves room for `max_new_tokens` inside the budget and falls back to visual-only processing if audio would exceed the configured token budget.
|
||||||
|
|
||||||
<Warning>
|
A request may carry several videos, and videos can be mixed with image and audio parts. Each video is flattened independently under the same per-request budget, and the flattened frames and audio segments are spliced back at the position of their `video_url` part, so the modality ordering of the prompt is preserved.
|
||||||
Native video currently supports one video per request, and a native video cannot be mixed with separate image or audio inputs in the same request.
|
|
||||||
</Warning>
|
|
||||||
|
|
||||||
### 3.2 Image and audio input
|
### 3.2 Image and audio input
|
||||||
|
|
||||||
Outside the native-video path, images and audio clips use the standard OpenAI multimodal message format and SGLang's multimodal serving (`--enable-multimodal` is in every cell). The vision and audio towers run in-process, so no extra server is needed.
|
Outside the native-video path, images and audio clips use the standard OpenAI multimodal message format. `--enable-multimodal` is in every cell; the vision and audio towers run in-process, so no extra server is needed.
|
||||||
|
|
||||||
|
<Accordion title="Image Example (Python)">
|
||||||
|
|
||||||
|
```python Example
|
||||||
|
from openai import OpenAI
|
||||||
|
|
||||||
|
client = OpenAI(base_url="http://localhost:30000/v1", api_key="EMPTY")
|
||||||
|
|
||||||
|
response = client.chat.completions.create(
|
||||||
|
model="dots3.note",
|
||||||
|
messages=[
|
||||||
|
{
|
||||||
|
"role": "user",
|
||||||
|
"content": [
|
||||||
|
{
|
||||||
|
"type": "image_url",
|
||||||
|
"image_url": {"url": "https://example.com/sample.jpg"},
|
||||||
|
},
|
||||||
|
{"type": "text", "text": "Describe this image."},
|
||||||
|
],
|
||||||
|
}
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
|
print(response.choices[0].message.content)
|
||||||
|
```
|
||||||
|
|
||||||
|
</Accordion>
|
||||||
|
|
||||||
|
<Accordion title="Example Output">
|
||||||
|
|
||||||
|
```text Output
|
||||||
|
Pending update...
|
||||||
|
```
|
||||||
|
|
||||||
|
</Accordion>
|
||||||
|
|
||||||
### 3.3 Tool Calling
|
### 3.3 Tool Calling
|
||||||
|
|
||||||
The cells launch with `--tool-call-parser dots`, so structured tool calls surface via `message.tool_calls` out of the box.
|
Toggle **Tool Call Parser** (`--tool-call-parser dots`) and **Reasoning Parser** (`--reasoning-parser dots`) in the **Parsers** card of the [Playground above](#playground). Structured tool calls then surface via `message.tool_calls`.
|
||||||
|
|
||||||
<Accordion title="Tool Calling Example (Python)">
|
<Accordion title="Tool Calling Example (Python)">
|
||||||
|
|
||||||
@@ -195,7 +275,7 @@ Pending update...
|
|||||||
|
|
||||||
### 3.4 Encoder/LLM Disaggregation (EPD)
|
### 3.4 Encoder/LLM Disaggregation (EPD)
|
||||||
|
|
||||||
`Dot3NoteForCausalLM` supports both roles of an encoder/LLM-disaggregated deployment:
|
`Dots3NoteForCausalLM` supports both roles of an encoder/LLM-disaggregated deployment:
|
||||||
|
|
||||||
- **Encoder role** — serve with `--encoder-only`; the instance runs only the vision and audio towers.
|
- **Encoder role** — serve with `--encoder-only`; the instance runs only the vision and audio towers.
|
||||||
- **Language role** — serve with `--language-only`; the instance skips tower construction, leaving the memory to the language model.
|
- **Language role** — serve with `--language-only`; the instance skips tower construction, leaving the memory to the language model.
|
||||||
|
|||||||
@@ -1,13 +1,9 @@
|
|||||||
// Dots3-Note cookbook config. Consumed by _deployment.jsx + _playground.jsx.
|
// Dots3-Note cookbook config. Consumed by _deployment.jsx + _playground.jsx.
|
||||||
// Single `export const config` literal - no spreads/calls/IIFE (Mintlify re-evals at hydration).
|
// Single `export const config` literal — no spreads/calls/IIFE (Mintlify re-evals at hydration).
|
||||||
|
|
||||||
export const config = {
|
export const config = {
|
||||||
modelName: "Dots3-Note",
|
modelName: "Dots3-Note",
|
||||||
|
|
||||||
// No Playground on this page — the only extra knob (the dots tool-call parser)
|
|
||||||
// is already baked into the cells.
|
|
||||||
showPlaygroundLink: false,
|
|
||||||
|
|
||||||
// Hopper only for now — no Blackwell support.
|
// Hopper only for now — no Blackwell support.
|
||||||
supportedHardware: ["h200", "h100"],
|
supportedHardware: ["h200", "h100"],
|
||||||
|
|
||||||
@@ -50,14 +46,78 @@ export const config = {
|
|||||||
{"type": "video_url", "video_url": {"url": "https://example.com/sample.mp4"}},
|
{"type": "video_url", "video_url": {"url": "https://example.com/sample.mp4"}},
|
||||||
{"type": "text", "text": "Summarize what happens in this video."}
|
{"type": "text", "text": "Summarize what happens in this video."}
|
||||||
]
|
]
|
||||||
}]
|
}],
|
||||||
|
"video_config": {
|
||||||
|
"seq": 131072,
|
||||||
|
"audio_cap": 0.5,
|
||||||
|
"audio_sr": 16000,
|
||||||
|
"k_mode": "eval_ek"
|
||||||
|
}
|
||||||
}'`,
|
}'`,
|
||||||
|
|
||||||
dockerImages: {
|
dockerImages: {
|
||||||
h200: "lmsysorg/sglang:dev",
|
h200: "lmsysorg/sglang:dev-dots3-note",
|
||||||
h100: "lmsysorg/sglang:dev",
|
h100: "lmsysorg/sglang:dev-dots3-note",
|
||||||
},
|
},
|
||||||
|
|
||||||
|
github: {
|
||||||
|
cookbookModel: "dots-studio/dots3-note-prev",
|
||||||
|
},
|
||||||
|
|
||||||
|
playgroundFeatures: {
|
||||||
|
attention: {
|
||||||
|
knobs: [
|
||||||
|
{ id: "tp", label: "TP", values: [null, 4, 8] },
|
||||||
|
{
|
||||||
|
id: "dpAttn",
|
||||||
|
label: "DP-Attention",
|
||||||
|
values: [null, false, 4, 8],
|
||||||
|
labels: { "auto": "Auto", "false": "Off" },
|
||||||
|
},
|
||||||
|
],
|
||||||
|
},
|
||||||
|
moe: {
|
||||||
|
backend: {
|
||||||
|
options: [
|
||||||
|
{ id: null, label: "Inherited" },
|
||||||
|
{ id: "deepep", label: "DeepEP", flags: ["--moe-a2a-backend deepep"] },
|
||||||
|
],
|
||||||
|
},
|
||||||
|
ep: { label: "EP", values: [null, 4, 8] },
|
||||||
|
},
|
||||||
|
parsers: {
|
||||||
|
items: [
|
||||||
|
{
|
||||||
|
id: "reasoning",
|
||||||
|
label: "Reasoning Parser",
|
||||||
|
flag: "--reasoning-parser dots",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
id: "toolCall",
|
||||||
|
label: "Tool Call Parser",
|
||||||
|
flag: "--tool-call-parser dots",
|
||||||
|
},
|
||||||
|
],
|
||||||
|
},
|
||||||
|
speculative: {
|
||||||
|
options: [
|
||||||
|
{ id: "current", label: "Inherited from base" },
|
||||||
|
{ id: "off", label: "Off (greedy)" },
|
||||||
|
{
|
||||||
|
id: "nextn-314",
|
||||||
|
label: "NEXTN / MTP 3-1-4",
|
||||||
|
flags: [
|
||||||
|
"--speculative-algorithm NEXTN",
|
||||||
|
"--speculative-num-steps 3",
|
||||||
|
"--speculative-eagle-topk 1",
|
||||||
|
"--speculative-num-draft-tokens 4",
|
||||||
|
"--speculative-draft-model-path {{MODEL_NAME}}",
|
||||||
|
"--speculative-draft-attention-backend fa3",
|
||||||
|
],
|
||||||
|
},
|
||||||
|
],
|
||||||
|
},
|
||||||
|
},
|
||||||
|
|
||||||
cells: [
|
cells: [
|
||||||
{
|
{
|
||||||
@@ -74,20 +134,25 @@ export const config = {
|
|||||||
],
|
],
|
||||||
flags: [
|
flags: [
|
||||||
"--model-path {{MODEL_NAME}}",
|
"--model-path {{MODEL_NAME}}",
|
||||||
"--context-length 524288",
|
"--trust-remote-code",
|
||||||
"--enable-dp-attention",
|
"--enable-dp-attention",
|
||||||
"--dp-size 8",
|
"--tp 8",
|
||||||
"--tp-size 8",
|
"--dp 8",
|
||||||
"--ep-size 8",
|
"--ep 8",
|
||||||
|
"--moe-dense-tp-size 1",
|
||||||
|
"--moe-a2a-backend deepep",
|
||||||
|
"--moe-runner-backend deep_gemm",
|
||||||
|
"--deepep-dispatcher-output-dtype bf16",
|
||||||
|
"--deepep-mode auto",
|
||||||
|
"--enable-nccl-nvls",
|
||||||
|
"--context-length 524288",
|
||||||
"--mem-fraction-static 0.87",
|
"--mem-fraction-static 0.87",
|
||||||
"--max-running-requests 256",
|
"--max-running-requests 256",
|
||||||
"--chunked-prefill-size 16384",
|
"--chunked-prefill-size 16384",
|
||||||
"--trust-remote-code",
|
|
||||||
"--swa-full-tokens-ratio 0.03",
|
"--swa-full-tokens-ratio 0.03",
|
||||||
"--prefill-attention-backend fa3",
|
"--prefill-attention-backend fa3",
|
||||||
"--decode-attention-backend fa3",
|
"--decode-attention-backend fa3",
|
||||||
"--page-size 64",
|
"--page-size 64",
|
||||||
"--moe-dense-tp-size 1",
|
|
||||||
"--cuda-graph-backend-decode full",
|
"--cuda-graph-backend-decode full",
|
||||||
"--cuda-graph-backend-prefill disabled",
|
"--cuda-graph-backend-prefill disabled",
|
||||||
"--cuda-graph-max-bs-decode 32",
|
"--cuda-graph-max-bs-decode 32",
|
||||||
@@ -97,15 +162,8 @@ export const config = {
|
|||||||
"--speculative-num-draft-tokens 4",
|
"--speculative-num-draft-tokens 4",
|
||||||
"--speculative-draft-model-path {{MODEL_NAME}}",
|
"--speculative-draft-model-path {{MODEL_NAME}}",
|
||||||
"--speculative-draft-attention-backend fa3",
|
"--speculative-draft-attention-backend fa3",
|
||||||
"--moe-a2a-backend deepep",
|
|
||||||
"--moe-runner-backend deep_gemm",
|
|
||||||
"--deepep-dispatcher-output-dtype bf16",
|
|
||||||
"--deepep-mode auto",
|
|
||||||
"--enable-nccl-nvls",
|
|
||||||
"--enable-multimodal",
|
"--enable-multimodal",
|
||||||
"--enable-metrics",
|
"--enable-metrics",
|
||||||
"--tool-call-parser dots",
|
|
||||||
"--reasoning-parser qwen3",
|
|
||||||
"--watchdog-timeout 1800",
|
"--watchdog-timeout 1800",
|
||||||
"--host {{HOST_IP}}",
|
"--host {{HOST_IP}}",
|
||||||
"--port {{PORT}}",
|
"--port {{PORT}}",
|
||||||
@@ -125,20 +183,25 @@ export const config = {
|
|||||||
],
|
],
|
||||||
flags: [
|
flags: [
|
||||||
"--model-path {{MODEL_NAME}}",
|
"--model-path {{MODEL_NAME}}",
|
||||||
"--context-length 524288",
|
"--trust-remote-code",
|
||||||
"--enable-dp-attention",
|
"--enable-dp-attention",
|
||||||
"--dp-size 8",
|
"--tp 8",
|
||||||
"--tp-size 8",
|
"--dp 8",
|
||||||
"--ep-size 8",
|
"--ep 8",
|
||||||
|
"--moe-dense-tp-size 1",
|
||||||
|
"--moe-a2a-backend deepep",
|
||||||
|
"--moe-runner-backend auto",
|
||||||
|
"--deepep-dispatcher-output-dtype auto",
|
||||||
|
"--deepep-mode auto",
|
||||||
|
"--enable-nccl-nvls",
|
||||||
|
"--context-length 524288",
|
||||||
"--mem-fraction-static 0.87",
|
"--mem-fraction-static 0.87",
|
||||||
"--max-running-requests 256",
|
"--max-running-requests 256",
|
||||||
"--chunked-prefill-size 16384",
|
"--chunked-prefill-size 16384",
|
||||||
"--trust-remote-code",
|
|
||||||
"--swa-full-tokens-ratio 0.03",
|
"--swa-full-tokens-ratio 0.03",
|
||||||
"--prefill-attention-backend fa3",
|
"--prefill-attention-backend fa3",
|
||||||
"--decode-attention-backend fa3",
|
"--decode-attention-backend fa3",
|
||||||
"--page-size 64",
|
"--page-size 64",
|
||||||
"--moe-dense-tp-size 1",
|
|
||||||
"--cuda-graph-backend-decode full",
|
"--cuda-graph-backend-decode full",
|
||||||
"--cuda-graph-backend-prefill disabled",
|
"--cuda-graph-backend-prefill disabled",
|
||||||
"--cuda-graph-max-bs-decode 32",
|
"--cuda-graph-max-bs-decode 32",
|
||||||
@@ -148,15 +211,8 @@ export const config = {
|
|||||||
"--speculative-num-draft-tokens 4",
|
"--speculative-num-draft-tokens 4",
|
||||||
"--speculative-draft-model-path {{MODEL_NAME}}",
|
"--speculative-draft-model-path {{MODEL_NAME}}",
|
||||||
"--speculative-draft-attention-backend fa3",
|
"--speculative-draft-attention-backend fa3",
|
||||||
"--moe-a2a-backend deepep",
|
|
||||||
"--moe-runner-backend auto",
|
|
||||||
"--deepep-dispatcher-output-dtype auto",
|
|
||||||
"--deepep-mode auto",
|
|
||||||
"--enable-nccl-nvls",
|
|
||||||
"--enable-multimodal",
|
"--enable-multimodal",
|
||||||
"--enable-metrics",
|
"--enable-metrics",
|
||||||
"--reasoning-parser qwen3",
|
|
||||||
"--tool-call-parser dots",
|
|
||||||
"--watchdog-timeout 1800",
|
"--watchdog-timeout 1800",
|
||||||
"--host {{HOST_IP}}",
|
"--host {{HOST_IP}}",
|
||||||
"--port {{PORT}}",
|
"--port {{PORT}}",
|
||||||
@@ -176,20 +232,25 @@ export const config = {
|
|||||||
],
|
],
|
||||||
flags: [
|
flags: [
|
||||||
"--model-path {{MODEL_NAME}}",
|
"--model-path {{MODEL_NAME}}",
|
||||||
"--context-length 524288",
|
"--trust-remote-code",
|
||||||
"--enable-dp-attention",
|
"--enable-dp-attention",
|
||||||
"--dp-size 8",
|
"--tp 8",
|
||||||
"--tp-size 8",
|
"--dp 8",
|
||||||
"--ep-size 8",
|
"--ep 8",
|
||||||
|
"--moe-dense-tp-size 1",
|
||||||
|
"--moe-a2a-backend deepep",
|
||||||
|
"--moe-runner-backend deep_gemm",
|
||||||
|
"--deepep-dispatcher-output-dtype bf16",
|
||||||
|
"--deepep-mode auto",
|
||||||
|
"--enable-nccl-nvls",
|
||||||
|
"--context-length 524288",
|
||||||
"--mem-fraction-static 0.87",
|
"--mem-fraction-static 0.87",
|
||||||
"--max-running-requests 256",
|
"--max-running-requests 256",
|
||||||
"--chunked-prefill-size 16384",
|
"--chunked-prefill-size 16384",
|
||||||
"--trust-remote-code",
|
|
||||||
"--swa-full-tokens-ratio 0.03",
|
"--swa-full-tokens-ratio 0.03",
|
||||||
"--prefill-attention-backend fa3",
|
"--prefill-attention-backend fa3",
|
||||||
"--decode-attention-backend fa3",
|
"--decode-attention-backend fa3",
|
||||||
"--page-size 64",
|
"--page-size 64",
|
||||||
"--moe-dense-tp-size 1",
|
|
||||||
"--cuda-graph-backend-decode full",
|
"--cuda-graph-backend-decode full",
|
||||||
"--cuda-graph-backend-prefill disabled",
|
"--cuda-graph-backend-prefill disabled",
|
||||||
"--cuda-graph-max-bs-decode 32",
|
"--cuda-graph-max-bs-decode 32",
|
||||||
@@ -199,15 +260,8 @@ export const config = {
|
|||||||
"--speculative-num-draft-tokens 4",
|
"--speculative-num-draft-tokens 4",
|
||||||
"--speculative-draft-model-path {{MODEL_NAME}}",
|
"--speculative-draft-model-path {{MODEL_NAME}}",
|
||||||
"--speculative-draft-attention-backend fa3",
|
"--speculative-draft-attention-backend fa3",
|
||||||
"--moe-a2a-backend deepep",
|
|
||||||
"--moe-runner-backend deep_gemm",
|
|
||||||
"--deepep-dispatcher-output-dtype bf16",
|
|
||||||
"--deepep-mode auto",
|
|
||||||
"--enable-nccl-nvls",
|
|
||||||
"--enable-multimodal",
|
"--enable-multimodal",
|
||||||
"--enable-metrics",
|
"--enable-metrics",
|
||||||
"--tool-call-parser dots",
|
|
||||||
"--reasoning-parser qwen3",
|
|
||||||
"--watchdog-timeout 1800",
|
"--watchdog-timeout 1800",
|
||||||
"--host {{HOST_IP}}",
|
"--host {{HOST_IP}}",
|
||||||
"--port {{PORT}}",
|
"--port {{PORT}}",
|
||||||
@@ -227,20 +281,25 @@ export const config = {
|
|||||||
],
|
],
|
||||||
flags: [
|
flags: [
|
||||||
"--model-path {{MODEL_NAME}}",
|
"--model-path {{MODEL_NAME}}",
|
||||||
"--context-length 524288",
|
"--trust-remote-code",
|
||||||
"--enable-dp-attention",
|
"--enable-dp-attention",
|
||||||
"--dp-size 8",
|
"--tp 8",
|
||||||
"--tp-size 8",
|
"--dp 8",
|
||||||
"--ep-size 8",
|
"--ep 8",
|
||||||
|
"--moe-dense-tp-size 1",
|
||||||
|
"--moe-a2a-backend deepep",
|
||||||
|
"--moe-runner-backend auto",
|
||||||
|
"--deepep-dispatcher-output-dtype auto",
|
||||||
|
"--deepep-mode auto",
|
||||||
|
"--enable-nccl-nvls",
|
||||||
|
"--context-length 524288",
|
||||||
"--mem-fraction-static 0.87",
|
"--mem-fraction-static 0.87",
|
||||||
"--max-running-requests 256",
|
"--max-running-requests 256",
|
||||||
"--chunked-prefill-size 16384",
|
"--chunked-prefill-size 16384",
|
||||||
"--trust-remote-code",
|
|
||||||
"--swa-full-tokens-ratio 0.03",
|
"--swa-full-tokens-ratio 0.03",
|
||||||
"--prefill-attention-backend fa3",
|
"--prefill-attention-backend fa3",
|
||||||
"--decode-attention-backend fa3",
|
"--decode-attention-backend fa3",
|
||||||
"--page-size 64",
|
"--page-size 64",
|
||||||
"--moe-dense-tp-size 1",
|
|
||||||
"--cuda-graph-backend-decode full",
|
"--cuda-graph-backend-decode full",
|
||||||
"--cuda-graph-backend-prefill disabled",
|
"--cuda-graph-backend-prefill disabled",
|
||||||
"--cuda-graph-max-bs-decode 32",
|
"--cuda-graph-max-bs-decode 32",
|
||||||
@@ -250,15 +309,8 @@ export const config = {
|
|||||||
"--speculative-num-draft-tokens 4",
|
"--speculative-num-draft-tokens 4",
|
||||||
"--speculative-draft-model-path {{MODEL_NAME}}",
|
"--speculative-draft-model-path {{MODEL_NAME}}",
|
||||||
"--speculative-draft-attention-backend fa3",
|
"--speculative-draft-attention-backend fa3",
|
||||||
"--moe-a2a-backend deepep",
|
|
||||||
"--moe-runner-backend auto",
|
|
||||||
"--deepep-dispatcher-output-dtype auto",
|
|
||||||
"--deepep-mode auto",
|
|
||||||
"--enable-nccl-nvls",
|
|
||||||
"--enable-multimodal",
|
"--enable-multimodal",
|
||||||
"--enable-metrics",
|
"--enable-metrics",
|
||||||
"--reasoning-parser qwen3",
|
|
||||||
"--tool-call-parser dots",
|
|
||||||
"--watchdog-timeout 1800",
|
"--watchdog-timeout 1800",
|
||||||
"--host {{HOST_IP}}",
|
"--host {{HOST_IP}}",
|
||||||
"--port {{PORT}}",
|
"--port {{PORT}}",
|
||||||
|
|||||||
@@ -599,6 +599,7 @@ def _kimi_k3_moe_runner_overrides(server_args: Any, hf_config: Any) -> dict:
|
|||||||
"GlmMoeDsaForCausalLM",
|
"GlmMoeDsaForCausalLM",
|
||||||
"LongcatFlashForCausalLM",
|
"LongcatFlashForCausalLM",
|
||||||
"LongcatFlashForCausalLMNextN",
|
"LongcatFlashForCausalLMNextN",
|
||||||
|
"Dots3NoteForCausalLM",
|
||||||
)
|
)
|
||||||
def _deepseek_family_overrides(server_args: Any, hf_config: Any) -> dict:
|
def _deepseek_family_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||||
"""Order-safe declarations of the DeepSeek/DSA branch. The CP parallel
|
"""Order-safe declarations of the DeepSeek/DSA branch. The CP parallel
|
||||||
@@ -609,6 +610,7 @@ def _deepseek_family_overrides(server_args: Any, hf_config: Any) -> dict:
|
|||||||
from sglang.srt.configs.model_config import is_deepseek_dsa
|
from sglang.srt.configs.model_config import is_deepseek_dsa
|
||||||
|
|
||||||
overrides: Dict[str, Any] = {}
|
overrides: Dict[str, Any] = {}
|
||||||
|
|
||||||
if is_deepseek_dsa(hf_config): # DeepSeek 3.2/GLM 5
|
if is_deepseek_dsa(hf_config): # DeepSeek 3.2/GLM 5
|
||||||
# Set attention backend for DeepSeek
|
# Set attention backend for DeepSeek
|
||||||
if server_args.is_attention_backend_not_set():
|
if server_args.is_attention_backend_not_set():
|
||||||
@@ -1769,6 +1771,7 @@ _DEEPSEEK_FAMILY_ARCHS = frozenset(
|
|||||||
"GlmMoeDsaForCausalLM",
|
"GlmMoeDsaForCausalLM",
|
||||||
"LongcatFlashForCausalLM",
|
"LongcatFlashForCausalLM",
|
||||||
"LongcatFlashForCausalLMNextN",
|
"LongcatFlashForCausalLMNextN",
|
||||||
|
"Dots3NoteForCausalLM",
|
||||||
}
|
}
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -4,6 +4,7 @@ from sglang.srt.configs.chatglm import ChatGLMConfig
|
|||||||
from sglang.srt.configs.cohere2_moe import Cohere2MoeConfig
|
from sglang.srt.configs.cohere2_moe import Cohere2MoeConfig
|
||||||
from sglang.srt.configs.dbrx import DbrxConfig
|
from sglang.srt.configs.dbrx import DbrxConfig
|
||||||
from sglang.srt.configs.deepseekvl2 import DeepseekVL2Config
|
from sglang.srt.configs.deepseekvl2 import DeepseekVL2Config
|
||||||
|
from sglang.srt.configs.dots3 import Dots3Config
|
||||||
from sglang.srt.configs.dots_ocr import DotsOCRConfig
|
from sglang.srt.configs.dots_ocr import DotsOCRConfig
|
||||||
from sglang.srt.configs.dots_vlm import DotsVLMConfig
|
from sglang.srt.configs.dots_vlm import DotsVLMConfig
|
||||||
from sglang.srt.configs.exaone import ExaoneConfig
|
from sglang.srt.configs.exaone import ExaoneConfig
|
||||||
@@ -97,6 +98,7 @@ __all__ = [
|
|||||||
"InternS2MobiusVisionConfig",
|
"InternS2MobiusVisionConfig",
|
||||||
"DotsVLMConfig",
|
"DotsVLMConfig",
|
||||||
"DotsOCRConfig",
|
"DotsOCRConfig",
|
||||||
|
"Dots3Config",
|
||||||
"FalconH1Config",
|
"FalconH1Config",
|
||||||
"GraniteMoeHybridConfig",
|
"GraniteMoeHybridConfig",
|
||||||
"Lfm2Config",
|
"Lfm2Config",
|
||||||
|
|||||||
@@ -0,0 +1,243 @@
|
|||||||
|
from transformers import AutoTokenizer
|
||||||
|
from transformers.configuration_utils import PretrainedConfig
|
||||||
|
|
||||||
|
from sglang.srt.multimodal.customized_mm_processor_utils import (
|
||||||
|
register_customized_processor,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class DotsNoteOmniTokenizerProxy:
|
||||||
|
@classmethod
|
||||||
|
def from_pretrained(cls, model_path: str, *args, **kwargs):
|
||||||
|
kwargs.pop("use_fast", None)
|
||||||
|
return AutoTokenizer.from_pretrained(model_path, *args, **kwargs)
|
||||||
|
|
||||||
|
|
||||||
|
@register_customized_processor(DotsNoteOmniTokenizerProxy)
|
||||||
|
class Dots3Config(PretrainedConfig):
|
||||||
|
model_type = "dots3_note"
|
||||||
|
keys_to_ignore_at_inference = ["past_key_values"]
|
||||||
|
is_hybrid_swa = True
|
||||||
|
requires_draft_attention_wrapper = True
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
# General model parameters
|
||||||
|
vocab_size=152064,
|
||||||
|
hidden_size=2560,
|
||||||
|
hidden_act="silu",
|
||||||
|
intermediate_size=7168,
|
||||||
|
num_hidden_layers=30,
|
||||||
|
max_position_embeddings=8192,
|
||||||
|
initializer_range=0.02,
|
||||||
|
rms_norm_eps=1e-5,
|
||||||
|
use_cache=True,
|
||||||
|
pretraining_tp=1,
|
||||||
|
# Token IDs
|
||||||
|
pad_token_id=None,
|
||||||
|
bos_token_id=151643,
|
||||||
|
eos_token_id=151645,
|
||||||
|
tie_word_embeddings=False,
|
||||||
|
# Attention parameters
|
||||||
|
attention_bias=False,
|
||||||
|
attention_dropout=0.0,
|
||||||
|
apply_mla_qkv_lora_rescale=True,
|
||||||
|
# MLA (Multi-head Latent Attention) parameters
|
||||||
|
attention_gate_type="headwise",
|
||||||
|
kv_lora_rank=512,
|
||||||
|
q_lora_rank=512,
|
||||||
|
qk_nope_head_dim=128,
|
||||||
|
qk_rope_head_dim=64,
|
||||||
|
num_attention_heads=64,
|
||||||
|
num_key_value_heads=64,
|
||||||
|
v_head_dim=128,
|
||||||
|
# Dots3 uses one shared MTP layer for NEXTN decoding.
|
||||||
|
num_nextn_predict_layers=1,
|
||||||
|
# Sliding Window Attention (SWA) parameters
|
||||||
|
layer_types=None,
|
||||||
|
sliding_window_size=512,
|
||||||
|
swa_attention_gate_type="headwise",
|
||||||
|
swa_q_lora_rank=512,
|
||||||
|
swa_kv_lora_rank=512,
|
||||||
|
swa_qk_nope_head_dim=128,
|
||||||
|
swa_qk_rope_head_dim=64,
|
||||||
|
swa_rope_theta=None,
|
||||||
|
swa_num_attention_heads=32,
|
||||||
|
swa_num_key_value_heads=32,
|
||||||
|
swa_v_head_dim=128,
|
||||||
|
# MoE (Mixture of Experts) parameters
|
||||||
|
moe_intermediate_size=1024,
|
||||||
|
n_shared_experts=1,
|
||||||
|
n_routed_experts=128,
|
||||||
|
num_experts_per_tok=6,
|
||||||
|
moe_layer_freq=1,
|
||||||
|
first_k_dense_replace=1,
|
||||||
|
routed_scaling_factor=1.0,
|
||||||
|
norm_topk_prob=True,
|
||||||
|
scoring_func="sigmoid",
|
||||||
|
n_group=1,
|
||||||
|
topk_method="noaux_tc",
|
||||||
|
topk_group=1,
|
||||||
|
# RoPE parameters
|
||||||
|
rope_theta=50000.0,
|
||||||
|
rope_scaling=None,
|
||||||
|
# Optional DSA indexer parameters.
|
||||||
|
index_n_heads=None,
|
||||||
|
index_head_dim=None,
|
||||||
|
index_topk=None,
|
||||||
|
language_only=False,
|
||||||
|
# Multimodal special tokens
|
||||||
|
im_start_token="<|img|>",
|
||||||
|
im_token="<|imgpad|>",
|
||||||
|
im_end_token="<|endofimg|>",
|
||||||
|
audio_start_token="<|audio_comp_start|>",
|
||||||
|
audio_token="<|audio_comp_pad|>",
|
||||||
|
audio_end_token="<|audio_comp_end|>",
|
||||||
|
video_token="<|video_pad|>",
|
||||||
|
**kwargs,
|
||||||
|
):
|
||||||
|
# General model parameters
|
||||||
|
self.vocab_size = vocab_size
|
||||||
|
self.hidden_size = hidden_size
|
||||||
|
self.hidden_act = hidden_act
|
||||||
|
self.intermediate_size = intermediate_size
|
||||||
|
self.num_hidden_layers = num_hidden_layers
|
||||||
|
self.max_position_embeddings = max_position_embeddings
|
||||||
|
self.initializer_range = initializer_range
|
||||||
|
self.rms_norm_eps = rms_norm_eps
|
||||||
|
self.use_cache = use_cache
|
||||||
|
self.pretraining_tp = pretraining_tp
|
||||||
|
|
||||||
|
# Attention parameters
|
||||||
|
self.attention_bias = attention_bias
|
||||||
|
self.attention_dropout = attention_dropout
|
||||||
|
self.apply_mla_qkv_lora_rescale = apply_mla_qkv_lora_rescale
|
||||||
|
|
||||||
|
# MLA (Multi-head Latent Attention) parameters
|
||||||
|
self.attention_gate_type = attention_gate_type
|
||||||
|
self.kv_lora_rank = kv_lora_rank
|
||||||
|
self.q_lora_rank = q_lora_rank
|
||||||
|
self.qk_nope_head_dim = qk_nope_head_dim
|
||||||
|
self.qk_rope_head_dim = qk_rope_head_dim
|
||||||
|
self.num_attention_heads = num_attention_heads
|
||||||
|
self.num_key_value_heads = num_key_value_heads
|
||||||
|
self.v_head_dim = v_head_dim
|
||||||
|
|
||||||
|
# MTP / NextN
|
||||||
|
self.num_nextn_predict_layers = num_nextn_predict_layers
|
||||||
|
|
||||||
|
# Sliding Window Attention (SWA) parameters
|
||||||
|
self.layer_types = layer_types
|
||||||
|
self.sliding_window_size = sliding_window_size
|
||||||
|
self.swa_attention_gate_type = swa_attention_gate_type
|
||||||
|
self.swa_q_lora_rank = swa_q_lora_rank
|
||||||
|
self.swa_kv_lora_rank = swa_kv_lora_rank
|
||||||
|
self.swa_qk_nope_head_dim = swa_qk_nope_head_dim
|
||||||
|
self.swa_qk_rope_head_dim = swa_qk_rope_head_dim
|
||||||
|
self.swa_rope_theta = rope_theta if swa_rope_theta is None else swa_rope_theta
|
||||||
|
self.swa_num_attention_heads = swa_num_attention_heads
|
||||||
|
self.swa_num_key_value_heads = swa_num_key_value_heads
|
||||||
|
self.swa_v_head_dim = swa_v_head_dim
|
||||||
|
# Runtime cache geometry for the SWA attention path.
|
||||||
|
self.swa_head_dim = swa_qk_nope_head_dim + swa_qk_rope_head_dim
|
||||||
|
|
||||||
|
# MoE (Mixture of Experts) parameters
|
||||||
|
self.moe_intermediate_size = moe_intermediate_size
|
||||||
|
self.n_shared_experts = n_shared_experts
|
||||||
|
self.n_routed_experts = n_routed_experts
|
||||||
|
self.num_experts_per_tok = num_experts_per_tok
|
||||||
|
self.moe_layer_freq = moe_layer_freq
|
||||||
|
self.first_k_dense_replace = first_k_dense_replace
|
||||||
|
self.routed_scaling_factor = routed_scaling_factor
|
||||||
|
self.norm_topk_prob = norm_topk_prob
|
||||||
|
self.scoring_func = scoring_func
|
||||||
|
self.n_group = n_group
|
||||||
|
self.topk_method = topk_method
|
||||||
|
self.topk_group = topk_group
|
||||||
|
|
||||||
|
# RoPE parameters
|
||||||
|
self.rope_theta = rope_theta
|
||||||
|
self.rope_scaling = rope_scaling
|
||||||
|
self._rope_scaling_validation()
|
||||||
|
|
||||||
|
# NSA (Native Sparse Attention) parameters
|
||||||
|
self.index_n_heads = index_n_heads
|
||||||
|
self.index_head_dim = index_head_dim
|
||||||
|
self.index_topk = index_topk
|
||||||
|
self.language_only = language_only
|
||||||
|
|
||||||
|
self.im_start_token = im_start_token
|
||||||
|
self.im_token = im_token
|
||||||
|
self.im_end_token = im_end_token
|
||||||
|
self.audio_start_token = audio_start_token
|
||||||
|
self.audio_token = audio_token
|
||||||
|
self.audio_end_token = audio_end_token
|
||||||
|
# The chat template renders a video content part as this single token,
|
||||||
|
# which the processor replaces with the flattened frames and audio.
|
||||||
|
self.video_token = video_token
|
||||||
|
|
||||||
|
super().__init__(
|
||||||
|
pad_token_id=pad_token_id,
|
||||||
|
bos_token_id=bos_token_id,
|
||||||
|
eos_token_id=eos_token_id,
|
||||||
|
tie_word_embeddings=tie_word_embeddings,
|
||||||
|
**kwargs,
|
||||||
|
)
|
||||||
|
|
||||||
|
def configure_draft_model(self) -> str:
|
||||||
|
"""Configure the recursively shared MTP layer with SWA geometry."""
|
||||||
|
self.num_nextn_predict_layers = 1
|
||||||
|
self.layer_types = ["sliding_attention"]
|
||||||
|
self.attention_gate_type = self.swa_attention_gate_type
|
||||||
|
self.kv_lora_rank = self.swa_kv_lora_rank
|
||||||
|
self.q_lora_rank = self.swa_q_lora_rank
|
||||||
|
self.qk_nope_head_dim = self.swa_qk_nope_head_dim
|
||||||
|
self.qk_rope_head_dim = self.swa_qk_rope_head_dim
|
||||||
|
self.num_attention_heads = self.swa_num_attention_heads
|
||||||
|
self.num_key_value_heads = self.swa_num_key_value_heads
|
||||||
|
self.v_head_dim = self.swa_v_head_dim
|
||||||
|
return "Dots3NoteForCausalLMNextN"
|
||||||
|
|
||||||
|
def wrap_attention_backend(self, runner, full_attn_backend):
|
||||||
|
from sglang.srt.layers.attention.dots_hybrid_backend import (
|
||||||
|
wrap_dots_attention_backend,
|
||||||
|
)
|
||||||
|
|
||||||
|
return wrap_dots_attention_backend(runner, full_attn_backend)
|
||||||
|
|
||||||
|
def wrap_draft_decode_attention_backend(self, backend):
|
||||||
|
from sglang.srt.layers.attention.dots_hybrid_backend import (
|
||||||
|
wrap_dots_draft_decode_backend,
|
||||||
|
)
|
||||||
|
|
||||||
|
return wrap_dots_draft_decode_backend(backend)
|
||||||
|
|
||||||
|
def _rope_scaling_validation(self):
|
||||||
|
"""
|
||||||
|
Validate the `rope_scaling` configuration.
|
||||||
|
"""
|
||||||
|
if self.rope_scaling is None:
|
||||||
|
return
|
||||||
|
|
||||||
|
if not isinstance(self.rope_scaling, dict):
|
||||||
|
raise ValueError(
|
||||||
|
f"`rope_scaling` must be a dictionary, got {self.rope_scaling}"
|
||||||
|
)
|
||||||
|
rope_scaling_type = self.rope_scaling.get("type", None)
|
||||||
|
rope_scaling_factor = self.rope_scaling.get("factor", None)
|
||||||
|
if rope_scaling_type is None or rope_scaling_type not in [
|
||||||
|
"linear",
|
||||||
|
"dynamic",
|
||||||
|
"yarn",
|
||||||
|
]:
|
||||||
|
raise ValueError(
|
||||||
|
f"`rope_scaling`'s type field must be one of ['linear', 'dynamic', 'yarn'], got {rope_scaling_type}"
|
||||||
|
)
|
||||||
|
if (
|
||||||
|
rope_scaling_factor is None
|
||||||
|
or not isinstance(rope_scaling_factor, (int, float))
|
||||||
|
or rope_scaling_factor <= 1.0
|
||||||
|
):
|
||||||
|
raise ValueError(
|
||||||
|
f"`rope_scaling`'s factor field must be a number > 1, got {rope_scaling_factor}"
|
||||||
|
)
|
||||||
@@ -124,6 +124,8 @@ def is_deepseek_dsa(config) -> bool:
|
|||||||
"GlmMoeDsaForCausalLMNextN",
|
"GlmMoeDsaForCausalLMNextN",
|
||||||
"LongcatFlashForCausalLM",
|
"LongcatFlashForCausalLM",
|
||||||
"LongcatFlashForCausalLMNextN",
|
"LongcatFlashForCausalLMNextN",
|
||||||
|
"Dots3NoteForCausalLM",
|
||||||
|
"Dots3NoteForCausalLMNextN",
|
||||||
)
|
)
|
||||||
and _hf_attr(config, "index_topk") is not None
|
and _hf_attr(config, "index_topk") is not None
|
||||||
)
|
)
|
||||||
@@ -630,6 +632,13 @@ class ModelConfig:
|
|||||||
def _config_draft_model(self):
|
def _config_draft_model(self):
|
||||||
is_draft_model = self.is_draft_model
|
is_draft_model = self.is_draft_model
|
||||||
|
|
||||||
|
from sglang.srt.configs.dots3 import Dots3Config
|
||||||
|
|
||||||
|
if is_draft_model and isinstance(self.hf_text_config, Dots3Config):
|
||||||
|
self.hf_config.architectures[0] = (
|
||||||
|
self.hf_text_config.configure_draft_model()
|
||||||
|
)
|
||||||
|
|
||||||
if is_draft_model and self.hf_config.architectures[0] in [
|
if is_draft_model and self.hf_config.architectures[0] in [
|
||||||
"DeepseekV3ForCausalLM",
|
"DeepseekV3ForCausalLM",
|
||||||
"DeepseekV32ForCausalLM",
|
"DeepseekV32ForCausalLM",
|
||||||
@@ -866,6 +875,8 @@ class ModelConfig:
|
|||||||
self.hf_config.context_len = self.context_len
|
self.hf_config.context_len = self.context_len
|
||||||
|
|
||||||
def _derive_model_shapes(self):
|
def _derive_model_shapes(self):
|
||||||
|
from sglang.srt.configs.dots3 import Dots3Config
|
||||||
|
|
||||||
# Unify the config keys for hf_text_config
|
# Unify the config keys for hf_text_config
|
||||||
self.head_dim = getattr(self.hf_text_config, "head_dim", None)
|
self.head_dim = getattr(self.hf_text_config, "head_dim", None)
|
||||||
if self.head_dim is None:
|
if self.head_dim is None:
|
||||||
@@ -902,6 +913,8 @@ class ModelConfig:
|
|||||||
or "LongcatFlashForCausalLM" in self.hf_config.architectures
|
or "LongcatFlashForCausalLM" in self.hf_config.architectures
|
||||||
or "LongcatFlashForCausalLMNextN" in self.hf_config.architectures
|
or "LongcatFlashForCausalLMNextN" in self.hf_config.architectures
|
||||||
or "DotsVLMForCausalLM" in self.hf_config.architectures
|
or "DotsVLMForCausalLM" in self.hf_config.architectures
|
||||||
|
or "Dots3NoteForCausalLM" in self.hf_config.architectures
|
||||||
|
or "Dots3NoteForCausalLMNextN" in self.hf_config.architectures
|
||||||
or "MistralLarge3ForCausalLM" in self.hf_config.architectures
|
or "MistralLarge3ForCausalLM" in self.hf_config.architectures
|
||||||
or (
|
or (
|
||||||
"PixtralForConditionalGeneration" in self.hf_config.architectures
|
"PixtralForConditionalGeneration" in self.hf_config.architectures
|
||||||
@@ -917,6 +930,12 @@ class ModelConfig:
|
|||||||
self.qk_nope_head_dim = self.hf_text_config.qk_nope_head_dim
|
self.qk_nope_head_dim = self.hf_text_config.qk_nope_head_dim
|
||||||
self.qk_rope_head_dim = self.hf_text_config.qk_rope_head_dim
|
self.qk_rope_head_dim = self.hf_text_config.qk_rope_head_dim
|
||||||
self.v_head_dim = self.hf_text_config.v_head_dim
|
self.v_head_dim = self.hf_text_config.v_head_dim
|
||||||
|
if isinstance(self.hf_text_config, Dots3Config):
|
||||||
|
self.swa_kv_lora_rank = self.hf_text_config.swa_kv_lora_rank
|
||||||
|
self.swa_qk_rope_head_dim = self.hf_text_config.swa_qk_rope_head_dim
|
||||||
|
else:
|
||||||
|
self.swa_kv_lora_rank = self.kv_lora_rank
|
||||||
|
self.swa_qk_rope_head_dim = self.qk_rope_head_dim
|
||||||
self.index_head_dim = (
|
self.index_head_dim = (
|
||||||
get_dsa_index_head_dim(self.hf_text_config)
|
get_dsa_index_head_dim(self.hf_text_config)
|
||||||
if is_deepseek_dsa(self.hf_text_config)
|
if is_deepseek_dsa(self.hf_text_config)
|
||||||
@@ -1893,6 +1912,7 @@ multimodal_model_archs = [
|
|||||||
"Step3VLForConditionalGeneration",
|
"Step3VLForConditionalGeneration",
|
||||||
"POINTSV15ChatModel",
|
"POINTSV15ChatModel",
|
||||||
"DotsVLMForCausalLM",
|
"DotsVLMForCausalLM",
|
||||||
|
"Dots3NoteForCausalLM",
|
||||||
"DotsOCRForCausalLM",
|
"DotsOCRForCausalLM",
|
||||||
"Sarashina2VisionForCausalLM",
|
"Sarashina2VisionForCausalLM",
|
||||||
"NVILAForConditionalGeneration",
|
"NVILAForConditionalGeneration",
|
||||||
@@ -2118,7 +2138,10 @@ def get_hybrid_layer_ids(
|
|||||||
full_attention_layer_ids = [
|
full_attention_layer_ids = [
|
||||||
i for i in range(num_hidden_layers) if (i + 1) % 4 == 0
|
i for i in range(num_hidden_layers) if (i + 1) % 4 == 0
|
||||||
]
|
]
|
||||||
elif any(arch in SWA_SINK_ARCHS for arch in model_architectures):
|
elif any(arch in SWA_SINK_ARCHS for arch in model_architectures) or any(
|
||||||
|
arch in ("Dots3NoteForCausalLM", "Dots3NoteForCausalLMNextN")
|
||||||
|
for arch in model_architectures
|
||||||
|
):
|
||||||
layer_types = getattr(hf_text_config, "layer_types", [])
|
layer_types = getattr(hf_text_config, "layer_types", [])
|
||||||
swa_attention_layer_ids = [
|
swa_attention_layer_ids = [
|
||||||
i for i, x in enumerate(layer_types) if x == "sliding_attention"
|
i for i, x in enumerate(layer_types) if x == "sliding_attention"
|
||||||
|
|||||||
@@ -915,6 +915,7 @@ class ChatCompletionRequest(BaseModel):
|
|||||||
use_audio_in_video: bool = False
|
use_audio_in_video: bool = False
|
||||||
|
|
||||||
images_config: Optional[Dict] = None
|
images_config: Optional[Dict] = None
|
||||||
|
video_config: Optional[Dict] = None
|
||||||
|
|
||||||
# Custom logit processor for advanced sampling control
|
# Custom logit processor for advanced sampling control
|
||||||
custom_logit_processor: Optional[Union[List[Optional[str]], str]] = None
|
custom_logit_processor: Optional[Union[List[Optional[str]], str]] = None
|
||||||
|
|||||||
@@ -38,7 +38,10 @@ from jsonschema import Draft202012Validator, SchemaError
|
|||||||
|
|
||||||
from sglang.srt.entrypoints.openai import chat_encoding, encoding_dsv4, encoding_dsv32
|
from sglang.srt.entrypoints.openai import chat_encoding, encoding_dsv4, encoding_dsv32
|
||||||
from sglang.srt.entrypoints.openai.protocol import (
|
from sglang.srt.entrypoints.openai.protocol import (
|
||||||
|
ChatCompletionMessageContentTextPart,
|
||||||
|
ChatCompletionMessageContentVideoPart,
|
||||||
ChatCompletionMessageGenericParam,
|
ChatCompletionMessageGenericParam,
|
||||||
|
ChatCompletionMessageUserParam,
|
||||||
ChatCompletionRequest,
|
ChatCompletionRequest,
|
||||||
ChatCompletionResponse,
|
ChatCompletionResponse,
|
||||||
ChatCompletionResponseChoice,
|
ChatCompletionResponseChoice,
|
||||||
@@ -205,6 +208,38 @@ def neutralize_kimi_k3_image_placeholder_value(value: Any) -> Any:
|
|||||||
return value
|
return value
|
||||||
|
|
||||||
|
|
||||||
|
def _extract_video_question(request: ChatCompletionRequest) -> Optional[str]:
|
||||||
|
"""Return text paired with a video in the last user turn."""
|
||||||
|
for message in reversed(request.messages or []):
|
||||||
|
if not isinstance(message, ChatCompletionMessageUserParam):
|
||||||
|
continue
|
||||||
|
content = message.content
|
||||||
|
if not isinstance(content, list):
|
||||||
|
continue
|
||||||
|
has_video = any(
|
||||||
|
isinstance(part, ChatCompletionMessageContentVideoPart) for part in content
|
||||||
|
)
|
||||||
|
if not has_video:
|
||||||
|
continue
|
||||||
|
return "".join(
|
||||||
|
part.text
|
||||||
|
for part in content
|
||||||
|
if isinstance(part, ChatCompletionMessageContentTextPart)
|
||||||
|
)
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
def _build_video_config(request: ChatCompletionRequest) -> Optional[Dict[str, Any]]:
|
||||||
|
"""Build request-scoped video processor config without model-specific fields."""
|
||||||
|
config = dict(request.video_config or {})
|
||||||
|
question = _extract_video_question(request)
|
||||||
|
if question is not None:
|
||||||
|
# Internal metadata derived from the message must not be overridden by
|
||||||
|
# a model-specific public processor option.
|
||||||
|
config["_question"] = question
|
||||||
|
return config or None
|
||||||
|
|
||||||
|
|
||||||
class OpenAIServingChat(OpenAIServingBase):
|
class OpenAIServingChat(OpenAIServingBase):
|
||||||
"""Handler for /v1/chat/completions requests"""
|
"""Handler for /v1/chat/completions requests"""
|
||||||
|
|
||||||
@@ -1045,6 +1080,7 @@ class OpenAIServingChat(OpenAIServingBase):
|
|||||||
custom_labels=custom_labels,
|
custom_labels=custom_labels,
|
||||||
custom_logit_processor=request.custom_logit_processor,
|
custom_logit_processor=request.custom_logit_processor,
|
||||||
images_config=getattr(request, "images_config", None),
|
images_config=getattr(request, "images_config", None),
|
||||||
|
video_config=_build_video_config(request),
|
||||||
image_max_dynamic_patch=img_max_dynamic_patch,
|
image_max_dynamic_patch=img_max_dynamic_patch,
|
||||||
video_max_dynamic_patch=vid_max_dynamic_patch,
|
video_max_dynamic_patch=vid_max_dynamic_patch,
|
||||||
max_dynamic_patch=getattr(request, "max_dynamic_patch", None),
|
max_dynamic_patch=getattr(request, "max_dynamic_patch", None),
|
||||||
|
|||||||
@@ -0,0 +1,353 @@
|
|||||||
|
import json
|
||||||
|
import logging
|
||||||
|
import re
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
try:
|
||||||
|
import json_repair
|
||||||
|
except ImportError:
|
||||||
|
json_repair = None
|
||||||
|
|
||||||
|
from sglang.srt.entrypoints.openai.protocol import Tool
|
||||||
|
from sglang.srt.function_call.base_format_detector import BaseFormatDetector
|
||||||
|
from sglang.srt.function_call.core_types import (
|
||||||
|
StreamingParseResult,
|
||||||
|
StructureInfo,
|
||||||
|
ToolCallItem,
|
||||||
|
_GetInfoFunc,
|
||||||
|
)
|
||||||
|
from sglang.srt.function_call.utils import _is_complete_json
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
class DotsToolDetector(BaseFormatDetector):
|
||||||
|
"""Detector for the dots function-call format.
|
||||||
|
|
||||||
|
The canonical format contains one or more XML ``invoke`` elements inside a
|
||||||
|
``dots_function_call`` block::
|
||||||
|
|
||||||
|
<dots_function_call>
|
||||||
|
<invoke name="search">
|
||||||
|
<parameter name="query">weather in Shanghai</parameter>
|
||||||
|
</invoke>
|
||||||
|
</dots_function_call>
|
||||||
|
|
||||||
|
A JSON object with ``name`` and ``arguments`` is accepted as a fallback.
|
||||||
|
Multiple wrapper blocks and multiple invokes in one block are supported.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self):
|
||||||
|
super().__init__()
|
||||||
|
self.bot_token = "<dots_function_call>"
|
||||||
|
self.eot_token = "</dots_function_call>"
|
||||||
|
self.func_call_regex = re.compile(
|
||||||
|
rf"{re.escape(self.bot_token)}\s*(.*?)\s*{re.escape(self.eot_token)}",
|
||||||
|
re.DOTALL,
|
||||||
|
)
|
||||||
|
self.invoke_regex = re.compile(
|
||||||
|
r"<invoke\s+name\s*=\s*(?P<name>[^>]+)>(?P<body>.*?)</invoke>",
|
||||||
|
re.DOTALL,
|
||||||
|
)
|
||||||
|
self.parameter_regex = re.compile(
|
||||||
|
r"<parameter\s+name\s*=\s*(?P<name>[^>]+)>(?P<value>.*?)</parameter>",
|
||||||
|
re.DOTALL,
|
||||||
|
)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _extract_name(value: str) -> str:
|
||||||
|
value = value.strip()
|
||||||
|
if len(value) >= 2 and value[0] == value[-1] and value[0] in {'"', "'"}:
|
||||||
|
return value[1:-1]
|
||||||
|
return value
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _load_json(value: str) -> Any:
|
||||||
|
try:
|
||||||
|
return json.loads(value)
|
||||||
|
except (json.JSONDecodeError, ValueError):
|
||||||
|
if json_repair is None:
|
||||||
|
raise
|
||||||
|
return json_repair.loads(value)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def _convert_param_value(cls, value: str, param_type: Any) -> Any:
|
||||||
|
if value.lower() == "null":
|
||||||
|
return None
|
||||||
|
|
||||||
|
if isinstance(param_type, list):
|
||||||
|
param_type = next((item for item in param_type if item != "null"), "string")
|
||||||
|
if not isinstance(param_type, str):
|
||||||
|
param_type = str(param_type)
|
||||||
|
param_type = param_type.lower()
|
||||||
|
|
||||||
|
if param_type in {"string", "str", "text"}:
|
||||||
|
return value
|
||||||
|
if param_type in {"integer", "int"}:
|
||||||
|
try:
|
||||||
|
return int(value)
|
||||||
|
except (TypeError, ValueError):
|
||||||
|
return value
|
||||||
|
if param_type in {"number", "float"}:
|
||||||
|
try:
|
||||||
|
number = float(value)
|
||||||
|
return int(number) if number.is_integer() else number
|
||||||
|
except (TypeError, ValueError):
|
||||||
|
return value
|
||||||
|
if param_type in {"boolean", "bool"}:
|
||||||
|
return value.lower() in {"true", "1"}
|
||||||
|
|
||||||
|
try:
|
||||||
|
return cls._load_json(value)
|
||||||
|
except (json.JSONDecodeError, ValueError, TypeError):
|
||||||
|
return value
|
||||||
|
|
||||||
|
def _resolve_param_type(
|
||||||
|
self, schema: Any, defs: dict[str, Any], depth: int = 0
|
||||||
|
) -> Any | None:
|
||||||
|
"""Resolve a parameter type through local refs and schema compositions."""
|
||||||
|
if not isinstance(schema, dict) or depth > 10:
|
||||||
|
return None
|
||||||
|
if "type" in schema:
|
||||||
|
return schema["type"]
|
||||||
|
|
||||||
|
ref = schema.get("$ref")
|
||||||
|
if isinstance(ref, str) and ref.startswith("#/$defs/"):
|
||||||
|
return self._resolve_param_type(
|
||||||
|
defs.get(ref.rsplit("/", 1)[-1]), defs, depth + 1
|
||||||
|
)
|
||||||
|
|
||||||
|
for keyword in ("anyOf", "oneOf", "allOf"):
|
||||||
|
alternatives = schema.get(keyword)
|
||||||
|
if not isinstance(alternatives, list):
|
||||||
|
continue
|
||||||
|
for alternative in alternatives:
|
||||||
|
if isinstance(alternative, dict) and alternative.get("type") == "null":
|
||||||
|
continue
|
||||||
|
resolved = self._resolve_param_type(alternative, defs, depth + 1)
|
||||||
|
if resolved is not None:
|
||||||
|
return resolved
|
||||||
|
return None
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _tool_schema(name: str, tools: list[Tool]) -> tuple[dict, dict]:
|
||||||
|
for tool in tools:
|
||||||
|
if tool.function.name != name:
|
||||||
|
continue
|
||||||
|
schema = tool.function.parameters
|
||||||
|
if not isinstance(schema, dict):
|
||||||
|
break
|
||||||
|
properties = schema.get("properties", {})
|
||||||
|
defs = schema.get("$defs", {})
|
||||||
|
return (
|
||||||
|
properties if isinstance(properties, dict) else {},
|
||||||
|
defs if isinstance(defs, dict) else {},
|
||||||
|
)
|
||||||
|
return {}, {}
|
||||||
|
|
||||||
|
def _parse_xml_invoke(self, match: re.Match, tools: list[Tool]) -> dict[str, Any]:
|
||||||
|
name = self._extract_name(match.group("name"))
|
||||||
|
properties, defs = self._tool_schema(name, tools)
|
||||||
|
arguments: dict[str, Any] = {}
|
||||||
|
|
||||||
|
for parameter in self.parameter_regex.finditer(match.group("body")):
|
||||||
|
param_name = self._extract_name(parameter.group("name"))
|
||||||
|
value = parameter.group("value").strip()
|
||||||
|
param_type: Any = "string"
|
||||||
|
if param_name in properties:
|
||||||
|
param_type = (
|
||||||
|
self._resolve_param_type(properties[param_name], defs) or "string"
|
||||||
|
)
|
||||||
|
arguments[param_name] = self._convert_param_value(value, param_type)
|
||||||
|
|
||||||
|
return {"name": name, "arguments": arguments}
|
||||||
|
|
||||||
|
def _parse_block(self, content: str, tools: list[Tool]) -> list[dict[str, Any]]:
|
||||||
|
content = content.strip()
|
||||||
|
if content.startswith("<invoke"):
|
||||||
|
return [
|
||||||
|
self._parse_xml_invoke(match, tools)
|
||||||
|
for match in self.invoke_regex.finditer(content)
|
||||||
|
]
|
||||||
|
|
||||||
|
parsed = self._load_json(content)
|
||||||
|
if not isinstance(parsed, dict):
|
||||||
|
raise TypeError("dots JSON tool call must be an object")
|
||||||
|
return [parsed]
|
||||||
|
|
||||||
|
def has_tool_call(self, text: str) -> bool:
|
||||||
|
return self.bot_token in text
|
||||||
|
|
||||||
|
def detect_and_parse(self, text: str, tools: list[Tool]) -> StreamingParseResult:
|
||||||
|
marker_index = text.find(self.bot_token)
|
||||||
|
if marker_index == -1:
|
||||||
|
return StreamingParseResult(normal_text=text)
|
||||||
|
|
||||||
|
calls: list[ToolCallItem] = []
|
||||||
|
for block in self.func_call_regex.finditer(text):
|
||||||
|
try:
|
||||||
|
for parsed in self._parse_block(block.group(1), tools):
|
||||||
|
calls.extend(self.parse_base_json(parsed, tools))
|
||||||
|
except (json.JSONDecodeError, ValueError, TypeError) as exc:
|
||||||
|
logger.warning("Failed to parse dots tool call: %s", exc)
|
||||||
|
|
||||||
|
return StreamingParseResult(
|
||||||
|
normal_text=text[:marker_index].strip(), calls=calls
|
||||||
|
)
|
||||||
|
|
||||||
|
def _append_stream_call(
|
||||||
|
self, parsed: dict[str, Any], item: ToolCallItem
|
||||||
|
) -> ToolCallItem:
|
||||||
|
self.current_tool_id += 1
|
||||||
|
arguments = parsed.get("arguments", parsed.get("parameters", {})) or {}
|
||||||
|
serialized = json.dumps(arguments, ensure_ascii=False)
|
||||||
|
self.prev_tool_call_arr.append(
|
||||||
|
{"name": parsed.get("name"), "arguments": arguments}
|
||||||
|
)
|
||||||
|
self.streamed_args_for_tool.append(serialized)
|
||||||
|
item.tool_index = self.current_tool_id
|
||||||
|
item.parameters = serialized
|
||||||
|
return item
|
||||||
|
|
||||||
|
def parse_streaming_increment(
|
||||||
|
self, new_text: str, tools: list[Tool]
|
||||||
|
) -> StreamingParseResult:
|
||||||
|
"""Buffer incomplete XML and emit every complete call in the new data."""
|
||||||
|
self._buffer += new_text
|
||||||
|
normal_parts: list[str] = []
|
||||||
|
calls: list[ToolCallItem] = []
|
||||||
|
|
||||||
|
while self._buffer:
|
||||||
|
marker_index = self._buffer.find(self.bot_token)
|
||||||
|
if marker_index == -1:
|
||||||
|
partial_len = self._ends_with_partial_token(
|
||||||
|
self._buffer, self.bot_token
|
||||||
|
)
|
||||||
|
if partial_len:
|
||||||
|
normal_parts.append(self._buffer[:-partial_len])
|
||||||
|
self._buffer = self._buffer[-partial_len:]
|
||||||
|
else:
|
||||||
|
normal_parts.append(self._buffer)
|
||||||
|
self._buffer = ""
|
||||||
|
normal_parts = [
|
||||||
|
part.replace(self.eot_token, "") for part in normal_parts
|
||||||
|
]
|
||||||
|
break
|
||||||
|
|
||||||
|
if marker_index > 0:
|
||||||
|
normal_parts.append(self._buffer[:marker_index])
|
||||||
|
self._buffer = self._buffer[marker_index:]
|
||||||
|
|
||||||
|
end_index = self._buffer.find(self.eot_token, len(self.bot_token))
|
||||||
|
if end_index == -1:
|
||||||
|
self._stream_complete_json_body(tools, calls)
|
||||||
|
break
|
||||||
|
|
||||||
|
content = self._buffer[len(self.bot_token) : end_index]
|
||||||
|
self._buffer = self._buffer[end_index + len(self.eot_token) :]
|
||||||
|
try:
|
||||||
|
parsed_calls = self._parse_block(content, tools)
|
||||||
|
if not parsed_calls:
|
||||||
|
raise ValueError("dots tool-call block contains no invoke")
|
||||||
|
block_calls: list[ToolCallItem] = []
|
||||||
|
for index, parsed in enumerate(parsed_calls):
|
||||||
|
validated = self.parse_base_json(parsed, tools)
|
||||||
|
if index == 0 and self.current_tool_name_sent and validated:
|
||||||
|
item = validated[0]
|
||||||
|
arguments = item.parameters or ""
|
||||||
|
streamed = self.streamed_args_for_tool[self.current_tool_id]
|
||||||
|
remaining = arguments.removeprefix(streamed)
|
||||||
|
if remaining:
|
||||||
|
block_calls.append(
|
||||||
|
ToolCallItem(
|
||||||
|
tool_index=self.current_tool_id,
|
||||||
|
name=None,
|
||||||
|
parameters=remaining,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
self.prev_tool_call_arr[self.current_tool_id] = parsed
|
||||||
|
self.streamed_args_for_tool[self.current_tool_id] = arguments
|
||||||
|
else:
|
||||||
|
block_calls.extend(
|
||||||
|
self._append_stream_call(parsed, item) for item in validated
|
||||||
|
)
|
||||||
|
if block_calls:
|
||||||
|
calls.extend(block_calls)
|
||||||
|
elif not self.current_tool_name_sent:
|
||||||
|
normal_parts.append(content.strip())
|
||||||
|
except (json.JSONDecodeError, ValueError, TypeError) as exc:
|
||||||
|
logger.warning("Failed to parse streamed dots tool call: %s", exc)
|
||||||
|
normal_parts.append(content.strip())
|
||||||
|
|
||||||
|
self.current_tool_name_sent = False
|
||||||
|
|
||||||
|
return StreamingParseResult(normal_text="".join(normal_parts), calls=calls)
|
||||||
|
|
||||||
|
def _stream_complete_json_body(
|
||||||
|
self, tools: list[Tool], calls: list[ToolCallItem]
|
||||||
|
) -> None:
|
||||||
|
"""Emit a complete JSON body while its closing XML tag is pending."""
|
||||||
|
content = self._buffer[len(self.bot_token) :].strip()
|
||||||
|
if not content or not _is_complete_json(content):
|
||||||
|
return
|
||||||
|
|
||||||
|
try:
|
||||||
|
parsed = json.loads(content)
|
||||||
|
except (json.JSONDecodeError, ValueError):
|
||||||
|
return
|
||||||
|
if not isinstance(parsed, dict):
|
||||||
|
return
|
||||||
|
|
||||||
|
validated = self.parse_base_json(parsed, tools)
|
||||||
|
if not validated:
|
||||||
|
return
|
||||||
|
|
||||||
|
item = validated[0]
|
||||||
|
arguments = item.parameters or ""
|
||||||
|
if not self.current_tool_name_sent:
|
||||||
|
self.current_tool_id += 1
|
||||||
|
calls.append(
|
||||||
|
ToolCallItem(
|
||||||
|
tool_index=self.current_tool_id,
|
||||||
|
name=item.name,
|
||||||
|
parameters="",
|
||||||
|
)
|
||||||
|
)
|
||||||
|
self.prev_tool_call_arr.append(
|
||||||
|
{"name": item.name, "arguments": parsed.get("arguments", {})}
|
||||||
|
)
|
||||||
|
self.streamed_args_for_tool.append("")
|
||||||
|
self.current_tool_name_sent = True
|
||||||
|
|
||||||
|
streamed = self.streamed_args_for_tool[self.current_tool_id]
|
||||||
|
argument_diff = arguments.removeprefix(streamed)
|
||||||
|
if argument_diff:
|
||||||
|
calls.append(
|
||||||
|
ToolCallItem(
|
||||||
|
tool_index=self.current_tool_id,
|
||||||
|
name=None,
|
||||||
|
parameters=argument_diff,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
self.streamed_args_for_tool[self.current_tool_id] += argument_diff
|
||||||
|
|
||||||
|
def flush_pending_normal_text(self) -> str:
|
||||||
|
"""Flush a partial opening marker as plain text at end of stream."""
|
||||||
|
if not self._buffer or self.bot_token in self._buffer:
|
||||||
|
return ""
|
||||||
|
|
||||||
|
normal_text = self._buffer.replace(self.eot_token, "")
|
||||||
|
self._buffer = ""
|
||||||
|
return normal_text
|
||||||
|
|
||||||
|
def supports_structural_tag(self) -> bool:
|
||||||
|
return False
|
||||||
|
|
||||||
|
def structure_info(self) -> _GetInfoFunc:
|
||||||
|
# Kept for the detector interface. It is not used while structural tags
|
||||||
|
# are disabled for dots' mixed XML/JSON format.
|
||||||
|
return lambda name: StructureInfo(
|
||||||
|
begin=f'{self.bot_token}{{"name": "{name}", "arguments": ',
|
||||||
|
end=f"}}{self.eot_token}",
|
||||||
|
trigger=self.bot_token,
|
||||||
|
)
|
||||||
@@ -19,6 +19,7 @@ from sglang.srt.function_call.deepseekv3_detector import DeepSeekV3Detector
|
|||||||
from sglang.srt.function_call.deepseekv4_detector import DeepSeekV4Detector
|
from sglang.srt.function_call.deepseekv4_detector import DeepSeekV4Detector
|
||||||
from sglang.srt.function_call.deepseekv31_detector import DeepSeekV31Detector
|
from sglang.srt.function_call.deepseekv31_detector import DeepSeekV31Detector
|
||||||
from sglang.srt.function_call.deepseekv32_detector import DeepSeekV32Detector
|
from sglang.srt.function_call.deepseekv32_detector import DeepSeekV32Detector
|
||||||
|
from sglang.srt.function_call.dots_detector import DotsToolDetector
|
||||||
from sglang.srt.function_call.gemma4_detector import Gemma4Detector
|
from sglang.srt.function_call.gemma4_detector import Gemma4Detector
|
||||||
from sglang.srt.function_call.gigachat3_detector import GigaChat3Detector
|
from sglang.srt.function_call.gigachat3_detector import GigaChat3Detector
|
||||||
from sglang.srt.function_call.glm4_moe_detector import Glm4MoeDetector
|
from sglang.srt.function_call.glm4_moe_detector import Glm4MoeDetector
|
||||||
@@ -68,6 +69,7 @@ class FunctionCallParser:
|
|||||||
"deepseekv31": DeepSeekV31Detector,
|
"deepseekv31": DeepSeekV31Detector,
|
||||||
"deepseekv32": DeepSeekV32Detector,
|
"deepseekv32": DeepSeekV32Detector,
|
||||||
"deepseekv4": DeepSeekV4Detector,
|
"deepseekv4": DeepSeekV4Detector,
|
||||||
|
"dots": DotsToolDetector,
|
||||||
"glm": Glm4MoeDetector,
|
"glm": Glm4MoeDetector,
|
||||||
"glm45": Glm4MoeDetector,
|
"glm45": Glm4MoeDetector,
|
||||||
"glm47": Glm47MoeDetector,
|
"glm47": Glm47MoeDetector,
|
||||||
|
|||||||
@@ -299,13 +299,28 @@ def attn_backend_wrapper_for_draft_extend(
|
|||||||
the mamba hybrids whose MTP draft is all softmax attention. Inkling's draft has
|
the mamba hybrids whose MTP draft is all softmax attention. Inkling's draft has
|
||||||
its own short convs, so it must expose ``conv_state_metadata`` too.
|
its own short convs, so it must expose ``conv_state_metadata`` too.
|
||||||
"""
|
"""
|
||||||
|
from sglang.srt.configs.dots3 import Dots3Config
|
||||||
from sglang.srt.configs.inkling import InklingMMConfig, InklingModelConfig
|
from sglang.srt.configs.inkling import InklingMMConfig, InklingModelConfig
|
||||||
|
|
||||||
if isinstance(runner.model_config.hf_config, (InklingModelConfig, InklingMMConfig)):
|
if isinstance(runner.model_config.hf_config, (InklingModelConfig, InklingMMConfig)):
|
||||||
return attn_backend_wrapper(runner, full_attn_backend)
|
return attn_backend_wrapper(runner, full_attn_backend)
|
||||||
|
if isinstance(runner.model_config.hf_text_config, Dots3Config):
|
||||||
|
return attn_backend_wrapper(runner, full_attn_backend)
|
||||||
return full_attn_backend
|
return full_attn_backend
|
||||||
|
|
||||||
|
|
||||||
|
def attn_backend_wrapper_for_draft_decode(runner: "ModelRunner", backend):
|
||||||
|
"""Apply the Dots model wrapper to per-step draft backends."""
|
||||||
|
from sglang.srt.configs.dots3 import Dots3Config
|
||||||
|
|
||||||
|
if not hasattr(runner, "model_config"):
|
||||||
|
return backend
|
||||||
|
hf_text_config = runner.model_config.hf_text_config
|
||||||
|
if isinstance(hf_text_config, Dots3Config):
|
||||||
|
return hf_text_config.wrap_draft_decode_attention_backend(backend)
|
||||||
|
return backend
|
||||||
|
|
||||||
|
|
||||||
def attn_backend_wrapper(runner: "ModelRunner", full_attn_backend: "AttentionBackend"):
|
def attn_backend_wrapper(runner: "ModelRunner", full_attn_backend: "AttentionBackend"):
|
||||||
"""
|
"""
|
||||||
Wrapper for special models like hybrid GDN, so we don't
|
Wrapper for special models like hybrid GDN, so we don't
|
||||||
@@ -315,8 +330,14 @@ def attn_backend_wrapper(runner: "ModelRunner", full_attn_backend: "AttentionBac
|
|||||||
hybrid_gdn_config(runner.model_config) is not None and runner.use_mla_backend
|
hybrid_gdn_config(runner.model_config) is not None and runner.use_mla_backend
|
||||||
), "hybrid_gdn can only be used with non-MLA models."
|
), "hybrid_gdn can only be used with non-MLA models."
|
||||||
|
|
||||||
|
from sglang.srt.configs.dots3 import Dots3Config
|
||||||
from sglang.srt.configs.model_config import is_minimax_sparse
|
from sglang.srt.configs.model_config import is_minimax_sparse
|
||||||
|
|
||||||
|
if isinstance(runner.model_config.hf_text_config, Dots3Config):
|
||||||
|
return runner.model_config.hf_text_config.wrap_attention_backend(
|
||||||
|
runner, full_attn_backend
|
||||||
|
)
|
||||||
|
|
||||||
if is_minimax_sparse(runner.model_config.hf_config):
|
if is_minimax_sparse(runner.model_config.hf_config):
|
||||||
from sglang.srt.layers.attention.minimax_sparse_backend import (
|
from sglang.srt.layers.attention.minimax_sparse_backend import (
|
||||||
MiniMaxHybridAttnBackend,
|
MiniMaxHybridAttnBackend,
|
||||||
|
|||||||
@@ -0,0 +1,571 @@
|
|||||||
|
"""Layer-wise DSA/SWA attention dispatch for dots.note.omni."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from contextlib import contextmanager
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from sglang.kernels.ops.attention.flash_attention import flash_attn_varlen_func
|
||||||
|
from sglang.srt.layers.attention.base_attn_backend import (
|
||||||
|
AttentionBackend,
|
||||||
|
SharedReadEnds,
|
||||||
|
)
|
||||||
|
from sglang.srt.layers.attention.hybrid_attn_backend import HybridAttnBackend
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from sglang.srt.layers.radix_attention import RadixAttention
|
||||||
|
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
|
||||||
|
from sglang.srt.speculative.spec_info import SpecInput
|
||||||
|
|
||||||
|
|
||||||
|
def _normalize_page_table_rows(
|
||||||
|
page_table: torch.Tensor, batch_size: int
|
||||||
|
) -> torch.Tensor:
|
||||||
|
"""Match Dots' pre-planned SWA table to the live DP-padded batch."""
|
||||||
|
if page_table.shape[0] >= batch_size:
|
||||||
|
return page_table[:batch_size]
|
||||||
|
return torch.cat(
|
||||||
|
[
|
||||||
|
page_table,
|
||||||
|
page_table.new_zeros(
|
||||||
|
(batch_size - page_table.shape[0], page_table.shape[1])
|
||||||
|
),
|
||||||
|
],
|
||||||
|
dim=0,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _normalize_cache_seqlens_rows(
|
||||||
|
cache_seqlens: torch.Tensor,
|
||||||
|
seq_lens: torch.Tensor,
|
||||||
|
batch_size: int,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
"""Preserve planned rows and fill only newly DP-padded dummy rows."""
|
||||||
|
planned_bs = cache_seqlens.shape[0]
|
||||||
|
if planned_bs >= batch_size:
|
||||||
|
return cache_seqlens[:batch_size]
|
||||||
|
|
||||||
|
dummy_seqlens = seq_lens[planned_bs:batch_size].to(
|
||||||
|
device=cache_seqlens.device,
|
||||||
|
dtype=cache_seqlens.dtype,
|
||||||
|
non_blocking=True,
|
||||||
|
)
|
||||||
|
return torch.cat([cache_seqlens, dummy_seqlens], dim=0)
|
||||||
|
|
||||||
|
|
||||||
|
def _metadata_mismatches_dp_padded_batch(metadata, forward_batch) -> bool:
|
||||||
|
"""True when pre-planned attention metadata no longer matches the live batch.
|
||||||
|
|
||||||
|
EAGLE plans draft metadata before ModelRunner runs DP/MLP padding. Dummy
|
||||||
|
request rows then change ``batch_size`` / ``out_cache_loc``, leaving
|
||||||
|
page tables and SWA write targets short. Rebuilding is required; slicing
|
||||||
|
or zero-padding the stale tensors is not enough for DSA + SWA.
|
||||||
|
"""
|
||||||
|
if metadata is None:
|
||||||
|
return False
|
||||||
|
bs = forward_batch.batch_size
|
||||||
|
from sglang.srt.layers.attention.flashattention_backend import (
|
||||||
|
FlashAttentionMetadata,
|
||||||
|
)
|
||||||
|
|
||||||
|
if isinstance(metadata, FlashAttentionMetadata):
|
||||||
|
if metadata.page_table is not None and metadata.page_table.shape[0] != bs:
|
||||||
|
return True
|
||||||
|
if (
|
||||||
|
metadata.swa_page_table is not None
|
||||||
|
and metadata.swa_page_table.shape[0] != bs
|
||||||
|
):
|
||||||
|
return True
|
||||||
|
if (
|
||||||
|
metadata.cache_seqlens_int32 is not None
|
||||||
|
and metadata.cache_seqlens_int32.shape[0] != bs
|
||||||
|
):
|
||||||
|
return True
|
||||||
|
swa_loc = metadata.swa_out_cache_loc
|
||||||
|
out_loc = forward_batch.out_cache_loc
|
||||||
|
return (
|
||||||
|
swa_loc is not None
|
||||||
|
and out_loc is not None
|
||||||
|
and swa_loc.shape[0] != out_loc.shape[0]
|
||||||
|
)
|
||||||
|
|
||||||
|
from sglang.srt.layers.attention.dsa_backend import DSAMetadata
|
||||||
|
|
||||||
|
if isinstance(metadata, DSAMetadata):
|
||||||
|
return metadata.cache_seqlens_int32.shape[0] != bs
|
||||||
|
return False
|
||||||
|
|
||||||
|
|
||||||
|
def _dp_padding_changed_batch_size(forward_batch) -> bool:
|
||||||
|
original_bs = forward_batch._original_batch_size
|
||||||
|
return original_bs is not None and original_bs != forward_batch.batch_size
|
||||||
|
|
||||||
|
|
||||||
|
def _maybe_rebuild_dots_metadata(backend, forward_batch) -> None:
|
||||||
|
"""Eager-only rebuild when DP padding invalidated a Dots pre-plan."""
|
||||||
|
from sglang.srt.model_executor.runner_utils.capture_mode import (
|
||||||
|
get_is_capture_mode,
|
||||||
|
)
|
||||||
|
|
||||||
|
if get_is_capture_mode():
|
||||||
|
return
|
||||||
|
if backend._dp_rebuilt_batch_id == id(forward_batch):
|
||||||
|
return
|
||||||
|
stale = _metadata_mismatches_dp_padded_batch(
|
||||||
|
backend.forward_metadata, forward_batch
|
||||||
|
) or _dp_padding_changed_batch_size(forward_batch)
|
||||||
|
if not stale:
|
||||||
|
return
|
||||||
|
backend.init_forward_metadata(forward_batch)
|
||||||
|
backend._dp_rebuilt_batch_id = id(forward_batch)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class DotsSWAMLAPrefillMetadata:
|
||||||
|
kv_indices: torch.Tensor
|
||||||
|
cu_seqlens_q: torch.Tensor
|
||||||
|
cu_seqlens_k: torch.Tensor
|
||||||
|
max_seq_len_q: int
|
||||||
|
max_seq_len_k: int
|
||||||
|
|
||||||
|
|
||||||
|
class DotsSWAMLAAttnBackend(AttentionBackend):
|
||||||
|
"""Add Dots latent-cache SWA support around a FlashAttention backend."""
|
||||||
|
|
||||||
|
def __init__(self, backend: AttentionBackend):
|
||||||
|
self.backend = backend
|
||||||
|
self._active_backend = backend
|
||||||
|
self.token_to_kv_pool = backend.token_to_kv_pool
|
||||||
|
self.req_to_token_pool = backend.req_to_token_pool
|
||||||
|
self.needs_cpu_seq_lens = True
|
||||||
|
self._prefill_metadata: DotsSWAMLAPrefillMetadata | None = None
|
||||||
|
self._dp_rebuilt_batch_id: int | None = None
|
||||||
|
|
||||||
|
@property
|
||||||
|
def forward_metadata(self):
|
||||||
|
return self._active_backend.forward_metadata
|
||||||
|
|
||||||
|
@forward_metadata.setter
|
||||||
|
def forward_metadata(self, value):
|
||||||
|
self._active_backend.forward_metadata = value
|
||||||
|
|
||||||
|
@property
|
||||||
|
def verify_mask(self):
|
||||||
|
return self.backend.verify_mask
|
||||||
|
|
||||||
|
def shared_read_ends(self, fm: ForwardMode) -> SharedReadEnds:
|
||||||
|
return self.backend.shared_read_ends(fm)
|
||||||
|
|
||||||
|
def draft_extend_metadata_captured_in_graph(self) -> bool:
|
||||||
|
return self.backend.draft_extend_metadata_captured_in_graph()
|
||||||
|
|
||||||
|
def selected_backend(self, forward_batch: ForwardBatch) -> AttentionBackend:
|
||||||
|
return (
|
||||||
|
self.backend._select_backend(forward_batch.forward_mode)
|
||||||
|
if isinstance(self.backend, HybridAttnBackend)
|
||||||
|
else self.backend
|
||||||
|
)
|
||||||
|
|
||||||
|
def uses_flash_attention(self, forward_batch: ForwardBatch) -> bool:
|
||||||
|
from sglang.srt.layers.attention.flashattention_backend import (
|
||||||
|
FlashAttentionBackend,
|
||||||
|
)
|
||||||
|
|
||||||
|
return isinstance(self.selected_backend(forward_batch), FlashAttentionBackend)
|
||||||
|
|
||||||
|
def maybe_rebuild_metadata_after_dp_padding(
|
||||||
|
self, forward_batch: ForwardBatch
|
||||||
|
) -> None:
|
||||||
|
"""Rebuild FA + SWA-prefill metadata after eager DP dummy-row padding."""
|
||||||
|
self._active_backend = self.selected_backend(forward_batch)
|
||||||
|
_maybe_rebuild_dots_metadata(self, forward_batch)
|
||||||
|
|
||||||
|
def select_draft_step_out_cache_loc(self, forward_batch: ForwardBatch):
|
||||||
|
"""Return this draft step's write locations from a combined SWA buffer."""
|
||||||
|
from sglang.srt.layers.attention.flashattention_backend import (
|
||||||
|
FlashAttentionBackend,
|
||||||
|
)
|
||||||
|
|
||||||
|
out_cache_loc = forward_batch.out_cache_loc
|
||||||
|
backend = self._active_backend
|
||||||
|
if not isinstance(backend, FlashAttentionBackend):
|
||||||
|
return out_cache_loc
|
||||||
|
if (
|
||||||
|
out_cache_loc is not None
|
||||||
|
and forward_batch.forward_mode.is_decode_or_idle()
|
||||||
|
and forward_batch.spec_info is not None
|
||||||
|
and backend.speculative_num_steps > 0
|
||||||
|
and out_cache_loc.numel()
|
||||||
|
== forward_batch.batch_size * backend.topk * backend.speculative_num_steps
|
||||||
|
):
|
||||||
|
return out_cache_loc.view(
|
||||||
|
forward_batch.batch_size,
|
||||||
|
backend.topk,
|
||||||
|
backend.speculative_num_steps,
|
||||||
|
)[:, :, backend.speculative_step_id].reshape(-1)
|
||||||
|
return out_cache_loc
|
||||||
|
|
||||||
|
@contextmanager
|
||||||
|
def _use_draft_step_out_cache_loc(self, forward_batch: ForwardBatch):
|
||||||
|
"""Expose only this backend's draft-step write locations to FA."""
|
||||||
|
original = forward_batch.out_cache_loc
|
||||||
|
forward_batch.out_cache_loc = self.select_draft_step_out_cache_loc(
|
||||||
|
forward_batch
|
||||||
|
)
|
||||||
|
try:
|
||||||
|
yield
|
||||||
|
finally:
|
||||||
|
forward_batch.out_cache_loc = original
|
||||||
|
|
||||||
|
def init_forward_metadata(self, forward_batch: ForwardBatch):
|
||||||
|
self._active_backend = self.selected_backend(forward_batch)
|
||||||
|
with self._use_draft_step_out_cache_loc(forward_batch):
|
||||||
|
self.backend.init_forward_metadata(forward_batch)
|
||||||
|
self._init_prefill_metadata(forward_batch)
|
||||||
|
|
||||||
|
def init_forward_metadata_out_graph(
|
||||||
|
self, forward_batch: ForwardBatch, in_capture: bool = False
|
||||||
|
):
|
||||||
|
self._active_backend = self.selected_backend(forward_batch)
|
||||||
|
with self._use_draft_step_out_cache_loc(forward_batch):
|
||||||
|
self.backend.init_forward_metadata_out_graph(
|
||||||
|
forward_batch, in_capture=in_capture
|
||||||
|
)
|
||||||
|
self._init_prefill_metadata(forward_batch)
|
||||||
|
|
||||||
|
def init_forward_metadata_in_graph(self, forward_batch: ForwardBatch):
|
||||||
|
self.backend.init_forward_metadata_in_graph(forward_batch)
|
||||||
|
|
||||||
|
def init_cuda_graph_state(self, max_bs: int, max_num_tokens: int):
|
||||||
|
self.backend.init_cuda_graph_state(max_bs, max_num_tokens)
|
||||||
|
|
||||||
|
def get_cuda_graph_seq_len_fill_value(self):
|
||||||
|
return self.backend.get_cuda_graph_seq_len_fill_value()
|
||||||
|
|
||||||
|
def on_after_cuda_graph_warmup(self):
|
||||||
|
self.backend.on_after_cuda_graph_warmup()
|
||||||
|
|
||||||
|
def update_verify_buffers_to_fill_after_draft(
|
||||||
|
self, spec_info: SpecInput, cuda_graph_bs: int | None
|
||||||
|
):
|
||||||
|
return self.backend.update_verify_buffers_to_fill_after_draft(
|
||||||
|
spec_info, cuda_graph_bs
|
||||||
|
)
|
||||||
|
|
||||||
|
def forward(self, q, k, v, layer, forward_batch, save_kv_cache=True, **kwargs):
|
||||||
|
self.maybe_rebuild_metadata_after_dp_padding(forward_batch)
|
||||||
|
return self.backend.forward(
|
||||||
|
q, k, v, layer, forward_batch, save_kv_cache, **kwargs
|
||||||
|
)
|
||||||
|
|
||||||
|
def forward_extend(
|
||||||
|
self, q, k, v, layer, forward_batch, save_kv_cache=True, **kwargs
|
||||||
|
):
|
||||||
|
return self.backend.forward_extend(
|
||||||
|
q, k, v, layer, forward_batch, save_kv_cache, **kwargs
|
||||||
|
)
|
||||||
|
|
||||||
|
def forward_decode(
|
||||||
|
self, q, k, v, layer, forward_batch, save_kv_cache=True, **kwargs
|
||||||
|
):
|
||||||
|
return self.backend.forward_decode(
|
||||||
|
q, k, v, layer, forward_batch, save_kv_cache, **kwargs
|
||||||
|
)
|
||||||
|
|
||||||
|
def init_mha_chunk_metadata(self, forward_batch: ForwardBatch):
|
||||||
|
self.backend.init_mha_chunk_metadata(forward_batch)
|
||||||
|
|
||||||
|
def _init_prefill_metadata(self, forward_batch: ForwardBatch):
|
||||||
|
if not forward_batch.forward_mode.is_extend_without_speculative():
|
||||||
|
self._prefill_metadata = None
|
||||||
|
return
|
||||||
|
|
||||||
|
metadata = self._active_backend.forward_metadata
|
||||||
|
assert forward_batch.seq_lens_cpu is not None
|
||||||
|
batch_kv_indices = self._active_backend.req_to_token[
|
||||||
|
forward_batch.req_pool_indices, :
|
||||||
|
]
|
||||||
|
sliced_indices = []
|
||||||
|
kv_lens = []
|
||||||
|
for i in range(forward_batch.batch_size):
|
||||||
|
q_len = int(forward_batch.extend_seq_lens_cpu[i])
|
||||||
|
kv_len = int(forward_batch.seq_lens_cpu[i])
|
||||||
|
tail_len = min(q_len + self._active_backend.sliding_window_size, kv_len)
|
||||||
|
sliced_indices.append(batch_kv_indices[i, kv_len - tail_len : kv_len])
|
||||||
|
kv_lens.append(tail_len)
|
||||||
|
|
||||||
|
full_kv_indices = torch.cat(sliced_indices)
|
||||||
|
kv_indices = self.token_to_kv_pool.translate_loc_from_full_to_swa(
|
||||||
|
full_kv_indices
|
||||||
|
).to(torch.int32)
|
||||||
|
lens_cpu = torch.tensor([0, *kv_lens], dtype=torch.int32, pin_memory=True)
|
||||||
|
self._prefill_metadata = DotsSWAMLAPrefillMetadata(
|
||||||
|
kv_indices=kv_indices,
|
||||||
|
cu_seqlens_q=metadata.cu_seqlens_q,
|
||||||
|
cu_seqlens_k=torch.cumsum(
|
||||||
|
lens_cpu.to(device=forward_batch.seq_lens.device, non_blocking=True),
|
||||||
|
dim=0,
|
||||||
|
dtype=torch.int32,
|
||||||
|
),
|
||||||
|
max_seq_len_q=metadata.max_seq_len_q,
|
||||||
|
max_seq_len_k=max(kv_lens),
|
||||||
|
)
|
||||||
|
|
||||||
|
def get_swa_mla_prefill_latent_cache(
|
||||||
|
self, forward_batch: ForwardBatch, layer_id: int
|
||||||
|
):
|
||||||
|
assert self._prefill_metadata is not None
|
||||||
|
return self.token_to_kv_pool.get_key_buffer(layer_id)[
|
||||||
|
self._prefill_metadata.kv_indices
|
||||||
|
]
|
||||||
|
|
||||||
|
def forward_swa_mla_expanded(self, q, k, v, layer, forward_batch=None):
|
||||||
|
"""Run dense SWA after Dots expands its compact MLA cache."""
|
||||||
|
metadata = self._prefill_metadata
|
||||||
|
assert metadata is not None
|
||||||
|
q = q.view(-1, layer.tp_q_head_num, layer.head_dim)
|
||||||
|
k = k.view(-1, layer.tp_k_head_num, layer.head_dim).to(q.dtype)
|
||||||
|
v = v.view(-1, layer.tp_k_head_num, layer.v_head_dim).to(q.dtype)
|
||||||
|
|
||||||
|
# FA3 requires equal QK/V widths when QK exceeds 192.
|
||||||
|
pad_v_to_qk = layer.head_dim > 192 and layer.v_head_dim != layer.head_dim
|
||||||
|
if pad_v_to_qk:
|
||||||
|
v = torch.nn.functional.pad(v, (0, layer.head_dim - layer.v_head_dim))
|
||||||
|
|
||||||
|
output = flash_attn_varlen_func(
|
||||||
|
q=q,
|
||||||
|
k=k,
|
||||||
|
v=v,
|
||||||
|
cu_seqlens_q=metadata.cu_seqlens_q,
|
||||||
|
cu_seqlens_k=metadata.cu_seqlens_k,
|
||||||
|
max_seqlen_q=metadata.max_seq_len_q,
|
||||||
|
max_seqlen_k=metadata.max_seq_len_k,
|
||||||
|
softmax_scale=layer.scaling,
|
||||||
|
causal=True,
|
||||||
|
window_size=(layer.sliding_window_size, 0),
|
||||||
|
ver=self._active_backend.fa_impl_ver,
|
||||||
|
)
|
||||||
|
if pad_v_to_qk:
|
||||||
|
output = output[..., : layer.v_head_dim]
|
||||||
|
return output.reshape(-1, layer.tp_q_head_num * layer.v_head_dim)
|
||||||
|
|
||||||
|
def forward_swa_mla_absorbed(self, q, layer, forward_batch):
|
||||||
|
"""Run decode directly against the page64 latent SWA cache."""
|
||||||
|
from sglang.srt.layers.attention.swa_mla_fallback.forward import (
|
||||||
|
forward_dense_kvlora_swa_torch_fallback,
|
||||||
|
)
|
||||||
|
|
||||||
|
backend = self.selected_backend(forward_batch)
|
||||||
|
if backend.page_size != 64:
|
||||||
|
raise RuntimeError(
|
||||||
|
"Dots SWA latent decode requires page_size=64, "
|
||||||
|
f"got {backend.page_size}."
|
||||||
|
)
|
||||||
|
|
||||||
|
self.maybe_rebuild_metadata_after_dp_padding(forward_batch)
|
||||||
|
metadata = backend.forward_metadata
|
||||||
|
block_table = metadata.swa_page_table
|
||||||
|
if block_table is None:
|
||||||
|
raise RuntimeError("Dots SWA latent decode requires an SWA page table.")
|
||||||
|
bs = forward_batch.batch_size
|
||||||
|
block_table = _normalize_page_table_rows(block_table, bs)
|
||||||
|
cache_seqlens = _normalize_cache_seqlens_rows(
|
||||||
|
metadata.cache_seqlens_int32,
|
||||||
|
forward_batch.seq_lens,
|
||||||
|
bs,
|
||||||
|
)
|
||||||
|
reshape_q = q.view(bs, -1, layer.tp_q_head_num, layer.head_dim)
|
||||||
|
k_cache = self.token_to_kv_pool.get_key_buffer(layer.layer_id)
|
||||||
|
output = forward_dense_kvlora_swa_torch_fallback(
|
||||||
|
reshape_q=reshape_q,
|
||||||
|
k_cache=k_cache,
|
||||||
|
block_table=block_table,
|
||||||
|
cache_seqlens=cache_seqlens,
|
||||||
|
layer=layer,
|
||||||
|
kv_cache_dim=layer.head_dim,
|
||||||
|
head_dim_v=layer.v_head_dim,
|
||||||
|
window_size=layer.sliding_window_size + 1,
|
||||||
|
)
|
||||||
|
return output.view(-1, layer.tp_q_head_num * layer.v_head_dim)
|
||||||
|
|
||||||
|
|
||||||
|
class DotsHybridAttnBackend(AttentionBackend):
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
dsa_backend: AttentionBackend,
|
||||||
|
swa_backend: AttentionBackend,
|
||||||
|
):
|
||||||
|
self.dsa_backend = dsa_backend
|
||||||
|
# Keep DSA on its radix-aware MLA path.
|
||||||
|
self.dsa_backend.supports_mha_one_shot = False
|
||||||
|
self.swa_backend = swa_backend
|
||||||
|
self.token_to_kv_pool = swa_backend.token_to_kv_pool
|
||||||
|
self.req_to_token_pool = swa_backend.req_to_token_pool
|
||||||
|
# SWA latent expansion uses host sequence-length mirrors.
|
||||||
|
self.needs_cpu_seq_lens = True
|
||||||
|
self._dp_rebuilt_batch_id: int | None = None
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _is_swa_layer(layer: RadixAttention) -> bool:
|
||||||
|
return layer.sliding_window_size is not None and layer.sliding_window_size > -1
|
||||||
|
|
||||||
|
def backend_for_layer(self, layer: RadixAttention) -> AttentionBackend:
|
||||||
|
return self.swa_backend if self._is_swa_layer(layer) else self.dsa_backend
|
||||||
|
|
||||||
|
def selected_swa_backend(self, forward_batch: ForwardBatch) -> AttentionBackend:
|
||||||
|
return (
|
||||||
|
self.swa_backend._select_backend(forward_batch.forward_mode)
|
||||||
|
if isinstance(self.swa_backend, HybridAttnBackend)
|
||||||
|
else self.swa_backend
|
||||||
|
)
|
||||||
|
|
||||||
|
def maybe_rebuild_metadata_after_dp_padding(
|
||||||
|
self, forward_batch: ForwardBatch
|
||||||
|
) -> None:
|
||||||
|
"""Rebuild both DSA and SWA plans after eager DP dummy-row padding."""
|
||||||
|
from sglang.srt.model_executor.runner_utils.capture_mode import (
|
||||||
|
get_is_capture_mode,
|
||||||
|
)
|
||||||
|
|
||||||
|
if get_is_capture_mode():
|
||||||
|
return
|
||||||
|
if self._dp_rebuilt_batch_id == id(forward_batch):
|
||||||
|
return
|
||||||
|
dsa_stale = _metadata_mismatches_dp_padded_batch(
|
||||||
|
self.dsa_backend.forward_metadata, forward_batch
|
||||||
|
)
|
||||||
|
swa_backend = self.selected_swa_backend(forward_batch)
|
||||||
|
swa_stale = _metadata_mismatches_dp_padded_batch(
|
||||||
|
swa_backend.forward_metadata, forward_batch
|
||||||
|
)
|
||||||
|
if dsa_stale or swa_stale or _dp_padding_changed_batch_size(forward_batch):
|
||||||
|
self.init_forward_metadata(forward_batch)
|
||||||
|
self._dp_rebuilt_batch_id = id(forward_batch)
|
||||||
|
|
||||||
|
def init_forward_metadata(self, forward_batch: ForwardBatch):
|
||||||
|
self.dsa_backend.init_forward_metadata(forward_batch)
|
||||||
|
self.swa_backend.init_forward_metadata(forward_batch)
|
||||||
|
|
||||||
|
def init_forward_metadata_out_graph(
|
||||||
|
self, forward_batch: ForwardBatch, in_capture: bool = False
|
||||||
|
):
|
||||||
|
self.dsa_backend.init_forward_metadata_out_graph(
|
||||||
|
forward_batch, in_capture=in_capture
|
||||||
|
)
|
||||||
|
self.swa_backend.init_forward_metadata_out_graph(
|
||||||
|
forward_batch, in_capture=in_capture
|
||||||
|
)
|
||||||
|
|
||||||
|
def init_forward_metadata_in_graph(self, forward_batch: ForwardBatch):
|
||||||
|
self.dsa_backend.init_forward_metadata_in_graph(forward_batch)
|
||||||
|
self.swa_backend.init_forward_metadata_in_graph(forward_batch)
|
||||||
|
|
||||||
|
def init_cuda_graph_state(self, max_bs: int, max_num_tokens: int):
|
||||||
|
self.dsa_backend.init_cuda_graph_state(max_bs, max_num_tokens)
|
||||||
|
self.swa_backend.init_cuda_graph_state(max_bs, max_num_tokens)
|
||||||
|
|
||||||
|
def get_cuda_graph_seq_len_fill_value(self):
|
||||||
|
return self.swa_backend.get_cuda_graph_seq_len_fill_value()
|
||||||
|
|
||||||
|
def on_after_cuda_graph_warmup(self):
|
||||||
|
self.dsa_backend.on_after_cuda_graph_warmup()
|
||||||
|
self.swa_backend.on_after_cuda_graph_warmup()
|
||||||
|
|
||||||
|
def forward(
|
||||||
|
self,
|
||||||
|
q: torch.Tensor,
|
||||||
|
k: torch.Tensor,
|
||||||
|
v: torch.Tensor,
|
||||||
|
layer: RadixAttention,
|
||||||
|
forward_batch: ForwardBatch,
|
||||||
|
save_kv_cache: bool = True,
|
||||||
|
**kwargs,
|
||||||
|
):
|
||||||
|
self.maybe_rebuild_metadata_after_dp_padding(forward_batch)
|
||||||
|
return self.backend_for_layer(layer).forward(
|
||||||
|
q, k, v, layer, forward_batch, save_kv_cache, **kwargs
|
||||||
|
)
|
||||||
|
|
||||||
|
def forward_extend(
|
||||||
|
self, q, k, v, layer, forward_batch, save_kv_cache=True, **kwargs
|
||||||
|
):
|
||||||
|
return self.backend_for_layer(layer).forward_extend(
|
||||||
|
q, k, v, layer, forward_batch, save_kv_cache, **kwargs
|
||||||
|
)
|
||||||
|
|
||||||
|
def forward_decode(
|
||||||
|
self, q, k, v, layer, forward_batch, save_kv_cache=True, **kwargs
|
||||||
|
):
|
||||||
|
return self.backend_for_layer(layer).forward_decode(
|
||||||
|
q, k, v, layer, forward_batch, save_kv_cache, **kwargs
|
||||||
|
)
|
||||||
|
|
||||||
|
def get_indexer_metadata(self, layer_id: int, forward_batch: ForwardBatch):
|
||||||
|
return self.dsa_backend.get_indexer_metadata(layer_id, forward_batch)
|
||||||
|
|
||||||
|
def get_swa_mla_prefill_latent_cache(
|
||||||
|
self, forward_batch: ForwardBatch, layer_id: int
|
||||||
|
):
|
||||||
|
backend = self.selected_swa_backend(forward_batch)
|
||||||
|
return backend.get_swa_mla_prefill_latent_cache(forward_batch, layer_id)
|
||||||
|
|
||||||
|
def forward_swa_mla_expanded(self, q, k, v, layer, forward_batch):
|
||||||
|
backend = self.selected_swa_backend(forward_batch)
|
||||||
|
return backend.forward_swa_mla_expanded(q, k, v, layer, forward_batch)
|
||||||
|
|
||||||
|
def forward_swa_mla_absorbed(self, q, layer, forward_batch):
|
||||||
|
backend = self.selected_swa_backend(forward_batch)
|
||||||
|
return backend.forward_swa_mla_absorbed(q, layer, forward_batch)
|
||||||
|
|
||||||
|
def init_mha_chunk_metadata(self, forward_batch: ForwardBatch):
|
||||||
|
backend = self.selected_swa_backend(forward_batch)
|
||||||
|
backend.init_mha_chunk_metadata(forward_batch)
|
||||||
|
|
||||||
|
|
||||||
|
def _wrap_dots_swa_backend(backend: AttentionBackend) -> AttentionBackend:
|
||||||
|
"""Add latent-cache SWA behavior when a backend uses FlashAttention."""
|
||||||
|
from sglang.srt.layers.attention.flashattention_backend import (
|
||||||
|
FlashAttentionBackend,
|
||||||
|
)
|
||||||
|
|
||||||
|
if isinstance(backend, FlashAttentionBackend) or (
|
||||||
|
isinstance(backend, HybridAttnBackend)
|
||||||
|
and (
|
||||||
|
isinstance(backend.prefill_backend, FlashAttentionBackend)
|
||||||
|
or isinstance(backend.decode_backend, FlashAttentionBackend)
|
||||||
|
)
|
||||||
|
):
|
||||||
|
return DotsSWAMLAAttnBackend(backend)
|
||||||
|
return backend
|
||||||
|
|
||||||
|
|
||||||
|
def wrap_dots_draft_decode_backend(backend: AttentionBackend) -> AttentionBackend:
|
||||||
|
"""Wrap each per-step backend used by the Dots NextN draft container."""
|
||||||
|
backend.attn_backends = [
|
||||||
|
_wrap_dots_swa_backend(child) for child in backend.attn_backends
|
||||||
|
]
|
||||||
|
return backend
|
||||||
|
|
||||||
|
|
||||||
|
def wrap_dots_attention_backend(runner, full_attn_backend: AttentionBackend):
|
||||||
|
"""Construct the Dots target or draft attention backend."""
|
||||||
|
if runner.model_config.is_draft_model:
|
||||||
|
return _wrap_dots_swa_backend(full_attn_backend)
|
||||||
|
|
||||||
|
if runner.model_config.hf_text_config.index_topk is None:
|
||||||
|
return DotsSWAMLAAttnBackend(full_attn_backend)
|
||||||
|
|
||||||
|
from sglang.srt.layers.attention.attention_registry import create_dsa_backend
|
||||||
|
|
||||||
|
swa_backend = (
|
||||||
|
full_attn_backend.prefill_backend
|
||||||
|
if isinstance(full_attn_backend, HybridAttnBackend)
|
||||||
|
else full_attn_backend
|
||||||
|
)
|
||||||
|
return DotsHybridAttnBackend(
|
||||||
|
dsa_backend=create_dsa_backend(runner),
|
||||||
|
swa_backend=DotsSWAMLAAttnBackend(swa_backend),
|
||||||
|
)
|
||||||
@@ -575,7 +575,13 @@ class Indexer(DSANPUIndexerMixin, BaseFusedOp):
|
|||||||
key, [self.rope_head_dim, self.head_dim - self.rope_head_dim], dim=-1
|
key, [self.rope_head_dim, self.head_dim - self.rope_head_dim], dim=-1
|
||||||
)
|
)
|
||||||
|
|
||||||
_, k_rope = self.rotary_emb(positions, k_rope, k_rope)
|
# Rotary may update both inputs in place, so the K-only path must not
|
||||||
|
# alias its dummy query with the key.
|
||||||
|
if _is_cuda or _is_hip or _is_xpu:
|
||||||
|
dummy_q_rope = torch.empty_like(k_rope)
|
||||||
|
else:
|
||||||
|
dummy_q_rope = k_rope
|
||||||
|
_, k_rope = self.rotary_emb(positions, dummy_q_rope, k_rope)
|
||||||
self._update_rope_guarded(key[..., : self.rope_head_dim], k_rope)
|
self._update_rope_guarded(key[..., : self.rope_head_dim], k_rope)
|
||||||
key = rotate_activation(key)
|
key = rotate_activation(key)
|
||||||
|
|
||||||
|
|||||||
@@ -333,6 +333,7 @@ class DeepseekSparseAttnBackend(
|
|||||||
self.req_to_token = model_runner.req_to_token_pool.req_to_token
|
self.req_to_token = model_runner.req_to_token_pool.req_to_token
|
||||||
|
|
||||||
self.use_mha: bool = False
|
self.use_mha: bool = False
|
||||||
|
self.supports_mha_one_shot: bool = True
|
||||||
self.dsa_prefill_impl: _DSA_IMPL_T = (
|
self.dsa_prefill_impl: _DSA_IMPL_T = (
|
||||||
model_runner.server_args.dsa_prefill_backend
|
model_runner.server_args.dsa_prefill_backend
|
||||||
)
|
)
|
||||||
@@ -3325,7 +3326,8 @@ class DeepseekSparseAttnBackend(
|
|||||||
|
|
||||||
# Requirements: H200/B200/MI355X, short sequences, supported dtype, fits in chunk
|
# Requirements: H200/B200/MI355X, short sequences, supported dtype, fits in chunk
|
||||||
self.use_mha = (
|
self.use_mha = (
|
||||||
(
|
self.supports_mha_one_shot
|
||||||
|
and (
|
||||||
device_sm == 90
|
device_sm == 90
|
||||||
or (device_sm >= 100 and device_sm < 110)
|
or (device_sm >= 100 and device_sm < 110)
|
||||||
or _IS_GFX95
|
or _IS_GFX95
|
||||||
|
|||||||
@@ -0,0 +1 @@
|
|||||||
|
"""Fallback operations for sliding-window MLA attention paths."""
|
||||||
@@ -0,0 +1,77 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from sglang.srt.layers.attention.swa_mla_fallback.ops import (
|
||||||
|
apply_swa_score_mask,
|
||||||
|
gather_page64_kv_latent,
|
||||||
|
)
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from sglang.srt.layers.radix_attention import RadixAttention
|
||||||
|
|
||||||
|
|
||||||
|
def forward_dense_kvlora_swa_torch_fallback(
|
||||||
|
reshape_q: torch.Tensor,
|
||||||
|
k_cache: torch.Tensor,
|
||||||
|
block_table: torch.Tensor,
|
||||||
|
cache_seqlens: torch.Tensor,
|
||||||
|
layer: RadixAttention,
|
||||||
|
kv_cache_dim: int,
|
||||||
|
head_dim_v: int,
|
||||||
|
window_size: int,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
"""Page-64 SWA fallback for dense KV-LoRA decode."""
|
||||||
|
if layer.tp_k_head_num != 1:
|
||||||
|
raise RuntimeError(
|
||||||
|
"SWA MLA torch fallback currently supports MLA with one "
|
||||||
|
f"KV head, got tp_k_head_num={layer.tp_k_head_num}."
|
||||||
|
)
|
||||||
|
|
||||||
|
bs, s_q, num_heads, qk_dim = reshape_q.shape
|
||||||
|
if qk_dim != kv_cache_dim:
|
||||||
|
raise RuntimeError(
|
||||||
|
f"SWA MLA torch fallback got q dim {qk_dim}, "
|
||||||
|
f"expected kv_cache_dim {kv_cache_dim}."
|
||||||
|
)
|
||||||
|
if s_q not in (1, 4):
|
||||||
|
raise RuntimeError(
|
||||||
|
"SWA MLA torch fallback mask is specialized for s_q=1 "
|
||||||
|
f"or s_q=4, got s_q={s_q}."
|
||||||
|
)
|
||||||
|
|
||||||
|
# Include the full union of causal windows and align it for BMM.
|
||||||
|
kv_latent, kv_valid = gather_page64_kv_latent(
|
||||||
|
k_cache,
|
||||||
|
block_table,
|
||||||
|
cache_seqlens,
|
||||||
|
window_size,
|
||||||
|
s_q,
|
||||||
|
kv_cache_dim,
|
||||||
|
)
|
||||||
|
gather_len = kv_latent.shape[1]
|
||||||
|
|
||||||
|
# Keep the output in [bs, s_q, num_heads, head_dim_v] order.
|
||||||
|
q_for_scores = reshape_q.reshape(bs, s_q * num_heads, qk_dim)
|
||||||
|
scores = torch.bmm(q_for_scores, kv_latent.transpose(1, 2)).view(
|
||||||
|
bs, s_q, num_heads, gather_len
|
||||||
|
)
|
||||||
|
scores = scores.float()
|
||||||
|
scores.mul_(layer.scaling)
|
||||||
|
|
||||||
|
apply_swa_score_mask(
|
||||||
|
scores.transpose(1, 2),
|
||||||
|
cache_seqlens,
|
||||||
|
kv_valid,
|
||||||
|
num_heads,
|
||||||
|
window_size,
|
||||||
|
s_q,
|
||||||
|
)
|
||||||
|
|
||||||
|
probs = torch.softmax(scores, dim=-1).to(reshape_q.dtype)
|
||||||
|
return torch.bmm(
|
||||||
|
probs.reshape(bs, s_q * num_heads, gather_len),
|
||||||
|
kv_latent[..., :head_dim_v],
|
||||||
|
).view(bs, s_q, num_heads, head_dim_v)
|
||||||
@@ -0,0 +1,205 @@
|
|||||||
|
import torch
|
||||||
|
import triton
|
||||||
|
import triton.language as tl
|
||||||
|
|
||||||
|
PAGE_SIZE = 64
|
||||||
|
|
||||||
|
|
||||||
|
@triton.jit
|
||||||
|
def _gather_page64_kv_latent_kernel(
|
||||||
|
k_cache_ptr,
|
||||||
|
block_table_ptr,
|
||||||
|
cache_seqlens_ptr,
|
||||||
|
kv_out_ptr,
|
||||||
|
valid_out_ptr,
|
||||||
|
k_cache_stride_t,
|
||||||
|
k_cache_stride_d,
|
||||||
|
block_table_stride_b,
|
||||||
|
block_table_num_pages: tl.constexpr,
|
||||||
|
k_cache_num_tokens: tl.constexpr,
|
||||||
|
kv_out_stride_b,
|
||||||
|
kv_out_stride_t,
|
||||||
|
kv_out_stride_d,
|
||||||
|
valid_out_stride_b,
|
||||||
|
valid_out_stride_t,
|
||||||
|
GATHER_LEN: tl.constexpr,
|
||||||
|
KV_DIM: tl.constexpr,
|
||||||
|
BLOCK_T: tl.constexpr,
|
||||||
|
BLOCK_D: tl.constexpr,
|
||||||
|
PAGE_SIZE_: tl.constexpr,
|
||||||
|
):
|
||||||
|
bid = tl.program_id(0)
|
||||||
|
tid = tl.program_id(1)
|
||||||
|
did = tl.program_id(2)
|
||||||
|
|
||||||
|
cache_seqlen = tl.load(cache_seqlens_ptr + bid).to(tl.int32)
|
||||||
|
gather_start = tl.maximum(cache_seqlen - GATHER_LEN, 0)
|
||||||
|
offs_t = tid * BLOCK_T + tl.arange(0, BLOCK_T)
|
||||||
|
logical_token = gather_start + offs_t
|
||||||
|
logical_valid = (offs_t < GATHER_LEN) & (logical_token < cache_seqlen)
|
||||||
|
|
||||||
|
logical_page = logical_token // PAGE_SIZE_
|
||||||
|
intra_page = logical_token - logical_page * PAGE_SIZE_
|
||||||
|
page_table_valid = logical_valid & (logical_page < block_table_num_pages)
|
||||||
|
physical_page = tl.load(
|
||||||
|
block_table_ptr + bid * block_table_stride_b + logical_page,
|
||||||
|
mask=page_table_valid,
|
||||||
|
other=-1,
|
||||||
|
).to(tl.int32)
|
||||||
|
physical_token = physical_page * PAGE_SIZE_ + intra_page
|
||||||
|
physical_valid = (
|
||||||
|
page_table_valid & (physical_page >= 0) & (physical_token < k_cache_num_tokens)
|
||||||
|
)
|
||||||
|
|
||||||
|
offs_d = did * BLOCK_D + tl.arange(0, BLOCK_D)
|
||||||
|
values = tl.load(
|
||||||
|
k_cache_ptr
|
||||||
|
+ physical_token[:, None] * k_cache_stride_t
|
||||||
|
+ offs_d[None, :] * k_cache_stride_d,
|
||||||
|
mask=physical_valid[:, None] & (offs_d[None, :] < KV_DIM),
|
||||||
|
other=0.0,
|
||||||
|
)
|
||||||
|
tl.store(
|
||||||
|
kv_out_ptr
|
||||||
|
+ bid * kv_out_stride_b
|
||||||
|
+ offs_t[:, None] * kv_out_stride_t
|
||||||
|
+ offs_d[None, :] * kv_out_stride_d,
|
||||||
|
values,
|
||||||
|
mask=(offs_t[:, None] < GATHER_LEN) & (offs_d[None, :] < KV_DIM),
|
||||||
|
)
|
||||||
|
tl.store(
|
||||||
|
valid_out_ptr + bid * valid_out_stride_b + offs_t * valid_out_stride_t,
|
||||||
|
physical_valid,
|
||||||
|
mask=(did == 0) & (offs_t < GATHER_LEN),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def gather_page64_kv_latent(
|
||||||
|
k_cache: torch.Tensor,
|
||||||
|
block_table: torch.Tensor,
|
||||||
|
cache_seqlens: torch.Tensor,
|
||||||
|
window_size: int,
|
||||||
|
s_q: int,
|
||||||
|
kv_cache_dim: int,
|
||||||
|
):
|
||||||
|
bs = cache_seqlens.shape[0]
|
||||||
|
assert block_table.shape[0] == bs
|
||||||
|
|
||||||
|
gather_len = ((window_size + s_q - 1 + 7) // 8) * 8
|
||||||
|
kv_latent = torch.empty(
|
||||||
|
(bs, gather_len, kv_cache_dim),
|
||||||
|
dtype=k_cache.dtype,
|
||||||
|
device=k_cache.device,
|
||||||
|
)
|
||||||
|
kv_valid = torch.empty((bs, gather_len), dtype=torch.bool, device=k_cache.device)
|
||||||
|
|
||||||
|
block_t = 8
|
||||||
|
block_d = 128
|
||||||
|
_gather_page64_kv_latent_kernel[
|
||||||
|
(bs, triton.cdiv(gather_len, block_t), triton.cdiv(kv_cache_dim, block_d))
|
||||||
|
](
|
||||||
|
k_cache,
|
||||||
|
block_table,
|
||||||
|
cache_seqlens,
|
||||||
|
kv_latent,
|
||||||
|
kv_valid,
|
||||||
|
k_cache.stride(0),
|
||||||
|
k_cache.stride(2),
|
||||||
|
block_table.stride(0),
|
||||||
|
block_table.shape[1],
|
||||||
|
k_cache.shape[0],
|
||||||
|
kv_latent.stride(0),
|
||||||
|
kv_latent.stride(1),
|
||||||
|
kv_latent.stride(2),
|
||||||
|
kv_valid.stride(0),
|
||||||
|
kv_valid.stride(1),
|
||||||
|
GATHER_LEN=gather_len,
|
||||||
|
KV_DIM=kv_cache_dim,
|
||||||
|
BLOCK_T=block_t,
|
||||||
|
BLOCK_D=block_d,
|
||||||
|
PAGE_SIZE_=PAGE_SIZE,
|
||||||
|
num_warps=4,
|
||||||
|
)
|
||||||
|
return kv_latent, kv_valid
|
||||||
|
|
||||||
|
|
||||||
|
@triton.jit
|
||||||
|
def _apply_swa_score_mask_kernel(
|
||||||
|
scores_ptr,
|
||||||
|
cache_seqlens_ptr,
|
||||||
|
valid_ptr,
|
||||||
|
scores_stride_b,
|
||||||
|
scores_stride_h,
|
||||||
|
scores_stride_q,
|
||||||
|
scores_stride_t,
|
||||||
|
valid_stride_b,
|
||||||
|
valid_stride_t,
|
||||||
|
GATHER_LEN: tl.constexpr,
|
||||||
|
WINDOW_SIZE: tl.constexpr,
|
||||||
|
S_Q: tl.constexpr,
|
||||||
|
BLOCK_T: tl.constexpr,
|
||||||
|
):
|
||||||
|
bid = tl.program_id(0)
|
||||||
|
hid = tl.program_id(1)
|
||||||
|
qid = tl.program_id(2)
|
||||||
|
|
||||||
|
offs_t = tl.arange(0, BLOCK_T)
|
||||||
|
cache_seqlen = tl.load(cache_seqlens_ptr + bid).to(tl.int32)
|
||||||
|
gather_start = tl.maximum(cache_seqlen - GATHER_LEN, 0)
|
||||||
|
kv_pos = gather_start + offs_t
|
||||||
|
q_pos = cache_seqlen - S_Q + qid
|
||||||
|
page_valid = tl.load(
|
||||||
|
valid_ptr + bid * valid_stride_b + offs_t * valid_stride_t,
|
||||||
|
mask=offs_t < GATHER_LEN,
|
||||||
|
other=0,
|
||||||
|
).to(tl.int1)
|
||||||
|
valid = (
|
||||||
|
(offs_t < GATHER_LEN)
|
||||||
|
& page_valid
|
||||||
|
& (kv_pos <= q_pos)
|
||||||
|
& (kv_pos >= q_pos - WINDOW_SIZE + 1)
|
||||||
|
& (q_pos >= 0)
|
||||||
|
)
|
||||||
|
tl.store(
|
||||||
|
scores_ptr
|
||||||
|
+ bid * scores_stride_b
|
||||||
|
+ hid * scores_stride_h
|
||||||
|
+ qid * scores_stride_q
|
||||||
|
+ offs_t * scores_stride_t,
|
||||||
|
tl.full((BLOCK_T,), -3.4028234663852886e38, tl.float32),
|
||||||
|
mask=(offs_t < GATHER_LEN) & ~valid,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def apply_swa_score_mask(
|
||||||
|
scores: torch.Tensor,
|
||||||
|
cache_seqlens: torch.Tensor,
|
||||||
|
kv_valid: torch.Tensor,
|
||||||
|
num_heads: int,
|
||||||
|
window_size: int,
|
||||||
|
s_q: int,
|
||||||
|
):
|
||||||
|
bs = cache_seqlens.shape[0]
|
||||||
|
assert scores.shape[0] == bs
|
||||||
|
assert scores.shape[1] == num_heads
|
||||||
|
assert scores.shape[2] == s_q
|
||||||
|
assert kv_valid.shape == (bs, scores.shape[3])
|
||||||
|
gather_len = scores.shape[3]
|
||||||
|
mask_block_t = triton.next_power_of_2(gather_len)
|
||||||
|
_apply_swa_score_mask_kernel[(bs, num_heads, s_q)](
|
||||||
|
scores,
|
||||||
|
cache_seqlens,
|
||||||
|
kv_valid,
|
||||||
|
scores.stride(0),
|
||||||
|
scores.stride(1),
|
||||||
|
scores.stride(2),
|
||||||
|
scores.stride(3),
|
||||||
|
kv_valid.stride(0),
|
||||||
|
kv_valid.stride(1),
|
||||||
|
GATHER_LEN=gather_len,
|
||||||
|
WINDOW_SIZE=window_size,
|
||||||
|
S_Q=s_q,
|
||||||
|
BLOCK_T=mask_block_t,
|
||||||
|
num_warps=4,
|
||||||
|
)
|
||||||
|
return scores
|
||||||
@@ -128,6 +128,8 @@ class FusedMoEMethodBase(QuantizeMethodBase):
|
|||||||
class QuantizationConfig(ABC):
|
class QuantizationConfig(ABC):
|
||||||
"""Base class for quantization configs."""
|
"""Base class for quantization configs."""
|
||||||
|
|
||||||
|
weight_block_size: Optional[List[int]] = None
|
||||||
|
|
||||||
def __init__(self):
|
def __init__(self):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
# mapping is updated by models as they initialize
|
# mapping is updated by models as they initialize
|
||||||
|
|||||||
@@ -206,6 +206,8 @@ class GenerateReqInput:
|
|||||||
] = None
|
] = None
|
||||||
# Whether to extract and process audio from video inputs.
|
# Whether to extract and process audio from video inputs.
|
||||||
use_audio_in_video: bool = False
|
use_audio_in_video: bool = False
|
||||||
|
# Optional request-scoped video processor configuration.
|
||||||
|
video_config: Optional[Dict[str, Any]] = None
|
||||||
# The sampling_params. See descriptions below.
|
# The sampling_params. See descriptions below.
|
||||||
sampling_params: Optional[Union[List[Dict[str, Any]], Dict[str, Any]]] = None
|
sampling_params: Optional[Union[List[Dict[str, Any]], Dict[str, Any]]] = None
|
||||||
# Whether to return logprobs.
|
# Whether to return logprobs.
|
||||||
|
|||||||
@@ -738,6 +738,8 @@ class KVCacheConfigurator:
|
|||||||
get_exec().kernel.attention_backend == "ascend" and not self.mambaish_config
|
get_exec().kernel.attention_backend == "ascend" and not self.mambaish_config
|
||||||
):
|
):
|
||||||
unsupported_pool_family = "NPU/Ascend KV pool"
|
unsupported_pool_family = "NPU/Ascend KV pool"
|
||||||
|
elif self.use_mla_backend and self.is_hybrid_swa:
|
||||||
|
unsupported_pool_family = "hybrid DSA/MLA-SWA KV pool"
|
||||||
elif self.use_mla_backend and is_dsa_model:
|
elif self.use_mla_backend and is_dsa_model:
|
||||||
unsupported_pool_family = "DSA/MLA KV pool"
|
unsupported_pool_family = "DSA/MLA KV pool"
|
||||||
elif self.use_mla_backend and not self.mambaish_config:
|
elif self.use_mla_backend and not self.mambaish_config:
|
||||||
@@ -1013,6 +1015,12 @@ class KVCacheConfigurator:
|
|||||||
token_to_kv_pool = self._build_ascend_mha_kv_pool(
|
token_to_kv_pool = self._build_ascend_mha_kv_pool(
|
||||||
max_total_num_tokens=sizes.max_total_num_tokens,
|
max_total_num_tokens=sizes.max_total_num_tokens,
|
||||||
)
|
)
|
||||||
|
elif self.use_mla_backend and self.is_hybrid_swa:
|
||||||
|
token_to_kv_pool = self._build_hybrid_mla_swa_kv_pool(
|
||||||
|
full_max_total_num_tokens=sizes.full_max_total_num_tokens,
|
||||||
|
swa_max_total_num_tokens=sizes.swa_max_total_num_tokens,
|
||||||
|
is_dsa_model=is_dsa_model,
|
||||||
|
)
|
||||||
elif self.use_mla_backend and is_dsa_model:
|
elif self.use_mla_backend and is_dsa_model:
|
||||||
token_to_kv_pool = self._build_dsa_kv_pool(
|
token_to_kv_pool = self._build_dsa_kv_pool(
|
||||||
max_total_num_tokens=sizes.max_total_num_tokens,
|
max_total_num_tokens=sizes.max_total_num_tokens,
|
||||||
@@ -1358,6 +1366,60 @@ class KVCacheConfigurator:
|
|||||||
)
|
)
|
||||||
return token_to_kv_pool
|
return token_to_kv_pool
|
||||||
|
|
||||||
|
def _build_hybrid_mla_swa_kv_pool(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
full_max_total_num_tokens: int,
|
||||||
|
swa_max_total_num_tokens: int,
|
||||||
|
is_dsa_model: bool,
|
||||||
|
) -> KVCache:
|
||||||
|
"""Build a hybrid MLA pool with independent full/SWA cache geometries.
|
||||||
|
|
||||||
|
Full-attention layers may use either MLA or DSA storage, while sliding
|
||||||
|
layers use MLA storage. The returned ``SWAKVPool`` exposes the common
|
||||||
|
MLA and optional DSA-index interfaces independent of model type.
|
||||||
|
"""
|
||||||
|
full_pool_class = DSATokenToKVPool if is_dsa_model else MLATokenToKVPool
|
||||||
|
common = {
|
||||||
|
"page_size": self.server_args.page_size,
|
||||||
|
"device": self.device,
|
||||||
|
"enable_memory_saver": False,
|
||||||
|
}
|
||||||
|
full_pool_kwargs = {
|
||||||
|
**common,
|
||||||
|
"kv_lora_rank": self.model_config.kv_lora_rank,
|
||||||
|
"qk_rope_head_dim": self.model_config.qk_rope_head_dim,
|
||||||
|
}
|
||||||
|
if is_dsa_model:
|
||||||
|
full_pool_kwargs.update(
|
||||||
|
index_head_dim=get_dsa_index_head_dim(self.model_config.hf_config),
|
||||||
|
kv_cache_dim=calculate_mla_kv_cache_dim(
|
||||||
|
model_config=self.model_config,
|
||||||
|
kv_cache_dtype=self.kv_cache_dtype,
|
||||||
|
server_args=self.server_args,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
return SWAKVPool(
|
||||||
|
size=full_max_total_num_tokens,
|
||||||
|
size_swa=swa_max_total_num_tokens,
|
||||||
|
page_size=self.server_args.page_size,
|
||||||
|
dtype=self.kv_cache_dtype,
|
||||||
|
head_num=0,
|
||||||
|
head_dim=0,
|
||||||
|
swa_attention_layer_ids=self.model_config.swa_attention_layer_ids,
|
||||||
|
full_attention_layer_ids=self.model_config.full_attention_layer_ids,
|
||||||
|
device=self.device,
|
||||||
|
full_kv_pool_class=full_pool_class,
|
||||||
|
swa_kv_pool_class=MLATokenToKVPool,
|
||||||
|
full_kv_pool_kwargs=full_pool_kwargs,
|
||||||
|
swa_kv_pool_kwargs={
|
||||||
|
**common,
|
||||||
|
"kv_lora_rank": self.model_config.swa_kv_lora_rank,
|
||||||
|
"qk_rope_head_dim": self.model_config.swa_qk_rope_head_dim,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
def _build_mla_fp4_kv_pool(self, *, max_total_num_tokens: int) -> KVCache:
|
def _build_mla_fp4_kv_pool(self, *, max_total_num_tokens: int) -> KVCache:
|
||||||
token_to_kv_pool = MLATokenToKVPoolFP4(
|
token_to_kv_pool = MLATokenToKVPoolFP4(
|
||||||
max_total_num_tokens,
|
max_total_num_tokens,
|
||||||
|
|||||||
@@ -4125,10 +4125,13 @@ class MLATokenToKVPool(KVCache):
|
|||||||
loc_info,
|
loc_info,
|
||||||
cache_k: torch.Tensor,
|
cache_k: torch.Tensor,
|
||||||
cache_v: torch.Tensor,
|
cache_v: torch.Tensor,
|
||||||
|
layer_id_override: Optional[int] = None,
|
||||||
):
|
):
|
||||||
loc, _, _ = unwrap_write_loc(loc_info)
|
loc, _, _ = unwrap_write_loc(loc_info)
|
||||||
maybe_detect_oob(loc, 0, self.size + self.page_size, "set_kv_buffer (MLA)")
|
maybe_detect_oob(loc, 0, self.size + self.page_size, "set_kv_buffer (MLA)")
|
||||||
layer_id = layer.layer_id
|
layer_id = (
|
||||||
|
layer_id_override if layer_id_override is not None else layer.layer_id
|
||||||
|
)
|
||||||
assert not self.dsa_kv_cache_store_fp8
|
assert not self.dsa_kv_cache_store_fp8
|
||||||
parallel = get_parallel()
|
parallel = get_parallel()
|
||||||
if parallel.dcp_enabled:
|
if parallel.dcp_enabled:
|
||||||
@@ -4201,6 +4204,7 @@ class MLATokenToKVPool(KVCache):
|
|||||||
loc: torch.Tensor,
|
loc: torch.Tensor,
|
||||||
cache_k_nope: torch.Tensor,
|
cache_k_nope: torch.Tensor,
|
||||||
cache_k_rope: torch.Tensor,
|
cache_k_rope: torch.Tensor,
|
||||||
|
layer_id_override: Optional[int] = None,
|
||||||
):
|
):
|
||||||
# loc is widened under DCP; the kernel divides by the world size itself.
|
# loc is widened under DCP; the kernel divides by the world size itself.
|
||||||
maybe_detect_oob(
|
maybe_detect_oob(
|
||||||
@@ -4209,7 +4213,9 @@ class MLATokenToKVPool(KVCache):
|
|||||||
(self.size + self.page_size) * get_parallel().attn_dcp_size,
|
(self.size + self.page_size) * get_parallel().attn_dcp_size,
|
||||||
"set_mla_kv_buffer (MLA)",
|
"set_mla_kv_buffer (MLA)",
|
||||||
)
|
)
|
||||||
layer_id = layer.layer_id
|
layer_id = (
|
||||||
|
layer_id_override if layer_id_override is not None else layer.layer_id
|
||||||
|
)
|
||||||
self._write_mla_kv_buffer(
|
self._write_mla_kv_buffer(
|
||||||
self.kv_buffer[layer_id - self.start_layer],
|
self.kv_buffer[layer_id - self.start_layer],
|
||||||
loc,
|
loc,
|
||||||
|
|||||||
@@ -6,8 +6,10 @@ import torch
|
|||||||
from sglang.srt.layers.radix_attention import RadixAttention
|
from sglang.srt.layers.radix_attention import RadixAttention
|
||||||
from sglang.srt.mem_cache.base_swa_memory_pool import BaseSWAKVPool
|
from sglang.srt.mem_cache.base_swa_memory_pool import BaseSWAKVPool
|
||||||
from sglang.srt.mem_cache.memory_pool import (
|
from sglang.srt.mem_cache.memory_pool import (
|
||||||
|
DSATokenToKVPool,
|
||||||
KVCache,
|
KVCache,
|
||||||
MHATokenToKVPool,
|
MHATokenToKVPool,
|
||||||
|
MLATokenToKVPool,
|
||||||
unwrap_write_loc,
|
unwrap_write_loc,
|
||||||
)
|
)
|
||||||
from sglang.srt.mem_cache.utils import maybe_init_custom_mem_pool
|
from sglang.srt.mem_cache.utils import maybe_init_custom_mem_pool
|
||||||
@@ -17,7 +19,12 @@ GB = 1024 * 1024 * 1024
|
|||||||
|
|
||||||
|
|
||||||
class SWAKVPool(BaseSWAKVPool):
|
class SWAKVPool(BaseSWAKVPool):
|
||||||
"""KV cache with separate pools for full and SWA attention layers."""
|
"""Hybrid full/SWA cache composed from independently configurable pools.
|
||||||
|
|
||||||
|
The default remains two MHA pools. Supplying ``full_kv_pool_class`` and
|
||||||
|
``swa_kv_pool_class`` enables other KV cache families, including MLA/DSA,
|
||||||
|
without adding model-specific behavior to the pool selector.
|
||||||
|
"""
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
@@ -31,6 +38,10 @@ class SWAKVPool(BaseSWAKVPool):
|
|||||||
full_attention_layer_ids: List[int],
|
full_attention_layer_ids: List[int],
|
||||||
device: str,
|
device: str,
|
||||||
token_to_kv_pool_class: KVCache = MHATokenToKVPool,
|
token_to_kv_pool_class: KVCache = MHATokenToKVPool,
|
||||||
|
full_kv_pool_class: Optional[type] = None,
|
||||||
|
swa_kv_pool_class: Optional[type] = None,
|
||||||
|
full_kv_pool_kwargs: Optional[dict] = None,
|
||||||
|
swa_kv_pool_kwargs: Optional[dict] = None,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
):
|
):
|
||||||
self.size = size
|
self.size = size
|
||||||
@@ -46,35 +57,58 @@ class SWAKVPool(BaseSWAKVPool):
|
|||||||
self.page_size = page_size
|
self.page_size = page_size
|
||||||
self.layer_transfer_counter = None
|
self.layer_transfer_counter = None
|
||||||
|
|
||||||
kwargs["page_size"] = page_size
|
|
||||||
kwargs["enable_memory_saver"] = False
|
|
||||||
kwargs["head_num"] = head_num
|
|
||||||
kwargs["head_dim"] = head_dim
|
|
||||||
kwargs["device"] = device
|
|
||||||
|
|
||||||
# for disagg with nvlink
|
# for disagg with nvlink
|
||||||
self.enable_custom_mem_pool, self.custom_mem_pool, _ = (
|
self.enable_custom_mem_pool, self.custom_mem_pool, _ = (
|
||||||
maybe_init_custom_mem_pool(device=self.device)
|
maybe_init_custom_mem_pool(device=self.device)
|
||||||
)
|
)
|
||||||
|
|
||||||
full_pool_kwargs = kwargs.copy()
|
full_kv_pool_class = full_kv_pool_class or token_to_kv_pool_class
|
||||||
full_pool_kwargs.pop("swa_head_num", None)
|
swa_kv_pool_class = swa_kv_pool_class or token_to_kv_pool_class
|
||||||
full_pool_kwargs.pop("swa_head_dim", None)
|
common_kwargs = {
|
||||||
full_pool_kwargs.pop("swa_v_head_dim", None)
|
"page_size": page_size,
|
||||||
self.full_kv_pool = token_to_kv_pool_class(
|
"enable_memory_saver": False,
|
||||||
|
"device": device,
|
||||||
|
}
|
||||||
|
if full_kv_pool_kwargs is None:
|
||||||
|
full_kv_pool_kwargs = {
|
||||||
|
**common_kwargs,
|
||||||
|
"head_num": head_num,
|
||||||
|
"head_dim": head_dim,
|
||||||
|
"allocation_label": "Full",
|
||||||
|
**kwargs,
|
||||||
|
}
|
||||||
|
full_kv_pool_kwargs.pop("swa_head_num", None)
|
||||||
|
full_kv_pool_kwargs.pop("swa_head_dim", None)
|
||||||
|
full_kv_pool_kwargs.pop("swa_v_head_dim", None)
|
||||||
|
if swa_kv_pool_kwargs is None:
|
||||||
|
swa_kv_pool_kwargs = {
|
||||||
|
**common_kwargs,
|
||||||
|
"head_num": head_num,
|
||||||
|
"head_dim": head_dim,
|
||||||
|
"allocation_label": "SWA",
|
||||||
|
**kwargs,
|
||||||
|
}
|
||||||
|
|
||||||
|
self.full_kv_pool = full_kv_pool_class(
|
||||||
size=size,
|
size=size,
|
||||||
dtype=dtype,
|
dtype=dtype,
|
||||||
layer_num=self.full_layer_nums,
|
layer_num=self.full_layer_nums,
|
||||||
allocation_label="Full",
|
**full_kv_pool_kwargs,
|
||||||
**full_pool_kwargs,
|
|
||||||
)
|
)
|
||||||
self.swa_kv_pool = token_to_kv_pool_class(
|
self.swa_kv_pool = swa_kv_pool_class(
|
||||||
size=size_swa,
|
size=size_swa,
|
||||||
dtype=dtype,
|
dtype=dtype,
|
||||||
layer_num=self.swa_layer_nums,
|
layer_num=self.swa_layer_nums,
|
||||||
allocation_label="SWA",
|
**swa_kv_pool_kwargs,
|
||||||
**kwargs,
|
|
||||||
)
|
)
|
||||||
|
self.dsa_kv_cache_store_fp8 = False
|
||||||
|
self.kv_cache_dim = None
|
||||||
|
self.index_head_dim = None
|
||||||
|
if isinstance(self.full_kv_pool, MLATokenToKVPool):
|
||||||
|
self.dsa_kv_cache_store_fp8 = self.full_kv_pool.dsa_kv_cache_store_fp8
|
||||||
|
self.kv_cache_dim = self.full_kv_pool.kv_cache_dim
|
||||||
|
if isinstance(self.full_kv_pool, DSATokenToKVPool):
|
||||||
|
self.index_head_dim = self.full_kv_pool.index_head_dim
|
||||||
# {layer_id: (index, is_swa_layer)}
|
# {layer_id: (index, is_swa_layer)}
|
||||||
self.layers_mapping: Dict[int, Tuple[int, bool]] = {}
|
self.layers_mapping: Dict[int, Tuple[int, bool]] = {}
|
||||||
for full_attn_layer_id, global_layer_id in enumerate(full_attention_layer_ids):
|
for full_attn_layer_id, global_layer_id in enumerate(full_attention_layer_ids):
|
||||||
@@ -123,10 +157,19 @@ class SWAKVPool(BaseSWAKVPool):
|
|||||||
self.layer_transfer_counter.wait_until(layer_id - self.start_layer)
|
self.layer_transfer_counter.wait_until(layer_id - self.start_layer)
|
||||||
|
|
||||||
def get_kv_size_bytes(self):
|
def get_kv_size_bytes(self):
|
||||||
k_size, v_size = self.full_kv_pool.get_kv_size_bytes()
|
def split_size(pool):
|
||||||
k_size_swa, v_size_swa = self.swa_kv_pool.get_kv_size_bytes()
|
size = pool.get_kv_size_bytes()
|
||||||
|
return size if isinstance(size, tuple) else (size, 0)
|
||||||
|
|
||||||
|
k_size, v_size = split_size(self.full_kv_pool)
|
||||||
|
k_size_swa, v_size_swa = split_size(self.swa_kv_pool)
|
||||||
return k_size + k_size_swa, v_size + v_size_swa
|
return k_size + k_size_swa, v_size + v_size_swa
|
||||||
|
|
||||||
|
def is_mla(self) -> bool:
|
||||||
|
return isinstance(self.full_kv_pool, MLATokenToKVPool) and isinstance(
|
||||||
|
self.swa_kv_pool, MLATokenToKVPool
|
||||||
|
)
|
||||||
|
|
||||||
def get_contiguous_buf_infos(self):
|
def get_contiguous_buf_infos(self):
|
||||||
full_kv_data_ptrs, full_kv_data_lens, full_kv_item_lens = (
|
full_kv_data_ptrs, full_kv_data_lens, full_kv_item_lens = (
|
||||||
self.full_kv_pool.get_contiguous_buf_infos()
|
self.full_kv_pool.get_contiguous_buf_infos()
|
||||||
@@ -200,21 +243,22 @@ class SWAKVPool(BaseSWAKVPool):
|
|||||||
loc, swa_loc, _ = unwrap_write_loc(loc_info)
|
loc, swa_loc, _ = unwrap_write_loc(loc_info)
|
||||||
layer_id = layer.layer_id
|
layer_id = layer.layer_id
|
||||||
layer_id_pool, is_swa_layer = self.layers_mapping[layer_id]
|
layer_id_pool, is_swa_layer = self.layers_mapping[layer_id]
|
||||||
|
pool = self.swa_kv_pool if is_swa_layer else self.full_kv_pool
|
||||||
if is_swa_layer:
|
if is_swa_layer:
|
||||||
# swa_loc is the full->SWA translation, computed once per forward by
|
# swa_loc is the full->SWA translation, computed once per forward by
|
||||||
# the attention backend; set_kv_buffer never translates internally.
|
# the attention backend; set_kv_buffer never translates internally.
|
||||||
assert swa_loc is not None
|
assert swa_loc is not None
|
||||||
self.swa_kv_pool.set_kv_buffer(
|
loc = swa_loc
|
||||||
|
if isinstance(pool, MLATokenToKVPool):
|
||||||
|
pool.set_kv_buffer(
|
||||||
None,
|
None,
|
||||||
swa_loc,
|
loc,
|
||||||
cache_k,
|
cache_k,
|
||||||
cache_v,
|
cache_v,
|
||||||
k_scale,
|
|
||||||
v_scale,
|
|
||||||
layer_id_override=layer_id_pool,
|
layer_id_override=layer_id_pool,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
self.full_kv_pool.set_kv_buffer(
|
pool.set_kv_buffer(
|
||||||
None,
|
None,
|
||||||
loc,
|
loc,
|
||||||
cache_k,
|
cache_k,
|
||||||
@@ -224,6 +268,60 @@ class SWAKVPool(BaseSWAKVPool):
|
|||||||
layer_id_override=layer_id_pool,
|
layer_id_override=layer_id_pool,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def set_mla_kv_buffer(
|
||||||
|
self,
|
||||||
|
layer: RadixAttention,
|
||||||
|
loc_info,
|
||||||
|
cache_k_nope: torch.Tensor,
|
||||||
|
cache_k_rope: torch.Tensor,
|
||||||
|
):
|
||||||
|
loc, swa_loc, _ = unwrap_write_loc(loc_info)
|
||||||
|
layer_id_pool, is_swa_layer = self.layers_mapping[layer.layer_id]
|
||||||
|
pool = self.swa_kv_pool if is_swa_layer else self.full_kv_pool
|
||||||
|
if is_swa_layer:
|
||||||
|
assert swa_loc is not None
|
||||||
|
loc = swa_loc
|
||||||
|
if not isinstance(pool, MLATokenToKVPool):
|
||||||
|
raise TypeError(f"Layer {layer.layer_id} is not backed by an MLA KV pool")
|
||||||
|
pool.set_mla_kv_buffer(
|
||||||
|
None,
|
||||||
|
loc,
|
||||||
|
cache_k_nope,
|
||||||
|
cache_k_rope,
|
||||||
|
layer_id_override=layer_id_pool,
|
||||||
|
)
|
||||||
|
|
||||||
|
def get_index_k_with_scale_buffer(self, layer_id: int) -> torch.Tensor:
|
||||||
|
layer_id_pool, is_swa_layer = self.layers_mapping[layer_id]
|
||||||
|
assert not is_swa_layer
|
||||||
|
return self.full_kv_pool.get_index_k_with_scale_buffer(layer_id_pool)
|
||||||
|
|
||||||
|
def get_index_k_continuous(self, layer_id: int, *args, **kwargs):
|
||||||
|
layer_id_pool, is_swa_layer = self.layers_mapping[layer_id]
|
||||||
|
assert not is_swa_layer
|
||||||
|
return self.full_kv_pool.get_index_k_continuous(layer_id_pool, *args, **kwargs)
|
||||||
|
|
||||||
|
def get_index_k_scale_continuous(self, layer_id: int, *args, **kwargs):
|
||||||
|
layer_id_pool, is_swa_layer = self.layers_mapping[layer_id]
|
||||||
|
assert not is_swa_layer
|
||||||
|
return self.full_kv_pool.get_index_k_scale_continuous(
|
||||||
|
layer_id_pool, *args, **kwargs
|
||||||
|
)
|
||||||
|
|
||||||
|
def get_index_k_scale_buffer(self, layer_id: int, *args, **kwargs):
|
||||||
|
layer_id_pool, is_swa_layer = self.layers_mapping[layer_id]
|
||||||
|
assert not is_swa_layer
|
||||||
|
return self.full_kv_pool.get_index_k_scale_buffer(
|
||||||
|
layer_id_pool, *args, **kwargs
|
||||||
|
)
|
||||||
|
|
||||||
|
def set_index_k_scale_buffer(self, layer_id: int, *args, **kwargs):
|
||||||
|
layer_id_pool, is_swa_layer = self.layers_mapping[layer_id]
|
||||||
|
assert not is_swa_layer
|
||||||
|
return self.full_kv_pool.set_index_k_scale_buffer(
|
||||||
|
layer_id_pool, *args, **kwargs
|
||||||
|
)
|
||||||
|
|
||||||
def move_kv_cache(self, tgt_loc: torch.Tensor, src_loc: torch.Tensor):
|
def move_kv_cache(self, tgt_loc: torch.Tensor, src_loc: torch.Tensor):
|
||||||
self.full_kv_pool.move_kv_cache(tgt_loc, src_loc)
|
self.full_kv_pool.move_kv_cache(tgt_loc, src_loc)
|
||||||
tgt_loc_swa = self.translate_loc_from_full_to_swa(tgt_loc)
|
tgt_loc_swa = self.translate_loc_from_full_to_swa(tgt_loc)
|
||||||
|
|||||||
@@ -128,11 +128,19 @@ class ForwardBatchDeepSeekMHAMixin:
|
|||||||
HybridLinearKVPool,
|
HybridLinearKVPool,
|
||||||
MLATokenToKVPool,
|
MLATokenToKVPool,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool
|
||||||
|
|
||||||
token_to_kv_pool = get_token_to_kv_pool()
|
token_to_kv_pool = get_token_to_kv_pool()
|
||||||
assert isinstance(token_to_kv_pool, MLATokenToKVPool) or (
|
assert (
|
||||||
|
isinstance(token_to_kv_pool, MLATokenToKVPool)
|
||||||
|
or (
|
||||||
isinstance(token_to_kv_pool, HybridLinearKVPool)
|
isinstance(token_to_kv_pool, HybridLinearKVPool)
|
||||||
and isinstance(token_to_kv_pool.full_kv_pool, MLATokenToKVPool)
|
and isinstance(token_to_kv_pool.full_kv_pool, MLATokenToKVPool)
|
||||||
|
)
|
||||||
|
or (
|
||||||
|
isinstance(token_to_kv_pool, SWAKVPool)
|
||||||
|
and isinstance(token_to_kv_pool.full_kv_pool, MLATokenToKVPool)
|
||||||
|
)
|
||||||
), "Currently chunked prefix cache can only be used by Deepseek models"
|
), "Currently chunked prefix cache can only be used by Deepseek models"
|
||||||
|
|
||||||
if not any(self.extend_prefix_lens_cpu):
|
if not any(self.extend_prefix_lens_cpu):
|
||||||
|
|||||||
@@ -514,7 +514,6 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
|
|||||||
# Has to be None when cuda graph is captured.
|
# Has to be None when cuda graph is captured.
|
||||||
global_num_tokens_for_logprob_cpu: Optional[List[int]] = None
|
global_num_tokens_for_logprob_cpu: Optional[List[int]] = None
|
||||||
global_num_tokens_for_logprob_gpu: Optional[torch.Tensor] = None
|
global_num_tokens_for_logprob_gpu: Optional[torch.Tensor] = None
|
||||||
|
|
||||||
# For padding
|
# For padding
|
||||||
num_token_non_padded: Optional[torch.Tensor] = None # scalar tensor
|
num_token_non_padded: Optional[torch.Tensor] = None # scalar tensor
|
||||||
num_token_non_padded_cpu: int = None
|
num_token_non_padded_cpu: int = None
|
||||||
@@ -1535,9 +1534,35 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
|
|||||||
dim=1,
|
dim=1,
|
||||||
)
|
)
|
||||||
|
|
||||||
# TODO: check if we need to pad other tensors
|
# Draft-extend padding uses the fixed per-request token width.
|
||||||
|
dummy_extend_len = 0
|
||||||
|
if (
|
||||||
|
self.spec_info is not None
|
||||||
|
and self.forward_mode.is_draft_extend_v2()
|
||||||
|
and self.spec_info.num_tokens_per_req > 0
|
||||||
|
):
|
||||||
|
dummy_extend_len = self.spec_info.num_tokens_per_req
|
||||||
|
|
||||||
if self.extend_seq_lens is not None:
|
if self.extend_seq_lens is not None:
|
||||||
self.extend_seq_lens = self._pad_tensor_to_size(self.extend_seq_lens, bs)
|
self.extend_seq_lens = self._pad_tensor_to_size(
|
||||||
|
self.extend_seq_lens, bs, value=dummy_extend_len
|
||||||
|
)
|
||||||
|
if self.extend_prefix_lens is not None:
|
||||||
|
self.extend_prefix_lens = self._pad_tensor_to_size(
|
||||||
|
self.extend_prefix_lens, bs
|
||||||
|
)
|
||||||
|
if self.extend_seq_lens_cpu is not None:
|
||||||
|
self.extend_seq_lens_cpu.extend(
|
||||||
|
[dummy_extend_len] * (bs - len(self.extend_seq_lens_cpu))
|
||||||
|
)
|
||||||
|
if self.extend_prefix_lens_cpu is not None:
|
||||||
|
self.extend_prefix_lens_cpu.extend(
|
||||||
|
[0] * (bs - len(self.extend_prefix_lens_cpu))
|
||||||
|
)
|
||||||
|
if self.extend_logprob_start_lens_cpu is not None:
|
||||||
|
self.extend_logprob_start_lens_cpu.extend(
|
||||||
|
[0] * (bs - len(self.extend_logprob_start_lens_cpu))
|
||||||
|
)
|
||||||
|
|
||||||
if self.rids_int is not None:
|
if self.rids_int is not None:
|
||||||
self.rids_int = self._pad_tensor_to_size(self.rids_int, bs)
|
self.rids_int = self._pad_tensor_to_size(self.rids_int, bs)
|
||||||
|
|||||||
@@ -31,6 +31,8 @@ def _map_muse_target_layer_ids(*, target_hf_config, draft_hf_config, layer_ids):
|
|||||||
class SpecAuxHiddenStateConfig(msgspec.Struct, kw_only=True):
|
class SpecAuxHiddenStateConfig(msgspec.Struct, kw_only=True):
|
||||||
eagle_use_aux_hidden_state: bool = False
|
eagle_use_aux_hidden_state: bool = False
|
||||||
eagle_draft_num_layers: Optional[int] = None
|
eagle_draft_num_layers: Optional[int] = None
|
||||||
|
# Draft layers whose KV cache uses the target SWA pool capacity.
|
||||||
|
eagle_draft_swa_num_layers: Optional[int] = None
|
||||||
eagle_aux_hidden_state_layer_ids: Any = None
|
eagle_aux_hidden_state_layer_ids: Any = None
|
||||||
dflash_use_aux_hidden_state: bool = False
|
dflash_use_aux_hidden_state: bool = False
|
||||||
dflash_draft_num_layers: Optional[int] = None
|
dflash_draft_num_layers: Optional[int] = None
|
||||||
@@ -93,6 +95,14 @@ def _resolve_eagle_aux_hidden_state(
|
|||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
|
if (
|
||||||
|
draft_model_config.is_hybrid_swa
|
||||||
|
and not draft_model_config.is_deepseek_v4_arch
|
||||||
|
):
|
||||||
|
config.eagle_draft_swa_num_layers = len(
|
||||||
|
draft_model_config.swa_attention_layer_ids
|
||||||
|
)
|
||||||
|
|
||||||
if spec_algorithm.is_eagle3():
|
if spec_algorithm.is_eagle3():
|
||||||
config.eagle_use_aux_hidden_state = True
|
config.eagle_use_aux_hidden_state = True
|
||||||
try:
|
try:
|
||||||
|
|||||||
@@ -21,6 +21,7 @@ import torch
|
|||||||
|
|
||||||
from sglang.srt.configs.hybrid_arch import mambaish_config
|
from sglang.srt.configs.hybrid_arch import mambaish_config
|
||||||
from sglang.srt.configs.model_config import (
|
from sglang.srt.configs.model_config import (
|
||||||
|
AttentionArch,
|
||||||
dsa_layer_skips_topk,
|
dsa_layer_skips_topk,
|
||||||
get_dsa_index_head_dim,
|
get_dsa_index_head_dim,
|
||||||
get_minimax_sparse_attention_config,
|
get_minimax_sparse_attention_config,
|
||||||
@@ -431,7 +432,7 @@ class DefaultPoolConfigurator(MemoryPoolConfigurator):
|
|||||||
|
|
||||||
|
|
||||||
class HybridSWAPoolConfigurator(MemoryPoolConfigurator):
|
class HybridSWAPoolConfigurator(MemoryPoolConfigurator):
|
||||||
"""Configurator for hybrid sliding window attention models (Gemma2, Command-R, MiMo).
|
"""Configurator for MHA or MLA models with sliding-window layers.
|
||||||
|
|
||||||
Splits available memory between full attention and SWA pools.
|
Splits available memory between full attention and SWA pools.
|
||||||
Does NOT inherit DefaultPoolConfigurator — different coeff model.
|
Does NOT inherit DefaultPoolConfigurator — different coeff model.
|
||||||
@@ -454,6 +455,33 @@ class HybridSWAPoolConfigurator(MemoryPoolConfigurator):
|
|||||||
self._sliding_window_size = kvc.sliding_window_size
|
self._sliding_window_size = kvc.sliding_window_size
|
||||||
self._page_size = kvc.page_size
|
self._page_size = kvc.page_size
|
||||||
|
|
||||||
|
if model_config.attention_arch == AttentionArch.MLA:
|
||||||
|
# MLA pool sizing uses latent dimensions rather than MHA heads.
|
||||||
|
from sglang.srt.mem_cache.kv_cache_configurator import (
|
||||||
|
calculate_mla_kv_cache_dim,
|
||||||
|
)
|
||||||
|
|
||||||
|
self._full_per_token = (
|
||||||
|
calculate_mla_kv_cache_dim(
|
||||||
|
model_config=model_config,
|
||||||
|
kv_cache_dtype=kv_cache_dtype,
|
||||||
|
server_args=kvc.server_args,
|
||||||
|
)
|
||||||
|
* kv_size
|
||||||
|
)
|
||||||
|
if is_deepseek_dsa(model_config.hf_config):
|
||||||
|
index_head_dim = get_dsa_index_head_dim(model_config.hf_config)
|
||||||
|
index_elements = (
|
||||||
|
index_head_dim
|
||||||
|
+ index_head_dim // DSATokenToKVPool.quant_block_size * 4
|
||||||
|
)
|
||||||
|
self._full_per_token += index_elements * torch._utils._element_size(
|
||||||
|
DSATokenToKVPool.index_k_with_scale_buffer_dtype
|
||||||
|
)
|
||||||
|
self._swa_per_token = (
|
||||||
|
model_config.swa_kv_lora_rank + model_config.swa_qk_rope_head_dim
|
||||||
|
) * kv_size
|
||||||
|
else:
|
||||||
# Full layer per-token memory (bytes)
|
# Full layer per-token memory (bytes)
|
||||||
self._full_per_token = (
|
self._full_per_token = (
|
||||||
model_config.get_num_kv_heads(tp_size)
|
model_config.get_num_kv_heads(tp_size)
|
||||||
@@ -479,11 +507,9 @@ class HybridSWAPoolConfigurator(MemoryPoolConfigurator):
|
|||||||
* (model_config.swa_head_dim + model_config.swa_v_head_dim)
|
* (model_config.swa_head_dim + model_config.swa_v_head_dim)
|
||||||
) // scale_block_size
|
) // scale_block_size
|
||||||
|
|
||||||
# EAGLE/STANDALONE draft KV pool inherits max_total tokens with its
|
# Draft KV tensors use full, SWA, or full-capacity SWA geometry.
|
||||||
# full-attn layers; budget into the full term. A banded MTP depth
|
|
||||||
# (Inkling mtp_local_layer_ids) instead allocates an swa-geometry ring
|
|
||||||
# at FULL draft capacity, so budget those depths at swa_per_token.
|
|
||||||
self._draft_full_layers_num = 0
|
self._draft_full_layers_num = 0
|
||||||
|
self._draft_swa_layers_num = 0
|
||||||
self._draft_swa_full_layers_num = 0
|
self._draft_swa_full_layers_num = 0
|
||||||
if (
|
if (
|
||||||
kvc.spec_algorithm.is_eagle() or kvc.spec_algorithm.is_standalone()
|
kvc.spec_algorithm.is_eagle() or kvc.spec_algorithm.is_standalone()
|
||||||
@@ -504,7 +530,17 @@ class HybridSWAPoolConfigurator(MemoryPoolConfigurator):
|
|||||||
]
|
]
|
||||||
)
|
)
|
||||||
self._draft_swa_full_layers_num = banded_depths
|
self._draft_swa_full_layers_num = banded_depths
|
||||||
self._draft_full_layers_num = draft_layers - banded_depths
|
else:
|
||||||
|
draft_swa_layers = kvc.spec_aux_config.eagle_draft_swa_num_layers
|
||||||
|
if draft_swa_layers is not None:
|
||||||
|
self._draft_swa_layers_num = min(
|
||||||
|
max(int(draft_swa_layers), 0), draft_layers
|
||||||
|
)
|
||||||
|
self._draft_full_layers_num = (
|
||||||
|
draft_layers
|
||||||
|
- self._draft_swa_layers_num
|
||||||
|
- self._draft_swa_full_layers_num
|
||||||
|
)
|
||||||
|
|
||||||
self._draft_cell_size = _dflash_draft_cell_size(kvc)
|
self._draft_cell_size = _dflash_draft_cell_size(kvc)
|
||||||
|
|
||||||
@@ -521,6 +557,7 @@ class HybridSWAPoolConfigurator(MemoryPoolConfigurator):
|
|||||||
self._cell_size = (
|
self._cell_size = (
|
||||||
self._swa_per_token * self._swa_layers_num
|
self._swa_per_token * self._swa_layers_num
|
||||||
+ self._full_per_token * self._draft_full_layers_num
|
+ self._full_per_token * self._draft_full_layers_num
|
||||||
|
+ self._swa_per_token * self._draft_swa_layers_num
|
||||||
+ self._swa_per_token * self._draft_swa_full_layers_num
|
+ self._swa_per_token * self._draft_swa_full_layers_num
|
||||||
+ self._draft_cell_size
|
+ self._draft_cell_size
|
||||||
)
|
)
|
||||||
@@ -531,7 +568,7 @@ class HybridSWAPoolConfigurator(MemoryPoolConfigurator):
|
|||||||
+ self._swa_per_token * self._draft_swa_full_layers_num
|
+ self._swa_per_token * self._draft_swa_full_layers_num
|
||||||
+ self._swa_full_tokens_ratio
|
+ self._swa_full_tokens_ratio
|
||||||
* self._swa_per_token
|
* self._swa_per_token
|
||||||
* self._swa_layers_num
|
* (self._swa_layers_num + self._draft_swa_layers_num)
|
||||||
+ self._draft_cell_size
|
+ self._draft_cell_size
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -667,7 +704,11 @@ class SWAChunkCapPoolConfigurator(HybridSWAPoolConfigurator):
|
|||||||
) -> MemoryPoolConfig:
|
) -> MemoryPoolConfig:
|
||||||
# SWA pool sized tightly from the cap; the rest of the budget goes to full.
|
# SWA pool sized tightly from the cap; the rest of the budget goes to full.
|
||||||
swa_tokens = ceil_align(self._swa_cap, page_size)
|
swa_tokens = ceil_align(self._swa_cap, page_size)
|
||||||
fixed_swa_bytes = swa_tokens * self._swa_per_token * self._swa_layers_num
|
fixed_swa_bytes = (
|
||||||
|
swa_tokens
|
||||||
|
* self._swa_per_token
|
||||||
|
* (self._swa_layers_num + self._draft_swa_layers_num)
|
||||||
|
)
|
||||||
full_cell_size = (
|
full_cell_size = (
|
||||||
self._full_per_token * (self._full_layers_num + self._draft_full_layers_num)
|
self._full_per_token * (self._full_layers_num + self._draft_full_layers_num)
|
||||||
+ self._swa_per_token * self._draft_swa_full_layers_num
|
+ self._swa_per_token * self._draft_swa_full_layers_num
|
||||||
|
|||||||
@@ -78,6 +78,14 @@ def _clone_if_runai_streamed_tensor(tensor: torch.Tensor) -> torch.Tensor:
|
|||||||
return tensor
|
return tensor
|
||||||
|
|
||||||
|
|
||||||
|
def _get_indexer_weight_block_size(
|
||||||
|
quant_config: Optional[QuantizationConfig],
|
||||||
|
) -> List[int]:
|
||||||
|
if quant_config is not None and quant_config.weight_block_size is not None:
|
||||||
|
return quant_config.weight_block_size
|
||||||
|
return [128, 128]
|
||||||
|
|
||||||
|
|
||||||
def _load_fused_indexer_wk(
|
def _load_fused_indexer_wk(
|
||||||
name: str,
|
name: str,
|
||||||
loaded_weight: torch.Tensor,
|
loaded_weight: torch.Tensor,
|
||||||
@@ -99,10 +107,25 @@ def _load_fused_indexer_wk(
|
|||||||
return False
|
return False
|
||||||
|
|
||||||
if ".indexer.weights_proj." in name:
|
if ".indexer.weights_proj." in name:
|
||||||
|
is_scale = name.endswith(".weight_scale_inv")
|
||||||
|
if not is_scale and loaded_weight.dtype != torch.float8_e4m3fn:
|
||||||
w = _clone_if_runai_streamed_tensor(loaded_weight)
|
w = _clone_if_runai_streamed_tensor(loaded_weight)
|
||||||
fused_param.data[-w.shape[0] :].copy_(w)
|
fused_param.data[-w.shape[0] :].copy_(w)
|
||||||
return True
|
return True
|
||||||
|
|
||||||
|
entry = pending.setdefault(fused_name + ".weights_proj", {})
|
||||||
|
entry["scale" if is_scale else "weight"] = _clone_if_runai_streamed_tensor(
|
||||||
|
loaded_weight
|
||||||
|
)
|
||||||
|
if "weight" in entry and "scale" in entry:
|
||||||
|
pending.pop(fused_name + ".weights_proj")
|
||||||
|
block_size = _get_indexer_weight_block_size(quant_config)
|
||||||
|
weights_bf16 = block_quant_dequant(
|
||||||
|
entry["weight"], entry["scale"], block_size, torch.bfloat16
|
||||||
|
)
|
||||||
|
fused_param.data[-weights_bf16.shape[0] :].copy_(weights_bf16)
|
||||||
|
return True
|
||||||
|
|
||||||
# wk: a bf16 checkpoint copies straight in; block-fp8 needs weight + scale.
|
# wk: a bf16 checkpoint copies straight in; block-fp8 needs weight + scale.
|
||||||
is_scale = name.endswith(".weight_scale_inv")
|
is_scale = name.endswith(".weight_scale_inv")
|
||||||
if not is_scale and loaded_weight.dtype != torch.float8_e4m3fn:
|
if not is_scale and loaded_weight.dtype != torch.float8_e4m3fn:
|
||||||
@@ -116,7 +139,7 @@ def _load_fused_indexer_wk(
|
|||||||
)
|
)
|
||||||
if "weight" in entry and "scale" in entry:
|
if "weight" in entry and "scale" in entry:
|
||||||
pending.pop(fused_name)
|
pending.pop(fused_name)
|
||||||
block_size = getattr(quant_config, "weight_block_size", None) or [128, 128]
|
block_size = _get_indexer_weight_block_size(quant_config)
|
||||||
wk_bf16 = block_quant_dequant(
|
wk_bf16 = block_quant_dequant(
|
||||||
entry["weight"], entry["scale"], block_size, torch.bfloat16
|
entry["weight"], entry["scale"], block_size, torch.bfloat16
|
||||||
)
|
)
|
||||||
@@ -547,8 +570,10 @@ class DeepseekV2WeightLoaderMixin:
|
|||||||
)
|
)
|
||||||
if selected_quant_config is None:
|
if selected_quant_config is None:
|
||||||
selected_quant_config = self.quant_config
|
selected_quant_config = self.quant_config
|
||||||
weight_block_size = getattr(
|
weight_block_size = (
|
||||||
selected_quant_config, "weight_block_size", None
|
selected_quant_config.weight_block_size
|
||||||
|
if selected_quant_config is not None
|
||||||
|
else None
|
||||||
)
|
)
|
||||||
if weight_block_size is not None:
|
if weight_block_size is not None:
|
||||||
assert hasattr(self_attn.kv_b_proj, "weight_scale_inv") or hasattr(
|
assert hasattr(self_attn.kv_b_proj, "weight_scale_inv") or hasattr(
|
||||||
@@ -571,8 +596,10 @@ class DeepseekV2WeightLoaderMixin:
|
|||||||
# In multiple weight loading scenarios (e.g. RL), we need to inverse the scale of the weights after the requantization happened at the first loading.
|
# In multiple weight loading scenarios (e.g. RL), we need to inverse the scale of the weights after the requantization happened at the first loading.
|
||||||
if (
|
if (
|
||||||
should_deepgemm_weight_requant_ue8m0(
|
should_deepgemm_weight_requant_ue8m0(
|
||||||
weight_block_size=getattr(
|
weight_block_size=(
|
||||||
self.quant_config, "weight_block_size", None
|
self.quant_config.weight_block_size
|
||||||
|
if self.quant_config is not None
|
||||||
|
else None
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
and weight_scale.format_ue8m0
|
and weight_scale.format_ue8m0
|
||||||
@@ -624,16 +651,19 @@ class DeepseekV2WeightLoaderMixin:
|
|||||||
self_attn.w_scale = scale
|
self_attn.w_scale = scale
|
||||||
|
|
||||||
if w.dtype == torch.int8:
|
if w.dtype == torch.int8:
|
||||||
if hasattr(self.quant_config, "weight_block_size"):
|
weight_block_size = (
|
||||||
# block-wise int8 need it
|
self.quant_config.weight_block_size
|
||||||
weight_block_size = self.quant_config.weight_block_size
|
if self.quant_config is not None
|
||||||
|
else None
|
||||||
|
)
|
||||||
if weight_block_size is not None:
|
if weight_block_size is not None:
|
||||||
|
# block-wise int8 need it
|
||||||
assert hasattr(self_attn.kv_b_proj, "weight_scale_inv")
|
assert hasattr(self_attn.kv_b_proj, "weight_scale_inv")
|
||||||
weight = w
|
weight = w
|
||||||
weight_scale = self_attn.kv_b_proj.weight_scale_inv
|
weight_scale = self_attn.kv_b_proj.weight_scale_inv
|
||||||
w = int8_block_dequant(
|
w = int8_block_dequant(weight, weight_scale, weight_block_size).to(
|
||||||
weight, weight_scale, weight_block_size
|
torch.bfloat16
|
||||||
).to(torch.bfloat16)
|
)
|
||||||
else:
|
else:
|
||||||
# channel-wise int8 need it
|
# channel-wise int8 need it
|
||||||
w = w.to(torch.bfloat16) * self_attn.kv_b_proj.weight_scale.to(
|
w = w.to(torch.bfloat16) * self_attn.kv_b_proj.weight_scale.to(
|
||||||
|
|||||||
@@ -0,0 +1,37 @@
|
|||||||
|
# SPDX-License-Identifier: Apache-2.0
|
||||||
|
# Copyright 2023-2024 SGLang Team
|
||||||
|
|
||||||
|
"""Registry entry point for the Dots3 model."""
|
||||||
|
|
||||||
|
from sglang.srt.models.dots3_common.modeling import (
|
||||||
|
Dots3AttentionMLA,
|
||||||
|
Dots3AttnForwardMethod,
|
||||||
|
Dots3DecoderLayer,
|
||||||
|
Dots3LanguageModelForCausalLM,
|
||||||
|
Dots3MLP,
|
||||||
|
Dots3Model,
|
||||||
|
Dots3MoE,
|
||||||
|
Dots3MoEGate,
|
||||||
|
Dots3NoteForCausalLM,
|
||||||
|
DotsNoteOmniForConditionalGeneration,
|
||||||
|
DotsNoteOmniThinkerForConditionalGeneration,
|
||||||
|
get_attention_sliding_window_size,
|
||||||
|
)
|
||||||
|
|
||||||
|
EntryClass = [Dots3NoteForCausalLM]
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
"Dots3AttentionMLA",
|
||||||
|
"Dots3AttnForwardMethod",
|
||||||
|
"Dots3DecoderLayer",
|
||||||
|
"Dots3LanguageModelForCausalLM",
|
||||||
|
"Dots3MLP",
|
||||||
|
"Dots3Model",
|
||||||
|
"Dots3MoE",
|
||||||
|
"Dots3MoEGate",
|
||||||
|
"Dots3NoteForCausalLM",
|
||||||
|
"DotsNoteOmniForConditionalGeneration",
|
||||||
|
"DotsNoteOmniThinkerForConditionalGeneration",
|
||||||
|
"EntryClass",
|
||||||
|
"get_attention_sliding_window_size",
|
||||||
|
]
|
||||||
@@ -0,0 +1 @@
|
|||||||
|
"""Shared implementation modules for Dots3 models."""
|
||||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,240 @@
|
|||||||
|
"""In-process vision/audio towers for dots.note.omni."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import json
|
||||||
|
import math
|
||||||
|
from collections.abc import Iterable
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
import numpy as np
|
||||||
|
import torch
|
||||||
|
from PIL import Image
|
||||||
|
|
||||||
|
from sglang.srt.models.dots3_common.dots_omni_audio import (
|
||||||
|
OmniAudioConfig,
|
||||||
|
OmniAudioModel,
|
||||||
|
compute_audio_token_length,
|
||||||
|
)
|
||||||
|
from sglang.srt.models.dots3_common.dots_omni_vision import (
|
||||||
|
DotsMoEVitConfig,
|
||||||
|
DotsMoEVitModel,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _read_json(path: Path) -> dict:
|
||||||
|
with path.open() as file:
|
||||||
|
return json.load(file)
|
||||||
|
|
||||||
|
|
||||||
|
def load_omni_component_config(model_dir: Path, component: str) -> dict:
|
||||||
|
"""Read a tower config from a flat dots.note.omni publish."""
|
||||||
|
config = _read_json(model_dir / "config.json")
|
||||||
|
nested_name = f"{component}_config"
|
||||||
|
if nested_name not in config:
|
||||||
|
raise KeyError(f"Missing {nested_name!r} in {model_dir / 'config.json'}")
|
||||||
|
return config[nested_name]
|
||||||
|
|
||||||
|
|
||||||
|
class DotsNoteOmniVisionEncoder(DotsMoEVitModel):
|
||||||
|
"""Native MoE ViT used by dots.note.omni."""
|
||||||
|
|
||||||
|
def __init__(self, model_dir: str):
|
||||||
|
model_dir = Path(model_dir)
|
||||||
|
config = DotsMoEVitConfig(**load_omni_component_config(model_dir, "vision"))
|
||||||
|
super().__init__(config)
|
||||||
|
self.to(torch.bfloat16)
|
||||||
|
|
||||||
|
def load_converted_state(self, state: dict[str, torch.Tensor]):
|
||||||
|
missing, unexpected = self.load_state_dict(state, strict=False)
|
||||||
|
if missing:
|
||||||
|
raise RuntimeError(f"Dots vision tower missing weights: {missing[:8]}")
|
||||||
|
if unexpected:
|
||||||
|
raise RuntimeError(
|
||||||
|
f"Dots vision tower has unexpected weights: {unexpected[:8]}"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class DotsNoteOmniAudioEncoder(OmniAudioModel):
|
||||||
|
"""Native Dots speech encoder and adapter."""
|
||||||
|
|
||||||
|
def __init__(self, model_dir: str):
|
||||||
|
model_dir = Path(model_dir)
|
||||||
|
config = OmniAudioConfig(**load_omni_component_config(model_dir, "audio"))
|
||||||
|
super().__init__(config)
|
||||||
|
self.to(torch.bfloat16)
|
||||||
|
|
||||||
|
@property
|
||||||
|
def dtype(self):
|
||||||
|
return next(self.parameters()).dtype
|
||||||
|
|
||||||
|
def load_converted_state(self, state: dict[str, torch.Tensor]):
|
||||||
|
missing, unexpected = self.load_state_dict(state, strict=True)
|
||||||
|
if missing or unexpected:
|
||||||
|
raise RuntimeError(
|
||||||
|
"Dots audio tower weight mismatch: "
|
||||||
|
f"missing={missing[:8]}, unexpected={unexpected[:8]}"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class DotsNoteOmniImagePreprocessor:
|
||||||
|
"""CPU image preprocessing matching the converted native ViT."""
|
||||||
|
|
||||||
|
def __init__(self, model_dir: str):
|
||||||
|
model_dir = Path(model_dir)
|
||||||
|
config = _read_json(model_dir / "preprocessor_config.json")
|
||||||
|
config = config["vision_config"]
|
||||||
|
self.min_pixels = config["min_pixels"]
|
||||||
|
self.max_pixels = config["max_pixels"]
|
||||||
|
self.patch_size = config["patch_size"]
|
||||||
|
self.temporal_patch_size = config["temporal_patch_size"]
|
||||||
|
self.merge_size = config["merge_size"]
|
||||||
|
self.pre_pixel_shuffle = config.get("pre_pixel_shuffle", True)
|
||||||
|
self.image_mean = np.asarray(config["image_mean"], dtype=np.float32)
|
||||||
|
self.image_std = np.asarray(config["image_std"], dtype=np.float32)
|
||||||
|
image_detail_path = model_dir / "image_detail.json"
|
||||||
|
self.image_detail_config = (
|
||||||
|
_read_json(image_detail_path).get("image_details", {})
|
||||||
|
if image_detail_path.is_file()
|
||||||
|
else {}
|
||||||
|
)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _round_by_factor(value: int, factor: int) -> int:
|
||||||
|
return round(value / factor) * factor
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _ceil_by_factor(value: float, factor: int) -> int:
|
||||||
|
return math.ceil(value / factor) * factor
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _floor_by_factor(value: float, factor: int) -> int:
|
||||||
|
return math.floor(value / factor) * factor
|
||||||
|
|
||||||
|
def _resized_size(
|
||||||
|
self,
|
||||||
|
width: int,
|
||||||
|
height: int,
|
||||||
|
min_pixels: int,
|
||||||
|
max_pixels: int,
|
||||||
|
target_height=None,
|
||||||
|
target_width=None,
|
||||||
|
):
|
||||||
|
height = target_height or height
|
||||||
|
width = target_width or width
|
||||||
|
factor = self.patch_size * self.merge_size
|
||||||
|
if min(height, width) < factor // 4:
|
||||||
|
raise ValueError(
|
||||||
|
f"Image height/width must be at least {factor // 4}, "
|
||||||
|
f"got {height}x{width}"
|
||||||
|
)
|
||||||
|
if max(height, width) / min(height, width) > 200:
|
||||||
|
raise ValueError("Image aspect ratio must be smaller than 200")
|
||||||
|
resized_h = max(factor, self._round_by_factor(height, factor))
|
||||||
|
resized_w = max(factor, self._round_by_factor(width, factor))
|
||||||
|
if resized_h * resized_w > max_pixels:
|
||||||
|
beta = math.sqrt(height * width / max_pixels)
|
||||||
|
resized_h = max(factor, self._floor_by_factor(height / beta, factor))
|
||||||
|
resized_w = max(factor, self._floor_by_factor(width / beta, factor))
|
||||||
|
elif resized_h * resized_w < min_pixels:
|
||||||
|
beta = math.sqrt(min_pixels / (height * width))
|
||||||
|
resized_h = self._ceil_by_factor(height * beta, factor)
|
||||||
|
resized_w = self._ceil_by_factor(width * beta, factor)
|
||||||
|
if resized_h * resized_w > max_pixels:
|
||||||
|
beta = math.sqrt(resized_h * resized_w / max_pixels)
|
||||||
|
resized_h = max(factor, self._floor_by_factor(resized_h / beta, factor))
|
||||||
|
resized_w = max(factor, self._floor_by_factor(resized_w / beta, factor))
|
||||||
|
return resized_h, resized_w
|
||||||
|
|
||||||
|
def _process_image(self, image, detail="auto"):
|
||||||
|
if not isinstance(image, Image.Image):
|
||||||
|
raise TypeError(f"Expected a PIL image, got {type(image)}")
|
||||||
|
if image.mode == "RGBA":
|
||||||
|
background = Image.new("RGB", image.size, (255, 255, 255))
|
||||||
|
background.paste(image, mask=image.getchannel("A"))
|
||||||
|
image = background
|
||||||
|
elif image.mode != "RGB":
|
||||||
|
image = image.convert("RGB")
|
||||||
|
|
||||||
|
detail_config = self.image_detail_config.get(detail, {})
|
||||||
|
resized_h, resized_w = self._resized_size(
|
||||||
|
*image.size,
|
||||||
|
min_pixels=detail_config.get("min_pixels", self.min_pixels),
|
||||||
|
max_pixels=detail_config.get("max_pixels", self.max_pixels),
|
||||||
|
target_height=detail_config.get("target_height"),
|
||||||
|
target_width=detail_config.get("target_width"),
|
||||||
|
)
|
||||||
|
image = image.resize((resized_w, resized_h), Image.Resampling.BICUBIC)
|
||||||
|
array = np.asarray(image, dtype=np.float32) / 255.0
|
||||||
|
array = (array - self.image_mean) / self.image_std
|
||||||
|
patches = array.transpose(2, 0, 1)[None]
|
||||||
|
if patches.shape[0] == 1:
|
||||||
|
patches = np.tile(patches, (self.temporal_patch_size, 1, 1, 1))
|
||||||
|
channel = patches.shape[1]
|
||||||
|
grid_t = patches.shape[0] // self.temporal_patch_size
|
||||||
|
grid_h = resized_h // self.patch_size
|
||||||
|
grid_w = resized_w // self.patch_size
|
||||||
|
if self.pre_pixel_shuffle:
|
||||||
|
patches = patches.reshape(
|
||||||
|
grid_t,
|
||||||
|
self.temporal_patch_size,
|
||||||
|
channel,
|
||||||
|
grid_h // self.merge_size,
|
||||||
|
self.merge_size,
|
||||||
|
self.patch_size,
|
||||||
|
grid_w // self.merge_size,
|
||||||
|
self.merge_size,
|
||||||
|
self.patch_size,
|
||||||
|
)
|
||||||
|
patches = patches.transpose(0, 3, 6, 4, 7, 2, 1, 5, 8)
|
||||||
|
else:
|
||||||
|
patches = patches.reshape(
|
||||||
|
grid_t,
|
||||||
|
self.temporal_patch_size,
|
||||||
|
channel,
|
||||||
|
grid_h,
|
||||||
|
self.patch_size,
|
||||||
|
grid_w,
|
||||||
|
self.patch_size,
|
||||||
|
)
|
||||||
|
patches = patches.transpose(0, 3, 5, 2, 1, 4, 6)
|
||||||
|
pixel_values = torch.from_numpy(
|
||||||
|
patches.reshape(
|
||||||
|
grid_t * grid_h * grid_w,
|
||||||
|
channel * self.temporal_patch_size * self.patch_size * self.patch_size,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
return {
|
||||||
|
"pixel_values": pixel_values,
|
||||||
|
"image_grid_thw": torch.tensor([[grid_t, grid_h, grid_w]]),
|
||||||
|
}
|
||||||
|
|
||||||
|
def _get_image_token_str(self, token_count: int):
|
||||||
|
return "<|img|>" + "<|imgpad|>" * token_count + "<|endofimg|>"
|
||||||
|
|
||||||
|
def process_images(self, images: Iterable, details=None):
|
||||||
|
images = list(images)
|
||||||
|
details = details or ["auto"] * len(images)
|
||||||
|
pixel_values = []
|
||||||
|
grids = []
|
||||||
|
token_strings = []
|
||||||
|
for image, detail in zip(images, details):
|
||||||
|
processed = self._process_image(image, detail)
|
||||||
|
grid = processed["image_grid_thw"]
|
||||||
|
token_count = int(grid.prod().item()) // self.merge_size**2
|
||||||
|
pixel_values.append(processed["pixel_values"])
|
||||||
|
grids.append(grid)
|
||||||
|
token_strings.append(self._get_image_token_str(token_count))
|
||||||
|
return pixel_values, grids, token_strings
|
||||||
|
|
||||||
|
|
||||||
|
def get_audio_token_string(num_samples: int, config: OmniAudioConfig) -> str:
|
||||||
|
count = compute_audio_token_length(
|
||||||
|
num_samples,
|
||||||
|
chunk_seconds=config.chunk_seconds,
|
||||||
|
conv_temporal_stride=config.conv_temporal_stride,
|
||||||
|
merge_factor=config.merge_factor,
|
||||||
|
)
|
||||||
|
return (
|
||||||
|
config.audio_comp_start + config.audio_comp_span * count + config.audio_comp_end
|
||||||
|
)
|
||||||
@@ -0,0 +1,769 @@
|
|||||||
|
import math
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
import torch
|
||||||
|
import torch.nn.functional as F
|
||||||
|
from torch import nn
|
||||||
|
from torch.nn import LayerNorm
|
||||||
|
from transformers.configuration_utils import PretrainedConfig
|
||||||
|
from transformers.modeling_utils import PreTrainedModel
|
||||||
|
|
||||||
|
from sglang.srt.layers.activation import SiluAndMul
|
||||||
|
from sglang.srt.layers.attention.vision import VisionAttention as SGLVisionAttention
|
||||||
|
from sglang.srt.layers.conv import Conv2dLayer
|
||||||
|
from sglang.srt.layers.layernorm import RMSNorm
|
||||||
|
|
||||||
|
|
||||||
|
class VisionRotaryEmbedding(nn.Module):
|
||||||
|
"""2D vision RoPE frequency table with optional caching."""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
dim: int,
|
||||||
|
theta: float = 10000.0,
|
||||||
|
cache_seq_len: int | None = None,
|
||||||
|
) -> None:
|
||||||
|
super().__init__()
|
||||||
|
inv_freq = 1.0 / (theta ** (torch.arange(0, dim, 2, dtype=torch.float32) / dim))
|
||||||
|
self.register_buffer("inv_freq", inv_freq, persistent=False)
|
||||||
|
self._cache_seq_len = cache_seq_len
|
||||||
|
if cache_seq_len is not None:
|
||||||
|
self.register_buffer(
|
||||||
|
"freqs_cache", self._compute_freqs(cache_seq_len), persistent=False
|
||||||
|
)
|
||||||
|
|
||||||
|
def _compute_freqs(self, seqlen: int) -> torch.Tensor:
|
||||||
|
seq = torch.arange(
|
||||||
|
seqlen, device=self.inv_freq.device, dtype=self.inv_freq.dtype
|
||||||
|
)
|
||||||
|
return torch.outer(seq, self.inv_freq)
|
||||||
|
|
||||||
|
def forward(self, seqlen: int) -> torch.Tensor:
|
||||||
|
if self._cache_seq_len is None:
|
||||||
|
return self._compute_freqs(seqlen)
|
||||||
|
if seqlen > self.freqs_cache.shape[0]:
|
||||||
|
self.freqs_cache = self._compute_freqs(seqlen)
|
||||||
|
return self.freqs_cache[:seqlen]
|
||||||
|
|
||||||
|
|
||||||
|
class VisionAttention(SGLVisionAttention):
|
||||||
|
"""Dots checkpoint compatibility wrapper around SGLang vision attention."""
|
||||||
|
|
||||||
|
def __init__(self, config: Any) -> None:
|
||||||
|
dim = config.embed_dim
|
||||||
|
super().__init__(
|
||||||
|
embed_dim=dim,
|
||||||
|
num_heads=config.num_attention_heads,
|
||||||
|
projection_size=dim,
|
||||||
|
use_qkv_parallel=True,
|
||||||
|
flatten_batch=True,
|
||||||
|
use_data_parallel=True,
|
||||||
|
qkv_bias=config.use_bias,
|
||||||
|
proj_bias=config.use_bias,
|
||||||
|
qk_normalization_by_head_size=config.use_qk_norm,
|
||||||
|
layer_norm_eps=config.rms_norm_eps,
|
||||||
|
)
|
||||||
|
self.register_load_state_dict_pre_hook(VisionAttention._map_qkv_weight)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _map_qkv_weight(
|
||||||
|
module: "VisionAttention",
|
||||||
|
state_dict: dict[str, torch.Tensor],
|
||||||
|
prefix: str,
|
||||||
|
*args,
|
||||||
|
) -> None:
|
||||||
|
for suffix in ("weight", "bias"):
|
||||||
|
source = f"{prefix}qkv.{suffix}"
|
||||||
|
target = f"{prefix}qkv_proj.{suffix}"
|
||||||
|
if source in state_dict and target not in state_dict:
|
||||||
|
state_dict[target] = state_dict.pop(source)
|
||||||
|
|
||||||
|
def forward(
|
||||||
|
self,
|
||||||
|
hidden_states: torch.Tensor,
|
||||||
|
cu_seqlens: torch.Tensor,
|
||||||
|
max_seqlen: int,
|
||||||
|
rotary_pos_emb: torch.Tensor,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
output = super().forward(
|
||||||
|
hidden_states,
|
||||||
|
cu_seqlens=cu_seqlens,
|
||||||
|
position_embeddings=(rotary_pos_emb.cos(), rotary_pos_emb.sin()),
|
||||||
|
max_seqlen=max_seqlen,
|
||||||
|
)
|
||||||
|
return output.squeeze(0)
|
||||||
|
|
||||||
|
|
||||||
|
class DotsMoEVitConfig(PretrainedConfig):
|
||||||
|
model_type: str = "dots_moe_vit"
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
embed_dim: int = 1536,
|
||||||
|
hidden_size: int = 2048,
|
||||||
|
intermediate_size: int = 4224,
|
||||||
|
moe_intermediate_size: int = 2112,
|
||||||
|
num_hidden_layers: int = 42,
|
||||||
|
num_attention_heads: int = 24,
|
||||||
|
num_channels: int = 3,
|
||||||
|
patch_size: int = 14,
|
||||||
|
spatial_merge_size: int = 2,
|
||||||
|
temporal_patch_size: int = 1,
|
||||||
|
rms_norm_eps: float = 1e-5,
|
||||||
|
use_bias: bool = False,
|
||||||
|
use_qk_norm: bool = True,
|
||||||
|
attn_implementation="flash_attention_3",
|
||||||
|
initializer_range=0.02,
|
||||||
|
is_causal=False,
|
||||||
|
post_norm=True,
|
||||||
|
gradient_checkpointing=False,
|
||||||
|
pyramid_num_routed: list[int] | None = None,
|
||||||
|
capacity_factor: float = 2.0,
|
||||||
|
router_scoring_func: str = "sigmoid",
|
||||||
|
router_scale: float = 1.0,
|
||||||
|
adapter_in_dim: int = 1536,
|
||||||
|
adapter_out_dim: int = 2048,
|
||||||
|
adapter_merge_size: int = 2,
|
||||||
|
# Adapter used to merge each 2x2 patch group.
|
||||||
|
adapter_type: str = "pixel_shuffle_mlp",
|
||||||
|
# Whether input patches and RoPE positions are already 2x2-grouped.
|
||||||
|
pre_pixel_shuffle: bool = False,
|
||||||
|
# If True, use FP8 MoE implementation
|
||||||
|
enable_fp8_moe: bool = True,
|
||||||
|
**kwargs: Any,
|
||||||
|
):
|
||||||
|
super().__init__(**kwargs)
|
||||||
|
self.embed_dim = embed_dim
|
||||||
|
self.hidden_size = hidden_size
|
||||||
|
self.intermediate_size = intermediate_size
|
||||||
|
self.moe_intermediate_size = moe_intermediate_size
|
||||||
|
self.num_hidden_layers = num_hidden_layers
|
||||||
|
self.num_attention_heads = num_attention_heads
|
||||||
|
self.num_channels = num_channels
|
||||||
|
self.patch_size = patch_size
|
||||||
|
self.spatial_merge_size = spatial_merge_size
|
||||||
|
self.temporal_patch_size = temporal_patch_size
|
||||||
|
self.rms_norm_eps = rms_norm_eps
|
||||||
|
self.use_bias = use_bias
|
||||||
|
self.use_qk_norm = use_qk_norm
|
||||||
|
self.attn_implementation = attn_implementation
|
||||||
|
self.initializer_range = initializer_range
|
||||||
|
self.is_causal = is_causal
|
||||||
|
self.post_norm = post_norm
|
||||||
|
self.gradient_checkpointing = gradient_checkpointing
|
||||||
|
self.pyramid_num_routed = pyramid_num_routed or []
|
||||||
|
self.capacity_factor = capacity_factor
|
||||||
|
self.router_scoring_func = router_scoring_func
|
||||||
|
self.router_scale = router_scale
|
||||||
|
self.adapter_in_dim = adapter_in_dim
|
||||||
|
self.adapter_out_dim = adapter_out_dim
|
||||||
|
self.adapter_merge_size = adapter_merge_size
|
||||||
|
if adapter_type not in ("pixel_shuffle_mlp", "patch_merger"):
|
||||||
|
raise ValueError(
|
||||||
|
f"adapter_type must be 'pixel_shuffle_mlp' or 'patch_merger', got {adapter_type!r}"
|
||||||
|
)
|
||||||
|
self.adapter_type = adapter_type
|
||||||
|
self.pre_pixel_shuffle = pre_pixel_shuffle
|
||||||
|
self.enable_fp8_moe = enable_fp8_moe
|
||||||
|
|
||||||
|
|
||||||
|
# ---- FFN modules ----
|
||||||
|
|
||||||
|
|
||||||
|
class DotsSwiGLUFFN(nn.Module):
|
||||||
|
def __init__(self, in_features, hidden_features, bias=False):
|
||||||
|
super().__init__()
|
||||||
|
self.fc13 = nn.Linear(in_features, hidden_features * 2, bias=bias)
|
||||||
|
self.fc2 = nn.Linear(hidden_features, in_features, bias=bias)
|
||||||
|
self.act = SiluAndMul()
|
||||||
|
self.register_load_state_dict_pre_hook(
|
||||||
|
DotsSwiGLUFFN._load_fused_fc13_from_split
|
||||||
|
)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _load_fused_fc13_from_split(
|
||||||
|
module: "DotsSwiGLUFFN",
|
||||||
|
state_dict: dict[str, torch.Tensor],
|
||||||
|
prefix: str,
|
||||||
|
local_metadata: dict[str, Any],
|
||||||
|
strict: bool,
|
||||||
|
missing_keys: list[str],
|
||||||
|
unexpected_keys: list[str],
|
||||||
|
error_msgs: list[str],
|
||||||
|
) -> None:
|
||||||
|
fc13_weight_key = prefix + "fc13.weight"
|
||||||
|
fc1_weight_key = prefix + "fc1.weight"
|
||||||
|
fc3_weight_key = prefix + "fc3.weight"
|
||||||
|
fc1_weight = state_dict.get(fc1_weight_key)
|
||||||
|
fc3_weight = state_dict.get(fc3_weight_key)
|
||||||
|
if fc1_weight is not None and fc3_weight is not None:
|
||||||
|
if fc13_weight_key not in state_dict:
|
||||||
|
state_dict[fc13_weight_key] = torch.cat((fc1_weight, fc3_weight), dim=0)
|
||||||
|
state_dict.pop(fc1_weight_key)
|
||||||
|
state_dict.pop(fc3_weight_key)
|
||||||
|
|
||||||
|
fc13_bias_key = prefix + "fc13.bias"
|
||||||
|
fc1_bias_key = prefix + "fc1.bias"
|
||||||
|
fc3_bias_key = prefix + "fc3.bias"
|
||||||
|
fc1_bias = state_dict.get(fc1_bias_key)
|
||||||
|
fc3_bias = state_dict.get(fc3_bias_key)
|
||||||
|
if fc1_bias is not None and fc3_bias is not None:
|
||||||
|
if module.fc13.bias is not None and fc13_bias_key not in state_dict:
|
||||||
|
state_dict[fc13_bias_key] = torch.cat((fc1_bias, fc3_bias), dim=0)
|
||||||
|
state_dict.pop(fc1_bias_key)
|
||||||
|
state_dict.pop(fc3_bias_key)
|
||||||
|
|
||||||
|
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||||
|
return self.fc2(self.act(self.fc13(x)))
|
||||||
|
|
||||||
|
|
||||||
|
def _ceil_to_multiple(v: int, multiple: int) -> int:
|
||||||
|
return ((v + multiple - 1) // multiple) * multiple
|
||||||
|
|
||||||
|
|
||||||
|
def _per_block_cast_to_fp8_padded(
|
||||||
|
x: torch.Tensor,
|
||||||
|
*,
|
||||||
|
use_ue8m0: bool = False,
|
||||||
|
gran_k: int = 128,
|
||||||
|
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||||
|
"""`per_block_cast_to_fp8` with zero-padding for non-divisible 2D tensors.
|
||||||
|
|
||||||
|
DeepGEMM block FP8 path expects block-aligned dimensions. When `x.shape` is
|
||||||
|
not divisible by `gran_k`, this helper pads zeros on both dims to the nearest
|
||||||
|
multiple and then calls `per_block_cast_to_fp8`.
|
||||||
|
"""
|
||||||
|
if x.dim() != 2:
|
||||||
|
raise ValueError(f"expected 2D tensor, got shape={tuple(x.shape)}")
|
||||||
|
if gran_k <= 0:
|
||||||
|
raise ValueError(f"gran_k must be positive, got {gran_k}")
|
||||||
|
|
||||||
|
from deep_gemm import per_block_cast_to_fp8
|
||||||
|
|
||||||
|
m, n = int(x.shape[0]), int(x.shape[1])
|
||||||
|
m_pad = _ceil_to_multiple(m, gran_k)
|
||||||
|
n_pad = _ceil_to_multiple(n, gran_k)
|
||||||
|
|
||||||
|
if m_pad == m and n_pad == n:
|
||||||
|
return per_block_cast_to_fp8(x.contiguous(), use_ue8m0=use_ue8m0, gran_k=gran_k)
|
||||||
|
|
||||||
|
x_pad = torch.zeros((m_pad, n_pad), dtype=x.dtype, device=x.device)
|
||||||
|
x_pad[:m, :n] = x
|
||||||
|
return per_block_cast_to_fp8(x_pad.contiguous(), use_ue8m0=use_ue8m0, gran_k=gran_k)
|
||||||
|
|
||||||
|
|
||||||
|
class MoESwiGLUFFN(nn.Module):
|
||||||
|
"""MoE FFN with per-expert SwiGLU experts, sigmoid/softmax gating, top-k routing."""
|
||||||
|
|
||||||
|
def __init__(self, config: DotsMoEVitConfig, layer_number: int):
|
||||||
|
super().__init__()
|
||||||
|
self.config = config
|
||||||
|
self.layer_number = layer_number
|
||||||
|
self.hidden_size = config.embed_dim
|
||||||
|
self.num_routed = config.pyramid_num_routed[layer_number]
|
||||||
|
self.capacity_factor = config.capacity_factor
|
||||||
|
self.router_scoring_func = config.router_scoring_func
|
||||||
|
self.router_scale = config.router_scale
|
||||||
|
|
||||||
|
self.register_buffer(
|
||||||
|
"router_bias", torch.zeros(self.num_routed, dtype=torch.float32)
|
||||||
|
)
|
||||||
|
|
||||||
|
self.experts = nn.ModuleList(
|
||||||
|
[
|
||||||
|
DotsSwiGLUFFN(
|
||||||
|
self.hidden_size, config.moe_intermediate_size, bias=config.use_bias
|
||||||
|
)
|
||||||
|
for _ in range(self.num_routed)
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
self.gate_weight = nn.Parameter(
|
||||||
|
torch.empty((self.num_routed, self.hidden_size), dtype=torch.float32)
|
||||||
|
)
|
||||||
|
nn.init.kaiming_uniform_(self.gate_weight, a=math.sqrt(5))
|
||||||
|
|
||||||
|
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||||
|
# Keep routing and top-k selection in FP32 to avoid BF16 ties.
|
||||||
|
epsilon = 1e-9
|
||||||
|
x_flat = x.contiguous()
|
||||||
|
num_tokens = x_flat.shape[0]
|
||||||
|
|
||||||
|
gate_logits = F.linear(x_flat.float(), self.gate_weight.float())
|
||||||
|
|
||||||
|
if self.router_scoring_func == "sigmoid":
|
||||||
|
gating_prob = torch.sigmoid(gate_logits)
|
||||||
|
else:
|
||||||
|
gating_prob = torch.softmax(gate_logits, dim=-1, dtype=torch.float32)
|
||||||
|
|
||||||
|
aggregated_output = torch.zeros_like(x_flat)
|
||||||
|
aggregated_gate = torch.zeros(num_tokens, dtype=x.dtype, device=x.device)
|
||||||
|
|
||||||
|
topk = min(int(self.capacity_factor), self.num_routed)
|
||||||
|
|
||||||
|
gating_with_bias = gating_prob + self.router_bias.to(torch.float32).unsqueeze(0)
|
||||||
|
_, topk_indices = torch.topk(gating_with_bias, k=topk, dim=-1, sorted=False)
|
||||||
|
|
||||||
|
routed_weights = gating_prob.gather(1, topk_indices)
|
||||||
|
if self.router_scoring_func == "sigmoid" and topk > 1:
|
||||||
|
routed_weights = routed_weights / (
|
||||||
|
routed_weights.sum(dim=-1, keepdim=True) + epsilon
|
||||||
|
)
|
||||||
|
routed_weights = (routed_weights * self.router_scale).to(x_flat.dtype)
|
||||||
|
|
||||||
|
for expert_idx in range(self.num_routed):
|
||||||
|
selected_mask = topk_indices == expert_idx
|
||||||
|
if selected_mask.sum() == 0:
|
||||||
|
continue
|
||||||
|
n_idx, top = torch.where(selected_mask)
|
||||||
|
# Fancy indexing can yield non-contiguous rows; cuBLAS bf16 GEMM may then fail
|
||||||
|
# with ``CUBLAS_STATUS_INVALID_VALUE`` inside ``F.linear``.
|
||||||
|
x_selected = x_flat[n_idx].contiguous()
|
||||||
|
expert_output = self.experts[expert_idx](x_selected)
|
||||||
|
contrib = expert_output * routed_weights[n_idx, top].unsqueeze(-1)
|
||||||
|
aggregated_output[n_idx] = aggregated_output[n_idx] + contrib
|
||||||
|
aggregated_gate[n_idx] = aggregated_gate[n_idx] + routed_weights[n_idx, top]
|
||||||
|
|
||||||
|
aggregated_output = aggregated_output / (
|
||||||
|
aggregated_gate.unsqueeze(-1) + epsilon
|
||||||
|
)
|
||||||
|
return aggregated_output
|
||||||
|
|
||||||
|
|
||||||
|
class MoESwiGLUFFNFP8(MoESwiGLUFFN):
|
||||||
|
"""FP8 variant of :class:`MoESwiGLUFFN` using fused expert kernels."""
|
||||||
|
|
||||||
|
def __init__(self, config: DotsMoEVitConfig, layer_number: int):
|
||||||
|
super().__init__(config, layer_number)
|
||||||
|
# ``MoESwiGLUFFN`` already builds ``DotsSwiGLUFFN`` experts.
|
||||||
|
from sglang.srt.layers.moe.moe_runner import MoeRunnerConfig
|
||||||
|
|
||||||
|
self._moe_runner_config = MoeRunnerConfig(inplace=False)
|
||||||
|
self.register_buffer("_fused_w13_fp8", None, persistent=False)
|
||||||
|
self.register_buffer("_fused_w13_scale", None, persistent=False)
|
||||||
|
self.register_buffer("_fused_w2_fp8", None, persistent=False)
|
||||||
|
self.register_buffer("_fused_w2_scale", None, persistent=False)
|
||||||
|
self.register_load_state_dict_post_hook(
|
||||||
|
MoESwiGLUFFNFP8._post_load_pack_fused_fp8
|
||||||
|
)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _post_load_pack_fused_fp8(
|
||||||
|
module: "MoESwiGLUFFNFP8", _incompatible_keys
|
||||||
|
) -> None:
|
||||||
|
module._pack_fused_fp8_weights()
|
||||||
|
|
||||||
|
@torch.no_grad()
|
||||||
|
def _pack_fused_fp8_weights(self) -> None:
|
||||||
|
"""Stack gate+up per expert, block-quantize with DeepGEMM (128×128), layout for ``fused_moe``."""
|
||||||
|
e_list = list(self.experts)
|
||||||
|
if not e_list:
|
||||||
|
return
|
||||||
|
w13_chunks: list[torch.Tensor] = []
|
||||||
|
s13_chunks: list[torch.Tensor] = []
|
||||||
|
w2_chunks: list[torch.Tensor] = []
|
||||||
|
s2_chunks: list[torch.Tensor] = []
|
||||||
|
for ex in e_list:
|
||||||
|
# Block-quantize gate and up separately, then stack for fused MoE w13.
|
||||||
|
w1_weight, w3_weight = ex.fc13.weight.detach().chunk(2, dim=0)
|
||||||
|
w1_bf16 = w1_weight.to(torch.bfloat16)
|
||||||
|
w3_bf16 = w3_weight.to(torch.bfloat16)
|
||||||
|
q1, s1 = _per_block_cast_to_fp8_padded(w1_bf16, use_ue8m0=False, gran_k=128)
|
||||||
|
q3, s3 = _per_block_cast_to_fp8_padded(w3_bf16, use_ue8m0=False, gran_k=128)
|
||||||
|
w13_fp8 = torch.cat([q1, q3], dim=0).contiguous()
|
||||||
|
s13 = torch.cat([s1, s3], dim=0).contiguous()
|
||||||
|
w13_chunks.append(w13_fp8)
|
||||||
|
s13_chunks.append(s13)
|
||||||
|
|
||||||
|
w2_bf16 = ex.fc2.weight.detach().to(torch.bfloat16)
|
||||||
|
q2, s2 = _per_block_cast_to_fp8_padded(w2_bf16, use_ue8m0=False, gran_k=128)
|
||||||
|
w2_chunks.append(q2.contiguous())
|
||||||
|
s2_chunks.append(s2)
|
||||||
|
|
||||||
|
self._fused_w13_fp8 = torch.stack(w13_chunks, dim=0).contiguous()
|
||||||
|
self._fused_w13_scale = torch.stack(s13_chunks, dim=0).contiguous()
|
||||||
|
self._fused_w2_fp8 = torch.stack(w2_chunks, dim=0).contiguous()
|
||||||
|
self._fused_w2_scale = torch.stack(s2_chunks, dim=0).contiguous()
|
||||||
|
|
||||||
|
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||||
|
from sglang.srt.layers.moe.moe_runner.triton_utils.fused_moe import fused_moe
|
||||||
|
from sglang.srt.layers.moe.topk import StandardTopKOutput
|
||||||
|
|
||||||
|
# Keep routing decisions in FP32 to avoid BF16 top-k ties.
|
||||||
|
epsilon = 1e-9
|
||||||
|
x_flat = x.contiguous()
|
||||||
|
|
||||||
|
gate_logits = F.linear(x_flat.float(), self.gate_weight.float())
|
||||||
|
|
||||||
|
if self.router_scoring_func == "sigmoid":
|
||||||
|
gating_prob = torch.sigmoid(gate_logits)
|
||||||
|
else:
|
||||||
|
gating_prob = torch.softmax(gate_logits, dim=-1, dtype=torch.float32)
|
||||||
|
|
||||||
|
topk = min(int(self.capacity_factor), self.num_routed)
|
||||||
|
gating_with_bias = gating_prob + self.router_bias.to(torch.float32).unsqueeze(0)
|
||||||
|
|
||||||
|
_, topk_indices = torch.topk(gating_with_bias, k=topk, dim=-1, sorted=False)
|
||||||
|
|
||||||
|
routed_weights = gating_prob.gather(1, topk_indices)
|
||||||
|
if self.router_scoring_func == "sigmoid" and topk > 1:
|
||||||
|
routed_weights = routed_weights / (
|
||||||
|
routed_weights.sum(dim=-1, keepdim=True) + epsilon
|
||||||
|
)
|
||||||
|
routed_weights = routed_weights * float(self.router_scale)
|
||||||
|
|
||||||
|
topk_ids = topk_indices.to(torch.int32)
|
||||||
|
topk_output = StandardTopKOutput(routed_weights, topk_ids, gate_logits)
|
||||||
|
|
||||||
|
if self._fused_w13_fp8 is None:
|
||||||
|
self._pack_fused_fp8_weights()
|
||||||
|
|
||||||
|
b1 = b2 = None
|
||||||
|
if self.config.use_bias:
|
||||||
|
b1_list = []
|
||||||
|
b2_list = []
|
||||||
|
for ex in self.experts:
|
||||||
|
b1_list.append(ex.fc13.bias.detach().to(x.dtype))
|
||||||
|
b2_list.append(ex.fc2.bias.detach().to(x.dtype))
|
||||||
|
b1 = torch.stack(b1_list, dim=0).contiguous()
|
||||||
|
b2 = torch.stack(b2_list, dim=0).contiguous()
|
||||||
|
|
||||||
|
fused_out = fused_moe(
|
||||||
|
x_flat,
|
||||||
|
self._fused_w13_fp8,
|
||||||
|
self._fused_w2_fp8,
|
||||||
|
topk_output,
|
||||||
|
moe_runner_config=self._moe_runner_config,
|
||||||
|
b1=b1,
|
||||||
|
b2=b2,
|
||||||
|
use_fp8_w8a8=True,
|
||||||
|
w1_scale=self._fused_w13_scale,
|
||||||
|
w2_scale=self._fused_w2_scale,
|
||||||
|
block_shape=[128, 128],
|
||||||
|
)
|
||||||
|
denom = routed_weights.sum(dim=-1, keepdim=True).clamp_min(epsilon)
|
||||||
|
return (fused_out / denom).type_as(x)
|
||||||
|
|
||||||
|
|
||||||
|
# ---- PatchEmbed ----
|
||||||
|
|
||||||
|
|
||||||
|
class DotsPatchEmbed(nn.Module):
|
||||||
|
def __init__(self, config: DotsMoEVitConfig):
|
||||||
|
super().__init__()
|
||||||
|
self.num_channels = config.num_channels
|
||||||
|
self.patch_size = config.patch_size
|
||||||
|
self.temporal_patch_size = config.temporal_patch_size
|
||||||
|
self.embed_dim = config.embed_dim
|
||||||
|
self.proj = Conv2dLayer(
|
||||||
|
config.num_channels,
|
||||||
|
config.embed_dim,
|
||||||
|
kernel_size=(config.patch_size, config.patch_size),
|
||||||
|
stride=(config.patch_size, config.patch_size),
|
||||||
|
)
|
||||||
|
self.norm = RMSNorm(config.embed_dim, eps=config.rms_norm_eps)
|
||||||
|
|
||||||
|
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||||
|
x = x.view(
|
||||||
|
-1,
|
||||||
|
self.num_channels,
|
||||||
|
self.temporal_patch_size,
|
||||||
|
self.patch_size,
|
||||||
|
self.patch_size,
|
||||||
|
)[:, :, 0]
|
||||||
|
x = self.proj(x).view(-1, self.embed_dim)
|
||||||
|
x = self.norm(x)
|
||||||
|
return x
|
||||||
|
|
||||||
|
|
||||||
|
# ---- Block ----
|
||||||
|
|
||||||
|
|
||||||
|
class MoEVisionBlock(nn.Module):
|
||||||
|
def __init__(self, config: DotsMoEVitConfig, layer_number: int):
|
||||||
|
super().__init__()
|
||||||
|
self.attn = VisionAttention(config)
|
||||||
|
self.norm_1 = RMSNorm(config.embed_dim, eps=config.rms_norm_eps)
|
||||||
|
self.norm_2 = RMSNorm(config.embed_dim, eps=config.rms_norm_eps)
|
||||||
|
|
||||||
|
is_moe = (
|
||||||
|
config.pyramid_num_routed
|
||||||
|
and layer_number < len(config.pyramid_num_routed)
|
||||||
|
and config.pyramid_num_routed[layer_number] > 0
|
||||||
|
)
|
||||||
|
if is_moe and config.enable_fp8_moe:
|
||||||
|
self.mlp = MoESwiGLUFFNFP8(config, layer_number)
|
||||||
|
elif is_moe:
|
||||||
|
self.mlp = MoESwiGLUFFN(config, layer_number)
|
||||||
|
else:
|
||||||
|
self.mlp = DotsSwiGLUFFN(
|
||||||
|
config.embed_dim, config.intermediate_size, bias=config.use_bias
|
||||||
|
)
|
||||||
|
|
||||||
|
def forward(
|
||||||
|
self,
|
||||||
|
hidden_states,
|
||||||
|
cu_seqlens,
|
||||||
|
rotary_pos_emb,
|
||||||
|
max_seqlen: int,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
hidden_states = hidden_states + self.attn(
|
||||||
|
self.norm_1(hidden_states), cu_seqlens, max_seqlen, rotary_pos_emb
|
||||||
|
)
|
||||||
|
hidden_states = hidden_states + self.mlp(self.norm_2(hidden_states))
|
||||||
|
return hidden_states
|
||||||
|
|
||||||
|
|
||||||
|
# ---- Adapter (pixel_shuffle + MLP) ----
|
||||||
|
|
||||||
|
|
||||||
|
def _pixel_shuffle(x, scale_factor=0.5):
|
||||||
|
if x.size(1) % 2 == 1:
|
||||||
|
x = torch.cat([x[:, :1], x], dim=1)
|
||||||
|
if x.size(2) % 2 == 1:
|
||||||
|
x = torch.cat([x[:, :, :1], x], dim=2)
|
||||||
|
n, h, w, c = x.size()
|
||||||
|
x = x.reshape(n, h, int(w * scale_factor), int(c / scale_factor))
|
||||||
|
x = x.permute(0, 2, 1, 3).contiguous()
|
||||||
|
x = x.reshape(
|
||||||
|
n,
|
||||||
|
int(w * scale_factor),
|
||||||
|
int(h * scale_factor),
|
||||||
|
int(c / (scale_factor * scale_factor)),
|
||||||
|
)
|
||||||
|
x = x.permute(0, 2, 1, 3).contiguous()
|
||||||
|
return x
|
||||||
|
|
||||||
|
|
||||||
|
class PixelShuffleAdapter(nn.Module):
|
||||||
|
"""Legacy adapter: NHWC pixel-shuffle spatial merge + LayerNorm + 2-layer MLP.
|
||||||
|
|
||||||
|
Mirrors ``cybertron`` ``FCAdapter(pool_kind='pixel_shuffle', proj_kind='mlp2x_ln_gelu')``.
|
||||||
|
State-dict keys: ``proj.0`` (LayerNorm of in_dim*merge**2), ``proj.1`` / ``proj.3`` (Linear).
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, config: DotsMoEVitConfig):
|
||||||
|
super().__init__()
|
||||||
|
in_dim = config.adapter_in_dim
|
||||||
|
out_dim = config.adapter_out_dim
|
||||||
|
merge_size = config.adapter_merge_size
|
||||||
|
merged_dim = in_dim * merge_size**2
|
||||||
|
self.proj = nn.Sequential(
|
||||||
|
LayerNorm(merged_dim),
|
||||||
|
nn.Linear(merged_dim, out_dim),
|
||||||
|
nn.GELU(),
|
||||||
|
nn.Linear(out_dim, out_dim),
|
||||||
|
)
|
||||||
|
|
||||||
|
def forward(
|
||||||
|
self,
|
||||||
|
patch_embed: torch.Tensor,
|
||||||
|
grid_thw: torch.Tensor,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
assert patch_embed.dim() == 2 and grid_thw is not None
|
||||||
|
image_features = []
|
||||||
|
token_index = 0
|
||||||
|
for i in range(grid_thw.shape[0]):
|
||||||
|
grid_t, grid_h, grid_w = grid_thw[i]
|
||||||
|
images_token_length = grid_t * grid_h * grid_w
|
||||||
|
_pe = patch_embed[token_index : token_index + images_token_length]
|
||||||
|
token_index += images_token_length
|
||||||
|
if grid_t == 1:
|
||||||
|
_pe = _pe.reshape(int(grid_h), int(grid_w), -1).unsqueeze(0)
|
||||||
|
else:
|
||||||
|
_pe = _pe.reshape(int(grid_t), int(grid_h), int(grid_w), -1)
|
||||||
|
_pe = _pixel_shuffle(_pe, scale_factor=0.5)
|
||||||
|
if grid_t == 1:
|
||||||
|
_pe = _pe.squeeze(0)
|
||||||
|
else:
|
||||||
|
_pe = _pe.reshape(-1, _pe.shape[-1])
|
||||||
|
image_features.append(_pe.reshape(-1, _pe.shape[-1]))
|
||||||
|
out = torch.cat(image_features, dim=0)
|
||||||
|
out = self.proj(out)
|
||||||
|
return out
|
||||||
|
|
||||||
|
|
||||||
|
class PatchMergerAdapter(nn.Module):
|
||||||
|
"""Cybertron ``PatchMerger`` (``pool_kind='patch_merger', proj_kind='identity'``).
|
||||||
|
|
||||||
|
Assumes the encoder output is already laid out in ``merge_size``x``merge_size`` groups
|
||||||
|
(qwen ``pre_pixel_shuffle`` preprocessor + RoPE grouped accordingly), so merging is a
|
||||||
|
simple ``view(-1, merge**2 * in_dim)`` of consecutive tokens. State-dict layout matches
|
||||||
|
cybertron's ``PatchMerger`` (``ln_q`` over the per-token dim, ``mlp.0`` / ``mlp.2`` Linear).
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, config: DotsMoEVitConfig):
|
||||||
|
super().__init__()
|
||||||
|
in_dim = config.adapter_in_dim
|
||||||
|
out_dim = config.adapter_out_dim
|
||||||
|
merge_size = config.adapter_merge_size
|
||||||
|
merged_dim = in_dim * merge_size**2
|
||||||
|
self.merge_size = merge_size
|
||||||
|
self.merged_dim = merged_dim
|
||||||
|
self.ln_q = LayerNorm(in_dim, eps=1e-6)
|
||||||
|
self.mlp = nn.Sequential(
|
||||||
|
nn.Linear(merged_dim, merged_dim),
|
||||||
|
nn.GELU(),
|
||||||
|
nn.Linear(merged_dim, out_dim),
|
||||||
|
)
|
||||||
|
|
||||||
|
def forward(
|
||||||
|
self,
|
||||||
|
patch_embed: torch.Tensor,
|
||||||
|
grid_thw: torch.Tensor,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
assert patch_embed.dim() == 2 and grid_thw is not None
|
||||||
|
x = self.ln_q(patch_embed)
|
||||||
|
x = x.reshape(-1, self.merged_dim)
|
||||||
|
return self.mlp(x)
|
||||||
|
|
||||||
|
|
||||||
|
_ADAPTER_CLASSES = {
|
||||||
|
"pixel_shuffle_mlp": PixelShuffleAdapter,
|
||||||
|
"patch_merger": PatchMergerAdapter,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
# ---- Full Model ----
|
||||||
|
|
||||||
|
|
||||||
|
class DotsMoEVitModel(PreTrainedModel):
|
||||||
|
config_class = DotsMoEVitConfig
|
||||||
|
|
||||||
|
def __init__(self, config: DotsMoEVitConfig) -> None:
|
||||||
|
super().__init__(config)
|
||||||
|
self.config = config
|
||||||
|
self.spatial_merge_size = config.spatial_merge_size
|
||||||
|
|
||||||
|
self.patch_embed = DotsPatchEmbed(config)
|
||||||
|
|
||||||
|
head_dim = config.embed_dim // config.num_attention_heads
|
||||||
|
self.rotary_pos_emb = VisionRotaryEmbedding(head_dim // 2, cache_seq_len=100000)
|
||||||
|
|
||||||
|
self.blocks = nn.ModuleList(
|
||||||
|
[MoEVisionBlock(config, i) for i in range(config.num_hidden_layers)]
|
||||||
|
)
|
||||||
|
|
||||||
|
if config.post_norm:
|
||||||
|
self.post_trunk_norm = RMSNorm(config.embed_dim, eps=config.rms_norm_eps)
|
||||||
|
|
||||||
|
adapter_cls = _ADAPTER_CLASSES.get(config.adapter_type)
|
||||||
|
if adapter_cls is None:
|
||||||
|
raise ValueError(f"Unknown adapter_type {config.adapter_type!r}")
|
||||||
|
self.adapter = adapter_cls(config)
|
||||||
|
|
||||||
|
self.gradient_checkpointing = False
|
||||||
|
self._gradient_checkpointing_func = torch.utils.checkpoint.checkpoint
|
||||||
|
|
||||||
|
@property
|
||||||
|
def dtype(self) -> torch.dtype:
|
||||||
|
mlp = self.blocks[0].mlp
|
||||||
|
if isinstance(mlp, DotsSwiGLUFFN):
|
||||||
|
return mlp.fc13.weight.dtype
|
||||||
|
expert = mlp.experts[0]
|
||||||
|
return expert.fc13.weight.dtype
|
||||||
|
|
||||||
|
@property
|
||||||
|
def device(self) -> torch.device:
|
||||||
|
return self.patch_embed.proj.weight.device
|
||||||
|
|
||||||
|
def get_pos_ids_by_grid(self, grid_thw):
|
||||||
|
# Mirrors ``cybertron`` ``AIMv2NativeModel.rot_pos_emb``: when ``pre_pixel_shuffle``
|
||||||
|
# is set, RoPE positions follow the qwen ``merge_size`` grouped layout (default 2x2);
|
||||||
|
# otherwise positions are flat row-major regardless of ``spatial_merge_size``.
|
||||||
|
if self.config.pre_pixel_shuffle:
|
||||||
|
rope_merge_size = (
|
||||||
|
self.spatial_merge_size if self.spatial_merge_size > 1 else 2
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
rope_merge_size = 1
|
||||||
|
pos_ids = []
|
||||||
|
for t, h, w in grid_thw:
|
||||||
|
hpos_ids = torch.arange(h).unsqueeze(1).expand(-1, w)
|
||||||
|
hpos_ids = hpos_ids.reshape(
|
||||||
|
h // rope_merge_size,
|
||||||
|
rope_merge_size,
|
||||||
|
w // rope_merge_size,
|
||||||
|
rope_merge_size,
|
||||||
|
)
|
||||||
|
hpos_ids = hpos_ids.permute(0, 2, 1, 3).flatten()
|
||||||
|
|
||||||
|
wpos_ids = torch.arange(w).unsqueeze(0).expand(h, -1)
|
||||||
|
wpos_ids = wpos_ids.reshape(
|
||||||
|
h // rope_merge_size,
|
||||||
|
rope_merge_size,
|
||||||
|
w // rope_merge_size,
|
||||||
|
rope_merge_size,
|
||||||
|
)
|
||||||
|
wpos_ids = wpos_ids.permute(0, 2, 1, 3).flatten()
|
||||||
|
pos_ids.append(torch.stack([hpos_ids, wpos_ids], dim=-1).repeat(t, 1))
|
||||||
|
return pos_ids
|
||||||
|
|
||||||
|
def rot_pos_emb(self, grid_thw):
|
||||||
|
pos_ids = self.get_pos_ids_by_grid(grid_thw)
|
||||||
|
pos_ids = torch.cat(pos_ids, dim=0)
|
||||||
|
max_grid_size = grid_thw[:, 1:].max()
|
||||||
|
rotary_pos_emb_full = self.rotary_pos_emb(max_grid_size)
|
||||||
|
rotary_pos_emb = rotary_pos_emb_full[pos_ids].flatten(1)
|
||||||
|
return rotary_pos_emb
|
||||||
|
|
||||||
|
def _build_cu_seqlens_from_grid(self, grid_thw: torch.Tensor):
|
||||||
|
cu_seqlens = torch.repeat_interleave(
|
||||||
|
grid_thw[:, 1] * grid_thw[:, 2], grid_thw[:, 0]
|
||||||
|
).cumsum(
|
||||||
|
dim=0,
|
||||||
|
dtype=grid_thw.dtype if torch.jit.is_tracing() else torch.int32,
|
||||||
|
)
|
||||||
|
cu_seqlens = F.pad(cu_seqlens, (1, 0), value=0)
|
||||||
|
# Same as (cu_seqlens[1:] - cu_seqlens[:-1]).max(); computed here to avoid D2H inside attention.
|
||||||
|
max_seqlen = int((grid_thw[:, 1] * grid_thw[:, 2]).max().item())
|
||||||
|
return cu_seqlens, max_seqlen
|
||||||
|
|
||||||
|
def _build_single_temporal_cu_seqlens_from_grid(self, grid_thw: torch.Tensor):
|
||||||
|
seq_lens = grid_thw[:, 1] * grid_thw[:, 2]
|
||||||
|
cu_seqlens = torch.empty(
|
||||||
|
(grid_thw.shape[0] + 1,), device=grid_thw.device, dtype=torch.int32
|
||||||
|
)
|
||||||
|
cu_seqlens[0] = 0
|
||||||
|
torch.cumsum(seq_lens, dim=0, dtype=torch.int32, out=cu_seqlens[1:])
|
||||||
|
max_seqlen = int(seq_lens.max().item())
|
||||||
|
return cu_seqlens, max_seqlen
|
||||||
|
|
||||||
|
def forward(
|
||||||
|
self, hidden_states: torch.Tensor, grid_thw: torch.Tensor, bf16=True
|
||||||
|
) -> torch.Tensor:
|
||||||
|
if bf16:
|
||||||
|
hidden_states = hidden_states.bfloat16()
|
||||||
|
hidden_states = self.patch_embed(hidden_states)
|
||||||
|
|
||||||
|
rotary_pos_emb = self.rot_pos_emb(grid_thw)
|
||||||
|
|
||||||
|
if grid_thw[:, 0].sum().item() == grid_thw.shape[0]:
|
||||||
|
cu_seqlens, max_seqlen = self._build_single_temporal_cu_seqlens_from_grid(
|
||||||
|
grid_thw
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
cu_seqlens, max_seqlen = self._build_cu_seqlens_from_grid(grid_thw)
|
||||||
|
|
||||||
|
for blk in self.blocks:
|
||||||
|
if self.gradient_checkpointing and self.training:
|
||||||
|
hidden_states = self._gradient_checkpointing_func(
|
||||||
|
blk.__call__,
|
||||||
|
hidden_states,
|
||||||
|
cu_seqlens,
|
||||||
|
rotary_pos_emb,
|
||||||
|
max_seqlen,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
hidden_states = blk(
|
||||||
|
hidden_states,
|
||||||
|
cu_seqlens,
|
||||||
|
rotary_pos_emb,
|
||||||
|
max_seqlen,
|
||||||
|
)
|
||||||
|
|
||||||
|
if self.config.post_norm:
|
||||||
|
hidden_states = self.post_trunk_norm(hidden_states)
|
||||||
|
|
||||||
|
hidden_states = self.adapter(hidden_states, grid_thw)
|
||||||
|
return hidden_states
|
||||||
@@ -0,0 +1,112 @@
|
|||||||
|
"""Dots-specific FP8 helpers for absorbed MLA batched matmuls."""
|
||||||
|
|
||||||
|
from typing import Tuple
|
||||||
|
|
||||||
|
import torch
|
||||||
|
import triton
|
||||||
|
import triton.language as tl
|
||||||
|
|
||||||
|
from sglang.kernels.ops.quantization.fp8_kernel import is_fp8_fnuz
|
||||||
|
from sglang.srt.utils import ceil_align
|
||||||
|
|
||||||
|
_FP8_MAX = 224.0 if is_fp8_fnuz() else torch.finfo(torch.float8_e4m3fn).max
|
||||||
|
|
||||||
|
|
||||||
|
@triton.jit
|
||||||
|
def _per_token_group_quant_einsum_fp8(
|
||||||
|
x_ptr,
|
||||||
|
x_q_ptr,
|
||||||
|
x_s_ptr,
|
||||||
|
group_size,
|
||||||
|
num_b,
|
||||||
|
num_k,
|
||||||
|
total_rows,
|
||||||
|
x_stride_m,
|
||||||
|
x_stride_b,
|
||||||
|
x_q_stride_m,
|
||||||
|
x_q_stride_b,
|
||||||
|
x_s_stride_m,
|
||||||
|
x_s_stride_b,
|
||||||
|
x_s_stride_g,
|
||||||
|
eps,
|
||||||
|
quant_min,
|
||||||
|
quant_max,
|
||||||
|
BLOCK_M: tl.constexpr,
|
||||||
|
BLOCK_K: tl.constexpr,
|
||||||
|
):
|
||||||
|
row_ids = tl.program_id(0) * BLOCK_M + tl.arange(0, BLOCK_M)
|
||||||
|
group_id = tl.program_id(1)
|
||||||
|
m_ids = row_ids // num_b
|
||||||
|
b_ids = row_ids - m_ids * num_b
|
||||||
|
k_offsets = tl.arange(0, BLOCK_K)
|
||||||
|
k_ids = group_id * group_size + k_offsets
|
||||||
|
mask = (row_ids[:, None] < total_rows) & (
|
||||||
|
(k_offsets[None, :] < group_size) & (k_ids[None, :] < num_k)
|
||||||
|
)
|
||||||
|
x_ptrs = (
|
||||||
|
x_ptr
|
||||||
|
+ m_ids[:, None] * x_stride_m
|
||||||
|
+ b_ids[:, None] * x_stride_b
|
||||||
|
+ k_ids[None, :]
|
||||||
|
)
|
||||||
|
x = tl.load(x_ptrs, mask=mask, other=0.0).to(tl.float32)
|
||||||
|
absmax = tl.maximum(tl.max(tl.abs(x), axis=1), eps)
|
||||||
|
scale = absmax / quant_max
|
||||||
|
quant = tl.clamp(x / scale[:, None], quant_min, quant_max).to(
|
||||||
|
x_q_ptr.dtype.element_ty
|
||||||
|
)
|
||||||
|
q_ptrs = (
|
||||||
|
x_q_ptr
|
||||||
|
+ m_ids[:, None] * x_q_stride_m
|
||||||
|
+ b_ids[:, None] * x_q_stride_b
|
||||||
|
+ k_ids[None, :]
|
||||||
|
)
|
||||||
|
s_ptrs = (
|
||||||
|
x_s_ptr + m_ids * x_s_stride_m + b_ids * x_s_stride_b + group_id * x_s_stride_g
|
||||||
|
)
|
||||||
|
tl.store(q_ptrs, quant, mask=mask)
|
||||||
|
tl.store(s_ptrs, scale, mask=row_ids < total_rows)
|
||||||
|
|
||||||
|
|
||||||
|
def per_token_group_quant_einsum_fp8(
|
||||||
|
x: torch.Tensor,
|
||||||
|
group_size: int = 128,
|
||||||
|
eps: float = 1e-12,
|
||||||
|
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||||
|
"""Quantize ``[m, b, k]`` in the scale layout required by FP8 einsum."""
|
||||||
|
assert x.ndim == 3 and x.stride(-1) == 1
|
||||||
|
assert group_size == 128
|
||||||
|
m, b, k = x.shape
|
||||||
|
num_groups = (k + group_size - 1) // group_size
|
||||||
|
aligned_m = ceil_align(m, 4)
|
||||||
|
x_q = x.new_empty((m, b, k), dtype=torch.float8_e4m3fn)
|
||||||
|
scale_storage = x.new_empty((b, num_groups, aligned_m), dtype=torch.float32)
|
||||||
|
x_s = scale_storage.permute(2, 0, 1)[:m]
|
||||||
|
if m == 0 or b == 0 or num_groups == 0:
|
||||||
|
return x_q, x_s
|
||||||
|
block_m = 16
|
||||||
|
block_k = triton.next_power_of_2(group_size)
|
||||||
|
_per_token_group_quant_einsum_fp8[(triton.cdiv(m * b, block_m), num_groups)](
|
||||||
|
x,
|
||||||
|
x_q,
|
||||||
|
x_s,
|
||||||
|
group_size,
|
||||||
|
b,
|
||||||
|
k,
|
||||||
|
m * b,
|
||||||
|
x.stride(0),
|
||||||
|
x.stride(1),
|
||||||
|
x_q.stride(0),
|
||||||
|
x_q.stride(1),
|
||||||
|
x_s.stride(0),
|
||||||
|
x_s.stride(1),
|
||||||
|
x_s.stride(2),
|
||||||
|
eps,
|
||||||
|
-_FP8_MAX,
|
||||||
|
_FP8_MAX,
|
||||||
|
block_m,
|
||||||
|
block_k,
|
||||||
|
num_warps=4,
|
||||||
|
num_stages=1,
|
||||||
|
)
|
||||||
|
return x_q, x_s
|
||||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,204 @@
|
|||||||
|
"""Inference-only full-sharing Dots3 MTP / NextN draft model."""
|
||||||
|
|
||||||
|
import logging
|
||||||
|
from collections.abc import Iterable
|
||||||
|
|
||||||
|
import torch
|
||||||
|
from torch import nn
|
||||||
|
from transformers import PretrainedConfig
|
||||||
|
|
||||||
|
from sglang.srt.distributed import get_pp_group
|
||||||
|
from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder
|
||||||
|
from sglang.srt.layers.dp_attention import is_dp_attention_enabled
|
||||||
|
from sglang.srt.layers.layernorm import RMSNorm
|
||||||
|
from sglang.srt.layers.linear import ReplicatedLinear
|
||||||
|
from sglang.srt.layers.logits_processor import LogitsProcessor
|
||||||
|
from sglang.srt.layers.quantization.base_config import QuantizationConfig
|
||||||
|
from sglang.srt.layers.vocab_parallel_embedding import (
|
||||||
|
ParallelLMHead,
|
||||||
|
VocabParallelEmbedding,
|
||||||
|
)
|
||||||
|
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||||
|
from sglang.srt.models.dots3_common.modeling import (
|
||||||
|
Dots3DecoderLayer,
|
||||||
|
Dots3LanguageModelForCausalLM,
|
||||||
|
)
|
||||||
|
from sglang.srt.runtime_context import get_parallel
|
||||||
|
from sglang.srt.utils import BumpAllocator, add_prefix
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
class Dots3MTPHead(nn.Module):
|
||||||
|
"""The single MTP layer, recursively reused by every draft step."""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
config: PretrainedConfig,
|
||||||
|
quant_config: QuantizationConfig | None,
|
||||||
|
prefix: str,
|
||||||
|
) -> None:
|
||||||
|
super().__init__()
|
||||||
|
self.enorm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
|
||||||
|
self.hnorm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
|
||||||
|
self.eh_proj = ReplicatedLinear(
|
||||||
|
2 * config.hidden_size,
|
||||||
|
config.hidden_size,
|
||||||
|
bias=False,
|
||||||
|
quant_config=quant_config,
|
||||||
|
prefix=add_prefix("eh_proj", prefix),
|
||||||
|
)
|
||||||
|
self.decoder = Dots3DecoderLayer(
|
||||||
|
config,
|
||||||
|
layer_id=0,
|
||||||
|
quant_config=quant_config,
|
||||||
|
is_nextn=True,
|
||||||
|
prefix=add_prefix("decoder", prefix),
|
||||||
|
)
|
||||||
|
self.shared_head = nn.Module()
|
||||||
|
self.shared_head.norm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
|
||||||
|
|
||||||
|
|
||||||
|
class Dot3NoteModelNextN(nn.Module):
|
||||||
|
"""Text-only draft model containing one full-sharing MTP layer."""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
config: PretrainedConfig,
|
||||||
|
quant_config: QuantizationConfig | None = None,
|
||||||
|
prefix: str = "",
|
||||||
|
) -> None:
|
||||||
|
super().__init__()
|
||||||
|
if config.num_nextn_predict_layers != 1:
|
||||||
|
raise ValueError(
|
||||||
|
"Dots3 MTP currently supports one full-sharing layer only."
|
||||||
|
)
|
||||||
|
if list(config.layer_types) != ["sliding_attention"]:
|
||||||
|
raise ValueError("Dots3 MTP full-sharing layer must use sliding_attention.")
|
||||||
|
|
||||||
|
self.vocab_size = config.vocab_size
|
||||||
|
self.embed_tokens = VocabParallelEmbedding(
|
||||||
|
config.vocab_size,
|
||||||
|
config.hidden_size,
|
||||||
|
enable_tp=not is_dp_attention_enabled(),
|
||||||
|
prefix=add_prefix("embed_tokens", prefix),
|
||||||
|
)
|
||||||
|
# The weight loader maps the shared MTP layer to heads.0.
|
||||||
|
self.heads = nn.ModuleList(
|
||||||
|
[Dots3MTPHead(config, quant_config, add_prefix("heads.0", prefix))]
|
||||||
|
)
|
||||||
|
|
||||||
|
def forward(
|
||||||
|
self,
|
||||||
|
input_ids: torch.Tensor,
|
||||||
|
positions: torch.Tensor,
|
||||||
|
forward_batch: ForwardBatch,
|
||||||
|
input_embeds: torch.Tensor | None = None,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
device = input_embeds.device if input_embeds is not None else input_ids.device
|
||||||
|
zero_allocator = BumpAllocator(
|
||||||
|
buffer_size=2, dtype=torch.float32, device=device
|
||||||
|
)
|
||||||
|
hidden_states = (
|
||||||
|
self._embed_input_ids(input_ids) if input_embeds is None else input_embeds
|
||||||
|
)
|
||||||
|
head = self.heads[0]
|
||||||
|
if hidden_states.shape[0] > 0:
|
||||||
|
hidden_states, _ = head.eh_proj(
|
||||||
|
torch.cat(
|
||||||
|
(
|
||||||
|
head.enorm(hidden_states),
|
||||||
|
head.hnorm(forward_batch.spec_info.hidden_states),
|
||||||
|
),
|
||||||
|
dim=-1,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
residual = None
|
||||||
|
with get_global_expert_distribution_recorder().disable_this_region():
|
||||||
|
hidden_states, residual = head.decoder(
|
||||||
|
positions, hidden_states, forward_batch, residual, zero_allocator
|
||||||
|
)
|
||||||
|
|
||||||
|
if not forward_batch.forward_mode.is_idle():
|
||||||
|
if residual is None:
|
||||||
|
hidden_states = head.shared_head.norm(hidden_states)
|
||||||
|
else:
|
||||||
|
hidden_states, _ = head.shared_head.norm(hidden_states, residual)
|
||||||
|
return hidden_states
|
||||||
|
|
||||||
|
def _embed_input_ids(self, input_ids: torch.Tensor) -> torch.Tensor:
|
||||||
|
# Multimodal sentinels use target hidden states, so clamp their unused
|
||||||
|
# draft embedding indices to the vocabulary.
|
||||||
|
return self.embed_tokens(input_ids.clamp(min=0, max=self.vocab_size - 1))
|
||||||
|
|
||||||
|
|
||||||
|
class Dots3NoteForCausalLMNextN(Dots3LanguageModelForCausalLM):
|
||||||
|
"""Full-sharing Dots3 MTP draft registered for NEXTN decoding."""
|
||||||
|
|
||||||
|
fused_shared_experts_architecture = "Dots3NoteForCausalLMNextN"
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
config: PretrainedConfig,
|
||||||
|
quant_config: QuantizationConfig | None = None,
|
||||||
|
prefix: str = "",
|
||||||
|
) -> None:
|
||||||
|
nn.Module.__init__(self)
|
||||||
|
self.config = config
|
||||||
|
self.tp_size = get_parallel().tp_size
|
||||||
|
self.quant_config = quant_config
|
||||||
|
self.pp_group = get_pp_group()
|
||||||
|
self.fuse_qkv_a_g_proj = True
|
||||||
|
self.packed_modules_mapping = {
|
||||||
|
"fused_qkv_a_g_proj_with_mqa": [
|
||||||
|
"q_a_proj",
|
||||||
|
"kv_a_proj_with_mqa",
|
||||||
|
"g_proj",
|
||||||
|
]
|
||||||
|
}
|
||||||
|
self.determine_num_fused_shared_experts()
|
||||||
|
|
||||||
|
self.model = Dot3NoteModelNextN(
|
||||||
|
config, quant_config, prefix=add_prefix("model", prefix)
|
||||||
|
)
|
||||||
|
self.lm_head = ParallelLMHead(
|
||||||
|
config.vocab_size,
|
||||||
|
config.hidden_size,
|
||||||
|
quant_config=quant_config,
|
||||||
|
prefix=add_prefix("model.shared_head.head", prefix),
|
||||||
|
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||||
|
)
|
||||||
|
self.logits_processor = LogitsProcessor(config)
|
||||||
|
self._mtp_loaded_embed = False
|
||||||
|
|
||||||
|
@torch.no_grad()
|
||||||
|
def forward(
|
||||||
|
self,
|
||||||
|
input_ids: torch.Tensor,
|
||||||
|
positions: torch.Tensor,
|
||||||
|
forward_batch: ForwardBatch,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
hidden_states = self.model(input_ids, positions, forward_batch)
|
||||||
|
return self.logits_processor(
|
||||||
|
input_ids, hidden_states, self.lm_head, forward_batch
|
||||||
|
)
|
||||||
|
|
||||||
|
def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]):
|
||||||
|
weights = list(weights)
|
||||||
|
self._mtp_loaded_embed = any(
|
||||||
|
name.startswith("model.mtp.embed_tokens.") for name, _ in weights
|
||||||
|
)
|
||||||
|
super().load_weights(weights, is_nextn=True)
|
||||||
|
|
||||||
|
def set_embed_and_head(self, embed, head):
|
||||||
|
# Preserve a checkpoint-provided MTP embedding; share the output head.
|
||||||
|
if not self._mtp_loaded_embed:
|
||||||
|
del self.model.embed_tokens.weight
|
||||||
|
self.model.embed_tokens.weight = embed
|
||||||
|
else:
|
||||||
|
logger.info("Keeping the checkpoint's MTP-specific input embedding.")
|
||||||
|
del self.lm_head.weight
|
||||||
|
self.lm_head.weight = head
|
||||||
|
torch.cuda.empty_cache()
|
||||||
|
torch.cuda.synchronize()
|
||||||
@@ -0,0 +1,19 @@
|
|||||||
|
# SPDX-License-Identifier: Apache-2.0
|
||||||
|
# Copyright 2023-2024 SGLang Team
|
||||||
|
|
||||||
|
"""Registry entry point for the Dots3 next-N model."""
|
||||||
|
|
||||||
|
from sglang.srt.models.dots3_common.nextn import (
|
||||||
|
Dot3NoteModelNextN,
|
||||||
|
Dots3MTPHead,
|
||||||
|
Dots3NoteForCausalLMNextN,
|
||||||
|
)
|
||||||
|
|
||||||
|
EntryClass = [Dots3NoteForCausalLMNextN]
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
"Dot3NoteModelNextN",
|
||||||
|
"Dots3MTPHead",
|
||||||
|
"Dots3NoteForCausalLMNextN",
|
||||||
|
"EntryClass",
|
||||||
|
]
|
||||||
@@ -0,0 +1,565 @@
|
|||||||
|
import asyncio
|
||||||
|
import base64
|
||||||
|
import hashlib
|
||||||
|
import logging
|
||||||
|
import os
|
||||||
|
import re
|
||||||
|
import time
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Any, ClassVar
|
||||||
|
|
||||||
|
import numpy as np
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from sglang.srt.managers.io_struct import GenerateReqInput
|
||||||
|
from sglang.srt.managers.schedule_batch import (
|
||||||
|
Modality,
|
||||||
|
MultimodalDataItem,
|
||||||
|
MultimodalProcessorOutput,
|
||||||
|
)
|
||||||
|
from sglang.srt.models.dots3 import Dots3NoteForCausalLM
|
||||||
|
from sglang.srt.models.dots3_common.dots_omni_towers import (
|
||||||
|
DotsNoteOmniImagePreprocessor,
|
||||||
|
OmniAudioConfig,
|
||||||
|
get_audio_token_string,
|
||||||
|
load_omni_component_config,
|
||||||
|
)
|
||||||
|
from sglang.srt.multimodal.processors.base_processor import (
|
||||||
|
BaseMultimodalProcessor,
|
||||||
|
MultimodalSpecialTokens,
|
||||||
|
)
|
||||||
|
from sglang.srt.utils import VideoData, get_video_bytes
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
_VIDEO_TOKEN_RE = re.compile(r"(<image_\d+>|<audio_\d+>)")
|
||||||
|
_EXPANDED_VIDEO_MEDIA_RE = re.compile(
|
||||||
|
r"<\|sglang_dots_video_(?P<video>\d+)_(?P<modality>image|audio)_(?P<item>\d+)\|>"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _build_video_cfg(
|
||||||
|
*,
|
||||||
|
seq: int,
|
||||||
|
audio_cap: float,
|
||||||
|
audio_sr: int,
|
||||||
|
max_new_tokens: int,
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
if seq <= 0:
|
||||||
|
raise ValueError(f"seq must be positive, got {seq}")
|
||||||
|
if max_new_tokens < 0:
|
||||||
|
raise ValueError(f"max_new_tokens must be non-negative, got {max_new_tokens}")
|
||||||
|
if max_new_tokens >= seq:
|
||||||
|
raise ValueError(
|
||||||
|
"max_new_tokens must leave room for input: "
|
||||||
|
f"max_new_tokens={max_new_tokens}, seq={seq}"
|
||||||
|
)
|
||||||
|
if audio_cap < 0:
|
||||||
|
raise ValueError(f"audio_cap must be non-negative, got {audio_cap}")
|
||||||
|
if audio_sr <= 0:
|
||||||
|
raise ValueError(f"audio_sr must be positive, got {audio_sr}")
|
||||||
|
|
||||||
|
return {
|
||||||
|
"process_audio": audio_cap > 0,
|
||||||
|
"seq_length": seq - max_new_tokens,
|
||||||
|
"reserve_interleave": True,
|
||||||
|
"audio_token_ratio_cap": float(audio_cap),
|
||||||
|
"audio_sample_rate": int(audio_sr),
|
||||||
|
"video_jpeg_quality": int(os.environ.get("XHS_VIDEO_JPEG_QUALITY", "85")),
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def _video_payload(raw_video) -> tuple[bytes, str]:
|
||||||
|
if isinstance(raw_video, VideoData):
|
||||||
|
raw_video = raw_video.url
|
||||||
|
raw_url = raw_video.get("url") if isinstance(raw_video, dict) else raw_video
|
||||||
|
video_bytes = get_video_bytes(raw_url)
|
||||||
|
return video_bytes, hashlib.sha1(video_bytes).hexdigest()
|
||||||
|
|
||||||
|
|
||||||
|
def _cfg_for_pure_visual(cfg: dict[str, Any]) -> dict[str, Any]:
|
||||||
|
cfg = dict(cfg)
|
||||||
|
cfg["process_audio"] = False
|
||||||
|
return cfg
|
||||||
|
|
||||||
|
|
||||||
|
def _flat_video_to_content(flat: dict[str, Any]) -> list[dict[str, Any]]:
|
||||||
|
meta = flat.get("meta", {})
|
||||||
|
user_value = next(
|
||||||
|
(
|
||||||
|
conv.get("value", "")
|
||||||
|
for conv in flat.get("conversations", [])
|
||||||
|
if (conv.get("from") or conv.get("role")) == "user"
|
||||||
|
),
|
||||||
|
"",
|
||||||
|
)
|
||||||
|
content: list[dict[str, Any]] = []
|
||||||
|
last = 0
|
||||||
|
for match in _VIDEO_TOKEN_RE.finditer(user_value):
|
||||||
|
if match.start() > last:
|
||||||
|
content.append({"type": "text", "text": user_value[last : match.start()]})
|
||||||
|
key = match.group(1)[1:-1]
|
||||||
|
encoded = meta.get(key)
|
||||||
|
if encoded and key.startswith("image_"):
|
||||||
|
content.append(
|
||||||
|
{
|
||||||
|
"type": "image_url",
|
||||||
|
"image_url": {"url": f"data:image/jpeg;base64,{encoded}"},
|
||||||
|
}
|
||||||
|
)
|
||||||
|
elif encoded:
|
||||||
|
content.append(
|
||||||
|
{
|
||||||
|
"type": "audio_url",
|
||||||
|
"audio_url": {"url": f"data:audio/wav;base64,{encoded}"},
|
||||||
|
}
|
||||||
|
)
|
||||||
|
last = match.end()
|
||||||
|
if last < len(user_value):
|
||||||
|
content.append({"type": "text", "text": user_value[last:]})
|
||||||
|
return content
|
||||||
|
|
||||||
|
|
||||||
|
def preprocess_dots_video(
|
||||||
|
raw_video,
|
||||||
|
question: str,
|
||||||
|
*,
|
||||||
|
tokenizer,
|
||||||
|
seq: int = 131072,
|
||||||
|
audio_cap: float = 1.0,
|
||||||
|
audio_sr: int = 16000,
|
||||||
|
k_mode: str = "eval_ek",
|
||||||
|
max_new_tokens: int = 0,
|
||||||
|
) -> list[dict[str, Any]]:
|
||||||
|
"""Return in-memory timestamp/image/audio content using the server tokenizer."""
|
||||||
|
if not k_mode:
|
||||||
|
raise ValueError("k_mode must not be empty")
|
||||||
|
video_bytes, video_id = _video_payload(raw_video)
|
||||||
|
cfg = _build_video_cfg(
|
||||||
|
seq=seq,
|
||||||
|
audio_cap=audio_cap,
|
||||||
|
audio_sr=audio_sr,
|
||||||
|
max_new_tokens=max_new_tokens,
|
||||||
|
)
|
||||||
|
from sglang.srt.multimodal.processors.dots_note_omni_video_core import (
|
||||||
|
flatten_runner,
|
||||||
|
)
|
||||||
|
from sglang.srt.multimodal.processors.dots_note_omni_video_core import (
|
||||||
|
preprocess as pp,
|
||||||
|
)
|
||||||
|
|
||||||
|
video_b64 = base64.b64encode(video_bytes).decode()
|
||||||
|
sample = {
|
||||||
|
"meta": {"video_0": video_b64},
|
||||||
|
"conversations": [{"from": "user", "value": f"<video_0>{question}"}],
|
||||||
|
}
|
||||||
|
record_key = hashlib.sha1(f"{video_id}|{question}".encode()).hexdigest()
|
||||||
|
|
||||||
|
def run(run_cfg):
|
||||||
|
new_meta, conversations = pp.process_sample_video(
|
||||||
|
sample, run_cfg, tokenizer=tokenizer
|
||||||
|
)
|
||||||
|
plan = flatten_runner.build_plan(
|
||||||
|
new_meta,
|
||||||
|
conversations,
|
||||||
|
record_key,
|
||||||
|
k_mode=k_mode,
|
||||||
|
process_audio=run_cfg["process_audio"],
|
||||||
|
)
|
||||||
|
return _flat_video_to_content(flatten_runner.render_flat(plan))
|
||||||
|
|
||||||
|
try:
|
||||||
|
return run(cfg)
|
||||||
|
except pp.SkipSample as exc:
|
||||||
|
if "audio_token_ratio_exceed" not in str(exc):
|
||||||
|
raise
|
||||||
|
return run(_cfg_for_pure_visual(cfg))
|
||||||
|
|
||||||
|
|
||||||
|
class DotsNoteOmniProcessor(BaseMultimodalProcessor):
|
||||||
|
"""Native image/audio processor for dots.note.omni."""
|
||||||
|
|
||||||
|
models: ClassVar[list] = [Dots3NoteForCausalLM]
|
||||||
|
gpu_image_decode = False
|
||||||
|
|
||||||
|
def __init__(self, hf_config, server_args, processor, transport_mode, **kwargs):
|
||||||
|
self.image_start_token = hf_config.im_start_token
|
||||||
|
self.image_token = hf_config.im_token
|
||||||
|
self.image_end_token = hf_config.im_end_token
|
||||||
|
self.audio_start_token = hf_config.audio_start_token
|
||||||
|
self.audio_token = hf_config.audio_token
|
||||||
|
self.audio_end_token = hf_config.audio_end_token
|
||||||
|
self.video_placeholder_regex = re.compile(re.escape(hf_config.video_token))
|
||||||
|
self.mm_tokens = MultimodalSpecialTokens(
|
||||||
|
image_token=self.image_token,
|
||||||
|
image_token_id=self._token_id(processor, self.image_token),
|
||||||
|
image_token_regex=re.compile(
|
||||||
|
re.escape(
|
||||||
|
self.image_start_token + self.image_token + self.image_end_token
|
||||||
|
)
|
||||||
|
),
|
||||||
|
audio_token=self.audio_token,
|
||||||
|
audio_token_id=self._token_id(processor, self.audio_token),
|
||||||
|
audio_token_regex=re.compile(
|
||||||
|
re.escape(
|
||||||
|
self.audio_start_token + self.audio_token + self.audio_end_token
|
||||||
|
)
|
||||||
|
),
|
||||||
|
).build(processor)
|
||||||
|
self.mm_token_ids = {
|
||||||
|
"im_start_id": self._token_id(processor, self.image_start_token),
|
||||||
|
"im_token_id": self._token_id(processor, self.image_token),
|
||||||
|
"im_end_id": self._token_id(processor, self.image_end_token),
|
||||||
|
"audio_start_id": self._token_id(processor, self.audio_start_token),
|
||||||
|
"audio_token_id": self._token_id(processor, self.audio_token),
|
||||||
|
"audio_end_id": self._token_id(processor, self.audio_end_token),
|
||||||
|
}
|
||||||
|
|
||||||
|
model_dir = Path(hf_config._name_or_path)
|
||||||
|
self.image_preprocessor = DotsNoteOmniImagePreprocessor(str(model_dir))
|
||||||
|
self.audio_processor_config = OmniAudioConfig(
|
||||||
|
**load_omni_component_config(model_dir, "audio")
|
||||||
|
)
|
||||||
|
super().__init__(hf_config, server_args, processor, transport_mode, **kwargs)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _token_id(processor, token: str) -> int:
|
||||||
|
token_ids = processor.encode(token, add_special_tokens=False)
|
||||||
|
if len(token_ids) != 1:
|
||||||
|
raise ValueError(
|
||||||
|
f"Dots omni special token {token!r} must encode to one id, got "
|
||||||
|
f"{token_ids}"
|
||||||
|
)
|
||||||
|
return token_ids[0]
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _normalize_audio(audio) -> torch.Tensor:
|
||||||
|
if isinstance(audio, torch.Tensor):
|
||||||
|
waveform = audio
|
||||||
|
elif isinstance(audio, np.ndarray):
|
||||||
|
waveform = torch.from_numpy(audio)
|
||||||
|
else:
|
||||||
|
waveform = torch.as_tensor(audio)
|
||||||
|
waveform = waveform.float().squeeze()
|
||||||
|
if waveform.ndim != 1:
|
||||||
|
raise ValueError(
|
||||||
|
f"Dots omni audio must be mono, got shape={tuple(waveform.shape)}"
|
||||||
|
)
|
||||||
|
return waveform.contiguous()
|
||||||
|
|
||||||
|
def _render_video_content(
|
||||||
|
self,
|
||||||
|
input_text: str,
|
||||||
|
question: str,
|
||||||
|
video_index: int,
|
||||||
|
content: list[dict],
|
||||||
|
) -> tuple[str, dict[str, tuple[Modality, str]]]:
|
||||||
|
"""Insert one expanded video while retaining its media ordering."""
|
||||||
|
rendered = []
|
||||||
|
media = {}
|
||||||
|
for item in content:
|
||||||
|
item_type = item.get("type")
|
||||||
|
if item_type == "text":
|
||||||
|
rendered.append(item.get("text", ""))
|
||||||
|
elif item_type == "image_url":
|
||||||
|
marker = f"<|sglang_dots_video_{video_index}_image_{len(media)}|>"
|
||||||
|
media[marker] = (Modality.IMAGE, item["image_url"]["url"])
|
||||||
|
rendered.append(marker)
|
||||||
|
elif item_type == "audio_url":
|
||||||
|
marker = f"<|sglang_dots_video_{video_index}_audio_{len(media)}|>"
|
||||||
|
media[marker] = (Modality.AUDIO, item["audio_url"]["url"])
|
||||||
|
rendered.append(marker)
|
||||||
|
else:
|
||||||
|
raise ValueError(f"Unsupported preprocessed video item: {item_type}")
|
||||||
|
|
||||||
|
expanded = "".join(rendered)
|
||||||
|
# The adapter appends the question to every flattened video. Keep the
|
||||||
|
# question already rendered by the chat template so multiple videos do
|
||||||
|
# not duplicate it.
|
||||||
|
if question:
|
||||||
|
question_pos = expanded.rfind(question)
|
||||||
|
if question_pos >= 0:
|
||||||
|
expanded = (
|
||||||
|
expanded[:question_pos] + expanded[question_pos + len(question) :]
|
||||||
|
)
|
||||||
|
|
||||||
|
placeholder = self.video_placeholder_regex.search(input_text)
|
||||||
|
if placeholder is not None:
|
||||||
|
input_text = (
|
||||||
|
input_text[: placeholder.start()]
|
||||||
|
+ expanded
|
||||||
|
+ input_text[placeholder.end() :]
|
||||||
|
)
|
||||||
|
elif question and question in input_text:
|
||||||
|
input_text = input_text.replace(question, expanded + question, 1)
|
||||||
|
else:
|
||||||
|
# The normal dots template starts the user turn with <|user|>. Keep
|
||||||
|
# system text ahead of video media if a custom template is used.
|
||||||
|
user_marker = "<|user|>"
|
||||||
|
pos = input_text.rfind(user_marker)
|
||||||
|
insert_at = pos + len(user_marker) if pos >= 0 else 0
|
||||||
|
input_text = input_text[:insert_at] + expanded + input_text[insert_at:]
|
||||||
|
return input_text, media
|
||||||
|
|
||||||
|
def _merge_video_media(
|
||||||
|
self,
|
||||||
|
input_text: str,
|
||||||
|
image_data: list | None,
|
||||||
|
audio_data: list | None,
|
||||||
|
video_media: dict[str, tuple[Modality, str]],
|
||||||
|
) -> tuple[str, list, list]:
|
||||||
|
"""Resolve native and video-derived media in final prompt order."""
|
||||||
|
native_images = iter(image_data or [])
|
||||||
|
native_audios = iter(audio_data or [])
|
||||||
|
ordered_images = []
|
||||||
|
ordered_audios = []
|
||||||
|
native_pattern = self.mm_tokens.get_combined_regex()
|
||||||
|
pattern = re.compile(
|
||||||
|
f"({native_pattern.pattern}|{_EXPANDED_VIDEO_MEDIA_RE.pattern})"
|
||||||
|
)
|
||||||
|
rendered = []
|
||||||
|
last = 0
|
||||||
|
|
||||||
|
for match in pattern.finditer(input_text):
|
||||||
|
rendered.append(input_text[last : match.start()])
|
||||||
|
marker = match.group(0)
|
||||||
|
expanded_media = video_media.get(marker)
|
||||||
|
if expanded_media is not None:
|
||||||
|
modality, value = expanded_media
|
||||||
|
else:
|
||||||
|
modality = self.mm_tokens.get_modality_of_token(marker)
|
||||||
|
if modality == Modality.IMAGE:
|
||||||
|
try:
|
||||||
|
value = next(native_images)
|
||||||
|
except StopIteration as exc:
|
||||||
|
raise ValueError(
|
||||||
|
"Image placeholder count does not match image_data"
|
||||||
|
) from exc
|
||||||
|
elif modality == Modality.AUDIO:
|
||||||
|
try:
|
||||||
|
value = next(native_audios)
|
||||||
|
except StopIteration as exc:
|
||||||
|
raise ValueError(
|
||||||
|
"Audio placeholder count does not match audio_data"
|
||||||
|
) from exc
|
||||||
|
else:
|
||||||
|
raise ValueError(f"Unsupported dots omni media marker: {marker}")
|
||||||
|
|
||||||
|
if modality == Modality.IMAGE:
|
||||||
|
ordered_images.append(value)
|
||||||
|
rendered.append(
|
||||||
|
self.image_start_token + self.image_token + self.image_end_token
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
ordered_audios.append(value)
|
||||||
|
rendered.append(
|
||||||
|
self.audio_start_token + self.audio_token + self.audio_end_token
|
||||||
|
)
|
||||||
|
last = match.end()
|
||||||
|
|
||||||
|
rendered.append(input_text[last:])
|
||||||
|
try:
|
||||||
|
next(native_images)
|
||||||
|
raise ValueError("Image placeholder count does not match image_data")
|
||||||
|
except StopIteration:
|
||||||
|
pass
|
||||||
|
try:
|
||||||
|
next(native_audios)
|
||||||
|
raise ValueError("Audio placeholder count does not match audio_data")
|
||||||
|
except StopIteration:
|
||||||
|
pass
|
||||||
|
return "".join(rendered), ordered_images, ordered_audios
|
||||||
|
|
||||||
|
async def process_mm_data_async(
|
||||||
|
self,
|
||||||
|
input_text: list[int] | str,
|
||||||
|
request_obj: GenerateReqInput,
|
||||||
|
max_req_input_len: int,
|
||||||
|
*args,
|
||||||
|
image_data: list | None = None,
|
||||||
|
audio_data: list | None = None,
|
||||||
|
video_data=None,
|
||||||
|
**kwargs,
|
||||||
|
):
|
||||||
|
video_data = request_obj.video_data or video_data
|
||||||
|
if not image_data and not audio_data and not video_data:
|
||||||
|
return None
|
||||||
|
if not isinstance(input_text, str):
|
||||||
|
raise ValueError( # noqa: TRY004 - preserve the processor API contract
|
||||||
|
"Dots note omni requires a text prompt for multimodal requests"
|
||||||
|
)
|
||||||
|
|
||||||
|
request_videos = len(video_data) if video_data else 0
|
||||||
|
request_images = len(image_data) if image_data else 0
|
||||||
|
request_audios = len(audio_data) if audio_data else 0
|
||||||
|
logger.info(
|
||||||
|
"[dots_mm] rid=%s request videos=%d images=%d audios=%d",
|
||||||
|
request_obj.rid,
|
||||||
|
request_videos,
|
||||||
|
request_images,
|
||||||
|
request_audios,
|
||||||
|
)
|
||||||
|
|
||||||
|
if video_data:
|
||||||
|
video_config = dict(request_obj.video_config or {})
|
||||||
|
question = video_config.pop("_question", "") or ""
|
||||||
|
seq = video_config.pop("seq", 131072)
|
||||||
|
audio_cap = video_config.pop("audio_cap", 1.0)
|
||||||
|
audio_sr = video_config.pop("audio_sr", 16000)
|
||||||
|
k_mode = video_config.pop("k_mode", "eval_ek")
|
||||||
|
if video_config:
|
||||||
|
raise ValueError(
|
||||||
|
"Unsupported dots note omni video_config fields: "
|
||||||
|
+ ", ".join(sorted(video_config))
|
||||||
|
)
|
||||||
|
sampling_params = request_obj.sampling_params or {}
|
||||||
|
if not isinstance(sampling_params, dict):
|
||||||
|
raise ValueError(
|
||||||
|
"Dots note omni video preprocessing requires one request's "
|
||||||
|
"sampling_params as a dictionary."
|
||||||
|
)
|
||||||
|
max_new_tokens = sampling_params.get("max_new_tokens") or 0
|
||||||
|
loop = asyncio.get_running_loop()
|
||||||
|
preprocess_started = time.perf_counter()
|
||||||
|
video_media = {}
|
||||||
|
total_content_items = 0
|
||||||
|
total_frames = 0
|
||||||
|
total_audio_segments = 0
|
||||||
|
for video_index, video in enumerate(video_data):
|
||||||
|
content = await loop.run_in_executor(
|
||||||
|
self.io_executor,
|
||||||
|
lambda video=video: preprocess_dots_video(
|
||||||
|
video,
|
||||||
|
question,
|
||||||
|
tokenizer=self._tokenizer,
|
||||||
|
seq=seq,
|
||||||
|
audio_cap=audio_cap,
|
||||||
|
audio_sr=audio_sr,
|
||||||
|
k_mode=k_mode,
|
||||||
|
max_new_tokens=max_new_tokens,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
total_content_items += len(content)
|
||||||
|
total_frames += sum(item.get("type") == "image_url" for item in content)
|
||||||
|
total_audio_segments += sum(
|
||||||
|
item.get("type") == "audio_url" for item in content
|
||||||
|
)
|
||||||
|
input_text, media = self._render_video_content(
|
||||||
|
input_text, question, video_index, content
|
||||||
|
)
|
||||||
|
video_media.update(media)
|
||||||
|
|
||||||
|
leftover = self.video_placeholder_regex.search(input_text)
|
||||||
|
if leftover is not None:
|
||||||
|
raise ValueError(
|
||||||
|
"Video placeholder count does not match video_data: "
|
||||||
|
f"{len(video_data)} video(s) given"
|
||||||
|
)
|
||||||
|
input_text, image_data, audio_data = self._merge_video_media(
|
||||||
|
input_text, image_data, audio_data, video_media
|
||||||
|
)
|
||||||
|
preprocess_elapsed = time.perf_counter() - preprocess_started
|
||||||
|
logger.info(
|
||||||
|
"[dots_mm] rid=%s video_preprocess elapsed=%.3fs "
|
||||||
|
"expanded_frames=%d expanded_audio_segments=%d content_items=%d "
|
||||||
|
"after_preprocess images=%d audios=%d",
|
||||||
|
request_obj.rid,
|
||||||
|
preprocess_elapsed,
|
||||||
|
total_frames,
|
||||||
|
total_audio_segments,
|
||||||
|
total_content_items,
|
||||||
|
len(image_data),
|
||||||
|
len(audio_data),
|
||||||
|
)
|
||||||
|
|
||||||
|
base_output = await self.load_mm_data(
|
||||||
|
prompt=input_text,
|
||||||
|
image_data=image_data,
|
||||||
|
audio_data=audio_data,
|
||||||
|
video_data=None,
|
||||||
|
multimodal_tokens=self.mm_tokens,
|
||||||
|
audio_sample_rate=self.audio_processor_config.sampling_rate,
|
||||||
|
)
|
||||||
|
|
||||||
|
pattern = self.mm_tokens.get_combined_regex()
|
||||||
|
parts = re.split(pattern, base_output.input_text)
|
||||||
|
modality_order = [
|
||||||
|
modality
|
||||||
|
for part in parts
|
||||||
|
if (modality := self.mm_tokens.get_modality_of_token(part)) is not None
|
||||||
|
]
|
||||||
|
if modality_order.count(Modality.IMAGE) != len(base_output.images):
|
||||||
|
raise ValueError("Image placeholder count does not match image_data")
|
||||||
|
if modality_order.count(Modality.AUDIO) != len(base_output.audios):
|
||||||
|
raise ValueError("Audio placeholder count does not match audio_data")
|
||||||
|
|
||||||
|
image_features, image_grids, image_token_strings = (
|
||||||
|
self.image_preprocessor.process_images(base_output.images)
|
||||||
|
if base_output.images
|
||||||
|
else ([], [], [])
|
||||||
|
)
|
||||||
|
audio_features = [self._normalize_audio(audio) for audio in base_output.audios]
|
||||||
|
audio_token_strings = [
|
||||||
|
get_audio_token_string(waveform.numel(), self.audio_processor_config)
|
||||||
|
for waveform in audio_features
|
||||||
|
]
|
||||||
|
feature_iters = {
|
||||||
|
Modality.IMAGE: iter(zip(image_features, image_grids, image_token_strings)),
|
||||||
|
Modality.AUDIO: iter(zip(audio_features, audio_token_strings)),
|
||||||
|
}
|
||||||
|
|
||||||
|
input_ids = []
|
||||||
|
mm_items = []
|
||||||
|
add_special_tokens = True
|
||||||
|
for part in parts:
|
||||||
|
modality = self.mm_tokens.get_modality_of_token(part)
|
||||||
|
if modality is None:
|
||||||
|
input_ids.extend(
|
||||||
|
self._tokenizer.encode(part, add_special_tokens=add_special_tokens)
|
||||||
|
)
|
||||||
|
add_special_tokens = False
|
||||||
|
continue
|
||||||
|
|
||||||
|
if modality == Modality.IMAGE:
|
||||||
|
feature, grid_thw, expanded_token_string = next(feature_iters[modality])
|
||||||
|
pad_token_id = self.mm_token_ids["im_token_id"]
|
||||||
|
model_specific_data = {"image_grid_thw": grid_thw.reshape(-1, 3)}
|
||||||
|
else:
|
||||||
|
feature, expanded_token_string = next(feature_iters[modality])
|
||||||
|
pad_token_id = self.mm_token_ids["audio_token_id"]
|
||||||
|
model_specific_data = {}
|
||||||
|
|
||||||
|
item_token_ids = self._tokenizer.encode(
|
||||||
|
expanded_token_string, add_special_tokens=False
|
||||||
|
)
|
||||||
|
item_start = len(input_ids)
|
||||||
|
input_ids.extend(item_token_ids)
|
||||||
|
local_offsets = self.get_mm_items_offset(
|
||||||
|
torch.tensor(item_token_ids), pad_token_id
|
||||||
|
)
|
||||||
|
offsets = [
|
||||||
|
(item_start + start, item_start + end) for start, end in local_offsets
|
||||||
|
]
|
||||||
|
item = MultimodalDataItem(
|
||||||
|
modality=modality,
|
||||||
|
feature=feature,
|
||||||
|
offsets=offsets,
|
||||||
|
model_specific_data=model_specific_data,
|
||||||
|
)
|
||||||
|
item.set_pad_value()
|
||||||
|
mm_items.append(item)
|
||||||
|
|
||||||
|
if len(input_ids) > max_req_input_len:
|
||||||
|
raise ValueError(
|
||||||
|
"Dots note omni expanded prompt is too long: "
|
||||||
|
f"{len(input_ids)} > {max_req_input_len}"
|
||||||
|
)
|
||||||
|
padded_input_ids = MultimodalProcessorOutput.build_padded_input_ids(
|
||||||
|
input_ids, mm_items
|
||||||
|
)
|
||||||
|
return MultimodalProcessorOutput(
|
||||||
|
mm_items=mm_items,
|
||||||
|
input_ids=input_ids,
|
||||||
|
padded_input_ids=padded_input_ids,
|
||||||
|
**self.mm_token_ids,
|
||||||
|
)
|
||||||
@@ -0,0 +1,180 @@
|
|||||||
|
"""Build and render the train-compatible dots video interleave plan."""
|
||||||
|
|
||||||
|
import hashlib
|
||||||
|
import random
|
||||||
|
|
||||||
|
import numpy as np
|
||||||
|
|
||||||
|
from .video_qa_flattener import (
|
||||||
|
_VIDEO_KEY_RE,
|
||||||
|
_VIDEO_MARKER_RE,
|
||||||
|
VideoQAFlattener,
|
||||||
|
_format_timestamp,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _derive_seed(record_key: str) -> int:
|
||||||
|
digest = hashlib.sha1(f"42|flatten|{record_key}".encode()).hexdigest()
|
||||||
|
return int(digest[:8], 16)
|
||||||
|
|
||||||
|
|
||||||
|
def _normalize_conversations(conversations):
|
||||||
|
if conversations is None:
|
||||||
|
return []
|
||||||
|
if isinstance(conversations, np.ndarray):
|
||||||
|
conversations = conversations.tolist()
|
||||||
|
normalized = []
|
||||||
|
for conversation in conversations:
|
||||||
|
if isinstance(conversation, np.ndarray):
|
||||||
|
conversation = conversation.tolist()
|
||||||
|
normalized.append(dict(conversation))
|
||||||
|
return normalized
|
||||||
|
|
||||||
|
|
||||||
|
def _plan_interleaved_emissions(flattener, frames, timestamps, audio_b64, video_dict):
|
||||||
|
"""Plan frame/audio emissions without coupling the policy to an output schema."""
|
||||||
|
duration = float(video_dict.get("audio_duration", 0) or 0)
|
||||||
|
if duration <= 0:
|
||||||
|
duration = float(timestamps[-1]) if timestamps else float(len(frames))
|
||||||
|
|
||||||
|
bounds = flattener._decide_group_bounds(len(frames), duration)
|
||||||
|
try:
|
||||||
|
pcm, sample_rate = flattener._decode_wav_b64(audio_b64)
|
||||||
|
except Exception: # noqa: BLE001 - malformed audio falls back to one block
|
||||||
|
emissions = [
|
||||||
|
{"kind": "frame", "ts": timestamp, "b64": frame}
|
||||||
|
for frame, timestamp in zip(frames, timestamps)
|
||||||
|
]
|
||||||
|
emissions.append({"kind": "audio", "b64": audio_b64, "dur": duration})
|
||||||
|
return emissions
|
||||||
|
|
||||||
|
emissions = []
|
||||||
|
for group in range(len(bounds) - 1):
|
||||||
|
frame_start, frame_end = bounds[group : group + 2]
|
||||||
|
if frame_end <= frame_start:
|
||||||
|
continue
|
||||||
|
time_start = 0.0 if group == 0 else float(timestamps[frame_start])
|
||||||
|
time_end = (
|
||||||
|
duration
|
||||||
|
if group == len(bounds) - 2
|
||||||
|
else float(timestamps[bounds[group + 1]])
|
||||||
|
)
|
||||||
|
if time_end <= time_start:
|
||||||
|
time_end = time_start + duration / max(len(bounds) - 1, 1)
|
||||||
|
emissions.extend(
|
||||||
|
{
|
||||||
|
"kind": "frame",
|
||||||
|
"ts": timestamps[index],
|
||||||
|
"b64": frames[index],
|
||||||
|
}
|
||||||
|
for index in range(frame_start, frame_end)
|
||||||
|
)
|
||||||
|
sample_start = max(0, round(time_start * sample_rate))
|
||||||
|
sample_end = min(len(pcm), round(time_end * sample_rate))
|
||||||
|
if sample_end > sample_start:
|
||||||
|
segment = pcm[sample_start:sample_end]
|
||||||
|
emissions.append(
|
||||||
|
{
|
||||||
|
"kind": "audio",
|
||||||
|
"b64": flattener._encode_wav_b64(segment, sample_rate),
|
||||||
|
"dur": len(segment) / sample_rate,
|
||||||
|
}
|
||||||
|
)
|
||||||
|
return emissions
|
||||||
|
|
||||||
|
|
||||||
|
def build_plan(
|
||||||
|
meta,
|
||||||
|
conversations,
|
||||||
|
record_key: str,
|
||||||
|
*,
|
||||||
|
k_mode: str,
|
||||||
|
process_audio: bool,
|
||||||
|
):
|
||||||
|
"""Create a deterministic intermediate plan for one request."""
|
||||||
|
rng = random.Random(_derive_seed(record_key))
|
||||||
|
flattener = VideoQAFlattener(
|
||||||
|
time_format="hms",
|
||||||
|
audio_interleave=process_audio,
|
||||||
|
ai_k_mode=k_mode,
|
||||||
|
rng=rng,
|
||||||
|
)
|
||||||
|
old_meta = dict(meta) if meta else {}
|
||||||
|
video_pairs = []
|
||||||
|
for key, video in old_meta.items():
|
||||||
|
match = _VIDEO_KEY_RE.fullmatch(key) if isinstance(key, str) else None
|
||||||
|
if match and isinstance(video, dict):
|
||||||
|
video_pairs.append((int(match.group(1)), video))
|
||||||
|
video_pairs.sort()
|
||||||
|
passthrough = {
|
||||||
|
key: value
|
||||||
|
for key, value in old_meta.items()
|
||||||
|
if not (isinstance(key, str) and _VIDEO_KEY_RE.fullmatch(key))
|
||||||
|
}
|
||||||
|
|
||||||
|
videos = []
|
||||||
|
for video_index, video in video_pairs:
|
||||||
|
frames, timestamps, audio_b64 = flattener._subsample_one_video(video)
|
||||||
|
if flattener.audio_interleave and audio_b64 is not None and frames:
|
||||||
|
emissions = _plan_interleaved_emissions(
|
||||||
|
flattener, frames, timestamps, audio_b64, video
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
emissions = [
|
||||||
|
{"kind": "frame", "ts": timestamp, "b64": frame}
|
||||||
|
for frame, timestamp in zip(frames, timestamps)
|
||||||
|
]
|
||||||
|
if audio_b64 is not None:
|
||||||
|
emissions.append({"kind": "audio", "b64": audio_b64})
|
||||||
|
videos.append(
|
||||||
|
{
|
||||||
|
"index": video_index,
|
||||||
|
"time_format": flattener.time_format,
|
||||||
|
"emissions": emissions,
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
return {
|
||||||
|
"seconds_decimals": flattener.seconds_decimals,
|
||||||
|
"videos": videos,
|
||||||
|
"passthrough_meta": passthrough,
|
||||||
|
"conversations": _normalize_conversations(conversations),
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def render_flat(plan):
|
||||||
|
"""Render a plan as globally numbered image/audio markers."""
|
||||||
|
new_meta = dict(plan["passthrough_meta"])
|
||||||
|
next_image = 0
|
||||||
|
next_audio = 0
|
||||||
|
replacements = {}
|
||||||
|
decimals = plan["seconds_decimals"]
|
||||||
|
for video in plan["videos"]:
|
||||||
|
parts = []
|
||||||
|
for emission in video["emissions"]:
|
||||||
|
if emission["kind"] == "frame":
|
||||||
|
key = f"image_{next_image}"
|
||||||
|
new_meta[key] = emission["b64"]
|
||||||
|
timestamp = _format_timestamp(
|
||||||
|
emission["ts"],
|
||||||
|
fmt=video["time_format"],
|
||||||
|
seconds_decimals=decimals,
|
||||||
|
)
|
||||||
|
parts.append(f"<{timestamp}><{key}>")
|
||||||
|
next_image += 1
|
||||||
|
else:
|
||||||
|
key = f"audio_{next_audio}"
|
||||||
|
new_meta[key] = emission["b64"]
|
||||||
|
parts.append(f"<{key}>")
|
||||||
|
next_audio += 1
|
||||||
|
replacements[video["index"]] = "".join(parts)
|
||||||
|
|
||||||
|
conversations = []
|
||||||
|
for conversation in plan["conversations"]:
|
||||||
|
updated = dict(conversation)
|
||||||
|
value = conversation.get("value", "") or ""
|
||||||
|
updated["value"] = _VIDEO_MARKER_RE.sub(
|
||||||
|
lambda match: replacements.get(int(match.group(1)), ""), value
|
||||||
|
)
|
||||||
|
conversations.append(updated)
|
||||||
|
return {"meta": new_meta, "conversations": conversations, "data_type": "mm"}
|
||||||
@@ -0,0 +1,338 @@
|
|||||||
|
"""Train-compatible v2 preprocessing for dots video serving."""
|
||||||
|
|
||||||
|
import base64
|
||||||
|
import io
|
||||||
|
import math
|
||||||
|
import re
|
||||||
|
import wave
|
||||||
|
|
||||||
|
import numpy as np
|
||||||
|
from PIL import Image
|
||||||
|
from torchcodec.decoders import VideoDecoder
|
||||||
|
|
||||||
|
from .v2core import (
|
||||||
|
ALIGN,
|
||||||
|
INTERLEAVE_SEG_MIN_SEC,
|
||||||
|
V2_FPS_CAP,
|
||||||
|
V2_FPS_MIN,
|
||||||
|
V2_OVH,
|
||||||
|
V2_PF_CEIL,
|
||||||
|
V2_PF_FLOOR,
|
||||||
|
compute_target_size,
|
||||||
|
v2_solve_degrade,
|
||||||
|
v2_split_visual_budget,
|
||||||
|
)
|
||||||
|
|
||||||
|
DEFAULT_SYSTEM_PROMPT = "You are a helpful assistant."
|
||||||
|
DEFAULT_AUDIO_SAMPLE_RATE = 16000
|
||||||
|
DEFAULT_AUDIO_SAMPLES_PER_TOKEN = 1280
|
||||||
|
DEFAULT_AUDIO_CHUNK_SEC = 30
|
||||||
|
|
||||||
|
TOK_SYS_START = "<|system|>"
|
||||||
|
TOK_SYS_END = "<|endofsystem|>"
|
||||||
|
AUDIO_WRAP_TOKENS = 2
|
||||||
|
ROLE_WRAP_TOKENS = 2
|
||||||
|
_OVERHEAD_MARGIN = 64
|
||||||
|
|
||||||
|
|
||||||
|
class SkipSample(Exception):
|
||||||
|
"""Raised when a video cannot fit or cannot be decoded."""
|
||||||
|
|
||||||
|
|
||||||
|
def tokenize_len(text, tokenizer):
|
||||||
|
"""Count tokens without adding tokenizer special tokens."""
|
||||||
|
if not text:
|
||||||
|
return 0
|
||||||
|
return len(tokenizer(text, add_special_tokens=False)["input_ids"])
|
||||||
|
|
||||||
|
|
||||||
|
def system_block_tokens(prompt=DEFAULT_SYSTEM_PROMPT, tokenizer=None):
|
||||||
|
if tokenizer is None:
|
||||||
|
raise ValueError("dots video preprocessing requires a tokenizer")
|
||||||
|
return tokenize_len(f"{TOK_SYS_START}{prompt}{TOK_SYS_END}\n", tokenizer=tokenizer)
|
||||||
|
|
||||||
|
|
||||||
|
def audio_block_tokens(
|
||||||
|
duration_sec,
|
||||||
|
samples_per_token=DEFAULT_AUDIO_SAMPLES_PER_TOKEN,
|
||||||
|
chunk_sec=DEFAULT_AUDIO_CHUNK_SEC,
|
||||||
|
sr=DEFAULT_AUDIO_SAMPLE_RATE,
|
||||||
|
):
|
||||||
|
"""Estimate tokens for one independently encoded audio segment."""
|
||||||
|
if duration_sec is None or duration_sec <= 0:
|
||||||
|
return 0
|
||||||
|
total_samples = int(duration_sec * sr)
|
||||||
|
chunk_samples = chunk_sec * sr
|
||||||
|
padding_tokens = 0
|
||||||
|
for position in range(0, total_samples, chunk_samples):
|
||||||
|
length = min(chunk_samples, total_samples - position)
|
||||||
|
padding_tokens += math.ceil(length / samples_per_token)
|
||||||
|
return AUDIO_WRAP_TOKENS + padding_tokens
|
||||||
|
|
||||||
|
|
||||||
|
def conversation_tokens(conversations, tokenizer=None):
|
||||||
|
"""Count conversation text while preserving video-marker boundaries."""
|
||||||
|
if tokenizer is None:
|
||||||
|
raise ValueError("dots video preprocessing requires a tokenizer")
|
||||||
|
total = 0
|
||||||
|
for conversation in conversations:
|
||||||
|
total += ROLE_WRAP_TOKENS
|
||||||
|
parts = re.split(r"<video_(\d+)>", conversation.get("value", "") or "")
|
||||||
|
for index, part in enumerate(parts):
|
||||||
|
if index % 2 == 0:
|
||||||
|
total += tokenize_len(part, tokenizer=tokenizer)
|
||||||
|
else:
|
||||||
|
video_index = int(part)
|
||||||
|
total += tokenize_len(f"<video_{video_index}>", tokenizer=tokenizer)
|
||||||
|
return total
|
||||||
|
|
||||||
|
|
||||||
|
def _make_video_decoder(video_bytes):
|
||||||
|
try:
|
||||||
|
return VideoDecoder(
|
||||||
|
video_bytes,
|
||||||
|
dimension_order="NHWC",
|
||||||
|
num_ffmpeg_threads=1,
|
||||||
|
seek_mode="approximate",
|
||||||
|
)
|
||||||
|
except TypeError:
|
||||||
|
return VideoDecoder(video_bytes, dimension_order="NHWC", num_ffmpeg_threads=1)
|
||||||
|
|
||||||
|
|
||||||
|
def extract_frames_v2(
|
||||||
|
decoder,
|
||||||
|
seq_length,
|
||||||
|
visual_budget,
|
||||||
|
*,
|
||||||
|
pf_floor=V2_PF_FLOOR,
|
||||||
|
pf_ceil=V2_PF_CEIL,
|
||||||
|
fps_cap=V2_FPS_CAP,
|
||||||
|
fps_min=V2_FPS_MIN,
|
||||||
|
overhead=V2_OVH,
|
||||||
|
jpeg_quality=85,
|
||||||
|
):
|
||||||
|
"""Decode, resize, and JPEG-encode frames selected by the v2 policy."""
|
||||||
|
metadata = decoder.metadata
|
||||||
|
duration = float(metadata.duration_seconds or 0)
|
||||||
|
original_height = int(metadata.height)
|
||||||
|
original_width = int(metadata.width)
|
||||||
|
total_frames = int(metadata.num_frames or 0)
|
||||||
|
if duration <= 0 or original_height <= 0 or original_width <= 0:
|
||||||
|
raise ValueError(
|
||||||
|
f"bad metadata: dur={duration} h={original_height} w={original_width}"
|
||||||
|
)
|
||||||
|
original_fps = float(metadata.average_fps or 0) or 25.0
|
||||||
|
if total_frames <= 0:
|
||||||
|
total_frames = max(1, int(duration * original_fps))
|
||||||
|
|
||||||
|
aligned_height = max(ALIGN, round(original_height / ALIGN) * ALIGN)
|
||||||
|
aligned_width = max(ALIGN, round(original_width / ALIGN) * ALIGN)
|
||||||
|
original_patches = (aligned_height // ALIGN) * (aligned_width // ALIGN)
|
||||||
|
num_frames, _, target_patches = v2_solve_degrade(
|
||||||
|
visual_budget,
|
||||||
|
duration,
|
||||||
|
original_patches,
|
||||||
|
original_fps,
|
||||||
|
seq_length,
|
||||||
|
fps_cap=fps_cap,
|
||||||
|
fps_min=fps_min,
|
||||||
|
pf_floor=pf_floor,
|
||||||
|
pf_ceil=pf_ceil,
|
||||||
|
ovh=overhead,
|
||||||
|
orig_h=original_height,
|
||||||
|
orig_w=original_width,
|
||||||
|
)
|
||||||
|
max_pixels = min(target_patches, original_patches) * ALIGN * ALIGN
|
||||||
|
target_height, target_width = compute_target_size(
|
||||||
|
original_height,
|
||||||
|
original_width,
|
||||||
|
pf_floor * ALIGN * ALIGN,
|
||||||
|
max_pixels,
|
||||||
|
)
|
||||||
|
|
||||||
|
num_frames = max(4, min(num_frames, total_frames))
|
||||||
|
if num_frames == 1:
|
||||||
|
indices = [0]
|
||||||
|
else:
|
||||||
|
step = (total_frames - 1) / (num_frames - 1)
|
||||||
|
indices = sorted(
|
||||||
|
{
|
||||||
|
max(0, min(round(index * step), total_frames - 1))
|
||||||
|
for index in range(num_frames)
|
||||||
|
}
|
||||||
|
)
|
||||||
|
try:
|
||||||
|
frames = decoder.get_frames_at(indices=indices).data
|
||||||
|
except (IndexError, RuntimeError):
|
||||||
|
safe_indices = list(indices)
|
||||||
|
while safe_indices and safe_indices[-1] > 0:
|
||||||
|
safe_indices.pop()
|
||||||
|
try:
|
||||||
|
frames = decoder.get_frames_at(indices=safe_indices).data
|
||||||
|
indices = safe_indices
|
||||||
|
break
|
||||||
|
except (IndexError, RuntimeError):
|
||||||
|
continue
|
||||||
|
else:
|
||||||
|
raise
|
||||||
|
|
||||||
|
timestamps = [round(index / original_fps, 3) for index in indices[: len(frames)]]
|
||||||
|
encoded_frames = []
|
||||||
|
for timestamp, frame in zip(timestamps, frames):
|
||||||
|
array = frame.numpy()
|
||||||
|
image = Image.fromarray(array)
|
||||||
|
if image.size != (target_width, target_height):
|
||||||
|
image = image.resize((target_width, target_height), Image.BICUBIC)
|
||||||
|
buffer = io.BytesIO()
|
||||||
|
image.save(buffer, format="JPEG", quality=jpeg_quality)
|
||||||
|
encoded_frames.append(
|
||||||
|
(timestamp, base64.b64encode(buffer.getvalue()).decode("ascii"))
|
||||||
|
)
|
||||||
|
|
||||||
|
actual_fps = round(len(encoded_frames) / duration, 4)
|
||||||
|
return encoded_frames, duration, target_height, target_width, actual_fps
|
||||||
|
|
||||||
|
|
||||||
|
def extract_audio_wav(
|
||||||
|
video_bytes, sample_rate=DEFAULT_AUDIO_SAMPLE_RATE, max_duration_sec=None
|
||||||
|
):
|
||||||
|
"""Extract mono 16-bit WAV audio; return ``(None, 0.0)`` when unavailable."""
|
||||||
|
from torchcodec.decoders import AudioDecoder
|
||||||
|
|
||||||
|
try:
|
||||||
|
samples = AudioDecoder(
|
||||||
|
video_bytes, sample_rate=sample_rate, num_channels=1
|
||||||
|
).get_all_samples()
|
||||||
|
except Exception: # noqa: BLE001 - videos without decodable audio are valid
|
||||||
|
return None, 0.0
|
||||||
|
waveform = samples.data
|
||||||
|
sample_rate = int(samples.sample_rate)
|
||||||
|
if waveform is None or waveform.numel() == 0:
|
||||||
|
return None, 0.0
|
||||||
|
if waveform.dim() == 2:
|
||||||
|
waveform = waveform.mean(dim=0) if waveform.shape[0] > 1 else waveform[0]
|
||||||
|
if max_duration_sec and max_duration_sec > 0:
|
||||||
|
waveform = waveform[: int(max_duration_sec * sample_rate)]
|
||||||
|
pcm = (np.clip(waveform.numpy(), -1.0, 1.0) * 32767.0).astype(np.int16)
|
||||||
|
|
||||||
|
buffer = io.BytesIO()
|
||||||
|
with wave.open(buffer, "wb") as wav:
|
||||||
|
wav.setnchannels(1)
|
||||||
|
wav.setsampwidth(2)
|
||||||
|
wav.setframerate(sample_rate)
|
||||||
|
wav.writeframes(pcm.tobytes())
|
||||||
|
duration = len(pcm) / sample_rate if sample_rate else 0.0
|
||||||
|
return base64.b64encode(buffer.getvalue()).decode("ascii"), duration
|
||||||
|
|
||||||
|
|
||||||
|
def _find_video_keys(conversations):
|
||||||
|
text = " ".join(conversation.get("value", "") for conversation in conversations)
|
||||||
|
return sorted({int(match) for match in re.findall(r"<video_(\d+)>", text)})
|
||||||
|
|
||||||
|
|
||||||
|
def _normalize_sample(sample):
|
||||||
|
meta = dict(sample.get("meta") or {})
|
||||||
|
conversations = sample.get("conversations") or []
|
||||||
|
return meta, [dict(conversation) for conversation in conversations]
|
||||||
|
|
||||||
|
|
||||||
|
def process_sample_video(sample, cfg, *, tokenizer):
|
||||||
|
"""Convert an in-memory video sample into nested frame/audio metadata."""
|
||||||
|
meta, conversations = _normalize_sample(sample)
|
||||||
|
if not conversations:
|
||||||
|
raise SkipSample("empty_conversations")
|
||||||
|
video_keys = _find_video_keys(conversations)
|
||||||
|
if not video_keys:
|
||||||
|
raise SkipSample("no_video_marker_in_conversations")
|
||||||
|
seq_length = cfg["seq_length"]
|
||||||
|
if seq_length <= 0:
|
||||||
|
raise SkipSample("invalid_seq_length")
|
||||||
|
|
||||||
|
video_data = {}
|
||||||
|
audio_data = {}
|
||||||
|
audio_tokens = {}
|
||||||
|
durations = {}
|
||||||
|
for video_index in video_keys:
|
||||||
|
encoded = meta.get(f"video_{video_index}")
|
||||||
|
if not encoded:
|
||||||
|
raise SkipSample(f"no_video_data:video_{video_index}")
|
||||||
|
try:
|
||||||
|
video_bytes = base64.b64decode(encoded)
|
||||||
|
decoder = _make_video_decoder(video_bytes)
|
||||||
|
duration = float(decoder.metadata.duration_seconds or 0)
|
||||||
|
except Exception as exc:
|
||||||
|
raise SkipSample(
|
||||||
|
f"bad_video_meta:video_{video_index}:{type(exc).__name__}:{exc}"
|
||||||
|
) from exc
|
||||||
|
if duration <= 0:
|
||||||
|
raise SkipSample(f"bad_video_duration:video_{video_index}:{duration}")
|
||||||
|
video_data[video_index] = (video_bytes, decoder)
|
||||||
|
durations[video_index] = duration
|
||||||
|
audio_tokens[video_index] = 0
|
||||||
|
|
||||||
|
if cfg["process_audio"]:
|
||||||
|
audio_b64, audio_duration = extract_audio_wav(
|
||||||
|
video_bytes, sample_rate=cfg["audio_sample_rate"]
|
||||||
|
)
|
||||||
|
if audio_b64 and audio_duration > 0:
|
||||||
|
audio_data[video_index] = (audio_b64, audio_duration)
|
||||||
|
tokens = audio_block_tokens(audio_duration, sr=cfg["audio_sample_rate"])
|
||||||
|
if cfg["reserve_interleave"]:
|
||||||
|
frame_upper_bound = max(1, int(duration * V2_FPS_CAP))
|
||||||
|
max_groups = min(
|
||||||
|
frame_upper_bound,
|
||||||
|
max(1, int(audio_duration // INTERLEAVE_SEG_MIN_SEC)),
|
||||||
|
)
|
||||||
|
tokens += 3 * max_groups
|
||||||
|
audio_tokens[video_index] = tokens
|
||||||
|
|
||||||
|
fixed_tokens = (
|
||||||
|
system_block_tokens(tokenizer=tokenizer)
|
||||||
|
+ conversation_tokens(conversations, tokenizer=tokenizer)
|
||||||
|
+ _OVERHEAD_MARGIN
|
||||||
|
)
|
||||||
|
total_audio_tokens = sum(audio_tokens.values())
|
||||||
|
pure_audio_tokens = sum(
|
||||||
|
audio_block_tokens(duration, sr=cfg["audio_sample_rate"])
|
||||||
|
for _, duration in audio_data.values()
|
||||||
|
)
|
||||||
|
audio_cap = cfg["audio_token_ratio_cap"]
|
||||||
|
if audio_cap > 0 and pure_audio_tokens > audio_cap * seq_length:
|
||||||
|
raise SkipSample("audio_token_ratio_exceed")
|
||||||
|
minimum_visual_tokens = len(video_keys) * 4 * (V2_PF_FLOOR + V2_OVH)
|
||||||
|
if fixed_tokens + total_audio_tokens + minimum_visual_tokens > seq_length:
|
||||||
|
reason = "audio_token_ratio_exceed" if audio_data else "input_budget_exhausted"
|
||||||
|
raise SkipSample(reason)
|
||||||
|
|
||||||
|
visual_total = seq_length - fixed_tokens - total_audio_tokens
|
||||||
|
budgets = v2_split_visual_budget(
|
||||||
|
[durations[index] for index in video_keys], visual_total
|
||||||
|
)
|
||||||
|
|
||||||
|
new_meta = {}
|
||||||
|
for video_index, visual_budget in zip(video_keys, budgets):
|
||||||
|
_, decoder = video_data[video_index]
|
||||||
|
try:
|
||||||
|
frames, _, _, _, actual_fps = extract_frames_v2(
|
||||||
|
decoder,
|
||||||
|
seq_length,
|
||||||
|
visual_budget,
|
||||||
|
jpeg_quality=cfg["video_jpeg_quality"],
|
||||||
|
)
|
||||||
|
except Exception as exc:
|
||||||
|
raise SkipSample(
|
||||||
|
f"video_decode_fail:video_{video_index}:{type(exc).__name__}:{exc}"
|
||||||
|
) from exc
|
||||||
|
nested = {"fps": actual_fps}
|
||||||
|
nested.update(
|
||||||
|
{f"image_{index}": encoded for index, (_, encoded) in enumerate(frames)}
|
||||||
|
)
|
||||||
|
if video_index in audio_data:
|
||||||
|
audio_b64, audio_duration = audio_data[video_index]
|
||||||
|
nested.update(
|
||||||
|
audio_0=audio_b64,
|
||||||
|
audio_duration=audio_duration,
|
||||||
|
audio_sample_rate=cfg["audio_sample_rate"],
|
||||||
|
)
|
||||||
|
new_meta[f"video_{video_index}"] = nested
|
||||||
|
return new_meta, conversations
|
||||||
@@ -0,0 +1,126 @@
|
|||||||
|
"""Dependency-free token-budget algorithms for dots v2 video packing."""
|
||||||
|
|
||||||
|
import math
|
||||||
|
|
||||||
|
PATCH_SIZE = 14
|
||||||
|
MERGE_SIZE = 2
|
||||||
|
ALIGN = PATCH_SIZE * MERGE_SIZE
|
||||||
|
|
||||||
|
FPS_MIN_FRAMES = 4
|
||||||
|
TIMESTAMP_TOKENS = 13
|
||||||
|
IMG_WRAP_TOKENS = 2
|
||||||
|
|
||||||
|
V2_FPS_CAP = 1.0
|
||||||
|
V2_FPS_MIN = 0.2
|
||||||
|
V2_PF_FLOOR = 128
|
||||||
|
V2_PF_CEIL = 1024
|
||||||
|
V2_OVH = TIMESTAMP_TOKENS + IMG_WRAP_TOKENS
|
||||||
|
|
||||||
|
INTERLEAVE_SEG_MIN_SEC = 1.0
|
||||||
|
|
||||||
|
|
||||||
|
def compute_target_size(orig_h, orig_w, min_pixels, max_pixels):
|
||||||
|
"""Resize proportionally to aligned dimensions within the pixel budget."""
|
||||||
|
h = max(ALIGN, round(orig_h / ALIGN) * ALIGN)
|
||||||
|
w = max(ALIGN, round(orig_w / ALIGN) * ALIGN)
|
||||||
|
if h * w > max_pixels:
|
||||||
|
beta = math.sqrt(orig_h * orig_w / max_pixels)
|
||||||
|
h = max(ALIGN, math.floor(orig_h / beta / ALIGN) * ALIGN)
|
||||||
|
w = max(ALIGN, math.floor(orig_w / beta / ALIGN) * ALIGN)
|
||||||
|
elif h * w < min_pixels:
|
||||||
|
beta = math.sqrt(min_pixels / max(1, orig_h * orig_w))
|
||||||
|
h = math.ceil(orig_h * beta / ALIGN) * ALIGN
|
||||||
|
w = math.ceil(orig_w * beta / ALIGN) * ALIGN
|
||||||
|
if h * w > max_pixels: # max_pixels first to control the token length
|
||||||
|
beta = math.sqrt(h * w / max_pixels)
|
||||||
|
h = max(ALIGN, math.floor(h / beta / ALIGN) * ALIGN)
|
||||||
|
w = max(ALIGN, math.floor(w / beta / ALIGN) * ALIGN)
|
||||||
|
return int(h), int(w)
|
||||||
|
|
||||||
|
|
||||||
|
def v2_frame_hardcap(seq_length, pf_floor=V2_PF_FLOOR, ovh=V2_OVH):
|
||||||
|
"""Return a power-of-two frame cap derived from the sequence budget."""
|
||||||
|
need = max(1, (seq_length - 2240) // (pf_floor + ovh))
|
||||||
|
if need <= 1024:
|
||||||
|
return 1024
|
||||||
|
p = 1
|
||||||
|
while p < need:
|
||||||
|
p <<= 1
|
||||||
|
return p
|
||||||
|
|
||||||
|
|
||||||
|
def real_patches_at(orig_h, orig_w, pf_cap, pf_floor=V2_PF_FLOOR):
|
||||||
|
"""Return the actual patch count after aligned resizing."""
|
||||||
|
eff_max_px = pf_cap * ALIGN * ALIGN
|
||||||
|
th, tw = compute_target_size(orig_h, orig_w, pf_floor * ALIGN * ALIGN, eff_max_px)
|
||||||
|
return (th // ALIGN) * (tw // ALIGN)
|
||||||
|
|
||||||
|
|
||||||
|
def v2_solve_degrade(
|
||||||
|
visual_budget,
|
||||||
|
duration,
|
||||||
|
orig_max_pf,
|
||||||
|
orig_fps,
|
||||||
|
seq_length,
|
||||||
|
fps_cap=V2_FPS_CAP,
|
||||||
|
fps_min=V2_FPS_MIN,
|
||||||
|
pf_floor=V2_PF_FLOOR,
|
||||||
|
pf_ceil=V2_PF_CEIL,
|
||||||
|
ovh=V2_OVH,
|
||||||
|
orig_h=None,
|
||||||
|
orig_w=None,
|
||||||
|
):
|
||||||
|
"""Jointly reduce frame rate and resolution to fit the visual budget."""
|
||||||
|
fps_cap_eff = min(fps_cap, max(orig_fps, 1e-6))
|
||||||
|
pf_ceil_eff = min(pf_ceil, max(orig_max_pf, pf_floor))
|
||||||
|
hardcap = v2_frame_hardcap(seq_length, pf_floor, ovh)
|
||||||
|
_use_real = orig_h is not None and orig_w is not None
|
||||||
|
|
||||||
|
def _patch_of(pf):
|
||||||
|
if _use_real:
|
||||||
|
return real_patches_at(orig_h, orig_w, max(pf_floor, round(pf)), pf_floor)
|
||||||
|
return round(pf)
|
||||||
|
|
||||||
|
def usage(r):
|
||||||
|
fps = fps_min + r * (fps_cap_eff - fps_min)
|
||||||
|
pf = pf_floor + r * (pf_ceil_eff - pf_floor)
|
||||||
|
nf = max(FPS_MIN_FRAMES, min(round(duration * fps), hardcap))
|
||||||
|
return nf * (_patch_of(pf) + ovh), fps, pf, nf
|
||||||
|
|
||||||
|
if usage(1.0)[0] <= visual_budget:
|
||||||
|
_, fps, pf, nf = usage(1.0)
|
||||||
|
return nf, fps, round(pf)
|
||||||
|
if usage(0.0)[0] > visual_budget:
|
||||||
|
_floor_cost = _patch_of(pf_floor) + ovh
|
||||||
|
nf = max(FPS_MIN_FRAMES, min(visual_budget // _floor_cost, hardcap))
|
||||||
|
fps = nf / max(duration, 1e-6)
|
||||||
|
return nf, round(fps, 4), pf_floor
|
||||||
|
lo, hi = 0.0, 1.0
|
||||||
|
for _ in range(50):
|
||||||
|
mid = (lo + hi) / 2
|
||||||
|
if usage(mid)[0] <= visual_budget:
|
||||||
|
lo = mid
|
||||||
|
else:
|
||||||
|
hi = mid
|
||||||
|
_, fps, pf, nf = usage(lo)
|
||||||
|
return nf, fps, round(pf)
|
||||||
|
|
||||||
|
|
||||||
|
def v2_split_visual_budget(durations, visual_total, pf_floor=V2_PF_FLOOR, ovh=V2_OVH):
|
||||||
|
"""Allocate a shared visual budget by duration with a per-video floor."""
|
||||||
|
n = len(durations)
|
||||||
|
if n <= 1:
|
||||||
|
return [max(pf_floor + ovh, int(visual_total))]
|
||||||
|
floor_each = FPS_MIN_FRAMES * (pf_floor + ovh)
|
||||||
|
floor_total = floor_each * n
|
||||||
|
remain = max(0, int(visual_total) - floor_total)
|
||||||
|
tot_dur = sum(max(0.0, d) for d in durations)
|
||||||
|
out = []
|
||||||
|
if tot_dur <= 0:
|
||||||
|
share = remain // n
|
||||||
|
for _ in range(n):
|
||||||
|
out.append(floor_each + share)
|
||||||
|
else:
|
||||||
|
for d in durations:
|
||||||
|
out.append(floor_each + int(remain * (max(0.0, d) / tot_dur)))
|
||||||
|
return out
|
||||||
+119
@@ -0,0 +1,119 @@
|
|||||||
|
"""Frame sampling and audio interleaving helpers for dots video requests."""
|
||||||
|
|
||||||
|
import base64
|
||||||
|
import io
|
||||||
|
import math
|
||||||
|
import random
|
||||||
|
import re
|
||||||
|
import wave
|
||||||
|
|
||||||
|
import numpy as np
|
||||||
|
|
||||||
|
_VIDEO_KEY_RE = re.compile(r"video_(\d+)")
|
||||||
|
_VIDEO_MARKER_RE = re.compile(r"<video_(\d+)>")
|
||||||
|
|
||||||
|
|
||||||
|
def _format_timestamp(sec: float, fmt: str = "hms", seconds_decimals: int = 1) -> str:
|
||||||
|
"""Format a non-negative timestamp in the model's training format."""
|
||||||
|
sec = max(0.0, sec)
|
||||||
|
if fmt == "seconds":
|
||||||
|
return f"{sec:.{seconds_decimals}f} seconds"
|
||||||
|
total_cs = round(sec * 100)
|
||||||
|
hours = total_cs // (3600 * 100)
|
||||||
|
minutes = (total_cs // (60 * 100)) % 60
|
||||||
|
seconds = (total_cs // 100) % 60
|
||||||
|
centiseconds = total_cs % 100
|
||||||
|
return f"{hours:02d}:{minutes:02d}:{seconds:02d}.{centiseconds:02d}"
|
||||||
|
|
||||||
|
|
||||||
|
def _sorted_image_keys(video_dict: dict) -> list[str]:
|
||||||
|
"""Return image keys ordered by their numeric suffix."""
|
||||||
|
pairs = []
|
||||||
|
for key in video_dict:
|
||||||
|
match = re.fullmatch(r"image_(\d+)", key)
|
||||||
|
if match:
|
||||||
|
pairs.append((int(match.group(1)), key))
|
||||||
|
return [key for _, key in sorted(pairs)]
|
||||||
|
|
||||||
|
|
||||||
|
class VideoQAFlattener:
|
||||||
|
"""Apply the frame and audio sampling policy used during training."""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
time_format: str = "random",
|
||||||
|
seconds_decimals: int = 1,
|
||||||
|
audio_interleave: bool = False,
|
||||||
|
ai_seg_min_sec: float = 1.0,
|
||||||
|
ai_k_mode: str = "eval30",
|
||||||
|
rng: random.Random | None = None,
|
||||||
|
):
|
||||||
|
if time_format not in ("hms", "seconds", "random"):
|
||||||
|
raise ValueError(f"unsupported time format: {time_format!r}")
|
||||||
|
if ai_k_mode not in ("logk", "eval30", "eval_ek", "whole"):
|
||||||
|
raise ValueError(f"unsupported audio interleave mode: {ai_k_mode!r}")
|
||||||
|
|
||||||
|
self.time_format = time_format
|
||||||
|
self.seconds_decimals = max(0, int(seconds_decimals))
|
||||||
|
self.audio_interleave = bool(audio_interleave)
|
||||||
|
self.ai_seg_min_sec = max(1e-6, float(ai_seg_min_sec))
|
||||||
|
self.ai_k_mode = ai_k_mode
|
||||||
|
self.rng = rng or random.Random()
|
||||||
|
|
||||||
|
def _subsample_one_video(self, video_dict: dict):
|
||||||
|
"""Return ordered frames, timestamps, and optional WAV data."""
|
||||||
|
original_fps = float(video_dict.get("fps", 1.0)) or 1.0
|
||||||
|
image_keys = _sorted_image_keys(video_dict)
|
||||||
|
frames = [video_dict[key] for key in image_keys]
|
||||||
|
timestamps = [round(i / original_fps, 3) for i in range(len(image_keys))]
|
||||||
|
return frames, timestamps, video_dict.get("audio_0") or None
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _decode_wav_b64(audio_b64: str):
|
||||||
|
"""Decode a base64 WAV into mono int16 PCM."""
|
||||||
|
with wave.open(io.BytesIO(base64.b64decode(audio_b64)), "rb") as wav:
|
||||||
|
sample_rate = wav.getframerate()
|
||||||
|
channels = wav.getnchannels()
|
||||||
|
pcm = np.frombuffer(wav.readframes(wav.getnframes()), dtype=np.int16)
|
||||||
|
if channels > 1:
|
||||||
|
pcm = pcm.reshape(-1, channels).mean(axis=1).astype(np.int16)
|
||||||
|
return pcm, sample_rate
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _encode_wav_b64(pcm, sample_rate: int):
|
||||||
|
"""Encode mono int16 PCM as a base64 WAV."""
|
||||||
|
buffer = io.BytesIO()
|
||||||
|
with wave.open(buffer, "wb") as wav:
|
||||||
|
wav.setnchannels(1)
|
||||||
|
wav.setsampwidth(2)
|
||||||
|
wav.setframerate(sample_rate)
|
||||||
|
wav.writeframes(pcm.tobytes())
|
||||||
|
return base64.b64encode(buffer.getvalue()).decode("ascii")
|
||||||
|
|
||||||
|
def _decide_group_bounds(self, n_frames: int, duration: float):
|
||||||
|
"""Split frames into the configured number of audio windows."""
|
||||||
|
if n_frames <= 1 or duration <= 0:
|
||||||
|
return [0, n_frames] if n_frames else [0, 0]
|
||||||
|
k_max = min(n_frames, max(1, int(duration // self.ai_seg_min_sec)))
|
||||||
|
if self.ai_k_mode == "whole" or k_max <= 1:
|
||||||
|
groups = 1
|
||||||
|
elif self.ai_k_mode == "eval30":
|
||||||
|
groups = round(math.sqrt(k_max))
|
||||||
|
elif self.ai_k_mode == "eval_ek":
|
||||||
|
groups = round((k_max - 1) / math.log(k_max))
|
||||||
|
else:
|
||||||
|
groups = round(math.exp(random.uniform(0.0, math.log(k_max))))
|
||||||
|
groups = max(1, min(k_max, groups))
|
||||||
|
if groups == 1:
|
||||||
|
return [0, n_frames]
|
||||||
|
if self.ai_k_mode == "logk":
|
||||||
|
cuts = sorted(self.rng.sample(range(1, n_frames), groups - 1))
|
||||||
|
else:
|
||||||
|
cuts = sorted(
|
||||||
|
{
|
||||||
|
cut
|
||||||
|
for i in range(1, groups)
|
||||||
|
if 0 < (cut := round(i * n_frames / groups)) < n_frames
|
||||||
|
}
|
||||||
|
)
|
||||||
|
return [0, *cuts, n_frames]
|
||||||
@@ -1884,6 +1884,7 @@ class ReasoningParser:
|
|||||||
"deepseek-r1": DeepSeekR1Detector,
|
"deepseek-r1": DeepSeekR1Detector,
|
||||||
"deepseek-v3": _DeepSeekV3Detector,
|
"deepseek-v3": _DeepSeekV3Detector,
|
||||||
"deepseek-v4": DeepSeekV4Detector,
|
"deepseek-v4": DeepSeekV4Detector,
|
||||||
|
"dots": Qwen3Detector,
|
||||||
"glm45": Glm45Detector,
|
"glm45": Glm45Detector,
|
||||||
"hunyuan": HunyuanDetector,
|
"hunyuan": HunyuanDetector,
|
||||||
"gpt-oss": GptOssDetector,
|
"gpt-oss": GptOssDetector,
|
||||||
|
|||||||
@@ -5371,6 +5371,7 @@ class ServerArgs:
|
|||||||
"PixtralForConditionalGeneration",
|
"PixtralForConditionalGeneration",
|
||||||
"GlmMoeDsaForCausalLM",
|
"GlmMoeDsaForCausalLM",
|
||||||
"LongcatFlashForCausalLM",
|
"LongcatFlashForCausalLM",
|
||||||
|
"Dots3NoteForCausalLM",
|
||||||
]:
|
]:
|
||||||
# Set attention backend for DeepSeek
|
# Set attention backend for DeepSeek
|
||||||
if is_deepseek_dsa(hf_config): # DeepSeek 3.2/GLM 5
|
if is_deepseek_dsa(hf_config): # DeepSeek 3.2/GLM 5
|
||||||
@@ -7855,6 +7856,7 @@ class ServerArgs:
|
|||||||
"Qwen3OmniMoeForConditionalGeneration",
|
"Qwen3OmniMoeForConditionalGeneration",
|
||||||
"Qwen2AudioForConditionalGeneration",
|
"Qwen2AudioForConditionalGeneration",
|
||||||
"Qwen2_5OmniForConditionalGeneration",
|
"Qwen2_5OmniForConditionalGeneration",
|
||||||
|
"Dots3NoteForCausalLM",
|
||||||
"KimiVLForConditionalGeneration",
|
"KimiVLForConditionalGeneration",
|
||||||
"KimiK25ForConditionalGeneration",
|
"KimiK25ForConditionalGeneration",
|
||||||
"KimiK3ForConditionalGeneration",
|
"KimiK3ForConditionalGeneration",
|
||||||
@@ -7862,7 +7864,8 @@ class ServerArgs:
|
|||||||
]:
|
]:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
f"Model type {model_arch} is not supported for encoder disaggregation. "
|
f"Model type {model_arch} is not supported for encoder disaggregation. "
|
||||||
f"Supported architectures: Qwen2VL, Qwen3VL, Qwen3.5, InternS2, Qwen2Audio, Qwen2.5Omni, Kimi, MiMoV2."
|
f"Supported architectures: Qwen2VL, Qwen3VL, Qwen3.5, InternS2, "
|
||||||
|
f"Qwen2Audio, Qwen2.5Omni, Dots3-Note, Kimi, MiMoV2."
|
||||||
)
|
)
|
||||||
|
|
||||||
def _validate_ib_devices(self, device_str: Optional[str]) -> Optional[str]:
|
def _validate_ib_devices(self, device_str: Optional[str]) -> Optional[str]:
|
||||||
|
|||||||
@@ -60,6 +60,14 @@ class DraftBackendFactory:
|
|||||||
|
|
||||||
stamp, backend = backend_map[backend_type]()
|
stamp, backend = backend_map[backend_type]()
|
||||||
if backend is not None:
|
if backend is not None:
|
||||||
|
if stamps_children:
|
||||||
|
from sglang.srt.layers.attention.attention_registry import (
|
||||||
|
attn_backend_wrapper_for_draft_decode,
|
||||||
|
)
|
||||||
|
|
||||||
|
backend = attn_backend_wrapper_for_draft_decode(
|
||||||
|
self.draft_model_runner, backend
|
||||||
|
)
|
||||||
backend.prefill_attention_backend_str = stamp
|
backend.prefill_attention_backend_str = stamp
|
||||||
backend.decode_attention_backend_str = stamp
|
backend.decode_attention_backend_str = stamp
|
||||||
if stamps_children:
|
if stamps_children:
|
||||||
|
|||||||
@@ -620,10 +620,15 @@ class EAGLEDraftCudaGraphRunner(DecodeCudaGraphRunner):
|
|||||||
raw_seq_lens_sum = forward_batch.seq_lens_sum
|
raw_seq_lens_sum = forward_batch.seq_lens_sum
|
||||||
|
|
||||||
if bs != raw_bs:
|
if bs != raw_bs:
|
||||||
|
raw_out_cache_loc = forward_batch.out_cache_loc
|
||||||
forward_batch.batch_size = bs
|
forward_batch.batch_size = bs
|
||||||
forward_batch.seq_lens = buffers.seq_lens[:bs]
|
forward_batch.seq_lens = buffers.seq_lens[:bs]
|
||||||
forward_batch.req_pool_indices = buffers.req_pool_indices[:bs]
|
forward_batch.req_pool_indices = buffers.req_pool_indices[:bs]
|
||||||
forward_batch.positions = buffers.positions[:num_tokens]
|
forward_batch.positions = buffers.positions[:num_tokens]
|
||||||
|
# Match out_cache_loc to the padded graph batch for metadata replay.
|
||||||
|
forward_batch.out_cache_loc = buffers.out_cache_loc[
|
||||||
|
: num_tokens * self.speculative_num_steps
|
||||||
|
]
|
||||||
if raw_seq_lens_sum is not None:
|
if raw_seq_lens_sum is not None:
|
||||||
forward_batch.seq_lens_sum = (
|
forward_batch.seq_lens_sum = (
|
||||||
raw_seq_lens_sum + (bs - raw_bs) * self.seq_len_fill_value
|
raw_seq_lens_sum + (bs - raw_bs) * self.seq_len_fill_value
|
||||||
@@ -674,5 +679,6 @@ class EAGLEDraftCudaGraphRunner(DecodeCudaGraphRunner):
|
|||||||
if forward_batch.seq_lens_cpu is not None:
|
if forward_batch.seq_lens_cpu is not None:
|
||||||
forward_batch.seq_lens_cpu = buffers.seq_lens_cpu[:raw_bs]
|
forward_batch.seq_lens_cpu = buffers.seq_lens_cpu[:raw_bs]
|
||||||
forward_batch.seq_lens_sum = raw_seq_lens_sum
|
forward_batch.seq_lens_sum = raw_seq_lens_sum
|
||||||
|
forward_batch.out_cache_loc = raw_out_cache_loc
|
||||||
|
|
||||||
return out
|
return out
|
||||||
|
|||||||
@@ -564,6 +564,9 @@ class EagleDraftWorker(EagleDraftWorkerBase):
|
|||||||
def draft_forward(self, forward_batch: ForwardBatch):
|
def draft_forward(self, forward_batch: ForwardBatch):
|
||||||
# Parse args
|
# Parse args
|
||||||
spec_info: EagleDraftInput = forward_batch.spec_info
|
spec_info: EagleDraftInput = forward_batch.spec_info
|
||||||
|
if forward_batch.forward_mode.is_idle():
|
||||||
|
return self._draft_forward_idle(forward_batch, spec_info)
|
||||||
|
|
||||||
out_cache_loc = forward_batch.out_cache_loc
|
out_cache_loc = forward_batch.out_cache_loc
|
||||||
topk_p, topk_index, hidden_states = (
|
topk_p, topk_index, hidden_states = (
|
||||||
spec_info.topk_p,
|
spec_info.topk_p,
|
||||||
@@ -730,6 +733,38 @@ class EagleDraftWorker(EagleDraftWorkerBase):
|
|||||||
|
|
||||||
return parent_list, top_scores_index, draft_tokens, draft_probs
|
return parent_list, top_scores_index, draft_tokens, draft_probs
|
||||||
|
|
||||||
|
def _draft_forward_idle(
|
||||||
|
self, forward_batch: ForwardBatch, spec_info: EagleDraftInput
|
||||||
|
):
|
||||||
|
"""Run eager idle-rank collectives without materializing draft state."""
|
||||||
|
input_ids = forward_batch.input_ids
|
||||||
|
out_cache_loc = forward_batch.out_cache_loc
|
||||||
|
hidden_states = spec_info.hidden_states
|
||||||
|
|
||||||
|
# ModelRunner pads and unpads the empty batch on every call. Avoid the
|
||||||
|
# normal tree/cache-layout path: idle outputs are discarded when the
|
||||||
|
# verify input is built, but every rank must still enter each forward.
|
||||||
|
for i in range(self.speculative_num_steps - 1):
|
||||||
|
forward_batch.input_ids = input_ids
|
||||||
|
forward_batch.out_cache_loc = out_cache_loc
|
||||||
|
spec_info.hidden_states = hidden_states
|
||||||
|
canary_index_ctx = (
|
||||||
|
c.with_active_single_forward_manager(i)
|
||||||
|
if (c := self.draft_runner.canary_manager) is not None
|
||||||
|
else contextlib.nullcontext()
|
||||||
|
)
|
||||||
|
with (
|
||||||
|
forward_context(
|
||||||
|
ForwardContext(
|
||||||
|
attn_backend=self.draft_attn_backend.attn_backends[i]
|
||||||
|
)
|
||||||
|
),
|
||||||
|
canary_index_ctx,
|
||||||
|
):
|
||||||
|
self.draft_runner.forward(forward_batch)
|
||||||
|
|
||||||
|
return None, None, None, None
|
||||||
|
|
||||||
def draft_extend(self):
|
def draft_extend(self):
|
||||||
pass
|
pass
|
||||||
|
|
||||||
|
|||||||
@@ -27,6 +27,7 @@ from sglang.srt.configs import (
|
|||||||
ChatGLMConfig,
|
ChatGLMConfig,
|
||||||
DbrxConfig,
|
DbrxConfig,
|
||||||
DeepseekVL2Config,
|
DeepseekVL2Config,
|
||||||
|
Dots3Config,
|
||||||
DotsOCRConfig,
|
DotsOCRConfig,
|
||||||
DotsVLMConfig,
|
DotsVLMConfig,
|
||||||
ExaoneConfig,
|
ExaoneConfig,
|
||||||
@@ -112,6 +113,7 @@ _CONFIG_REGISTRY: Dict[str, Type[PretrainedConfig]] = {
|
|||||||
GraniteMoeHybridConfig,
|
GraniteMoeHybridConfig,
|
||||||
DotsVLMConfig,
|
DotsVLMConfig,
|
||||||
DotsOCRConfig,
|
DotsOCRConfig,
|
||||||
|
Dots3Config,
|
||||||
NemotronH_Nano_VL_V2_Config,
|
NemotronH_Nano_VL_V2_Config,
|
||||||
NemotronH_Nano_Omni_Reasoning_V3_Config,
|
NemotronH_Nano_Omni_Reasoning_V3_Config,
|
||||||
NemotronHConfig,
|
NemotronHConfig,
|
||||||
@@ -139,6 +141,7 @@ _CONFIG_REGISTRY: Dict[str, Type[PretrainedConfig]] = {
|
|||||||
]
|
]
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
# DeepSeek V3.2 / V4 reuse the V3 config schema. Subclass the upstream
|
# DeepSeek V3.2 / V4 reuse the V3 config schema. Subclass the upstream
|
||||||
# transformers class with each model_type so AutoConfig.register passes its
|
# transformers class with each model_type so AutoConfig.register passes its
|
||||||
# consistency check (which requires class.model_type == registered key).
|
# consistency check (which requires class.model_type == registered key).
|
||||||
|
|||||||
@@ -8,11 +8,19 @@ cuda-graph buffer plumbing is covered by the backend SWA integration tests.
|
|||||||
|
|
||||||
import unittest
|
import unittest
|
||||||
from types import SimpleNamespace
|
from types import SimpleNamespace
|
||||||
|
from unittest.mock import MagicMock
|
||||||
|
|
||||||
|
from sglang.test.test_utils import maybe_stub_sgl_kernel
|
||||||
|
|
||||||
|
maybe_stub_sgl_kernel()
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
from sglang.srt.mem_cache.memory_pool import KVWriteLoc
|
from sglang.srt.layers.attention.dots_hybrid_backend import DotsSWAMLAAttnBackend
|
||||||
|
from sglang.srt.layers.attention.flashattention_backend import FlashAttentionBackend
|
||||||
|
from sglang.srt.mem_cache.memory_pool import KVWriteLoc, MLATokenToKVPool
|
||||||
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool
|
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool
|
||||||
|
from sglang.srt.model_executor.forward_batch_info import ForwardMode
|
||||||
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
|
||||||
|
|
||||||
@@ -78,6 +86,78 @@ class TestSWAKVPoolSetKVBuffer(CustomTestCase):
|
|||||||
)
|
)
|
||||||
self.assertIs(recorded["full_loc"], loc)
|
self.assertIs(recorded["full_loc"], loc)
|
||||||
|
|
||||||
|
def test_composed_mla_pools_route_local_layer_ids(self):
|
||||||
|
pool = object.__new__(SWAKVPool)
|
||||||
|
pool.layers_mapping = {7: (1, False), 8: (2, True)}
|
||||||
|
recorded = {}
|
||||||
|
|
||||||
|
def make_mla_pool(name):
|
||||||
|
mla_pool = object.__new__(MLATokenToKVPool)
|
||||||
|
|
||||||
|
def set_kv(layer, loc, k, v, layer_id_override=None):
|
||||||
|
recorded[f"{name}_kv"] = (layer, loc, layer_id_override)
|
||||||
|
|
||||||
|
def set_mla(layer, loc, k_nope, k_rope, layer_id_override=None):
|
||||||
|
recorded[f"{name}_mla"] = (layer, loc, layer_id_override)
|
||||||
|
|
||||||
|
mla_pool.set_kv_buffer = set_kv
|
||||||
|
mla_pool.set_mla_kv_buffer = set_mla
|
||||||
|
return mla_pool
|
||||||
|
|
||||||
|
pool.full_kv_pool = make_mla_pool("full")
|
||||||
|
pool.swa_kv_pool = make_mla_pool("swa")
|
||||||
|
full_loc = torch.tensor([3, 4])
|
||||||
|
swa_loc = torch.tensor([7, 8])
|
||||||
|
|
||||||
|
pool.set_kv_buffer(
|
||||||
|
SimpleNamespace(layer_id=7), KVWriteLoc(full_loc, swa_loc), None, None
|
||||||
|
)
|
||||||
|
pool.set_mla_kv_buffer(
|
||||||
|
SimpleNamespace(layer_id=8), KVWriteLoc(full_loc, swa_loc), None, None
|
||||||
|
)
|
||||||
|
|
||||||
|
full_layer, recorded_full_loc, full_layer_id = recorded["full_kv"]
|
||||||
|
swa_layer, recorded_swa_loc, swa_layer_id = recorded["swa_mla"]
|
||||||
|
self.assertIsNone(full_layer)
|
||||||
|
self.assertIs(recorded_full_loc, full_loc)
|
||||||
|
self.assertEqual(full_layer_id, 1)
|
||||||
|
self.assertIsNone(swa_layer)
|
||||||
|
self.assertIs(recorded_swa_loc, swa_loc)
|
||||||
|
self.assertEqual(swa_layer_id, 2)
|
||||||
|
|
||||||
|
|
||||||
|
class TestDotsDraftSWAOutCacheLoc(CustomTestCase):
|
||||||
|
def test_metadata_sees_only_current_step_and_forward_batch_is_restored(self):
|
||||||
|
backend = object.__new__(FlashAttentionBackend)
|
||||||
|
backend.topk = 2
|
||||||
|
backend.speculative_num_steps = 3
|
||||||
|
backend.speculative_step_id = 1
|
||||||
|
|
||||||
|
seen = []
|
||||||
|
backend.init_forward_metadata_out_graph = MagicMock(
|
||||||
|
side_effect=lambda forward_batch, in_capture=False: seen.append(
|
||||||
|
forward_batch.out_cache_loc.clone()
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
wrapper = object.__new__(DotsSWAMLAAttnBackend)
|
||||||
|
wrapper.backend = backend
|
||||||
|
wrapper._active_backend = backend
|
||||||
|
wrapper._prefill_metadata = None
|
||||||
|
|
||||||
|
original = torch.arange(12)
|
||||||
|
forward_batch = SimpleNamespace(
|
||||||
|
batch_size=2,
|
||||||
|
forward_mode=ForwardMode.DECODE,
|
||||||
|
out_cache_loc=original,
|
||||||
|
spec_info=object(),
|
||||||
|
)
|
||||||
|
|
||||||
|
wrapper.init_forward_metadata_out_graph(forward_batch)
|
||||||
|
|
||||||
|
torch.testing.assert_close(seen[0], torch.tensor([1, 4, 7, 10]))
|
||||||
|
self.assertIs(forward_batch.out_cache_loc, original)
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
@@ -0,0 +1,214 @@
|
|||||||
|
import json
|
||||||
|
import unittest
|
||||||
|
|
||||||
|
from sglang.srt.entrypoints.openai.protocol import Function, Tool
|
||||||
|
from sglang.srt.function_call.dots_detector import DotsToolDetector
|
||||||
|
from sglang.srt.function_call.function_call_parser import FunctionCallParser
|
||||||
|
from sglang.srt.parser.reasoning_parser import Qwen3Detector, ReasoningParser
|
||||||
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
|
|
||||||
|
register_cpu_ci(est_time=5, suite="base-a-test-cpu")
|
||||||
|
|
||||||
|
|
||||||
|
def _tool(name: str, properties: dict) -> Tool:
|
||||||
|
return Tool(
|
||||||
|
type="function",
|
||||||
|
function=Function(
|
||||||
|
name=name,
|
||||||
|
description="test tool",
|
||||||
|
parameters={"type": "object", "properties": properties},
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class TestDotsToolDetector(unittest.TestCase):
|
||||||
|
def test_dots_parsers_are_registered(self):
|
||||||
|
self.assertIs(ReasoningParser.DetectorMap["dots"], Qwen3Detector)
|
||||||
|
self.assertIs(FunctionCallParser.ToolCallParserEnum["dots"], DotsToolDetector)
|
||||||
|
|
||||||
|
def test_dots_reasoning_uses_qwen3_format(self):
|
||||||
|
parser = ReasoningParser("dots", stream_reasoning=False, force_reasoning=True)
|
||||||
|
reasoning, content = parser.parse_non_stream(
|
||||||
|
"Need to inspect inputs.</think>Final answer"
|
||||||
|
)
|
||||||
|
self.assertEqual(reasoning, "Need to inspect inputs.")
|
||||||
|
self.assertEqual(content, "Final answer")
|
||||||
|
|
||||||
|
def test_non_stream_xml_converts_schema_types_and_resolves_ref(self):
|
||||||
|
tool = Tool(
|
||||||
|
type="function",
|
||||||
|
function=Function(
|
||||||
|
name="set_location",
|
||||||
|
description="Set location",
|
||||||
|
parameters={
|
||||||
|
"type": "object",
|
||||||
|
"properties": {
|
||||||
|
"location": {"$ref": "#/$defs/Location"},
|
||||||
|
"days": {"type": "integer"},
|
||||||
|
"include_weather": {"type": "boolean"},
|
||||||
|
},
|
||||||
|
"$defs": {
|
||||||
|
"Location": {
|
||||||
|
"type": "object",
|
||||||
|
"properties": {"city": {"type": "string"}},
|
||||||
|
}
|
||||||
|
},
|
||||||
|
},
|
||||||
|
),
|
||||||
|
)
|
||||||
|
parser = FunctionCallParser([tool], "dots")
|
||||||
|
text = (
|
||||||
|
"ok<dots_function_call>"
|
||||||
|
'<invoke name="set_location">'
|
||||||
|
'<parameter name="location">{"city": "Shanghai"}</parameter>'
|
||||||
|
'<parameter name="days">3</parameter>'
|
||||||
|
'<parameter name="include_weather">true</parameter>'
|
||||||
|
"</invoke>"
|
||||||
|
"</dots_function_call>"
|
||||||
|
)
|
||||||
|
|
||||||
|
normal_text, calls = parser.parse_non_stream(text)
|
||||||
|
|
||||||
|
self.assertEqual(normal_text, "ok")
|
||||||
|
self.assertEqual(len(calls), 1)
|
||||||
|
self.assertEqual(calls[0].name, "set_location")
|
||||||
|
self.assertEqual(
|
||||||
|
json.loads(calls[0].parameters),
|
||||||
|
{
|
||||||
|
"location": {"city": "Shanghai"},
|
||||||
|
"days": 3,
|
||||||
|
"include_weather": True,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_non_stream_supports_multiple_invokes_and_json_fallback(self):
|
||||||
|
tools = [
|
||||||
|
_tool("search", {"query": {"type": "string"}}),
|
||||||
|
_tool("open", {"id": {"type": "integer"}}),
|
||||||
|
]
|
||||||
|
parser = FunctionCallParser(tools, "dots")
|
||||||
|
text = (
|
||||||
|
"<dots_function_call>"
|
||||||
|
'<invoke name="search"><parameter name="query">chairs</parameter></invoke>'
|
||||||
|
'<invoke name="open"><parameter name="id">7</parameter></invoke>'
|
||||||
|
"</dots_function_call>"
|
||||||
|
'<dots_function_call>{"name":"search","arguments":{"query":"tables"}}'
|
||||||
|
"</dots_function_call>"
|
||||||
|
)
|
||||||
|
|
||||||
|
_, calls = parser.parse_non_stream(text)
|
||||||
|
|
||||||
|
self.assertEqual([call.name for call in calls], ["search", "open", "search"])
|
||||||
|
self.assertEqual(
|
||||||
|
[json.loads(call.parameters) for call in calls],
|
||||||
|
[
|
||||||
|
{"query": "chairs"},
|
||||||
|
{"id": 7},
|
||||||
|
{"query": "tables"},
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_streaming_buffers_partial_marker_and_emits_all_complete_calls(self):
|
||||||
|
tools = [_tool("search", {"query": {"type": "string"}})]
|
||||||
|
detector = DotsToolDetector()
|
||||||
|
chunks = [
|
||||||
|
"visible<dots_func",
|
||||||
|
(
|
||||||
|
"tion_call>"
|
||||||
|
'<invoke name="search"><parameter name="query">chairs</parameter></invoke>'
|
||||||
|
"</dots_function_call>"
|
||||||
|
"<dots_function_call>"
|
||||||
|
'<invoke name="search"><parameter name="query">tables</parameter></invoke>'
|
||||||
|
"</dots_function_call>"
|
||||||
|
),
|
||||||
|
]
|
||||||
|
|
||||||
|
results = [detector.parse_streaming_increment(chunk, tools) for chunk in chunks]
|
||||||
|
|
||||||
|
self.assertEqual("".join(result.normal_text for result in results), "visible")
|
||||||
|
calls = [call for result in results for call in result.calls]
|
||||||
|
self.assertEqual([call.tool_index for call in calls], [0, 1])
|
||||||
|
self.assertEqual(
|
||||||
|
[json.loads(call.parameters) for call in calls],
|
||||||
|
[
|
||||||
|
{"query": "chairs"},
|
||||||
|
{"query": "tables"},
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_streaming_filters_unknown_tools_and_surfaces_the_content(self):
|
||||||
|
tools = [_tool("search", {"query": {"type": "string"}})]
|
||||||
|
detector = DotsToolDetector()
|
||||||
|
text = (
|
||||||
|
"<dots_function_call>"
|
||||||
|
'<invoke name="ghost"><parameter name="query">chairs</parameter></invoke>'
|
||||||
|
"</dots_function_call>"
|
||||||
|
)
|
||||||
|
|
||||||
|
result = detector.parse_streaming_increment(text, tools)
|
||||||
|
|
||||||
|
self.assertEqual(result.calls, [])
|
||||||
|
self.assertIn("ghost", result.normal_text)
|
||||||
|
self.assertEqual(detector._buffer, "")
|
||||||
|
|
||||||
|
def test_streaming_malformed_block_does_not_block_a_later_valid_call(self):
|
||||||
|
tools = [_tool("search", {"query": {"type": "string"}})]
|
||||||
|
detector = DotsToolDetector()
|
||||||
|
|
||||||
|
malformed = detector.parse_streaming_increment(
|
||||||
|
"<dots_function_call>garbage</dots_function_call>", tools
|
||||||
|
)
|
||||||
|
valid = detector.parse_streaming_increment(
|
||||||
|
"<dots_function_call>"
|
||||||
|
'<invoke name="search"><parameter name="query">chairs</parameter></invoke>'
|
||||||
|
"</dots_function_call>",
|
||||||
|
tools,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(malformed.calls, [])
|
||||||
|
self.assertEqual(malformed.normal_text, "garbage")
|
||||||
|
self.assertEqual([call.name for call in valid.calls], ["search"])
|
||||||
|
|
||||||
|
def test_streaming_strips_stray_end_marker_from_normal_text(self):
|
||||||
|
detector = DotsToolDetector()
|
||||||
|
|
||||||
|
result = detector.parse_streaming_increment(
|
||||||
|
"some text </dots_function_call>", []
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(result.calls, [])
|
||||||
|
self.assertEqual(result.normal_text, "some text ")
|
||||||
|
|
||||||
|
def test_streaming_flushes_partial_opening_marker_at_eof(self):
|
||||||
|
tools = [_tool("search", {"query": {"type": "string"}})]
|
||||||
|
detector = DotsToolDetector()
|
||||||
|
|
||||||
|
result = detector.parse_streaming_increment("answer <dots_func", tools)
|
||||||
|
|
||||||
|
self.assertEqual(result.calls, [])
|
||||||
|
self.assertEqual(result.normal_text, "answer ")
|
||||||
|
self.assertEqual(detector.flush_pending_normal_text(), "<dots_func")
|
||||||
|
self.assertEqual(detector.flush_pending_normal_text(), "")
|
||||||
|
|
||||||
|
def test_streaming_emits_complete_json_body_before_end_marker_without_duplication(
|
||||||
|
self,
|
||||||
|
):
|
||||||
|
tools = [_tool("search", {"query": {"type": "string"}})]
|
||||||
|
detector = DotsToolDetector()
|
||||||
|
|
||||||
|
opening = detector.parse_streaming_increment("<dots_function_call>", tools)
|
||||||
|
body = detector.parse_streaming_increment(
|
||||||
|
'{"name":"search","arguments":{"query":"chairs"}}', tools
|
||||||
|
)
|
||||||
|
closing = detector.parse_streaming_increment("</dots_function_call>", tools)
|
||||||
|
|
||||||
|
self.assertEqual(opening.calls, [])
|
||||||
|
self.assertEqual([call.name for call in body.calls], ["search", None])
|
||||||
|
self.assertEqual(
|
||||||
|
"".join(call.parameters for call in body.calls), '{"query": "chairs"}'
|
||||||
|
)
|
||||||
|
self.assertEqual(closing.calls, [])
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
@@ -0,0 +1,135 @@
|
|||||||
|
import sys
|
||||||
|
from types import SimpleNamespace
|
||||||
|
from unittest.mock import MagicMock
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from sglang.test.test_utils import maybe_stub_sgl_kernel
|
||||||
|
|
||||||
|
maybe_stub_sgl_kernel()
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from sglang.srt.layers.attention.dots_hybrid_backend import (
|
||||||
|
DotsHybridAttnBackend,
|
||||||
|
DotsSWAMLAAttnBackend,
|
||||||
|
_metadata_mismatches_dp_padded_batch,
|
||||||
|
_normalize_cache_seqlens_rows,
|
||||||
|
)
|
||||||
|
from sglang.srt.layers.attention.flashattention_backend import FlashAttentionMetadata
|
||||||
|
from sglang.srt.layers.attention.swa_mla_fallback.ops import (
|
||||||
|
gather_page64_kv_latent,
|
||||||
|
)
|
||||||
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
|
|
||||||
|
register_cpu_ci(est_time=2, suite="base-a-test-cpu")
|
||||||
|
|
||||||
|
|
||||||
|
def _batch(*, bs: int, num_tokens: int, original_bs: int | None = None):
|
||||||
|
return SimpleNamespace(
|
||||||
|
batch_size=bs,
|
||||||
|
out_cache_loc=torch.zeros(num_tokens, dtype=torch.int64),
|
||||||
|
forward_mode=SimpleNamespace(),
|
||||||
|
_original_batch_size=original_bs,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _fa_metadata(*, bs: int, num_tokens: int):
|
||||||
|
return FlashAttentionMetadata(
|
||||||
|
page_table=torch.zeros((bs, 4), dtype=torch.int32),
|
||||||
|
swa_page_table=torch.zeros((bs, 4), dtype=torch.int32),
|
||||||
|
cache_seqlens_int32=torch.ones(bs, dtype=torch.int32),
|
||||||
|
swa_out_cache_loc=torch.zeros(num_tokens, dtype=torch.int64),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_mismatch_detects_short_page_table_and_swa_loc():
|
||||||
|
metadata = _fa_metadata(bs=1, num_tokens=1)
|
||||||
|
assert _metadata_mismatches_dp_padded_batch(metadata, _batch(bs=2, num_tokens=2))
|
||||||
|
|
||||||
|
metadata = _fa_metadata(bs=2, num_tokens=2)
|
||||||
|
assert not _metadata_mismatches_dp_padded_batch(
|
||||||
|
metadata, _batch(bs=2, num_tokens=2)
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_swa_backend_rebuilds_when_dp_padding_changes_rows():
|
||||||
|
inner = SimpleNamespace(
|
||||||
|
forward_metadata=_fa_metadata(bs=1, num_tokens=1),
|
||||||
|
init_forward_metadata=MagicMock(),
|
||||||
|
)
|
||||||
|
backend = object.__new__(DotsSWAMLAAttnBackend)
|
||||||
|
backend.backend = inner
|
||||||
|
backend._active_backend = inner
|
||||||
|
backend._prefill_metadata = None
|
||||||
|
backend._dp_rebuilt_batch_id = None
|
||||||
|
backend.init_forward_metadata = MagicMock()
|
||||||
|
|
||||||
|
backend.maybe_rebuild_metadata_after_dp_padding(_batch(bs=2, num_tokens=2))
|
||||||
|
backend.init_forward_metadata.assert_called_once()
|
||||||
|
|
||||||
|
|
||||||
|
def test_hybrid_rebuilds_when_dp_padding_changes_batch_size():
|
||||||
|
matching = _fa_metadata(bs=2, num_tokens=2)
|
||||||
|
hybrid = object.__new__(DotsHybridAttnBackend)
|
||||||
|
hybrid.dsa_backend = SimpleNamespace(forward_metadata=matching)
|
||||||
|
hybrid.swa_backend = SimpleNamespace(forward_metadata=matching)
|
||||||
|
hybrid._dp_rebuilt_batch_id = None
|
||||||
|
hybrid.init_forward_metadata = MagicMock()
|
||||||
|
|
||||||
|
hybrid.maybe_rebuild_metadata_after_dp_padding(
|
||||||
|
_batch(bs=2, num_tokens=2, original_bs=1)
|
||||||
|
)
|
||||||
|
hybrid.init_forward_metadata.assert_called_once()
|
||||||
|
|
||||||
|
|
||||||
|
def test_normalize_cache_seqlens_preserves_planned_rows_and_pads_dummy_rows():
|
||||||
|
cache_seqlens = torch.tensor([17, 23], dtype=torch.int32)
|
||||||
|
seq_lens = torch.tensor([100, 200, 300, 400], dtype=torch.int64)
|
||||||
|
|
||||||
|
normalized = _normalize_cache_seqlens_rows(cache_seqlens, seq_lens, 4)
|
||||||
|
|
||||||
|
assert torch.equal(normalized, torch.tensor([17, 23, 300, 400], dtype=torch.int32))
|
||||||
|
|
||||||
|
|
||||||
|
def test_normalize_cache_seqlens_truncates_extra_rows():
|
||||||
|
cache_seqlens = torch.tensor([17, 23, 29], dtype=torch.int32)
|
||||||
|
seq_lens = torch.tensor([100, 200], dtype=torch.int64)
|
||||||
|
|
||||||
|
normalized = _normalize_cache_seqlens_rows(cache_seqlens, seq_lens, 2)
|
||||||
|
|
||||||
|
assert torch.equal(normalized, torch.tensor([17, 23], dtype=torch.int32))
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA")
|
||||||
|
def test_page64_gather_masks_out_of_range_page_table_entries():
|
||||||
|
kv_cache_dim = 128
|
||||||
|
k_cache = torch.arange(64 * kv_cache_dim, dtype=torch.float32, device="cuda").view(
|
||||||
|
64, 1, kv_cache_dim
|
||||||
|
)
|
||||||
|
# Row 0 has a sequence longer than its one-page table. Row 1 points past
|
||||||
|
# the physical KV pool. Both can occur transiently when DP padding changes
|
||||||
|
# the live batch after speculative metadata was planned.
|
||||||
|
block_table = torch.tensor([[0], [9]], dtype=torch.int32, device="cuda")
|
||||||
|
cache_seqlens = torch.tensor([130, 64], dtype=torch.int32, device="cuda")
|
||||||
|
|
||||||
|
gathered, valid = gather_page64_kv_latent(
|
||||||
|
k_cache=k_cache,
|
||||||
|
block_table=block_table,
|
||||||
|
cache_seqlens=cache_seqlens,
|
||||||
|
window_size=128,
|
||||||
|
s_q=1,
|
||||||
|
kv_cache_dim=kv_cache_dim,
|
||||||
|
)
|
||||||
|
torch.cuda.synchronize()
|
||||||
|
|
||||||
|
assert valid[0, :62].all()
|
||||||
|
assert not valid[0, 62:].any()
|
||||||
|
assert not valid[1].any()
|
||||||
|
torch.testing.assert_close(gathered[0, :62], k_cache[2:64, 0])
|
||||||
|
assert not gathered[0, 62:].any()
|
||||||
|
assert not gathered[1].any()
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
sys.exit(pytest.main([__file__]))
|
||||||
@@ -170,6 +170,38 @@ class TestMlpSyncPadUnpad(CustomTestCase):
|
|||||||
# row count must match the real request count.
|
# row count must match the real request count.
|
||||||
self.assertEqual((fb.seq_lens - 1).shape[0], fb.batch_size)
|
self.assertEqual((fb.seq_lens - 1).shape[0], fb.batch_size)
|
||||||
|
|
||||||
|
def test_draft_extend_dummy_request_pads_cpu_and_gpu_lens(self):
|
||||||
|
spec_info = MagicMock()
|
||||||
|
spec_info.num_tokens_per_req = 4
|
||||||
|
spec_info.is_draft_input.return_value = False
|
||||||
|
fb = ForwardBatch(
|
||||||
|
forward_mode=ForwardMode.DRAFT_EXTEND_V2,
|
||||||
|
batch_size=1,
|
||||||
|
input_ids=torch.empty(0, dtype=torch.int64),
|
||||||
|
req_pool_indices=torch.empty(0, dtype=torch.int64),
|
||||||
|
seq_lens=torch.empty(0, dtype=torch.int64),
|
||||||
|
seq_lens_sum=0,
|
||||||
|
out_cache_loc=torch.empty(0, dtype=torch.int64),
|
||||||
|
positions=torch.empty(0, dtype=torch.int64),
|
||||||
|
seq_lens_cpu=torch.empty(0, dtype=torch.int64),
|
||||||
|
extend_seq_lens=torch.empty(0, dtype=torch.int32),
|
||||||
|
extend_prefix_lens=torch.empty(0, dtype=torch.int64),
|
||||||
|
extend_seq_lens_cpu=[],
|
||||||
|
extend_prefix_lens_cpu=[],
|
||||||
|
extend_logprob_start_lens_cpu=[],
|
||||||
|
spec_info=spec_info,
|
||||||
|
)
|
||||||
|
|
||||||
|
fb._pad_inputs_to_size(_mock_model_runner(), num_tokens=4, bs=1)
|
||||||
|
|
||||||
|
torch.testing.assert_close(
|
||||||
|
fb.extend_seq_lens, torch.tensor([4], dtype=torch.int32)
|
||||||
|
)
|
||||||
|
torch.testing.assert_close(fb.extend_prefix_lens, torch.tensor([0]))
|
||||||
|
self.assertEqual(fb.extend_seq_lens_cpu, [4])
|
||||||
|
self.assertEqual(fb.extend_prefix_lens_cpu, [0])
|
||||||
|
self.assertEqual(fb.extend_logprob_start_lens_cpu, [0])
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
@@ -10,6 +10,7 @@ import unittest
|
|||||||
from types import SimpleNamespace
|
from types import SimpleNamespace
|
||||||
from unittest.mock import MagicMock, patch
|
from unittest.mock import MagicMock, patch
|
||||||
|
|
||||||
|
from sglang.srt.configs.model_config import AttentionArch
|
||||||
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
|
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
|
||||||
from sglang.srt.runtime_context import get_memory, get_parallel, get_server_args
|
from sglang.srt.runtime_context import get_memory, get_parallel, get_server_args
|
||||||
from sglang.test.ci.ci_register import register_cpu_ci
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
@@ -79,6 +80,8 @@ def _make_model_runner(
|
|||||||
disaggregation_decode_extra_slots=0,
|
disaggregation_decode_extra_slots=0,
|
||||||
kv_lora_rank=512,
|
kv_lora_rank=512,
|
||||||
qk_rope_head_dim=64,
|
qk_rope_head_dim=64,
|
||||||
|
swa_kv_lora_rank=128,
|
||||||
|
swa_qk_rope_head_dim=32,
|
||||||
):
|
):
|
||||||
"""Create a mock ModelRunner with the fields configurators need."""
|
"""Create a mock ModelRunner with the fields configurators need."""
|
||||||
mr = MagicMock()
|
mr = MagicMock()
|
||||||
@@ -99,6 +102,9 @@ def _make_model_runner(
|
|||||||
mc.v_head_dim = v_head_dim
|
mc.v_head_dim = v_head_dim
|
||||||
mc.kv_lora_rank = kv_lora_rank
|
mc.kv_lora_rank = kv_lora_rank
|
||||||
mc.qk_rope_head_dim = qk_rope_head_dim
|
mc.qk_rope_head_dim = qk_rope_head_dim
|
||||||
|
mc.swa_kv_lora_rank = swa_kv_lora_rank
|
||||||
|
mc.swa_qk_rope_head_dim = swa_qk_rope_head_dim
|
||||||
|
mc.attention_arch = AttentionArch.MLA if use_mla_backend else AttentionArch.MHA
|
||||||
mc.is_hybrid_swa = is_hybrid_swa
|
mc.is_hybrid_swa = is_hybrid_swa
|
||||||
mc.full_attention_layer_ids = (
|
mc.full_attention_layer_ids = (
|
||||||
full_attention_layer_ids
|
full_attention_layer_ids
|
||||||
@@ -158,7 +164,9 @@ def _make_model_runner(
|
|||||||
mr.ps = ParallelState.trivial()
|
mr.ps = ParallelState.trivial()
|
||||||
mr.pp_group = SimpleNamespace(rank_in_group=0)
|
mr.pp_group = SimpleNamespace(rank_in_group=0)
|
||||||
mr.spec_aux_config = SimpleNamespace(
|
mr.spec_aux_config = SimpleNamespace(
|
||||||
eagle_draft_num_layers=None, dflash_draft_num_layers=None
|
eagle_draft_num_layers=None,
|
||||||
|
eagle_draft_swa_num_layers=None,
|
||||||
|
dflash_draft_num_layers=None,
|
||||||
)
|
)
|
||||||
|
|
||||||
return mr
|
return mr
|
||||||
@@ -314,6 +322,58 @@ class TestHybridSWAConfigurator(CustomTestCase):
|
|||||||
self.assertLessEqual(used, available)
|
self.assertLessEqual(used, available)
|
||||||
self.assertGreater(used, available * 0.99)
|
self.assertGreater(used, available * 0.99)
|
||||||
|
|
||||||
|
@patch(
|
||||||
|
"sglang.srt.mem_cache.kv_cache_configurator.calculate_mla_kv_cache_dim",
|
||||||
|
return_value=576,
|
||||||
|
)
|
||||||
|
def test_mla_uses_full_and_swa_latent_geometry(
|
||||||
|
self,
|
||||||
|
mock_calculate_mla_kv_cache_dim,
|
||||||
|
):
|
||||||
|
"""Hybrid MLA pools must not be sized from MHA head geometry."""
|
||||||
|
available = 10_000_000
|
||||||
|
full_layers = 2
|
||||||
|
swa_layers = 3
|
||||||
|
swa_kv_lora_rank = 128
|
||||||
|
swa_qk_rope_head_dim = 32
|
||||||
|
mr = _make_model_runner(
|
||||||
|
self,
|
||||||
|
num_kv_heads=32,
|
||||||
|
head_dim=256,
|
||||||
|
v_head_dim=256,
|
||||||
|
use_mla_backend=True,
|
||||||
|
is_hybrid_swa=True,
|
||||||
|
full_attention_layer_ids=list(range(full_layers)),
|
||||||
|
swa_attention_layer_ids=list(range(full_layers, full_layers + swa_layers)),
|
||||||
|
swa_num_kv_heads=16,
|
||||||
|
swa_head_dim=128,
|
||||||
|
swa_v_head_dim=128,
|
||||||
|
swa_kv_lora_rank=swa_kv_lora_rank,
|
||||||
|
swa_qk_rope_head_dim=swa_qk_rope_head_dim,
|
||||||
|
swa_full_tokens_ratio=0.5,
|
||||||
|
)
|
||||||
|
|
||||||
|
with mock_cpu_env(kv_size=2):
|
||||||
|
from sglang.srt.model_executor.pool_configurator import (
|
||||||
|
create_memory_pool_configurator,
|
||||||
|
)
|
||||||
|
|
||||||
|
cfg = create_memory_pool_configurator(mr)
|
||||||
|
config = cfg.calculate_pool_sizes(available, page_size=1)
|
||||||
|
|
||||||
|
expected_full_per_token = 576 * 2
|
||||||
|
expected_swa_per_token = (swa_kv_lora_rank + swa_qk_rope_head_dim) * 2
|
||||||
|
self.assertEqual(cfg._full_per_token, expected_full_per_token)
|
||||||
|
self.assertEqual(cfg._swa_per_token, expected_swa_per_token)
|
||||||
|
mock_calculate_mla_kv_cache_dim.assert_called_once()
|
||||||
|
|
||||||
|
used = (
|
||||||
|
config.full_max_total_num_tokens * expected_full_per_token * full_layers
|
||||||
|
+ config.swa_max_total_num_tokens * expected_swa_per_token * swa_layers
|
||||||
|
)
|
||||||
|
self.assertLessEqual(used, available)
|
||||||
|
self.assertGreater(used, available * 0.99)
|
||||||
|
|
||||||
def test_ratio_respected(self):
|
def test_ratio_respected(self):
|
||||||
"""swa_tokens ~= full_tokens * ratio (within page alignment)"""
|
"""swa_tokens ~= full_tokens * ratio (within page alignment)"""
|
||||||
available = 10_000_000
|
available = 10_000_000
|
||||||
@@ -443,6 +503,40 @@ class TestHybridSWAConfigurator(CustomTestCase):
|
|||||||
self.assertEqual(config.swa_max_total_num_tokens, 91)
|
self.assertEqual(config.swa_max_total_num_tokens, 91)
|
||||||
self.assertLessEqual(_actual_memory_used(mr, config), available)
|
self.assertLessEqual(_actual_memory_used(mr, config), available)
|
||||||
|
|
||||||
|
def test_chunk_cache_cap_accounts_for_draft_swa_layers(self):
|
||||||
|
"""Draft SWA tensors consume the same fixed-capacity pool as target SWA."""
|
||||||
|
available = 1_000_000
|
||||||
|
mr = _make_model_runner(
|
||||||
|
self,
|
||||||
|
is_hybrid_swa=True,
|
||||||
|
full_attention_layer_ids=[0],
|
||||||
|
swa_attention_layer_ids=[1],
|
||||||
|
swa_num_kv_heads=4,
|
||||||
|
disable_radix_cache=True,
|
||||||
|
chunked_prefill_size=4,
|
||||||
|
sliding_window_size=8,
|
||||||
|
page_size=1,
|
||||||
|
max_running_requests=2,
|
||||||
|
)
|
||||||
|
mr.spec_algorithm.is_eagle.return_value = True
|
||||||
|
mr.spec_algorithm.is_none.return_value = False
|
||||||
|
mr.spec_aux_config.eagle_draft_num_layers = 1
|
||||||
|
mr.spec_aux_config.eagle_draft_swa_num_layers = 1
|
||||||
|
|
||||||
|
with mock_cpu_env():
|
||||||
|
from sglang.srt.model_executor.pool_configurator import (
|
||||||
|
create_memory_pool_configurator,
|
||||||
|
)
|
||||||
|
|
||||||
|
cfg = create_memory_pool_configurator(mr)
|
||||||
|
config = cfg.calculate_pool_sizes(available, page_size=1)
|
||||||
|
|
||||||
|
full_tokens = config.full_max_total_num_tokens
|
||||||
|
swa_tokens = config.swa_max_total_num_tokens
|
||||||
|
used = full_tokens * _full_per_token(mr) + swa_tokens * _swa_per_token(mr) * 2
|
||||||
|
self.assertLessEqual(used, available)
|
||||||
|
self.assertGreater(used, available * 0.99)
|
||||||
|
|
||||||
def test_chunk_cache_cap_drops_prefill_for_disagg_decode(self):
|
def test_chunk_cache_cap_drops_prefill_for_disagg_decode(self):
|
||||||
available = 1_000_000
|
available = 1_000_000
|
||||||
mr = _make_model_runner(
|
mr = _make_model_runner(
|
||||||
@@ -658,6 +752,45 @@ class TestEagleConfigurator(CustomTestCase):
|
|||||||
available,
|
available,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def test_hybrid_swa_draft_uses_swa_geometry_and_capacity(self):
|
||||||
|
"""SWA draft layers use SWA KV geometry and capacity."""
|
||||||
|
available = 10_000_000
|
||||||
|
ratio = 0.25
|
||||||
|
mr = _make_model_runner(
|
||||||
|
self,
|
||||||
|
num_kv_heads=8,
|
||||||
|
head_dim=64,
|
||||||
|
v_head_dim=64,
|
||||||
|
num_layers=4,
|
||||||
|
is_hybrid_swa=True,
|
||||||
|
full_attention_layer_ids=[0, 1],
|
||||||
|
swa_attention_layer_ids=[2, 3],
|
||||||
|
swa_num_kv_heads=2,
|
||||||
|
swa_head_dim=32,
|
||||||
|
swa_v_head_dim=32,
|
||||||
|
swa_full_tokens_ratio=ratio,
|
||||||
|
)
|
||||||
|
mr.spec_algorithm.is_eagle.return_value = True
|
||||||
|
mr.spec_algorithm.is_none.return_value = False
|
||||||
|
mr.spec_aux_config.eagle_draft_num_layers = 1
|
||||||
|
mr.spec_aux_config.eagle_draft_swa_num_layers = 1
|
||||||
|
|
||||||
|
with mock_cpu_env():
|
||||||
|
from sglang.srt.model_executor.pool_configurator import (
|
||||||
|
create_memory_pool_configurator,
|
||||||
|
)
|
||||||
|
|
||||||
|
cfg = create_memory_pool_configurator(mr)
|
||||||
|
config = cfg.calculate_pool_sizes(available, page_size=1)
|
||||||
|
|
||||||
|
full_tokens = config.full_max_total_num_tokens
|
||||||
|
swa_tokens = config.swa_max_total_num_tokens
|
||||||
|
full_pt = _full_per_token(mr)
|
||||||
|
swa_pt = _swa_per_token(mr)
|
||||||
|
used = full_tokens * full_pt * 2 + swa_tokens * swa_pt * 3
|
||||||
|
self.assertLessEqual(used, available)
|
||||||
|
self.assertGreater(used, available * 0.99)
|
||||||
|
|
||||||
|
|
||||||
class TestDSAIndexerAllocationPolicy(CustomTestCase):
|
class TestDSAIndexerAllocationPolicy(CustomTestCase):
|
||||||
@patch(
|
@patch(
|
||||||
|
|||||||
@@ -0,0 +1,285 @@
|
|||||||
|
import asyncio
|
||||||
|
import concurrent.futures
|
||||||
|
import re
|
||||||
|
import types
|
||||||
|
import unittest
|
||||||
|
from unittest.mock import patch
|
||||||
|
|
||||||
|
from sglang.test.test_utils import maybe_stub_sgl_kernel
|
||||||
|
|
||||||
|
maybe_stub_sgl_kernel()
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from sglang.srt.managers.schedule_batch import Modality
|
||||||
|
from sglang.srt.multimodal.processors.dots_note_omni import DotsNoteOmniProcessor
|
||||||
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
|
from sglang.test.test_utils import CustomTestCase
|
||||||
|
|
||||||
|
register_cpu_ci(est_time=5, suite="base-a-test-cpu")
|
||||||
|
|
||||||
|
_IM_TOKEN_ID = 100
|
||||||
|
_AUDIO_TOKEN_ID = 200
|
||||||
|
# The dots.note chat template renders a video content part as this single token.
|
||||||
|
_VIDEO_PLACEHOLDER = "<|video_pad|>"
|
||||||
|
|
||||||
|
|
||||||
|
class _FakeMultimodalTokens:
|
||||||
|
_pattern = re.compile(r"(<image>|<audio>)")
|
||||||
|
|
||||||
|
def get_combined_regex(self):
|
||||||
|
return self._pattern
|
||||||
|
|
||||||
|
def get_modality_of_token(self, token):
|
||||||
|
return {
|
||||||
|
"<image>": Modality.IMAGE,
|
||||||
|
"<audio>": Modality.AUDIO,
|
||||||
|
}.get(token)
|
||||||
|
|
||||||
|
|
||||||
|
class _FakeTokenizer:
|
||||||
|
"""Maps each media marker to one pad id and other text to per-char ids."""
|
||||||
|
|
||||||
|
def encode(self, text, add_special_tokens=False):
|
||||||
|
ids = []
|
||||||
|
for part in _FakeMultimodalTokens._pattern.split(text):
|
||||||
|
if part == "<image>":
|
||||||
|
ids.append(_IM_TOKEN_ID)
|
||||||
|
elif part == "<audio>":
|
||||||
|
ids.append(_AUDIO_TOKEN_ID)
|
||||||
|
else:
|
||||||
|
ids.extend(ord(char) for char in part)
|
||||||
|
return ids
|
||||||
|
|
||||||
|
|
||||||
|
def _fake_preprocess_dots_video(raw_video, question, **kwargs):
|
||||||
|
"""Each video flattens into one frame, one audio segment and the question."""
|
||||||
|
return [
|
||||||
|
{"type": "image_url", "image_url": {"url": f"{raw_video}-frame"}},
|
||||||
|
{"type": "audio_url", "audio_url": {"url": f"{raw_video}-audio"}},
|
||||||
|
{"type": "text", "text": question},
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
class TestDotsNoteOmniVideoMixing(CustomTestCase):
|
||||||
|
def setUp(self):
|
||||||
|
self.processor = DotsNoteOmniProcessor.__new__(DotsNoteOmniProcessor)
|
||||||
|
self.processor.image_start_token = ""
|
||||||
|
self.processor.image_token = "<image>"
|
||||||
|
self.processor.image_end_token = ""
|
||||||
|
self.processor.audio_start_token = ""
|
||||||
|
self.processor.audio_token = "<audio>"
|
||||||
|
self.processor.audio_end_token = ""
|
||||||
|
self.processor.mm_tokens = _FakeMultimodalTokens()
|
||||||
|
self.processor.video_placeholder_regex = re.compile(
|
||||||
|
re.escape(_VIDEO_PLACEHOLDER)
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_multiple_videos_and_native_media_keep_prompt_order(self):
|
||||||
|
prompt = f"{_VIDEO_PLACEHOLDER}<image>between{_VIDEO_PLACEHOLDER}question"
|
||||||
|
all_video_media = {}
|
||||||
|
contents = [
|
||||||
|
[
|
||||||
|
{"type": "image_url", "image_url": {"url": "video-0-frame"}},
|
||||||
|
{"type": "text", "text": "question"},
|
||||||
|
],
|
||||||
|
[
|
||||||
|
{"type": "audio_url", "audio_url": {"url": "video-1-audio"}},
|
||||||
|
{"type": "text", "text": "question"},
|
||||||
|
],
|
||||||
|
]
|
||||||
|
|
||||||
|
for index, content in enumerate(contents):
|
||||||
|
prompt, video_media = self.processor._render_video_content(
|
||||||
|
prompt, "question", index, content
|
||||||
|
)
|
||||||
|
all_video_media.update(video_media)
|
||||||
|
|
||||||
|
prompt, images, audios = self.processor._merge_video_media(
|
||||||
|
prompt,
|
||||||
|
image_data=["native-image"],
|
||||||
|
audio_data=None,
|
||||||
|
video_media=all_video_media,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(prompt, "<image><image>between<audio>question")
|
||||||
|
self.assertEqual(images, ["video-0-frame", "native-image"])
|
||||||
|
self.assertEqual(audios, ["video-1-audio"])
|
||||||
|
|
||||||
|
def test_template_without_video_placeholders_inserts_each_video_once(self):
|
||||||
|
prompt = "<|user|>question"
|
||||||
|
all_video_media = {}
|
||||||
|
|
||||||
|
for index in range(2):
|
||||||
|
prompt, video_media = self.processor._render_video_content(
|
||||||
|
prompt,
|
||||||
|
"question",
|
||||||
|
index,
|
||||||
|
[
|
||||||
|
{
|
||||||
|
"type": "image_url",
|
||||||
|
"image_url": {"url": f"video-{index}-frame"},
|
||||||
|
},
|
||||||
|
{"type": "text", "text": "question"},
|
||||||
|
],
|
||||||
|
)
|
||||||
|
all_video_media.update(video_media)
|
||||||
|
|
||||||
|
prompt, images, audios = self.processor._merge_video_media(
|
||||||
|
prompt, None, None, all_video_media
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(prompt, "<|user|><image><image>question")
|
||||||
|
self.assertEqual(images, ["video-0-frame", "video-1-frame"])
|
||||||
|
self.assertEqual(audios, [])
|
||||||
|
|
||||||
|
|
||||||
|
class TestDotsNoteOmniProcessMmDataAsync(CustomTestCase):
|
||||||
|
"""Drive the request entry point that used to reject mixed video inputs."""
|
||||||
|
|
||||||
|
def setUp(self):
|
||||||
|
self.processor = DotsNoteOmniProcessor.__new__(DotsNoteOmniProcessor)
|
||||||
|
self.processor.image_start_token = ""
|
||||||
|
self.processor.image_token = "<image>"
|
||||||
|
self.processor.image_end_token = ""
|
||||||
|
self.processor.audio_start_token = ""
|
||||||
|
self.processor.audio_token = "<audio>"
|
||||||
|
self.processor.audio_end_token = ""
|
||||||
|
self.processor.mm_tokens = _FakeMultimodalTokens()
|
||||||
|
self.processor.video_placeholder_regex = re.compile(
|
||||||
|
re.escape(_VIDEO_PLACEHOLDER)
|
||||||
|
)
|
||||||
|
self.processor.mm_token_ids = {
|
||||||
|
"im_start_id": 98,
|
||||||
|
"im_token_id": _IM_TOKEN_ID,
|
||||||
|
"im_end_id": 99,
|
||||||
|
"audio_start_id": 198,
|
||||||
|
"audio_token_id": _AUDIO_TOKEN_ID,
|
||||||
|
"audio_end_id": 199,
|
||||||
|
}
|
||||||
|
self.processor._tokenizer = _FakeTokenizer()
|
||||||
|
self.processor.audio_processor_config = types.SimpleNamespace(
|
||||||
|
sampling_rate=16000
|
||||||
|
)
|
||||||
|
self.processor.image_preprocessor = types.SimpleNamespace(
|
||||||
|
process_images=lambda images: (
|
||||||
|
[torch.tensor([float(index)]) for index in range(len(images))],
|
||||||
|
[torch.tensor([1, 1, 4]) for _ in images],
|
||||||
|
["<image>" for _ in images],
|
||||||
|
)
|
||||||
|
)
|
||||||
|
self.processor.io_executor = concurrent.futures.ThreadPoolExecutor(
|
||||||
|
max_workers=2
|
||||||
|
)
|
||||||
|
self.addCleanup(self.processor.io_executor.shutdown)
|
||||||
|
|
||||||
|
self.loaded = {}
|
||||||
|
|
||||||
|
async def fake_load_mm_data(prompt, image_data=None, audio_data=None, **kwargs):
|
||||||
|
self.loaded["prompt"] = prompt
|
||||||
|
self.loaded["image_data"] = list(image_data or [])
|
||||||
|
self.loaded["audio_data"] = list(audio_data or [])
|
||||||
|
return types.SimpleNamespace(
|
||||||
|
input_text=prompt,
|
||||||
|
images=list(image_data or []),
|
||||||
|
audios=[torch.zeros(4) for _ in audio_data or []],
|
||||||
|
)
|
||||||
|
|
||||||
|
self.processor.load_mm_data = fake_load_mm_data
|
||||||
|
|
||||||
|
def _run(self, request_obj, **kwargs):
|
||||||
|
with (
|
||||||
|
patch(
|
||||||
|
"sglang.srt.multimodal.processors.dots_note_omni.preprocess_dots_video",
|
||||||
|
_fake_preprocess_dots_video,
|
||||||
|
),
|
||||||
|
patch(
|
||||||
|
"sglang.srt.multimodal.processors.dots_note_omni.get_audio_token_string",
|
||||||
|
lambda *args, **kwargs: "<audio>",
|
||||||
|
),
|
||||||
|
):
|
||||||
|
return asyncio.run(
|
||||||
|
self.processor.process_mm_data_async(
|
||||||
|
request_obj.text,
|
||||||
|
request_obj,
|
||||||
|
max_req_input_len=4096,
|
||||||
|
**kwargs,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _request(text, video_data):
|
||||||
|
return types.SimpleNamespace(
|
||||||
|
text=text,
|
||||||
|
video_data=video_data,
|
||||||
|
video_config={
|
||||||
|
"_question": "question",
|
||||||
|
"seq": 131072,
|
||||||
|
"audio_cap": 1.0,
|
||||||
|
"audio_sr": 16000,
|
||||||
|
"k_mode": "eval_ek",
|
||||||
|
},
|
||||||
|
sampling_params={"max_new_tokens": 16},
|
||||||
|
rid="test-rid",
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_video_mixed_with_image_and_audio(self):
|
||||||
|
request_obj = self._request(
|
||||||
|
f"<image>{_VIDEO_PLACEHOLDER}middle<audio>question", ["video-0"]
|
||||||
|
)
|
||||||
|
|
||||||
|
output = self._run(
|
||||||
|
request_obj, image_data=["native-image"], audio_data=["native-audio"]
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(
|
||||||
|
self.loaded["prompt"], "<image><image><audio>middle<audio>question"
|
||||||
|
)
|
||||||
|
self.assertEqual(self.loaded["image_data"], ["native-image", "video-0-frame"])
|
||||||
|
self.assertEqual(self.loaded["audio_data"], ["video-0-audio", "native-audio"])
|
||||||
|
self.assertEqual(
|
||||||
|
[item.modality for item in output.mm_items],
|
||||||
|
[
|
||||||
|
Modality.IMAGE,
|
||||||
|
Modality.IMAGE,
|
||||||
|
Modality.AUDIO,
|
||||||
|
Modality.AUDIO,
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_multiple_videos_mixed_with_image(self):
|
||||||
|
request_obj = self._request(
|
||||||
|
f"{_VIDEO_PLACEHOLDER}<image>middle{_VIDEO_PLACEHOLDER}question",
|
||||||
|
["video-0", "video-1"],
|
||||||
|
)
|
||||||
|
|
||||||
|
output = self._run(request_obj, image_data=["native-image"])
|
||||||
|
|
||||||
|
self.assertEqual(
|
||||||
|
self.loaded["prompt"],
|
||||||
|
"<image><audio><image>middle<image><audio>question",
|
||||||
|
)
|
||||||
|
self.assertEqual(
|
||||||
|
self.loaded["image_data"],
|
||||||
|
["video-0-frame", "native-image", "video-1-frame"],
|
||||||
|
)
|
||||||
|
self.assertEqual(self.loaded["audio_data"], ["video-0-audio", "video-1-audio"])
|
||||||
|
self.assertEqual(len(output.mm_items), 5)
|
||||||
|
|
||||||
|
def test_extra_native_image_without_placeholder_is_rejected(self):
|
||||||
|
request_obj = self._request(f"{_VIDEO_PLACEHOLDER}question", ["video-0"])
|
||||||
|
|
||||||
|
with self.assertRaisesRegex(ValueError, "Image placeholder count"):
|
||||||
|
self._run(request_obj, image_data=["native-image"])
|
||||||
|
|
||||||
|
def test_unconsumed_video_placeholder_is_rejected(self):
|
||||||
|
request_obj = self._request(
|
||||||
|
f"{_VIDEO_PLACEHOLDER}{_VIDEO_PLACEHOLDER}question", ["video-0"]
|
||||||
|
)
|
||||||
|
|
||||||
|
with self.assertRaisesRegex(ValueError, "Video placeholder count"):
|
||||||
|
self._run(request_obj)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
@@ -6,6 +6,7 @@ tests check that the pre-allocated `parent_list` / `top_scores_index` match the
|
|||||||
slow path (`organize_draft_results`) for num_steps in {1, 2, 3, 4}.
|
slow path (`organize_draft_results`) for num_steps in {1, 2, 3, 4}.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
import contextlib
|
||||||
import unittest
|
import unittest
|
||||||
from types import SimpleNamespace
|
from types import SimpleNamespace
|
||||||
from unittest.mock import MagicMock, patch
|
from unittest.mock import MagicMock, patch
|
||||||
@@ -134,6 +135,31 @@ class TestEagleWorkerV2Topk1FastPath(CustomTestCase):
|
|||||||
with self.assertRaises(AssertionError):
|
with self.assertRaises(AssertionError):
|
||||||
worker._rebuild_topk1_chain_buffers()
|
worker._rebuild_topk1_chain_buffers()
|
||||||
|
|
||||||
|
def test_idle_draft_runs_each_eager_forward_without_tree_layout(self):
|
||||||
|
worker = object.__new__(EagleDraftWorker)
|
||||||
|
worker.speculative_num_steps = 3
|
||||||
|
worker.draft_attn_backend = SimpleNamespace(attn_backends=[object(), object()])
|
||||||
|
worker.draft_runner = SimpleNamespace(
|
||||||
|
canary_manager=None,
|
||||||
|
forward=MagicMock(),
|
||||||
|
)
|
||||||
|
spec_info = SimpleNamespace(hidden_states=torch.empty((0, 8), device=DEVICE))
|
||||||
|
forward_batch = SimpleNamespace(
|
||||||
|
forward_mode=ForwardMode.IDLE,
|
||||||
|
input_ids=torch.empty((0,), dtype=torch.long, device=DEVICE),
|
||||||
|
out_cache_loc=torch.empty((0,), dtype=torch.long, device=DEVICE),
|
||||||
|
spec_info=spec_info,
|
||||||
|
)
|
||||||
|
|
||||||
|
with patch(
|
||||||
|
"sglang.srt.speculative.eagle_worker_v2.forward_context",
|
||||||
|
side_effect=lambda *_args, **_kwargs: contextlib.nullcontext(),
|
||||||
|
):
|
||||||
|
result = worker.draft_forward(forward_batch)
|
||||||
|
|
||||||
|
self.assertEqual(result, (None, None, None, None))
|
||||||
|
self.assertEqual(worker.draft_runner.forward.call_count, 2)
|
||||||
|
|
||||||
|
|
||||||
class TestEagleWorkerV2BackendFallback(CustomTestCase):
|
class TestEagleWorkerV2BackendFallback(CustomTestCase):
|
||||||
def setUp(self):
|
def setUp(self):
|
||||||
|
|||||||
Reference in New Issue
Block a user