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:
co-authored by
Claude Opus 4.6
parent
f3970b17ef
commit
9f7792415a
@@ -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],
|
||||||
|
|||||||
@@ -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:
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
Reference in New Issue
Block a user