Add bootstrap_room validation to detect metadata corruption in PD disaggregation (#17430)
Co-authored-by: 继优 <jiyou.ljy@alibaba-inc.com> Co-authored-by: Shangming Cai <csmthu@gmail.com>
This commit is contained in:
co-authored by
继优
Shangming Cai
parent
c84cd4b5ff
commit
8ed35df204
@@ -722,7 +722,12 @@ class DecodeTransferQueue:
|
|||||||
def extend(self, decode_reqs: List[DecodeRequest]) -> None:
|
def extend(self, decode_reqs: List[DecodeRequest]) -> None:
|
||||||
self.queue.extend(decode_reqs)
|
self.queue.extend(decode_reqs)
|
||||||
|
|
||||||
def _commit_transfer_to_req(self, decode_req: DecodeRequest) -> None:
|
def _commit_transfer_to_req(self, decode_req: DecodeRequest) -> bool:
|
||||||
|
"""
|
||||||
|
Returns:
|
||||||
|
True if the request should be removed from the queue (success or corruption)
|
||||||
|
False if metadata not ready yet (keep in queue for next poll)
|
||||||
|
"""
|
||||||
idx = decode_req.metadata_buffer_index
|
idx = decode_req.metadata_buffer_index
|
||||||
(
|
(
|
||||||
output_id,
|
output_id,
|
||||||
@@ -734,8 +739,48 @@ class DecodeTransferQueue:
|
|||||||
output_topk_p,
|
output_topk_p,
|
||||||
output_topk_index,
|
output_topk_index,
|
||||||
output_hidden_states,
|
output_hidden_states,
|
||||||
|
output_bootstrap_room,
|
||||||
) = self.metadata_buffers.get_buf(idx)
|
) = self.metadata_buffers.get_buf(idx)
|
||||||
|
|
||||||
|
# Validate bootstrap_room to detect context corruption
|
||||||
|
actual_room = output_bootstrap_room[0].item()
|
||||||
|
expected_room = (
|
||||||
|
decode_req.req.bootstrap_room
|
||||||
|
if decode_req.req.bootstrap_room is not None
|
||||||
|
else 0
|
||||||
|
)
|
||||||
|
|
||||||
|
if decode_req.req.bootstrap_host == FAKE_BOOTSTRAP_HOST or (
|
||||||
|
decode_req.req.bootstrap_host is None
|
||||||
|
and self.scheduler.server_args.disaggregation_decode_enable_fake_auto
|
||||||
|
):
|
||||||
|
# Warm up or fake transfer mode
|
||||||
|
pass
|
||||||
|
elif actual_room == 0:
|
||||||
|
# Case 1: Metadata not ready yet (actual_room == 0)
|
||||||
|
# Keep request in queue and wait for next poll
|
||||||
|
return False
|
||||||
|
elif actual_room != expected_room:
|
||||||
|
# Case 2: Real corruption detected (mismatch)
|
||||||
|
# Abort the request and remove from the queue
|
||||||
|
error_msg = (
|
||||||
|
f"Context corruption detected: Request {decode_req.req.rid} "
|
||||||
|
f"(bootstrap_room={expected_room}) received metadata from "
|
||||||
|
f"bootstrap_room={actual_room}. "
|
||||||
|
f"Metadata buffer index: {idx}. "
|
||||||
|
f"This indicates metadata buffer index collision."
|
||||||
|
)
|
||||||
|
logger.error(error_msg)
|
||||||
|
prepare_abort(
|
||||||
|
decode_req.req,
|
||||||
|
"Metadata corruption detected - bootstrap_room mismatch",
|
||||||
|
status_code=HTTPStatus.INTERNAL_SERVER_ERROR,
|
||||||
|
)
|
||||||
|
decode_req.kv_receiver.clear()
|
||||||
|
decode_req.kv_receiver = None
|
||||||
|
return True
|
||||||
|
|
||||||
|
# Case 3: Success - commit the transfer
|
||||||
decode_req.req.output_ids.append(output_id[0].item())
|
decode_req.req.output_ids.append(output_id[0].item())
|
||||||
decode_req.req.cached_tokens = cached_tokens[0].item()
|
decode_req.req.cached_tokens = cached_tokens[0].item()
|
||||||
if not self.spec_algorithm.is_none():
|
if not self.spec_algorithm.is_none():
|
||||||
@@ -765,6 +810,7 @@ class DecodeTransferQueue:
|
|||||||
auto_next_anon=True,
|
auto_next_anon=True,
|
||||||
)
|
)
|
||||||
decode_req.req.time_stats.wait_queue_entry_time = time.perf_counter()
|
decode_req.req.time_stats.wait_queue_entry_time = time.perf_counter()
|
||||||
|
return True
|
||||||
|
|
||||||
def pop_transferred(self, rids_to_check: Optional[List[str]] = None) -> List[Req]:
|
def pop_transferred(self, rids_to_check: Optional[List[str]] = None) -> List[Req]:
|
||||||
if not self.queue:
|
if not self.queue:
|
||||||
@@ -800,8 +846,20 @@ class DecodeTransferQueue:
|
|||||||
self.scheduler.metrics_collector.increment_transfer_failed_reqs()
|
self.scheduler.metrics_collector.increment_transfer_failed_reqs()
|
||||||
continue
|
continue
|
||||||
elif poll == KVPoll.Success:
|
elif poll == KVPoll.Success:
|
||||||
self._commit_transfer_to_req(decode_req)
|
should_remove = self._commit_transfer_to_req(decode_req)
|
||||||
|
if should_remove:
|
||||||
indices_to_remove.add(i)
|
indices_to_remove.add(i)
|
||||||
|
# Check if request was aborted due to corruption
|
||||||
|
if isinstance(decode_req.req.finished_reason, FINISH_ABORT):
|
||||||
|
self.scheduler.stream_output(
|
||||||
|
[decode_req.req], decode_req.req.return_logprob
|
||||||
|
)
|
||||||
|
release_kv_cache(
|
||||||
|
decode_req.req, self.tree_cache, is_insert=False
|
||||||
|
)
|
||||||
|
if self.scheduler.enable_metrics:
|
||||||
|
self.scheduler.metrics_collector.increment_transfer_failed_reqs()
|
||||||
|
else:
|
||||||
transferred_reqs.append(decode_req.req)
|
transferred_reqs.append(decode_req.req)
|
||||||
elif poll in [
|
elif poll in [
|
||||||
KVPoll.Bootstrapping,
|
KVPoll.Bootstrapping,
|
||||||
|
|||||||
@@ -132,6 +132,10 @@ class MetadataBuffers:
|
|||||||
self.output_hidden_states = torch.zeros(
|
self.output_hidden_states = torch.zeros(
|
||||||
(size, hidden_size), dtype=hidden_states_dtype, device=device
|
(size, hidden_size), dtype=hidden_states_dtype, device=device
|
||||||
)
|
)
|
||||||
|
# Request validation: store bootstrap_room to detect metadata corruption
|
||||||
|
self.bootstrap_room = torch.zeros(
|
||||||
|
(size, 8), dtype=torch.int64, device=device
|
||||||
|
)
|
||||||
|
|
||||||
def get_buf_infos(self):
|
def get_buf_infos(self):
|
||||||
ptrs = [
|
ptrs = [
|
||||||
@@ -144,6 +148,7 @@ class MetadataBuffers:
|
|||||||
self.output_topk_p.data_ptr(),
|
self.output_topk_p.data_ptr(),
|
||||||
self.output_topk_index.data_ptr(),
|
self.output_topk_index.data_ptr(),
|
||||||
self.output_hidden_states.data_ptr(),
|
self.output_hidden_states.data_ptr(),
|
||||||
|
self.bootstrap_room.data_ptr(),
|
||||||
]
|
]
|
||||||
data_lens = [
|
data_lens = [
|
||||||
self.output_ids.nbytes,
|
self.output_ids.nbytes,
|
||||||
@@ -155,6 +160,7 @@ class MetadataBuffers:
|
|||||||
self.output_topk_p.nbytes,
|
self.output_topk_p.nbytes,
|
||||||
self.output_topk_index.nbytes,
|
self.output_topk_index.nbytes,
|
||||||
self.output_hidden_states.nbytes,
|
self.output_hidden_states.nbytes,
|
||||||
|
self.bootstrap_room.nbytes,
|
||||||
]
|
]
|
||||||
item_lens = [
|
item_lens = [
|
||||||
self.output_ids[0].nbytes,
|
self.output_ids[0].nbytes,
|
||||||
@@ -166,6 +172,7 @@ class MetadataBuffers:
|
|||||||
self.output_topk_p[0].nbytes,
|
self.output_topk_p[0].nbytes,
|
||||||
self.output_topk_index[0].nbytes,
|
self.output_topk_index[0].nbytes,
|
||||||
self.output_hidden_states[0].nbytes,
|
self.output_hidden_states[0].nbytes,
|
||||||
|
self.bootstrap_room[0].nbytes,
|
||||||
]
|
]
|
||||||
return ptrs, data_lens, item_lens
|
return ptrs, data_lens, item_lens
|
||||||
|
|
||||||
@@ -180,6 +187,7 @@ class MetadataBuffers:
|
|||||||
self.output_topk_p[idx],
|
self.output_topk_p[idx],
|
||||||
self.output_topk_index[idx],
|
self.output_topk_index[idx],
|
||||||
self.output_hidden_states[idx],
|
self.output_hidden_states[idx],
|
||||||
|
self.bootstrap_room[idx],
|
||||||
)
|
)
|
||||||
|
|
||||||
def set_buf(self, req: Req):
|
def set_buf(self, req: Req):
|
||||||
@@ -222,6 +230,10 @@ class MetadataBuffers:
|
|||||||
self.output_hidden_states[req.metadata_buffer_index].copy_(
|
self.output_hidden_states[req.metadata_buffer_index].copy_(
|
||||||
req.hidden_states_tensor
|
req.hidden_states_tensor
|
||||||
)
|
)
|
||||||
|
# Store bootstrap_room for validation on decode side
|
||||||
|
self.bootstrap_room[req.metadata_buffer_index, 0] = (
|
||||||
|
req.bootstrap_room if req.bootstrap_room is not None else 0
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
#########################
|
#########################
|
||||||
|
|||||||
Reference in New Issue
Block a user