1111import warnings
1212from collections import deque
1313from dataclasses import dataclass , field
14- from http .server import BaseHTTPRequestHandler , HTTPServer
14+ from http .server import BaseHTTPRequestHandler , ThreadingHTTPServer
1515from pathlib import Path
1616from queue import Empty as QueueEmpty
1717from queue import Queue
3232from huggingface_hub import scan_cache_dir
3333
3434from ._version import __version__
35- from .generate import stream_generate
35+ from .generate import BatchGenerator , stream_generate
3636from .models .cache import can_trim_prompt_cache , make_prompt_cache , trim_prompt_cache
3737from .sample_utils import make_logits_processors , make_sampler
3838from .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+
361370class 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