Clean up detokenizer and remove dead multimodal_gen code (#21588)

Co-authored-by: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
Lianmin Zheng
2026-03-28 21:44:40 -07:00
committed by GitHub
co-authored by Claude Opus 4.6
parent 516cff97a3
commit ba6b501f3a
13 changed files with 53 additions and 103 deletions
-14
View File
@@ -178,12 +178,6 @@ class ModelConfig:
self.is_multimodal = enable_multimodal and is_multimodal_model( self.is_multimodal = enable_multimodal and is_multimodal_model(
self.hf_config.architectures self.hf_config.architectures
) )
self.is_multimodal_gen = enable_multimodal and is_multimodal_gen_model(
self.hf_config.architectures
)
self.is_image_gen = enable_multimodal and is_image_gen_model(
self.hf_config.architectures
)
self.is_audio_model = enable_multimodal and is_audio_model( self.is_audio_model = enable_multimodal and is_audio_model(
self.hf_config.architectures self.hf_config.architectures
) )
@@ -1351,14 +1345,6 @@ def is_multimodal_model(model_architectures: List[str]):
return False return False
def is_multimodal_gen_model(model_architectures: List[str]):
return False
def is_image_gen_model(model_architectures: List[str]):
return False
def is_audio_model(model_architectures: List[str]): def is_audio_model(model_architectures: List[str]):
models = [ models = [
"WhisperForConditionalGeneration", "WhisperForConditionalGeneration",
+2
View File
@@ -8,3 +8,5 @@ GPU_MEMORY_ALL_TYPES = [
GPU_MEMORY_TYPE_WEIGHTS, GPU_MEMORY_TYPE_WEIGHTS,
GPU_MEMORY_TYPE_CUDA_GRAPH, GPU_MEMORY_TYPE_CUDA_GRAPH,
] ]
HEALTH_CHECK_RID_PREFIX = "HEALTH_CHECK"
+2 -1
View File
@@ -59,6 +59,7 @@ from fastapi.exceptions import RequestValidationError
from fastapi.middleware.cors import CORSMiddleware from fastapi.middleware.cors import CORSMiddleware
from fastapi.responses import ORJSONResponse, Response, StreamingResponse from fastapi.responses import ORJSONResponse, Response, StreamingResponse
from sglang.srt.constants import HEALTH_CHECK_RID_PREFIX
from sglang.srt.disaggregation.utils import FAKE_BOOTSTRAP_HOST, DisaggregationMode from sglang.srt.disaggregation.utils import FAKE_BOOTSTRAP_HOST, DisaggregationMode
from sglang.srt.entrypoints.anthropic.protocol import ( from sglang.srt.entrypoints.anthropic.protocol import (
AnthropicCountTokensRequest, AnthropicCountTokensRequest,
@@ -509,7 +510,7 @@ async def health_generate(request: Request) -> Response:
return Response(status_code=200) return Response(status_code=200)
sampling_params = {"max_new_tokens": 1, "temperature": 0.0} sampling_params = {"max_new_tokens": 1, "temperature": 0.0}
rid = f"HEALTH_CHECK_{time.time()}" rid = f"{HEALTH_CHECK_RID_PREFIX}_{time.time()}"
if _global_state.tokenizer_manager.is_image_gen: if _global_state.tokenizer_manager.is_image_gen:
gri = _global_state.tokenizer_manager.get_image_gen_health_check_request( gri = _global_state.tokenizer_manager.get_image_gen_health_check_request(
@@ -25,10 +25,10 @@ import pybase64
import setproctitle import setproctitle
import zmq import zmq
from sglang.srt.constants import HEALTH_CHECK_RID_PREFIX
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.managers.io_struct import ( from sglang.srt.managers.io_struct import (
BatchEmbeddingOutput, BatchEmbeddingOutput,
BatchMultimodalDecodeReq,
BatchStrOutput, BatchStrOutput,
BatchTokenIDOutput, BatchTokenIDOutput,
FreezeGCReq, FreezeGCReq,
@@ -88,16 +88,9 @@ class DetokenizerManager(MultiHttpWorkerDetokenizerMixin):
# Init running status # Init running status
self.init_running_status(server_args) self.init_running_status(server_args)
if server_args.enable_metrics:
start_cpu_monitor_thread("detokenizer")
# Init dispatcher # Init dispatcher
self.init_request_dispatcher() self.init_request_dispatcher()
@staticmethod
def is_health_check_request(rid: Optional[str]) -> bool:
return isinstance(rid, str) and rid.startswith("HEALTH_CHECK")
def init_ipc_channels(self, port_args: PortArgs): def init_ipc_channels(self, port_args: PortArgs):
context = zmq.Context(2) context = zmq.Context(2)
self.recv_from_scheduler = get_zmq_socket( self.recv_from_scheduler = get_zmq_socket(
@@ -120,9 +113,8 @@ class DetokenizerManager(MultiHttpWorkerDetokenizerMixin):
def init_running_status(self, server_args: ServerArgs): def init_running_status(self, server_args: ServerArgs):
self.decode_status = LimitedCapacityDict(capacity=DETOKENIZER_MAX_STATES) self.decode_status = LimitedCapacityDict(capacity=DETOKENIZER_MAX_STATES)
self.is_dummy = False
self.is_tool_call_parser_gpt_oss = server_args.tool_call_parser == "gpt-oss"
self.disable_tokenizer_batch_decode = server_args.disable_tokenizer_batch_decode self.disable_tokenizer_batch_decode = server_args.disable_tokenizer_batch_decode
self.is_tool_call_parser_gpt_oss = server_args.tool_call_parser == "gpt-oss"
self.soft_watchdog = Watchdog.create( self.soft_watchdog = Watchdog.create(
debug_name="DetokenizerManager", debug_name="DetokenizerManager",
@@ -131,12 +123,14 @@ class DetokenizerManager(MultiHttpWorkerDetokenizerMixin):
test_stuck_time=envs.SGLANG_TEST_STUCK_DETOKENIZER.get(), test_stuck_time=envs.SGLANG_TEST_STUCK_DETOKENIZER.get(),
) )
if server_args.enable_metrics:
start_cpu_monitor_thread("detokenizer")
def init_request_dispatcher(self): def init_request_dispatcher(self):
self._request_dispatcher = TypeBasedDispatcher( self._request_dispatcher = TypeBasedDispatcher(
[ [
(BatchEmbeddingOutput, self.handle_batch_embedding_out), (BatchEmbeddingOutput, self.handle_batch_embedding_out),
(BatchTokenIDOutput, self.handle_batch_token_id_out), (BatchTokenIDOutput, self.handle_batch_token_id_out),
(BatchMultimodalDecodeReq, self.handle_multimodal_decode_req),
(FreezeGCReq, self.handle_freeze_gc_req), (FreezeGCReq, self.handle_freeze_gc_req),
] ]
) )
@@ -190,8 +184,6 @@ class DetokenizerManager(MultiHttpWorkerDetokenizerMixin):
) -> List[str]: ) -> List[str]:
"""Batch decode with grouping by (skip_special_tokens, spaces_between_special_tokens).""" """Batch decode with grouping by (skip_special_tokens, spaces_between_special_tokens)."""
assert self.tokenizer is not None
# fast path # fast path
first_skip, first_space = skip_list[0], space_list[0] first_skip, first_space = skip_list[0], space_list[0]
if all(s == first_skip for s in skip_list) and all( if all(s == first_skip for s in skip_list) and all(
@@ -236,8 +228,6 @@ class DetokenizerManager(MultiHttpWorkerDetokenizerMixin):
surr_offset=0, surr_offset=0,
read_offset=recv_obj.read_offsets[i], read_offset=recv_obj.read_offsets[i],
) )
if not self.is_health_check_request(rid):
# for health check requests, we do not store the decode status
self.decode_status[rid] = s self.decode_status[rid] = s
else: else:
s = self.decode_status[rid] s = self.decode_status[rid]
@@ -254,8 +244,6 @@ class DetokenizerManager(MultiHttpWorkerDetokenizerMixin):
# Decode token ids to strings # Decode token ids to strings
if not self.disable_tokenizer_batch_decode: if not self.disable_tokenizer_batch_decode:
if not self.is_dummy:
# Run normal batch decode
surr_texts = self._grouped_batch_decode( surr_texts = self._grouped_batch_decode(
surr_ids, surr_ids,
recv_obj.skip_special_tokens, recv_obj.skip_special_tokens,
@@ -266,10 +254,6 @@ class DetokenizerManager(MultiHttpWorkerDetokenizerMixin):
recv_obj.skip_special_tokens, recv_obj.skip_special_tokens,
recv_obj.spaces_between_special_tokens, recv_obj.spaces_between_special_tokens,
) )
else:
# If it is dummy weights, just return dummy strings to prevent potential detokenization edge cases
surr_texts = ["dog" for _ in surr_ids]
read_texts = ["cat" for _ in read_ids]
else: else:
# Do not use batch decode to prevent some detokenization edge cases (e.g., gpt-oss). # Do not use batch decode to prevent some detokenization edge cases (e.g., gpt-oss).
surr_texts = [ surr_texts = [
@@ -297,14 +281,6 @@ class DetokenizerManager(MultiHttpWorkerDetokenizerMixin):
output_strs = [] output_strs = []
for i in range(bs): for i in range(bs):
rid = recv_obj.rids[i] rid = recv_obj.rids[i]
if self.is_health_check_request(rid):
s = DecodeStatus(
decoded_text=recv_obj.decoded_texts[i],
decode_ids=recv_obj.decode_ids[i],
surr_offset=0,
read_offset=recv_obj.read_offsets[i],
)
else:
try: try:
s = self.decode_status[rid] s = self.decode_status[rid]
except KeyError: except KeyError:
@@ -335,6 +311,7 @@ class DetokenizerManager(MultiHttpWorkerDetokenizerMixin):
recv_obj.finished_reasons[i], recv_obj.finished_reasons[i],
recv_obj.no_stop_trim[i], recv_obj.no_stop_trim[i],
) )
# Incrementally send text. # Incrementally send text.
incremental_output = output_str[s.sent_offset :] incremental_output = output_str[s.sent_offset :]
s.sent_offset = len(output_str) s.sent_offset = len(output_str)
@@ -404,14 +381,15 @@ class DetokenizerManager(MultiHttpWorkerDetokenizerMixin):
time_stats=recv_obj.time_stats, time_stats=recv_obj.time_stats,
) )
def handle_multimodal_decode_req(self, recv_obj: BatchMultimodalDecodeReq):
raise NotImplementedError()
def handle_freeze_gc_req(self, recv_req: FreezeGCReq): def handle_freeze_gc_req(self, recv_req: FreezeGCReq):
freeze_gc("Detokenizer Manager") freeze_gc("Detokenizer Manager")
return None return None
def is_health_check_request(rid: Optional[str]) -> bool:
return isinstance(rid, str) and rid.startswith(HEALTH_CHECK_RID_PREFIX)
class LimitedCapacityDict(OrderedDict): class LimitedCapacityDict(OrderedDict):
def __init__(self, capacity: int, *args, **kwargs): def __init__(self, capacity: int, *args, **kwargs):
super().__init__(*args, **kwargs) super().__init__(*args, **kwargs)
+2 -1
View File
@@ -39,6 +39,7 @@ from torch.distributed import barrier
from sglang.jit_kernel.ngram_embedding import update_token_table from sglang.jit_kernel.ngram_embedding import update_token_table
from sglang.srt.configs.model_config import ModelConfig from sglang.srt.configs.model_config import ModelConfig
from sglang.srt.constants import HEALTH_CHECK_RID_PREFIX
from sglang.srt.constrained.grammar_manager import GrammarManager from sglang.srt.constrained.grammar_manager import GrammarManager
from sglang.srt.disaggregation.decode import ( from sglang.srt.disaggregation.decode import (
DecodePreallocQueue, DecodePreallocQueue,
@@ -3412,7 +3413,7 @@ class IdleSleeper:
def is_health_check_generate_req(recv_req): def is_health_check_generate_req(recv_req):
rid = getattr(recv_req, "rid", None) rid = getattr(recv_req, "rid", None)
return rid is not None and rid.startswith("HEALTH_CHECK") return rid is not None and rid.startswith(HEALTH_CHECK_RID_PREFIX)
def is_work_request(recv_req): def is_work_request(recv_req):
@@ -953,10 +953,6 @@ class SchedulerOutputProcessorMixin:
if req is skip_req: if req is skip_req:
continue continue
# Multimodal partial stream chunks break the detokenizer, so drop aborted requests here.
if self.model_config.is_multimodal_gen and req.to_finish:
continue
if req.finished(): if req.finished():
if req.finished_output: if req.finished_output:
# With the overlap schedule, a request will try to output twice and hit this line twice # With the overlap schedule, a request will try to output twice and hit this line twice
@@ -975,8 +971,7 @@ class SchedulerOutputProcessorMixin:
# origin stream_interval logic # origin stream_interval logic
should_output = ( should_output = (
len(req.output_ids) % stream_interval == 1 len(req.output_ids) % stream_interval == 1
if not self.model_config.is_multimodal_gen if stream_interval > 1
and stream_interval > 1
else len(req.output_ids) % stream_interval == 0 else len(req.output_ids) % stream_interval == 0
) )
@@ -986,8 +981,6 @@ class SchedulerOutputProcessorMixin:
else: else:
should_output = ( should_output = (
len(req.output_ids) % DEFAULT_FORCE_STREAM_INTERVAL == 0 len(req.output_ids) % DEFAULT_FORCE_STREAM_INTERVAL == 0
if not self.model_config.is_multimodal_gen
else False
) )
if should_output: if should_output:
@@ -1003,9 +996,6 @@ class SchedulerOutputProcessorMixin:
decoded_texts.append(req.decoded_text) decoded_texts.append(req.decoded_text)
decode_ids, read_offset = req.init_incremental_detokenize() decode_ids, read_offset = req.init_incremental_detokenize()
if self.model_config.is_multimodal_gen:
decode_ids_list.append(decode_ids)
else:
decode_ids_list.append(decode_ids[req.send_decode_id_offset :]) decode_ids_list.append(decode_ids[req.send_decode_id_offset :])
# Exclude the tokens after stop condition # Exclude the tokens after stop condition
@@ -1132,8 +1122,6 @@ class SchedulerOutputProcessorMixin:
# Send to detokenizer # Send to detokenizer
if reqs or is_idle_batch: if reqs or is_idle_batch:
if self.model_config.is_multimodal_gen:
return
self.send_to_detokenizer.send_output( self.send_to_detokenizer.send_output(
BatchTokenIDOutput( BatchTokenIDOutput(
rids=rids, rids=rids,
@@ -231,7 +231,7 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
self.served_model_name = server_args.served_model_name self.served_model_name = server_args.served_model_name
self.model_config = model_config_class.from_server_args(server_args) self.model_config = model_config_class.from_server_args(server_args)
self.is_generation = self.model_config.is_generation self.is_generation = self.model_config.is_generation
self.is_image_gen = self.model_config.is_image_gen self.is_image_gen = getattr(self.model_config, "is_image_gen", False)
self.context_len = self.model_config.context_len self.context_len = self.model_config.context_len
self.image_token_id = self.model_config.image_token_id self.image_token_id = self.model_config.image_token_id
self.max_req_input_len = None # Will be set later in engine.py self.max_req_input_len = None # Will be set later in engine.py
@@ -1194,7 +1194,6 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
self.request_logger.log_finished_request( self.request_logger.log_finished_request(
obj, obj,
out, out,
is_multimodal_gen=self.model_config.is_multimodal_gen,
request=request, request=request,
) )
@@ -7,6 +7,7 @@ from abc import ABC, abstractmethod
from datetime import datetime from datetime import datetime
from typing import List, Optional, Union from typing import List, Optional, Union
from sglang.srt.constants import HEALTH_CHECK_RID_PREFIX
from sglang.srt.managers.io_struct import EmbeddingReqInput, GenerateReqInput from sglang.srt.managers.io_struct import EmbeddingReqInput, GenerateReqInput
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
@@ -128,7 +129,7 @@ class FileRequestMetricsExporter(RequestMetricsExporter):
self, obj: Union[GenerateReqInput, EmbeddingReqInput], out_dict: dict self, obj: Union[GenerateReqInput, EmbeddingReqInput], out_dict: dict
): ):
# Do not log health check requests, since they don't represent real user requests. # Do not log health check requests, since they don't represent real user requests.
if isinstance(obj.rid, str) and "HEALTH_CHECK" in obj.rid: if isinstance(obj.rid, str) and HEALTH_CHECK_RID_PREFIX in obj.rid:
return return
try: try:
+1 -4
View File
@@ -20,6 +20,7 @@ import asyncio
import builtins import builtins
import ctypes import ctypes
import functools import functools
import gc
import importlib import importlib
import inspect import inspect
import io import io
@@ -2945,8 +2946,6 @@ def configure_gc_warning(warn_threshold_secs):
def freeze_gc(context: str): def freeze_gc(context: str):
import gc
g0_before, g1_before, g2_before = gc_object_counts() g0_before, g1_before, g2_before = gc_object_counts()
gc.freeze() gc.freeze()
g0_after, g1_after, g2_after = gc_object_counts() g0_after, g1_after, g2_after = gc_object_counts()
@@ -2961,8 +2960,6 @@ def freeze_gc(context: str):
def configure_gc_logger(): def configure_gc_logger():
logger.info("Enable GC Logger") logger.info("Enable GC Logger")
import gc
gc_start_time = {} gc_start_time = {}
def gc_callback(phase, info): def gc_callback(phase, info):
+1 -7
View File
@@ -162,7 +162,6 @@ class RequestLogger:
self, self,
obj: Union["GenerateReqInput", "EmbeddingReqInput"], obj: Union["GenerateReqInput", "EmbeddingReqInput"],
out: Any, out: Any,
is_multimodal_gen: bool = False,
request: Optional["fastapi.Request"] = None, request: Optional["fastapi.Request"] = None,
) -> None: ) -> None:
if not self.log_requests: if not self.log_requests:
@@ -181,7 +180,6 @@ class RequestLogger:
} }
if headers: if headers:
log_data["headers"] = headers log_data["headers"] = headers
if not is_multimodal_gen:
log_data["out"] = _transform_data_for_logging( log_data["out"] = _transform_data_for_logging(
out, max_length, out_skip_names out, max_length, out_skip_names
) )
@@ -190,11 +188,7 @@ class RequestLogger:
obj_str = _dataclass_to_string_truncated( obj_str = _dataclass_to_string_truncated(
obj, max_length, skip_names=skip_names obj, max_length, skip_names=skip_names
) )
out_str = ( out_str = f", out={_dataclass_to_string_truncated(out, max_length, skip_names=out_skip_names)}"
""
if is_multimodal_gen
else f", out={_dataclass_to_string_truncated(out, max_length, skip_names=out_skip_names)}"
)
headers_str = f", headers={headers}" if headers else "" headers_str = f", headers={headers}" if headers else ""
self._log(f"Finish: obj={obj_str}{headers_str}{out_str}") self._log(f"Finish: obj={obj_str}{headers_str}{out_str}")
@@ -8,6 +8,7 @@ from pathlib import Path
from sglang.bench_serving import run_benchmark from sglang.bench_serving import run_benchmark
from sglang.benchmark.utils import parse_custom_headers from sglang.benchmark.utils import parse_custom_headers
from sglang.srt.constants import HEALTH_CHECK_RID_PREFIX
from sglang.srt.utils import kill_process_tree from sglang.srt.utils import kill_process_tree
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
from sglang.test.test_utils import ( from sglang.test.test_utils import (
@@ -80,7 +81,7 @@ class TestBenchServingFunctionality(CustomTestCase):
continue continue
text = obj.get("obj", {}).get("text") text = obj.get("obj", {}).get("text")
rid = obj.get("rid", "") rid = obj.get("rid", "")
if text and not rid.startswith("HEALTH_CHECK"): if text and not rid.startswith(HEALTH_CHECK_RID_PREFIX):
reqs.append(text) reqs.append(text)
self.assertGreaterEqual(len(reqs), NUM_CONVERSATIONS * NUM_TURNS) self.assertGreaterEqual(len(reqs), NUM_CONVERSATIONS * NUM_TURNS)
@@ -58,6 +58,7 @@ import tempfile
import unittest import unittest
from unittest.mock import MagicMock, patch from unittest.mock import MagicMock, patch
from sglang.srt.constants import HEALTH_CHECK_RID_PREFIX
from sglang.srt.observability.request_metrics_exporter import ( from sglang.srt.observability.request_metrics_exporter import (
FileRequestMetricsExporter, FileRequestMetricsExporter,
RequestMetricsExporter, RequestMetricsExporter,
@@ -243,7 +244,7 @@ class TestFileRequestMetricsExporter(unittest.TestCase):
def test_write_record_skips_health_check(self): def test_write_record_skips_health_check(self):
exporter = self._make_exporter() exporter = self._make_exporter()
obj = _GenerateReqInput(rid="HEALTH_CHECK_123", text="ping") obj = _GenerateReqInput(rid=f"{HEALTH_CHECK_RID_PREFIX}_123", text="ping")
asyncio.run(exporter.write_record(obj, {})) asyncio.run(exporter.write_record(obj, {}))
files = os.listdir(self.tmp_dir) files = os.listdir(self.tmp_dir)
+2 -1
View File
@@ -8,6 +8,7 @@ from pathlib import Path
import requests import requests
from sglang.srt.constants import HEALTH_CHECK_RID_PREFIX
from sglang.srt.utils import kill_process_tree from sglang.srt.utils import kill_process_tree
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
from sglang.test.test_utils import ( from sglang.test.test_utils import (
@@ -197,7 +198,7 @@ class TestRequestLoggerJson(BaseTestRequestLogger, CustomTestCase):
continue continue
rid = data.get("rid", "") rid = data.get("rid", "")
if rid.startswith("HEALTH_CHECK"): if rid.startswith(HEALTH_CHECK_RID_PREFIX):
continue continue
if data.get("event") == "request.received": if data.get("event") == "request.received":