fix(vlm): make EPD cache publication transactional (#36949)

This commit is contained in:
Mick
2026-09-05 20:33:23 +08:00
committed by GitHub
parent bd16c22a04
commit a18106bbc3
6 changed files with 746 additions and 139 deletions
@@ -1028,9 +1028,7 @@ async def _push_embedding_to_prefill(
if backend == "zmq_to_scheduler" and request.get("embedding_port") is None: if backend == "zmq_to_scheduler" and request.get("embedding_port") is None:
send_coro = enc.send_with_url(req_id=req_id) send_coro = enc.send_with_url(req_id=req_id)
if background_url_send: if background_url_send:
task = asyncio.create_task(send_coro) enc._create_background_task(send_coro)
enc.background_tasks.add(task)
task.add_done_callback(enc.background_tasks.discard)
else: else:
await send_coro await send_coro
return return
@@ -350,15 +350,8 @@ class MooncakeDelivery(EncoderDelivery):
async def release(self, state: ReqState) -> None: async def release(self, state: ReqState) -> None:
mm_data = state.embedding_data mm_data = state.embedding_data
if mm_data is not None and mm_data._mr_ptr is not None: if mm_data is not None:
try: self.encoder._deregister_shared_mr(mm_data)
self.encoder.engine.deregister(mm_data._mr_ptr)
except Exception as dereg_err:
logger.warning(
f"Shared-MR deregister failed for {state.req_id}: {dereg_err}"
)
finally:
mm_data._mr_ptr = None
class ZmqDelivery(EncoderDelivery): class ZmqDelivery(EncoderDelivery):
@@ -607,6 +600,21 @@ class MMEncoder:
logger.info(f"rank {rank} init finish ") logger.info(f"rank {rank} init finish ")
def _background_task_done(self, task: asyncio.Task) -> None:
self.background_tasks.discard(task)
try:
task.result()
except asyncio.CancelledError:
pass
except Exception:
logger.exception("MMEncoder background task failed")
def _create_background_task(self, awaitable: Awaitable[Any]) -> asyncio.Task:
task = asyncio.create_task(awaitable)
self.background_tasks.add(task)
task.add_done_callback(self._background_task_done)
return task
def supports_modality(self, modality: Modality) -> bool: def supports_modality(self, modality: Modality) -> bool:
return self.preprocessor.supports_modality(modality) return self.preprocessor.supports_modality(modality)
@@ -646,7 +654,7 @@ class MMEncoder:
if should_release: if should_release:
await self.release_request(state.req_id) await self.release_request(state.req_id)
def _stage_embedding(self, mm_data: EmbeddingData) -> None: def _embedding_state_for_stage(self, mm_data: EmbeddingData) -> ReqState:
state = self._require_active_encode_state(mm_data.req_id) state = self._require_active_encode_state(mm_data.req_id)
metadata = state.embedding_data metadata = state.embedding_data
if ( if (
@@ -660,9 +668,21 @@ class MMEncoder:
f"expected={metadata.shape}/{metadata.dtype}, " f"expected={metadata.shape}/{metadata.dtype}, "
f"actual={mm_data.shape}/{mm_data.dtype}" f"actual={mm_data.shape}/{mm_data.dtype}"
) )
return state
def _stage_embedding(self, mm_data: EmbeddingData) -> None:
state = self._embedding_state_for_stage(mm_data)
state.embedding_data = mm_data state.embedding_data = mm_data
state.embedding_ready.set() state.embedding_ready.set()
def _stage_embedding_batch(self, embeddings: List[EmbeddingData]) -> None:
"""Validate the whole fused batch before publishing any result."""
states = [self._embedding_state_for_stage(mm_data) for mm_data in embeddings]
for state, mm_data in zip(states, embeddings):
state.embedding_data = mm_data
for state in states:
state.embedding_ready.set()
async def _wait_for_embedding(self, state: ReqState) -> EmbeddingData: async def _wait_for_embedding(self, state: ReqState) -> EmbeddingData:
await state.embedding_ready.wait() await state.embedding_ready.wait()
if state.embedding_data is None: if state.embedding_data is None:
@@ -1046,7 +1066,17 @@ class MMEncoder:
ctx: EncodeContext, ctx: EncodeContext,
) -> Tuple[List[int], List[int]]: ) -> Tuple[List[int], List[int]]:
if self.rank == 0: if self.rank == 0:
exist_mask = await self.mm_global_cache.batch_is_exist(ctx.str_mm_hashes) try:
exist_mask = await self.mm_global_cache.batch_is_exist(
ctx.str_mm_hashes
)
except Exception:
logger.exception(
"Global multimodal cache lookup failed for req %s; "
"falling back to ViT",
ctx.req_id,
)
exist_mask = [False] * ctx.num_items
mask_tensor = torch.tensor( mask_tensor = torch.tensor(
[1 if e else 0 for e in exist_mask], dtype=torch.int32 [1 if e else 0 for e in exist_mask], dtype=torch.int32
) )
@@ -1064,27 +1094,44 @@ class MMEncoder:
self, self,
ctx: EncodeContext, ctx: EncodeContext,
hit_indices: List[int], hit_indices: List[int],
) -> List[str]: ) -> Tuple[List[str], bool]:
if self.rank != 0 or not hit_indices: if self.rank != 0 or not hit_indices:
return [] return [], False
hit_hashes = [ctx.str_mm_hashes[i] for i in hit_indices] hit_hashes = [ctx.str_mm_hashes[i] for i in hit_indices]
hit_tokens = [ctx.preprocess_result.token_counts[i] for i in hit_indices] hit_tokens = [ctx.preprocess_result.token_counts[i] for i in hit_indices]
self.mm_global_cache.prefetch(ctx.req_id, hit_hashes, hit_tokens, ctx.modality) try:
return hit_hashes self.mm_global_cache.prefetch(
ctx.req_id, hit_hashes, hit_tokens, ctx.modality
)
except Exception:
logger.exception(
"Global multimodal cache prefetch failed for req %s; "
"falling back to ViT",
ctx.req_id,
)
return [], True
return hit_hashes, False
async def _wait_global_cache_prefetch( async def _wait_global_cache_prefetch(
self, self,
ctx: EncodeContext, ctx: EncodeContext,
hit_indices: List[int], hit_indices: List[int],
hit_hashes: List[str], hit_hashes: List[str],
prefetch_failed: bool,
) -> List[int]: ) -> List[int]:
fallback_mask = torch.zeros(ctx.num_items, dtype=torch.int32) fallback_mask = torch.zeros(ctx.num_items, dtype=torch.int32)
if self.rank == 0 and hit_indices: if self.rank == 0 and hit_indices:
if prefetch_failed:
for idx in hit_indices:
fallback_mask[idx] = 1
else:
try: try:
async def _wait_prefetch(): async def _wait_prefetch():
while not self.mm_global_cache.check_prefetch_progress(ctx.req_id): while not self.mm_global_cache.check_prefetch_progress(
ctx.req_id
):
await asyncio.sleep(0.005) await asyncio.sleep(0.005)
await asyncio.wait_for(_wait_prefetch(), timeout=60.0) await asyncio.wait_for(_wait_prefetch(), timeout=60.0)
@@ -1098,7 +1145,7 @@ class MMEncoder:
f"Req {ctx.req_id}: {num_partial_fail}/{len(hit_indices)} " f"Req {ctx.req_id}: {num_partial_fail}/{len(hit_indices)} "
f"cache-hit items failed to load, falling back to ViT" f"cache-hit items failed to load, falling back to ViT"
) )
except (asyncio.TimeoutError, Exception) as e: except Exception as e:
logger.error( logger.error(
f"Prefetch failed for req {ctx.req_id}: {e}. " f"Prefetch failed for req {ctx.req_id}: {e}. "
f"Falling back to ViT for {len(hit_indices)} hit items." f"Falling back to ViT for {len(hit_indices)} hit items."
@@ -1112,6 +1159,28 @@ class MMEncoder:
] ]
return fallback_indices return fallback_indices
def _stage_global_cache_slices(
self,
ctx: EncodeContext,
indices: List[int],
slices: List[torch.Tensor],
) -> Tuple[List[str], List[Any]]:
"""Stage cache insert data without making cache failure fatal."""
if not slices:
return [], []
hashes = [ctx.str_mm_hashes[i] for i in indices]
try:
handles = self.mm_global_cache.store_to_pool_async(
hashes, slices, ctx.modality
)
except Exception:
logger.exception(
"Global multimodal cache staging failed for req %s; skipping insert",
ctx.req_id,
)
return [], []
return hashes, handles
def _launch_global_cache_insert( def _launch_global_cache_insert(
self, self,
ctx: EncodeContext, ctx: EncodeContext,
@@ -1122,6 +1191,7 @@ class MMEncoder:
return return
async def _background_insert(): async def _background_insert():
try:
await asyncio.to_thread( await asyncio.to_thread(
self.mm_global_cache.wait_store_to_pool, self.mm_global_cache.wait_store_to_pool,
d2h_handles, d2h_handles,
@@ -1131,10 +1201,12 @@ class MMEncoder:
hashes, hashes,
ctx.modality, ctx.modality,
) )
except Exception:
logger.exception(
"Global multimodal cache insert failed for req %s", ctx.req_id
)
task = asyncio.create_task(_background_insert()) self._create_background_task(_background_insert())
self.background_tasks.add(task)
task.add_done_callback(self.background_tasks.discard)
@staticmethod @staticmethod
def _as_2d_tensor(tensor: torch.Tensor) -> torch.Tensor: def _as_2d_tensor(tensor: torch.Tensor) -> torch.Tensor:
@@ -1262,7 +1334,7 @@ class MMEncoder:
) -> Optional[torch.Tensor]: ) -> Optional[torch.Tensor]:
"""Resolve cache hits, compute misses, assemble output, and insert misses.""" """Resolve cache hits, compute misses, assemble output, and insert misses."""
missing_indices, hit_indices = await self._lookup_global_cache(ctx) missing_indices, hit_indices = await self._lookup_global_cache(ctx)
hit_hashes = self._prefetch_global_cache_hits(ctx, hit_indices) hit_hashes, prefetch_failed = self._prefetch_global_cache_hits(ctx, hit_indices)
new_slices = [] new_slices = []
if missing_indices: if missing_indices:
@@ -1274,19 +1346,20 @@ class MMEncoder:
ctx.get_feature_fn, ctx.get_feature_fn,
) )
miss_hashes = []
miss_d2h_handles = [] miss_d2h_handles = []
# The CPU output path starts D2H staging before waiting for cache-hit loads. # The CPU output path starts D2H staging before waiting for cache-hit loads.
if self.rank == 0 and new_slices and not keep_on_gpu: if self.rank == 0 and new_slices and not keep_on_gpu:
miss_hashes = [ctx.str_mm_hashes[i] for i in missing_indices] miss_hashes, miss_d2h_handles = self._stage_global_cache_slices(
miss_d2h_handles = self.mm_global_cache.store_to_pool_async( ctx, missing_indices, new_slices
miss_hashes, new_slices, ctx.modality
) )
fallback_indices = await self._wait_global_cache_prefetch( fallback_indices = await self._wait_global_cache_prefetch(
ctx, hit_indices, hit_hashes ctx, hit_indices, hit_hashes, prefetch_failed
) )
fallback_slices = [] fallback_slices = []
fallback_hashes = []
fallback_d2h_handles = [] fallback_d2h_handles = []
if fallback_indices: if fallback_indices:
logger.info( logger.info(
@@ -1301,9 +1374,8 @@ class MMEncoder:
ctx.get_feature_fn, ctx.get_feature_fn,
) )
if self.rank == 0 and not keep_on_gpu: if self.rank == 0 and not keep_on_gpu:
fallback_hashes = [ctx.str_mm_hashes[i] for i in fallback_indices] fallback_hashes, fallback_d2h_handles = self._stage_global_cache_slices(
fallback_d2h_handles = self.mm_global_cache.store_to_pool_async( ctx, fallback_indices, fallback_slices
fallback_hashes, fallback_slices, ctx.modality
) )
if self.rank == 0: if self.rank == 0:
@@ -1311,14 +1383,14 @@ class MMEncoder:
# Start staging newly computed GPU slices into the CPU cache # Start staging newly computed GPU slices into the CPU cache
# pool asynchronously before assembling the GPU output. # pool asynchronously before assembling the GPU output.
if new_slices: if new_slices:
miss_hashes = [ctx.str_mm_hashes[i] for i in missing_indices] miss_hashes, miss_d2h_handles = self._stage_global_cache_slices(
miss_d2h_handles = self.mm_global_cache.store_to_pool_async( ctx, missing_indices, new_slices
miss_hashes, new_slices, ctx.modality
) )
if fallback_slices: if fallback_slices:
fallback_hashes = [ctx.str_mm_hashes[i] for i in fallback_indices] fallback_hashes, fallback_d2h_handles = (
fallback_d2h_handles = self.mm_global_cache.store_to_pool_async( self._stage_global_cache_slices(
fallback_hashes, fallback_slices, ctx.modality ctx, fallback_indices, fallback_slices
)
) )
mm_embedding = self._assemble_global_cache_gpu( mm_embedding = self._assemble_global_cache_gpu(
ctx, ctx,
@@ -1337,11 +1409,9 @@ class MMEncoder:
fallback_slices, fallback_slices,
) )
new_hashes = [ctx.str_mm_hashes[i] for i in missing_indices]
new_hashes += [ctx.str_mm_hashes[i] for i in fallback_indices]
self._launch_global_cache_insert( self._launch_global_cache_insert(
ctx, ctx,
new_hashes, miss_hashes + fallback_hashes,
miss_d2h_handles + fallback_d2h_handles, miss_d2h_handles + fallback_d2h_handles,
) )
return mm_embedding return mm_embedding
@@ -1401,6 +1471,17 @@ class MMEncoder:
time.perf_counter() - forward_start, modality=modality_str time.perf_counter() - forward_start, modality=modality_str
) )
try:
self._validate_embedding_token_count(ctx, mm_embedding)
except InternalError:
# Old releases could cache a malformed result before the
# outer validation ran. Do not make that entry permanently
# poison every request for the same media.
if cache_hit:
async with self.mm_cache_lock:
self.mm_cache.free(mm_hash, None)
raise
# Per-request cache hit metrics: tokens = embedding rows. # Per-request cache hit metrics: tokens = embedding rows.
if use_mm_cache and encoder_metrics_collector is not None: if use_mm_cache and encoder_metrics_collector is not None:
total_tokens = int(mm_embedding.shape[0]) total_tokens = int(mm_embedding.shape[0])
@@ -1479,18 +1560,21 @@ class MMEncoder:
mm_embedding = await self._compute_global_cache_embedding( mm_embedding = await self._compute_global_cache_embedding(
ctx, keep_on_gpu=keep_on_gpu ctx, keep_on_gpu=keep_on_gpu
) )
else: if mm_embedding is not None:
mm_embedding = await self._compute_direct_embedding( self._validate_embedding_token_count(ctx, mm_embedding)
ctx, keep_on_gpu=keep_on_gpu return mm_embedding
) return await self._compute_direct_embedding(ctx, keep_on_gpu=keep_on_gpu)
@staticmethod
def _validate_embedding_token_count(
ctx: EncodeContext, mm_embedding: torch.Tensor
) -> None:
expected_tokens = sum(ctx.preprocess_result.token_counts) expected_tokens = sum(ctx.preprocess_result.token_counts)
if mm_embedding is not None and mm_embedding.shape[0] != expected_tokens: if mm_embedding.shape[0] != expected_tokens:
raise InternalError( raise InternalError(
f"Encoder produced {mm_embedding.shape[0]} tokens, but " f"Encoder produced {mm_embedding.shape[0]} tokens, but "
f"preprocessor metadata expected {expected_tokens}" f"preprocessor metadata expected {expected_tokens}"
) )
return mm_embedding
async def _publish_preprocess_metadata( async def _publish_preprocess_metadata(
self, ctx: EncodeContext, requests: List[dict] self, ctx: EncodeContext, requests: List[dict]
@@ -1551,21 +1635,35 @@ class MMEncoder:
mr_already_registered = mm_data._mr_ptr == embedding.data_ptr() mr_already_registered = mm_data._mr_ptr == embedding.data_ptr()
if not mr_already_registered: if not mr_already_registered:
self.engine.register(embedding.data_ptr(), embedding.nbytes) self.engine.register(embedding.data_ptr(), embedding.nbytes)
transfer_error = None
try:
_t_xfer_start = time.monotonic() _t_xfer_start = time.monotonic()
xfer_ret = await asyncio.to_thread( xfer_ret = await self._run_mooncake_transfer(
self.engine.transfer_sync,
session_id, session_id,
embedding.data_ptr(), embedding.data_ptr(),
buffer_address, buffer_address,
embedding.nbytes, embedding.nbytes,
) )
except BaseException as error:
transfer_error = error
raise
finally:
if not mr_already_registered:
try:
self.engine.deregister(embedding.data_ptr())
except Exception:
if transfer_error is None:
raise
logger.exception(
"Per-send MR deregistration also failed for %s; "
"preserving the transfer error",
req_id,
)
xfer_ms = (time.monotonic() - _t_xfer_start) * 1000.0 xfer_ms = (time.monotonic() - _t_xfer_start) * 1000.0
if encoder_metrics_collector is not None: if encoder_metrics_collector is not None:
encoder_metrics_collector.observe_transfer( encoder_metrics_collector.observe_transfer(
xfer_ms / 1000.0, backend="mooncake" xfer_ms / 1000.0, backend="mooncake"
) )
if not mr_already_registered:
self.engine.deregister(embedding.data_ptr())
if xfer_ret < 0: if xfer_ret < 0:
raise InternalError( raise InternalError(
f"Mooncake transfer_sync failed for {req_id} " f"Mooncake transfer_sync failed for {req_id} "
@@ -1684,6 +1782,32 @@ class MMEncoder:
backend=get_disagg().encoder_transfer_backend, backend=get_disagg().encoder_transfer_backend,
) )
async def _run_mooncake_transfer(
self,
session_id,
source_address: int,
destination_address: int,
size: int,
) -> int:
"""Keep the send active until its blocking transfer stops using the MR."""
transfer_task = asyncio.create_task(
asyncio.to_thread(
self.engine.transfer_sync,
session_id,
source_address,
destination_address,
size,
)
)
try:
return await asyncio.shield(transfer_task)
except asyncio.CancelledError:
try:
await transfer_task
except Exception:
pass
raise
def _register_shared_mr(self, mm_data: EmbeddingData, embedding: torch.Tensor): def _register_shared_mr(self, mm_data: EmbeddingData, embedding: torch.Tensor):
"""Register one MR shared by every rank's /send; _send re-registers on failure.""" """Register one MR shared by every rank's /send; _send re-registers on failure."""
try: try:
@@ -1695,6 +1819,18 @@ class MMEncoder:
f"falling back to per-/send register: {reg_err}" f"falling back to per-/send register: {reg_err}"
) )
def _deregister_shared_mr(self, mm_data: EmbeddingData) -> None:
if mm_data._mr_ptr is None:
return
try:
self.engine.deregister(mm_data._mr_ptr)
except Exception as dereg_err:
logger.warning(
f"Shared-MR deregister failed for {mm_data.req_id}: {dereg_err}"
)
finally:
mm_data._mr_ptr = None
def _stage_embeddings( def _stage_embeddings(
self, self,
ctx: EncodeContext, ctx: EncodeContext,
@@ -1715,11 +1851,14 @@ class MMEncoder:
results = [] results = []
staged_embeddings = [] staged_embeddings = []
try:
item_offset = 0 item_offset = 0
token_offset = 0 token_offset = 0
for req, num_items in zip(requests, ctx.items_per_req): for req, num_items in zip(requests, ctx.items_per_req):
item_end = item_offset + num_items item_end = item_offset + num_items
num_tokens = sum(ctx.preprocess_result.token_counts[item_offset:item_end]) num_tokens = sum(
ctx.preprocess_result.token_counts[item_offset:item_end]
)
embedding = mm_embedding[token_offset : token_offset + num_tokens] embedding = mm_embedding[token_offset : token_offset + num_tokens]
if keep_on_gpu and len(requests) > 1: if keep_on_gpu and len(requests) > 1:
# A view would pin the whole batch tensor until the last transfer. # A view would pin the whole batch tensor until the last transfer.
@@ -1743,7 +1882,13 @@ class MMEncoder:
self._register_shared_mr(mm_data, embedding) self._register_shared_mr(mm_data, embedding)
staged_embeddings.append(mm_data) staged_embeddings.append(mm_data)
results.append( results.append(
(embedding.nbytes, embedding.shape[0], embedding.shape[1], None, None) (
embedding.nbytes,
embedding.shape[0],
embedding.shape[1],
None,
None,
)
) )
item_offset = item_end item_offset = item_end
token_offset += num_tokens token_offset += num_tokens
@@ -1752,9 +1897,12 @@ class MMEncoder:
# per-request clones) must land before /send reads the buffers. # per-request clones) must land before /send reads the buffers.
if keep_on_gpu and mm_embedding.is_cuda: if keep_on_gpu and mm_embedding.is_cuda:
torch.cuda.current_stream(mm_embedding.device).synchronize() torch.cuda.current_stream(mm_embedding.device).synchronize()
for mm_data in staged_embeddings: self._stage_embedding_batch(staged_embeddings)
self._stage_embedding(mm_data)
return results return results
except BaseException:
for mm_data in staged_embeddings:
self._deregister_shared_mr(mm_data)
raise
def _stage_errors( def _stage_errors(
self, requests: List[dict], modality: Modality, exc: Exception self, requests: List[dict], modality: Modality, exc: Exception
@@ -143,25 +143,25 @@ class MooncakeTransferEngine:
self.hostname, self.engine.get_rpc_port() self.hostname, self.engine.get_rpc_port()
).to_host_port_str() ).to_host_port_str()
def register(self, ptr, length): def register(self, ptr, length) -> None:
try: try:
ret_value = self.engine.register_memory(ptr, length) ret_value = self.engine.register_memory(ptr, length)
except Exception: except Exception as exc:
# Mark register as failed raise RuntimeError("Mooncake memory registration failed") from exc
ret_value = -1
if ret_value != 0: if ret_value != 0:
logger.debug("Mooncake memory registration %s failed.", ptr) raise RuntimeError(f"Mooncake memory registration failed (ret={ret_value})")
def deregister(self, ptr): def deregister(self, ptr) -> None:
try: try:
ret_value = self.engine.unregister_memory(ptr) ret_value = self.engine.unregister_memory(ptr)
except Exception: except Exception as exc:
# Mark deregister as failed raise RuntimeError("Mooncake memory deregistration failed") from exc
ret_value = -1
if ret_value != 0: if ret_value != 0:
logger.debug("Mooncake memory deregistration %s failed.", ptr) raise RuntimeError(
f"Mooncake memory deregistration failed (ret={ret_value})"
)
def batch_register(self, ptrs: List[int], lengths: List[int]) -> int: def batch_register(self, ptrs: List[int], lengths: List[int]) -> int:
"""Batch register multiple memory regions.""" """Batch register multiple memory regions."""
@@ -979,13 +979,32 @@ class EmbeddingCacheController:
handles: List[Tuple["EmbeddingCacheEntry", "AsyncCopyHandle"]], handles: List[Tuple["EmbeddingCacheEntry", "AsyncCopyHandle"]],
): ):
"""Wait for async D2H copies and mark entries READY.""" """Wait for async D2H copies and mark entries READY."""
completed = []
errors = []
for entry, handle in handles: for entry, handle in handles:
try:
handle.wait() handle.wait()
completed.append(True)
except BaseException as error:
completed.append(False)
errors.append(error)
with self.lock: with self.lock:
for entry, handle in handles: for (entry, _), success in zip(handles, completed, strict=True):
current = self.entries.get(entry.hash) current = self.entries.get(entry.hash)
if current is entry and current.state == EntryState.FILLING: if current is entry and current.state == EntryState.FILLING:
if success:
self._mark_ready(current) self._mark_ready(current)
else:
self._evict_entry(entry.hash)
if errors:
first_error = errors[0]
if len(errors) > 1:
first_error.add_note(
f"{len(errors) - 1} additional embedding copy operation(s) failed"
)
raise first_error
def _copy_tensor_to_pool( def _copy_tensor_to_pool(
self, tensor: torch.Tensor, entry: EmbeddingCacheEntry, pool: EmbeddingPool self, tensor: torch.Tensor, entry: EmbeddingCacheEntry, pool: EmbeddingPool
@@ -1,8 +1,9 @@
import asyncio import asyncio
import pickle import pickle
import threading
import unittest import unittest
from types import SimpleNamespace from types import SimpleNamespace
from unittest.mock import AsyncMock, patch from unittest.mock import AsyncMock, Mock, patch
import numpy as np import numpy as np
import torch import torch
@@ -23,7 +24,14 @@ from sglang.srt.disaggregation.encoder.server import (
rid_to_receive_count, rid_to_receive_count,
rid_to_receive_endpoint, rid_to_receive_endpoint,
) )
from sglang.srt.distributed.device_communicators.mooncake_transfer_engine import (
MooncakeTransferEngine,
)
from sglang.srt.managers.schedule_batch import Modality from sglang.srt.managers.schedule_batch import Modality
from sglang.srt.mem_cache.multimodal_cache import (
EmbeddingResult,
MultiModalStaticCache,
)
from sglang.srt.utils.common import safe_pickle_loads from sglang.srt.utils.common import safe_pickle_loads
from sglang.test.ci.ci_register import register_cpu_ci from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase from sglang.test.test_utils import CustomTestCase
@@ -130,6 +138,129 @@ class TestEncoderPreprocessorKimiGrid(CustomTestCase):
class TestEncoderDelivery(CustomTestCase): class TestEncoderDelivery(CustomTestCase):
@staticmethod
def _global_cache_context(num_items=2):
return SimpleNamespace(
req_id="req",
num_items=num_items,
str_mm_hashes=[f"hash-{i}" for i in range(num_items)],
modality=Modality.IMAGE,
preprocess_result=SimpleNamespace(token_counts=[2] * num_items),
)
@staticmethod
def _make_prefix_cache_encoder_and_context(get_feature_fn):
encoder = MMEncoder.__new__(MMEncoder)
encoder.mm_cache = MultiModalStaticCache(1024 * 1024)
encoder.mm_cache_lock = asyncio.Lock()
item = SimpleNamespace(hash=123, set_pad_value=lambda: None)
encoder._build_model_mm_items = Mock(return_value=[item])
ctx = SimpleNamespace(
req_id="req",
modality=Modality.IMAGE,
num_items=1,
mm_feature=None,
preprocess_result=SimpleNamespace(token_counts=[2], mm_inputs={}),
get_feature_fn=get_feature_fn,
is_health_check=False,
items_per_req=[1],
aux_data={},
)
return encoder, ctx
def test_invalid_fresh_embedding_is_not_cached(self):
async def run():
for actual_tokens in (1, 3):
with self.subTest(actual_tokens=actual_tokens):
get_feature_fn = Mock(return_value=torch.zeros((actual_tokens, 4)))
encoder, ctx = self._make_prefix_cache_encoder_and_context(
get_feature_fn
)
with (
patch(
"sglang.srt.disaggregation.encoder.server.get_mm",
return_value=SimpleNamespace(enable_prefix_mm_cache=True),
),
self.assertRaisesRegex(
InternalError,
f"Encoder produced {actual_tokens} tokens, but "
"preprocessor metadata expected 2",
),
):
await encoder._compute_direct_embedding(ctx, keep_on_gpu=False)
self.assertEqual(len(encoder.mm_cache), 0)
get_feature_fn.assert_called_once()
asyncio.run(run())
def test_valid_fresh_embedding_is_cached_and_reused(self):
async def run():
get_feature_fn = Mock(return_value=torch.zeros((2, 4)))
encoder, ctx = self._make_prefix_cache_encoder_and_context(get_feature_fn)
with patch(
"sglang.srt.disaggregation.encoder.server.get_mm",
return_value=SimpleNamespace(enable_prefix_mm_cache=True),
):
first = await encoder._compute_direct_embedding(ctx, keep_on_gpu=False)
second = await encoder._compute_direct_embedding(ctx, keep_on_gpu=False)
torch.testing.assert_close(first, second)
self.assertEqual(len(encoder.mm_cache), 1)
get_feature_fn.assert_called_once()
asyncio.run(run())
def test_invalid_cached_embedding_is_evicted(self):
async def run():
get_feature_fn = Mock()
encoder, ctx = self._make_prefix_cache_encoder_and_context(get_feature_fn)
mm_hash = MultiModalStaticCache.combine_hashes([123])
encoder.mm_cache.set(
mm_hash,
EmbeddingResult(embedding=torch.zeros((1, 4))),
)
with (
patch(
"sglang.srt.disaggregation.encoder.server.get_mm",
return_value=SimpleNamespace(enable_prefix_mm_cache=True),
),
self.assertRaisesRegex(
InternalError,
"Encoder produced 1 tokens, but preprocessor metadata expected 2",
),
):
await encoder._compute_direct_embedding(ctx, keep_on_gpu=False)
self.assertEqual(len(encoder.mm_cache), 0)
get_feature_fn.assert_not_called()
asyncio.run(run())
def test_background_task_failure_is_observed(self):
async def run():
encoder = MMEncoder.__new__(MMEncoder)
encoder.background_tasks = set()
async def fail():
raise RuntimeError("background failure")
with patch(
"sglang.srt.disaggregation.encoder.server.logger.exception"
) as log_exception:
task = encoder._create_background_task(fail())
await asyncio.sleep(0)
await asyncio.sleep(0)
self.assertTrue(task.done())
self.assertNotIn(task, encoder.background_tasks)
log_exception.assert_called_once_with("MMEncoder background task failed")
asyncio.run(run())
def test_contract_has_two_direct_implementations(self): def test_contract_has_two_direct_implementations(self):
self.assertEqual(EncoderDelivery.__abstractmethods__, {"send", "release"}) self.assertEqual(EncoderDelivery.__abstractmethods__, {"send", "release"})
self.assertEqual( self.assertEqual(
@@ -140,6 +271,116 @@ class TestEncoderDelivery(CustomTestCase):
}, },
) )
@staticmethod
def _make_mooncake_send(engine):
embedding = torch.zeros((2, 4), dtype=torch.float32)
mm_data = EmbeddingData(
"req",
1,
0,
None,
Modality.IMAGE,
embedding=embedding,
)
encoder = MMEncoder.__new__(MMEncoder)
encoder._element_size = embedding.element_size()
encoder.engine = engine
return encoder, embedding, mm_data
def test_mooncake_fallback_registration_is_released_after_transfer_error(self):
async def run():
events = []
def register(*_):
events.append("register")
def transfer_sync(*_):
events.append("transfer")
raise RuntimeError("transfer failed")
def deregister(*_):
events.append("deregister")
engine = SimpleNamespace(
register=register,
transfer_sync=transfer_sync,
deregister=deregister,
)
encoder, embedding, mm_data = self._make_mooncake_send(engine)
with (
patch(
"sglang.srt.disaggregation.encoder.server.get_disagg",
return_value=SimpleNamespace(encoder_transfer_backend="mooncake"),
),
self.assertRaisesRegex(RuntimeError, "transfer failed"),
):
await encoder._send(
embedding,
mm_data,
session_id="session",
buffer_address=1,
)
self.assertEqual(events, ["register", "transfer", "deregister"])
asyncio.run(run())
def test_mooncake_cancel_waits_before_releasing_fallback_registration(self):
async def run():
events = []
transfer_started = threading.Event()
finish_transfer = threading.Event()
def register(*_):
events.append("register")
def transfer_sync(*_):
events.append("transfer-start")
transfer_started.set()
finish_transfer.wait(timeout=2)
events.append("transfer-finish")
return 0
def deregister(*_):
events.append("deregister")
engine = SimpleNamespace(
register=register,
transfer_sync=transfer_sync,
deregister=deregister,
)
encoder, embedding, mm_data = self._make_mooncake_send(engine)
with patch(
"sglang.srt.disaggregation.encoder.server.get_disagg",
return_value=SimpleNamespace(encoder_transfer_backend="mooncake"),
):
send_task = asyncio.create_task(
encoder._send(
embedding,
mm_data,
session_id="session",
buffer_address=1,
)
)
self.assertTrue(await asyncio.to_thread(transfer_started.wait, 1))
send_task.cancel()
await asyncio.sleep(0)
self.assertFalse(send_task.done())
self.assertNotIn("deregister", events)
finish_transfer.set()
with self.assertRaises(asyncio.CancelledError):
await send_task
self.assertEqual(
events,
["register", "transfer-start", "transfer-finish", "deregister"],
)
asyncio.run(run())
def test_zmq_delivery_cleanup_is_configurable(self): def test_zmq_delivery_cleanup_is_configurable(self):
async def run(): async def run():
req_id = "test-zmq-delivery-cleanup" req_id = "test-zmq-delivery-cleanup"
@@ -214,6 +455,90 @@ class TestEncoderDelivery(CustomTestCase):
asyncio.run(run()) asyncio.run(run())
def test_global_cache_lookup_failure_falls_back_to_all_misses(self):
async def run():
encoder = MMEncoder.__new__(MMEncoder)
encoder.rank = 0
encoder.mm_global_cache = SimpleNamespace(
batch_is_exist=AsyncMock(side_effect=RuntimeError("store down"))
)
encoder._broadcast_global_cache_mask = unittest.mock.Mock()
missing_indices, hit_indices = await encoder._lookup_global_cache(
self._global_cache_context()
)
self.assertEqual(missing_indices, [0, 1])
self.assertEqual(hit_indices, [])
torch.testing.assert_close(
encoder._broadcast_global_cache_mask.call_args.args[0],
torch.zeros(2, dtype=torch.int32),
)
asyncio.run(run())
def test_global_cache_prefetch_failure_immediately_falls_back(self):
async def run():
encoder = MMEncoder.__new__(MMEncoder)
encoder.rank = 0
encoder.mm_global_cache = SimpleNamespace(
prefetch=unittest.mock.Mock(side_effect=RuntimeError("store down"))
)
encoder._broadcast_global_cache_mask = unittest.mock.Mock()
ctx = self._global_cache_context()
hit_hashes, failed = encoder._prefetch_global_cache_hits(ctx, [0, 1])
fallback_indices = await encoder._wait_global_cache_prefetch(
ctx, [0, 1], hit_hashes, failed
)
self.assertTrue(failed)
self.assertEqual(hit_hashes, [])
self.assertEqual(fallback_indices, [0, 1])
torch.testing.assert_close(
encoder._broadcast_global_cache_mask.call_args.args[0],
torch.ones(2, dtype=torch.int32),
)
asyncio.run(run())
def test_global_cache_staging_failure_skips_insert(self):
encoder = MMEncoder.__new__(MMEncoder)
encoder.mm_global_cache = SimpleNamespace(
store_to_pool_async=unittest.mock.Mock(
side_effect=RuntimeError("pool full")
)
)
hashes, handles = encoder._stage_global_cache_slices(
self._global_cache_context(num_items=1),
[0],
[torch.ones((2, 4))],
)
self.assertEqual(hashes, [])
self.assertEqual(handles, [])
def test_global_cache_insert_failure_is_contained(self):
async def run():
encoder = MMEncoder.__new__(MMEncoder)
encoder.background_tasks = set()
encoder.mm_global_cache = SimpleNamespace(
wait_store_to_pool=unittest.mock.Mock(
side_effect=RuntimeError("store down")
),
insert_batch=unittest.mock.Mock(),
)
encoder._launch_global_cache_insert(
self._global_cache_context(num_items=1), ["hash-0"], [object()]
)
await asyncio.gather(*encoder.background_tasks)
encoder.mm_global_cache.insert_batch.assert_not_called()
asyncio.run(run())
def test_mooncake_embedding_is_ready_only_after_cuda_sync(self): def test_mooncake_embedding_is_ready_only_after_cuda_sync(self):
class FakeCudaEmbedding: class FakeCudaEmbedding:
shape = (2, 4) shape = (2, 4)
@@ -228,7 +553,9 @@ class TestEncoderDelivery(CustomTestCase):
encoder = MMEncoder.__new__(MMEncoder) encoder = MMEncoder.__new__(MMEncoder)
encoder.rank = 0 encoder.rank = 0
events = [] events = []
encoder._stage_embedding = lambda mm_data: events.append("ready") state = ReqState("req", active_encodes=1)
state.embedding_ready = SimpleNamespace(set=lambda: events.append("ready"))
encoder.req_states = {"req": state}
ctx = SimpleNamespace( ctx = SimpleNamespace(
req_id="req", req_id="req",
modality=Modality.IMAGE, modality=Modality.IMAGE,
@@ -252,6 +579,71 @@ class TestEncoderDelivery(CustomTestCase):
self.assertEqual(events, ["sync", "ready"]) self.assertEqual(events, ["sync", "ready"])
def test_shared_mr_registration_failure_keeps_send_fallback_enabled(self):
encoder = MMEncoder.__new__(MMEncoder)
encoder.engine = unittest.mock.Mock()
encoder.engine.register.side_effect = RuntimeError("register failed")
embedding = torch.ones((2, 4))
mm_data = EmbeddingData(
"req", 1, 0, [[1, 1, 1]], Modality.IMAGE, embedding=embedding
)
encoder._register_shared_mr(mm_data, embedding)
self.assertIsNone(mm_data._mr_ptr)
def test_fused_staging_rolls_back_mrs_before_any_result_is_published(self):
encoder = MMEncoder.__new__(MMEncoder)
encoder.rank = 0
encoder.engine = Mock()
first_state = ReqState("req-0", active_encodes=1)
second_metadata = EmbeddingData(
"req-1",
1,
0,
[[1, 1, 1]],
Modality.IMAGE,
embedding_shape=[2, 4],
dtype=torch.float32,
)
second_state = ReqState(
"req-1", embedding_data=second_metadata, active_encodes=1
)
encoder.req_states = {"req-0": first_state, "req-1": second_state}
ctx = SimpleNamespace(
req_id="req-0",
modality=Modality.IMAGE,
items_per_req=[1, 1],
preprocess_result=SimpleNamespace(
token_counts=[1, 1],
grid_thw=[[1, 1, 1], [1, 1, 1]],
),
aux_data={},
use_global_cache=False,
)
requests = [
{"req_id": "req-0", "num_parts": 1, "part_idx": 0},
{"req_id": "req-1", "num_parts": 1, "part_idx": 0},
]
with self.assertRaisesRegex(InternalError, "Embedding metadata mismatch"):
encoder._stage_embeddings(
ctx, requests, torch.ones((2, 4)), keep_on_gpu=True
)
registered_ptrs = [
call.args[0] for call in encoder.engine.register.call_args_list
]
deregistered_ptrs = [
call.args[0] for call in encoder.engine.deregister.call_args_list
]
self.assertEqual(len(registered_ptrs), 2)
self.assertCountEqual(deregistered_ptrs, registered_ptrs)
self.assertIsNone(first_state.embedding_data)
self.assertIs(second_state.embedding_data, second_metadata)
self.assertFalse(first_state.embedding_ready.is_set())
self.assertFalse(second_state.embedding_ready.is_set())
def test_stage_embedding_does_not_resurrect_missing_state(self): def test_stage_embedding_does_not_resurrect_missing_state(self):
encoder = MMEncoder.__new__(MMEncoder) encoder = MMEncoder.__new__(MMEncoder)
encoder.req_states = {} encoder.req_states = {}
@@ -535,5 +927,26 @@ class TestEncoderDelivery(CustomTestCase):
asyncio.run(run()) asyncio.run(run())
class TestMooncakeRegistration(CustomTestCase):
def setUp(self):
self.engine = MooncakeTransferEngine.__new__(MooncakeTransferEngine)
self.engine.engine = unittest.mock.Mock()
def test_register_raises_on_nonzero_status(self):
self.engine.engine.register_memory.return_value = -1
with self.assertRaisesRegex(RuntimeError, "registration failed.*ret=-1"):
self.engine.register(1234, 4096)
def test_deregister_preserves_backend_failure(self):
backend_error = OSError("backend failed")
self.engine.engine.unregister_memory.side_effect = backend_error
with self.assertRaisesRegex(RuntimeError, "deregistration failed") as ctx:
self.engine.deregister(1234)
self.assertIs(ctx.exception.__cause__, backend_error)
if __name__ == "__main__": if __name__ == "__main__":
unittest.main() unittest.main()
@@ -250,6 +250,35 @@ class TestStoreToPool(unittest.TestCase):
with self.assertRaises(ValueError): with self.assertRaises(ValueError):
ctrl.store_to_pool_async(["h"], [tensor], Modality.IMAGE) ctrl.store_to_pool_async(["h"], [tensor], Modality.IMAGE)
def test_wait_store_isolates_failed_copy_and_waits_for_the_rest(self):
ctrl = _make_controller(num_pages=4, dim=4, page_size=2)
entries = []
for mm_hash in ("first", "failed", "last"):
page_runs = ctrl.vision_pool.allocator.allocate(2, 2)
entry = EmbeddingCacheEntry(
hash=mm_hash,
modality=Modality.IMAGE,
num_tokens=2,
dim=4,
page_runs=page_runs,
state=EntryState.FILLING,
)
ctrl.entries[mm_hash] = entry
entries.append(entry)
handles = [MagicMock(), MagicMock(), MagicMock()]
handles[1].wait.side_effect = RuntimeError("D2H copy failed")
with self.assertRaisesRegex(RuntimeError, "D2H copy failed"):
ctrl.wait_store_to_pool(list(zip(entries, handles)))
for handle in handles:
handle.wait.assert_called_once_with()
self.assertEqual(ctrl.entries["first"].state, EntryState.READY)
self.assertNotIn("failed", ctrl.entries)
self.assertEqual(ctrl.entries["last"].state, EntryState.READY)
self.assertEqual(ctrl.vision_pool.allocator.free_pages, 2)
def _insert_ready_entry(ctrl, mm_hash, tensor, modality=Modality.IMAGE): def _insert_ready_entry(ctrl, mm_hash, tensor, modality=Modality.IMAGE):
"""Manually write tensor into pool pages and create a READY entry.""" """Manually write tensor into pool pages and create a READY entry."""