model: support baidu unlimited-ocr (#29186)
Co-authored-by: Mick <mickjagger19@icloud.com>
This commit is contained in:
@@ -166,6 +166,10 @@ def get_model_config(
|
|||||||
0 if disable_shared_experts_fusion or topk_ids_dir is None else 1
|
0 if disable_shared_experts_fusion or topk_ids_dir is None else 1
|
||||||
)
|
)
|
||||||
intermediate_size = config.intermediate_size
|
intermediate_size = config.intermediate_size
|
||||||
|
elif architecture == "UnlimitedOCRForCausalLM":
|
||||||
|
E = config.n_routed_experts // ep_size
|
||||||
|
topk = config.num_experts_per_tok
|
||||||
|
intermediate_size = config.moe_intermediate_size
|
||||||
else:
|
else:
|
||||||
# Default: Mixtral
|
# Default: Mixtral
|
||||||
E = config.num_local_experts // ep_size
|
E = config.num_local_experts // ep_size
|
||||||
|
|||||||
@@ -0,0 +1,6 @@
|
|||||||
|
<svg width="940" height="525" viewBox="0 0 940 525" fill="none" xmlns="http://www.w3.org/2000/svg">
|
||||||
|
<rect width="940" height="525" fill="none"/>
|
||||||
|
<text x="470" y="286" text-anchor="middle" font-family="Arial, Helvetica, sans-serif" font-size="128" font-weight="700" letter-spacing="0">
|
||||||
|
<tspan fill="#2B5BFF">Bai</tspan><tspan fill="#D5001C">du</tspan>
|
||||||
|
</text>
|
||||||
|
</svg>
|
||||||
|
After Width: | Height: | Size: 376 B |
@@ -0,0 +1,155 @@
|
|||||||
|
---
|
||||||
|
title: Unlimited-OCR
|
||||||
|
description: "Deploy Baidu Unlimited-OCR with SGLang for long document OCR using prefill-aware sliding-window attention."
|
||||||
|
tag: NEW
|
||||||
|
---
|
||||||
|
|
||||||
|
## Deployment
|
||||||
|
|
||||||
|
<a id="install" />
|
||||||
|
|
||||||
|
<Accordion title="Install SGLang">
|
||||||
|
|
||||||
|
Unlimited-OCR support is in [SGLang PR #29186](https://github.com/sgl-project/sglang/pull/29186). Until that PR is included in a tagged SGLang release, install from a build that contains the PR.
|
||||||
|
|
||||||
|
<Tabs>
|
||||||
|
|
||||||
|
<Tab title="Python (pip / uv)">
|
||||||
|
|
||||||
|
```bash Command
|
||||||
|
pip install -U uv
|
||||||
|
uv venv --python 3.12 && source .venv/bin/activate
|
||||||
|
|
||||||
|
git clone https://github.com/sgl-project/sglang.git
|
||||||
|
cd sglang
|
||||||
|
git fetch origin pull/29186/head && git checkout FETCH_HEAD
|
||||||
|
uv pip install -e python
|
||||||
|
```
|
||||||
|
|
||||||
|
Then run the **Python** output of the command panel below in that environment.
|
||||||
|
|
||||||
|
</Tab>
|
||||||
|
|
||||||
|
<Tab title="Docker">
|
||||||
|
|
||||||
|
```bash Command
|
||||||
|
docker pull lmsysorg/sglang:dev
|
||||||
|
```
|
||||||
|
|
||||||
|
For how to launch the image, see [Install → Method 3: Using Docker](../../../docs/get-started/install#method-3-using-docker). Substitute the inner `sglang serve ...` with what the command generator below produces.
|
||||||
|
|
||||||
|
</Tab>
|
||||||
|
|
||||||
|
</Tabs>
|
||||||
|
|
||||||
|
</Accordion>
|
||||||
|
|
||||||
|
Pick your hardware to generate the launch command. The recipe uses FlashAttention-3 with `--page-size 1`, which is required by the current prefill-aware sliding-window attention path. It also disables radix cache by default, which is the better fit for batch OCR workloads where each request usually contains a different image.
|
||||||
|
|
||||||
|
import { Deployment } from "/src/snippets/_deployment.jsx";
|
||||||
|
import { config } from "/src/snippets/configs/baidu/unlimited-ocr.jsx";
|
||||||
|
|
||||||
|
<Deployment config={config} />
|
||||||
|
|
||||||
|
## Playground
|
||||||
|
|
||||||
|
Use the Playground to adjust tensor parallelism on top of the selected deployment cell.
|
||||||
|
|
||||||
|
import { Playground } from "/src/snippets/_playground.jsx";
|
||||||
|
|
||||||
|
<Playground config={config} />
|
||||||
|
|
||||||
|
## 1. Model Introduction
|
||||||
|
|
||||||
|
[Unlimited-OCR](https://huggingface.co/baidu/Unlimited-OCR) is Baidu's multimodal OCR model for document parsing. It uses a sliding-window language backbone, but SGLang serves it with a prefill-aware sliding-window path so image and prompt tokens remain visible during long decode.
|
||||||
|
|
||||||
|
The SGLang integration loads the standalone Unlimited-OCR architecture with SAM and CLIP vision encoders plus a DeepSeek-style language backbone. It supports OpenAI-compatible image requests and model-specific image processing options through `images_config`.
|
||||||
|
|
||||||
|
**Resources:** [Hugging Face](https://huggingface.co/baidu/Unlimited-OCR) · [SGLang PR #29186](https://github.com/sgl-project/sglang/pull/29186)
|
||||||
|
|
||||||
|
## 2. Configuration Tips
|
||||||
|
|
||||||
|
- **Attention backend**: use `--attention-backend fa3 --page-size 1`. The prefill-aware SWA page table is built with token-level locations, so page size 1 is required.
|
||||||
|
- **Radix cache**: keep `--disable-radix-cache` for batch OCR over different documents. If your workload repeatedly asks about the same image and prompt, remove this flag to allow prefix reuse through `PureSWARadixCache`.
|
||||||
|
- **Long OCR generations**: keep the default prefill-aware SWA path enabled. It retains prompt and image KV while still applying a sliding window to generated text.
|
||||||
|
- **Custom logit processor**: keep `--enable-custom-logit-processor` in the launch command.
|
||||||
|
- **Image modes**: pass `images_config.image_mode` per request. Supported modes are `tiny`, `small`, `base`, `large`, and `gundam`. Multiple images are supported only for `tiny`, `small`, and `base`.
|
||||||
|
- **Default image mode**: when `images_config.image_mode` is omitted, SGLang uses `gundam`.
|
||||||
|
|
||||||
|
## 3. Advanced Usage
|
||||||
|
|
||||||
|
### 3.1 OCR request
|
||||||
|
|
||||||
|
<Accordion title="OCR 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="baidu/Unlimited-OCR",
|
||||||
|
messages=[
|
||||||
|
{
|
||||||
|
"role": "user",
|
||||||
|
"content": [
|
||||||
|
{"type": "text", "text": "document parsing."},
|
||||||
|
{
|
||||||
|
"type": "image_url",
|
||||||
|
"image_url": {
|
||||||
|
"url": "https://example.com/your_document.png"
|
||||||
|
},
|
||||||
|
},
|
||||||
|
],
|
||||||
|
}
|
||||||
|
],
|
||||||
|
max_tokens=2048,
|
||||||
|
temperature=0,
|
||||||
|
extra_body={"images_config": {"image_mode": "gundam"}},
|
||||||
|
)
|
||||||
|
|
||||||
|
print(response.choices[0].message.content)
|
||||||
|
```
|
||||||
|
|
||||||
|
</Accordion>
|
||||||
|
|
||||||
|
### 3.2 Choosing an image mode
|
||||||
|
|
||||||
|
Use lower modes to reduce prefill cost for simple images, and use `gundam` for high-detail document parsing.
|
||||||
|
|
||||||
|
<table style={{width: "100%", borderCollapse: "collapse", tableLayout: "fixed"}}>
|
||||||
|
<thead>
|
||||||
|
<tr style={{borderBottom: "2px solid #d55816"}}>
|
||||||
|
<th style={{textAlign: "left", padding: "10px 12px", fontWeight: 700}}>Mode</th>
|
||||||
|
<th style={{textAlign: "left", padding: "10px 12px", fontWeight: 700}}>Use</th>
|
||||||
|
<th style={{textAlign: "left", padding: "10px 12px", fontWeight: 700}}>Multiple images</th>
|
||||||
|
</tr>
|
||||||
|
</thead>
|
||||||
|
<tbody>
|
||||||
|
<tr>
|
||||||
|
<td style={{padding: "9px 12px"}}><code>tiny</code></td>
|
||||||
|
<td style={{padding: "9px 12px"}}>Lowest prefill cost.</td>
|
||||||
|
<td style={{padding: "9px 12px"}}>Yes</td>
|
||||||
|
</tr>
|
||||||
|
<tr>
|
||||||
|
<td style={{padding: "9px 12px"}}><code>small</code></td>
|
||||||
|
<td style={{padding: "9px 12px"}}>Lightweight OCR requests.</td>
|
||||||
|
<td style={{padding: "9px 12px"}}>Yes</td>
|
||||||
|
</tr>
|
||||||
|
<tr>
|
||||||
|
<td style={{padding: "9px 12px"}}><code>base</code></td>
|
||||||
|
<td style={{padding: "9px 12px"}}>Balanced quality and cost.</td>
|
||||||
|
<td style={{padding: "9px 12px"}}>Yes</td>
|
||||||
|
</tr>
|
||||||
|
<tr>
|
||||||
|
<td style={{padding: "9px 12px"}}><code>large</code></td>
|
||||||
|
<td style={{padding: "9px 12px"}}>Higher resolution single-image OCR.</td>
|
||||||
|
<td style={{padding: "9px 12px"}}>No</td>
|
||||||
|
</tr>
|
||||||
|
<tr>
|
||||||
|
<td style={{padding: "9px 12px"}}><code>gundam</code></td>
|
||||||
|
<td style={{padding: "9px 12px"}}>Default high-detail document parsing mode.</td>
|
||||||
|
<td style={{padding: "9px 12px"}}>No</td>
|
||||||
|
</tr>
|
||||||
|
</tbody>
|
||||||
|
</table>
|
||||||
@@ -67,6 +67,12 @@ metatags:
|
|||||||
href="/cookbook/autoregressive/NVIDIA/Nemotron3-Ultra"
|
href="/cookbook/autoregressive/NVIDIA/Nemotron3-Ultra"
|
||||||
img="/cards/logos/nvidia.png"
|
img="/cards/logos/nvidia.png"
|
||||||
/>
|
/>
|
||||||
|
<Card
|
||||||
|
title="Baidu"
|
||||||
|
mode="card"
|
||||||
|
href="/cookbook/autoregressive/Baidu/Unlimited-OCR"
|
||||||
|
img="/cards/logos/baidu.svg"
|
||||||
|
/>
|
||||||
<Card
|
<Card
|
||||||
title="Ernie"
|
title="Ernie"
|
||||||
mode="card"
|
mode="card"
|
||||||
|
|||||||
@@ -201,6 +201,10 @@
|
|||||||
"source": "/advanced_features/vlm_query.html",
|
"source": "/advanced_features/vlm_query.html",
|
||||||
"destination": "/docs/advanced_features/vlm_query"
|
"destination": "/docs/advanced_features/vlm_query"
|
||||||
},
|
},
|
||||||
|
{
|
||||||
|
"source": "/basic_usage/unlimited_ocr.html",
|
||||||
|
"destination": "/cookbook/autoregressive/Baidu/Unlimited-OCR"
|
||||||
|
},
|
||||||
{
|
{
|
||||||
"source": "/basic_usage/deepseek_ocr.html",
|
"source": "/basic_usage/deepseek_ocr.html",
|
||||||
"destination": "/cookbook/autoregressive/DeepSeek/DeepSeek-OCR"
|
"destination": "/cookbook/autoregressive/DeepSeek/DeepSeek-OCR"
|
||||||
@@ -1058,6 +1062,12 @@
|
|||||||
"cookbook/autoregressive/NVIDIA/Nemotron3-Super"
|
"cookbook/autoregressive/NVIDIA/Nemotron3-Super"
|
||||||
]
|
]
|
||||||
},
|
},
|
||||||
|
{
|
||||||
|
"group": "Baidu",
|
||||||
|
"pages": [
|
||||||
|
"cookbook/autoregressive/Baidu/Unlimited-OCR"
|
||||||
|
]
|
||||||
|
},
|
||||||
{
|
{
|
||||||
"group": "Ernie",
|
"group": "Ernie",
|
||||||
"pages": [
|
"pages": [
|
||||||
|
|||||||
@@ -0,0 +1,188 @@
|
|||||||
|
// Unlimited-OCR cookbook config. Consumed by _deployment.jsx + _playground.jsx.
|
||||||
|
|
||||||
|
export const config = {
|
||||||
|
modelName: "Unlimited-OCR",
|
||||||
|
|
||||||
|
supportedHardware: ["h100", "h200", "b200", "b300", "gb200", "gb300"],
|
||||||
|
|
||||||
|
variants: [{ id: "default", label: "Default" }],
|
||||||
|
quantizations: [{ id: "default", label: "Default" }],
|
||||||
|
strategies: [{ id: "balanced", label: "Balanced" }],
|
||||||
|
nodesOptions: [{ id: "single", label: "Single Node" }],
|
||||||
|
|
||||||
|
modelNames: {
|
||||||
|
"default|default": "baidu/Unlimited-OCR",
|
||||||
|
},
|
||||||
|
|
||||||
|
placeholders: {
|
||||||
|
HOST_IP: { target: "command", label: "Bind host", default: "0.0.0.0" },
|
||||||
|
PORT: { target: "command", label: "Bind port", default: "30000" },
|
||||||
|
HF_TOKEN: {
|
||||||
|
target: "command",
|
||||||
|
label: "HF token (Docker)",
|
||||||
|
default: "<your-hf-token>",
|
||||||
|
},
|
||||||
|
CURL_HOST: { target: "curl", label: "Server host", default: "localhost" },
|
||||||
|
CURL_PORT: { target: "curl", label: "Server port", default: "30000" },
|
||||||
|
},
|
||||||
|
|
||||||
|
curl: `curl http://{{CURL_HOST}}:{{CURL_PORT}}/v1/chat/completions \\
|
||||||
|
-H 'Content-Type: application/json' \\
|
||||||
|
-d '{
|
||||||
|
"model": "{{MODEL_NAME}}",
|
||||||
|
"messages": [{
|
||||||
|
"role": "user",
|
||||||
|
"content": [
|
||||||
|
{"type": "text", "text": "document parsing."},
|
||||||
|
{"type": "image_url", "image_url": {"url": "https://example.com/your_document.png"}}
|
||||||
|
]
|
||||||
|
}],
|
||||||
|
"images_config": {"image_mode": "gundam"},
|
||||||
|
"temperature": 0,
|
||||||
|
"max_tokens": 2048
|
||||||
|
}'`,
|
||||||
|
|
||||||
|
dockerImages: {
|
||||||
|
h100: "lmsysorg/sglang:dev",
|
||||||
|
h200: "lmsysorg/sglang:dev",
|
||||||
|
b200: "lmsysorg/sglang:dev",
|
||||||
|
b300: "lmsysorg/sglang:dev",
|
||||||
|
gb200: "lmsysorg/sglang:dev",
|
||||||
|
gb300: "lmsysorg/sglang:dev",
|
||||||
|
},
|
||||||
|
|
||||||
|
github: {
|
||||||
|
cookbookModel: "baidu/Unlimited-OCR",
|
||||||
|
},
|
||||||
|
|
||||||
|
playgroundFeatures: {
|
||||||
|
attention: {
|
||||||
|
knobs: [
|
||||||
|
{ id: "tp", label: "TP", values: [null, 1, 2, 4, 8] },
|
||||||
|
],
|
||||||
|
},
|
||||||
|
},
|
||||||
|
|
||||||
|
cells: [
|
||||||
|
{
|
||||||
|
match: {
|
||||||
|
hw: "h100",
|
||||||
|
variant: "default",
|
||||||
|
quant: "default",
|
||||||
|
strategy: "balanced",
|
||||||
|
nodes: "single",
|
||||||
|
},
|
||||||
|
env: [],
|
||||||
|
flags: [
|
||||||
|
"--model-path {{MODEL_NAME}}",
|
||||||
|
"--attention-backend fa3",
|
||||||
|
"--page-size 1",
|
||||||
|
"--context-length 32768",
|
||||||
|
"--enable-custom-logit-processor",
|
||||||
|
"--disable-radix-cache",
|
||||||
|
"--host {{HOST_IP}}",
|
||||||
|
"--port {{PORT}}",
|
||||||
|
],
|
||||||
|
},
|
||||||
|
{
|
||||||
|
match: {
|
||||||
|
hw: "h200",
|
||||||
|
variant: "default",
|
||||||
|
quant: "default",
|
||||||
|
strategy: "balanced",
|
||||||
|
nodes: "single",
|
||||||
|
},
|
||||||
|
env: [],
|
||||||
|
flags: [
|
||||||
|
"--model-path {{MODEL_NAME}}",
|
||||||
|
"--attention-backend fa3",
|
||||||
|
"--page-size 1",
|
||||||
|
"--context-length 32768",
|
||||||
|
"--enable-custom-logit-processor",
|
||||||
|
"--disable-radix-cache",
|
||||||
|
"--host {{HOST_IP}}",
|
||||||
|
"--port {{PORT}}",
|
||||||
|
],
|
||||||
|
},
|
||||||
|
{
|
||||||
|
match: {
|
||||||
|
hw: "b200",
|
||||||
|
variant: "default",
|
||||||
|
quant: "default",
|
||||||
|
strategy: "balanced",
|
||||||
|
nodes: "single",
|
||||||
|
},
|
||||||
|
env: [],
|
||||||
|
flags: [
|
||||||
|
"--model-path {{MODEL_NAME}}",
|
||||||
|
"--attention-backend fa3",
|
||||||
|
"--page-size 1",
|
||||||
|
"--context-length 32768",
|
||||||
|
"--enable-custom-logit-processor",
|
||||||
|
"--disable-radix-cache",
|
||||||
|
"--host {{HOST_IP}}",
|
||||||
|
"--port {{PORT}}",
|
||||||
|
],
|
||||||
|
},
|
||||||
|
{
|
||||||
|
match: {
|
||||||
|
hw: "b300",
|
||||||
|
variant: "default",
|
||||||
|
quant: "default",
|
||||||
|
strategy: "balanced",
|
||||||
|
nodes: "single",
|
||||||
|
},
|
||||||
|
env: [],
|
||||||
|
flags: [
|
||||||
|
"--model-path {{MODEL_NAME}}",
|
||||||
|
"--attention-backend fa3",
|
||||||
|
"--page-size 1",
|
||||||
|
"--context-length 32768",
|
||||||
|
"--enable-custom-logit-processor",
|
||||||
|
"--disable-radix-cache",
|
||||||
|
"--host {{HOST_IP}}",
|
||||||
|
"--port {{PORT}}",
|
||||||
|
],
|
||||||
|
},
|
||||||
|
{
|
||||||
|
match: {
|
||||||
|
hw: "gb200",
|
||||||
|
variant: "default",
|
||||||
|
quant: "default",
|
||||||
|
strategy: "balanced",
|
||||||
|
nodes: "single",
|
||||||
|
},
|
||||||
|
env: [],
|
||||||
|
flags: [
|
||||||
|
"--model-path {{MODEL_NAME}}",
|
||||||
|
"--attention-backend fa3",
|
||||||
|
"--page-size 1",
|
||||||
|
"--context-length 32768",
|
||||||
|
"--enable-custom-logit-processor",
|
||||||
|
"--disable-radix-cache",
|
||||||
|
"--host {{HOST_IP}}",
|
||||||
|
"--port {{PORT}}",
|
||||||
|
],
|
||||||
|
},
|
||||||
|
{
|
||||||
|
match: {
|
||||||
|
hw: "gb300",
|
||||||
|
variant: "default",
|
||||||
|
quant: "default",
|
||||||
|
strategy: "balanced",
|
||||||
|
nodes: "single",
|
||||||
|
},
|
||||||
|
env: [],
|
||||||
|
flags: [
|
||||||
|
"--model-path {{MODEL_NAME}}",
|
||||||
|
"--attention-backend fa3",
|
||||||
|
"--page-size 1",
|
||||||
|
"--context-length 32768",
|
||||||
|
"--enable-custom-logit-processor",
|
||||||
|
"--disable-radix-cache",
|
||||||
|
"--host {{HOST_IP}}",
|
||||||
|
"--port {{PORT}}",
|
||||||
|
],
|
||||||
|
},
|
||||||
|
],
|
||||||
|
};
|
||||||
@@ -39,6 +39,7 @@ from sglang.srt.configs.step3_vl import (
|
|||||||
)
|
)
|
||||||
from sglang.srt.configs.step3p5 import Step3p5Config
|
from sglang.srt.configs.step3p5 import Step3p5Config
|
||||||
from sglang.srt.configs.step3p7 import Step3p7Config
|
from sglang.srt.configs.step3p7 import Step3p7Config
|
||||||
|
from sglang.srt.configs.unlimited_ocr import UnlimitedVLConfig
|
||||||
from sglang.srt.configs.zaya import ZayaConfig
|
from sglang.srt.configs.zaya import ZayaConfig
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
@@ -81,5 +82,6 @@ __all__ = [
|
|||||||
"Step3p5Config",
|
"Step3p5Config",
|
||||||
"Step3p7Config",
|
"Step3p7Config",
|
||||||
"Qwen3ASRConfig",
|
"Qwen3ASRConfig",
|
||||||
|
"UnlimitedVLConfig",
|
||||||
"ZayaConfig",
|
"ZayaConfig",
|
||||||
]
|
]
|
||||||
|
|||||||
@@ -1707,6 +1707,7 @@ multimodal_model_archs = [
|
|||||||
"NVILAForConditionalGeneration",
|
"NVILAForConditionalGeneration",
|
||||||
"NVILALiteForConditionalGeneration",
|
"NVILALiteForConditionalGeneration",
|
||||||
"DeepseekOCRForCausalLM",
|
"DeepseekOCRForCausalLM",
|
||||||
|
"UnlimitedOCRForCausalLM",
|
||||||
"JetVLMForConditionalGeneration",
|
"JetVLMForConditionalGeneration",
|
||||||
"PaddleOCRVLForConditionalGeneration",
|
"PaddleOCRVLForConditionalGeneration",
|
||||||
"MiDashengLMModel",
|
"MiDashengLMModel",
|
||||||
@@ -1863,6 +1864,7 @@ def is_hybrid_swa_model(
|
|||||||
"Gemma4ForConditionalGeneration",
|
"Gemma4ForConditionalGeneration",
|
||||||
"Gemma4UnifiedForConditionalGeneration",
|
"Gemma4UnifiedForConditionalGeneration",
|
||||||
"LagunaForCausalLM",
|
"LagunaForCausalLM",
|
||||||
|
"UnlimitedOCRForCausalLM",
|
||||||
}
|
}
|
||||||
if any(arch in hybrid_swa_archs for arch in model_architectures):
|
if any(arch in hybrid_swa_archs for arch in model_architectures):
|
||||||
# Only treat Laguna as hybrid SWA when it actually has a sliding window.
|
# Only treat Laguna as hybrid SWA when it actually has a sliding window.
|
||||||
@@ -1949,6 +1951,9 @@ def get_hybrid_layer_ids(
|
|||||||
full_attention_layer_ids = [
|
full_attention_layer_ids = [
|
||||||
i for i, x in enumerate(layer_types) if x == "full_attention"
|
i for i, x in enumerate(layer_types) if x == "full_attention"
|
||||||
]
|
]
|
||||||
|
elif "UnlimitedOCRForCausalLM" in model_architectures:
|
||||||
|
swa_attention_layer_ids = list(range(num_hidden_layers))
|
||||||
|
full_attention_layer_ids = []
|
||||||
elif getattr(hf_text_config, "hybrid_layer_pattern", None) is not None:
|
elif getattr(hf_text_config, "hybrid_layer_pattern", None) is not None:
|
||||||
# Generic fallback for custom hybrid SWA models that opt in via
|
# Generic fallback for custom hybrid SWA models that opt in via
|
||||||
# hf_text_config.is_hybrid_swa and expose a hybrid_layer_pattern
|
# hf_text_config.is_hybrid_swa and expose a hybrid_layer_pattern
|
||||||
|
|||||||
@@ -0,0 +1,629 @@
|
|||||||
|
"""Standalone UNLIMITED-OCR configuration and HF processor."""
|
||||||
|
|
||||||
|
import math
|
||||||
|
from typing import Any, Dict, List, Tuple
|
||||||
|
|
||||||
|
import torch
|
||||||
|
from PIL import Image, ImageOps
|
||||||
|
from transformers import (
|
||||||
|
AutoConfig,
|
||||||
|
AutoProcessor,
|
||||||
|
PretrainedConfig,
|
||||||
|
PreTrainedTokenizerFast,
|
||||||
|
ProcessorMixin,
|
||||||
|
)
|
||||||
|
|
||||||
|
from sglang.srt.configs.deepseek_ocr import (
|
||||||
|
ImageTransform,
|
||||||
|
MlpProjectorConfig,
|
||||||
|
VisionEncoderConfig,
|
||||||
|
VLChatProcessorOutput,
|
||||||
|
find_closest_aspect_ratio,
|
||||||
|
)
|
||||||
|
from sglang.srt.multimodal.customized_mm_processor_utils import (
|
||||||
|
register_customized_processor,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def dynamic_preprocess(
|
||||||
|
image, min_num=2, max_num=32, image_size=640, use_thumbnail=False
|
||||||
|
):
|
||||||
|
"""Split an image into tiles based on the best-matching aspect ratio."""
|
||||||
|
orig_width, orig_height = image.size
|
||||||
|
aspect_ratio = orig_width / orig_height
|
||||||
|
|
||||||
|
target_ratios = set(
|
||||||
|
(i, j)
|
||||||
|
for n in range(min_num, max_num + 1)
|
||||||
|
for i in range(1, n + 1)
|
||||||
|
for j in range(1, n + 1)
|
||||||
|
if i * j <= max_num and i * j >= min_num
|
||||||
|
)
|
||||||
|
target_ratios = sorted(target_ratios, key=lambda x: x[0] * x[1])
|
||||||
|
|
||||||
|
target_aspect_ratio = find_closest_aspect_ratio(
|
||||||
|
aspect_ratio, target_ratios, orig_width, orig_height, image_size
|
||||||
|
)
|
||||||
|
|
||||||
|
target_width = image_size * target_aspect_ratio[0]
|
||||||
|
target_height = image_size * target_aspect_ratio[1]
|
||||||
|
blocks = target_aspect_ratio[0] * target_aspect_ratio[1]
|
||||||
|
|
||||||
|
resized_img = image.resize((target_width, target_height))
|
||||||
|
processed_images = []
|
||||||
|
for i in range(blocks):
|
||||||
|
box = (
|
||||||
|
(i % (target_width // image_size)) * image_size,
|
||||||
|
(i // (target_width // image_size)) * image_size,
|
||||||
|
((i % (target_width // image_size)) + 1) * image_size,
|
||||||
|
((i // (target_width // image_size)) + 1) * image_size,
|
||||||
|
)
|
||||||
|
split_img = resized_img.crop(box)
|
||||||
|
processed_images.append(split_img)
|
||||||
|
assert len(processed_images) == blocks
|
||||||
|
if use_thumbnail and len(processed_images) != 1:
|
||||||
|
thumbnail_img = image.resize((image_size, image_size))
|
||||||
|
processed_images.append(thumbnail_img)
|
||||||
|
return processed_images, target_aspect_ratio
|
||||||
|
|
||||||
|
|
||||||
|
class UnlimitedOCRHFProcessor(ProcessorMixin):
|
||||||
|
"""HuggingFace-style processor for UNLIMITED-OCR (OCR mode)."""
|
||||||
|
|
||||||
|
tokenizer_class = "PreTrainedTokenizerFast"
|
||||||
|
attributes = ["tokenizer"]
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
tokenizer: PreTrainedTokenizerFast,
|
||||||
|
candidate_resolutions: Tuple[Tuple[int, int]],
|
||||||
|
patch_size: int,
|
||||||
|
downsample_ratio: int,
|
||||||
|
image_mean: Tuple[float, float, float] = (0.5, 0.5, 0.5),
|
||||||
|
image_std: Tuple[float, float, float] = (0.5, 0.5, 0.5),
|
||||||
|
normalize: bool = True,
|
||||||
|
image_token: str = "<image>",
|
||||||
|
pad_token: str = "<|▁pad▁|>",
|
||||||
|
add_special_token: bool = False,
|
||||||
|
sft_format: str = "unlimitedocr",
|
||||||
|
mask_prompt: bool = True,
|
||||||
|
ignore_id: int = -100,
|
||||||
|
base_size: int = 1024,
|
||||||
|
image_size: int = 640,
|
||||||
|
crop_mode: bool = True,
|
||||||
|
**kwargs,
|
||||||
|
):
|
||||||
|
"""Initialize tokenizer, image transform, and special tokens."""
|
||||||
|
self.candidate_resolutions = candidate_resolutions
|
||||||
|
self.base_size = base_size
|
||||||
|
self.image_size = image_size
|
||||||
|
self.crop_mode = crop_mode
|
||||||
|
self.patch_size = patch_size
|
||||||
|
self.image_mean = image_mean
|
||||||
|
self.image_std = image_std
|
||||||
|
self.normalize = normalize
|
||||||
|
self.downsample_ratio = downsample_ratio
|
||||||
|
self.image_transform = ImageTransform(
|
||||||
|
mean=image_mean, std=image_std, normalize=normalize
|
||||||
|
)
|
||||||
|
if type(tokenizer) is not PreTrainedTokenizerFast:
|
||||||
|
tokenizer = PreTrainedTokenizerFast.from_pretrained(tokenizer.name_or_path)
|
||||||
|
self.tokenizer = tokenizer
|
||||||
|
self.tokenizer.padding_side = "left"
|
||||||
|
|
||||||
|
if tokenizer.pad_token is None:
|
||||||
|
self.tokenizer.add_special_tokens({"pad_token": pad_token})
|
||||||
|
|
||||||
|
image_token_id = self.tokenizer.vocab.get(image_token)
|
||||||
|
if image_token_id is None:
|
||||||
|
special_tokens = [image_token]
|
||||||
|
special_tokens_dict = {"additional_special_tokens": special_tokens}
|
||||||
|
self.tokenizer.add_special_tokens(special_tokens_dict)
|
||||||
|
self.image_token_id = self.tokenizer.vocab.get(image_token)
|
||||||
|
|
||||||
|
special_tokens = ["<|ref|>", "<|/ref|>", "<|det|>", "<|/det|>", "<|grounding|>"]
|
||||||
|
special_tokens_dict = {"additional_special_tokens": special_tokens}
|
||||||
|
self.tokenizer.add_special_tokens(special_tokens_dict)
|
||||||
|
|
||||||
|
special_tokens = ["<|User|>", "<|Assistant|>"]
|
||||||
|
special_tokens_dict = {"additional_special_tokens": special_tokens}
|
||||||
|
self.tokenizer.add_special_tokens(special_tokens_dict)
|
||||||
|
|
||||||
|
self.image_token = image_token
|
||||||
|
self.pad_token = pad_token
|
||||||
|
self.add_special_token = add_special_token
|
||||||
|
self.sft_format = sft_format
|
||||||
|
self.mask_prompt = mask_prompt
|
||||||
|
self.ignore_id = ignore_id
|
||||||
|
|
||||||
|
super().__init__(tokenizer, **kwargs)
|
||||||
|
|
||||||
|
def format_messages_v2(
|
||||||
|
self,
|
||||||
|
messages: str,
|
||||||
|
pil_images,
|
||||||
|
max_req_input_len=-1,
|
||||||
|
base_size: int = None,
|
||||||
|
image_size: int = None,
|
||||||
|
crop_mode: bool = None,
|
||||||
|
):
|
||||||
|
"""Tokenize messages with embedded images and return processed tensors."""
|
||||||
|
base_size = base_size or self.base_size
|
||||||
|
image_size = image_size or self.image_size
|
||||||
|
crop_mode = crop_mode if crop_mode is not None else self.crop_mode
|
||||||
|
|
||||||
|
tokenized_data = []
|
||||||
|
masked_tokenized_data = []
|
||||||
|
images_list = []
|
||||||
|
images_seq_mask = []
|
||||||
|
images_spatial_crop = []
|
||||||
|
|
||||||
|
image_index = 0
|
||||||
|
image_token_cnt = messages.count(self.image_token)
|
||||||
|
(
|
||||||
|
input_ids,
|
||||||
|
images,
|
||||||
|
images_crop,
|
||||||
|
seq_mask,
|
||||||
|
spatial_crop,
|
||||||
|
num_image_tokens,
|
||||||
|
image_shapes,
|
||||||
|
) = self.tokenize_with_images(
|
||||||
|
messages,
|
||||||
|
pil_images[image_index : image_index + image_token_cnt],
|
||||||
|
bos=True,
|
||||||
|
eos=True,
|
||||||
|
cropping=crop_mode,
|
||||||
|
base_size=base_size,
|
||||||
|
image_size=image_size,
|
||||||
|
)
|
||||||
|
|
||||||
|
image_index = image_token_cnt
|
||||||
|
images_list += images
|
||||||
|
images_seq_mask += seq_mask
|
||||||
|
images_spatial_crop = spatial_crop
|
||||||
|
|
||||||
|
return (
|
||||||
|
input_ids,
|
||||||
|
masked_tokenized_data,
|
||||||
|
images_list,
|
||||||
|
images_seq_mask,
|
||||||
|
images_spatial_crop,
|
||||||
|
images_crop,
|
||||||
|
)
|
||||||
|
|
||||||
|
@property
|
||||||
|
def bos_id(self):
|
||||||
|
"""Return the beginning-of-sequence token ID."""
|
||||||
|
return self.tokenizer.bos_token_id
|
||||||
|
|
||||||
|
@property
|
||||||
|
def eos_id(self):
|
||||||
|
"""Return the end-of-sequence token ID."""
|
||||||
|
return self.tokenizer.eos_token_id
|
||||||
|
|
||||||
|
@property
|
||||||
|
def pad_id(self):
|
||||||
|
"""Return the padding token ID."""
|
||||||
|
return self.tokenizer.pad_token_id
|
||||||
|
|
||||||
|
def encode(self, text: str, bos: bool = True, eos: bool = False):
|
||||||
|
"""Encode text into token IDs with optional BOS/EOS."""
|
||||||
|
t = self.tokenizer.encode(text, add_special_tokens=False)
|
||||||
|
if bos:
|
||||||
|
t = [self.bos_id] + t
|
||||||
|
if eos:
|
||||||
|
t = t + [self.eos_id]
|
||||||
|
return t
|
||||||
|
|
||||||
|
def decode(self, t: List[int], **kwargs) -> str:
|
||||||
|
"""Decode token IDs back into a string."""
|
||||||
|
return self.tokenizer.decode(t, **kwargs)
|
||||||
|
|
||||||
|
def process_one(
|
||||||
|
self,
|
||||||
|
prompt: str = None,
|
||||||
|
conversations: List[Dict[str, str]] = None,
|
||||||
|
images: List[Image.Image] = None,
|
||||||
|
apply_sft_format: bool = False,
|
||||||
|
inference_mode: bool = True,
|
||||||
|
system_prompt: str = "",
|
||||||
|
max_req_input_len: int = -1,
|
||||||
|
base_size: int = None,
|
||||||
|
image_size: int = None,
|
||||||
|
crop_mode: bool = None,
|
||||||
|
**kwargs,
|
||||||
|
):
|
||||||
|
"""Process a single prompt with images into model-ready tensors."""
|
||||||
|
base_size = base_size or self.base_size
|
||||||
|
image_size = image_size or self.image_size
|
||||||
|
crop_mode = crop_mode if crop_mode is not None else self.crop_mode
|
||||||
|
|
||||||
|
prompt = conversations or prompt
|
||||||
|
(
|
||||||
|
input_ids,
|
||||||
|
masked_tokenized_str,
|
||||||
|
images_list,
|
||||||
|
images_seq_mask,
|
||||||
|
images_spatial_crop,
|
||||||
|
images_crop,
|
||||||
|
) = self.format_messages_v2(
|
||||||
|
prompt,
|
||||||
|
images,
|
||||||
|
max_req_input_len,
|
||||||
|
base_size=base_size,
|
||||||
|
image_size=image_size,
|
||||||
|
crop_mode=crop_mode,
|
||||||
|
)
|
||||||
|
|
||||||
|
target_ids = torch.LongTensor(masked_tokenized_str)
|
||||||
|
|
||||||
|
has_images = len(images_list) > 0
|
||||||
|
has_local_crops = []
|
||||||
|
if len(images_spatial_crop) > 0:
|
||||||
|
has_local_crops = [
|
||||||
|
(crop[0] > 1 or crop[1] > 1).item() for crop in images_spatial_crop
|
||||||
|
]
|
||||||
|
|
||||||
|
if len(images_list) == 0:
|
||||||
|
images = torch.zeros((1, 3, image_size, image_size))
|
||||||
|
else:
|
||||||
|
images = torch.stack(images_list, dim=0)
|
||||||
|
|
||||||
|
images_spatial_crop = torch.stack([images_spatial_crop], dim=0)
|
||||||
|
|
||||||
|
prepare = VLChatProcessorOutput(
|
||||||
|
input_ids=input_ids,
|
||||||
|
target_ids=target_ids,
|
||||||
|
images_crop=images_crop,
|
||||||
|
pixel_values=images,
|
||||||
|
images_seq_mask=images_seq_mask,
|
||||||
|
images_spatial_crop=images_spatial_crop,
|
||||||
|
)
|
||||||
|
prepare.has_images = has_images
|
||||||
|
prepare.has_local_crops = has_local_crops
|
||||||
|
|
||||||
|
return prepare
|
||||||
|
|
||||||
|
def __call__(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
prompt: str = None,
|
||||||
|
conversations: List[Dict[str, str]] = None,
|
||||||
|
images: List[Image.Image] = None,
|
||||||
|
apply_sft_format: bool = False,
|
||||||
|
inference_mode: bool = True,
|
||||||
|
system_prompt: str = "",
|
||||||
|
max_req_input_len: int = -1,
|
||||||
|
text: list[str] = None,
|
||||||
|
base_size: int = None,
|
||||||
|
image_size: int = None,
|
||||||
|
crop_mode: bool = None,
|
||||||
|
**kwargs,
|
||||||
|
):
|
||||||
|
"""Call the processor to tokenize text and images for inference."""
|
||||||
|
assert text is None or isinstance(text, list)
|
||||||
|
if text is not None:
|
||||||
|
text = text[0]
|
||||||
|
|
||||||
|
prepare = self.process_one(
|
||||||
|
prompt=prompt or text,
|
||||||
|
conversations=conversations,
|
||||||
|
images=images,
|
||||||
|
apply_sft_format=apply_sft_format,
|
||||||
|
inference_mode=inference_mode,
|
||||||
|
system_prompt=system_prompt,
|
||||||
|
max_req_input_len=max_req_input_len,
|
||||||
|
base_size=base_size if base_size is not None else self.base_size,
|
||||||
|
image_size=image_size if image_size is not None else self.image_size,
|
||||||
|
crop_mode=crop_mode if crop_mode is not None else self.crop_mode,
|
||||||
|
)
|
||||||
|
return prepare
|
||||||
|
|
||||||
|
def find_all_indices(self, messages, target_value):
|
||||||
|
"""Return all indices where target_value appears in messages."""
|
||||||
|
indices = []
|
||||||
|
for index, item in enumerate(messages):
|
||||||
|
if item == target_value:
|
||||||
|
indices.append(index)
|
||||||
|
return indices
|
||||||
|
|
||||||
|
def tokenize_with_images(
|
||||||
|
self,
|
||||||
|
conversation: str,
|
||||||
|
images: List[Image.Image],
|
||||||
|
bos: bool = True,
|
||||||
|
eos: bool = True,
|
||||||
|
cropping: bool = True,
|
||||||
|
base_size: int = None,
|
||||||
|
image_size: int = None,
|
||||||
|
):
|
||||||
|
"""Tokenize text with <image> tags (OCR mode)."""
|
||||||
|
base_size = base_size or self.base_size
|
||||||
|
image_size = image_size or self.image_size
|
||||||
|
|
||||||
|
assert conversation.count(self.image_token) == len(images)
|
||||||
|
text_splits: list[str] = conversation.split(self.image_token)
|
||||||
|
images_list, images_crop_list, images_seq_mask, images_spatial_crop = (
|
||||||
|
[],
|
||||||
|
[],
|
||||||
|
[],
|
||||||
|
[],
|
||||||
|
)
|
||||||
|
image_shapes = []
|
||||||
|
num_image_tokens = []
|
||||||
|
tokenized_str = []
|
||||||
|
|
||||||
|
for text_sep, image in zip(text_splits, images):
|
||||||
|
tokenized_sep = self.encode(text_sep, bos=False, eos=False)
|
||||||
|
tokenized_str += tokenized_sep
|
||||||
|
images_seq_mask += [False] * len(tokenized_sep)
|
||||||
|
image_shapes.append(image.size)
|
||||||
|
|
||||||
|
if image.size[0] <= 640 and image.size[1] <= 640:
|
||||||
|
crop_ratio = [1, 1]
|
||||||
|
else:
|
||||||
|
if cropping:
|
||||||
|
images_crop_raw, crop_ratio = dynamic_preprocess(
|
||||||
|
image, image_size=image_size
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
crop_ratio = [1, 1]
|
||||||
|
|
||||||
|
if image_size <= 640 and not cropping:
|
||||||
|
image = image.resize((image_size, image_size))
|
||||||
|
if cropping:
|
||||||
|
pad_size = base_size
|
||||||
|
else:
|
||||||
|
pad_size = image_size
|
||||||
|
|
||||||
|
global_view = ImageOps.pad(
|
||||||
|
image,
|
||||||
|
(pad_size, pad_size),
|
||||||
|
color=tuple(int(x * 255) for x in self.image_transform.mean),
|
||||||
|
)
|
||||||
|
images_list.append(self.image_transform(global_view))
|
||||||
|
|
||||||
|
num_width_tiles, num_height_tiles = crop_ratio
|
||||||
|
images_spatial_crop.append([num_width_tiles, num_height_tiles])
|
||||||
|
|
||||||
|
if num_width_tiles > 1 or num_height_tiles > 1:
|
||||||
|
for i in range(len(images_crop_raw)):
|
||||||
|
images_crop_list.append(self.image_transform(images_crop_raw[i]))
|
||||||
|
|
||||||
|
num_queries = math.ceil(
|
||||||
|
(image_size // self.patch_size) / self.downsample_ratio
|
||||||
|
)
|
||||||
|
num_queries_base = math.ceil(
|
||||||
|
(base_size // self.patch_size) / self.downsample_ratio
|
||||||
|
)
|
||||||
|
if cropping:
|
||||||
|
tokenized_image = (
|
||||||
|
[self.image_token_id] * num_queries_base + [self.image_token_id]
|
||||||
|
) * num_queries_base
|
||||||
|
tokenized_image += [self.image_token_id]
|
||||||
|
if num_width_tiles > 1 or num_height_tiles > 1:
|
||||||
|
tokenized_image += (
|
||||||
|
[self.image_token_id] * (num_queries * num_width_tiles)
|
||||||
|
+ [self.image_token_id]
|
||||||
|
) * (num_queries * num_height_tiles)
|
||||||
|
else:
|
||||||
|
tokenized_image = (
|
||||||
|
[self.image_token_id] * num_queries + [self.image_token_id]
|
||||||
|
) * num_queries
|
||||||
|
tokenized_image += [self.image_token_id]
|
||||||
|
|
||||||
|
tokenized_str += tokenized_image
|
||||||
|
images_seq_mask += [True] * len(tokenized_image)
|
||||||
|
num_image_tokens.append(len(tokenized_image))
|
||||||
|
|
||||||
|
tokenized_sep = self.encode(text_splits[-1], bos=False, eos=False)
|
||||||
|
tokenized_str += tokenized_sep
|
||||||
|
images_seq_mask += [False] * len(tokenized_sep)
|
||||||
|
|
||||||
|
if bos:
|
||||||
|
tokenized_str = [self.bos_id] + tokenized_str
|
||||||
|
images_seq_mask = [False] + images_seq_mask
|
||||||
|
if eos:
|
||||||
|
tokenized_str = tokenized_str + [self.eos_id]
|
||||||
|
images_seq_mask = images_seq_mask + [False]
|
||||||
|
|
||||||
|
assert len(tokenized_str) == len(images_seq_mask)
|
||||||
|
|
||||||
|
masked_tokenized_str = []
|
||||||
|
for token_index in tokenized_str:
|
||||||
|
if token_index != self.image_token_id:
|
||||||
|
masked_tokenized_str.append(token_index)
|
||||||
|
else:
|
||||||
|
masked_tokenized_str.append(self.ignore_id)
|
||||||
|
|
||||||
|
assert len(tokenized_str) == len(images_seq_mask) == len(masked_tokenized_str)
|
||||||
|
|
||||||
|
input_ids = torch.LongTensor(tokenized_str)
|
||||||
|
target_ids = torch.LongTensor(masked_tokenized_str)
|
||||||
|
images_seq_mask = torch.tensor(images_seq_mask, dtype=torch.bool)
|
||||||
|
|
||||||
|
target_ids[(input_ids < 0) | (input_ids == self.image_token_id)] = (
|
||||||
|
self.ignore_id
|
||||||
|
)
|
||||||
|
input_ids[input_ids < 0] = self.pad_id
|
||||||
|
|
||||||
|
inference_mode = True
|
||||||
|
if inference_mode:
|
||||||
|
assert input_ids[-1] == self.eos_id
|
||||||
|
input_ids = input_ids[:-1]
|
||||||
|
target_ids = target_ids[:-1]
|
||||||
|
images_seq_mask = images_seq_mask[:-1]
|
||||||
|
|
||||||
|
if len(images_list) == 0:
|
||||||
|
pixel_values = torch.zeros((1, 3, base_size, base_size))
|
||||||
|
images_spatial_crop = torch.zeros((1, 1), dtype=torch.long)
|
||||||
|
images_crop = torch.zeros((1, 3, image_size, image_size)).unsqueeze(0)
|
||||||
|
else:
|
||||||
|
pixel_values = torch.stack(images_list, dim=0)
|
||||||
|
images_spatial_crop = torch.tensor(images_spatial_crop, dtype=torch.long)
|
||||||
|
if images_crop_list:
|
||||||
|
images_crop = torch.stack(images_crop_list, dim=0).unsqueeze(0)
|
||||||
|
else:
|
||||||
|
images_crop = torch.zeros(
|
||||||
|
(len(images_list), 3, image_size, image_size)
|
||||||
|
).unsqueeze(1)
|
||||||
|
|
||||||
|
input_ids = input_ids.unsqueeze(0)
|
||||||
|
return (
|
||||||
|
input_ids,
|
||||||
|
pixel_values,
|
||||||
|
images_crop,
|
||||||
|
images_seq_mask,
|
||||||
|
images_spatial_crop,
|
||||||
|
num_image_tokens,
|
||||||
|
image_shapes,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class UnlimitedLanguageConfig(PretrainedConfig):
|
||||||
|
"""Configuration for the UNLIMITED language model backbone."""
|
||||||
|
|
||||||
|
model_type = "unlimited_language"
|
||||||
|
keys_to_ignore_at_inference = ["past_key_values"]
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
vocab_size=102400,
|
||||||
|
hidden_size=4096,
|
||||||
|
intermediate_size=11008,
|
||||||
|
moe_intermediate_size=1407,
|
||||||
|
num_hidden_layers=30,
|
||||||
|
num_attention_heads=32,
|
||||||
|
num_key_value_heads=32,
|
||||||
|
n_shared_experts=None,
|
||||||
|
n_routed_experts=None,
|
||||||
|
ep_size=1,
|
||||||
|
routed_scaling_factor=1.0,
|
||||||
|
kv_lora_rank=512,
|
||||||
|
q_lora_rank=1536,
|
||||||
|
qk_rope_head_dim=64,
|
||||||
|
v_head_dim=128,
|
||||||
|
qk_nope_head_dim=128,
|
||||||
|
topk_method="gready",
|
||||||
|
n_group=None,
|
||||||
|
topk_group=None,
|
||||||
|
num_experts_per_tok=None,
|
||||||
|
moe_layer_freq=1,
|
||||||
|
first_k_dense_replace=0,
|
||||||
|
norm_topk_prob=False,
|
||||||
|
scoring_func="softmax",
|
||||||
|
aux_loss_alpha=0.001,
|
||||||
|
seq_aux=True,
|
||||||
|
hidden_act="silu",
|
||||||
|
max_position_embeddings=2048,
|
||||||
|
initializer_range=0.02,
|
||||||
|
rms_norm_eps=1e-6,
|
||||||
|
use_cache=True,
|
||||||
|
pad_token_id=None,
|
||||||
|
bos_token_id=100000,
|
||||||
|
eos_token_id=100001,
|
||||||
|
pretraining_tp=1,
|
||||||
|
tie_word_embeddings=False,
|
||||||
|
rope_theta=10000.0,
|
||||||
|
rope_scaling=None,
|
||||||
|
attention_bias=False,
|
||||||
|
attention_dropout=0.0,
|
||||||
|
use_mla=True,
|
||||||
|
**kwargs,
|
||||||
|
):
|
||||||
|
"""Initialize language model configuration parameters."""
|
||||||
|
self.vocab_size = vocab_size
|
||||||
|
self.max_position_embeddings = max_position_embeddings
|
||||||
|
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.n_shared_experts = n_shared_experts
|
||||||
|
self.n_routed_experts = n_routed_experts
|
||||||
|
self.ep_size = ep_size
|
||||||
|
self.routed_scaling_factor = routed_scaling_factor
|
||||||
|
self.kv_lora_rank = kv_lora_rank
|
||||||
|
self.q_lora_rank = q_lora_rank
|
||||||
|
self.qk_rope_head_dim = qk_rope_head_dim
|
||||||
|
self.v_head_dim = v_head_dim
|
||||||
|
self.qk_nope_head_dim = qk_nope_head_dim
|
||||||
|
self.topk_method = topk_method
|
||||||
|
self.n_group = n_group
|
||||||
|
self.topk_group = topk_group
|
||||||
|
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.norm_topk_prob = norm_topk_prob
|
||||||
|
self.scoring_func = scoring_func
|
||||||
|
self.aux_loss_alpha = aux_loss_alpha
|
||||||
|
self.seq_aux = seq_aux
|
||||||
|
if num_key_value_heads is None:
|
||||||
|
num_key_value_heads = num_attention_heads
|
||||||
|
self.num_key_value_heads = num_key_value_heads
|
||||||
|
self.hidden_act = hidden_act
|
||||||
|
self.initializer_range = initializer_range
|
||||||
|
self.rms_norm_eps = float(rms_norm_eps)
|
||||||
|
self.pretraining_tp = pretraining_tp
|
||||||
|
self.use_cache = use_cache
|
||||||
|
self.rope_theta = rope_theta
|
||||||
|
self.rope_scaling = rope_scaling
|
||||||
|
self.attention_bias = attention_bias
|
||||||
|
self.attention_dropout = attention_dropout
|
||||||
|
self.use_mla = use_mla
|
||||||
|
|
||||||
|
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,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@register_customized_processor(processor_class=UnlimitedOCRHFProcessor)
|
||||||
|
class UnlimitedVLConfig(PretrainedConfig):
|
||||||
|
"""Top-level vision-language config for UNLIMITED-OCR models."""
|
||||||
|
|
||||||
|
model_type = "unlimited-ocr"
|
||||||
|
vision_config: VisionEncoderConfig = None
|
||||||
|
projector_config: MlpProjectorConfig = None
|
||||||
|
|
||||||
|
tile_tag: str = "2D"
|
||||||
|
global_view_pos: str = "head"
|
||||||
|
candidate_resolutions: tuple[tuple[int, int]] = ((384, 384),)
|
||||||
|
customized_processor_type: type[Any] = UnlimitedOCRHFProcessor
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
tile_tag: str = "tile_tag",
|
||||||
|
global_view_pos: str = "head",
|
||||||
|
candidate_resolutions: tuple[tuple[int, int]] = ((384, 384),),
|
||||||
|
**kwargs,
|
||||||
|
):
|
||||||
|
"""Initialize UNLIMITED VL config with vision, projector, and language sub-configs."""
|
||||||
|
super().__init__(**kwargs)
|
||||||
|
|
||||||
|
vision_config = kwargs.get("vision_config", {})
|
||||||
|
self.vision_config = VisionEncoderConfig(**vision_config)
|
||||||
|
|
||||||
|
projector_config = kwargs.get("projector_config", {})
|
||||||
|
self.projector_config = MlpProjectorConfig(**projector_config)
|
||||||
|
|
||||||
|
language_config = kwargs.get("language_config", {})
|
||||||
|
self.text_config = UnlimitedLanguageConfig(**language_config)
|
||||||
|
|
||||||
|
self.tile_tag = tile_tag
|
||||||
|
self.global_view_pos = global_view_pos
|
||||||
|
self.candidate_resolutions = candidate_resolutions
|
||||||
|
self.vocab_size = self.text_config.vocab_size
|
||||||
|
self.hidden_size = self.text_config.hidden_size
|
||||||
|
|
||||||
|
|
||||||
|
AutoProcessor.register(UnlimitedVLConfig, UnlimitedOCRHFProcessor)
|
||||||
|
|
||||||
|
try:
|
||||||
|
AutoConfig.register("unlimited-ocr", UnlimitedVLConfig)
|
||||||
|
except ValueError:
|
||||||
|
pass
|
||||||
@@ -361,6 +361,8 @@ class CompletionRequest(BaseModel):
|
|||||||
custom_params: Optional[Dict] = None
|
custom_params: Optional[Dict] = None
|
||||||
custom_logit_processor: Optional[str] = None
|
custom_logit_processor: Optional[str] = None
|
||||||
|
|
||||||
|
images_config: Optional[Dict] = None
|
||||||
|
|
||||||
# For PD disaggregation
|
# For PD disaggregation
|
||||||
bootstrap_host: Optional[Union[List[str], str]] = None
|
bootstrap_host: Optional[Union[List[str], str]] = None
|
||||||
bootstrap_port: Optional[Union[List[Optional[int]], int]] = None
|
bootstrap_port: Optional[Union[List[Optional[int]], int]] = None
|
||||||
@@ -736,6 +738,8 @@ class ChatCompletionRequest(BaseModel):
|
|||||||
min_dynamic_patch: Optional[int] = None
|
min_dynamic_patch: Optional[int] = None
|
||||||
use_audio_in_video: bool = False
|
use_audio_in_video: bool = False
|
||||||
|
|
||||||
|
images_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
|
||||||
custom_params: Optional[Dict] = None
|
custom_params: Optional[Dict] = None
|
||||||
|
|||||||
@@ -612,6 +612,7 @@ class OpenAIServingChat(OpenAIServingBase):
|
|||||||
routing_key=self.extract_routing_key(raw_request),
|
routing_key=self.extract_routing_key(raw_request),
|
||||||
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),
|
||||||
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),
|
||||||
|
|||||||
@@ -130,6 +130,7 @@ class OpenAIServingCompletion(OpenAIServingBase):
|
|||||||
routing_key=self.extract_routing_key(raw_request),
|
routing_key=self.extract_routing_key(raw_request),
|
||||||
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),
|
||||||
)
|
)
|
||||||
|
|
||||||
return adapted_request, request
|
return adapted_request, request
|
||||||
|
|||||||
@@ -5,6 +5,8 @@ from typing import TYPE_CHECKING, Optional
|
|||||||
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import torch
|
import torch
|
||||||
|
import triton
|
||||||
|
import triton.language as tl
|
||||||
|
|
||||||
from sglang.srt.configs.model_config import AttentionArch
|
from sglang.srt.configs.model_config import AttentionArch
|
||||||
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
|
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
|
||||||
@@ -40,6 +42,95 @@ from sglang.jit_kernel.flash_attention import (
|
|||||||
from sglang.srt.model_executor.cuda_graph_config import cuda_graph_fully_disabled
|
from sglang.srt.model_executor.cuda_graph_config import cuda_graph_fully_disabled
|
||||||
|
|
||||||
|
|
||||||
|
@triton.jit
|
||||||
|
def _build_pa_page_table_kernel(
|
||||||
|
req_to_token_ptr,
|
||||||
|
req_pool_indices_ptr,
|
||||||
|
seq_lens_ptr,
|
||||||
|
prefill_lens_ptr,
|
||||||
|
dst_page_table_ptr,
|
||||||
|
kv_lens_ptr,
|
||||||
|
window_size: tl.constexpr,
|
||||||
|
req_to_token_stride,
|
||||||
|
dst_stride,
|
||||||
|
BLOCK_SIZE: tl.constexpr,
|
||||||
|
):
|
||||||
|
"""Build PA-SWA page_table directly from req_to_token.
|
||||||
|
|
||||||
|
For each request, dst row = [0..prefill_len) ∪ [decode_start..seq_len).
|
||||||
|
decode_start = max(prefill_len, seq_len - window_size)
|
||||||
|
|
||||||
|
prefill_lens_ptr is the full pool-sized buffer, prefill_len is loaded
|
||||||
|
via indirect indexing using req_idx.
|
||||||
|
"""
|
||||||
|
bid = tl.program_id(0)
|
||||||
|
req_idx = tl.load(req_pool_indices_ptr + bid)
|
||||||
|
sl = tl.load(seq_lens_ptr + bid).to(tl.int32)
|
||||||
|
pf = tl.load(prefill_lens_ptr + req_idx).to(tl.int32)
|
||||||
|
|
||||||
|
decode_start = tl.maximum(pf, sl - window_size)
|
||||||
|
gap = tl.where(decode_start > pf, decode_start - pf, 0)
|
||||||
|
kv_len = sl - gap
|
||||||
|
|
||||||
|
tl.store(kv_lens_ptr + bid, kv_len)
|
||||||
|
|
||||||
|
src_base = req_idx * req_to_token_stride
|
||||||
|
dst_base = bid * dst_stride
|
||||||
|
|
||||||
|
for start in tl.range(0, kv_len, BLOCK_SIZE):
|
||||||
|
offs = start + tl.arange(0, BLOCK_SIZE)
|
||||||
|
mask = offs < kv_len
|
||||||
|
pos = tl.where(offs < pf, offs, offs + gap)
|
||||||
|
kv_loc = tl.load(
|
||||||
|
req_to_token_ptr + src_base + pos,
|
||||||
|
mask=mask,
|
||||||
|
other=0,
|
||||||
|
)
|
||||||
|
tl.store(dst_page_table_ptr + dst_base + offs, kv_loc.to(tl.int32), mask=mask)
|
||||||
|
|
||||||
|
|
||||||
|
def _build_pa_page_table(
|
||||||
|
req_to_token: torch.Tensor,
|
||||||
|
req_pool_indices: torch.Tensor,
|
||||||
|
seq_lens: torch.Tensor,
|
||||||
|
prefill_lens: torch.Tensor,
|
||||||
|
window_size: int,
|
||||||
|
bs: int,
|
||||||
|
pa_max_len: int,
|
||||||
|
device: torch.device,
|
||||||
|
dst_page_table: Optional[torch.Tensor] = None,
|
||||||
|
dst_kv_lens: Optional[torch.Tensor] = None,
|
||||||
|
):
|
||||||
|
"""Build prefill-aware page_table from req_to_token.
|
||||||
|
|
||||||
|
When dst_page_table/dst_kv_lens are None, allocates new tensors (non-CUDA-graph).
|
||||||
|
When provided, writes in-place into existing buffers (CUDA-graph replay).
|
||||||
|
|
||||||
|
prefill_lens is the full pool-sized buffer; the kernel indexes it via
|
||||||
|
req_pool_indices values (indirect indexing, avoids external gather).
|
||||||
|
|
||||||
|
Returns (page_table, kv_lens).
|
||||||
|
"""
|
||||||
|
if dst_page_table is None:
|
||||||
|
dst_page_table = torch.zeros(bs, pa_max_len, dtype=torch.int32, device=device)
|
||||||
|
if dst_kv_lens is None:
|
||||||
|
dst_kv_lens = torch.empty(bs, dtype=torch.int32, device=device)
|
||||||
|
if bs > 0 and pa_max_len > 0:
|
||||||
|
_build_pa_page_table_kernel[(bs,)](
|
||||||
|
req_to_token,
|
||||||
|
req_pool_indices.contiguous(),
|
||||||
|
seq_lens.to(torch.int32),
|
||||||
|
prefill_lens,
|
||||||
|
dst_page_table,
|
||||||
|
dst_kv_lens,
|
||||||
|
window_size,
|
||||||
|
req_to_token.stride(0),
|
||||||
|
dst_page_table.stride(0),
|
||||||
|
BLOCK_SIZE=256,
|
||||||
|
)
|
||||||
|
return dst_page_table, dst_kv_lens
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class FlashAttentionMetadata:
|
class FlashAttentionMetadata:
|
||||||
"""Metadata to be init once in the model forward pass,
|
"""Metadata to be init once in the model forward pass,
|
||||||
@@ -68,6 +159,9 @@ class FlashAttentionMetadata:
|
|||||||
page_table: torch.Tensor = None
|
page_table: torch.Tensor = None
|
||||||
# Page table for Sliding Window Attention
|
# Page table for Sliding Window Attention
|
||||||
swa_page_table: torch.Tensor = None
|
swa_page_table: torch.Tensor = None
|
||||||
|
|
||||||
|
pa_swa_page_table: torch.Tensor = None
|
||||||
|
pa_swa_cache_seqlens: torch.Tensor = None
|
||||||
# full->SWA translated out_cache_loc (SWA KV-store write target)
|
# full->SWA translated out_cache_loc (SWA KV-store write target)
|
||||||
swa_out_cache_loc: torch.Tensor = None
|
swa_out_cache_loc: torch.Tensor = None
|
||||||
# Precomputed FA3 scheduler metadata (avoids per-layer prepare_varlen_num_blocks)
|
# Precomputed FA3 scheduler metadata (avoids per-layer prepare_varlen_num_blocks)
|
||||||
@@ -178,6 +272,18 @@ class FlashAttentionBackend(AttentionBackend):
|
|||||||
self.sliding_window_size is not None and self.sliding_window_size > -1
|
self.sliding_window_size is not None and self.sliding_window_size > -1
|
||||||
)
|
)
|
||||||
|
|
||||||
|
self.is_prefill_aware_swa = getattr(model_runner, "prefill_aware_swa", False)
|
||||||
|
if self.is_prefill_aware_swa:
|
||||||
|
assert self.page_size == 1, (
|
||||||
|
"Prefill-aware SWA requires page_size=1, "
|
||||||
|
f"got page_size={self.page_size}"
|
||||||
|
)
|
||||||
|
max_bs = model_runner.req_to_token_pool.size
|
||||||
|
self._pa_swa_prefill_lens = torch.zeros(
|
||||||
|
max_bs, dtype=torch.int32, device=model_runner.device
|
||||||
|
)
|
||||||
|
self._pa_swa_max_prefill_len = 0
|
||||||
|
|
||||||
# Select version
|
# Select version
|
||||||
self.fa_impl_ver = fa_impl_ver
|
self.fa_impl_ver = fa_impl_ver
|
||||||
if self.fa_impl_ver == 3:
|
if self.fa_impl_ver == 3:
|
||||||
@@ -483,6 +589,25 @@ class FlashAttentionBackend(AttentionBackend):
|
|||||||
metadata.page_table = self.req_to_token_pool.req_to_token[
|
metadata.page_table = self.req_to_token_pool.req_to_token[
|
||||||
forward_batch.req_pool_indices, : metadata.max_seq_len_k
|
forward_batch.req_pool_indices, : metadata.max_seq_len_k
|
||||||
]
|
]
|
||||||
|
|
||||||
|
if self.is_prefill_aware_swa and self.has_swa:
|
||||||
|
pa_max_len = min(
|
||||||
|
self._pa_swa_max_prefill_len + self.sliding_window_size,
|
||||||
|
metadata.max_seq_len_k,
|
||||||
|
)
|
||||||
|
|
||||||
|
pa_page_table, pa_kv_lens = _build_pa_page_table(
|
||||||
|
self.req_to_token,
|
||||||
|
forward_batch.req_pool_indices[:batch_size],
|
||||||
|
forward_batch.seq_lens,
|
||||||
|
self._pa_swa_prefill_lens,
|
||||||
|
self.sliding_window_size,
|
||||||
|
batch_size,
|
||||||
|
pa_max_len,
|
||||||
|
device,
|
||||||
|
)
|
||||||
|
metadata.pa_swa_page_table = pa_page_table
|
||||||
|
metadata.pa_swa_cache_seqlens = pa_kv_lens
|
||||||
# Precompute FA3 scheduler metadata to avoid per-layer
|
# Precompute FA3 scheduler metadata to avoid per-layer
|
||||||
# prepare_varlen_num_blocks kernel calls
|
# prepare_varlen_num_blocks kernel calls
|
||||||
metadata.scheduler_metadata = self._compute_scheduler_metadata(
|
metadata.scheduler_metadata = self._compute_scheduler_metadata(
|
||||||
@@ -668,6 +793,14 @@ class FlashAttentionBackend(AttentionBackend):
|
|||||||
if forward_batch.forward_mode == ForwardMode.EXTEND:
|
if forward_batch.forward_mode == ForwardMode.EXTEND:
|
||||||
self._maybe_init_local_attn_metadata(forward_batch, metadata, device)
|
self._maybe_init_local_attn_metadata(forward_batch, metadata, device)
|
||||||
|
|
||||||
|
if self.is_prefill_aware_swa:
|
||||||
|
self._pa_swa_prefill_lens[
|
||||||
|
forward_batch.req_pool_indices[:batch_size]
|
||||||
|
] = forward_batch.seq_lens[:batch_size].to(torch.int32)
|
||||||
|
max_pf = int(forward_batch.seq_lens_cpu[:batch_size].max().item())
|
||||||
|
if max_pf > self._pa_swa_max_prefill_len:
|
||||||
|
self._pa_swa_max_prefill_len = max_pf
|
||||||
|
|
||||||
# Encoder metadata for cross attention. Supports per-request varlen
|
# Encoder metadata for cross attention. Supports per-request varlen
|
||||||
# encoder lengths (e.g. MossVL with different image sizes per request).
|
# encoder lengths (e.g. MossVL with different image sizes per request).
|
||||||
if forward_batch.encoder_lens is not None:
|
if forward_batch.encoder_lens is not None:
|
||||||
@@ -1453,6 +1586,14 @@ class FlashAttentionBackend(AttentionBackend):
|
|||||||
)
|
)
|
||||||
cache_seqlens = metadata.cache_seqlens_int32
|
cache_seqlens = metadata.cache_seqlens_int32
|
||||||
max_seqlen_q = metadata.max_seq_len_q
|
max_seqlen_q = metadata.max_seq_len_q
|
||||||
|
|
||||||
|
pa_swa_active = False
|
||||||
|
if self.is_prefill_aware_swa and metadata.pa_swa_page_table is not None:
|
||||||
|
page_table = metadata.pa_swa_page_table
|
||||||
|
cache_seqlens = metadata.pa_swa_cache_seqlens
|
||||||
|
window_size = (-1, -1)
|
||||||
|
pa_swa_active = True
|
||||||
|
|
||||||
q_reshaped = q.contiguous().view(
|
q_reshaped = q.contiguous().view(
|
||||||
-1, layer.tp_q_head_num, layer.head_dim
|
-1, layer.tp_q_head_num, layer.head_dim
|
||||||
)
|
)
|
||||||
@@ -1465,6 +1606,7 @@ class FlashAttentionBackend(AttentionBackend):
|
|||||||
metadata.scheduler_metadata is not None
|
metadata.scheduler_metadata is not None
|
||||||
and not is_swa_layer
|
and not is_swa_layer
|
||||||
and not use_cascade_attn
|
and not use_cascade_attn
|
||||||
|
and not pa_swa_active
|
||||||
):
|
):
|
||||||
sched_meta = metadata.scheduler_metadata
|
sched_meta = metadata.scheduler_metadata
|
||||||
result = flash_attn_with_kvcache(
|
result = flash_attn_with_kvcache(
|
||||||
@@ -1995,6 +2137,9 @@ class FlashAttentionBackend(AttentionBackend):
|
|||||||
metadata.page_table = self.decode_cuda_graph_metadata["page_table"][
|
metadata.page_table = self.decode_cuda_graph_metadata["page_table"][
|
||||||
:bs, :
|
:bs, :
|
||||||
]
|
]
|
||||||
|
if self.is_prefill_aware_swa:
|
||||||
|
metadata.pa_swa_page_table = metadata.page_table
|
||||||
|
metadata.pa_swa_cache_seqlens = metadata.cache_seqlens_int32
|
||||||
if self.use_sliding_window_kv_pool:
|
if self.use_sliding_window_kv_pool:
|
||||||
metadata.swa_page_table = self.decode_cuda_graph_metadata[
|
metadata.swa_page_table = self.decode_cuda_graph_metadata[
|
||||||
"swa_page_table"
|
"swa_page_table"
|
||||||
@@ -2253,20 +2398,43 @@ class FlashAttentionBackend(AttentionBackend):
|
|||||||
metadata.page_table.shape[1],
|
metadata.page_table.shape[1],
|
||||||
"FA3 decode page_table",
|
"FA3 decode page_table",
|
||||||
)
|
)
|
||||||
normal_decode_set_metadata(
|
if self.is_prefill_aware_swa:
|
||||||
metadata.cache_seqlens_int32,
|
pa_max_len = min(
|
||||||
metadata.cu_seqlens_k,
|
self._pa_swa_max_prefill_len + self.sliding_window_size,
|
||||||
metadata.page_table,
|
max_len,
|
||||||
self.req_to_token,
|
)
|
||||||
req_pool_indices,
|
if pa_max_len > 0:
|
||||||
self.decode_cuda_graph_metadata["strided_indices"],
|
_build_pa_page_table(
|
||||||
max_seq_pages,
|
self.req_to_token,
|
||||||
seq_lens,
|
req_pool_indices,
|
||||||
0,
|
seq_lens,
|
||||||
self.page_size,
|
self._pa_swa_prefill_lens,
|
||||||
metadata.swa_page_table,
|
self.sliding_window_size,
|
||||||
self.token_to_kv_pool if self.use_sliding_window_kv_pool else None,
|
bs,
|
||||||
)
|
pa_max_len,
|
||||||
|
device,
|
||||||
|
dst_page_table=metadata.page_table,
|
||||||
|
dst_kv_lens=metadata.cache_seqlens_int32,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
normal_decode_set_metadata(
|
||||||
|
metadata.cache_seqlens_int32,
|
||||||
|
metadata.cu_seqlens_k,
|
||||||
|
metadata.page_table,
|
||||||
|
self.req_to_token,
|
||||||
|
req_pool_indices,
|
||||||
|
self.decode_cuda_graph_metadata["strided_indices"],
|
||||||
|
max_seq_pages,
|
||||||
|
seq_lens,
|
||||||
|
0,
|
||||||
|
self.page_size,
|
||||||
|
metadata.swa_page_table,
|
||||||
|
(
|
||||||
|
self.token_to_kv_pool
|
||||||
|
if self.use_sliding_window_kv_pool
|
||||||
|
else None
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
self._maybe_update_local_attn_metadata_for_replay(
|
self._maybe_update_local_attn_metadata_for_replay(
|
||||||
metadata,
|
metadata,
|
||||||
|
|||||||
@@ -292,6 +292,9 @@ class GenerateReqInput:
|
|||||||
image_max_dynamic_patch: Optional[int] = None
|
image_max_dynamic_patch: Optional[int] = None
|
||||||
video_max_dynamic_patch: Optional[int] = None
|
video_max_dynamic_patch: Optional[int] = None
|
||||||
|
|
||||||
|
# For Unlimited-OCR
|
||||||
|
images_config: Optional[dict] = None
|
||||||
|
|
||||||
# Pre-computed delimiter indices for multi-item scoring.
|
# Pre-computed delimiter indices for multi-item scoring.
|
||||||
# Batch-level: List[List[int]] (one per request). After __getitem__: List[int].
|
# Batch-level: List[List[int]] (one per request). After __getitem__: List[int].
|
||||||
multi_item_delimiter_indices: Optional[Union[List[List[int]], List[int]]] = None
|
multi_item_delimiter_indices: Optional[Union[List[List[int]], List[int]]] = None
|
||||||
|
|||||||
@@ -740,7 +740,7 @@ class Req(ReqDllmMixin):
|
|||||||
self.kv_committed_freed = False
|
self.kv_committed_freed = False
|
||||||
self.kv_overallocated_freed = False
|
self.kv_overallocated_freed = False
|
||||||
|
|
||||||
# for corss-endoder model
|
# for cross-encoder model
|
||||||
self.token_type_ids = token_type_ids
|
self.token_type_ids = token_type_ids
|
||||||
|
|
||||||
# The length of KV that have been removed in swa cache.
|
# The length of KV that have been removed in swa cache.
|
||||||
@@ -749,6 +749,9 @@ class Req(ReqDllmMixin):
|
|||||||
# `ScheduleBatch.maybe_evict_swa`; KV in range [0, cache_protected_len) is freed during radix cache eviction.
|
# `ScheduleBatch.maybe_evict_swa`; KV in range [0, cache_protected_len) is freed during radix cache eviction.
|
||||||
# - Chunk cache: KV in range [0, swa_evicted_seqlen) is freed manually in `ScheduleBatch.maybe_evict_swa`.
|
# - Chunk cache: KV in range [0, swa_evicted_seqlen) is freed manually in `ScheduleBatch.maybe_evict_swa`.
|
||||||
self.swa_evicted_seqlen = 0
|
self.swa_evicted_seqlen = 0
|
||||||
|
# Tokens in [0, swa_evict_floor) are protected from SWA window eviction.
|
||||||
|
# This is used by prefill-aware SWA models such as Unlimited-OCR to keep prompt/image KV visible during decode.
|
||||||
|
self.swa_evict_floor: int = 0
|
||||||
|
|
||||||
# The index of the extend / decode batch
|
# The index of the extend / decode batch
|
||||||
self.extend_batch_idx = 0
|
self.extend_batch_idx = 0
|
||||||
|
|||||||
@@ -41,7 +41,10 @@ from sglang.srt.managers.schedule_batch import Req, ScheduleBatch
|
|||||||
from sglang.srt.mem_cache.allocator.hisparse import (
|
from sglang.srt.mem_cache.allocator.hisparse import (
|
||||||
DeepSeekV4HiSparseTokenToKVPoolAllocator,
|
DeepSeekV4HiSparseTokenToKVPoolAllocator,
|
||||||
)
|
)
|
||||||
from sglang.srt.mem_cache.allocator.swa import SWATokenToKVPoolAllocator
|
from sglang.srt.mem_cache.allocator.swa import (
|
||||||
|
PureSWATokenToKVPoolAllocator,
|
||||||
|
SWATokenToKVPoolAllocator,
|
||||||
|
)
|
||||||
from sglang.srt.mem_cache.base_prefix_cache import (
|
from sglang.srt.mem_cache.base_prefix_cache import (
|
||||||
BasePrefixCache,
|
BasePrefixCache,
|
||||||
InitLoadBackParams,
|
InitLoadBackParams,
|
||||||
@@ -483,6 +486,9 @@ class PrefillAdder:
|
|||||||
self.token_to_kv_pool_allocator,
|
self.token_to_kv_pool_allocator,
|
||||||
(SWATokenToKVPoolAllocator, DeepSeekV4HiSparseTokenToKVPoolAllocator),
|
(SWATokenToKVPoolAllocator, DeepSeekV4HiSparseTokenToKVPoolAllocator),
|
||||||
)
|
)
|
||||||
|
self.is_all_swa = isinstance(
|
||||||
|
self.token_to_kv_pool_allocator, PureSWATokenToKVPoolAllocator
|
||||||
|
)
|
||||||
self.is_hybrid_ssm_cache = self.tree_cache.supports_mamba()
|
self.is_hybrid_ssm_cache = self.tree_cache.supports_mamba()
|
||||||
|
|
||||||
self.rem_swa_token_offset = 0
|
self.rem_swa_token_offset = 0
|
||||||
@@ -517,7 +523,12 @@ class PrefillAdder:
|
|||||||
|
|
||||||
@property
|
@property
|
||||||
def rem_total_tokens(self):
|
def rem_total_tokens(self):
|
||||||
if self.is_hybrid_swa:
|
if self.is_all_swa:
|
||||||
|
available_and_evictable = (
|
||||||
|
self.token_to_kv_pool_allocator.swa_available_size()
|
||||||
|
+ self.tree_cache.swa_evictable_size()
|
||||||
|
)
|
||||||
|
elif self.is_hybrid_swa:
|
||||||
available_and_evictable = (
|
available_and_evictable = (
|
||||||
self.token_to_kv_pool_allocator.full_available_size()
|
self.token_to_kv_pool_allocator.full_available_size()
|
||||||
+ self.tree_cache.full_evictable_size()
|
+ self.tree_cache.full_evictable_size()
|
||||||
@@ -544,7 +555,12 @@ class PrefillAdder:
|
|||||||
|
|
||||||
@property
|
@property
|
||||||
def cur_rem_tokens(self):
|
def cur_rem_tokens(self):
|
||||||
if self.is_hybrid_swa:
|
if self.is_all_swa:
|
||||||
|
available_and_evictable = (
|
||||||
|
self.token_to_kv_pool_allocator.swa_available_size()
|
||||||
|
+ self.tree_cache.swa_evictable_size()
|
||||||
|
)
|
||||||
|
elif self.is_hybrid_swa:
|
||||||
available_and_evictable = (
|
available_and_evictable = (
|
||||||
self.token_to_kv_pool_allocator.full_available_size()
|
self.token_to_kv_pool_allocator.full_available_size()
|
||||||
+ self.tree_cache.full_evictable_size()
|
+ self.tree_cache.full_evictable_size()
|
||||||
|
|||||||
@@ -2972,6 +2972,10 @@ class Scheduler(
|
|||||||
|
|
||||||
new_batch.prepare_for_extend()
|
new_batch.prepare_for_extend()
|
||||||
|
|
||||||
|
if self.tp_worker.model_runner.prefill_aware_swa:
|
||||||
|
for req in can_run_list:
|
||||||
|
req.swa_evict_floor = req.fill_len
|
||||||
|
|
||||||
# Record prefill stats for logging after forward.
|
# Record prefill stats for logging after forward.
|
||||||
new_batch.prefill_stats = PrefillStats.from_adder(
|
new_batch.prefill_stats = PrefillStats.from_adder(
|
||||||
adder,
|
adder,
|
||||||
|
|||||||
@@ -82,6 +82,8 @@ class SchedulerInvariantChecker:
|
|||||||
return leak, msg
|
return leak, msg
|
||||||
|
|
||||||
def _check_full_pool(self, ps: PoolStats, uncached: int = 0) -> Tuple[bool, str]:
|
def _check_full_pool(self, ps: PoolStats, uncached: int = 0) -> Tuple[bool, str]:
|
||||||
|
if self.is_hybrid_swa and not self.full_tokens_per_layer:
|
||||||
|
return False, ""
|
||||||
if self.is_hybrid_swa:
|
if self.is_hybrid_swa:
|
||||||
protected = self.tree_cache.full_protected_size()
|
protected = self.tree_cache.full_protected_size()
|
||||||
session_held = self.pool_stats_observer.session_held_full_tokens()
|
session_held = self.pool_stats_observer.session_held_full_tokens()
|
||||||
|
|||||||
@@ -300,7 +300,12 @@ class SchedulerPoolStatsObserver:
|
|||||||
if self.enable_hisparse:
|
if self.enable_hisparse:
|
||||||
full_num_used = max(0, full_num_used)
|
full_num_used = max(0, full_num_used)
|
||||||
swa_num_used = max(0, swa_num_used)
|
swa_num_used = max(0, swa_num_used)
|
||||||
full_token_usage = full_num_used / self.full_tokens_per_layer
|
if not self.full_tokens_per_layer:
|
||||||
|
full_num_used = 0
|
||||||
|
full_available_size = 0
|
||||||
|
full_token_usage = 0.0
|
||||||
|
else:
|
||||||
|
full_token_usage = full_num_used / self.full_tokens_per_layer
|
||||||
swa_token_usage = swa_num_used / self.swa_tokens_per_layer
|
swa_token_usage = swa_num_used / self.swa_tokens_per_layer
|
||||||
|
|
||||||
return PoolStats(
|
return PoolStats(
|
||||||
|
|||||||
@@ -385,3 +385,120 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
|
|||||||
return self._kvcache.load_cpu_copy(
|
return self._kvcache.load_cpu_copy(
|
||||||
kv_cache_cpu, indices, mamba_indices=mamba_indices
|
kv_cache_cpu, indices, mamba_indices=mamba_indices
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class PureSWATokenToKVPoolAllocator(SWATokenToKVPoolAllocator):
|
||||||
|
"""Single-pool allocator for models whose every layer is sliding-window attention."""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
size_swa: int,
|
||||||
|
page_size: int,
|
||||||
|
dtype: torch.dtype,
|
||||||
|
device: str,
|
||||||
|
kvcache: BaseSWAKVPool,
|
||||||
|
need_sort: bool,
|
||||||
|
):
|
||||||
|
assert page_size == 1
|
||||||
|
assert isinstance(kvcache, BaseSWAKVPool)
|
||||||
|
|
||||||
|
self.page_size = page_size
|
||||||
|
self.dtype = dtype
|
||||||
|
self.device = device
|
||||||
|
self.need_sort = need_sort
|
||||||
|
self._size_full = self._size_swa = size_swa
|
||||||
|
|
||||||
|
self.swa_attn_allocator = TokenToKVPoolAllocator(
|
||||||
|
size_swa,
|
||||||
|
dtype,
|
||||||
|
device,
|
||||||
|
kvcache.swa_kv_pool,
|
||||||
|
need_sort,
|
||||||
|
)
|
||||||
|
self.full_attn_allocator = self.swa_attn_allocator
|
||||||
|
|
||||||
|
self.full_to_swa_index_mapping = torch.cat(
|
||||||
|
[
|
||||||
|
torch.arange(size_swa + page_size, dtype=torch.int64, device=device),
|
||||||
|
torch.tensor([-1], dtype=torch.int64, device=device),
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
self.free_pages = None
|
||||||
|
self.release_pages = None
|
||||||
|
self.is_not_in_free_group = True
|
||||||
|
self.free_group = []
|
||||||
|
|
||||||
|
self._kvcache = kvcache
|
||||||
|
self.swa_attn_allocator.clear()
|
||||||
|
self._kvcache.register_mapping(self.full_to_swa_index_mapping)
|
||||||
|
|
||||||
|
def available_size(self):
|
||||||
|
return self.swa_attn_allocator.available_size()
|
||||||
|
|
||||||
|
def full_available_size(self):
|
||||||
|
return self.swa_attn_allocator.available_size()
|
||||||
|
|
||||||
|
def swa_available_size(self):
|
||||||
|
return self.swa_attn_allocator.available_size()
|
||||||
|
|
||||||
|
def new_pages_available(self, num_full_pages: int, num_swa_pages: int) -> bool:
|
||||||
|
avail = self.swa_attn_allocator.available_size() // self.page_size
|
||||||
|
return num_full_pages <= avail and num_swa_pages <= avail
|
||||||
|
|
||||||
|
def translate_loc_from_full_to_swa(self, kv_indices: torch.Tensor):
|
||||||
|
return kv_indices
|
||||||
|
|
||||||
|
def alloc(self, need_size: int):
|
||||||
|
assert self.page_size == 1
|
||||||
|
return self.swa_attn_allocator.alloc(need_size)
|
||||||
|
|
||||||
|
def alloc_extend(self, *args, **kwargs):
|
||||||
|
raise NotImplementedError(
|
||||||
|
"PureSWATokenToKVPoolAllocator does not support page_size > 1."
|
||||||
|
)
|
||||||
|
|
||||||
|
def alloc_decode(self, *args, **kwargs):
|
||||||
|
raise NotImplementedError(
|
||||||
|
"PureSWATokenToKVPoolAllocator does not support page_size > 1."
|
||||||
|
)
|
||||||
|
|
||||||
|
def alloc_extend_swa_tail(self, *args, **kwargs):
|
||||||
|
raise NotImplementedError(
|
||||||
|
"PureSWATokenToKVPoolAllocator does not support page_size > 1."
|
||||||
|
)
|
||||||
|
|
||||||
|
def free(self, free_index: torch.Tensor):
|
||||||
|
if free_index.numel() == 0:
|
||||||
|
return
|
||||||
|
if self.is_not_in_free_group:
|
||||||
|
self.swa_attn_allocator.free(free_index[free_index > 0])
|
||||||
|
else:
|
||||||
|
self.free_group.append(free_index)
|
||||||
|
assert self.swa_attn_allocator.available_size() <= self.swa_attn_allocator.size
|
||||||
|
|
||||||
|
def free_swa(self, free_index: torch.Tensor):
|
||||||
|
if free_index.numel() == 0:
|
||||||
|
return
|
||||||
|
self.swa_attn_allocator.free(free_index[free_index > 0])
|
||||||
|
|
||||||
|
def free_group_begin(self):
|
||||||
|
self.is_not_in_free_group = False
|
||||||
|
self.free_group = []
|
||||||
|
|
||||||
|
def free_group_end(self):
|
||||||
|
self.is_not_in_free_group = True
|
||||||
|
if self.free_group:
|
||||||
|
self.free(torch.cat(self.free_group))
|
||||||
|
self.free_group = []
|
||||||
|
|
||||||
|
def backup_state(self):
|
||||||
|
return self.swa_attn_allocator.backup_state()
|
||||||
|
|
||||||
|
def restore_state(self, state):
|
||||||
|
self.swa_attn_allocator.restore_state(state)
|
||||||
|
|
||||||
|
def clear(self):
|
||||||
|
self.swa_attn_allocator.clear()
|
||||||
|
self.is_not_in_free_group = True
|
||||||
|
self.free_group = []
|
||||||
|
|||||||
@@ -135,3 +135,35 @@ class SWAChunkCache(ChunkCache):
|
|||||||
|
|
||||||
def evict(self, params: EvictParams) -> EvictResult:
|
def evict(self, params: EvictParams) -> EvictResult:
|
||||||
return EvictResult()
|
return EvictResult()
|
||||||
|
|
||||||
|
|
||||||
|
class PureSWAChunkCache(SWAChunkCache):
|
||||||
|
"""ChunkCache for all-SWA models (no full attention layers).
|
||||||
|
|
||||||
|
For hybrid models, full_to_swa_index_mapping prevents SWA double-free.
|
||||||
|
All-SWA models lack this mapping, so on request completion we must
|
||||||
|
explicitly skip the range already freed by ``free_swa_out_of_window_slots``
|
||||||
|
(a.k.a. _evict_swa) during decode.
|
||||||
|
|
||||||
|
``req.swa_evict_floor`` only protects the prompt/image KV while the request
|
||||||
|
is active. ChunkCache does not retain finished prefixes, so the protected
|
||||||
|
prefix is released here when the request finishes.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def cache_finished_req(self, req: Req, is_insert: bool = True):
|
||||||
|
kv_committed_len = req.pop_committed_kv_cache()
|
||||||
|
kv_indices = self.req_to_token_pool.req_to_token[
|
||||||
|
req.req_pool_idx, :kv_committed_len
|
||||||
|
]
|
||||||
|
evict_floor = req.swa_evict_floor
|
||||||
|
evicted_seqlen = req.swa_evicted_seqlen
|
||||||
|
if evicted_seqlen > evict_floor:
|
||||||
|
parts = []
|
||||||
|
if evict_floor > 0:
|
||||||
|
parts.append(kv_indices[:evict_floor])
|
||||||
|
if evicted_seqlen < kv_committed_len:
|
||||||
|
parts.append(kv_indices[evicted_seqlen:kv_committed_len])
|
||||||
|
if parts:
|
||||||
|
self.token_to_kv_pool_allocator.free(torch.cat(parts))
|
||||||
|
else:
|
||||||
|
self.token_to_kv_pool_allocator.free(kv_indices)
|
||||||
|
|||||||
@@ -81,7 +81,10 @@ def free_swa_out_of_window_slots(
|
|||||||
assert (
|
assert (
|
||||||
req.cache_protected_len % page_size == 0
|
req.cache_protected_len % page_size == 0
|
||||||
), "cache_protected_len must be page aligned"
|
), "cache_protected_len must be page aligned"
|
||||||
req.swa_evicted_seqlen = max(req.swa_evicted_seqlen, req.cache_protected_len)
|
evict_floor = max(req.cache_protected_len, getattr(req, "swa_evict_floor", 0))
|
||||||
|
if page_size > 1 and evict_floor > req.cache_protected_len:
|
||||||
|
evict_floor = -(-evict_floor // page_size) * page_size
|
||||||
|
req.swa_evicted_seqlen = max(req.swa_evicted_seqlen, evict_floor)
|
||||||
|
|
||||||
# Subtract an extra page_size so the eviction frontier never reaches the
|
# Subtract an extra page_size so the eviction frontier never reaches the
|
||||||
# radix tree insert boundary (page_floor(seq_len)). This keeps at least one
|
# radix tree insert boundary (page_floor(seq_len)). This keeps at least one
|
||||||
|
|||||||
@@ -232,6 +232,7 @@ def build_kv_cache(
|
|||||||
server_args=server_args,
|
server_args=server_args,
|
||||||
params=params,
|
params=params,
|
||||||
is_hybrid_swa=is_hybrid_swa,
|
is_hybrid_swa=is_hybrid_swa,
|
||||||
|
full_tokens_per_layer=full_tokens_per_layer,
|
||||||
is_hybrid_ssm=is_hybrid_ssm,
|
is_hybrid_ssm=is_hybrid_ssm,
|
||||||
enable_hierarchical_cache=enable_hierarchical_cache,
|
enable_hierarchical_cache=enable_hierarchical_cache,
|
||||||
disable_radix_cache=disable_radix_cache,
|
disable_radix_cache=disable_radix_cache,
|
||||||
|
|||||||
@@ -0,0 +1,148 @@
|
|||||||
|
"""Radix cache for all-SWA models (every layer is sliding-window attention)."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import logging
|
||||||
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from sglang.srt.mem_cache.base_prefix_cache import (
|
||||||
|
EvictParams,
|
||||||
|
EvictResult,
|
||||||
|
InsertParams,
|
||||||
|
)
|
||||||
|
from sglang.srt.mem_cache.cache_init_params import CacheInitParams
|
||||||
|
from sglang.srt.mem_cache.radix_cache import RadixCache, RadixKey
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from sglang.srt.managers.schedule_batch import Req
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
class PureSWARadixCache(RadixCache):
|
||||||
|
"""Radix cache for all-SWA models (no full attention layers).
|
||||||
|
|
||||||
|
Extends RadixCache with SWA semantics. Only caches the prefill portion
|
||||||
|
[0, evict_floor) on request completion. Window-range KV is freed.
|
||||||
|
No tombstone mechanism needed.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, params: CacheInitParams):
|
||||||
|
super().__init__(params)
|
||||||
|
self.sliding_window_size = params.sliding_window_size
|
||||||
|
|
||||||
|
def supports_swa(self) -> bool:
|
||||||
|
assert (
|
||||||
|
self.sliding_window_size is not None
|
||||||
|
), "sliding_window_size must be set for PureSWARadixCache"
|
||||||
|
return True
|
||||||
|
|
||||||
|
def swa_evictable_size(self):
|
||||||
|
return self.evictable_size_
|
||||||
|
|
||||||
|
def swa_protected_size(self):
|
||||||
|
return self.protected_size_
|
||||||
|
|
||||||
|
def full_evictable_size(self):
|
||||||
|
return 0
|
||||||
|
|
||||||
|
def full_protected_size(self):
|
||||||
|
return 0
|
||||||
|
|
||||||
|
def sanity_check(self):
|
||||||
|
"""No-op: PureSWARadixCache uses RadixCache's simple tree structure
|
||||||
|
which doesn't need the dual-LRU sanity checks of SWARadixCache."""
|
||||||
|
pass
|
||||||
|
|
||||||
|
def evict(self, params: EvictParams) -> EvictResult:
|
||||||
|
"""For all-SWA models, evict_from_tree_cache passes swa_num_tokens
|
||||||
|
(with num_tokens=0). Use whichever is non-zero."""
|
||||||
|
num_tokens = max(params.num_tokens, params.swa_num_tokens)
|
||||||
|
return super().evict(EvictParams(num_tokens=num_tokens))
|
||||||
|
|
||||||
|
def cache_finished_req(self, req: Req, is_insert: bool = True):
|
||||||
|
"""Cache request when it finishes.
|
||||||
|
|
||||||
|
Only inserts the prefill portion [0, evict_floor) into the radix tree.
|
||||||
|
The window portion [swa_evicted_seqlen, committed_len) is freed back
|
||||||
|
to the allocator. The range [evict_floor, swa_evicted_seqlen) was already
|
||||||
|
freed by _evict_swa during decode — we skip it to avoid double-free.
|
||||||
|
"""
|
||||||
|
if self.disable_finished_insert:
|
||||||
|
is_insert = False
|
||||||
|
|
||||||
|
kv_committed_len = req.pop_committed_kv_cache()
|
||||||
|
if self.disable:
|
||||||
|
kv_indices = self.req_to_token_pool.req_to_token[
|
||||||
|
req.req_pool_idx, :kv_committed_len
|
||||||
|
]
|
||||||
|
self.token_to_kv_pool_allocator.free(kv_indices)
|
||||||
|
return
|
||||||
|
|
||||||
|
token_ids = (req.origin_input_ids + req.output_ids)[:kv_committed_len]
|
||||||
|
kv_indices = self.req_to_token_pool.req_to_token[
|
||||||
|
req.req_pool_idx, :kv_committed_len
|
||||||
|
]
|
||||||
|
|
||||||
|
radix_key = RadixKey(
|
||||||
|
token_ids, req.extra_key, is_bigram=self.is_eagle
|
||||||
|
).page_aligned(self.page_size)
|
||||||
|
keys_len = len(radix_key)
|
||||||
|
|
||||||
|
old_prefix_len = req.cache_protected_len
|
||||||
|
swa_evict_floor = req.swa_evict_floor
|
||||||
|
swa_evicted_seqlen = req.swa_evicted_seqlen
|
||||||
|
|
||||||
|
if self.page_size > 1 and swa_evict_floor > 0:
|
||||||
|
swa_evict_floor = -(-swa_evict_floor // self.page_size) * self.page_size
|
||||||
|
|
||||||
|
if swa_evict_floor > 0:
|
||||||
|
insert_end = min(swa_evict_floor, keys_len)
|
||||||
|
else:
|
||||||
|
insert_end = keys_len
|
||||||
|
|
||||||
|
if is_insert and insert_end > 0:
|
||||||
|
insert_values = kv_indices[:insert_end].to(dtype=torch.int64, copy=True)
|
||||||
|
result = self.insert(
|
||||||
|
InsertParams(key=radix_key[:insert_end], value=insert_values)
|
||||||
|
)
|
||||||
|
new_prefix_len = result.prefix_len
|
||||||
|
if new_prefix_len > old_prefix_len:
|
||||||
|
self.token_to_kv_pool_allocator.free(
|
||||||
|
kv_indices[old_prefix_len:new_prefix_len]
|
||||||
|
)
|
||||||
|
alive_start = max(swa_evicted_seqlen, insert_end)
|
||||||
|
if alive_start < keys_len:
|
||||||
|
self.token_to_kv_pool_allocator.free(kv_indices[alive_start:keys_len])
|
||||||
|
else:
|
||||||
|
free_end = (
|
||||||
|
min(swa_evict_floor, keys_len) if swa_evict_floor > 0 else keys_len
|
||||||
|
)
|
||||||
|
if free_end > old_prefix_len:
|
||||||
|
self.token_to_kv_pool_allocator.free(
|
||||||
|
kv_indices[old_prefix_len:free_end]
|
||||||
|
)
|
||||||
|
alive_start = max(swa_evicted_seqlen, old_prefix_len)
|
||||||
|
if swa_evicted_seqlen > 0 and alive_start < keys_len:
|
||||||
|
self.token_to_kv_pool_allocator.free(kv_indices[alive_start:keys_len])
|
||||||
|
|
||||||
|
self.token_to_kv_pool_allocator.free(kv_indices[keys_len:])
|
||||||
|
|
||||||
|
if req.last_node is not None:
|
||||||
|
self.dec_lock_ref(req.last_node)
|
||||||
|
|
||||||
|
def cache_unfinished_req(self, req: Req, chunked=False):
|
||||||
|
"""During chunked prefill, swa_evicted_seqlen is 0 and no SWA eviction
|
||||||
|
has happened yet, so standard RadixCache logic is correct."""
|
||||||
|
super().cache_unfinished_req(req, chunked=chunked)
|
||||||
|
|
||||||
|
def available_and_evictable_str(self) -> str:
|
||||||
|
allocator = self.token_to_kv_pool_allocator
|
||||||
|
swa_available = allocator.swa_available_size()
|
||||||
|
swa_evictable = self.swa_evictable_size()
|
||||||
|
return (
|
||||||
|
f"SWA available tokens: {swa_available + swa_evictable} "
|
||||||
|
f"({swa_available=} + {swa_evictable=})\n"
|
||||||
|
)
|
||||||
@@ -42,6 +42,7 @@ class TreeCacheBuildContext:
|
|||||||
tp_size: int
|
tp_size: int
|
||||||
tp_rank: int
|
tp_rank: int
|
||||||
tp_group: Any
|
tp_group: Any
|
||||||
|
full_tokens_per_layer: Optional[int] = None
|
||||||
|
|
||||||
|
|
||||||
RadixCacheFactory = Callable[[TreeCacheBuildContext], BasePrefixCache]
|
RadixCacheFactory = Callable[[TreeCacheBuildContext], BasePrefixCache]
|
||||||
@@ -84,6 +85,10 @@ def default_radix_cache_factory(ctx: TreeCacheBuildContext) -> BasePrefixCache:
|
|||||||
from sglang.srt.mem_cache.chunk_cache import ChunkCache
|
from sglang.srt.mem_cache.chunk_cache import ChunkCache
|
||||||
|
|
||||||
return ChunkCache(params)
|
return ChunkCache(params)
|
||||||
|
if ctx.full_tokens_per_layer == 0:
|
||||||
|
from sglang.srt.mem_cache.chunk_cache import PureSWAChunkCache
|
||||||
|
|
||||||
|
return PureSWAChunkCache(params)
|
||||||
from sglang.srt.mem_cache.chunk_cache import SWAChunkCache
|
from sglang.srt.mem_cache.chunk_cache import SWAChunkCache
|
||||||
|
|
||||||
return SWAChunkCache(params)
|
return SWAChunkCache(params)
|
||||||
@@ -112,6 +117,10 @@ def default_radix_cache_factory(ctx: TreeCacheBuildContext) -> BasePrefixCache:
|
|||||||
return cache
|
return cache
|
||||||
|
|
||||||
if ctx.is_hybrid_swa:
|
if ctx.is_hybrid_swa:
|
||||||
|
if ctx.full_tokens_per_layer == 0:
|
||||||
|
from sglang.srt.mem_cache.pure_swa_radix_cache import PureSWARadixCache
|
||||||
|
|
||||||
|
return PureSWARadixCache(params=params)
|
||||||
from sglang.srt.mem_cache.swa_radix_cache import SWARadixCache
|
from sglang.srt.mem_cache.swa_radix_cache import SWARadixCache
|
||||||
|
|
||||||
return SWARadixCache(params=params)
|
return SWARadixCache(params=params)
|
||||||
|
|||||||
@@ -1501,6 +1501,11 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
f"Setting sliding_window_size to be attention_chunk_size: {self.sliding_window_size}"
|
f"Setting sliding_window_size to be attention_chunk_size: {self.sliding_window_size}"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
self.prefill_aware_swa = (
|
||||||
|
hasattr(self.model, "is_prefill_aware_swa")
|
||||||
|
and self.model.is_prefill_aware_swa()
|
||||||
|
)
|
||||||
|
|
||||||
self.dtype = self.model_config.dtype
|
self.dtype = self.model_config.dtype
|
||||||
|
|
||||||
after_avail_memory = get_available_gpu_memory(self.device, self.gpu_id)
|
after_avail_memory = get_available_gpu_memory(self.device, self.gpu_id)
|
||||||
@@ -2318,7 +2323,7 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
def max_token_pool_size(self):
|
def max_token_pool_size(self):
|
||||||
"""Return the max token pool size considering hybrid swa settings."""
|
"""Return the max token pool size considering hybrid swa settings."""
|
||||||
if self.is_hybrid_swa:
|
if self.is_hybrid_swa:
|
||||||
return self.full_max_total_num_tokens
|
return self.full_max_total_num_tokens or self.swa_max_total_num_tokens
|
||||||
else:
|
else:
|
||||||
return self.max_total_num_tokens
|
return self.max_total_num_tokens
|
||||||
|
|
||||||
|
|||||||
@@ -24,7 +24,10 @@ from sglang.srt.mem_cache.allocator.hisparse import (
|
|||||||
DeepSeekV4HiSparseTokenToKVPoolAllocator,
|
DeepSeekV4HiSparseTokenToKVPoolAllocator,
|
||||||
HiSparseTokenToKVPoolAllocator,
|
HiSparseTokenToKVPoolAllocator,
|
||||||
)
|
)
|
||||||
from sglang.srt.mem_cache.allocator.swa import SWATokenToKVPoolAllocator
|
from sglang.srt.mem_cache.allocator.swa import (
|
||||||
|
PureSWATokenToKVPoolAllocator,
|
||||||
|
SWATokenToKVPoolAllocator,
|
||||||
|
)
|
||||||
from sglang.srt.mem_cache.common import get_req_to_token_extra_context_len
|
from sglang.srt.mem_cache.common import get_req_to_token_extra_context_len
|
||||||
from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool
|
from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool
|
||||||
from sglang.srt.mem_cache.hisparse_memory_pool import HiSparseDSATokenToKVPool
|
from sglang.srt.mem_cache.hisparse_memory_pool import HiSparseDSATokenToKVPool
|
||||||
@@ -861,7 +864,16 @@ class ModelRunnerKVCacheMixin:
|
|||||||
need_sort=need_sort,
|
need_sort=need_sort,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
if self.is_hybrid_swa:
|
if self.is_hybrid_swa and self.full_max_total_num_tokens == 0:
|
||||||
|
self.token_to_kv_pool_allocator = PureSWATokenToKVPoolAllocator(
|
||||||
|
self.swa_max_total_num_tokens,
|
||||||
|
page_size=self.page_size,
|
||||||
|
dtype=self.kv_cache_dtype,
|
||||||
|
device=self.device,
|
||||||
|
kvcache=self.token_to_kv_pool,
|
||||||
|
need_sort=need_sort,
|
||||||
|
)
|
||||||
|
elif self.is_hybrid_swa:
|
||||||
self.token_to_kv_pool_allocator = SWATokenToKVPoolAllocator(
|
self.token_to_kv_pool_allocator = SWATokenToKVPoolAllocator(
|
||||||
self.full_max_total_num_tokens,
|
self.full_max_total_num_tokens,
|
||||||
self.swa_max_total_num_tokens,
|
self.swa_max_total_num_tokens,
|
||||||
|
|||||||
@@ -0,0 +1,438 @@
|
|||||||
|
"""Standalone UNLIMITED-OCR model (SAM + CLIP vision encoders, Deepseek backbone)."""
|
||||||
|
|
||||||
|
import logging
|
||||||
|
from typing import Iterable, List, Optional, Set, Tuple, TypeAlias, Union
|
||||||
|
|
||||||
|
import torch
|
||||||
|
from torch import Tensor, nn
|
||||||
|
|
||||||
|
from sglang.srt.configs.unlimited_ocr import UnlimitedVLConfig
|
||||||
|
from sglang.srt.layers.quantization import QuantizationConfig
|
||||||
|
from sglang.srt.managers.mm_utils import (
|
||||||
|
MultiModalityDataPaddingPatternMultimodalTokens,
|
||||||
|
general_mm_embed_routine,
|
||||||
|
)
|
||||||
|
from sglang.srt.managers.schedule_batch import MultimodalDataItem, MultimodalInputs
|
||||||
|
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||||
|
from sglang.srt.model_loader.weight_utils import default_weight_loader
|
||||||
|
from sglang.srt.models.deepseek import DeepseekForCausalLM
|
||||||
|
from sglang.srt.models.deepseek_ocr import (
|
||||||
|
MlpProjector,
|
||||||
|
build_clip_l,
|
||||||
|
build_sam_vit_b,
|
||||||
|
merge_multimodal_embeddings,
|
||||||
|
)
|
||||||
|
from sglang.srt.models.transformers import maybe_prefix
|
||||||
|
from sglang.srt.utils import cpu_has_amx_support, is_cpu
|
||||||
|
|
||||||
|
_is_cpu_amx_available = cpu_has_amx_support()
|
||||||
|
_is_cpu = is_cpu()
|
||||||
|
|
||||||
|
NestedTensors: TypeAlias = Union[
|
||||||
|
list["NestedTensors"],
|
||||||
|
list["torch.Tensor"],
|
||||||
|
"torch.Tensor",
|
||||||
|
tuple["torch.Tensor", ...],
|
||||||
|
]
|
||||||
|
|
||||||
|
MultiModalEmbeddings: TypeAlias = list[Tensor] | Tensor | tuple[Tensor, ...]
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
class UnlimitedOCRForCausalLM(nn.Module):
|
||||||
|
"""Standalone UNLIMITED-OCR model (SAM + CLIP ViT) with prefill-aware SWA."""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
config: UnlimitedVLConfig,
|
||||||
|
quant_config: Optional[QuantizationConfig] = None,
|
||||||
|
prefix: str = "",
|
||||||
|
):
|
||||||
|
"""Initialize UnlimitedOCRForCausalLM with vision encoders, projector, and LM."""
|
||||||
|
super().__init__()
|
||||||
|
|
||||||
|
self.config = config
|
||||||
|
self.vision_config = config.vision_config
|
||||||
|
self.projector_config = config.projector_config
|
||||||
|
self.text_config = config.text_config
|
||||||
|
|
||||||
|
n_embed = getattr(self.projector_config, "n_embed", 1280)
|
||||||
|
|
||||||
|
self.tile_tag = config.tile_tag
|
||||||
|
self.global_view_pos = config.global_view_pos
|
||||||
|
|
||||||
|
embed_std = 1 / torch.sqrt(torch.tensor(n_embed, dtype=torch.float32))
|
||||||
|
if self.tile_tag == "2D":
|
||||||
|
self.view_seperator = nn.Parameter(torch.randn(n_embed) * embed_std)
|
||||||
|
self.image_newline = nn.Parameter(torch.randn(n_embed) * embed_std)
|
||||||
|
else:
|
||||||
|
raise ValueError(
|
||||||
|
f"Only 2D tile_tag is supported currently, got: {self.tile_tag}"
|
||||||
|
)
|
||||||
|
|
||||||
|
self.model = DeepseekForCausalLM(
|
||||||
|
config=config.text_config,
|
||||||
|
quant_config=quant_config,
|
||||||
|
prefix=maybe_prefix(prefix, "language"),
|
||||||
|
)
|
||||||
|
|
||||||
|
self.sam_model = build_sam_vit_b()
|
||||||
|
self.vision_model = build_clip_l()
|
||||||
|
|
||||||
|
self.projector = MlpProjector(
|
||||||
|
projector_type=self.projector_config.projector_type,
|
||||||
|
input_dim=self.projector_config.input_dim,
|
||||||
|
n_embed=n_embed,
|
||||||
|
depth=self.projector_config.depth,
|
||||||
|
mlp_ratio=self.projector_config.mlp_ratio,
|
||||||
|
downsample_ratio=self.projector_config.downsample_ratio,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.image_token_id = None
|
||||||
|
|
||||||
|
def get_attention_sliding_window_size(self) -> Optional[int]:
|
||||||
|
"""Return the sliding window size from the model config, or None."""
|
||||||
|
return getattr(self.config, "sliding_window_size", None)
|
||||||
|
|
||||||
|
def is_prefill_aware_swa(self) -> bool:
|
||||||
|
"""Prefill tokens are always retained in KV cache during decode."""
|
||||||
|
return True
|
||||||
|
|
||||||
|
def _encode_ocr1_features(self, images: torch.Tensor) -> torch.Tensor:
|
||||||
|
"""Encode images through SAM and CLIP encoders, then project features."""
|
||||||
|
features_1 = self.sam_model(images)
|
||||||
|
features_2 = self.vision_model(images, features_1)
|
||||||
|
features = torch.cat(
|
||||||
|
(
|
||||||
|
features_2[:, 1:],
|
||||||
|
features_1.flatten(2).permute(0, 2, 1),
|
||||||
|
),
|
||||||
|
dim=-1,
|
||||||
|
)
|
||||||
|
return self.projector(features)
|
||||||
|
|
||||||
|
def _format_ocr1_global_features(self, features: torch.Tensor) -> torch.Tensor:
|
||||||
|
"""Reshape global features into a flat sequence with newline tokens."""
|
||||||
|
_, hw, n_dim = features.shape
|
||||||
|
h = w = int(hw**0.5)
|
||||||
|
features = features.view(h, w, n_dim)
|
||||||
|
features = torch.cat(
|
||||||
|
[features, self.image_newline[None, None, :].expand(h, 1, n_dim)],
|
||||||
|
dim=1,
|
||||||
|
)
|
||||||
|
return features.view(-1, n_dim)
|
||||||
|
|
||||||
|
def _format_ocr1_local_features(
|
||||||
|
self, features: torch.Tensor, crop_shape: torch.Tensor
|
||||||
|
) -> torch.Tensor:
|
||||||
|
"""Reshape local crop features into a flat sequence with newline tokens."""
|
||||||
|
_, hw2, n_dim2 = features.shape
|
||||||
|
h2 = w2 = int(hw2**0.5)
|
||||||
|
width_crop_num, height_crop_num = int(crop_shape[0]), int(crop_shape[1])
|
||||||
|
features = (
|
||||||
|
features.view(height_crop_num, width_crop_num, h2, w2, n_dim2)
|
||||||
|
.permute(0, 2, 1, 3, 4)
|
||||||
|
.reshape(height_crop_num * h2, width_crop_num * w2, n_dim2)
|
||||||
|
)
|
||||||
|
features = torch.cat(
|
||||||
|
[
|
||||||
|
features,
|
||||||
|
self.image_newline[None, None, :].expand(
|
||||||
|
height_crop_num * h2, 1, n_dim2
|
||||||
|
),
|
||||||
|
],
|
||||||
|
dim=1,
|
||||||
|
)
|
||||||
|
return features.view(-1, n_dim2)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _collect_mm_flag(
|
||||||
|
items: List[MultimodalDataItem], flag_name: str
|
||||||
|
) -> Optional[List[bool]]:
|
||||||
|
"""Collect a boolean multimodal flag from all data items."""
|
||||||
|
values = []
|
||||||
|
for item in items:
|
||||||
|
value = getattr(item, flag_name, None)
|
||||||
|
if value is None:
|
||||||
|
return None
|
||||||
|
if isinstance(value, list):
|
||||||
|
values.extend(value)
|
||||||
|
else:
|
||||||
|
values.append(bool(value))
|
||||||
|
return values
|
||||||
|
|
||||||
|
def _parse_and_validate_image_input(self, **kwargs: object):
|
||||||
|
"""Parse and validate pixel values, spatial crops, and image crops."""
|
||||||
|
pixel_values = kwargs.pop("pixel_values", None)
|
||||||
|
images_spatial_crop = kwargs.pop("images_spatial_crop", None)
|
||||||
|
images_crop = kwargs.pop("images_crop", None)
|
||||||
|
has_images = kwargs.pop("has_images", None)
|
||||||
|
|
||||||
|
if pixel_values is None:
|
||||||
|
return None
|
||||||
|
if has_images is not None:
|
||||||
|
if not has_images:
|
||||||
|
return None
|
||||||
|
elif torch.sum(pixel_values).item() == 0:
|
||||||
|
return None
|
||||||
|
|
||||||
|
if pixel_values is not None:
|
||||||
|
if not isinstance(pixel_values, (torch.Tensor, list)):
|
||||||
|
raise ValueError(
|
||||||
|
"Incorrect type of pixel values. " f"Got type: {type(pixel_values)}"
|
||||||
|
)
|
||||||
|
if not isinstance(images_spatial_crop, (torch.Tensor, list)):
|
||||||
|
raise ValueError(
|
||||||
|
"Incorrect type of image sizes. "
|
||||||
|
f"Got type: {type(images_spatial_crop)}"
|
||||||
|
)
|
||||||
|
if not isinstance(images_crop, (torch.Tensor, list)):
|
||||||
|
raise ValueError(
|
||||||
|
"Incorrect type of image crop. " f"Got type: {type(images_crop)}"
|
||||||
|
)
|
||||||
|
return [pixel_values, images_crop, images_spatial_crop]
|
||||||
|
|
||||||
|
raise AssertionError("This line should be unreachable.")
|
||||||
|
|
||||||
|
def _pixel_values_to_embedding(
|
||||||
|
self,
|
||||||
|
pixel_values: torch.Tensor,
|
||||||
|
images_crop: torch.Tensor,
|
||||||
|
images_spatial_crop: torch.Tensor,
|
||||||
|
has_local_crops: Optional[List[bool]] = None,
|
||||||
|
) -> NestedTensors:
|
||||||
|
"""Encode pixel values into per-image embedding sequences."""
|
||||||
|
images_in_this_batch = []
|
||||||
|
|
||||||
|
with torch.no_grad():
|
||||||
|
for jdx in range(images_spatial_crop.size(0)):
|
||||||
|
patches = images_crop[jdx][0].to(torch.bfloat16)
|
||||||
|
image_ori = pixel_values[jdx]
|
||||||
|
crop_shape = images_spatial_crop[jdx][0]
|
||||||
|
use_local_crops = (
|
||||||
|
has_local_crops[jdx]
|
||||||
|
if has_local_crops is not None
|
||||||
|
else torch.sum(patches).item() != 0
|
||||||
|
)
|
||||||
|
|
||||||
|
global_features = self._encode_ocr1_features(image_ori)
|
||||||
|
global_features = self._format_ocr1_global_features(global_features)
|
||||||
|
|
||||||
|
if use_local_crops:
|
||||||
|
local_features = self._encode_ocr1_features(patches)
|
||||||
|
local_features = self._format_ocr1_local_features(
|
||||||
|
local_features, crop_shape
|
||||||
|
)
|
||||||
|
global_local_features = torch.cat(
|
||||||
|
[
|
||||||
|
local_features,
|
||||||
|
global_features,
|
||||||
|
self.view_seperator[None, :],
|
||||||
|
],
|
||||||
|
dim=0,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
global_local_features = torch.cat(
|
||||||
|
[global_features, self.view_seperator[None, :]], dim=0
|
||||||
|
)
|
||||||
|
|
||||||
|
images_in_this_batch.append(global_local_features)
|
||||||
|
|
||||||
|
return images_in_this_batch
|
||||||
|
|
||||||
|
def _process_image_input(self, mm_items: List[MultimodalDataItem]) -> torch.Tensor:
|
||||||
|
"""Process multimodal data items into concatenated vision features."""
|
||||||
|
target_dtype = self.vision_model.dtype
|
||||||
|
has_local_crops = self._collect_mm_flag(mm_items, "has_local_crops")
|
||||||
|
pixel_values = torch.stack([item.feature for item in mm_items], dim=0).type(
|
||||||
|
target_dtype
|
||||||
|
)
|
||||||
|
|
||||||
|
images_crop = (
|
||||||
|
torch.stack([item.images_crop for item in mm_items], dim=0)
|
||||||
|
.type(target_dtype)
|
||||||
|
.to(device=pixel_values.device)
|
||||||
|
)
|
||||||
|
images_spatial_crop = (
|
||||||
|
torch.cat([item.images_spatial_crop for item in mm_items], dim=0)
|
||||||
|
.type(torch.long)
|
||||||
|
.to(device=pixel_values.device)
|
||||||
|
)
|
||||||
|
pixel_values = pixel_values.view(
|
||||||
|
pixel_values.shape[0] * pixel_values.shape[1], 1, *pixel_values.shape[2:]
|
||||||
|
)
|
||||||
|
images_crop = images_crop.view(
|
||||||
|
images_crop.shape[0] * images_crop.shape[1], 1, *images_crop.shape[2:]
|
||||||
|
)
|
||||||
|
images_spatial_crop = images_spatial_crop.view(
|
||||||
|
images_spatial_crop.shape[0] * images_spatial_crop.shape[1],
|
||||||
|
1,
|
||||||
|
*images_spatial_crop.shape[2:],
|
||||||
|
)
|
||||||
|
|
||||||
|
assert images_crop.dim() == 6
|
||||||
|
assert images_spatial_crop.dim() == 3
|
||||||
|
|
||||||
|
vision_feature_lists = self._pixel_values_to_embedding(
|
||||||
|
pixel_values=pixel_values,
|
||||||
|
images_crop=images_crop,
|
||||||
|
images_spatial_crop=images_spatial_crop,
|
||||||
|
has_local_crops=has_local_crops,
|
||||||
|
)
|
||||||
|
vision_features = torch.cat(vision_feature_lists, dim=0).type(target_dtype)
|
||||||
|
return vision_features
|
||||||
|
|
||||||
|
def get_language_model(self) -> torch.nn.Module:
|
||||||
|
"""Return the underlying language model."""
|
||||||
|
return self.model
|
||||||
|
|
||||||
|
def get_multimodal_embeddings(
|
||||||
|
self, **kwargs: object
|
||||||
|
) -> Optional[MultiModalEmbeddings]:
|
||||||
|
"""Compute multimodal embeddings from image inputs, if present."""
|
||||||
|
image_input = self._parse_and_validate_image_input(**kwargs)
|
||||||
|
if image_input is None:
|
||||||
|
return None
|
||||||
|
vision_embeddings = self._process_image_input(image_input)
|
||||||
|
return vision_embeddings
|
||||||
|
|
||||||
|
def get_input_embeddings(
|
||||||
|
self,
|
||||||
|
input_ids: torch.Tensor,
|
||||||
|
multimodal_embeddings: Optional[MultiModalEmbeddings] = None,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
"""Get text embeddings and merge in multimodal embeddings if provided."""
|
||||||
|
inputs_embeds = self.model.get_input_embeddings(input_ids)
|
||||||
|
if multimodal_embeddings is not None:
|
||||||
|
inputs_embeds = merge_multimodal_embeddings(
|
||||||
|
input_ids, inputs_embeds, multimodal_embeddings, self.image_token_id
|
||||||
|
)
|
||||||
|
return inputs_embeds
|
||||||
|
|
||||||
|
def pad_input_ids(self, input_ids: List[int], mm_inputs: MultimodalInputs):
|
||||||
|
"""Pad input token IDs with multimodal placeholder tokens."""
|
||||||
|
pattern = MultiModalityDataPaddingPatternMultimodalTokens()
|
||||||
|
return pattern.pad_input_tokens(input_ids, mm_inputs)
|
||||||
|
|
||||||
|
def get_image_feature(self, items: List[MultimodalDataItem]) -> torch.Tensor:
|
||||||
|
"""Extract vision features from multimodal data items."""
|
||||||
|
vision_embeddings = self._process_image_input(items)
|
||||||
|
return vision_embeddings
|
||||||
|
|
||||||
|
def forward(
|
||||||
|
self,
|
||||||
|
input_ids: torch.Tensor,
|
||||||
|
positions: torch.Tensor,
|
||||||
|
forward_batch: ForwardBatch,
|
||||||
|
**kwargs: object,
|
||||||
|
):
|
||||||
|
"""Run the full multimodal forward pass (embed, encode, decode)."""
|
||||||
|
hidden_states = general_mm_embed_routine(
|
||||||
|
input_ids=input_ids,
|
||||||
|
forward_batch=forward_batch,
|
||||||
|
language_model=self.model,
|
||||||
|
multimodal_model=self,
|
||||||
|
positions=positions,
|
||||||
|
)
|
||||||
|
return hidden_states
|
||||||
|
|
||||||
|
def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]):
|
||||||
|
"""Load and remap checkpoint weights into the model parameters."""
|
||||||
|
stacked_params_mapping = [
|
||||||
|
(".qkv_proj", ".q_proj", "q"),
|
||||||
|
(".qkv_proj", ".k_proj", "k"),
|
||||||
|
(".qkv_proj", ".v_proj", "v"),
|
||||||
|
(".gate_up_proj", ".gate_proj", 0),
|
||||||
|
(".gate_up_proj", ".up_proj", 1),
|
||||||
|
]
|
||||||
|
|
||||||
|
params_dict = dict(self.named_parameters())
|
||||||
|
loaded_params: Set[str] = set()
|
||||||
|
for name, loaded_weight in weights:
|
||||||
|
if "rotary_emb.inv_freq" in name:
|
||||||
|
continue
|
||||||
|
if name == "lm_head.weight":
|
||||||
|
name = "model.lm_head.weight"
|
||||||
|
elif name.startswith("model."):
|
||||||
|
if (
|
||||||
|
"image_newline" in name
|
||||||
|
or ".projector" in name
|
||||||
|
or "vision_model" in name
|
||||||
|
or "sam_model" in name
|
||||||
|
or "view_seperator" in name
|
||||||
|
):
|
||||||
|
name = name[len("model.") :]
|
||||||
|
elif not (
|
||||||
|
".projector" in name
|
||||||
|
or "vision_model" in name
|
||||||
|
or "sam_model" in name
|
||||||
|
or "image_newline" in name
|
||||||
|
):
|
||||||
|
name = name.replace("model.", "model.model.")
|
||||||
|
|
||||||
|
for param_name, weight_name, shard_id in stacked_params_mapping:
|
||||||
|
if weight_name not in name:
|
||||||
|
continue
|
||||||
|
name = name.replace(weight_name, param_name)
|
||||||
|
if name.endswith(".bias") and name not in params_dict:
|
||||||
|
continue
|
||||||
|
if (
|
||||||
|
"mlp.experts." in name or "mlp.shared_experts." in name
|
||||||
|
) and name not in params_dict:
|
||||||
|
continue
|
||||||
|
param = params_dict[name]
|
||||||
|
weight_loader = param.weight_loader
|
||||||
|
weight_loader(param, loaded_weight, shard_id)
|
||||||
|
break
|
||||||
|
else:
|
||||||
|
if name.endswith(".bias") and name not in params_dict:
|
||||||
|
continue
|
||||||
|
if (
|
||||||
|
"mlp.experts." in name or "mlp.shared_experts." in name
|
||||||
|
) and name not in params_dict:
|
||||||
|
continue
|
||||||
|
param = params_dict[name]
|
||||||
|
weight_loader = getattr(param, "weight_loader", default_weight_loader)
|
||||||
|
weight_loader(param, loaded_weight)
|
||||||
|
loaded_params.add(name)
|
||||||
|
unloaded_params = params_dict.keys() - loaded_params
|
||||||
|
if unloaded_params:
|
||||||
|
raise RuntimeError(
|
||||||
|
f"Some weights are not initialized from checkpoints: {unloaded_params}"
|
||||||
|
)
|
||||||
|
self.post_load_weights()
|
||||||
|
|
||||||
|
def post_load_weights(self):
|
||||||
|
"""Apply post-loading weight transformations (e.g., AMX repacking on CPU)."""
|
||||||
|
if _is_cpu and _is_cpu_amx_available:
|
||||||
|
from sglang.srt.layers.amx_utils import _amx_process_weight_after_loading
|
||||||
|
|
||||||
|
layer_ids = int(self.config.num_hidden_layers)
|
||||||
|
first_k_dense_replace_id = (
|
||||||
|
self.config.first_k_dense_replace
|
||||||
|
if hasattr(self.config, "first_k_dense_replace")
|
||||||
|
else -1
|
||||||
|
)
|
||||||
|
moe_layer_freq_id = (
|
||||||
|
self.config.moe_layer_freq
|
||||||
|
if hasattr(self.config, "moe_layer_freq")
|
||||||
|
else 1
|
||||||
|
)
|
||||||
|
for layer_id in range(0, layer_ids):
|
||||||
|
if (
|
||||||
|
layer_id >= first_k_dense_replace_id
|
||||||
|
and layer_id % moe_layer_freq_id == 0
|
||||||
|
):
|
||||||
|
if (
|
||||||
|
hasattr(self.model, "model")
|
||||||
|
and hasattr(self.model.model, "layers")
|
||||||
|
and hasattr(self.model.model.layers[layer_id], "mlp")
|
||||||
|
):
|
||||||
|
self_moe = self.model.model.layers[layer_id].mlp
|
||||||
|
if hasattr(self_moe, "w1") and hasattr(self_moe, "w2"):
|
||||||
|
_amx_process_weight_after_loading(self_moe, ["w1", "w2"])
|
||||||
|
|
||||||
|
|
||||||
|
EntryClass = [UnlimitedOCRForCausalLM]
|
||||||
@@ -0,0 +1,119 @@
|
|||||||
|
"""Standalone UNLIMITED-OCR processor."""
|
||||||
|
|
||||||
|
import hashlib
|
||||||
|
import logging
|
||||||
|
from typing import List, Union
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
from sglang.srt.managers.schedule_batch import MultimodalProcessorOutput
|
||||||
|
from sglang.srt.models.unlimited_ocr import UnlimitedOCRForCausalLM
|
||||||
|
from sglang.srt.multimodal.processors.base_processor import (
|
||||||
|
BaseMultimodalProcessor,
|
||||||
|
MultimodalSpecialTokens,
|
||||||
|
)
|
||||||
|
|
||||||
|
_IMAGE_MODE_PRESETS = {
|
||||||
|
"tiny": (512, 512, False),
|
||||||
|
"small": (640, 640, False),
|
||||||
|
"base": (1024, 1024, False),
|
||||||
|
"large": (1280, 1280, False),
|
||||||
|
"gundam": (1024, 640, True),
|
||||||
|
}
|
||||||
|
_DEFAULT_MODE = "gundam"
|
||||||
|
|
||||||
|
|
||||||
|
def _resolve_mode(images_config, num_images: int = 1) -> dict:
|
||||||
|
"""Return processor kwargs from images_config (or default)."""
|
||||||
|
mode = _DEFAULT_MODE
|
||||||
|
if images_config:
|
||||||
|
mode = images_config.get("image_mode", _DEFAULT_MODE)
|
||||||
|
key = mode.strip().lower()
|
||||||
|
preset = _IMAGE_MODE_PRESETS.get(key)
|
||||||
|
if preset is None:
|
||||||
|
logger.error(
|
||||||
|
f"Unknown image_mode '{mode}'. Supported: {', '.join(_IMAGE_MODE_PRESETS)}"
|
||||||
|
)
|
||||||
|
raise ValueError(
|
||||||
|
f"Unknown image_mode '{mode}'. "
|
||||||
|
f"Supported: {', '.join(_IMAGE_MODE_PRESETS)}"
|
||||||
|
)
|
||||||
|
_MULTI_IMAGE_ALLOWED = ("tiny", "small", "base")
|
||||||
|
base_size, image_size, crop_mode = preset
|
||||||
|
if num_images > 1 and key not in _MULTI_IMAGE_ALLOWED:
|
||||||
|
raise ValueError(
|
||||||
|
f"image_mode='{mode}' is not supported with multiple images "
|
||||||
|
f"(got {num_images} images). "
|
||||||
|
f"Please use one of: {list(_MULTI_IMAGE_ALLOWED)}"
|
||||||
|
)
|
||||||
|
return dict(zip(("base_size", "image_size", "crop_mode"), preset))
|
||||||
|
|
||||||
|
|
||||||
|
class UnlimitedOCRProcessor(BaseMultimodalProcessor):
|
||||||
|
"""Multimodal processor for UNLIMITED-OCR model."""
|
||||||
|
|
||||||
|
models = [UnlimitedOCRForCausalLM]
|
||||||
|
gpu_image_decode = False
|
||||||
|
|
||||||
|
def __init__(self, hf_config, server_args, _processor, *args, **kwargs):
|
||||||
|
"""Initialize UnlimitedOCRProcessor."""
|
||||||
|
super().__init__(hf_config, server_args, _processor, *args, **kwargs)
|
||||||
|
self.mm_tokens = MultimodalSpecialTokens(
|
||||||
|
image_token="<image>", image_token_id=self._processor.image_token_id
|
||||||
|
).build(_processor)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _mix_config_into_hash(mm_items, processor_kwargs):
|
||||||
|
"""Mix images_config into mm_item hashes so that different configs
|
||||||
|
produce different pad_values, avoiding radix/embedding cache collisions."""
|
||||||
|
from sglang.srt.managers.mm_utils import hash_feature
|
||||||
|
|
||||||
|
config_bytes = str(sorted(processor_kwargs.items())).encode()
|
||||||
|
for item in mm_items:
|
||||||
|
if item.feature is not None:
|
||||||
|
base_hash = hash_feature(item.feature)
|
||||||
|
elif item.precomputed_embeddings is not None:
|
||||||
|
base_hash = hash_feature(item.precomputed_embeddings)
|
||||||
|
else:
|
||||||
|
continue
|
||||||
|
combined = hashlib.sha256(
|
||||||
|
base_hash.to_bytes(8, byteorder="big") + config_bytes
|
||||||
|
).digest()[:8]
|
||||||
|
item.hash = int.from_bytes(combined, byteorder="big", signed=False)
|
||||||
|
|
||||||
|
async def process_mm_data_async(
|
||||||
|
self, image_data: List[Union[str, bytes]], input_text, *args, **kwargs
|
||||||
|
):
|
||||||
|
"""Process multimodal data asynchronously."""
|
||||||
|
request_obj = kwargs.get("request_obj")
|
||||||
|
images_config = (
|
||||||
|
getattr(request_obj, "images_config", None) if request_obj else None
|
||||||
|
)
|
||||||
|
processor_kwargs = _resolve_mode(images_config, num_images=len(image_data))
|
||||||
|
|
||||||
|
prefix = images_config.get("prefix", "") if images_config else ""
|
||||||
|
|
||||||
|
base_output = await self.load_mm_data(
|
||||||
|
prompt=input_text,
|
||||||
|
multimodal_tokens=self.mm_tokens,
|
||||||
|
image_data=image_data,
|
||||||
|
)
|
||||||
|
mm_items, input_ids, _ = self.process_and_combine_mm_data(
|
||||||
|
base_output, self.mm_tokens, **processor_kwargs
|
||||||
|
)
|
||||||
|
|
||||||
|
if prefix:
|
||||||
|
prefix_ids = self._tokenizer.encode(prefix, add_special_tokens=False)
|
||||||
|
input_ids = torch.cat(
|
||||||
|
[input_ids, torch.tensor(prefix_ids, dtype=input_ids.dtype)]
|
||||||
|
)
|
||||||
|
|
||||||
|
self._mix_config_into_hash(mm_items, processor_kwargs)
|
||||||
|
|
||||||
|
return MultimodalProcessorOutput(
|
||||||
|
mm_items=mm_items,
|
||||||
|
input_ids=input_ids.tolist(),
|
||||||
|
im_token_id=self.mm_tokens.image_token_id,
|
||||||
|
)
|
||||||
@@ -67,6 +67,7 @@ class SeparatorStyle(IntEnum):
|
|||||||
GEMMA3 = auto()
|
GEMMA3 = auto()
|
||||||
MPT = auto()
|
MPT = auto()
|
||||||
PADDLE_OCR = auto()
|
PADDLE_OCR = auto()
|
||||||
|
UNLIMITED_OCR = auto()
|
||||||
|
|
||||||
|
|
||||||
@dataclasses.dataclass
|
@dataclasses.dataclass
|
||||||
@@ -398,6 +399,18 @@ class Conversation:
|
|||||||
else:
|
else:
|
||||||
ret += role + ": " # must be end with a space
|
ret += role + ": " # must be end with a space
|
||||||
return ret
|
return ret
|
||||||
|
elif self.sep_style == SeparatorStyle.UNLIMITED_OCR:
|
||||||
|
seps = [self.sep, self.sep2]
|
||||||
|
if system_prompt == "" or system_prompt is None:
|
||||||
|
ret = ""
|
||||||
|
else:
|
||||||
|
ret = system_prompt + seps[0]
|
||||||
|
for i, (role, message) in enumerate(self.messages):
|
||||||
|
if message:
|
||||||
|
ret += role + message + seps[i % 2]
|
||||||
|
else:
|
||||||
|
ret += role
|
||||||
|
return ret
|
||||||
else:
|
else:
|
||||||
raise ValueError(f"Invalid style: {self.sep_style}")
|
raise ValueError(f"Invalid style: {self.sep_style}")
|
||||||
|
|
||||||
@@ -643,7 +656,7 @@ def generate_chat_conv(
|
|||||||
conv.modalities.append(content.modalities)
|
conv.modalities.append(content.modalities)
|
||||||
image_token = (
|
image_token = (
|
||||||
conv.image_token + "\n"
|
conv.image_token + "\n"
|
||||||
if conv.name not in ("qwen2-vl", "moss-vl")
|
if conv.name not in ("qwen2-vl", "moss-vl", "unlimited-ocr")
|
||||||
else conv.image_token
|
else conv.image_token
|
||||||
)
|
)
|
||||||
add_token_as_needed: bool = (
|
add_token_as_needed: bool = (
|
||||||
@@ -656,7 +669,7 @@ def generate_chat_conv(
|
|||||||
video_token = conv.video_token
|
video_token = conv.video_token
|
||||||
for content in message.content:
|
for content in message.content:
|
||||||
if content.type == "text":
|
if content.type == "text":
|
||||||
if num_image_url > 16:
|
if num_image_url > 16 and conv.name not in ("unlimited-ocr",):
|
||||||
real_content += "\n" # for video
|
real_content += "\n" # for video
|
||||||
real_content += content.text
|
real_content += content.text
|
||||||
elif content.type == "image_url":
|
elif content.type == "image_url":
|
||||||
@@ -887,6 +900,22 @@ register_conv_template(
|
|||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
|
register_conv_template(
|
||||||
|
Conversation(
|
||||||
|
name="unlimited-ocr",
|
||||||
|
system_template="{system_message}",
|
||||||
|
system_message="",
|
||||||
|
roles=("", ""),
|
||||||
|
messages=(),
|
||||||
|
offset=0,
|
||||||
|
sep_style=SeparatorStyle.UNLIMITED_OCR,
|
||||||
|
sep="",
|
||||||
|
sep2="",
|
||||||
|
image_token="<image>",
|
||||||
|
image_token_at_prefix=True,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
register_conv_template(
|
register_conv_template(
|
||||||
Conversation(
|
Conversation(
|
||||||
name="paddle-ocr",
|
name="paddle-ocr",
|
||||||
@@ -1076,6 +1105,7 @@ MODEL_TYPE_TO_TEMPLATE = {
|
|||||||
"minicpmo": "minicpmo",
|
"minicpmo": "minicpmo",
|
||||||
"moss_vl": "moss-vl",
|
"moss_vl": "moss-vl",
|
||||||
"deepseek-ocr": "deepseek-ocr",
|
"deepseek-ocr": "deepseek-ocr",
|
||||||
|
"unlimited-ocr": "unlimited-ocr",
|
||||||
"paddleocr_vl": "paddle-ocr",
|
"paddleocr_vl": "paddle-ocr",
|
||||||
"whisper": "whisper",
|
"whisper": "whisper",
|
||||||
}
|
}
|
||||||
@@ -1182,6 +1212,17 @@ def match_deepseek_ocr(model_path: str):
|
|||||||
return MODEL_TYPE_TO_TEMPLATE.get(model_type)
|
return MODEL_TYPE_TO_TEMPLATE.get(model_type)
|
||||||
|
|
||||||
|
|
||||||
|
@register_conv_template_matching_function
|
||||||
|
def match_unlimited_ocr(model_path: str):
|
||||||
|
"""Match unlimited-ocr model by path or model type."""
|
||||||
|
if "unlimited" in model_path.lower():
|
||||||
|
return "unlimited-ocr"
|
||||||
|
model_type = get_model_type(model_path)
|
||||||
|
if model_type == "unlimited-ocr":
|
||||||
|
return "unlimited-ocr"
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
@register_conv_template_matching_function
|
@register_conv_template_matching_function
|
||||||
def match_paddle_ocr(model_path: str):
|
def match_paddle_ocr(model_path: str):
|
||||||
if "paddleocr" in model_path.lower():
|
if "paddleocr" in model_path.lower():
|
||||||
|
|||||||
@@ -0,0 +1,48 @@
|
|||||||
|
"""Unit tests for all-SWA ChunkCache release semantics."""
|
||||||
|
|
||||||
|
import unittest
|
||||||
|
from types import SimpleNamespace
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from sglang.srt.mem_cache.chunk_cache import PureSWAChunkCache
|
||||||
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
|
from sglang.test.test_utils import CustomTestCase
|
||||||
|
|
||||||
|
register_cpu_ci(est_time=3, suite="base-a-test-cpu")
|
||||||
|
|
||||||
|
|
||||||
|
class _FakeAllocator:
|
||||||
|
def __init__(self):
|
||||||
|
self.freed = []
|
||||||
|
|
||||||
|
def free(self, indices):
|
||||||
|
self.freed.append(indices.detach().cpu().clone())
|
||||||
|
|
||||||
|
|
||||||
|
class _FakeReq:
|
||||||
|
req_pool_idx = 0
|
||||||
|
swa_evict_floor = 3
|
||||||
|
swa_evicted_seqlen = 6
|
||||||
|
|
||||||
|
def pop_committed_kv_cache(self):
|
||||||
|
return 8
|
||||||
|
|
||||||
|
|
||||||
|
class TestPureSWAChunkCache(CustomTestCase):
|
||||||
|
def test_finished_req_skips_already_evicted_swa_range(self):
|
||||||
|
cache = PureSWAChunkCache.__new__(PureSWAChunkCache)
|
||||||
|
cache.req_to_token_pool = SimpleNamespace(
|
||||||
|
req_to_token=torch.arange(10, dtype=torch.int64).unsqueeze(0)
|
||||||
|
)
|
||||||
|
cache.token_to_kv_pool_allocator = _FakeAllocator()
|
||||||
|
|
||||||
|
cache.cache_finished_req(_FakeReq())
|
||||||
|
|
||||||
|
self.assertEqual(len(cache.token_to_kv_pool_allocator.freed), 1)
|
||||||
|
freed = cache.token_to_kv_pool_allocator.freed[0]
|
||||||
|
self.assertTrue(torch.equal(freed, torch.tensor([0, 1, 2, 6, 7])))
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
@@ -29,6 +29,7 @@ def _make_ctx(
|
|||||||
enable_hierarchical_cache=False,
|
enable_hierarchical_cache=False,
|
||||||
disable_radix_cache=False,
|
disable_radix_cache=False,
|
||||||
effective_chunked_prefill_size=None,
|
effective_chunked_prefill_size=None,
|
||||||
|
full_tokens_per_layer=None,
|
||||||
):
|
):
|
||||||
server_args = MagicMock()
|
server_args = MagicMock()
|
||||||
server_args.radix_cache_backend = backend
|
server_args.radix_cache_backend = backend
|
||||||
@@ -47,6 +48,7 @@ def _make_ctx(
|
|||||||
tp_size=1,
|
tp_size=1,
|
||||||
tp_rank=0,
|
tp_rank=0,
|
||||||
tp_group=MagicMock(),
|
tp_group=MagicMock(),
|
||||||
|
full_tokens_per_layer=full_tokens_per_layer,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@@ -172,6 +174,21 @@ class TestDefaultRadixCacheFactory(CustomTestCase):
|
|||||||
SWAChunkCache.assert_called_once_with(ctx.params)
|
SWAChunkCache.assert_called_once_with(ctx.params)
|
||||||
self.assertIs(result, SWAChunkCache.return_value)
|
self.assertIs(result, SWAChunkCache.return_value)
|
||||||
|
|
||||||
|
def test_pure_swa_chunk_cache_when_chunked_prefill_disable_and_all_swa(self):
|
||||||
|
ctx = _make_ctx(
|
||||||
|
effective_chunked_prefill_size=512,
|
||||||
|
disable_radix_cache=True,
|
||||||
|
is_hybrid_swa=True,
|
||||||
|
full_tokens_per_layer=0,
|
||||||
|
)
|
||||||
|
with patch(
|
||||||
|
"sglang.srt.mem_cache.chunk_cache.PureSWAChunkCache"
|
||||||
|
) as PureSWAChunkCache:
|
||||||
|
PureSWAChunkCache.return_value = MagicMock()
|
||||||
|
result = default_radix_cache_factory(ctx)
|
||||||
|
PureSWAChunkCache.assert_called_once_with(ctx.params)
|
||||||
|
self.assertIs(result, PureSWAChunkCache.return_value)
|
||||||
|
|
||||||
def test_cpp_radix_cache_when_env_flag_set(self):
|
def test_cpp_radix_cache_when_env_flag_set(self):
|
||||||
ctx = _make_ctx()
|
ctx = _make_ctx()
|
||||||
# `radix_cache_cpp` requires ninja + C++ extension to import, so
|
# `radix_cache_cpp` requires ninja + C++ extension to import, so
|
||||||
@@ -281,6 +298,16 @@ class TestDefaultRadixCacheFactory(CustomTestCase):
|
|||||||
SWA.assert_called_once_with(params=ctx.params)
|
SWA.assert_called_once_with(params=ctx.params)
|
||||||
self.assertIs(result, SWA.return_value)
|
self.assertIs(result, SWA.return_value)
|
||||||
|
|
||||||
|
def test_pure_swa_radix_cache_when_all_swa(self):
|
||||||
|
ctx = _make_ctx(is_hybrid_swa=True, full_tokens_per_layer=0)
|
||||||
|
with patch(
|
||||||
|
"sglang.srt.mem_cache.pure_swa_radix_cache.PureSWARadixCache"
|
||||||
|
) as PureSWA:
|
||||||
|
PureSWA.return_value = MagicMock()
|
||||||
|
result = default_radix_cache_factory(ctx)
|
||||||
|
PureSWA.assert_called_once_with(params=ctx.params)
|
||||||
|
self.assertIs(result, PureSWA.return_value)
|
||||||
|
|
||||||
def test_mamba_radix_cache_when_hybrid_ssm(self):
|
def test_mamba_radix_cache_when_hybrid_ssm(self):
|
||||||
ctx = _make_ctx(is_hybrid_ssm=True)
|
ctx = _make_ctx(is_hybrid_ssm=True)
|
||||||
with patch("sglang.srt.mem_cache.mamba_radix_cache.MambaRadixCache") as Mamba:
|
with patch("sglang.srt.mem_cache.mamba_radix_cache.MambaRadixCache") as Mamba:
|
||||||
|
|||||||
Reference in New Issue
Block a user