[Misc] Fix metrics, weight update lock, request logging (#2543)
This commit is contained in:
@@ -22,7 +22,7 @@ import signal
|
||||
import sys
|
||||
import time
|
||||
import uuid
|
||||
from typing import Any, Dict, List, Optional, Union
|
||||
from typing import Any, Awaitable, Dict, List, Optional, Tuple, Union
|
||||
|
||||
import fastapi
|
||||
import uvloop
|
||||
@@ -30,6 +30,7 @@ import zmq
|
||||
import zmq.asyncio
|
||||
from fastapi import BackgroundTasks
|
||||
|
||||
from sglang.srt.aio_rwlock import RWLock
|
||||
from sglang.srt.configs.model_config import ModelConfig
|
||||
from sglang.srt.hf_transformers_utils import get_processor, get_tokenizer
|
||||
from sglang.srt.managers.image_processor import (
|
||||
@@ -62,7 +63,11 @@ from sglang.srt.managers.io_struct import (
|
||||
from sglang.srt.metrics.collector import TokenizerMetricsCollector
|
||||
from sglang.srt.sampling.sampling_params import SamplingParams
|
||||
from sglang.srt.server_args import PortArgs, ServerArgs
|
||||
from sglang.srt.utils import get_zmq_socket, kill_process_tree
|
||||
from sglang.srt.utils import (
|
||||
dataclass_to_string_truncated,
|
||||
get_zmq_socket,
|
||||
kill_process_tree,
|
||||
)
|
||||
|
||||
asyncio.set_event_loop_policy(uvloop.EventLoopPolicy())
|
||||
|
||||
@@ -82,6 +87,9 @@ class ReqState:
|
||||
created_time: float
|
||||
first_token_time: Optional[float] = None
|
||||
|
||||
# For streaming output
|
||||
last_output_offset: int = 0
|
||||
|
||||
|
||||
class TokenizerManager:
|
||||
"""TokenizerManager is a process that tokenizes the text."""
|
||||
@@ -120,6 +128,7 @@ class TokenizerManager:
|
||||
|
||||
self.is_generation = self.model_config.is_generation
|
||||
self.context_len = self.model_config.context_len
|
||||
self.image_token_id = self.model_config.image_token_id
|
||||
|
||||
# Create image processor placeholder
|
||||
self.image_processor = get_dummy_image_processor()
|
||||
@@ -152,9 +161,12 @@ class TokenizerManager:
|
||||
self.to_create_loop = True
|
||||
self.rid_to_state: Dict[str, ReqState] = {}
|
||||
|
||||
# For update model weights
|
||||
self.model_update_lock = asyncio.Lock()
|
||||
self.model_update_result = None
|
||||
# The event to notify the weight sync is finished.
|
||||
self.model_update_lock = RWLock()
|
||||
self.model_update_result: Optional[Awaitable[UpdateWeightFromDiskReqOutput]] = (
|
||||
None
|
||||
)
|
||||
self.asyncio_tasks = set()
|
||||
|
||||
# For session info
|
||||
self.session_futures = {} # session_id -> asyncio event
|
||||
@@ -181,9 +193,6 @@ class TokenizerManager:
|
||||
if self.to_create_loop:
|
||||
self.create_handle_loop()
|
||||
|
||||
while self.model_update_lock.locked():
|
||||
await asyncio.sleep(0.001)
|
||||
|
||||
if isinstance(obj, EmbeddingReqInput) and self.is_generation:
|
||||
raise ValueError(
|
||||
"This model does not appear to be an embedding model by default. "
|
||||
@@ -191,17 +200,24 @@ class TokenizerManager:
|
||||
)
|
||||
|
||||
obj.normalize_batch_and_arguments()
|
||||
is_single = obj.is_single
|
||||
if is_single:
|
||||
tokenized_obj = await self._tokenize_one_request(obj)
|
||||
self.send_to_scheduler.send_pyobj(tokenized_obj)
|
||||
async for response in self._wait_one_response(obj, request, created_time):
|
||||
yield response
|
||||
else:
|
||||
async for response in self._handle_batch_request(
|
||||
obj, request, created_time
|
||||
):
|
||||
yield response
|
||||
|
||||
if self.server_args.log_requests:
|
||||
logger.info(f"Receive: obj={dataclass_to_string_truncated(obj)}")
|
||||
|
||||
async with self.model_update_lock.reader_lock:
|
||||
is_single = obj.is_single
|
||||
if is_single:
|
||||
tokenized_obj = await self._tokenize_one_request(obj)
|
||||
self.send_to_scheduler.send_pyobj(tokenized_obj)
|
||||
async for response in self._wait_one_response(
|
||||
obj, request, created_time
|
||||
):
|
||||
yield response
|
||||
else:
|
||||
async for response in self._handle_batch_request(
|
||||
obj, request, created_time
|
||||
):
|
||||
yield response
|
||||
|
||||
async def _tokenize_one_request(
|
||||
self,
|
||||
@@ -215,7 +231,7 @@ class TokenizerManager:
|
||||
if not self.server_args.disable_radix_cache:
|
||||
raise ValueError(
|
||||
"input_embeds is provided while disable_radix_cache is False. "
|
||||
"Please add `--disable-radix-cach` when you launch the server "
|
||||
"Please add `--disable-radix-cache` when you launch the server "
|
||||
"if you want to use input_embeds as inputs."
|
||||
)
|
||||
input_embeds = obj.input_embeds
|
||||
@@ -301,8 +317,8 @@ class TokenizerManager:
|
||||
state.out_list = []
|
||||
if state.finished:
|
||||
if self.server_args.log_requests:
|
||||
# Log requests
|
||||
logger.info(f"in={obj}, out={out}")
|
||||
msg = f"Finish: obj={dataclass_to_string_truncated(obj)}, out={dataclass_to_string_truncated(out)}"
|
||||
logger.info(msg)
|
||||
del self.rid_to_state[obj.rid]
|
||||
yield out
|
||||
break
|
||||
@@ -423,55 +439,52 @@ class TokenizerManager:
|
||||
self,
|
||||
obj: UpdateWeightFromDiskReqInput,
|
||||
request: Optional[fastapi.Request] = None,
|
||||
):
|
||||
) -> Tuple[bool, str]:
|
||||
if self.to_create_loop:
|
||||
self.create_handle_loop()
|
||||
|
||||
# default the load format to the server_args
|
||||
if obj.load_format is None:
|
||||
obj.load_format = self.server_args.load_format
|
||||
logger.info("Start update_weights. Load format=%s", obj.load_format)
|
||||
|
||||
if not self.model_update_lock.locked():
|
||||
if True:
|
||||
# Hold the lock if it is not async. This means that weight sync
|
||||
# cannot run while requests are in progress.
|
||||
async with self.model_update_lock.writer_lock:
|
||||
return await self._wait_for_model_update_from_disk(obj)
|
||||
|
||||
async with self.model_update_lock:
|
||||
# wait for the previous generation requests to finish
|
||||
for i in range(3):
|
||||
while len(self.rid_to_state) > 0:
|
||||
await asyncio.sleep(0.001)
|
||||
# FIXME: We add some sleep here to avoid some race conditions.
|
||||
# We can use a read-write lock as a better fix.
|
||||
await asyncio.sleep(0.01)
|
||||
self.send_to_scheduler.send_pyobj(obj)
|
||||
self.model_update_result = asyncio.Future()
|
||||
async def _wait_for_model_update_from_disk(
|
||||
self, obj: UpdateWeightFromDiskReqInput
|
||||
) -> Tuple[bool, str, int]:
|
||||
self.send_to_scheduler.send_pyobj(obj)
|
||||
self.model_update_result = asyncio.Future()
|
||||
if self.server_args.dp_size == 1:
|
||||
result = await self.model_update_result
|
||||
if result.success:
|
||||
self.served_model_name = obj.model_path
|
||||
self.server_args.model_path = obj.model_path
|
||||
self.server_args.load_format = obj.load_format
|
||||
self.model_path = obj.model_path
|
||||
return result.success, result.message
|
||||
else: # self.server_args.dp_size > 1
|
||||
self.model_update_tmp = []
|
||||
result = await self.model_update_result
|
||||
|
||||
if self.server_args.dp_size == 1:
|
||||
result = await self.model_update_result
|
||||
if result.success:
|
||||
self.server_args.model_path = obj.model_path
|
||||
self.server_args.load_format = obj.load_format
|
||||
self.model_path = obj.model_path
|
||||
return result.success, result.message
|
||||
else: # self.server_args.dp_size > 1
|
||||
self.model_update_tmp = []
|
||||
result = await self.model_update_result
|
||||
|
||||
all_success = all([r.success for r in result])
|
||||
if all_success is True:
|
||||
self.server_args.model_path = obj.model_path
|
||||
self.server_args.load_format = obj.load_format
|
||||
self.model_path = obj.model_path
|
||||
all_message = [r.message for r in result]
|
||||
all_message = " | ".join(all_message)
|
||||
return all_success, all_message
|
||||
|
||||
else:
|
||||
return False, "Another update is in progress. Please try again later."
|
||||
all_success = all([r.success for r in result])
|
||||
if all_success is True:
|
||||
self.server_args.model_path = obj.model_path
|
||||
self.server_args.load_format = obj.load_format
|
||||
self.model_path = obj.model_path
|
||||
all_message = [r.message for r in result]
|
||||
all_message = " | ".join(all_message)
|
||||
return all_success, all_message
|
||||
|
||||
async def init_weights_update_group(
|
||||
self,
|
||||
obj: InitWeightsUpdateGroupReqInput,
|
||||
request: Optional[fastapi.Request] = None,
|
||||
) -> bool:
|
||||
) -> Tuple[bool, str]:
|
||||
if self.to_create_loop:
|
||||
self.create_handle_loop()
|
||||
self.send_to_scheduler.send_pyobj(obj)
|
||||
@@ -487,25 +500,22 @@ class TokenizerManager:
|
||||
self,
|
||||
obj: UpdateWeightsFromDistributedReqInput,
|
||||
request: Optional[fastapi.Request] = None,
|
||||
):
|
||||
) -> Tuple[bool, str]:
|
||||
if self.to_create_loop:
|
||||
self.create_handle_loop()
|
||||
|
||||
if not self.model_update_lock.locked():
|
||||
async with self.model_update_lock:
|
||||
self.send_to_scheduler.send_pyobj(obj)
|
||||
self.parameter_update_result = asyncio.Future()
|
||||
assert (
|
||||
self.server_args.dp_size == 1
|
||||
), "dp_size must be for update weights from distributed"
|
||||
result = await self.parameter_update_result
|
||||
return result.success, result.message
|
||||
else:
|
||||
logger.error("Another parameter update is in progress in tokenizer manager")
|
||||
return (
|
||||
False,
|
||||
"Another parameter update is in progress. Please try again later.",
|
||||
)
|
||||
# This means that weight sync
|
||||
# cannot run while requests are in progress.
|
||||
async with self.model_update_lock.writer_lock:
|
||||
self.send_to_scheduler.send_pyobj(obj)
|
||||
self.parameter_update_result: Awaitable[
|
||||
UpdateWeightsFromDistributedReqOutput
|
||||
] = asyncio.Future()
|
||||
assert (
|
||||
self.server_args.dp_size == 1
|
||||
), "dp_size must be for update weights from distributed"
|
||||
result = await self.parameter_update_result
|
||||
return result.success, result.message
|
||||
|
||||
async def get_weights_by_name(
|
||||
self, obj: GetWeightsByNameReqInput, request: Optional[fastapi.Request] = None
|
||||
@@ -564,11 +574,11 @@ class TokenizerManager:
|
||||
|
||||
self.to_create_loop = False
|
||||
loop = asyncio.get_event_loop()
|
||||
loop.create_task(self.handle_loop())
|
||||
self.asyncio_tasks.add(loop.create_task(self.handle_loop()))
|
||||
|
||||
signal_handler = SignalHandler(self)
|
||||
loop.add_signal_handler(signal.SIGTERM, signal_handler.signal_handler)
|
||||
loop.create_task(self.sigterm_watchdog())
|
||||
self.asyncio_tasks.add(loop.create_task(self.sigterm_watchdog()))
|
||||
|
||||
async def sigterm_watchdog(self):
|
||||
while not self.gracefully_exit:
|
||||
|
||||
Reference in New Issue
Block a user