diff --git a/docs_new/docs/supported-models/multimodal_language_models.mdx b/docs_new/docs/supported-models/multimodal_language_models.mdx index b2131cc94..d6155af24 100644 --- a/docs_new/docs/supported-models/multimodal_language_models.mdx +++ b/docs_new/docs/supported-models/multimodal_language_models.mdx @@ -201,6 +201,12 @@ in the GitHub search bar. Liquid AI's vision-language model combining a SigLIP2 NaFlex vision encoder (variable resolution, native aspect ratio) with the LFM2 hybrid gated short conv + GQA language model. Supports multi-image inputs. + + LocateAnything (3B) + nvidia/LocateAnything-3B + NVIDIA's visual grounding/detection model (MoonViT vision encoder + Qwen2 backbone) that emits <ref>label</ref><box>...</box> outputs with coordinates normalized to [0, 1000]. Covers object detection, phrase grounding, scene-text detection, GUI grounding, and pointing. + Use --trust-remote-code. Set skip_special_tokens=false so the <ref>/<box> grounding tokens survive in the output. Constrained <box> decoding is opt-in and client-side: start the server with --enable-custom-logit-processor, then pass custom_logit_processor (a top-level request field) and custom_params (inside sampling_params) together — use LocateAnythingBoxGrammarLogitProcessor.build_sampling_params(config) to build both from the config token ids. + diff --git a/python/sglang/srt/configs/__init__.py b/python/sglang/srt/configs/__init__.py index 814992037..aa05b73be 100644 --- a/python/sglang/srt/configs/__init__.py +++ b/python/sglang/srt/configs/__init__.py @@ -21,6 +21,7 @@ from sglang.srt.configs.laguna import LagunaConfig from sglang.srt.configs.lfm2 import Lfm2Config from sglang.srt.configs.lfm2_moe import Lfm2MoeConfig from sglang.srt.configs.lfm2_vl import Lfm2VlConfig +from sglang.srt.configs.locate_anything import LocateAnythingConfig from sglang.srt.configs.longcat_flash import LongcatFlashConfig from sglang.srt.configs.minicpmv4_6 import MiniCPMV4_6Config, MiniCPMV4_6VisionConfig from sglang.srt.configs.nano_nemotron_vl import ( @@ -71,6 +72,7 @@ __all__ = [ "Lfm2Config", "Lfm2MoeConfig", "Lfm2VlConfig", + "LocateAnythingConfig", "MiniCPMV4_6Config", "MiniCPMV4_6VisionConfig", "NemotronHConfig", diff --git a/python/sglang/srt/configs/locate_anything.py b/python/sglang/srt/configs/locate_anything.py new file mode 100644 index 000000000..1024150d0 --- /dev/null +++ b/python/sglang/srt/configs/locate_anything.py @@ -0,0 +1,63 @@ +# SPDX-License-Identifier: Apache-2.0 +# Adapted from https://huggingface.co/nvidia/LocateAnything-3B/blob/main/configuration_locateanything.py +"""Config for nvidia/LocateAnything-3B. + +LocateAnything is a multimodal grounding/detection model composed of a MoonViT +vision encoder, an InternVL-style ``mlp1`` projector, and a Qwen2 language model +backbone. The config is a composite that wraps a ``MoonViTConfig`` (vision) and a +``Qwen2Config`` (text) plus the special token ids used for the grounding grammar +(````/````/coordinate tokens). +""" + +from typing import Optional, Union + +from transformers.configuration_utils import PretrainedConfig +from transformers.models.qwen2 import Qwen2Config + +from sglang.srt.configs.kimi_vl_moonvit import MoonViTConfig + + +class LocateAnythingConfig(PretrainedConfig): + model_type = "locateanything" + + def __init__( + self, + vision_config: Optional[Union[dict, MoonViTConfig]] = None, + text_config: Optional[Union[dict, Qwen2Config]] = None, + image_token_index: int = 151665, + box_start_token_id: int = 151668, + box_end_token_id: int = 151669, + ref_start_token_id: int = 151672, + ref_end_token_id: int = 151673, + coord_start_token_id: int = 151677, + coord_end_token_id: int = 152677, + none_token_id: int = 4064, + mlp_connector_layers: int = 2, + **kwargs, + ): + if vision_config is None: + vision_config = MoonViTConfig() + elif isinstance(vision_config, dict): + vision_config = MoonViTConfig(**vision_config) + self.vision_config = vision_config + + if text_config is None: + text_config = Qwen2Config() + elif isinstance(text_config, dict): + text_config = Qwen2Config(**text_config) + self.text_config = text_config + + self.image_token_index = image_token_index + self.box_start_token_id = box_start_token_id + self.box_end_token_id = box_end_token_id + # ref_*_token_id and mlp_connector_layers are kept for round-trip + # fidelity with the HF config; the box-grammar processor reads the box / + # coord / none ids, and the projector hardcodes its 2-layer structure. + self.ref_start_token_id = ref_start_token_id + self.ref_end_token_id = ref_end_token_id + self.coord_start_token_id = coord_start_token_id + self.coord_end_token_id = coord_end_token_id + self.none_token_id = none_token_id + self.mlp_connector_layers = mlp_connector_layers + + super().__init__(**kwargs) diff --git a/python/sglang/srt/configs/model_config.py b/python/sglang/srt/configs/model_config.py index d8b3bdcba..a0d0ce61c 100644 --- a/python/sglang/srt/configs/model_config.py +++ b/python/sglang/srt/configs/model_config.py @@ -1693,6 +1693,7 @@ multimodal_model_archs = [ "Qwen3ASRForConditionalGeneration", "Qwen3OmniMoeForConditionalGeneration", "KimiVLForConditionalGeneration", + "LocateAnythingForConditionalGeneration", "InternVLChatModel", "InternS1ForConditionalGeneration", "InternS1ProForConditionalGeneration", diff --git a/python/sglang/srt/managers/mm_utils.py b/python/sglang/srt/managers/mm_utils.py index d0ac4a988..c0eafff90 100644 --- a/python/sglang/srt/managers/mm_utils.py +++ b/python/sglang/srt/managers/mm_utils.py @@ -1326,6 +1326,22 @@ def _get_length(value): return None +def _is_rank2_grid(value): + """True if `value` is a rank-2 grid ([N, dims]) suitable for per-row prod. + + Tensors/arrays must have ndim == 2; nested lists/tuples must have each row + be a sequence. Anything flat (1-D / scalars) is rejected so callers fall + back to a simple split instead of mis-collapsing it with prod(dim=-1). + """ + if isinstance(value, (torch.Tensor, np.ndarray)): + return value.ndim == 2 + if isinstance(value, (list, tuple)): + return len(value) > 0 and all( + isinstance(row, (list, tuple, torch.Tensor, np.ndarray)) for row in value + ) + return False + + def _slice_value(value, start, end): if isinstance(value, torch.Tensor): return value[start:end] @@ -1490,7 +1506,15 @@ def get_new_expanded_mm_items(original_mm_items): num_items = len(item.offsets) if item.is_image(): + # MoonViT-style models (e.g. LocateAnything) carry per-image + # grids under `image_grid_hws` ([h, w]) rather than + # `image_grid_thw` ([t, h, w]); both encode dim-0 patch counts + # via prod over the last axis, so accept either key. (Use an + # explicit None check, not `a or b`: the value is a multi-element + # tensor whose truthiness is ambiguous.) image_grid_thw = item.model_specific_data.get("image_grid_thw") + if image_grid_thw is None: + image_grid_thw = item.model_specific_data.get("image_grid_hws") grid_len = _get_length(image_grid_thw) if image_grid_thw is None or grid_len != num_items: # No grid info — fall back to simple split by feature dim-0 @@ -1498,6 +1522,18 @@ def get_new_expanded_mm_items(original_mm_items): expanded_mm_items.append(item) continue + # The grid must be rank-2 ([N, dims]) so `prod` over the last + # axis yields one patch count per image. A flat 1-D grid (e.g. + # `tensor([h, w])` with num_items==2) would pass the length check + # above but `prod(dim=-1)` collapses it to a scalar and mis-splits. + # The HF processor always emits rank-2, so this only guards the + # degenerate case — fall back to simple split rather than corrupt + # the slice boundaries. + if not _is_rank2_grid(image_grid_thw): + if not _try_simple_split(item, num_items, expanded_mm_items): + expanded_mm_items.append(item) + continue + if isinstance(image_grid_thw, torch.Tensor): patches_per_item = ( torch.prod(image_grid_thw, dim=-1).long().tolist() diff --git a/python/sglang/srt/models/locate_anything.py b/python/sglang/srt/models/locate_anything.py new file mode 100644 index 000000000..565eb7162 --- /dev/null +++ b/python/sglang/srt/models/locate_anything.py @@ -0,0 +1,403 @@ +# SPDX-License-Identifier: Apache-2.0 +# Adapted from https://huggingface.co/nvidia/LocateAnything-3B/blob/main/modeling_locateanything.py +# and from vllm-project/vllm PR #44182. +"""Inference-only LocateAnything-3B model for SGLang. + +LocateAnything-3B is a multimodal grounding/detection model: + + * MoonViT vision encoder (reused unchanged from Kimi-VL) + * An InternVL-style ``mlp1`` projector (LayerNorm applied AFTER the 2x2 patch + merge, i.e. over ``hidden_size * merge_h * merge_w``) + * A Qwen2 language-model backbone + +The model emits structured grounding outputs such as +``object...`` when special tokens are preserved +(``skip_special_tokens=False``). An optional constrained-decoding logit +processor (:class:`LocateAnythingBoxGrammarLogitProcessor`) restricts the tokens +emitted inside a ``...`` block to a valid ``none`` / point / bbox +pattern. +""" + +import logging +from typing import Any, Dict, Iterable, List, Optional, Set, Tuple + +import torch +from torch import nn + +from sglang.srt.configs.kimi_vl_moonvit import MoonViTConfig +from sglang.srt.configs.locate_anything import LocateAnythingConfig +from sglang.srt.layers.quantization.base_config import QuantizationConfig +from sglang.srt.managers.mm_utils import ( + MultiModalityDataPaddingPatternMultimodalTokens, + general_mm_embed_routine, +) +from sglang.srt.managers.schedule_batch import ( + Modality, + 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.kimi_vl_moonvit import MoonVitPretrainedModel +from sglang.srt.models.qwen2 import Qwen2ForCausalLM +from sglang.srt.sampling.custom_logit_processor import CustomLogitProcessor +from sglang.srt.utils import add_prefix + +logger = logging.getLogger(__name__) + + +class LocateAnythingMultiModalProjector(nn.Module): + """InternVL-style ``mlp1`` projector. + + Unlike Kimi-VL's projector (which LayerNorms the per-patch features over + ``hidden_size`` *before* the 2x2 merge), LocateAnything merges first and then + LayerNorms over the merged width ``hidden_size * merge_h * merge_w``. + + HF checkpoint layout (``mlp1`` Sequential): + mlp1.0 = LayerNorm(merged_size) + mlp1.1 = Linear(merged_size, text_hidden) + mlp1.2 = GELU + mlp1.3 = Linear(text_hidden, text_hidden) + """ + + def __init__(self, config: LocateAnythingConfig): + super().__init__() + + merge = config.vision_config.merge_kernel_size + self.merged_size = config.vision_config.hidden_size * merge[0] * merge[1] + text_hidden = config.text_config.hidden_size + + self.pre_norm = nn.LayerNorm(self.merged_size, eps=1e-5) + self.linear_1 = nn.Linear(self.merged_size, text_hidden, bias=True) + # Plain (exact, erf-based) GELU to match the HF checkpoint's nn.GELU(). + self.act = nn.GELU() + self.linear_2 = nn.Linear(text_hidden, text_hidden, bias=True) + + def forward(self, image_features: torch.Tensor) -> torch.Tensor: + # MoonViT's patch_merger yields per-image tensors of shape + # (num_merged_tokens, merge_h * merge_w, hidden_size); concatenated and + # flattened to (num_merged_tokens, merged_size) the 4 sub-patches sit + # contiguously per token, matching the trained LayerNorm(merged_size). + # reshape (not view) since the concatenated input may be non-contiguous. + hidden_states = image_features.reshape(-1, self.merged_size) + hidden_states = self.pre_norm(hidden_states) + hidden_states = self.linear_1(hidden_states) + hidden_states = self.act(hidden_states) + hidden_states = self.linear_2(hidden_states) + return hidden_states + + +class LocateAnythingForConditionalGeneration(nn.Module): + def __init__( + self, + config: LocateAnythingConfig, + quant_config: Optional[QuantizationConfig] = None, + prefix: str = "", + **kwargs, + ) -> None: + super().__init__() + self.config = config + assert isinstance(config.vision_config, MoonViTConfig) + + self.vision_tower = MoonVitPretrainedModel(config.vision_config) + self.multi_modal_projector = LocateAnythingMultiModalProjector(config) + self.quant_config = quant_config + + self.language_model = Qwen2ForCausalLM( + config=config.text_config, + quant_config=quant_config, + prefix=add_prefix("language_model", prefix), + ) + + def get_image_feature(self, items: List[MultimodalDataItem]) -> torch.Tensor: + pixel_values = ( + torch.cat([item.feature for item in items], dim=0) + .type(self.vision_tower.dtype) + .to(self.vision_tower.device) + ) + + # Already-projected embeddings (e.g. precomputed) pass through. + if ( + pixel_values.dim() == 2 + and pixel_values.shape[-1] == self.config.text_config.hidden_size + ): + return pixel_values + + # image_grid_hws may arrive as numpy arrays from the HF image processor; + # coerce each to a tensor before concatenating. + image_grid_hws = torch.cat( + [torch.as_tensor(item.image_grid_hws) for item in items], dim=0 + ).to(self.vision_tower.device) + image_features = self.vision_tower(pixel_values, image_grid_hws) + assert isinstance(image_features, list) + return self.multi_modal_projector(torch.cat(image_features)) + + def pad_input_ids(self, input_ids: List[int], mm_inputs: MultimodalInputs): + pattern = MultiModalityDataPaddingPatternMultimodalTokens() + return pattern.pad_input_tokens(input_ids, mm_inputs) + + def forward( + self, + input_ids: torch.Tensor, + positions: torch.Tensor, + forward_batch: ForwardBatch, + get_embedding: bool = False, + ): + hidden_states = general_mm_embed_routine( + input_ids=input_ids, + forward_batch=forward_batch, + language_model=self.language_model, + data_embedding_funcs={ + Modality.IMAGE: self.get_image_feature, + }, + positions=positions, + ) + return hidden_states + + def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]) -> Set[str]: + # Remap HF checkpoint prefixes onto SGLang submodule names. + prefix_mapping = { + "vision_model.": "vision_tower.", + "mlp1.0.": "multi_modal_projector.pre_norm.", + "mlp1.1.": "multi_modal_projector.linear_1.", + "mlp1.3.": "multi_modal_projector.linear_2.", + } + + # Qwen2 packs qkv / gate-up; apply the same shard mapping for the LM part. + 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), + ] + + tie_word_embeddings = getattr( + self.config.text_config, "tie_word_embeddings", False + ) + + params_dict = dict(self.named_parameters()) + loaded_params: Set[str] = set() + + for name, loaded_weight in weights: + for src, dst in prefix_mapping.items(): + if name.startswith(src): + name = dst + name[len(src) :] + break + + if "rotary_emb.inv_freq" in name: + continue + if "rotary_emb.cos_cached" in name or "rotary_emb.sin_cached" in name: + continue + # Under tied embeddings the checkpoint's lm_head duplicates the input + # embedding and has no separate destination. + if tie_word_embeddings and name.startswith("language_model.lm_head."): + continue + + is_vision_weight = name.startswith("vision_tower.") or name.startswith( + "multi_modal_projector." + ) + + if is_vision_weight: + if name.endswith(".bias") and name not in params_dict: + continue + if name not in params_dict: + logger.warning(f"Parameter {name} not found 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) + continue + + # Language-model weights: apply Qwen2 stacked shard mapping. + for param_name, weight_name, shard_id in stacked_params_mapping: + if weight_name not in name: + continue + mapped = name.replace(weight_name, param_name) + if mapped.endswith(".bias") and mapped not in params_dict: + continue + if mapped not in params_dict: + continue + param = params_dict[mapped] + param.weight_loader(param, loaded_weight, shard_id) + loaded_params.add(mapped) + break + else: + if name.endswith(".bias") and name not in params_dict: + continue + if name not in params_dict: + logger.warning(f"Parameter {name} not found 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) + + # Reconcile: warn about any model parameter that never received a weight, + # so a partial/mismatched checkpoint is visible in the logs rather than + # silently serving garbage. Tied lm_head shares embed_tokens' storage and + # is loaded via it, so it is expected to be absent here. + missing = set(params_dict.keys()) - loaded_params + if tie_word_embeddings: + missing = { + n for n in missing if not n.startswith("language_model.lm_head.") + } + if missing: + logger.warning( + f"LocateAnything: {len(missing)} parameters did not receive " + f"weights, e.g. {sorted(missing)[:10]}" + ) + + return loaded_params + + +class LocateAnythingBoxGrammarLogitProcessor(CustomLogitProcessor): + """Constrained decoding for LocateAnything ``...`` blocks. + + Outside an open box the logits are untouched. Inside an open box (a + ``box_start`` with no matching ``box_end`` yet) the next token is restricted + so that the box body is one of: + + * ``none`` -> ``[none]`` + * a 2-coordinate point -> ``[c, c]`` + * a 4-coordinate bounding box -> ``[c, c, c, c]`` + + where ``c`` is any token in ``[coord_start_token_id, coord_end_token_id]``. + + Token ids are read per-request from ``custom_param_list[i]`` (keys + ``box_start_token_id``, ``box_end_token_id``, ``coord_start_token_id``, + ``coord_end_token_id``, ``none_token_id``) so the processor stays generic. + The ``__req__`` entry supplies the generated-so-far token ids. + + This processor is **opt-in**: it is never attached server-side, and the + server must be started with ``--enable-custom-logit-processor`` (off by + default) or the tokenizer rejects the request. A client enables it by + passing both the serialized processor and the matching token ids. + :meth:`build_sampling_params` wires both from a + :class:`LocateAnythingConfig` so callers don't hand-build the id dict. + + The two pieces live in **different** request fields, so do NOT spread them + both into ``sampling_params``: ``custom_logit_processor`` is a top-level + :class:`~sglang.srt.managers.io_struct.GenerateReqInput` field, while + ``custom_params`` is a :class:`SamplingParams` field. (Spreading both into + ``sampling_params`` raises ``TypeError: Unexpected keyword argument + 'custom_logit_processor'`` because ``SamplingParams`` is a strict + ``msgspec.Struct``.) Wire them like the OpenAI ``to_sampling_params`` path:: + + from sglang.srt.managers.io_struct import GenerateReqInput + from sglang.srt.models.locate_anything import ( + LocateAnythingBoxGrammarLogitProcessor, + ) + + extra = LocateAnythingBoxGrammarLogitProcessor.build_sampling_params(config) + req = GenerateReqInput( + text=prompt, + image_data=image, + sampling_params={ + "max_new_tokens": 8192, + "custom_params": extra["custom_params"], + }, + custom_logit_processor=extra["custom_logit_processor"], + ) + + Passing the processor without ``custom_params`` (or vice versa) silently + no-ops — both must be present together. + """ + + @classmethod + def build_sampling_params(cls, config: "LocateAnythingConfig") -> Dict[str, Any]: + """Build the two request fields needed to enable constrained decoding. + + Returns a dict with ``custom_logit_processor`` (the serialized + processor) and ``custom_params`` (the box/coord/none token ids read from + ``config``). These go to **different** request fields — put + ``custom_params`` inside ``sampling_params`` and pass + ``custom_logit_processor`` as a top-level ``GenerateReqInput`` field + (see the class docstring). The server also needs + ``--enable-custom-logit-processor``. + """ + return { + "custom_logit_processor": cls.to_str(), + "custom_params": { + "box_start_token_id": config.box_start_token_id, + "box_end_token_id": config.box_end_token_id, + "coord_start_token_id": config.coord_start_token_id, + "coord_end_token_id": config.coord_end_token_id, + "none_token_id": config.none_token_id, + }, + } + + def __call__( + self, + logits: torch.Tensor, + custom_param_list: Optional[List[Dict[str, Any]]] = None, + ) -> torch.Tensor: + if not custom_param_list: + return logits + + neg_inf = float("-inf") + for batch_idx, params in enumerate(custom_param_list): + if not params: + continue + req = params.get("__req__") + if req is None: + continue + + box_start = params.get("box_start_token_id") + box_end = params.get("box_end_token_id") + coord_start = params.get("coord_start_token_id") + coord_end = params.get("coord_end_token_id") + none_id = params.get("none_token_id") + if None in (box_start, box_end, coord_start, coord_end, none_id): + continue + + # Only the generated tokens are scanned (not origin_input_ids), + # which avoids an O(prompt_len) reverse scan per decode step over the + # long run. Assumes the prompt contains no *unclosed* + # : a closed ... in a few-shot / multi-turn prompt is + # harmless (last_open finds no open box here), but an unclosed + # left dangling in the prompt would not be constrained. + output_ids = list(req.output_ids) + + # Find the last box_start; if a box_end follows it, no box is open. + try: + last_open = len(output_ids) - 1 - output_ids[::-1].index(box_start) + except ValueError: + continue # no box opened yet + body = output_ids[last_open + 1 :] + if box_end in body: + continue # last box already closed + + num_coords = sum(1 for t in body if coord_start <= t <= coord_end) + has_none = none_id in body + + # Determine which token classes are allowed next. The coordinate + # range is contiguous, so it is masked as a slice rather than an + # enumerated set (the range can span ~1000 ids per decode step). + allow_coords = False + allow_scalars: Set[int] = set() + if has_none: + allow_scalars = {box_end} + elif num_coords == 0: + allow_coords, allow_scalars = True, {none_id} + elif num_coords in (1, 3): + allow_coords = True + elif num_coords == 2: + allow_coords, allow_scalars = True, {box_end} + else: # >= 4 coords -> must close + allow_scalars = {box_end} + + mask = torch.full_like(logits[batch_idx], neg_inf) + if allow_coords: + mask[coord_start : coord_end + 1] = logits[ + batch_idx, coord_start : coord_end + 1 + ] + for tok in allow_scalars: + mask[tok] = logits[batch_idx, tok] + logits[batch_idx] = mask + + return logits + + +EntryClass = [LocateAnythingForConditionalGeneration] diff --git a/python/sglang/srt/multimodal/processors/locate_anything.py b/python/sglang/srt/multimodal/processors/locate_anything.py new file mode 100644 index 000000000..1b2139fa6 --- /dev/null +++ b/python/sglang/srt/multimodal/processors/locate_anything.py @@ -0,0 +1,56 @@ +# SPDX-License-Identifier: Apache-2.0 +import re +from typing import Dict, List, Union + +from sglang.srt.managers.schedule_batch import MultimodalProcessorOutput +from sglang.srt.models.locate_anything import LocateAnythingForConditionalGeneration +from sglang.srt.multimodal.processors.base_processor import ( + BaseMultimodalProcessor as SGLangBaseProcessor, +) +from sglang.srt.multimodal.processors.base_processor import ( + MultimodalSpecialTokens, +) + + +# Compatible with LocateAnythingForConditionalGeneration +class LocateAnythingImageProcessor(SGLangBaseProcessor): + models = [LocateAnythingForConditionalGeneration] + # The LocateAnything HF processor is remote-code and does not support tensor inputs. + gpu_image_decode = False + + def __init__(self, hf_config, server_args, _processor, *args, **kwargs): + super().__init__(hf_config, server_args, _processor, *args, **kwargs) + # The model's chat template emits numbered ```` placeholders. + # The HF LocateAnythingProcessor expands each into + # ```` + N×```` + ```` and only the + # ```` (id 151665) run carries vision embeddings, so the + # offset/embedding token id is image_token_index while the prompt-level + # placeholder we split on is ````. + self.mm_tokens = MultimodalSpecialTokens( + image_token_id=hf_config.image_token_index, + image_token_regex=re.compile(r""), + ).build(_processor) + + async def process_mm_data_async( + self, + image_data: List[Union[str, bytes, Dict]], + input_text, + request_obj, + *args, + **kwargs, + ): + base_output = await self.load_mm_data( + prompt=input_text, + image_data=image_data, + multimodal_tokens=self.mm_tokens, + ) + + mm_items, input_ids, _ = self.process_and_combine_mm_data( + base_output, self.mm_tokens + ) + + return MultimodalProcessorOutput( + input_ids=input_ids.tolist(), + mm_items=mm_items, + im_token_id=self.mm_tokens.image_token_id, + ) diff --git a/python/sglang/srt/utils/hf_transformers/common.py b/python/sglang/srt/utils/hf_transformers/common.py index 1f93ae852..eef80dc7a 100644 --- a/python/sglang/srt/utils/hf_transformers/common.py +++ b/python/sglang/srt/utils/hf_transformers/common.py @@ -39,6 +39,7 @@ from sglang.srt.configs import ( KimiLinearConfig, KimiVLConfig, LagunaConfig, + LocateAnythingConfig, LongcatFlashConfig, MiniCPMV4_6Config, MiniCPMV4_6VisionConfig, @@ -84,6 +85,7 @@ _CONFIG_REGISTRY: Dict[str, Type[PretrainedConfig]] = { DeepseekVL2Config, MultiModalityConfig, KimiVLConfig, + LocateAnythingConfig, InternVLChatConfig, LagunaConfig, Step3VLConfig, diff --git a/test/registered/unit/configs/test_locate_anything_config.py b/test/registered/unit/configs/test_locate_anything_config.py new file mode 100644 index 000000000..35404cb8b --- /dev/null +++ b/test/registered/unit/configs/test_locate_anything_config.py @@ -0,0 +1,62 @@ +"""Unit tests for ``sglang.srt.configs.locate_anything.LocateAnythingConfig``.""" + +import unittest + +from transformers.models.qwen2 import Qwen2Config + +from sglang.srt.configs import LocateAnythingConfig +from sglang.srt.configs.kimi_vl_moonvit import MoonViTConfig +from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import CustomTestCase + +register_cpu_ci(est_time=5, suite="base-a-test-cpu") + + +class TestLocateAnythingConfig(CustomTestCase): + def test_default_fields(self): + """Defaults reflect the nvidia/LocateAnything-3B reference config.""" + cfg = LocateAnythingConfig() + self.assertEqual(cfg.model_type, "locateanything") + # Special token ids used by the grounding grammar. + self.assertEqual(cfg.image_token_index, 151665) + self.assertEqual(cfg.box_start_token_id, 151668) + self.assertEqual(cfg.box_end_token_id, 151669) + self.assertEqual(cfg.ref_start_token_id, 151672) + self.assertEqual(cfg.ref_end_token_id, 151673) + self.assertEqual(cfg.coord_start_token_id, 151677) + self.assertEqual(cfg.coord_end_token_id, 152677) + self.assertEqual(cfg.none_token_id, 4064) + self.assertEqual(cfg.mlp_connector_layers, 2) + + def test_composite_subconfigs_default(self): + cfg = LocateAnythingConfig() + self.assertIsInstance(cfg.vision_config, MoonViTConfig) + self.assertIsInstance(cfg.text_config, Qwen2Config) + + def test_subconfigs_from_dict(self): + cfg = LocateAnythingConfig( + vision_config={"hidden_size": 1152, "merge_kernel_size": [2, 2]}, + text_config={"hidden_size": 2048, "tie_word_embeddings": True}, + ) + self.assertIsInstance(cfg.vision_config, MoonViTConfig) + self.assertIsInstance(cfg.text_config, Qwen2Config) + self.assertEqual(cfg.vision_config.hidden_size, 1152) + self.assertEqual(cfg.text_config.hidden_size, 2048) + self.assertTrue(cfg.text_config.tie_word_embeddings) + + def test_subconfigs_passthrough_instances(self): + vision = MoonViTConfig(hidden_size=1152) + text = Qwen2Config(hidden_size=2048) + cfg = LocateAnythingConfig(vision_config=vision, text_config=text) + self.assertIs(cfg.vision_config, vision) + self.assertIs(cfg.text_config, text) + + def test_registered_in_config_registry(self): + """``model_type`` resolves to the config class via SGLang's registry.""" + from sglang.srt.utils.hf_transformers.common import _CONFIG_REGISTRY + + self.assertIs(_CONFIG_REGISTRY.get("locateanything"), LocateAnythingConfig) + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/unit/managers/test_mm_utils_split.py b/test/registered/unit/managers/test_mm_utils_split.py new file mode 100644 index 000000000..61ab59668 --- /dev/null +++ b/test/registered/unit/managers/test_mm_utils_split.py @@ -0,0 +1,145 @@ +"""Unit tests for ``get_new_expanded_mm_items`` per-image splitting. + +This is the load-bearing behavioral path for multi-image requests: a bundled +``MultimodalDataItem`` (one item carrying N image offsets + a concatenated +feature) must be split back into N per-image items so RadixAttention can cache +each image independently and chunked-prefill can encode them one at a time. + +The MoonViT-style models (e.g. nvidia/LocateAnything-3B) carry their per-image +grids under ``image_grid_hws`` rather than ``image_grid_thw``; the splitter must +recognize both keys, fall back cleanly when no usable grid is present, and not +mis-split a degenerate flat grid. No server / GPU / weight loading involved. +""" + +import unittest + +import numpy as np +import torch + +from sglang.srt.managers.mm_utils import get_new_expanded_mm_items +from sglang.srt.managers.schedule_batch import Modality, MultimodalDataItem +from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import CustomTestCase + +register_cpu_ci(est_time=10, suite="base-a-test-cpu") + + +def _bundled_item(grid_key=None, grid=None, feature_len=10, num_images=2): + """A bundled IMAGE item: `num_images` offsets, one concatenated feature.""" + model_specific_data = {} + if grid_key is not None: + model_specific_data[grid_key] = grid + # Distinct per-row values so slice boundaries are checkable. + feature = torch.arange(feature_len * 3, dtype=torch.float32).reshape(feature_len, 3) + offsets = [(0, 5), (5, feature_len)][:num_images] + return MultimodalDataItem( + modality=Modality.IMAGE, + offsets=offsets, + feature=feature, + model_specific_data=model_specific_data, + ) + + +class TestGetNewExpandedMMItems(CustomTestCase): + def test_image_grid_hws_splits_per_image(self): + # grid rows [[2,3],[4,1]] -> prod = [6, 4] patches -> feature_len 10. + item = _bundled_item( + grid_key="image_grid_hws", + grid=[[2, 3], [4, 1]], + feature_len=10, + ) + out = get_new_expanded_mm_items([item]) + + self.assertEqual(len(out), 2) + self.assertEqual([len(o.offsets) for o in out], [1, 1]) + self.assertEqual(out[0].offsets, [(0, 5)]) + self.assertEqual(out[1].offsets, [(5, 10)]) + # Feature sliced 0:6 and 6:10 along dim-0. + self.assertEqual(out[0].feature.shape[0], 6) + self.assertEqual(out[1].feature.shape[0], 4) + self.assertTrue(torch.equal(out[0].feature, item.feature[0:6])) + self.assertTrue(torch.equal(out[1].feature, item.feature[6:10])) + # Split items must re-hash (pad value is recomputed per image). + self.assertTrue(all(o.hash is None for o in out)) + + def test_image_grid_hws_tensor_splits_per_image(self): + # Same as above but the grid arrives as a rank-2 tensor (HF emits these). + item = _bundled_item( + grid_key="image_grid_hws", + grid=torch.tensor([[2, 3], [4, 1]], dtype=torch.long), + feature_len=10, + ) + out = get_new_expanded_mm_items([item]) + + self.assertEqual(len(out), 2) + self.assertTrue(torch.equal(out[0].feature, item.feature[0:6])) + self.assertTrue(torch.equal(out[1].feature, item.feature[6:10])) + + def test_image_grid_thw_still_splits(self): + # The pre-existing image_grid_thw path must keep working: + # [[1,2,3],[1,4,1]] -> [6,4]. + item = _bundled_item( + grid_key="image_grid_thw", + grid=[[1, 2, 3], [1, 4, 1]], + feature_len=10, + ) + out = get_new_expanded_mm_items([item]) + + self.assertEqual(len(out), 2) + self.assertTrue(torch.equal(out[0].feature, item.feature[0:6])) + self.assertTrue(torch.equal(out[1].feature, item.feature[6:10])) + + def test_missing_grid_falls_back_to_simple_split(self): + # No grid, but feature dim-0 == num offsets -> simple per-row split. + item = _bundled_item(grid_key=None, feature_len=2, num_images=2) + out = get_new_expanded_mm_items([item]) + + self.assertEqual(len(out), 2) + self.assertTrue(torch.equal(out[0].feature, item.feature[0:1])) + self.assertTrue(torch.equal(out[1].feature, item.feature[1:2])) + + def test_flat_1d_grid_does_not_mis_split(self): + # A flat 1-D grid (`tensor([2, 2])`) has length == num_items so it passes + # the length check, but prod(dim=-1) would collapse it to a scalar and + # corrupt the slice boundaries. The rank-2 guard must reject it. With + # feature_len != num_items, the simple-split fallback also declines, so + # the bundled item is passed through unchanged (never mis-sliced). + item = _bundled_item( + grid_key="image_grid_hws", + grid=torch.tensor([2, 2], dtype=torch.long), + feature_len=10, + ) + out = get_new_expanded_mm_items([item]) + + self.assertEqual(len(out), 1) + self.assertIs(out[0], item) + + def test_numpy_grid_splits_per_image(self): + # image_grid_hws can arrive as a numpy array from the HF image processor. + item = _bundled_item( + grid_key="image_grid_hws", + grid=np.array([[2, 3], [4, 1]], dtype=np.int64), + feature_len=10, + ) + out = get_new_expanded_mm_items([item]) + + self.assertEqual(len(out), 2) + self.assertTrue(torch.equal(out[0].feature, item.feature[0:6])) + self.assertTrue(torch.equal(out[1].feature, item.feature[6:10])) + + def test_non_bundled_item_passes_through(self): + # A single-image item (one offset) is not bundled and is returned as-is. + item = MultimodalDataItem( + modality=Modality.IMAGE, + offsets=[(0, 5)], + feature=torch.arange(18, dtype=torch.float32).reshape(6, 3), + model_specific_data={"image_grid_hws": [[2, 3]]}, + ) + out = get_new_expanded_mm_items([item]) + + self.assertEqual(len(out), 1) + self.assertIs(out[0], item) + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/unit/models/test_locate_anything.py b/test/registered/unit/models/test_locate_anything.py new file mode 100644 index 000000000..2af3d3935 --- /dev/null +++ b/test/registered/unit/models/test_locate_anything.py @@ -0,0 +1,340 @@ +"""Unit tests for srt/models/locate_anything.py — no server, no weight loading. + +Covers the InternVL-style ``mlp1`` projector shape and the optional box-grammar +logit processor's constrained-decoding state machine. +""" + +import unittest + +import numpy as np +import torch + +from sglang.srt.configs import LocateAnythingConfig +from sglang.srt.managers.schedule_batch import Modality, MultimodalDataItem +from sglang.srt.models.locate_anything import ( + LocateAnythingBoxGrammarLogitProcessor, + LocateAnythingForConditionalGeneration, + LocateAnythingMultiModalProjector, +) +from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import CustomTestCase + +register_cpu_ci(est_time=10, suite="base-a-test-cpu") + + +def _small_config(): + # Tiny dims keep the test fast and CPU-only. + return LocateAnythingConfig( + vision_config={"hidden_size": 8, "merge_kernel_size": [2, 2]}, + text_config={"hidden_size": 16}, + ) + + +class TestLocateAnythingProjector(CustomTestCase): + def test_merged_size_and_output_shape(self): + cfg = _small_config() + proj = LocateAnythingMultiModalProjector(cfg) + # merged_size = hidden_size * merge_h * merge_w = 8 * 2 * 2 = 32 + self.assertEqual(proj.merged_size, 32) + self.assertEqual(proj.pre_norm.normalized_shape, (32,)) + self.assertEqual(proj.linear_1.in_features, 32) + self.assertEqual(proj.linear_1.out_features, 16) + self.assertEqual(proj.linear_2.in_features, 16) + self.assertEqual(proj.linear_2.out_features, 16) + + def test_forward_flattens_merged_patches(self): + cfg = _small_config() + proj = LocateAnythingMultiModalProjector(cfg).eval() + # MoonViT patch_merger yields (num_merged_tokens, merge_h*merge_w, hidden). + num_tokens = 5 + feats = torch.randn(num_tokens, 4, 8) + with torch.no_grad(): + out = proj(feats) + # One projected vector of text_hidden width per merged token. + self.assertEqual(out.shape, (num_tokens, 16)) + + def test_forward_handles_noncontiguous_input(self): + cfg = _small_config() + proj = LocateAnythingMultiModalProjector(cfg).eval() + # A transposed/sliced tensor is non-contiguous; reshape (not view) must cope. + feats = torch.randn(4, 5, 8).transpose(0, 1) # (5, 4, 8), non-contiguous + self.assertFalse(feats.is_contiguous()) + with torch.no_grad(): + out = proj(feats) + self.assertEqual(out.shape, (5, 16)) + + +class _FakeReq: + def __init__(self, output_ids): + self.origin_input_ids = [1, 2, 3] + self.output_ids = output_ids + + +class TestBoxGrammarLogitProcessor(CustomTestCase): + # Token-id layout mirroring nvidia/LocateAnything-3B. + BOX_START = 151668 + BOX_END = 151669 + COORD_START = 151677 + COORD_END = 152677 + NONE = 4064 + VOCAB = 152681 + + def _params(self, output_ids): + return [ + { + "__req__": _FakeReq(output_ids), + "box_start_token_id": self.BOX_START, + "box_end_token_id": self.BOX_END, + "coord_start_token_id": self.COORD_START, + "coord_end_token_id": self.COORD_END, + "none_token_id": self.NONE, + } + ] + + def _allowed_ids(self, output_ids): + proc = LocateAnythingBoxGrammarLogitProcessor() + logits = torch.zeros(1, self.VOCAB) + out = proc(logits, self._params(output_ids)) + # Allowed ids are those left finite after masking. + return set(torch.nonzero(torch.isfinite(out[0])).flatten().tolist()) + + def test_no_box_open_is_untouched(self): + proc = LocateAnythingBoxGrammarLogitProcessor() + logits = torch.randn(1, self.VOCAB) + original = logits.clone() + out = proc(logits, self._params([42, 43])) # no box_start + self.assertTrue(torch.equal(out, original)) + + def test_just_after_box_start_allows_coords_or_none(self): + allowed = self._allowed_ids([self.BOX_START]) + self.assertIn(self.NONE, allowed) + self.assertIn(self.COORD_START, allowed) + self.assertIn(self.COORD_END, allowed) + self.assertNotIn(self.BOX_END, allowed) + + def test_after_none_must_close(self): + allowed = self._allowed_ids([self.BOX_START, self.NONE]) + self.assertEqual(allowed, {self.BOX_END}) + + def test_one_coord_forces_more_coords(self): + allowed = self._allowed_ids([self.BOX_START, self.COORD_START]) + self.assertNotIn(self.BOX_END, allowed) + self.assertNotIn(self.NONE, allowed) + self.assertIn(self.COORD_START, allowed) + + def test_two_coords_may_close_point_or_continue(self): + allowed = self._allowed_ids( + [self.BOX_START, self.COORD_START, self.COORD_START] + ) + self.assertIn(self.BOX_END, allowed) # 2-coord point can close + self.assertIn(self.COORD_START, allowed) # or continue toward a bbox + + def test_three_coords_forces_fourth(self): + allowed = self._allowed_ids([self.BOX_START] + [self.COORD_START] * 3) + self.assertNotIn(self.BOX_END, allowed) + self.assertIn(self.COORD_START, allowed) + + def test_four_coords_must_close(self): + allowed = self._allowed_ids([self.BOX_START] + [self.COORD_START] * 4) + self.assertEqual(allowed, {self.BOX_END}) + + def test_more_than_four_coords_must_close(self): + # The ">= 4 coords -> must close" branch must also fire if the model + # somehow emitted a 5th coordinate. + allowed = self._allowed_ids([self.BOX_START] + [self.COORD_START] * 5) + self.assertEqual(allowed, {self.BOX_END}) + + def test_coord_end_counts_as_a_coordinate(self): + # The coord range check is inclusive of coord_end (coord_start <= t <= + # coord_end); a body holding only coord_end must be treated as 1 coord. + allowed = self._allowed_ids([self.BOX_START, self.COORD_END]) + self.assertNotIn(self.BOX_END, allowed) # 1 coord -> need more + self.assertNotIn(self.NONE, allowed) + self.assertIn(self.COORD_START, allowed) + + def test_missing_token_id_is_noop(self): + # If a client passes custom_params missing one of the five ids, the + # processor must skip that request rather than crash or partially mask. + proc = LocateAnythingBoxGrammarLogitProcessor() + logits = torch.randn(1, self.VOCAB) + original = logits.clone() + params = self._params([self.BOX_START]) + del params[0]["none_token_id"] + out = proc(logits, params) + self.assertTrue(torch.equal(out, original)) + + def test_closed_box_is_untouched(self): + proc = LocateAnythingBoxGrammarLogitProcessor() + logits = torch.randn(1, self.VOCAB) + original = logits.clone() + # A fully-formed bbox that is already closed. + out = proc( + logits, + self._params([self.BOX_START] + [self.COORD_START] * 4 + [self.BOX_END]), + ) + self.assertTrue(torch.equal(out, original)) + + def test_empty_param_list_is_noop(self): + proc = LocateAnythingBoxGrammarLogitProcessor() + logits = torch.randn(1, self.VOCAB) + original = logits.clone() + self.assertTrue(torch.equal(proc(logits, None), original)) + + def test_build_sampling_params_wires_config_token_ids(self): + config = _small_config() + params = LocateAnythingBoxGrammarLogitProcessor.build_sampling_params(config) + # Serialized processor + the 5 token ids the processor reads per request. + self.assertIn("custom_logit_processor", params) + self.assertEqual( + params["custom_logit_processor"], + LocateAnythingBoxGrammarLogitProcessor.to_str(), + ) + self.assertEqual( + params["custom_params"], + { + "box_start_token_id": config.box_start_token_id, + "box_end_token_id": config.box_end_token_id, + "coord_start_token_id": config.coord_start_token_id, + "coord_end_token_id": config.coord_end_token_id, + "none_token_id": config.none_token_id, + }, + ) + + +class _StubVisionTower: + """Stand-in for MoonViT in get_image_feature. + + The real vision tower has its own tests (kimi_vl_moonvit); here we only need + it to (a) expose ``dtype``/``device`` and (b) return one ``(N, merge, hidden)`` + feature block per image so the projector + concat wiring is exercised with + real shapes. ``patches_per_image`` mirrors ``prod(image_grid_hws)``. + + To keep the oracle honest, ``__call__`` asserts that get_image_feature fed + it the inputs we expect — a ``(sum(patches), hidden)`` pixel tensor and a + rank-2 ``(num_images, 2)`` ``image_grid_hws`` whose per-row product matches + ``patches_per_image`` — so a regression in how the feature/grid are wired or + coerced fails here rather than passing on a fabricated shape. + """ + + def __init__(self, hidden, merge, patches_per_image): + self.dtype = torch.float32 + self.device = torch.device("cpu") + self._hidden = hidden + self._merge = merge + self._patches = patches_per_image + + def __call__(self, pixel_values, image_grid_hws): + # The concatenated raw patches across all images must line up. + assert pixel_values.shape == ( + sum(self._patches), + self._hidden, + ), pixel_values.shape + # image_grid_hws must be coerced to a rank-2 (num_images, 2) tensor whose + # rows multiply to the expected patch counts. + assert isinstance(image_grid_hws, torch.Tensor) + assert image_grid_hws.shape == (len(self._patches), 2), image_grid_hws.shape + assert image_grid_hws.prod(dim=-1).tolist() == list(self._patches) + # MoonViT yields a list of (num_merged_tokens, merge, hidden) per image. + return [ + torch.zeros(p // self._merge, self._merge, self._hidden) + for p in self._patches + ] + + +def _bare_model(config): + """A LocateAnythingForConditionalGeneration with a real projector but a + stubbed vision tower, bypassing the distributed Qwen2 __init__.""" + import torch.nn as nn + + model = LocateAnythingForConditionalGeneration.__new__( + LocateAnythingForConditionalGeneration + ) + nn.Module.__init__(model) + model.config = config + model.multi_modal_projector = LocateAnythingMultiModalProjector(config).eval() + return model + + +def _image_item(feature, grid_hws): + return MultimodalDataItem( + modality=Modality.IMAGE, + offsets=[(0, 1)], + feature=feature, + model_specific_data={"image_grid_hws": grid_hws}, + ) + + +class TestGetImageFeatureWiring(CustomTestCase): + """Forward-shape smoke test for get_image_feature. + + Guards the production path (pixel concat -> vision tower -> projector) and + the precomputed-embedding passthrough so a future change to the wiring or + the numpy->tensor image_grid_hws coercion doesn't silently regress. The + heavy MoonViT forward is stubbed (covered by its own tests); the projector + is real. + """ + + HIDDEN = 8 # vision hidden_size, must match _small_config() + MERGE = 4 # merge_h * merge_w = 2 * 2 + TEXT_HIDDEN = 16 # text_config hidden_size + + def test_single_image_projects_to_text_hidden(self): + cfg = _small_config() + model = _bare_model(cfg) + # grid [[2, 2]] -> prod = 4 patches. + model.vision_tower = _StubVisionTower(self.HIDDEN, self.MERGE, [4]) + feature = torch.randn(4, self.HIDDEN) # one image's raw patches + out = model.get_image_feature([_image_item(feature, [[2, 2]])]) + # 4 patches / merge(4) = 1 merged token, projected to text hidden width. + self.assertEqual(out.shape, (1, self.TEXT_HIDDEN)) + + def test_multi_image_features_concatenated_in_order(self): + cfg = _small_config() + model = _bare_model(cfg) + # Two images: [[2, 2]] -> 4 patches, [[4, 2]] -> 8 patches. + model.vision_tower = _StubVisionTower(self.HIDDEN, self.MERGE, [4, 8]) + items = [ + _image_item(torch.randn(4, self.HIDDEN), [[2, 2]]), + _image_item(torch.randn(8, self.HIDDEN), [[4, 2]]), + ] + out = model.get_image_feature(items) + # Merged tokens: 4/4 + 8/4 = 1 + 2 = 3, each projected to text hidden. + self.assertEqual(out.shape, (3, self.TEXT_HIDDEN)) + + def test_image_grid_hws_numpy_is_coerced(self): + # The HF image processor hands image_grid_hws back as a numpy array; + # get_image_feature must torch.as_tensor it before torch.cat (else the + # cat raises). A numpy grid must produce the same shape as a list grid. + cfg = _small_config() + model = _bare_model(cfg) + model.vision_tower = _StubVisionTower(self.HIDDEN, self.MERGE, [4]) + feature = torch.randn(4, self.HIDDEN) + grid = np.array([[2, 2]], dtype=np.int64) + out = model.get_image_feature([_image_item(feature, grid)]) + self.assertEqual(out.shape, (1, self.TEXT_HIDDEN)) + + def test_precomputed_embeddings_pass_through(self): + # Already-projected embeddings (dim==2, last dim == text hidden) must be + # returned untouched without invoking the vision tower forward. (dtype/ + # device are still read for the cast, so the stub exposes them but raises + # if its forward is actually called.) + cfg = _small_config() + model = _bare_model(cfg) + + class _NoCallTower: + dtype = torch.float32 + device = torch.device("cpu") + + def __call__(self, *args, **kwargs): + raise AssertionError( + "vision_tower forward should not run on precomputed embeds" + ) + + model.vision_tower = _NoCallTower() + embeds = torch.randn(5, self.TEXT_HIDDEN) + out = model.get_image_feature([_image_item(embeds, [[2, 2]])]) + self.assertTrue(torch.equal(out, embeds)) + + +if __name__ == "__main__": + unittest.main()