Clean up TokenizerManager: remove dead code and improve rid validation (#21639)

Co-authored-by: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
Lianmin Zheng
2026-03-29 15:12:49 -07:00
committed by GitHub
co-authored by Claude Opus 4.6
parent f3970b17ef
commit 9f7792415a
3 changed files with 34 additions and 32 deletions
+1 -5
View File
@@ -513,11 +513,7 @@ async def health_generate(request: Request) -> Response:
sampling_params = {"max_new_tokens": 1, "temperature": 0.0} sampling_params = {"max_new_tokens": 1, "temperature": 0.0}
rid = f"{HEALTH_CHECK_RID_PREFIX}_{time.time()}" rid = f"{HEALTH_CHECK_RID_PREFIX}_{time.time()}"
if _global_state.tokenizer_manager.is_image_gen: if _global_state.tokenizer_manager.is_generation:
gri = _global_state.tokenizer_manager.get_image_gen_health_check_request(
rid, sampling_params
)
elif _global_state.tokenizer_manager.is_generation:
gri = GenerateReqInput( gri = GenerateReqInput(
rid=rid, rid=rid,
input_ids=[0], input_ids=[0],
+14
View File
@@ -21,6 +21,7 @@ from __future__ import annotations
import copy import copy
import uuid import uuid
from abc import ABC from abc import ABC
from collections import Counter
from dataclasses import dataclass, field from dataclasses import dataclass, field
from enum import Enum from enum import Enum
from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Union from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Union
@@ -58,6 +59,15 @@ class BaseReq(ABC):
self.rid = uuid.uuid4().hex self.rid = uuid.uuid4().hex
return self.rid return self.rid
def _validate_rid_uniqueness(self):
"""Validate that request IDs within a batch are unique."""
if isinstance(self.rid, list) and len(set(self.rid)) != len(self.rid):
counts = Counter(self.rid)
duplicates = [rid for rid, count in counts.items() if count > 1]
raise ValueError(
f"Duplicate request IDs detected within the request: {duplicates}"
)
@dataclass @dataclass
class BaseBatchReq(ABC): class BaseBatchReq(ABC):
@@ -276,6 +286,8 @@ class GenerateReqInput(BaseReq):
else: else:
self._normalize_batch_inputs() self._normalize_batch_inputs()
self._validate_rid_uniqueness()
def _validate_inputs(self): def _validate_inputs(self):
"""Validate that the input configuration is valid.""" """Validate that the input configuration is valid."""
if ( if (
@@ -853,6 +865,8 @@ class EmbeddingReqInput(BaseReq):
self._normalize_lora_paths(self.batch_size) self._normalize_lora_paths(self.batch_size)
self._validate_rid_uniqueness()
def _normalize_lora_paths(self, num): def _normalize_lora_paths(self, num):
"""Normalize LoRA paths for batch processing.""" """Normalize LoRA paths for batch processing."""
if self.lora_path is not None: if self.lora_path is not None:
+19 -27
View File
@@ -132,6 +132,8 @@ class ReqState:
finished: bool finished: bool
event: asyncio.Event event: asyncio.Event
obj: Union[GenerateReqInput, EmbeddingReqInput] obj: Union[GenerateReqInput, EmbeddingReqInput]
# For performance metrics
time_stats: APIServerReqTimeStats time_stats: APIServerReqTimeStats
last_completion_tokens: int = 1 last_completion_tokens: int = 1
ttft_observed: bool = False ttft_observed: bool = False
@@ -219,9 +221,6 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
# Init metric collector and watchdog # Init metric collector and watchdog
self.init_metric_collector_watchdog() self.init_metric_collector_watchdog()
if self.enable_metrics:
start_cpu_monitor_thread("tokenizer")
# Init request dispatcher # Init request dispatcher
self.init_request_dispatcher() self.init_request_dispatcher()
@@ -234,7 +233,6 @@ 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 = 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
@@ -342,10 +340,6 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
self.gracefully_exit = False self.gracefully_exit = False
self.last_receive_tstamp = real_time() self.last_receive_tstamp = real_time()
# For load balancing
self.current_load = 0
self.current_load_lock = asyncio.Lock()
# Session # Session
self.session_futures = {} # session_id -> asyncio event self.session_futures = {} # session_id -> asyncio event
@@ -444,6 +438,8 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
collect_tokens_histogram=self.server_args.collect_tokens_histogram, collect_tokens_histogram=self.server_args.collect_tokens_histogram,
) )
start_cpu_monitor_thread("tokenizer")
if self.server_args.gc_warning_threshold_secs > 0.0: if self.server_args.gc_warning_threshold_secs > 0.0:
configure_gc_warning(self.server_args.gc_warning_threshold_secs) configure_gc_warning(self.server_args.gc_warning_threshold_secs)
self.soft_watchdog = Watchdog.create( self.soft_watchdog = Watchdog.create(
@@ -491,7 +487,7 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
# Normalize the request # Normalize the request
obj.normalize_batch_and_arguments() obj.normalize_batch_and_arguments()
self._set_default_priority(obj) self._set_default_priority(obj)
self._validate_rid(obj) self._validate_rid_not_in_flight(obj)
if isinstance(obj, GenerateReqInput) and obj.routed_dp_rank is not None: if isinstance(obj, GenerateReqInput) and obj.routed_dp_rank is not None:
dp_size = self.server_args.dp_size dp_size = self.server_args.dp_size
@@ -773,20 +769,16 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
obj, input_text, input_ids, input_embeds, mm_inputs, token_type_ids obj, input_text, input_ids, input_embeds, mm_inputs, token_type_ids
) )
def _validate_rid(self, obj: Union[GenerateReqInput, EmbeddingReqInput]) -> None: def _validate_rid_not_in_flight(
"""Validate the request ID (rid) uniqueness.""" self, obj: Union[GenerateReqInput, EmbeddingReqInput]
rid = obj.rid ) -> None:
if rid is None: """Validate that request IDs are not already in flight."""
if obj.rid is None:
return return
ids = rid if isinstance(rid, list) else [rid] rids = obj.rid if isinstance(obj.rid, list) else [obj.rid]
if len(ids) != len(set(ids)): conflicts = set(rids) & self.rid_to_state.keys()
raise ValueError( if conflicts:
f"Duplicate request IDs detected within the request: {ids}" raise ValueError(f"Duplicate request IDs detected: {list(conflicts)}")
)
for i in ids:
if i in self.rid_to_state:
raise ValueError(f"Duplicate request ID detected: {i}")
def _validate_one_request( def _validate_one_request(
self, obj: Union[GenerateReqInput, EmbeddingReqInput], input_ids: List[int] self, obj: Union[GenerateReqInput, EmbeddingReqInput], input_ids: List[int]
@@ -2314,14 +2306,14 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
external_trace_header = None external_trace_header = None
if self.server_args.enable_trace: if self.server_args.enable_trace:
if request: if obj.external_trace_header:
external_trace_header = extract_trace_headers(request.headers) # When the request comes from the rust grpc server or Engine there isn't a
obj.external_trace_header = external_trace_header
elif obj.external_trace_header:
# When the request comes form the rust grpc server or Engine there isn't a
# real request object but we still need to propagate the trace context from # real request object but we still need to propagate the trace context from
# the trace context that is explicitly passed in # the trace context that is explicitly passed in
external_trace_header = obj.external_trace_header external_trace_header = obj.external_trace_header
elif request:
external_trace_header = extract_trace_headers(request.headers)
obj.external_trace_header = external_trace_header
if not hasattr(obj, "is_single") or obj.is_single: if not hasattr(obj, "is_single") or obj.is_single:
time_stats = APIServerReqTimeStats(disagg_mode=self.disaggregation_mode) time_stats = APIServerReqTimeStats(disagg_mode=self.disaggregation_mode)