Improve type annotation (#1029)

This commit is contained in:
Lianmin Zheng
2024-08-11 02:44:59 -07:00
committed by GitHub
parent fcc0f5ed99
commit 9dae407812
6 changed files with 83 additions and 41 deletions
+28 -10
View File
@@ -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: