Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
29 changes: 20 additions & 9 deletions python/sglang/srt/mem_cache/buffer_mode/pipeline.py
Original file line number Diff line number Diff line change
Expand Up @@ -507,8 +507,13 @@ def _sweep_stale_backup_intents(self) -> dict[NodeId, BufferBackupState]:
return states

def flush_pending_writes(self) -> None:
"""Launch D2H transfers for admitted intents, head-of-line: device
locks and staging slots are taken only here, when capacity allows."""
"""Stage admitted intents, then submit their D2H as one operation.

Preparation must not allocate or free L1 KV slots: it only allocates
host staging, and buffer-mode evict_host is a no-op. Flush before
returning to scheduler admission, where L1 slots can be reused.
Each staged source stays locked until its D2H ack.
"""
if not self.pending_write_queue:
return
cc = self._cache.cache_controller
Expand Down Expand Up @@ -560,24 +565,29 @@ def flush_pending_writes(self) -> None:
# instead of failing the alloc inside cc.write; acks free
# aux staging, retry next round.
break
if not self._launch_backup_intent(intent, device_value, comp_xfers):
if not self._stage_backup_intent(intent, device_value, comp_xfers):
# Pool full of in-flight staging and nothing reclaimable
# (the tree never holds host values in buffer mode):
# defer, head-of-line; pending acks will free slots.
break
self.pending_write_queue.popleft()

def _launch_backup_intent(
# Submit earlier successes even if a later intent ran out of staging.
# Do not leave prepared copies deferred across scheduler admission.
cc.start_writing()

def _stage_backup_intent(
self,
intent: _UnifiedBackupIntent,
device_value: torch.Tensor,
comp_xfers: dict[ComponentType, list[PoolTransfer]],
) -> bool:
"""Launch one admitted intent's D2H (staging alloc + device lock +
async copy); the caller removes it from pending_write_queue. Returns
False when staging cannot be allocated. From a successful launch the
intent always reaches its storage-ack, so its content joins the
LAUNCHED cover consulted by admission."""
"""Allocate host staging and pin one admitted intent's source.

The caller removes successful intents from pending_write_queue and
submits them together before returning. Return False when staging
cannot be allocated.
"""
cache = self._cache
cc = cache.cache_controller
snapshot = intent.snapshot
Expand All @@ -591,6 +601,7 @@ def _launch_backup_intent(
device_value,
node_id=snapshot.node_id,
extra_pools=aux_xfers or None,
flush=False,
)
if host_indices is None:
return False
Expand Down
5 changes: 3 additions & 2 deletions test/registered/unit/mem_cache/test_buffer_mode_sidecar.py
Original file line number Diff line number Diff line change
Expand Up @@ -143,8 +143,9 @@ def test_write_stages_and_persists_dsv4_full_and_swa_sidecars(self):
)
}

def _write(device_value, *, node_id, extra_pools):
def _write(device_value, node_id, extra_pools, flush):
self.assertEqual(node_id, 7)
self.assertFalse(flush)
self.assertEqual(
[transfer.name for transfer in extra_pools],
[PoolName.SWA, *[transfer.name for transfer in sidecars]],
Expand Down Expand Up @@ -190,7 +191,7 @@ def _write(device_value, *, node_id, extra_pools):
intent = _UnifiedBackupIntent(snapshot=snapshot)

self.assertTrue(
pipeline._launch_backup_intent(
pipeline._stage_backup_intent(
intent,
device_indices,
comp_xfers={ComponentType.SWA: [swa]},
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,7 @@
import unittest
from array import array
from collections import defaultdict
from contextlib import ExitStack
from dataclasses import dataclass, replace
from types import SimpleNamespace
from typing import Optional
Expand Down Expand Up @@ -3980,7 +3981,66 @@ def test_buffer_only_write_path_roundtrip(self):
).last_device_node
chain = self._path_chain(cache, leaf)

self._buffer_backup_and_wait(cache, leaf)
pipeline = cache.buffer_pipeline
controller = cache.cache_controller
_write_backup(cache, leaf)
node_ids = [intent.snapshot.node_id for intent in pipeline.pending_write_queue]
self.assertGreaterEqual(len(node_ids), 2)

device_allocators = [allocator]
if self.cfg.has_swa:
device_allocators.extend(
[allocator.full_attn_allocator, allocator.swa_attn_allocator]
)
# A deferred D2H batch must not allocate or free L1 slots before submit.
with ExitStack() as guards:
for device_allocator in device_allocators:
for method in (
"alloc",
"alloc_extend",
"alloc_decode",
"free",
"free_segment",
"free_page_ids",
):
if hasattr(device_allocator, method):
guards.enter_context(
mock.patch.object(
device_allocator,
method,
side_effect=AssertionError(
"L1 mutation during backup preparation"
),
)
)
pipeline.flush_pending_writes()
self.assertFalse(pipeline.pending_write_queue)
self.assertFalse(controller.write_queue)
self.assertEqual(len(controller.ack_write_queue), 1)
ack = controller.ack_write_queue[0]
self.assertEqual(ack.node_ids, node_ids)
self.assertEqual(set(pipeline.ongoing_write_through), set(node_ids))
staged_avail = self._host_avail_sizes(cache)
self.assertNotEqual(staged_avail, avail0)

# GPU completion alone must not release any entry's device lock.
ack.finish_event.synchronize()
for node_id in node_ids:
self.assertGreater(_device_lock_ref(cache, node_id, ComponentType.FULL), 0)
cache.writing_check(finish_count=1)
self.assertFalse(pipeline.ongoing_write_through)
self.assertEqual(len(pipeline.ongoing_backup), len(node_ids))
for node_id in node_ids:
self.assertEqual(_device_lock_ref(cache, node_id, ComponentType.FULL), 0)
# Each storage write still owns its staging until its separate ACK.
self.assertEqual(self._host_avail_sizes(cache), staged_avail)
self._pump_hicache_until(
cache,
lambda: (
not pipeline.inflight_backup_node_ids and not pipeline.ongoing_backup
),
"merged buffer backup did not drain every storage write",
)
self.assertFalse(cache.tree_core.is_backuped(leaf))
self.assertEqual(_device_lock_ref(cache, leaf, ComponentType.FULL), 0)
self.assertEqual(self._host_avail_sizes(cache), avail0)
Expand Down
Loading