[Minor] Many cleanup (#1357)

This commit is contained in:
Lianmin Zheng
2024-09-09 04:14:11 -07:00
committed by GitHub
parent c9b75917d5
commit e4d68afcf0
24 changed files with 416 additions and 296 deletions
@@ -71,12 +71,10 @@ class ControllerMulti:
self,
server_args: ServerArgs,
port_args: PortArgs,
model_override_args,
):
# Parse args
self.server_args = server_args
self.port_args = port_args
self.model_override_args = model_override_args
self.load_balance_method = LoadBalanceMethod.from_str(
server_args.load_balance_method
)
@@ -114,7 +112,6 @@ class ControllerMulti:
self.server_args,
self.port_args,
pipe_controller_writer,
self.model_override_args,
True,
gpu_ids,
dp_worker_id,
@@ -189,14 +186,13 @@ def start_controller_process(
server_args: ServerArgs,
port_args: PortArgs,
pipe_writer,
model_override_args: dict,
):
"""Start a controller process."""
configure_logger(server_args)
try:
controller = ControllerMulti(server_args, port_args, model_override_args)
controller = ControllerMulti(server_args, port_args)
except Exception:
pipe_writer.send(get_exception_traceback())
raise
@@ -40,7 +40,6 @@ class ControllerSingle:
self,
server_args: ServerArgs,
port_args: PortArgs,
model_override_args: dict,
gpu_ids: List[int],
is_data_parallel_worker: bool,
dp_worker_id: int,
@@ -76,7 +75,6 @@ class ControllerSingle:
tp_rank_range,
server_args,
port_args.nccl_ports[dp_worker_id],
model_override_args,
)
# Launch tp rank 0
@@ -85,7 +83,6 @@ class ControllerSingle:
0,
server_args,
port_args.nccl_ports[dp_worker_id],
model_override_args,
)
self.tp_cpu_group = self.tp_server.model_runner.tp_group.cpu_group
@@ -126,7 +123,6 @@ def start_controller_process(
server_args: ServerArgs,
port_args: PortArgs,
pipe_writer: multiprocessing.connection.Connection,
model_override_args: dict,
is_data_parallel_worker: bool = False,
gpu_ids: List[int] = None,
dp_worker_id: int = None,
@@ -149,7 +145,6 @@ def start_controller_process(
controller = ControllerSingle(
server_args,
port_args,
model_override_args,
gpu_ids,
is_data_parallel_worker,
dp_worker_id,
@@ -18,6 +18,7 @@ limitations under the License.
import asyncio
import concurrent.futures
import dataclasses
import json
import logging
import multiprocessing as mp
import os
@@ -77,7 +78,6 @@ class TokenizerManager:
self,
server_args: ServerArgs,
port_args: PortArgs,
model_override_args: dict = None,
):
self.server_args = server_args
@@ -95,7 +95,7 @@ class TokenizerManager:
self.hf_config = get_config(
self.model_path,
trust_remote_code=server_args.trust_remote_code,
model_override_args=model_override_args,
model_override_args=json.loads(server_args.json_model_override_args),
)
self.is_generation = is_generation_model(
self.hf_config.architectures, self.server_args.is_embedding
+79 -76
View File
@@ -15,13 +15,14 @@ limitations under the License.
"""A tensor parallel worker."""
import json
import logging
import multiprocessing
import os
import pickle
import time
import warnings
from typing import Any, List, Optional, Union
from typing import Any, List, Optional
import torch
import torch.distributed
@@ -66,6 +67,7 @@ from sglang.utils import get_exception_traceback
logger = logging.getLogger(__name__)
# Crash on warning if we are running CI tests
crash_on_warning = os.getenv("SGLANG_IS_IN_CI", "false") == "true"
@@ -76,11 +78,10 @@ class ModelTpServer:
tp_rank: int,
server_args: ServerArgs,
nccl_port: int,
model_override_args: dict,
):
suppress_other_loggers()
# Copy arguments
# Parse arguments
self.gpu_id = gpu_id
self.tp_rank = tp_rank
self.tp_size = server_args.tp_size
@@ -93,9 +94,8 @@ class ModelTpServer:
server_args.model_path,
server_args.trust_remote_code,
context_length=server_args.context_length,
model_override_args=model_override_args,
model_override_args=json.loads(server_args.json_model_override_args),
)
self.model_runner = ModelRunner(
model_config=self.model_config,
mem_fraction_static=server_args.mem_fraction_static,
@@ -136,7 +136,7 @@ class ModelTpServer:
self.max_total_num_tokens - 1,
)
# Sync random seed
# Sync random seed across TP workers
server_args.random_seed = broadcast_recv_input(
[server_args.random_seed],
self.tp_rank,
@@ -144,7 +144,7 @@ class ModelTpServer:
)[0]
set_random_seed(server_args.random_seed)
# Print info
# Print debug info
logger.info(
f"max_total_num_tokens={self.max_total_num_tokens}, "
f"max_prefill_tokens={self.max_prefill_tokens}, "
@@ -181,7 +181,7 @@ class ModelTpServer:
self.num_generated_tokens = 0
self.last_stats_tic = time.time()
# Chunked prefill
# Init chunked prefill
self.chunked_prefill_size = server_args.chunked_prefill_size
self.current_inflight_req = None
self.is_mixed_chunk = (
@@ -197,16 +197,6 @@ class ModelTpServer:
"trust_remote_code": server_args.trust_remote_code,
},
skip_tokenizer_init=server_args.skip_tokenizer_init,
json_schema_mode=False,
)
self.json_fsm_cache = FSMCache(
server_args.tokenizer_path,
{
"tokenizer_mode": server_args.tokenizer_mode,
"trust_remote_code": server_args.trust_remote_code,
},
skip_tokenizer_init=server_args.skip_tokenizer_init,
json_schema_mode=True,
)
self.jump_forward_cache = JumpForwardCache()
@@ -227,11 +217,12 @@ class ModelTpServer:
try:
# Recv requests
for recv_req in recv_reqs:
if isinstance(
recv_req, (TokenizedGenerateReqInput, TokenizedEmbeddingReqInput)
):
if isinstance(recv_req, TokenizedGenerateReqInput):
self.handle_generate_request(recv_req)
self.do_not_get_new_batch = False
elif isinstance(recv_req, TokenizedEmbeddingReqInput):
self.handle_embedding_request(recv_req)
self.do_not_get_new_batch = False
elif isinstance(recv_req, FlushCacheReq):
self.flush_cache()
elif isinstance(recv_req, AbortReq):
@@ -331,57 +322,56 @@ class ModelTpServer:
def handle_generate_request(
self,
recv_req: Union[TokenizedGenerateReqInput, TokenizedEmbeddingReqInput],
recv_req: TokenizedGenerateReqInput,
):
req = Req(recv_req.rid, recv_req.input_text, recv_req.input_ids)
req.tokenizer = self.tokenizer
req.sampling_params = recv_req.sampling_params
if self.model_runner.is_generation:
req.pixel_values = recv_req.pixel_values
if req.pixel_values is not None:
# Use image hash as fake token_ids, which is then used
# for prefix matching
image_hash = hash(tuple(recv_req.image_hashes))
req.pad_value = [
(image_hash) % self.model_config.vocab_size,
(image_hash >> 16) % self.model_config.vocab_size,
(image_hash >> 32) % self.model_config.vocab_size,
(image_hash >> 64) % self.model_config.vocab_size,
]
req.image_sizes = recv_req.image_sizes
(
req.origin_input_ids,
req.image_offsets,
) = self.model_runner.model.pad_input_ids(
req.origin_input_ids_unpadded,
req.pad_value,
req.pixel_values,
req.image_sizes,
)
# Only when pixel values is not None we have modalities
req.modalities = recv_req.modalites
req.return_logprob = recv_req.return_logprob
req.logprob_start_len = recv_req.logprob_start_len
req.top_logprobs_num = recv_req.top_logprobs_num
req.stream = recv_req.stream
req.pixel_values = recv_req.pixel_values
if req.pixel_values is not None:
# Use image hash as fake token_ids, which is then used
# for prefix matching
image_hash = hash(tuple(recv_req.image_hashes))
req.pad_value = [
(image_hash) % self.model_config.vocab_size,
(image_hash >> 16) % self.model_config.vocab_size,
(image_hash >> 32) % self.model_config.vocab_size,
(image_hash >> 64) % self.model_config.vocab_size,
]
req.image_sizes = recv_req.image_sizes
(
req.origin_input_ids,
req.image_offsets,
) = self.model_runner.model.pad_input_ids(
req.origin_input_ids_unpadded,
req.pad_value,
req.pixel_values,
req.image_sizes,
)
# Only when pixel values is not None we have modalities
req.modalities = recv_req.modalites
req.return_logprob = recv_req.return_logprob
req.logprob_start_len = recv_req.logprob_start_len
req.top_logprobs_num = recv_req.top_logprobs_num
req.stream = recv_req.stream
# Init regex fsm fron json
# Init regex FSM
if (
req.sampling_params.json_schema is not None
or req.sampling_params.regex is not None
):
if req.sampling_params.json_schema is not None:
req.regex_fsm, computed_regex_string = self.json_fsm_cache.query(
req.sampling_params.json_schema
req.regex_fsm, computed_regex_string = self.regex_fsm_cache.query(
("json", req.sampling_params.json_schema)
)
if not self.disable_regex_jump_forward:
req.jump_forward_map = self.jump_forward_cache.query(
computed_regex_string
)
# Init regex fsm
elif req.sampling_params.regex is not None:
req.regex_fsm = self.regex_fsm_cache.query(req.sampling_params.regex)
if not self.disable_regex_jump_forward:
req.jump_forward_map = self.jump_forward_cache.query(
req.sampling_params.regex
)
req.regex_fsm, computed_regex_string = self.regex_fsm_cache.query(
("regex", req.sampling_params.regex)
)
if not self.disable_regex_jump_forward:
req.jump_forward_map = self.jump_forward_cache.query(
computed_regex_string
)
# Truncate prompts that are too long
if len(req.origin_input_ids) >= self.max_req_input_len:
@@ -390,16 +380,32 @@ class ModelTpServer:
"the max context length. Truncated!!!"
)
req.origin_input_ids = req.origin_input_ids[: self.max_req_input_len]
req.sampling_params.max_new_tokens = min(
(
req.sampling_params.max_new_tokens
if req.sampling_params.max_new_tokens is not None
else 1 << 30
),
self.max_req_input_len - 1 - len(req.origin_input_ids),
)
if self.model_runner.is_generation:
req.sampling_params.max_new_tokens = min(
(
req.sampling_params.max_new_tokens
if req.sampling_params.max_new_tokens is not None
else 1 << 30
),
self.max_req_input_len - 1 - len(req.origin_input_ids),
self.waiting_queue.append(req)
def handle_embedding_request(
self,
recv_req: TokenizedEmbeddingReqInput,
):
req = Req(recv_req.rid, recv_req.input_text, recv_req.input_ids)
req.tokenizer = self.tokenizer
req.sampling_params = recv_req.sampling_params
# Truncate prompts that are too long
if len(req.origin_input_ids) >= self.max_req_input_len:
logger.warn(
"Request length is longer than the KV cache pool size or "
"the max context length. Truncated!!!"
)
req.origin_input_ids = req.origin_input_ids[: self.max_req_input_len]
self.waiting_queue.append(req)
@@ -892,7 +898,6 @@ def run_tp_server(
tp_rank: int,
server_args: ServerArgs,
nccl_port: int,
model_override_args: dict,
):
"""Run a tensor parallel model server."""
configure_logger(server_args, prefix=f" TP{tp_rank}")
@@ -903,7 +908,6 @@ def run_tp_server(
tp_rank,
server_args,
nccl_port,
model_override_args,
)
tp_cpu_group = model_server.model_runner.tp_group.cpu_group
@@ -920,14 +924,13 @@ def launch_tp_servers(
tp_rank_range: List[int],
server_args: ServerArgs,
nccl_port: int,
model_override_args: dict,
):
"""Launch multiple tensor parallel servers."""
procs = []
for i in tp_rank_range:
proc = multiprocessing.Process(
target=run_tp_server,
args=(gpu_ids[i], i, server_args, nccl_port, model_override_args),
args=(gpu_ids[i], i, server_args, nccl_port),
)
proc.start()
procs.append(proc)