Skip to content

Commit 0d67672

Browse files
committed
fix: repair MLX Xavier recovery and prefix publication
1 parent 64f07ab commit 0d67672

8 files changed

Lines changed: 504 additions & 74 deletions

File tree

‎xinference/core/supervisor.py‎

Lines changed: 9 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -3095,6 +3095,14 @@ async def launch_builtin_model(
30953095
raise ValueError("NIXL requires explicit prefill and decode replica roles")
30963096
# Xavier-related
30973097
requested_xavier = bool(kwargs.pop("enable_xavier", False))
3098+
if (
3099+
requested_xavier
3100+
and not pd_enabled
3101+
and replica <= 1
3102+
and (model_engine or "").lower() in ("mlx", "sglang")
3103+
):
3104+
logger.warning("Enabling xavier when replica<=1 is meaningless.")
3105+
requested_xavier = False
30983106
mlx_xavier = (
30993107
(requested_xavier or pd_enabled)
31003108
and transport_backend == "xavier"
@@ -3108,7 +3116,7 @@ async def launch_builtin_model(
31083116
if (
31093117
model_type not in (None, "LLM")
31103118
or model_format != "mlx"
3111-
or quantization not in (None, "none")
3119+
or quantization not in (None, "none", "fp16", "bf16")
31123120
):
31133121
raise ValueError("MLX Xavier requires unquantized MLX text weights")
31143122
kwargs["_xavier_cache_config"] = {
@@ -3121,9 +3129,6 @@ async def launch_builtin_model(
31213129
and model_engine is not None
31223130
and model_engine.lower() == "sglang"
31233131
)
3124-
if sglang_xavier and not pd_enabled and replica <= 1:
3125-
logger.warning("Enabling xavier when replica<=1 is meaningless.")
3126-
sglang_xavier = False
31273132
sglang_nixl = (
31283133
pd_enabled
31293134
and transport_backend == "nixl"

‎xinference/core/tests/test_pd_launch.py‎

Lines changed: 61 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -69,10 +69,16 @@ def launch_kwargs():
6969

7070
@pytest.mark.asyncio
7171
@pytest.mark.parametrize("pd", [False, True])
72-
async def test_mlx_xavier_launch_and_cleanup(launch_runtime, pd):
72+
@pytest.mark.parametrize("quantization", [None, "none", "fp16", "bf16"])
73+
async def test_mlx_xavier_launch_and_cleanup(launch_runtime, pd, quantization):
7374
supervisor, workers, actors, destroy = launch_runtime
7475
kwargs = launch_kwargs()
75-
kwargs.update(model_engine="MLX", model_format="mlx", xavier_cache_bytes=123456)
76+
kwargs.update(
77+
model_engine="MLX",
78+
model_format="mlx",
79+
quantization=quantization,
80+
xavier_cache_bytes=123456,
81+
)
7682
if not pd:
7783
kwargs["enable_xavier"] = True
7884
for replica in kwargs["replica_config"]:
@@ -99,6 +105,59 @@ async def test_mlx_xavier_launch_and_cleanup(launch_runtime, pd):
99105
assert destroy.await_count == (2 if pd else 1)
100106

101107

108+
@pytest.mark.asyncio
109+
async def test_single_mlx_replica_disables_shared_xavier(launch_runtime, caplog):
110+
supervisor, workers, actors, _ = launch_runtime
111+
supervisor._resolve_replica_config.return_value = ([(workers[0], [0], 1)], {0: "p"})
112+
kwargs = launch_kwargs()
113+
kwargs.update(model_engine="MLX", model_format="mlx", enable_xavier=True, replica=1)
114+
kwargs["replica_config"] = kwargs["replica_config"][:1]
115+
kwargs["replica_config"][0].role = None
116+
await supervisor.launch_builtin_model(**kwargs)
117+
assert "replica<=1" in caplog.text
118+
assert not actors and not supervisor._xavier_cache_mapping
119+
assert (
120+
"_xavier_cache_config" not in workers[0].launch_builtin_model.call_args.kwargs
121+
)
122+
123+
124+
@pytest.mark.asyncio
125+
@pytest.mark.parametrize("role", ["prefill", "decode"])
126+
async def test_mlx_pd_recovery_keeps_bytes_cache_and_registers_replacement(role):
127+
import xoscar as xo
128+
129+
from ...model.llm.xavier.backends.bytes.storage import XavierBytesCacheActor
130+
from ...model.llm.xavier.tests.test_bytes_storage import contract
131+
from ..worker import WorkerActor
132+
133+
pool = await xo.create_actor_pool("127.0.0.1", n_process=0)
134+
async with pool:
135+
cache = await xo.create_actor(
136+
XavierBytesCacheActor, address=pool.external_address, uid="bytes-cache"
137+
)
138+
namespace = await cache.configure(contract().to_dict())
139+
worker = MagicMock()
140+
supervisor = AsyncMock()
141+
worker.get_supervisor_ref = AsyncMock(return_value=supervisor)
142+
worker.launch_builtin_model = AsyncMock(return_value="replacement:1234")
143+
worker.wait_for_load = AsyncMock()
144+
replacement = MagicMock()
145+
worker._model_uid_to_model = {"pd-rep0": replacement}
146+
config = dict(role=role, address=cache.address, uid=cache.uid)
147+
await WorkerActor.recover_model(
148+
worker, dict(model_uid="pd-rep0", _xavier_cache_config=config)
149+
)
150+
worker.launch_builtin_model.assert_awaited_once_with(
151+
model_uid="pd-rep0", _xavier_cache_config=config
152+
)
153+
worker.wait_for_load.assert_awaited_once_with("pd-rep0")
154+
supervisor.unregister_pd_replica.assert_awaited_once_with("pd", "pd-rep0")
155+
supervisor.register_pd_replica.assert_awaited_once_with(
156+
"pd", "pd-rep0", replacement
157+
)
158+
assert (await cache.get_stats())["namespace"] == namespace
159+
160+
102161
@pytest.mark.asyncio
103162
async def test_mlx_failed_launch_cleans_bytes_actor(launch_runtime):
104163
supervisor, workers, actors, destroy = launch_runtime

‎xinference/core/worker.py‎

Lines changed: 4 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -6756,7 +6756,10 @@ async def recover_model(self, launch_args: Dict[str, Any]):
67566756
).get("role") in ("prefill", "decode"):
67576757
await supervisor_ref.unregister_pd_replica(origin_uid, rep_model_uid)
67586758
cache_config = launch_args.get("_xavier_cache_config", {})
6759-
if cache_config.get("role") in ("prefill", "decode"):
6759+
if (
6760+
cache_config.get("role") in ("prefill", "decode")
6761+
and "rank" in cache_config
6762+
):
67606763
directory = await xo.actor_ref(
67616764
address=cache_config["address"], uid=cache_config["uid"]
67626765
)

‎xinference/model/llm/mlx/core.py‎

Lines changed: 23 additions & 13 deletions
Original file line numberDiff line numberDiff line change
@@ -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

Comments
 (0)