Ensure multi-node MM embedding cache consistency in insert_batch (#25959)
Signed-off-by: Michael Qiu <qiudayu.qdy@antgroup.com> Co-authored-by: Mike_Qiu <qiudayu.qdy@antgroup.com> Co-authored-by: Claude Opus 4.7 <noreply@anthropic.com> Co-authored-by: siyu <liusy58@linux.alibaba.com>
This commit is contained in:
co-authored by
Mike_Qiu
Claude Opus 4.7
siyu
parent
2dfbc3d781
commit
73c99e3361
@@ -190,18 +190,37 @@ class EmbeddingCacheController:
|
|||||||
def insert_batch(
|
def insert_batch(
|
||||||
self, image_hashes: List[str], embedding_tensors: List[torch.Tensor]
|
self, image_hashes: List[str], embedding_tensors: List[torch.Tensor]
|
||||||
):
|
):
|
||||||
"""Issues ONE batch PUT for all embeddings computed by this request."""
|
"""Issues ONE batch PUT for all embeddings computed by this request.
|
||||||
|
|
||||||
|
Note: Even if the embedding exists locally, we still push to Mooncake
|
||||||
|
to ensure multi-node cache consistency. Mooncake's batch_put has
|
||||||
|
built-in deduplication to avoid redundant transfers.
|
||||||
|
"""
|
||||||
keys, ptrs, sizes = [], [], []
|
keys, ptrs, sizes = [], [], []
|
||||||
|
local_hit_count = 0
|
||||||
|
new_count = 0
|
||||||
|
skipped_count = 0
|
||||||
|
|
||||||
with self.lock:
|
with self.lock:
|
||||||
for h, tensor in zip(image_hashes, embedding_tensors):
|
for h, tensor in zip(image_hashes, embedding_tensors):
|
||||||
if h in self.hash_to_metadata:
|
if h in self.hash_to_metadata:
|
||||||
|
# Local cache hit: ensure Mooncake has it
|
||||||
|
offset, num_tokens, dim, size_bytes = self.hash_to_metadata[h][:4]
|
||||||
|
|
||||||
|
# Still push to Mooncake for multi-node sharing
|
||||||
|
# (Mooncake batch_put will deduplicate if already exists)
|
||||||
|
keys.append(h)
|
||||||
|
ptrs.append(self.cpu_pool.data_ptr() + offset)
|
||||||
|
sizes.append(size_bytes)
|
||||||
|
local_hit_count += 1
|
||||||
continue
|
continue
|
||||||
|
|
||||||
|
# Local cache miss: allocate and copy
|
||||||
num_tokens, dim = tensor.shape[0], tensor.shape[1]
|
num_tokens, dim = tensor.shape[0], tensor.shape[1]
|
||||||
size_bytes = num_tokens * dim * self.element_size
|
size_bytes = num_tokens * dim * self.element_size
|
||||||
offset = self.allocator.allocate(size_bytes)
|
offset = self.allocator.allocate(size_bytes)
|
||||||
if offset is None:
|
if offset is None:
|
||||||
|
skipped_count += 1
|
||||||
continue
|
continue
|
||||||
|
|
||||||
# Copy to pinned pool for RDMA
|
# Copy to pinned pool for RDMA
|
||||||
@@ -216,10 +235,13 @@ class EmbeddingCacheController:
|
|||||||
keys.append(h)
|
keys.append(h)
|
||||||
ptrs.append(self.cpu_pool.data_ptr() + offset)
|
ptrs.append(self.cpu_pool.data_ptr() + offset)
|
||||||
sizes.append(size_bytes)
|
sizes.append(size_bytes)
|
||||||
|
new_count += 1
|
||||||
|
|
||||||
if keys:
|
if keys:
|
||||||
logger.info(
|
logger.info(
|
||||||
f"Global Cache: Inserting {len(keys)} new embeddings into Mooncake cluster."
|
f"Global Cache: Inserting {len(keys)} embeddings into Mooncake cluster "
|
||||||
|
f"({new_count} new, {local_hit_count} existing for replication, "
|
||||||
|
f"{skipped_count} skipped due to allocation failure)"
|
||||||
)
|
)
|
||||||
self.insert_queue.put(EmbeddingInsertOperation(keys, ptrs, sizes))
|
self.insert_queue.put(EmbeddingInsertOperation(keys, ptrs, sizes))
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user