Improve type annotation (#1029)
This commit is contained in:
@@ -21,15 +21,17 @@ import os
|
||||
import pickle
|
||||
import time
|
||||
import warnings
|
||||
from typing import List, Optional, Union
|
||||
from typing import Any, List, Optional, Union
|
||||
|
||||
import torch
|
||||
import torch.distributed
|
||||
import torch.distributed as dist
|
||||
|
||||
from sglang.global_config import global_config
|
||||
from sglang.srt.constrained.fsm_cache import FSMCache
|
||||
from sglang.srt.constrained.jump_forward import JumpForwardCache
|
||||
from sglang.srt.hf_transformers_utils import get_processor, get_tokenizer
|
||||
from sglang.srt.layers.logits_processor import LogitProcessorOutput
|
||||
from sglang.srt.managers.io_struct import (
|
||||
AbortReq,
|
||||
BatchEmbeddingOut,
|
||||
@@ -62,6 +64,10 @@ from sglang.utils import get_exception_traceback
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
# TODO: Rename "CI" to "SGLANG_IS_IN_CI".
|
||||
crash_on_warning = os.getenv("CI", "false") == "true"
|
||||
|
||||
|
||||
class ModelTpServer:
|
||||
def __init__(
|
||||
self,
|
||||
@@ -198,7 +204,7 @@ class ModelTpServer:
|
||||
self.new_token_ratio = self.min_new_token_ratio
|
||||
self.new_token_ratio_decay = global_config.new_token_ratio_decay
|
||||
|
||||
def exposed_step(self, recv_reqs):
|
||||
def exposed_step(self, recv_reqs: List):
|
||||
try:
|
||||
# Recv requests
|
||||
for recv_req in recv_reqs:
|
||||
@@ -247,7 +253,7 @@ class ModelTpServer:
|
||||
|
||||
# Print stats
|
||||
if self.tp_rank == 0 and self.decode_forward_ct % 40 == 0:
|
||||
self.print_stats()
|
||||
self.print_decode_stats()
|
||||
|
||||
if self.running_batch.is_empty():
|
||||
self.running_batch = None
|
||||
@@ -259,7 +265,7 @@ class ModelTpServer:
|
||||
self.check_memory()
|
||||
self.new_token_ratio = global_config.init_new_token_ratio
|
||||
|
||||
def print_stats(self):
|
||||
def print_decode_stats(self):
|
||||
num_used = self.max_total_num_tokens - (
|
||||
self.token_to_kv_pool.available_size() + self.tree_cache.evictable_size()
|
||||
)
|
||||
@@ -276,7 +282,6 @@ class ModelTpServer:
|
||||
)
|
||||
|
||||
def check_memory(self):
|
||||
crash = os.getenv("CI", "false") == "true"
|
||||
available_size = (
|
||||
self.token_to_kv_pool.available_size() + self.tree_cache.evictable_size()
|
||||
)
|
||||
@@ -286,7 +291,7 @@ class ModelTpServer:
|
||||
f"available_size={available_size}, max_total_num_tokens={self.max_total_num_tokens}\n"
|
||||
"KV cache pool leak detected!"
|
||||
)
|
||||
exit(1) if crash else None
|
||||
exit(1) if crash_on_warning else None
|
||||
|
||||
if len(self.req_to_token_pool.free_slots) != self.req_to_token_pool.size:
|
||||
warnings.warn(
|
||||
@@ -295,7 +300,7 @@ class ModelTpServer:
|
||||
f"total slots={self.req_to_token_pool.size}\n"
|
||||
"Memory pool leak detected!"
|
||||
)
|
||||
exit(1) if crash else None
|
||||
exit(1) if crash_on_warning else None
|
||||
|
||||
def handle_generate_request(
|
||||
self,
|
||||
@@ -511,7 +516,14 @@ class ModelTpServer:
|
||||
|
||||
self.handle_finished_requests(batch)
|
||||
|
||||
def add_logprob_return_values(self, i, req: Req, pt, next_token_ids, output):
|
||||
def add_logprob_return_values(
|
||||
self,
|
||||
i,
|
||||
req: Req,
|
||||
pt: int,
|
||||
next_token_ids: List[int],
|
||||
output: LogitProcessorOutput,
|
||||
):
|
||||
if req.normalized_prompt_logprob is None:
|
||||
req.normalized_prompt_logprob = output.normalized_prompt_logprobs[i]
|
||||
|
||||
@@ -786,7 +798,11 @@ def run_tp_server(
|
||||
|
||||
|
||||
def launch_tp_servers(
|
||||
gpu_ids, tp_rank_range, server_args, nccl_port, model_overide_args
|
||||
gpu_ids: List[int],
|
||||
tp_rank_range: List[int],
|
||||
server_args: ServerArgs,
|
||||
nccl_port: int,
|
||||
model_overide_args: dict,
|
||||
):
|
||||
"""Launch multiple tensor parallel servers."""
|
||||
procs = []
|
||||
@@ -801,7 +817,9 @@ def launch_tp_servers(
|
||||
return procs
|
||||
|
||||
|
||||
def broadcast_recv_input(data, rank, dist_group):
|
||||
def broadcast_recv_input(
|
||||
data: Any, rank: int, dist_group: torch.distributed.ProcessGroup
|
||||
):
|
||||
"""Broadcast inputs from rank=0 to all other ranks with torch.dist backend."""
|
||||
|
||||
if rank == 0:
|
||||
|
||||
Reference in New Issue
Block a user