@@ -226,6 +226,7 @@ def _get_or_create_generator(
226226 "active" : set (), # active uids
227227 "cache_boundaries" : {}, # uid -> stable prefix token count
228228 "xavier_prompts" : {},
229+ "xavier_writes" : {},
229230 "task" : None ,
230231 }
231232
@@ -345,11 +346,16 @@ def _store_xavier_prompt_caches(self, batch_generator, prompt_results, gen_dict)
345346 for result in prompt_results :
346347 if not result .end_of_prompt or result .uid not in prompts :
347348 continue
348- extracted = batch_generator .extract_cache ([result .uid ]).get (result .uid )
349- if extracted is not None :
350- cache , _ = extracted
351- tokens = prompts .pop (result .uid )
352- self ._xavier .publish (cache , tokens [:- 1 ])
349+ tokens , cached_tokens = prompts .pop (result .uid )
350+ try :
351+ extracted = batch_generator .extract_cache ([result .uid ]).get (result .uid )
352+ if extracted is not None :
353+ cache , _ = extracted
354+ gen_dict ["xavier_writes" ][result .uid ] = self ._xavier .publish (
355+ cache , tokens [:- 1 ], cached_tokens
356+ )
357+ except Exception :
358+ logger .warning ("MLX Xavier prefix publication failed" , exc_info = True )
353359
354360 async def _background_worker (self , gen_dict ):
355361 """Background worker that continuously calls next() and distributes results."""
@@ -425,7 +431,6 @@ async def generate_stream(
425431 request_id : Optional [str ] = None ,
426432 skip_special_tokens : bool = True ,
427433 prompt_cache_prefix_len : Optional [int ] = None ,
428- kv_transfer_params : Optional [dict ] = None ,
429434 prepared_cache = None ,
430435 prompt_token_ids : Optional [List [int ]] = None ,
431436 ) -> AsyncGenerator [CompletionChunk , None ]:
@@ -450,8 +455,6 @@ async def generate_stream(
450455 queue : asyncio .Queue = asyncio .Queue ()
451456
452457 external_cache = prepared_cache
453- if external_cache is None and getattr (self , "_xavier" , None ) is not None :
454- external_cache = await self ._xavier .fetch (prompt_tokens , kv_transfer_params )
455458 insert_args = (
456459 batch_generator ,
457460 prompt_tokens ,
@@ -463,11 +466,15 @@ async def generate_stream(
463466 if external_cache is not None
464467 else self ._insert_request (* insert_args )
465468 )
469+ remote_cached_tokens = external_cache [1 ] if external_cache is not None else 0
466470 if (
467471 getattr (self , "_xavier" , None ) is not None
468- and cached_prompt_tokens < input_echo_len - 1
472+ and remote_cached_tokens < input_echo_len - 1
469473 ):
470- gen_dict ["xavier_prompts" ][inserted_uid ] = prompt_tokens
474+ gen_dict ["xavier_prompts" ][inserted_uid ] = (
475+ prompt_tokens ,
476+ remote_cached_tokens ,
477+ )
471478 if cache_boundary is not None :
472479 gen_dict ["cache_boundaries" ][inserted_uid ] = cache_boundary
473480
@@ -606,7 +613,8 @@ async def generate_stream(
606613 gen_dict ["xavier_prompts" ].pop (inserted_uid , None )
607614 batch_generator .remove ([inserted_uid ])
608615 if getattr (self , "_xavier" , None ) is not None :
609- await self ._xavier .flush ()
616+ write = gen_dict ["xavier_writes" ].pop (inserted_uid , None )
617+ await self ._xavier .flush ([write ] if write is not None else [])
610618
611619 async def generate (
612620 self ,
@@ -618,7 +626,6 @@ async def generate(
618626 stream : bool = False ,
619627 skip_special_tokens : bool = True ,
620628 prompt_cache_prefix_len : Optional [int ] = None ,
621- kv_transfer_params : Optional [dict ] = None ,
622629 prepared_cache = None ,
623630 prompt_token_ids : Optional [List [int ]] = None ,
624631 ) -> Tuple [str , CompletionUsage ]:
@@ -640,7 +647,6 @@ async def generate(
640647 request_id = None ,
641648 skip_special_tokens = skip_special_tokens ,
642649 prompt_cache_prefix_len = prompt_cache_prefix_len ,
643- kv_transfer_params = kv_transfer_params ,
644650 prepared_cache = prepared_cache ,
645651 prompt_token_ids = prompt_token_ids ,
646652 ):
@@ -1119,6 +1125,10 @@ def wait_for_load(self):
11191125 raise ValueError (
11201126 "MLX Xavier requires the continuous batching text engine"
11211127 )
1128+ assert self ._loop is not None , "Service not started correctly"
1129+ asyncio .run_coroutine_threadsafe (
1130+ self ._xavier .initialize (), self ._loop
1131+ ).result ()
11221132
11231133 # Update allow_batch based on distributed inference
11241134 # Only enable continuous batching for non-distributed inference (single worker)
0 commit comments