fix(vlm): make EPD cache publication transactional (#36949)
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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,27 +1094,44 @@ 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:
|
||||
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):
|
||||
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)
|
||||
@@ -1098,7 +1145,7 @@ class MMEncoder:
|
||||
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:
|
||||
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."
|
||||
@@ -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,6 +1191,7 @@ class MMEncoder:
|
||||
return
|
||||
|
||||
async def _background_insert():
|
||||
try:
|
||||
await asyncio.to_thread(
|
||||
self.mm_global_cache.wait_store_to_pool,
|
||||
d2h_handles,
|
||||
@@ -1131,10 +1201,12 @@ class MMEncoder:
|
||||
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)
|
||||
transfer_error = None
|
||||
try:
|
||||
_t_xfer_start = time.monotonic()
|
||||
xfer_ret = await asyncio.to_thread(
|
||||
self.engine.transfer_sync,
|
||||
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,11 +1851,14 @@ class MMEncoder:
|
||||
|
||||
results = []
|
||||
staged_embeddings = []
|
||||
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])
|
||||
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.
|
||||
@@ -1743,7 +1882,13 @@ class MMEncoder:
|
||||
self._register_shared_mr(mm_data, embedding)
|
||||
staged_embeddings.append(mm_data)
|
||||
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
|
||||
token_offset += num_tokens
|
||||
@@ -1752,9 +1897,12 @@ class MMEncoder:
|
||||
# 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)
|
||||
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
|
||||
|
||||
@@ -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."""
|
||||
|
||||
@@ -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:
|
||||
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:
|
||||
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
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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."""
|
||||
|
||||
Reference in New Issue
Block a user