Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
10 changes: 10 additions & 0 deletions components/src/dynamo/sglang/register.py
Original file line number Diff line number Diff line change
Expand Up @@ -53,6 +53,7 @@
effective_gateway_workers,
gateway_engine_id,
)
from dynamo.sglang.video_routing import publish_sglang_qwen_video_processor_contract

SGLANG_HICACHE_MOONCAKE_RUNTIME_KEY = "sglang_hicache_mooncake"
SPEC_DECODE_RUNTIME_KEY = "spec_decode"
Expand Down Expand Up @@ -433,6 +434,15 @@ async def get_runtime_config(
# generation overflow handling to their downstream backend.
if engine is not None:
publish_token_budget(runtime_config, _get_token_budget(engine, server_args))
# Hash forwarding currently exists on SGLang's aggregated generation
# path. Do not advertise exact video routing to disaggregated workers,
# whose prefill/decode handlers would otherwise publish incompatible
# KV-event placeholder hashes.
if dynamo_args.frontend_decoding and server_args.disaggregation_mode in (
Comment thread
krishung5 marked this conversation as resolved.
None,
"null",
):
publish_sglang_qwen_video_processor_contract(runtime_config, engine)
# set reasoning parser and tool call parser
runtime_config.reasoning_parser = dynamo_args.dyn_reasoning_parser
runtime_config.tool_call_parser = dynamo_args.dyn_tool_call_parser
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -4,7 +4,7 @@
import asyncio
import logging
import time
from typing import Any, AsyncGenerator, AsyncIterator, Dict, List, Mapping, Optional
from typing import Any, AsyncGenerator, AsyncIterator, Dict, Mapping, Optional

import numpy as np
import sglang as sgl
Expand Down Expand Up @@ -47,6 +47,7 @@
VIDEO_URL_KEY,
build_disagg_mm_kwargs,
extract_media_urls,
extract_mm_hashes,
raise_if_unextracted_multimodal,
)
from dynamo.sglang.request_utils import request_cache_salt
Expand Down Expand Up @@ -433,32 +434,6 @@ def _resolve_mm_hashes_supported(engine: Any) -> bool:
probe = filter_supported_async_generate_kwargs(engine, {"mm_hashes": None})
return "mm_hashes" in probe

@staticmethod
def _extract_mm_hashes(request: Dict[str, Any]) -> Optional[List[str]]:
"""Pull the per-image hashes the Rust frontend forwards via extra_args.

Returns ``None`` when the field is absent or malformed; SGLang then
recomputes the hash internally via ``hash_feature()``.
"""
extra_args = request.get("extra_args")
if not isinstance(extra_args, dict):
return None
mm_hashes = extra_args.get("mm_hashes")
if not mm_hashes:
return None
if not isinstance(mm_hashes, list):
return None
# Fail closed if a non-string slipped into the list — downstream
# SGLang treats mm_hashes as List[str] and a bad element would
# crash the worker mid-request. Routing falls back to text-prefix.
if not all(isinstance(h, str) for h in mm_hashes):
logging.warning(
"extra_args.mm_hashes contained non-str entries; "
"ignoring routing-side hashes and letting SGLang recompute"
)
return None
return mm_hashes

def _metadata_uploader_from_request(
self, request: Dict[str, Any]
) -> MetadataUploader | None:
Expand Down Expand Up @@ -797,7 +772,7 @@ async def generate(

mm_hashes_kwargs: Dict[str, Any] = {}
if self._mm_hashes_supported:
forwarded = self._extract_mm_hashes(request)
forwarded = extract_mm_hashes(request)
if forwarded is not None:
mm_hashes_kwargs["mm_hashes"] = forwarded

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -20,6 +20,14 @@
_SUPPORTED_MULTIMODAL_CONTENT_TYPES = frozenset(
{IMAGE_URL_KEY, AUDIO_URL_KEY, VIDEO_URL_KEY}
)
# BaseMultiModalProcessorOutput.organize_results() builds SGLang's mm_items in
# this order, independent of their order in the original prompt.
_SGLANG_MM_ITEM_MODALITY_ORDER = ("image", "video", "audio")
_MM_DATA_KEY_BY_MODALITY = {
"image": IMAGE_URL_KEY,
"video": VIDEO_URL_KEY,
"audio": AUDIO_URL_KEY,
}


def _multi_modal_data(request: Dict[str, Any]) -> Dict[str, Any]:
Expand Down Expand Up @@ -124,6 +132,94 @@ def extract_media_urls(
return urls or None


def extract_mm_hashes(request: Dict[str, Any]) -> list[str] | None:
"""Return frontend MM hashes in SGLang's modality-grouped item order."""
extra_args = request.get("extra_args")
if not isinstance(extra_args, dict):
return None

grouped = extra_args.get("mm_hashes_by_modality")
if grouped is not None:
if not isinstance(grouped, dict):
logger.warning(
"extra_args.mm_hashes_by_modality is not an object; "
"ignoring routing-side hashes and letting SGLang recompute"
)
return None

unknown_modalities = {
str(modality)
for modality, hashes in grouped.items()
if modality not in _SGLANG_MM_ITEM_MODALITY_ORDER and hashes
}
if unknown_modalities:
logger.warning(
"extra_args.mm_hashes_by_modality contains unsupported "
"modalities %s; ignoring routing-side hashes and letting "
"SGLang recompute",
sorted(unknown_modalities),
)
return None

hashes_by_modality: dict[str, list[str]] = {}
for modality in _SGLANG_MM_ITEM_MODALITY_ORDER:
hashes = grouped.get(modality)
if hashes is None:
hashes = []
if not isinstance(hashes, list) or not all(
isinstance(value, str) for value in hashes
):
logger.warning(
"extra_args.mm_hashes_by_modality[%s] is not a string "
"list; ignoring routing-side hashes and letting SGLang recompute",
modality,
)
return None
hashes_by_modality[modality] = hashes

mm_data = request.get("multi_modal_data")
if not isinstance(mm_data, dict):
logger.warning(
"extra_args.mm_hashes_by_modality has no matching "
"multi_modal_data object; ignoring routing-side hashes and "
"letting SGLang recompute"
)
return None

flattened: list[str] = []
for modality in _SGLANG_MM_ITEM_MODALITY_ORDER:
hashes = hashes_by_modality[modality]
media_items = mm_data.get(_MM_DATA_KEY_BY_MODALITY[modality])
if media_items is None:
media_items = []
if not isinstance(media_items, list) or len(hashes) != len(media_items):
media_count = (
len(media_items) if isinstance(media_items, list) else None
)
logger.warning(
"extra_args.mm_hashes_by_modality[%s] count (%d) does not "
"match multi_modal_data count (%s); ignoring routing-side "
"hashes and letting SGLang recompute",
modality,
len(hashes),
media_count,
)
return None
flattened.extend(hashes)
return flattened or None

mm_hashes = extra_args.get("mm_hashes")
if not mm_hashes or not isinstance(mm_hashes, list):
return None
if not all(isinstance(value, str) for value in mm_hashes):
logger.warning(
"extra_args.mm_hashes contained non-str entries; ignoring "
"routing-side hashes and letting SGLang recompute"
)
return None
return mm_hashes


def build_disagg_mm_kwargs(request: Dict[str, Any]) -> Dict[str, Any]:
"""Build media kwargs for a disaggregated worker's ``async_generate`` call.

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -513,6 +513,37 @@ async def fake_async_generate(**kwargs):
assert captured["audio_data"] == ["https://example.com/a.wav"]


@pytest.mark.asyncio
async def test_aggregated_forwards_grouped_mm_hashes_in_sglang_item_order():
handler = _new_decode_handler(enable_frontend_decoding=False)
handler._mm_hashes_supported = True
captured: Dict[str, Any] = {}

async def fake_async_generate(**kwargs):
captured.update(kwargs)
return _empty_stream()

handler.engine = SimpleNamespace(async_generate=fake_async_generate)
request = {
"token_ids": [1, 2, 3],
"multi_modal_data": {
"image_url": ["https://example.com/a.jpg"],
"video_url": ["https://example.com/a.mp4"],
},
"extra_args": {
"mm_hashes_by_modality": {
"video": ["video-a"],
"image": ["image-a"],
}
},
}

async for _ in handler.generate(request, _Context()):
pass

assert captured["mm_hashes"] == ["image-a", "video-a"]


@pytest.mark.asyncio
async def test_aggregated_fd_on_loads_decoded_variants_to_pil():
"""With --frontend-decoding, Decoded items are loaded via ImageLoader and
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,7 @@
from dynamo.sglang.request_handlers.llm.mm_disagg_utils import (
build_disagg_mm_kwargs,
extract_media_urls,
extract_mm_hashes,
raise_if_unextracted_multimodal,
)
from dynamo.sglang.request_handlers.multimodal.worker_handler import StreamProcessor
Expand Down Expand Up @@ -79,6 +80,60 @@ def test_extract_media_urls_rejects_malformed_payloads():
extract_media_urls({"image_url": ""}, "image_url")


def test_extract_mm_hashes_preserves_legacy_image_protocol():
request = {"extra_args": {"mm_hashes": ["image-a", "image-b"]}}

assert extract_mm_hashes(request) == ["image-a", "image-b"]


def test_extract_mm_hashes_flattens_in_sglang_item_order():
request = {
"multi_modal_data": {
"image_url": ["image-a", "image-b"],
"video_url": ["video-a"],
},
"extra_args": {
"mm_hashes_by_modality": {
"video": ["video-a"],
"image": ["image-a", "image-b"],
},
},
}

assert extract_mm_hashes(request) == ["image-a", "image-b", "video-a"]


def test_extract_mm_hashes_rejects_per_modality_count_mismatch():
request = {
"multi_modal_data": {
"image_url": ["image-a", "image-b"],
"video_url": ["video-a"],
},
"extra_args": {
"mm_hashes_by_modality": {
# The total count still matches, but the modality association does not.
"image": ["image-a"],
"video": ["video-a", "video-b"],
}
},
}

assert extract_mm_hashes(request) is None


@pytest.mark.parametrize(
"grouped",
[
["not-an-object"],
{"video": "not-a-list"},
{"video": ["ok", 1]},
{"future_modality": ["hash"]},
],
)
def test_extract_mm_hashes_rejects_malformed_grouped_protocol(grouped):
assert extract_mm_hashes({"extra_args": {"mm_hashes_by_modality": grouped}}) is None


class TestMultimodalGuard:
"""Tests for multimodal guard when frontend extraction is missing."""

Expand Down
Loading
Loading