diff --git a/benchmark/kernels/fused_moe_triton/common_utils.py b/benchmark/kernels/fused_moe_triton/common_utils.py index cb7422d6a..ea669a1fa 100644 --- a/benchmark/kernels/fused_moe_triton/common_utils.py +++ b/benchmark/kernels/fused_moe_triton/common_utils.py @@ -166,6 +166,10 @@ def get_model_config( 0 if disable_shared_experts_fusion or topk_ids_dir is None else 1 ) 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: # Default: Mixtral E = config.num_local_experts // ep_size diff --git a/docs_new/cards/logos/baidu.svg b/docs_new/cards/logos/baidu.svg new file mode 100644 index 000000000..8520b583d --- /dev/null +++ b/docs_new/cards/logos/baidu.svg @@ -0,0 +1,6 @@ + + + + Baidu + + diff --git a/docs_new/cookbook/autoregressive/Baidu/Unlimited-OCR.mdx b/docs_new/cookbook/autoregressive/Baidu/Unlimited-OCR.mdx new file mode 100644 index 000000000..9e7b79780 --- /dev/null +++ b/docs_new/cookbook/autoregressive/Baidu/Unlimited-OCR.mdx @@ -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 + + + + + +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. + + + + + +```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. + + + + + +```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. + + + + + + + +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"; + + + +## Playground + +Use the Playground to adjust tensor parallelism on top of the selected deployment cell. + +import { Playground } from "/src/snippets/_playground.jsx"; + + + +## 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 + + + +```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) +``` + + + +### 3.2 Choosing an image mode + +Use lower modes to reduce prefill cost for simple images, and use `gundam` for high-detail document parsing. + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + +
ModeUseMultiple images
tinyLowest prefill cost.Yes
smallLightweight OCR requests.Yes
baseBalanced quality and cost.Yes
largeHigher resolution single-image OCR.No
gundamDefault high-detail document parsing mode.No
diff --git a/docs_new/cookbook/autoregressive/intro.mdx b/docs_new/cookbook/autoregressive/intro.mdx index c3d40f375..4cf117fca 100644 --- a/docs_new/cookbook/autoregressive/intro.mdx +++ b/docs_new/cookbook/autoregressive/intro.mdx @@ -67,6 +67,12 @@ metatags: href="/cookbook/autoregressive/NVIDIA/Nemotron3-Ultra" img="/cards/logos/nvidia.png" /> + ", + }, + 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}}", + ], + }, + ], +}; diff --git a/python/sglang/srt/configs/__init__.py b/python/sglang/srt/configs/__init__.py index 01c743828..814992037 100644 --- a/python/sglang/srt/configs/__init__.py +++ b/python/sglang/srt/configs/__init__.py @@ -39,6 +39,7 @@ from sglang.srt.configs.step3_vl import ( ) from sglang.srt.configs.step3p5 import Step3p5Config from sglang.srt.configs.step3p7 import Step3p7Config +from sglang.srt.configs.unlimited_ocr import UnlimitedVLConfig from sglang.srt.configs.zaya import ZayaConfig __all__ = [ @@ -81,5 +82,6 @@ __all__ = [ "Step3p5Config", "Step3p7Config", "Qwen3ASRConfig", + "UnlimitedVLConfig", "ZayaConfig", ] diff --git a/python/sglang/srt/configs/model_config.py b/python/sglang/srt/configs/model_config.py index d92874d2d..d8b3bdcba 100644 --- a/python/sglang/srt/configs/model_config.py +++ b/python/sglang/srt/configs/model_config.py @@ -1707,6 +1707,7 @@ multimodal_model_archs = [ "NVILAForConditionalGeneration", "NVILALiteForConditionalGeneration", "DeepseekOCRForCausalLM", + "UnlimitedOCRForCausalLM", "JetVLMForConditionalGeneration", "PaddleOCRVLForConditionalGeneration", "MiDashengLMModel", @@ -1863,6 +1864,7 @@ def is_hybrid_swa_model( "Gemma4ForConditionalGeneration", "Gemma4UnifiedForConditionalGeneration", "LagunaForCausalLM", + "UnlimitedOCRForCausalLM", } 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. @@ -1949,6 +1951,9 @@ def get_hybrid_layer_ids( full_attention_layer_ids = [ 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: # Generic fallback for custom hybrid SWA models that opt in via # hf_text_config.is_hybrid_swa and expose a hybrid_layer_pattern diff --git a/python/sglang/srt/configs/unlimited_ocr.py b/python/sglang/srt/configs/unlimited_ocr.py new file mode 100644 index 000000000..66be0f24e --- /dev/null +++ b/python/sglang/srt/configs/unlimited_ocr.py @@ -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 = "", + 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 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 diff --git a/python/sglang/srt/entrypoints/openai/protocol.py b/python/sglang/srt/entrypoints/openai/protocol.py index 8884132f9..ace007a0c 100644 --- a/python/sglang/srt/entrypoints/openai/protocol.py +++ b/python/sglang/srt/entrypoints/openai/protocol.py @@ -361,6 +361,8 @@ class CompletionRequest(BaseModel): custom_params: Optional[Dict] = None custom_logit_processor: Optional[str] = None + images_config: Optional[Dict] = None + # For PD disaggregation bootstrap_host: Optional[Union[List[str], str]] = None bootstrap_port: Optional[Union[List[Optional[int]], int]] = None @@ -736,6 +738,8 @@ class ChatCompletionRequest(BaseModel): min_dynamic_patch: Optional[int] = None use_audio_in_video: bool = False + images_config: Optional[Dict] = None + # Custom logit processor for advanced sampling control custom_logit_processor: Optional[Union[List[Optional[str]], str]] = None custom_params: Optional[Dict] = None diff --git a/python/sglang/srt/entrypoints/openai/serving_chat.py b/python/sglang/srt/entrypoints/openai/serving_chat.py index 3a400bcaa..bbcecfc15 100644 --- a/python/sglang/srt/entrypoints/openai/serving_chat.py +++ b/python/sglang/srt/entrypoints/openai/serving_chat.py @@ -612,6 +612,7 @@ class OpenAIServingChat(OpenAIServingBase): routing_key=self.extract_routing_key(raw_request), custom_labels=custom_labels, custom_logit_processor=request.custom_logit_processor, + images_config=getattr(request, "images_config", None), image_max_dynamic_patch=img_max_dynamic_patch, video_max_dynamic_patch=vid_max_dynamic_patch, max_dynamic_patch=getattr(request, "max_dynamic_patch", None), diff --git a/python/sglang/srt/entrypoints/openai/serving_completions.py b/python/sglang/srt/entrypoints/openai/serving_completions.py index 0a15aaf4d..72d5f0d9c 100644 --- a/python/sglang/srt/entrypoints/openai/serving_completions.py +++ b/python/sglang/srt/entrypoints/openai/serving_completions.py @@ -130,6 +130,7 @@ class OpenAIServingCompletion(OpenAIServingBase): routing_key=self.extract_routing_key(raw_request), custom_labels=custom_labels, custom_logit_processor=request.custom_logit_processor, + images_config=getattr(request, "images_config", None), ) return adapted_request, request diff --git a/python/sglang/srt/layers/attention/flashattention_backend.py b/python/sglang/srt/layers/attention/flashattention_backend.py index 0a50f0d60..8277a8d27 100644 --- a/python/sglang/srt/layers/attention/flashattention_backend.py +++ b/python/sglang/srt/layers/attention/flashattention_backend.py @@ -5,6 +5,8 @@ from typing import TYPE_CHECKING, Optional import numpy as np import torch +import triton +import triton.language as tl from sglang.srt.configs.model_config import AttentionArch 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 +@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 class FlashAttentionMetadata: """Metadata to be init once in the model forward pass, @@ -68,6 +159,9 @@ class FlashAttentionMetadata: page_table: torch.Tensor = None # Page table for Sliding Window Attention 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) swa_out_cache_loc: torch.Tensor = None # 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.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 self.fa_impl_ver = fa_impl_ver if self.fa_impl_ver == 3: @@ -483,6 +589,25 @@ class FlashAttentionBackend(AttentionBackend): metadata.page_table = self.req_to_token_pool.req_to_token[ 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 # prepare_varlen_num_blocks kernel calls metadata.scheduler_metadata = self._compute_scheduler_metadata( @@ -668,6 +793,14 @@ class FlashAttentionBackend(AttentionBackend): if forward_batch.forward_mode == ForwardMode.EXTEND: 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 lengths (e.g. MossVL with different image sizes per request). if forward_batch.encoder_lens is not None: @@ -1453,6 +1586,14 @@ class FlashAttentionBackend(AttentionBackend): ) cache_seqlens = metadata.cache_seqlens_int32 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( -1, layer.tp_q_head_num, layer.head_dim ) @@ -1465,6 +1606,7 @@ class FlashAttentionBackend(AttentionBackend): metadata.scheduler_metadata is not None and not is_swa_layer and not use_cascade_attn + and not pa_swa_active ): sched_meta = metadata.scheduler_metadata result = flash_attn_with_kvcache( @@ -1995,6 +2137,9 @@ class FlashAttentionBackend(AttentionBackend): metadata.page_table = self.decode_cuda_graph_metadata["page_table"][ :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: metadata.swa_page_table = self.decode_cuda_graph_metadata[ "swa_page_table" @@ -2253,20 +2398,43 @@ class FlashAttentionBackend(AttentionBackend): metadata.page_table.shape[1], "FA3 decode page_table", ) - 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, - ) + if self.is_prefill_aware_swa: + pa_max_len = min( + self._pa_swa_max_prefill_len + self.sliding_window_size, + max_len, + ) + if pa_max_len > 0: + _build_pa_page_table( + self.req_to_token, + req_pool_indices, + seq_lens, + self._pa_swa_prefill_lens, + self.sliding_window_size, + 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( metadata, diff --git a/python/sglang/srt/managers/io_struct.py b/python/sglang/srt/managers/io_struct.py index f8fdad62c..377206810 100644 --- a/python/sglang/srt/managers/io_struct.py +++ b/python/sglang/srt/managers/io_struct.py @@ -292,6 +292,9 @@ class GenerateReqInput: image_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. # Batch-level: List[List[int]] (one per request). After __getitem__: List[int]. multi_item_delimiter_indices: Optional[Union[List[List[int]], List[int]]] = None diff --git a/python/sglang/srt/managers/schedule_batch.py b/python/sglang/srt/managers/schedule_batch.py index 65348246c..021490c35 100755 --- a/python/sglang/srt/managers/schedule_batch.py +++ b/python/sglang/srt/managers/schedule_batch.py @@ -740,7 +740,7 @@ class Req(ReqDllmMixin): self.kv_committed_freed = False self.kv_overallocated_freed = False - # for corss-endoder model + # for cross-encoder model self.token_type_ids = token_type_ids # 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. # - Chunk cache: KV in range [0, swa_evicted_seqlen) is freed manually in `ScheduleBatch.maybe_evict_swa`. 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 self.extend_batch_idx = 0 diff --git a/python/sglang/srt/managers/schedule_policy.py b/python/sglang/srt/managers/schedule_policy.py index 23cebfd32..b37a78329 100644 --- a/python/sglang/srt/managers/schedule_policy.py +++ b/python/sglang/srt/managers/schedule_policy.py @@ -41,7 +41,10 @@ from sglang.srt.managers.schedule_batch import Req, ScheduleBatch from sglang.srt.mem_cache.allocator.hisparse import ( 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 ( BasePrefixCache, InitLoadBackParams, @@ -483,6 +486,9 @@ class PrefillAdder: self.token_to_kv_pool_allocator, (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.rem_swa_token_offset = 0 @@ -517,7 +523,12 @@ class PrefillAdder: @property 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 = ( self.token_to_kv_pool_allocator.full_available_size() + self.tree_cache.full_evictable_size() @@ -544,7 +555,12 @@ class PrefillAdder: @property 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 = ( self.token_to_kv_pool_allocator.full_available_size() + self.tree_cache.full_evictable_size() diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index 21c007ea3..09231e1b8 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -2972,6 +2972,10 @@ class Scheduler( 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. new_batch.prefill_stats = PrefillStats.from_adder( adder, diff --git a/python/sglang/srt/managers/scheduler_components/invariant_checker.py b/python/sglang/srt/managers/scheduler_components/invariant_checker.py index 9fe441668..9e625d404 100644 --- a/python/sglang/srt/managers/scheduler_components/invariant_checker.py +++ b/python/sglang/srt/managers/scheduler_components/invariant_checker.py @@ -82,6 +82,8 @@ class SchedulerInvariantChecker: return leak, msg 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: protected = self.tree_cache.full_protected_size() session_held = self.pool_stats_observer.session_held_full_tokens() diff --git a/python/sglang/srt/managers/scheduler_components/pool_stats_observer.py b/python/sglang/srt/managers/scheduler_components/pool_stats_observer.py index 7fcbda113..087e266fc 100644 --- a/python/sglang/srt/managers/scheduler_components/pool_stats_observer.py +++ b/python/sglang/srt/managers/scheduler_components/pool_stats_observer.py @@ -300,7 +300,12 @@ class SchedulerPoolStatsObserver: if self.enable_hisparse: full_num_used = max(0, full_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 return PoolStats( diff --git a/python/sglang/srt/mem_cache/allocator/swa.py b/python/sglang/srt/mem_cache/allocator/swa.py index 31745d645..8e27b731a 100644 --- a/python/sglang/srt/mem_cache/allocator/swa.py +++ b/python/sglang/srt/mem_cache/allocator/swa.py @@ -385,3 +385,120 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator): return self._kvcache.load_cpu_copy( 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 = [] diff --git a/python/sglang/srt/mem_cache/chunk_cache.py b/python/sglang/srt/mem_cache/chunk_cache.py index b2bf0445b..7bc00000a 100644 --- a/python/sglang/srt/mem_cache/chunk_cache.py +++ b/python/sglang/srt/mem_cache/chunk_cache.py @@ -135,3 +135,35 @@ class SWAChunkCache(ChunkCache): def evict(self, params: EvictParams) -> 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) diff --git a/python/sglang/srt/mem_cache/common.py b/python/sglang/srt/mem_cache/common.py index 0a8c83aa6..e50fcfc14 100644 --- a/python/sglang/srt/mem_cache/common.py +++ b/python/sglang/srt/mem_cache/common.py @@ -81,7 +81,10 @@ def free_swa_out_of_window_slots( assert ( req.cache_protected_len % page_size == 0 ), "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 # radix tree insert boundary (page_floor(seq_len)). This keeps at least one diff --git a/python/sglang/srt/mem_cache/kv_cache_builder.py b/python/sglang/srt/mem_cache/kv_cache_builder.py index 5f985de69..98eec2a86 100644 --- a/python/sglang/srt/mem_cache/kv_cache_builder.py +++ b/python/sglang/srt/mem_cache/kv_cache_builder.py @@ -232,6 +232,7 @@ def build_kv_cache( server_args=server_args, params=params, is_hybrid_swa=is_hybrid_swa, + full_tokens_per_layer=full_tokens_per_layer, is_hybrid_ssm=is_hybrid_ssm, enable_hierarchical_cache=enable_hierarchical_cache, disable_radix_cache=disable_radix_cache, diff --git a/python/sglang/srt/mem_cache/pure_swa_radix_cache.py b/python/sglang/srt/mem_cache/pure_swa_radix_cache.py new file mode 100644 index 000000000..309060fe4 --- /dev/null +++ b/python/sglang/srt/mem_cache/pure_swa_radix_cache.py @@ -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" + ) diff --git a/python/sglang/srt/mem_cache/registry.py b/python/sglang/srt/mem_cache/registry.py index 9e6e7fd7b..5b8b9e9de 100644 --- a/python/sglang/srt/mem_cache/registry.py +++ b/python/sglang/srt/mem_cache/registry.py @@ -42,6 +42,7 @@ class TreeCacheBuildContext: tp_size: int tp_rank: int tp_group: Any + full_tokens_per_layer: Optional[int] = None 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 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 return SWAChunkCache(params) @@ -112,6 +117,10 @@ def default_radix_cache_factory(ctx: TreeCacheBuildContext) -> BasePrefixCache: return cache 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 return SWARadixCache(params=params) diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index b03125d29..1afa7cfe6 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -1501,6 +1501,11 @@ class ModelRunner(ModelRunnerKVCacheMixin): 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 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): """Return the max token pool size considering hybrid swa settings.""" 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: return self.max_total_num_tokens diff --git a/python/sglang/srt/model_executor/model_runner_kv_cache_mixin.py b/python/sglang/srt/model_executor/model_runner_kv_cache_mixin.py index e98970256..c270db628 100644 --- a/python/sglang/srt/model_executor/model_runner_kv_cache_mixin.py +++ b/python/sglang/srt/model_executor/model_runner_kv_cache_mixin.py @@ -24,7 +24,10 @@ from sglang.srt.mem_cache.allocator.hisparse import ( DeepSeekV4HiSparseTokenToKVPoolAllocator, 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.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool from sglang.srt.mem_cache.hisparse_memory_pool import HiSparseDSATokenToKVPool @@ -861,7 +864,16 @@ class ModelRunnerKVCacheMixin: need_sort=need_sort, ) 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.full_max_total_num_tokens, self.swa_max_total_num_tokens, diff --git a/python/sglang/srt/models/unlimited_ocr.py b/python/sglang/srt/models/unlimited_ocr.py new file mode 100644 index 000000000..8c0cafd44 --- /dev/null +++ b/python/sglang/srt/models/unlimited_ocr.py @@ -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] diff --git a/python/sglang/srt/multimodal/processors/unlimited_ocr.py b/python/sglang/srt/multimodal/processors/unlimited_ocr.py new file mode 100644 index 000000000..67499a3c0 --- /dev/null +++ b/python/sglang/srt/multimodal/processors/unlimited_ocr.py @@ -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_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, + ) diff --git a/python/sglang/srt/parser/conversation.py b/python/sglang/srt/parser/conversation.py index 3197cf15d..e947737ab 100644 --- a/python/sglang/srt/parser/conversation.py +++ b/python/sglang/srt/parser/conversation.py @@ -67,6 +67,7 @@ class SeparatorStyle(IntEnum): GEMMA3 = auto() MPT = auto() PADDLE_OCR = auto() + UNLIMITED_OCR = auto() @dataclasses.dataclass @@ -398,6 +399,18 @@ class Conversation: else: ret += role + ": " # must be end with a space 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: raise ValueError(f"Invalid style: {self.sep_style}") @@ -643,7 +656,7 @@ def generate_chat_conv( conv.modalities.append(content.modalities) image_token = ( 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 ) add_token_as_needed: bool = ( @@ -656,7 +669,7 @@ def generate_chat_conv( video_token = conv.video_token for content in message.content: 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 += content.text 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_token_at_prefix=True, + ) +) + register_conv_template( Conversation( name="paddle-ocr", @@ -1076,6 +1105,7 @@ MODEL_TYPE_TO_TEMPLATE = { "minicpmo": "minicpmo", "moss_vl": "moss-vl", "deepseek-ocr": "deepseek-ocr", + "unlimited-ocr": "unlimited-ocr", "paddleocr_vl": "paddle-ocr", "whisper": "whisper", } @@ -1182,6 +1212,17 @@ def match_deepseek_ocr(model_path: str): 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 def match_paddle_ocr(model_path: str): if "paddleocr" in model_path.lower(): diff --git a/test/registered/unit/mem_cache/test_pure_swa_chunk_cache.py b/test/registered/unit/mem_cache/test_pure_swa_chunk_cache.py new file mode 100644 index 000000000..7c8d90f3e --- /dev/null +++ b/test/registered/unit/mem_cache/test_pure_swa_chunk_cache.py @@ -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() diff --git a/test/registered/unit/mem_cache/test_registry.py b/test/registered/unit/mem_cache/test_registry.py index 61b22390c..84b37820f 100644 --- a/test/registered/unit/mem_cache/test_registry.py +++ b/test/registered/unit/mem_cache/test_registry.py @@ -29,6 +29,7 @@ def _make_ctx( enable_hierarchical_cache=False, disable_radix_cache=False, effective_chunked_prefill_size=None, + full_tokens_per_layer=None, ): server_args = MagicMock() server_args.radix_cache_backend = backend @@ -47,6 +48,7 @@ def _make_ctx( tp_size=1, tp_rank=0, 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) 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): ctx = _make_ctx() # `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) 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): ctx = _make_ctx(is_hybrid_ssm=True) with patch("sglang.srt.mem_cache.mamba_radix_cache.MambaRadixCache") as Mamba: