Skip to content

Commit 68bb76e

Browse files
committed
Add batching and multithreaded server
1 parent 60c5782 commit 68bb76e

1 file changed

Lines changed: 218 additions & 43 deletions

File tree

mlx_lm/server.py

Lines changed: 218 additions & 43 deletions
Original file line numberDiff line numberDiff line change
@@ -11,7 +11,7 @@
1111
import warnings
1212
from collections import deque
1313
from dataclasses import dataclass, field
14-
from http.server import BaseHTTPRequestHandler, HTTPServer
14+
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
1515
from pathlib import Path
1616
from queue import Empty as QueueEmpty
1717
from queue import Queue
@@ -32,7 +32,7 @@
3232
from huggingface_hub import scan_cache_dir
3333

3434
from ._version import __version__
35-
from .generate import stream_generate
35+
from .generate import BatchGenerator, stream_generate
3636
from .models.cache import can_trim_prompt_cache, make_prompt_cache, trim_prompt_cache
3737
from .sample_utils import make_logits_processors, make_sampler
3838
from .utils import load
@@ -358,6 +358,15 @@ class GenerationContext:
358358
prompt: List[int]
359359

360360

361+
@dataclass
362+
class Response:
363+
text: str
364+
token: int
365+
logprob: float
366+
finish_reason: Optional[str]
367+
top_tokens: Optional[Tuple[int, float]]
368+
369+
361370
class ModelProvider:
362371
def __init__(self, cli_args: argparse.Namespace):
363372
"""Load models on demand and persist them across the whole process."""
@@ -447,44 +456,202 @@ def __init__(self, model_provider: ModelProvider, prompt_cache: LRUPromptCache):
447456
self._generation_thread = Thread(target=self._generate)
448457
self._generation_thread.start()
449458

450-
def _tokenize_chat(self, tokenizer, messages, tools=None, role_mapping=None):
451-
if tokenizer.chat_template:
452-
process_message_content(messages)
453-
prompt = tokenizer.apply_chat_template(
454-
messages,
455-
tools,
456-
add_generation_prompt=True,
457-
**self.model_provider.cli_args.chat_template_args,
458-
)
459+
def _tokenize(self, tokenizer, request):
460+
if request.request_type == "chat":
461+
messages = request.messages
462+
tools = request.tools
463+
role_mapping = request.role_mapping
464+
465+
if tokenizer.chat_template:
466+
process_message_content(messages)
467+
return tokenizer.apply_chat_template(
468+
messages,
469+
tools,
470+
add_generation_prompt=True,
471+
**self.model_provider.cli_args.chat_template_args,
472+
)
473+
else:
474+
return tokenizer.encode(convert_chat(messages, role_mapping))
459475
else:
460-
prompt = convert_chat(messages, role_mapping)
461-
prompt = tokenizer.encode(prompt)
476+
return tokenizer.encode(request.prompt)
462477

463-
return prompt
478+
def _is_batchable(self, args):
479+
if (
480+
args.model.draft != "default_model"
481+
or self.model_provider.cli_args.draft_model is not None
482+
):
483+
return False
484+
if args.logits.logit_bias is not None:
485+
return False
486+
if args.logits.repetition_penalty != 0:
487+
return False
488+
if args.logprobs > 0:
489+
return False
490+
if args.seed is not None:
491+
return False
492+
493+
return True
464494

465495
def _generate(self):
466496
current_model = None
467497
current_sampling = None
468-
current_logits_proc = None
498+
current_tokenizer = None
499+
current_model_key = None
500+
batch_generator = None
501+
drain_batch = False
502+
batch_results = {}
469503

470504
unprocessed_requests = []
471505

472506
def get_next_request():
473-
nonlocal unprocessed_requests
474-
475507
if unprocessed_requests:
476-
r, *unprocessed_requests = unprocessed_requests
477-
return r
508+
return unprocessed_requests.pop()
478509
else:
479510
try:
480511
return self.requests.get_nowait()
481512
except QueueEmpty:
482513
return None
483514

484515
while True:
485-
request = get_next_request()
516+
request = None
517+
if not drain_batch and len(batch_results) < 100:
518+
request = get_next_request()
519+
520+
# We got a request
486521
if request is not None:
487-
self._serve_single(request)
522+
rqueue, request, args = request
523+
524+
is_batchable = self._is_batchable(args)
525+
526+
# Can it be added to the current batch?
527+
if (
528+
batch_generator is not None
529+
and current_model == args.model
530+
and current_sampling == args.sampling
531+
and is_batchable
532+
):
533+
prompt = self._tokenize(current_tokenizer, request)
534+
ctx = GenerationContext(
535+
has_tool_calling=tokenizer.has_tool_calling,
536+
tool_call_start=tokenizer.tool_call_start,
537+
tool_call_end=tokenizer.tool_call_end,
538+
eos_token_id=tokenizer.eos_token_id,
539+
stop_token_sequences=[
540+
tokenizer.encode(stop_word, add_special_tokens=False)
541+
for stop_word in args.stop_words
542+
],
543+
prompt=prompt,
544+
)
545+
rqueue.put(ctx)
546+
547+
cache, rest = self.prompt_cache.fetch_nearest_cache(
548+
current_model_key, prompt
549+
)
550+
if cache is None:
551+
cache = make_prompt_cache(self.model_provider.model)
552+
553+
(uid,) = batch_generator.insert(
554+
[rest], args.max_tokens, caches=[cache]
555+
)
556+
batch_results[uid] = {
557+
"cache_key": prompt[:],
558+
"rqueue": rqueue,
559+
"detokenizer": tokenizer.detokenizer,
560+
}
561+
continue
562+
563+
# We have no batch and it actually is not a batchable request
564+
# so serve single sequence at a time.
565+
elif batch_generator is None and not is_batchable:
566+
self._serve_single((rqueue, request, args))
567+
continue
568+
569+
# No batch so make one and serve this batched
570+
elif batch_generator is None:
571+
try:
572+
model, tokenizer = self.model_provider.load(
573+
args.model.model, args.model.adapter, args.model.draft
574+
)
575+
except Exception as e:
576+
rqueue.put(e)
577+
continue
578+
579+
current_model = args.model
580+
current_sampling = args.sampling
581+
current_tokenizer = tokenizer
582+
current_model_key = self.model_provider.model_key
583+
batch_results = {}
584+
batch_generator = BatchGenerator(
585+
model,
586+
stop_tokens=tokenizer.eos_token_ids,
587+
sampler=make_sampler(
588+
args.sampling.temperature,
589+
top_p=args.sampling.top_p,
590+
top_k=args.sampling.top_k,
591+
min_p=args.sampling.min_p,
592+
xtc_probability=args.sampling.xtc_probability,
593+
xtc_threshold=args.sampling.xtc_threshold,
594+
xtc_special_tokens=[
595+
tokenizer.eos_token_id,
596+
tokenizer.encode("\n"),
597+
],
598+
),
599+
)
600+
unprocessed_requests.append((rqueue, request, args))
601+
continue
602+
603+
# We have a batch but this request cannot be added to the
604+
# batch so drain it to process the request.
605+
else:
606+
drain_batch = True
607+
unprocessed_requests.append((rqueue, request, args))
608+
continue
609+
610+
# No request so serve from the current batch
611+
elif batch_generator is not None:
612+
if len(batch_results) == 0:
613+
if drain_batch:
614+
current_model = None
615+
current_sampling = None
616+
current_tokenizer = None
617+
current_model_key = None
618+
batch_generator = None
619+
drain_batch = False
620+
continue
621+
622+
responses = batch_generator.next()
623+
for r in responses:
624+
result = batch_results[r.uid]
625+
result["cache_key"].append(r.token)
626+
result["detokenizer"].add_token(r.token)
627+
628+
top_tokens = None
629+
if args.logprobs > 0:
630+
sorted_indices = mx.argpartition(
631+
-gen.logprobs, kth=args.logprobs - 1
632+
)
633+
top_indices = sorted_indices[: args.logprobs]
634+
top_logprobs = gen.logprobs[top_indices]
635+
top_token_info = zip(
636+
top_indices.tolist(), top_logprobs.tolist()
637+
)
638+
top_tokens = tuple(top_token_info)
639+
result["rqueue"].put(
640+
Response(
641+
result["detokenizer"].last_segment,
642+
r.token,
643+
r.logprobs[r.token].item(),
644+
r.finish_reason,
645+
top_tokens,
646+
)
647+
)
648+
649+
if r.finish_reason is not None:
650+
result["rqueue"].put(None)
651+
self.prompt_cache.insert_cache(
652+
current_model_key, result["cache_key"], r.prompt_cache
653+
)
654+
del batch_results[r.uid]
488655

489656
def _serve_single(self, request):
490657
rqueue, request, args = request
@@ -497,12 +664,7 @@ def _serve_single(self, request):
497664
draft_model = self.model_provider.draft_model
498665

499666
# Prepare the prompt
500-
if request.request_type == "chat":
501-
prompt = self._tokenize_chat(
502-
tokenizer, request.messages, request.tools, request.role_mapping
503-
)
504-
else:
505-
prompt = tokenizer.encode(request.prompt)
667+
prompt = self._tokenize(tokenizer, request)
506668

507669
# Start the generation context
508670
ctx = GenerationContext(
@@ -545,14 +707,13 @@ def _serve_single(self, request):
545707
cache, rest = self.prompt_cache.fetch_nearest_cache(
546708
self.model_provider.model_key, prompt
547709
)
548-
cache_key = prompt[: len(prompt) - len(rest)]
710+
cache_key = prompt[:]
549711
if cache is None:
550712
cache = make_prompt_cache(self.model_provider.model)
551713
if self.model_provider.draft_model is not None:
552714
cache += make_prompt_cache(self.model_provider.draft_model)
553715

554716
# Process the prompt and generate tokens
555-
cache_key += rest
556717
for gen in stream_generate(
557718
model=model,
558719
tokenizer=tokenizer,
@@ -565,7 +726,25 @@ def _serve_single(self, request):
565726
num_draft_tokens=args.num_draft_tokens,
566727
# TODO: prompt progress callback
567728
):
568-
rqueue.put(gen)
729+
top_tokens = None
730+
if args.logprobs > 0:
731+
sorted_indices = mx.argpartition(
732+
-gen.logprobs, kth=args.logprobs - 1
733+
)
734+
top_indices = sorted_indices[: args.logprobs]
735+
top_logprobs = gen.logprobs[top_indices]
736+
top_token_info = zip(top_indices.tolist(), top_logprobs.tolist())
737+
top_tokens = tuple(top_token_info)
738+
739+
rqueue.put(
740+
Response(
741+
gen.text,
742+
gen.token,
743+
gen.logprobs[gen.token].item(),
744+
gen.finish_reason,
745+
top_tokens,
746+
)
747+
)
569748
cache_key.append(gen.token)
570749
rqueue.put(None)
571750

@@ -974,15 +1153,11 @@ def handle_completion(self, request: CompletionRequest, stop_words: List[str]):
9741153

9751154
# Save the token and its logprob
9761155
tokens.append(gen.token)
977-
token_logprobs.append(gen.logprobs[gen.token].item())
1156+
token_logprobs.append(gen.logprob)
9781157

9791158
# If requested save the k top logprobs
980-
if args.logprobs > 0:
981-
sorted_indices = mx.argpartition(-logprobs, kth=self.logprobs - 1)
982-
top_indices = sorted_indices[: self.logprobs]
983-
top_logprobs = logprobs[top_indices]
984-
top_token_info = zip(top_indices.tolist(), top_logprobs.tolist())
985-
top_tokens.append(tuple(top_token_info))
1159+
if gen.top_tokens is not None:
1160+
top_tokens.append(gen.top_tokens)
9861161

9871162
# Check if we should stop early
9881163
# TODO: This doesn't actually stop generation in the generation
@@ -993,8 +1168,8 @@ def handle_completion(self, request: CompletionRequest, stop_words: List[str]):
9931168
)
9941169
if stop_condition.stop_met:
9951170
finish_reason = "stop"
996-
tokens = tokens[: -stop_condition.trim_length]
997-
text = text[: -stop_condition.trim_text_length]
1171+
tokens = tokens[: len(tokens) - stop_condition.trim_length]
1172+
text = text[: len(text) - stop_condition.trim_text_length]
9981173
segment = ""
9991174
break
10001175

@@ -1020,9 +1195,9 @@ def handle_completion(self, request: CompletionRequest, stop_words: List[str]):
10201195
if gen.finish_reason is not None:
10211196
finish_reason = gen.finish_reason
10221197

1023-
logging.debug(f"Prompt: {gen.prompt_tps:.3f} tokens-per-sec")
1024-
logging.debug(f"Generation: {gen.generation_tps:.3f} tokens-per-sec")
1025-
logging.debug(f"Peak memory: {gen.peak_memory:.3f} GB")
1198+
# logging.debug(f"Prompt: {gen.prompt_tps:.3f} tokens-per-sec")
1199+
# logging.debug(f"Generation: {gen.generation_tps:.3f} tokens-per-sec")
1200+
# logging.debug(f"Peak memory: {gen.peak_memory:.3f} GB")
10261201

10271202
if self.stream:
10281203
response = self.generate_response(
@@ -1192,7 +1367,7 @@ def run(
11921367
host: str,
11931368
port: int,
11941369
model_provider: ModelProvider,
1195-
server_class=HTTPServer,
1370+
server_class=ThreadingHTTPServer,
11961371
handler_class=APIHandler,
11971372
):
11981373
server_address = (host, port)

0 commit comments

Comments
 (0)