Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
28 changes: 18 additions & 10 deletions python/sglang/backend/runtime_endpoint.py
Original file line number Diff line number Diff line change
Expand Up @@ -34,7 +34,7 @@ def __init__(
api_key=self.api_key,
verify=self.verify,
)
assert res.status_code == 200
self._assert_success(res)
self.model_info = res.json()

self.chat_template = get_chat_template_by_model_path(
Expand All @@ -50,14 +50,15 @@ def flush_cache(self):
auth_token=self.auth_token,
verify=self.verify,
)
return res.status_code == 200
self._assert_success(res)

def get_server_args(self):
res = http_request(
self.base_url + "/get_server_args",
auth_token=self.auth_token,
verify=self.verify,
)
self._assert_success(res)
return res.json()

def get_chat_template(self):
Expand All @@ -71,7 +72,7 @@ def cache_prefix(self, prefix_str: str):
api_key=self.api_key,
verify=self.verify,
)
assert res.status_code == 200
self._assert_success(res)

def commit_lazy_operations(self, s: StreamExecutor):
data = {"text": s.text_, "sampling_params": {"max_new_tokens": 0}}
Expand All @@ -83,7 +84,7 @@ def commit_lazy_operations(self, s: StreamExecutor):
api_key=self.api_key,
verify=self.verify,
)
assert res.status_code == 200
self._assert_success(res)

def fill_image(self, s: StreamExecutor):
data = {"text": s.text_, "sampling_params": {"max_new_tokens": 0}}
Expand All @@ -95,7 +96,7 @@ def fill_image(self, s: StreamExecutor):
api_key=self.api_key,
verify=self.verify,
)
assert res.status_code == 200
self._assert_success(res)

def generate(
self,
Expand Down Expand Up @@ -133,6 +134,8 @@ def generate(
api_key=self.api_key,
verify=self.verify,
)
self._assert_success(res)

obj = res.json()
comp = obj["text"]
return comp, obj["meta_info"]
Expand Down Expand Up @@ -167,18 +170,19 @@ def generate_stream(
data["stream"] = True
self._add_images(s, data)

response = http_request(
res = http_request(
self.base_url + "/generate",
json=data,
stream=True,
auth_token=self.auth_token,
api_key=self.api_key,
verify=self.verify,
)
self._assert_success(res)
pos = 0

incomplete_text = ""
for chunk in response.iter_lines(decode_unicode=False):
for chunk in res.iter_lines(decode_unicode=False):
chunk = chunk.decode("utf-8")
if chunk and chunk.startswith("data:"):
if chunk == "data: [DONE]":
Expand Down Expand Up @@ -211,7 +215,7 @@ def select(
api_key=self.api_key,
verify=self.verify,
)
assert res.status_code == 200
self._assert_success(res)
prompt_len = res.json()["meta_info"]["prompt_tokens"]

# Compute logprob
Expand All @@ -229,7 +233,7 @@ def select(
api_key=self.api_key,
verify=self.verify,
)
assert res.status_code == 200
self._assert_success(res)
obj = res.json()
normalized_prompt_logprobs = [
r["meta_info"]["normalized_prompt_logprob"] for r in obj
Expand All @@ -253,9 +257,13 @@ def concatenate_and_append(self, src_rids: List[str], dst_rid: str):
api_key=self.api_key,
verify=self.verify,
)
assert res.status_code == 200
self._assert_success(res)

def _add_images(self, s: StreamExecutor, data):
if s.images_:
assert len(s.images_) == 1, "Only support one image."
data["image_data"] = s.images_[0][1]

def _assert_success(self, res):
if res.status_code != 200:
raise RuntimeError(res.json())
10 changes: 7 additions & 3 deletions python/sglang/lang/interpreter.py
Original file line number Diff line number Diff line change
Expand Up @@ -191,7 +191,7 @@ def __init__(
self.variable_event = {} # Dict[name: str -> event: threading.Event]
self.meta_info = {} # Dict[name: str -> info: str]
self.is_finished = False
self.error = None
self.error_ = None

# For completion
self.text_ = "" # The full text
Expand Down Expand Up @@ -300,6 +300,10 @@ def messages(self):
self.sync()
return self.messages_

def error(self):
self.sync()
return self.error_

def end(self):
if self.use_thread:
if self.worker.is_alive():
Expand Down Expand Up @@ -338,7 +342,7 @@ def _thread_worker_func(self):
if self.stream_var_event:
for name in self.stream_var_event:
self.stream_var_event[name].set()
self.error = error
self.error_ = error

if self.stream_text_event:
self.stream_text_event.set()
Expand Down Expand Up @@ -713,7 +717,7 @@ def sync(self):
return self.stream_executor.sync()

def error(self):
return self.stream_executor.error
return self.stream_executor.error()

def text_iter(self, var_name: Optional[str] = None):
if self.stream_executor.stream:
Expand Down
17 changes: 10 additions & 7 deletions python/sglang/srt/managers/io_struct.py
Original file line number Diff line number Diff line change
Expand Up @@ -31,12 +31,9 @@ class GenerateReqInput:

def post_init(self):

if self.text is None:
assert (
self.input_ids is not None
), "Either text or input_ids should be provided"
else:
assert self.input_ids is None, "Either text or input_ids should be provided"
if ((self.text is None and self.input_ids is None) or
(self.text is not None and self.input_ids is not None)):
raise ValueError("Either text or input_ids should be provided.")

if self.text is not None:
is_single = isinstance(self.text, str)
Expand Down Expand Up @@ -71,7 +68,8 @@ def post_init(self):
if self.rid is None:
self.rid = [uuid.uuid4().hex for _ in range(num)]
else:
assert isinstance(self.rid, list)
if not isinstance(self.rid, list):
raise ValueError("The rid should be a list.")

if self.return_logprob is None:
self.return_logprob = [False] * num
Expand Down Expand Up @@ -129,6 +127,11 @@ class FlushCacheReq:
pass


@dataclass
class AbortReq:
rid: str


@dataclass
class DetokenizeReqInput:
input_ids: List[int]
80 changes: 49 additions & 31 deletions python/sglang/srt/managers/router/model_rpc.py
Original file line number Diff line number Diff line change
Expand Up @@ -20,6 +20,7 @@
from sglang.srt.constrained.jump_forward import JumpForwardCache
from sglang.srt.hf_transformers_utils import get_processor, get_tokenizer
from sglang.srt.managers.io_struct import (
AbortReq,
BatchTokenIDOut,
FlushCacheReq,
TokenizedGenerateReqInput,
Expand Down Expand Up @@ -110,6 +111,8 @@ def __init__(
get_int_token_logit_bias(self.tokenizer, self.model_config.vocab_size)
)
set_random_seed(server_args.random_seed)

# Print info
logger.info(
f"Rank {self.tp_rank}: "
f"max_total_num_token={self.max_total_num_token}, "
Expand Down Expand Up @@ -160,24 +163,6 @@ def __init__(
self.min_new_token_ratio = min(0.2 * server_args.schedule_conservativeness, 1.0)
self.new_token_ratio_step = (0.0001, 0.05) # (down, up)

def flush_cache(self):
if len(self.forward_queue) == 0 and (
self.running_batch is None or len(self.running_batch.reqs) == 0
):
self.tree_cache.reset()
self.tree_cache_metrics = {"total": 0, "hit": 0}
self.regex_fsm_cache.reset()
self.req_to_token_pool.clear()
self.token_to_kv_pool.clear()
torch.cuda.empty_cache()
logger.info("Cache flushed successfully!")
else:
warnings.warn(
f"Cache not flushed because there are pending requests. "
f"#queue-req: {len(self.forward_queue)}, "
f"#running-req: {0 if self.running_batch is None else len(self.running_batch.reqs)}"
)

def exposed_step(self, recv_reqs):
if self.tp_size != 1:
recv_reqs = obtain(recv_reqs)
Expand All @@ -189,6 +174,8 @@ def exposed_step(self, recv_reqs):
self.handle_generate_request(recv_req)
elif isinstance(recv_req, FlushCacheReq):
self.flush_cache()
elif isinstance(recv_req, AbortReq):
self.abort_request(recv_req)
else:
raise ValueError(f"Invalid request: {recv_req}")

Expand All @@ -207,9 +194,8 @@ def forward_step(self):
new_batch = self.get_new_fill_batch()

if new_batch is not None:
# Run new fill batch
# Run a new fill batch
self.forward_fill_batch(new_batch)

self.cache_filled_batch(new_batch)

if not new_batch.is_empty():
Expand All @@ -225,14 +211,8 @@ def forward_step(self):
self.num_generated_tokens += len(self.running_batch.reqs)
self.forward_decode_batch(self.running_batch)

if self.running_batch.is_empty():
self.running_batch = None
break

if self.out_pyobjs and self.running_batch.reqs[0].stream:
break

if self.running_batch is not None and self.tp_rank == 0:
# Print stats
if self.tp_rank == 0:
if self.decode_forward_ct % 40 == 0:
num_used = self.max_total_num_token - (
self.token_to_kv_pool.available_size()
Expand All @@ -250,8 +230,15 @@ def forward_step(self):
f"gen throughput (token/s): {throuhgput:.2f}, "
f"#queue-req: {len(self.forward_queue)}"
)

if self.running_batch.is_empty():
self.running_batch = None
break

if self.out_pyobjs and self.running_batch.reqs[0].stream:
break
else:
# check the available size
# Check the available size
available_size = (
self.token_to_kv_pool.available_size()
+ self.tree_cache.evictable_size()
Expand Down Expand Up @@ -295,7 +282,7 @@ def handle_generate_request(
req.sampling_params.regex
)

# Truncate long prompts
# Truncate prompts that are too long
req.input_ids = req.input_ids[: self.model_config.context_len - 1]
req.sampling_params.max_new_tokens = min(
req.sampling_params.max_new_tokens,
Expand All @@ -311,6 +298,7 @@ def get_new_fill_batch(self):
):
return None

# Compute matched prefix length
for req in self.forward_queue:
prefix_indices, last_node = self.tree_cache.match_prefix(req.input_ids)
if req.return_logprob:
Expand Down Expand Up @@ -383,6 +371,7 @@ def get_new_fill_batch(self):
if len(can_run_list) == 0:
return None

# Print stats
if self.tp_rank == 0:
running_req = (
0 if self.running_batch is None else len(self.running_batch.reqs)
Expand Down Expand Up @@ -410,6 +399,7 @@ def get_new_fill_batch(self):
# f"ff_cache_avg_init_time: {self.jump_forward_cache.get_avg_init_time():.2f}s. "
# )

# Return the new batch
new_batch = Batch.init_new(
can_run_list,
self.req_to_token_pool,
Expand Down Expand Up @@ -487,7 +477,7 @@ def forward_fill_batch(self, batch: Batch):
self.handle_finished_requests(batch)

def cache_filled_batch(self, batch: Batch):
req_pool_indices_cpu = batch.req_pool_indices.cpu().tolist()
req_pool_indices_cpu = batch.req_pool_indices.cpu().numpy()
for i, req in enumerate(batch.reqs):
new_prefix_indices, new_last_node = self.tree_cache.cache_req(
token_ids=tuple(req.input_ids + req.output_ids)[:-1],
Expand Down Expand Up @@ -671,6 +661,34 @@ def handle_finished_requests(self, batch: Batch):
else:
batch.reqs = []

def flush_cache(self):
if len(self.forward_queue) == 0 and (
self.running_batch is None or len(self.running_batch.reqs) == 0
):
self.tree_cache.reset()
self.tree_cache_metrics = {"total": 0, "hit": 0}
self.regex_fsm_cache.reset()
self.req_to_token_pool.clear()
self.token_to_kv_pool.clear()
torch.cuda.empty_cache()
logger.info("Cache flushed successfully!")
else:
warnings.warn(
f"Cache not flushed because there are pending requests. "
f"#queue-req: {len(self.forward_queue)}, "
f"#running-req: {0 if self.running_batch is None else len(self.running_batch.reqs)}"
)

def abort_request(self, recv_req):
to_del = None
for i, req in enumerate(self.forward_queue):
if req.rid == recv_req.rid:
to_del = i
break

if to_del is not None:
del self.forward_queue[to_del]


class ModelRpcService(rpyc.Service):
exposed_ModelRpcServer = ModelRpcServer
Expand Down
Loading