fix(staging): make empty staging rings reusable (#37166)
This commit is contained in:
@@ -243,8 +243,12 @@ class StagingAllocator:
|
|||||||
self.alloc_order.pop(0)
|
self.alloc_order.pop(0)
|
||||||
|
|
||||||
if not self.allocations:
|
if not self.allocations:
|
||||||
|
# An empty ring makes the entire prior round reusable. Start a
|
||||||
|
# fresh round at offset zero so the watermark cannot stay stale.
|
||||||
|
self.round += 1
|
||||||
|
self.head = 0
|
||||||
self.watermark_round = self.round
|
self.watermark_round = self.round
|
||||||
self.watermark_tail = self.head
|
self.watermark_tail = 0
|
||||||
elif self.alloc_order:
|
elif self.alloc_order:
|
||||||
off, _, rnd = self.allocations[self.alloc_order[0]]
|
off, _, rnd = self.allocations[self.alloc_order[0]]
|
||||||
self.watermark_round = rnd
|
self.watermark_round = rnd
|
||||||
|
|||||||
@@ -108,8 +108,17 @@ class DecodeStagingHandler:
|
|||||||
if receiver is None or not receiver.bootstrap_infos:
|
if receiver is None or not receiver.bootstrap_infos:
|
||||||
return
|
return
|
||||||
key = tuple(str(bi) for bi in receiver.bootstrap_infos)
|
key = tuple(str(bi) for bi in receiver.bootstrap_infos)
|
||||||
if key not in self._wm_subscribers:
|
if key in self._wm_subscribers:
|
||||||
|
return
|
||||||
|
|
||||||
self._wm_subscribers[key] = (receiver, session_id)
|
self._wm_subscribers[key] = (receiver, session_id)
|
||||||
|
# Watermark state is per prefill session. Send the current value so a
|
||||||
|
# new session's first allocation cannot wait on a missed update.
|
||||||
|
self._send_watermark(
|
||||||
|
receiver,
|
||||||
|
session_id,
|
||||||
|
self.staging_allocator.get_watermark(),
|
||||||
|
)
|
||||||
|
|
||||||
def num_writers_for(self, receiver) -> int:
|
def num_writers_for(self, receiver) -> int:
|
||||||
"""Compute all TP and PP writers expected for a staging chunk."""
|
"""Compute all TP and PP writers expected for a staging chunk."""
|
||||||
@@ -443,24 +452,26 @@ class DecodeStagingHandler:
|
|||||||
return True
|
return True
|
||||||
|
|
||||||
def _free_and_send_watermark(
|
def _free_and_send_watermark(
|
||||||
self, alloc_id: int, decode_req: DecodeRequest
|
self, alloc_id: int, _decode_req: DecodeRequest
|
||||||
) -> None:
|
) -> None:
|
||||||
"""Free a staging allocation and broadcast watermark to all prefills."""
|
"""Free a staging allocation and broadcast watermark to all prefills."""
|
||||||
self.staging_allocator.free(alloc_id)
|
self.staging_allocator.free(alloc_id)
|
||||||
post_wm = self.staging_allocator.get_watermark()
|
post_wm = self.staging_allocator.get_watermark()
|
||||||
room = decode_req.req.bootstrap_room
|
for receiver, session_id in list(self._wm_subscribers.values()):
|
||||||
wm_round, wm_tail = post_wm
|
self._send_watermark(receiver, session_id, post_wm)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _send_watermark(receiver, session_id: str, watermark) -> None:
|
||||||
|
"""Send one allocator watermark to a registered prefill session."""
|
||||||
|
wm_round, wm_tail = watermark
|
||||||
wm_round_b = str(wm_round).encode("ascii")
|
wm_round_b = str(wm_round).encode("ascii")
|
||||||
wm_tail_b = str(wm_tail).encode("ascii")
|
wm_tail_b = str(wm_tail).encode("ascii")
|
||||||
for _key, (receiver, session_id) in list(self._wm_subscribers.items()):
|
|
||||||
sid_b = session_id.encode("ascii")
|
sid_b = session_id.encode("ascii")
|
||||||
for bootstrap_info in receiver.bootstrap_infos:
|
for bootstrap_info in receiver.bootstrap_infos:
|
||||||
try:
|
try:
|
||||||
sock, lock = receiver._connect_to_bootstrap_server(bootstrap_info)
|
sock, lock = receiver._connect_to_bootstrap_server(bootstrap_info)
|
||||||
with lock:
|
with lock:
|
||||||
sock.send_multipart(
|
sock.send_multipart([b"WATERMARK", wm_round_b, wm_tail_b, sid_b])
|
||||||
[b"WATERMARK", wm_round_b, wm_tail_b, sid_b]
|
|
||||||
)
|
|
||||||
except Exception:
|
except Exception:
|
||||||
pass
|
pass
|
||||||
|
|
||||||
|
|||||||
@@ -9,7 +9,11 @@ import torch
|
|||||||
|
|
||||||
from sglang.srt.disaggregation.base.conn import KVArgs, StateType
|
from sglang.srt.disaggregation.base.conn import KVArgs, StateType
|
||||||
from sglang.srt.disaggregation.common.conn import CommonKVManager
|
from sglang.srt.disaggregation.common.conn import CommonKVManager
|
||||||
|
from sglang.srt.disaggregation.common.staging_buffer import (
|
||||||
|
StagingAllocator,
|
||||||
|
)
|
||||||
from sglang.srt.disaggregation.common.staging_handler import (
|
from sglang.srt.disaggregation.common.staging_handler import (
|
||||||
|
DecodeStagingHandler,
|
||||||
handle_staging_req,
|
handle_staging_req,
|
||||||
)
|
)
|
||||||
from sglang.srt.disaggregation.common.utils import (
|
from sglang.srt.disaggregation.common.utils import (
|
||||||
@@ -220,6 +224,37 @@ class TestGroupConcurrentContiguous(unittest.TestCase):
|
|||||||
group_concurrent_contiguous(self._arr([1, 2, 3]), self._arr([1, 2]))
|
group_concurrent_contiguous(self._arr([1, 2, 3]), self._arr([1, 2]))
|
||||||
|
|
||||||
|
|
||||||
|
class TestStagingWatermark(unittest.TestCase):
|
||||||
|
@patch("sglang.srt.disaggregation.common.staging_buffer.StagingBuffer")
|
||||||
|
def test_empty_ring_restarts_at_zero(self, staging_buffer):
|
||||||
|
staging_buffer.return_value.data_ptr = 0
|
||||||
|
allocator = StagingAllocator(100, "cpu", 0)
|
||||||
|
alloc_id, _, _ = allocator.assign(60)
|
||||||
|
|
||||||
|
allocator.free(alloc_id)
|
||||||
|
|
||||||
|
self.assertEqual(allocator.get_watermark(), (1, 0))
|
||||||
|
self.assertEqual(allocator.assign(70)[1:], (0, 1))
|
||||||
|
|
||||||
|
def test_new_watermark_subscriber_receives_current_allocator_state(self):
|
||||||
|
sock = Mock()
|
||||||
|
bootstrap_info = {"host": "prefill", "port": 7200}
|
||||||
|
receiver = Mock(
|
||||||
|
bootstrap_infos=[bootstrap_info],
|
||||||
|
)
|
||||||
|
receiver._connect_to_bootstrap_server.return_value = (sock, threading.Lock())
|
||||||
|
handler = object.__new__(DecodeStagingHandler)
|
||||||
|
handler.staging_allocator = Mock()
|
||||||
|
handler.staging_allocator.get_watermark.return_value = (3, 0)
|
||||||
|
handler._wm_subscribers = {}
|
||||||
|
|
||||||
|
handler.register_wm_subscriber(receiver, "session-new")
|
||||||
|
|
||||||
|
sock.send_multipart.assert_called_once_with(
|
||||||
|
[b"WATERMARK", b"3", b"0", b"session-new"]
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class TestMooncakePPStaging(unittest.TestCase):
|
class TestMooncakePPStaging(unittest.TestCase):
|
||||||
def test_staging_response_targets_requesting_pp_rank(self):
|
def test_staging_response_targets_requesting_pp_rank(self):
|
||||||
sock = Mock()
|
sock = Mock()
|
||||||
|
|||||||
Reference in New Issue
Block a user