diff --git a/python/sglang/srt/disaggregation/encoder/runtime.py b/python/sglang/srt/disaggregation/encoder/runtime.py index f8ee1250a..67b9b3a1c 100644 --- a/python/sglang/srt/disaggregation/encoder/runtime.py +++ b/python/sglang/srt/disaggregation/encoder/runtime.py @@ -1028,9 +1028,7 @@ async def _push_embedding_to_prefill( if backend == "zmq_to_scheduler" and request.get("embedding_port") is None: send_coro = enc.send_with_url(req_id=req_id) if background_url_send: - task = asyncio.create_task(send_coro) - enc.background_tasks.add(task) - task.add_done_callback(enc.background_tasks.discard) + enc._create_background_task(send_coro) else: await send_coro return diff --git a/python/sglang/srt/disaggregation/encoder/server.py b/python/sglang/srt/disaggregation/encoder/server.py index f76cef67a..bfc154f94 100644 --- a/python/sglang/srt/disaggregation/encoder/server.py +++ b/python/sglang/srt/disaggregation/encoder/server.py @@ -350,15 +350,8 @@ class MooncakeDelivery(EncoderDelivery): async def release(self, state: ReqState) -> None: mm_data = state.embedding_data - if mm_data is not None and mm_data._mr_ptr is not None: - try: - 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 + if mm_data is not None: + self.encoder._deregister_shared_mr(mm_data) class ZmqDelivery(EncoderDelivery): @@ -607,6 +600,21 @@ class MMEncoder: 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: return self.preprocessor.supports_modality(modality) @@ -646,7 +654,7 @@ class MMEncoder: if should_release: 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) metadata = state.embedding_data if ( @@ -660,9 +668,21 @@ class MMEncoder: f"expected={metadata.shape}/{metadata.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_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: await state.embedding_ready.wait() if state.embedding_data is None: @@ -1046,7 +1066,17 @@ class MMEncoder: ctx: EncodeContext, ) -> Tuple[List[int], List[int]]: 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( [1 if e else 0 for e in exist_mask], dtype=torch.int32 ) @@ -1064,47 +1094,64 @@ class MMEncoder: self, ctx: EncodeContext, hit_indices: List[int], - ) -> List[str]: + ) -> Tuple[List[str], bool]: if self.rank != 0 or not hit_indices: - return [] + return [], False 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] - self.mm_global_cache.prefetch(ctx.req_id, hit_hashes, hit_tokens, ctx.modality) - return hit_hashes + try: + 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( self, ctx: EncodeContext, hit_indices: List[int], hit_hashes: List[str], + prefetch_failed: bool, ) -> List[int]: fallback_mask = torch.zeros(ctx.num_items, dtype=torch.int32) if self.rank == 0 and hit_indices: - try: - - async def _wait_prefetch(): - while not self.mm_global_cache.check_prefetch_progress(ctx.req_id): - await asyncio.sleep(0.005) - - await asyncio.wait_for(_wait_prefetch(), timeout=60.0) - - for i, idx in enumerate(hit_indices): - if not self.mm_global_cache.has_local_embedding(hit_hashes[i]): - fallback_mask[idx] = 1 - num_partial_fail = int(fallback_mask.sum().item()) - if num_partial_fail > 0: - logger.warning( - f"Req {ctx.req_id}: {num_partial_fail}/{len(hit_indices)} " - f"cache-hit items failed to load, falling back to ViT" - ) - except (asyncio.TimeoutError, Exception) as e: - logger.error( - f"Prefetch failed for req {ctx.req_id}: {e}. " - f"Falling back to ViT for {len(hit_indices)} hit items." - ) + if prefetch_failed: for idx in hit_indices: fallback_mask[idx] = 1 + else: + try: + + async def _wait_prefetch(): + while not self.mm_global_cache.check_prefetch_progress( + ctx.req_id + ): + await asyncio.sleep(0.005) + + await asyncio.wait_for(_wait_prefetch(), timeout=60.0) + + for i, idx in enumerate(hit_indices): + if not self.mm_global_cache.has_local_embedding(hit_hashes[i]): + fallback_mask[idx] = 1 + num_partial_fail = int(fallback_mask.sum().item()) + if num_partial_fail > 0: + logger.warning( + f"Req {ctx.req_id}: {num_partial_fail}/{len(hit_indices)} " + f"cache-hit items failed to load, falling back to ViT" + ) + except Exception as e: + logger.error( + f"Prefetch failed for req {ctx.req_id}: {e}. " + f"Falling back to ViT for {len(hit_indices)} hit items." + ) + for idx in hit_indices: + fallback_mask[idx] = 1 self._broadcast_global_cache_mask(fallback_mask) fallback_indices = [ @@ -1112,6 +1159,28 @@ class MMEncoder: ] 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( self, ctx: EncodeContext, @@ -1122,19 +1191,22 @@ class MMEncoder: return async def _background_insert(): - await asyncio.to_thread( - self.mm_global_cache.wait_store_to_pool, - d2h_handles, - ) - await asyncio.to_thread( - self.mm_global_cache.insert_batch, - hashes, - ctx.modality, - ) + try: + await asyncio.to_thread( + self.mm_global_cache.wait_store_to_pool, + d2h_handles, + ) + await asyncio.to_thread( + self.mm_global_cache.insert_batch, + hashes, + ctx.modality, + ) + except Exception: + logger.exception( + "Global multimodal cache insert failed for req %s", ctx.req_id + ) - task = asyncio.create_task(_background_insert()) - self.background_tasks.add(task) - task.add_done_callback(self.background_tasks.discard) + self._create_background_task(_background_insert()) @staticmethod def _as_2d_tensor(tensor: torch.Tensor) -> torch.Tensor: @@ -1262,7 +1334,7 @@ class MMEncoder: ) -> Optional[torch.Tensor]: """Resolve cache hits, compute misses, assemble output, and insert misses.""" 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 = [] if missing_indices: @@ -1274,19 +1346,20 @@ class MMEncoder: ctx.get_feature_fn, ) + miss_hashes = [] miss_d2h_handles = [] # 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: - miss_hashes = [ctx.str_mm_hashes[i] for i in missing_indices] - miss_d2h_handles = self.mm_global_cache.store_to_pool_async( - miss_hashes, new_slices, ctx.modality + miss_hashes, miss_d2h_handles = self._stage_global_cache_slices( + ctx, missing_indices, new_slices ) fallback_indices = await self._wait_global_cache_prefetch( - ctx, hit_indices, hit_hashes + ctx, hit_indices, hit_hashes, prefetch_failed ) fallback_slices = [] + fallback_hashes = [] fallback_d2h_handles = [] if fallback_indices: logger.info( @@ -1301,9 +1374,8 @@ class MMEncoder: ctx.get_feature_fn, ) if self.rank == 0 and not keep_on_gpu: - fallback_hashes = [ctx.str_mm_hashes[i] for i in fallback_indices] - fallback_d2h_handles = self.mm_global_cache.store_to_pool_async( - fallback_hashes, fallback_slices, ctx.modality + fallback_hashes, fallback_d2h_handles = self._stage_global_cache_slices( + ctx, fallback_indices, fallback_slices ) if self.rank == 0: @@ -1311,14 +1383,14 @@ class MMEncoder: # Start staging newly computed GPU slices into the CPU cache # pool asynchronously before assembling the GPU output. if new_slices: - miss_hashes = [ctx.str_mm_hashes[i] for i in missing_indices] - miss_d2h_handles = self.mm_global_cache.store_to_pool_async( - miss_hashes, new_slices, ctx.modality + miss_hashes, miss_d2h_handles = self._stage_global_cache_slices( + ctx, missing_indices, new_slices ) if fallback_slices: - fallback_hashes = [ctx.str_mm_hashes[i] for i in fallback_indices] - fallback_d2h_handles = self.mm_global_cache.store_to_pool_async( - fallback_hashes, fallback_slices, ctx.modality + fallback_hashes, fallback_d2h_handles = ( + self._stage_global_cache_slices( + ctx, fallback_indices, fallback_slices + ) ) mm_embedding = self._assemble_global_cache_gpu( ctx, @@ -1337,11 +1409,9 @@ class MMEncoder: 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( ctx, - new_hashes, + miss_hashes + fallback_hashes, miss_d2h_handles + fallback_d2h_handles, ) return mm_embedding @@ -1401,6 +1471,17 @@ class MMEncoder: 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. if use_mm_cache and encoder_metrics_collector is not None: total_tokens = int(mm_embedding.shape[0]) @@ -1479,18 +1560,21 @@ class MMEncoder: mm_embedding = await self._compute_global_cache_embedding( ctx, keep_on_gpu=keep_on_gpu ) - else: - mm_embedding = await self._compute_direct_embedding( - ctx, keep_on_gpu=keep_on_gpu - ) + if mm_embedding is not None: + self._validate_embedding_token_count(ctx, mm_embedding) + 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) - if mm_embedding is not None and mm_embedding.shape[0] != expected_tokens: + if mm_embedding.shape[0] != expected_tokens: raise InternalError( f"Encoder produced {mm_embedding.shape[0]} tokens, but " f"preprocessor metadata expected {expected_tokens}" ) - return mm_embedding async def _publish_preprocess_metadata( self, ctx: EncodeContext, requests: List[dict] @@ -1551,21 +1635,35 @@ class MMEncoder: mr_already_registered = mm_data._mr_ptr == embedding.data_ptr() if not mr_already_registered: self.engine.register(embedding.data_ptr(), embedding.nbytes) - _t_xfer_start = time.monotonic() - xfer_ret = await asyncio.to_thread( - self.engine.transfer_sync, - session_id, - embedding.data_ptr(), - buffer_address, - embedding.nbytes, - ) + transfer_error = None + try: + _t_xfer_start = time.monotonic() + xfer_ret = await self._run_mooncake_transfer( + session_id, + embedding.data_ptr(), + buffer_address, + 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 if encoder_metrics_collector is not None: encoder_metrics_collector.observe_transfer( xfer_ms / 1000.0, backend="mooncake" ) - if not mr_already_registered: - self.engine.deregister(embedding.data_ptr()) if xfer_ret < 0: raise InternalError( f"Mooncake transfer_sync failed for {req_id} " @@ -1684,6 +1782,32 @@ class MMEncoder: 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): """Register one MR shared by every rank's /send; _send re-registers on failure.""" try: @@ -1695,6 +1819,18 @@ class MMEncoder: 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( self, ctx: EncodeContext, @@ -1715,46 +1851,58 @@ class MMEncoder: results = [] staged_embeddings = [] - item_offset = 0 - token_offset = 0 - for req, num_items in zip(requests, ctx.items_per_req): - item_end = item_offset + num_items - num_tokens = sum(ctx.preprocess_result.token_counts[item_offset:item_end]) - embedding = mm_embedding[token_offset : token_offset + num_tokens] - if keep_on_gpu and len(requests) > 1: - # A view would pin the whole batch tensor until the last transfer. - embedding = embedding.clone() - req_aux_data = dict(ctx.aux_data) - if ctx.aux_data.get("original_image_sizes") is not None: - req_aux_data["original_image_sizes"] = ctx.aux_data[ - "original_image_sizes" - ][item_offset:item_end] - mm_data = EmbeddingData( - req["req_id"], - req["num_parts"], - req["part_idx"], - ctx.preprocess_result.grid_thw[item_offset:item_end], - ctx.modality, - embedding, - **req_aux_data, - ) - # Global-cache embeddings keep registering per /send instead. - if keep_on_gpu and not ctx.use_global_cache: - self._register_shared_mr(mm_data, embedding) - staged_embeddings.append(mm_data) - results.append( - (embedding.nbytes, embedding.shape[0], embedding.shape[1], None, None) - ) - item_offset = item_end - token_offset += num_tokens + try: + item_offset = 0 + token_offset = 0 + for req, num_items in zip(requests, ctx.items_per_req): + item_end = item_offset + num_items + num_tokens = sum( + ctx.preprocess_result.token_counts[item_offset:item_end] + ) + embedding = mm_embedding[token_offset : token_offset + num_tokens] + if keep_on_gpu and len(requests) > 1: + # A view would pin the whole batch tensor until the last transfer. + embedding = embedding.clone() + req_aux_data = dict(ctx.aux_data) + if ctx.aux_data.get("original_image_sizes") is not None: + req_aux_data["original_image_sizes"] = ctx.aux_data[ + "original_image_sizes" + ][item_offset:item_end] + mm_data = EmbeddingData( + req["req_id"], + req["num_parts"], + req["part_idx"], + ctx.preprocess_result.grid_thw[item_offset:item_end], + ctx.modality, + embedding, + **req_aux_data, + ) + # Global-cache embeddings keep registering per /send instead. + if keep_on_gpu and not ctx.use_global_cache: + self._register_shared_mr(mm_data, embedding) + staged_embeddings.append(mm_data) + results.append( + ( + embedding.nbytes, + embedding.shape[0], + embedding.shape[1], + None, + None, + ) + ) + item_offset = item_end + token_offset += num_tokens - # transfer_sync bypasses CUDA streams, so GPU writes (forward and the - # per-request clones) must land before /send reads the buffers. - if keep_on_gpu and mm_embedding.is_cuda: - torch.cuda.current_stream(mm_embedding.device).synchronize() - for mm_data in staged_embeddings: - self._stage_embedding(mm_data) - return results + # transfer_sync bypasses CUDA streams, so GPU writes (forward and the + # per-request clones) must land before /send reads the buffers. + if keep_on_gpu and mm_embedding.is_cuda: + torch.cuda.current_stream(mm_embedding.device).synchronize() + self._stage_embedding_batch(staged_embeddings) + return results + except BaseException: + for mm_data in staged_embeddings: + self._deregister_shared_mr(mm_data) + raise def _stage_errors( self, requests: List[dict], modality: Modality, exc: Exception diff --git a/python/sglang/srt/distributed/device_communicators/mooncake_transfer_engine.py b/python/sglang/srt/distributed/device_communicators/mooncake_transfer_engine.py index efbcfeafe..9ae94543a 100644 --- a/python/sglang/srt/distributed/device_communicators/mooncake_transfer_engine.py +++ b/python/sglang/srt/distributed/device_communicators/mooncake_transfer_engine.py @@ -143,25 +143,25 @@ class MooncakeTransferEngine: self.hostname, self.engine.get_rpc_port() ).to_host_port_str() - def register(self, ptr, length): + def register(self, ptr, length) -> None: try: ret_value = self.engine.register_memory(ptr, length) - except Exception: - # Mark register as failed - ret_value = -1 + except Exception as exc: + raise RuntimeError("Mooncake memory registration failed") from exc 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: ret_value = self.engine.unregister_memory(ptr) - except Exception: - # Mark deregister as failed - ret_value = -1 + except Exception as exc: + raise RuntimeError("Mooncake memory deregistration failed") from exc 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: """Batch register multiple memory regions.""" diff --git a/python/sglang/srt/mem_cache/embedding_cache_controller.py b/python/sglang/srt/mem_cache/embedding_cache_controller.py index cf5d1fee1..4a95d3635 100644 --- a/python/sglang/srt/mem_cache/embedding_cache_controller.py +++ b/python/sglang/srt/mem_cache/embedding_cache_controller.py @@ -979,13 +979,32 @@ class EmbeddingCacheController: handles: List[Tuple["EmbeddingCacheEntry", "AsyncCopyHandle"]], ): """Wait for async D2H copies and mark entries READY.""" + completed = [] + errors = [] for entry, handle in handles: - handle.wait() + try: + handle.wait() + completed.append(True) + except BaseException as error: + completed.append(False) + errors.append(error) + with self.lock: - for entry, handle in handles: + for (entry, _), success in zip(handles, completed, strict=True): current = self.entries.get(entry.hash) if current is entry and current.state == EntryState.FILLING: - self._mark_ready(current) + if success: + 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( self, tensor: torch.Tensor, entry: EmbeddingCacheEntry, pool: EmbeddingPool diff --git a/test/registered/unit/disaggregation/test_encode_server.py b/test/registered/unit/disaggregation/test_encode_server.py index dfa7b4ae6..5994c2394 100644 --- a/test/registered/unit/disaggregation/test_encode_server.py +++ b/test/registered/unit/disaggregation/test_encode_server.py @@ -1,8 +1,9 @@ import asyncio import pickle +import threading import unittest from types import SimpleNamespace -from unittest.mock import AsyncMock, patch +from unittest.mock import AsyncMock, Mock, patch import numpy as np import torch @@ -23,7 +24,14 @@ from sglang.srt.disaggregation.encoder.server import ( rid_to_receive_count, 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.mem_cache.multimodal_cache import ( + EmbeddingResult, + MultiModalStaticCache, +) from sglang.srt.utils.common import safe_pickle_loads from sglang.test.ci.ci_register import register_cpu_ci from sglang.test.test_utils import CustomTestCase @@ -130,6 +138,129 @@ class TestEncoderPreprocessorKimiGrid(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): self.assertEqual(EncoderDelivery.__abstractmethods__, {"send", "release"}) 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): async def run(): req_id = "test-zmq-delivery-cleanup" @@ -214,6 +455,90 @@ class TestEncoderDelivery(CustomTestCase): 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): class FakeCudaEmbedding: shape = (2, 4) @@ -228,7 +553,9 @@ class TestEncoderDelivery(CustomTestCase): encoder = MMEncoder.__new__(MMEncoder) encoder.rank = 0 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( req_id="req", modality=Modality.IMAGE, @@ -252,6 +579,71 @@ class TestEncoderDelivery(CustomTestCase): 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): encoder = MMEncoder.__new__(MMEncoder) encoder.req_states = {} @@ -535,5 +927,26 @@ class TestEncoderDelivery(CustomTestCase): 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__": unittest.main() diff --git a/test/registered/unit/mem_cache/test_embedding_cache_controller.py b/test/registered/unit/mem_cache/test_embedding_cache_controller.py index 1fb4efd6e..fb87afb49 100644 --- a/test/registered/unit/mem_cache/test_embedding_cache_controller.py +++ b/test/registered/unit/mem_cache/test_embedding_cache_controller.py @@ -250,6 +250,35 @@ class TestStoreToPool(unittest.TestCase): with self.assertRaises(ValueError): 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): """Manually write tensor into pool pages and create a READY entry."""