[PD] Pack draft KV head slices for DCP transfers (#40500)
Co-authored-by: Qiaolin Yu <liin1211@outlook.com>
This commit is contained in:
@@ -1,4 +1,4 @@
|
|||||||
from typing import Sequence
|
from typing import Optional, Sequence
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
import triton
|
import triton
|
||||||
@@ -15,10 +15,11 @@ def _copy_mla_rows_into_pack_kernel(
|
|||||||
):
|
):
|
||||||
layer_id = tl.program_id(0)
|
layer_id = tl.program_id(0)
|
||||||
block_id = tl.program_id(1)
|
block_id = tl.program_id(1)
|
||||||
metadata_offset = layer_id * 3
|
metadata_offset = layer_id * 4
|
||||||
src = tl.load(src_metadata + metadata_offset).to(pack.dtype)
|
src = tl.load(src_metadata + metadata_offset).to(pack.dtype)
|
||||||
row_nbytes = tl.load(src_metadata + metadata_offset + 1)
|
row_nbytes = tl.load(src_metadata + metadata_offset + 1)
|
||||||
pack_offset = tl.load(src_metadata + metadata_offset + 2)
|
pack_offset = tl.load(src_metadata + metadata_offset + 2)
|
||||||
|
src_row_stride = tl.load(src_metadata + metadata_offset + 3)
|
||||||
|
|
||||||
offsets = block_id * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
|
offsets = block_id * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
|
||||||
layer_nbytes = num_rows * row_nbytes
|
layer_nbytes = num_rows * row_nbytes
|
||||||
@@ -26,7 +27,7 @@ def _copy_mla_rows_into_pack_kernel(
|
|||||||
row = offsets // row_nbytes
|
row = offsets // row_nbytes
|
||||||
byte = offsets % row_nbytes
|
byte = offsets % row_nbytes
|
||||||
src_row = tl.load(row_indices + row, mask=mask, other=0)
|
src_row = tl.load(row_indices + row, mask=mask, other=0)
|
||||||
values = tl.load(src + src_row * row_nbytes + byte, mask=mask)
|
values = tl.load(src + src_row * src_row_stride + byte, mask=mask)
|
||||||
tl.store(pack + pack_offset + offsets, values, mask=mask)
|
tl.store(pack + pack_offset + offsets, values, mask=mask)
|
||||||
|
|
||||||
|
|
||||||
@@ -35,11 +36,14 @@ def copy_mla_rows_into_pack(
|
|||||||
row_indices: torch.Tensor,
|
row_indices: torch.Tensor,
|
||||||
pack: torch.Tensor,
|
pack: torch.Tensor,
|
||||||
token_item_lens: Sequence[int],
|
token_item_lens: Sequence[int],
|
||||||
|
src_token_item_lens: Optional[Sequence[int]] = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
if len(kv_data_ptrs) != len(token_item_lens):
|
if src_token_item_lens is None:
|
||||||
|
src_token_item_lens = token_item_lens
|
||||||
|
if not (len(kv_data_ptrs) == len(token_item_lens) == len(src_token_item_lens)):
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
"kv_data_ptrs and token_item_lens length mismatch: "
|
"KV pointers, copy widths, and source strides length mismatch: "
|
||||||
f"{len(kv_data_ptrs)} vs {len(token_item_lens)}"
|
f"{len(kv_data_ptrs)}, {len(token_item_lens)}, {len(src_token_item_lens)}"
|
||||||
)
|
)
|
||||||
if not kv_data_ptrs:
|
if not kv_data_ptrs:
|
||||||
return
|
return
|
||||||
@@ -47,11 +51,13 @@ def copy_mla_rows_into_pack(
|
|||||||
n = int(row_indices.numel())
|
n = int(row_indices.numel())
|
||||||
metadata = []
|
metadata = []
|
||||||
offset = 0
|
offset = 0
|
||||||
for ptr, item_len in zip(kv_data_ptrs, token_item_lens):
|
for ptr, item_len, src_item_len in zip(
|
||||||
|
kv_data_ptrs, token_item_lens, src_token_item_lens
|
||||||
|
):
|
||||||
item_len = int(item_len)
|
item_len = int(item_len)
|
||||||
if item_len <= 0:
|
if item_len <= 0:
|
||||||
raise ValueError(f"MLA token item length must be positive, got {item_len}")
|
raise ValueError(f"MLA token item length must be positive, got {item_len}")
|
||||||
metadata.extend((int(ptr), item_len, offset))
|
metadata.extend((int(ptr), item_len, offset, int(src_item_len)))
|
||||||
offset += n * item_len
|
offset += n * item_len
|
||||||
|
|
||||||
src_metadata = torch.tensor(metadata, dtype=torch.int64, device=pack.device)
|
src_metadata = torch.tensor(metadata, dtype=torch.int64, device=pack.device)
|
||||||
|
|||||||
@@ -382,7 +382,9 @@ class CommonKVManager(BaseKVManager):
|
|||||||
f"{type(self).__name__} does not support staging memory registration"
|
f"{type(self).__name__} does not support staging memory registration"
|
||||||
)
|
)
|
||||||
|
|
||||||
def _init_dcp_pack_buffers_once(self, dcp_size: int) -> None:
|
def _init_dcp_pack_buffers_once(
|
||||||
|
self, dcp_size: int, *, include_draft: bool = False
|
||||||
|
) -> None:
|
||||||
if self._dcp_pack_buffers is not None:
|
if self._dcp_pack_buffers is not None:
|
||||||
return
|
return
|
||||||
if not self.kv_args.kv_item_lens:
|
if not self.kv_args.kv_item_lens:
|
||||||
@@ -397,6 +399,7 @@ class CommonKVManager(BaseKVManager):
|
|||||||
len(self.transfer_queues),
|
len(self.transfer_queues),
|
||||||
dcp_size,
|
dcp_size,
|
||||||
max_tokens,
|
max_tokens,
|
||||||
|
include_draft=include_draft,
|
||||||
)
|
)
|
||||||
self._dcp_pack_max_tokens = max_tokens
|
self._dcp_pack_max_tokens = max_tokens
|
||||||
|
|
||||||
|
|||||||
@@ -41,6 +41,7 @@ def try_pack_dcp_src(
|
|||||||
kv_data_ptrs: Sequence[int],
|
kv_data_ptrs: Sequence[int],
|
||||||
src_token_indices: npt.NDArray[np.integer],
|
src_token_indices: npt.NDArray[np.integer],
|
||||||
token_item_lens: Sequence[int],
|
token_item_lens: Sequence[int],
|
||||||
|
src_token_item_lens: Optional[Sequence[int]] = None,
|
||||||
pack_offset_bytes: int = 0,
|
pack_offset_bytes: int = 0,
|
||||||
pack_capacity_bytes: Optional[int] = None,
|
pack_capacity_bytes: Optional[int] = None,
|
||||||
) -> Optional[Tuple[List[int], npt.NDArray[np.int64]]]:
|
) -> Optional[Tuple[List[int], npt.NDArray[np.int64]]]:
|
||||||
@@ -75,7 +76,9 @@ def try_pack_dcp_src(
|
|||||||
gather_stream = pack_buffer.get_gather_stream()
|
gather_stream = pack_buffer.get_gather_stream()
|
||||||
gather_stream.wait_stream(torch.cuda.default_stream(pack.device))
|
gather_stream.wait_stream(torch.cuda.default_stream(pack.device))
|
||||||
with torch.cuda.stream(gather_stream):
|
with torch.cuda.stream(gather_stream):
|
||||||
copy_mla_rows_into_pack(kv_data_ptrs, row_indices, pack, token_item_lens)
|
copy_mla_rows_into_pack(
|
||||||
|
kv_data_ptrs, row_indices, pack, token_item_lens, src_token_item_lens
|
||||||
|
)
|
||||||
gather_stream.synchronize()
|
gather_stream.synchronize()
|
||||||
|
|
||||||
packed_ptrs: List[int] = []
|
packed_ptrs: List[int] = []
|
||||||
@@ -93,14 +96,16 @@ def init_dcp_pack_buffers(
|
|||||||
count: int,
|
count: int,
|
||||||
dcp_size: int,
|
dcp_size: int,
|
||||||
max_tokens: int,
|
max_tokens: int,
|
||||||
|
*,
|
||||||
|
include_draft: bool = False,
|
||||||
) -> List[StagingBuffer]:
|
) -> List[StagingBuffer]:
|
||||||
from sglang.srt.disaggregation.common.staging_handler import (
|
from sglang.srt.disaggregation.common.staging_handler import (
|
||||||
_get_custom_mem_pool,
|
_get_custom_mem_pool,
|
||||||
)
|
)
|
||||||
|
|
||||||
kv_item_lens = kv_args.kv_item_lens
|
kv_item_lens = kv_args.kv_item_lens
|
||||||
if kv_args.num_draft_entries > 0:
|
if not include_draft and kv_args.num_draft_entries:
|
||||||
kv_item_lens = kv_item_lens[: len(kv_item_lens) - kv_args.num_draft_entries]
|
kv_item_lens = kv_item_lens[: -kv_args.num_draft_entries]
|
||||||
# Note(kpham-sgl): size = dcp_size x ceil(max_tokens / dcp_size)
|
# Note(kpham-sgl): size = dcp_size x ceil(max_tokens / dcp_size)
|
||||||
# x sum(per-layer token bytes). At 32,768 tokens and 61 MLA layers
|
# x sum(per-layer token bytes). At 32,768 tokens and 61 MLA layers
|
||||||
# x 576 bf16 dims x 2 B: 2.14 GiB/buffer, 8.58 GiB for 4 queues.
|
# x 576 bf16 dims x 2 B: 2.14 GiB/buffer, 8.58 GiB for 4 queues.
|
||||||
|
|||||||
@@ -1173,6 +1173,7 @@ class MooncakeKVManager(StagingManagerMixin, CommonKVManager):
|
|||||||
for entry in range(num_target)
|
for entry in range(num_target)
|
||||||
]
|
]
|
||||||
sliced_draft_params = []
|
sliced_draft_params = []
|
||||||
|
draft_to_pack = []
|
||||||
if num_draft > 0 and plan.draft_src_token_indices.size:
|
if num_draft > 0 and plan.draft_src_token_indices.size:
|
||||||
if not dst_kv_item_lens and dst_attn_tp_size not in (
|
if not dst_kv_item_lens and dst_attn_tp_size not in (
|
||||||
None,
|
None,
|
||||||
@@ -1224,14 +1225,45 @@ class MooncakeKVManager(StagingManagerMixin, CommonKVManager):
|
|||||||
dst_rank = dst_tp_rank // max(1, dst_span // src_span)
|
dst_rank = dst_tp_rank // max(1, dst_span // src_span)
|
||||||
src_offset = (dst_rank * dst_width) % src_width
|
src_offset = (dst_rank * dst_width) % src_width
|
||||||
dst_offset = (src_rank * src_width) % dst_width
|
dst_offset = (src_rank * src_width) % dst_width
|
||||||
sliced_draft_params.append(
|
params = (
|
||||||
(
|
|
||||||
src_kv_ptrs[entry] + src_offset,
|
src_kv_ptrs[entry] + src_offset,
|
||||||
dst_kv_ptrs[entry] + dst_offset,
|
dst_kv_ptrs[entry] + dst_offset,
|
||||||
src_width,
|
src_width,
|
||||||
dst_width,
|
dst_width,
|
||||||
copy_width,
|
copy_width,
|
||||||
)
|
)
|
||||||
|
if pack_buffer is not None and src_width > dst_width:
|
||||||
|
draft_to_pack.append(params)
|
||||||
|
else:
|
||||||
|
sliced_draft_params.append(params)
|
||||||
|
|
||||||
|
if draft_to_pack:
|
||||||
|
from sglang.srt.disaggregation.common.dcp_pack import try_pack_dcp_src
|
||||||
|
|
||||||
|
draft_src_ptrs, draft_dst_ptrs, src_strides, _, copy_widths = zip(
|
||||||
|
*draft_to_pack
|
||||||
|
)
|
||||||
|
target_pack_bytes = plan.target_src_token_indices.size * sum(
|
||||||
|
dcp_token_item_lens[:num_target]
|
||||||
|
)
|
||||||
|
packed = try_pack_dcp_src(
|
||||||
|
pack_buffer=pack_buffer,
|
||||||
|
kv_data_ptrs=draft_src_ptrs,
|
||||||
|
src_token_indices=plan.draft_src_token_indices,
|
||||||
|
token_item_lens=copy_widths,
|
||||||
|
src_token_item_lens=src_strides,
|
||||||
|
pack_offset_bytes=target_pack_bytes,
|
||||||
|
)
|
||||||
|
if packed is None:
|
||||||
|
sliced_draft_params.extend(draft_to_pack)
|
||||||
|
else:
|
||||||
|
packed_ptrs, packed_indices = packed
|
||||||
|
packed_groups = group_concurrent_contiguous(
|
||||||
|
packed_indices, plan.draft_dst_token_indices
|
||||||
|
)
|
||||||
|
layers_params.extend(
|
||||||
|
(src, dst, width, packed_groups)
|
||||||
|
for src, dst, width in zip(packed_ptrs, draft_dst_ptrs, copy_widths)
|
||||||
)
|
)
|
||||||
|
|
||||||
def process_sliced_draft(params) -> int:
|
def process_sliced_draft(params) -> int:
|
||||||
@@ -1284,7 +1316,11 @@ class MooncakeKVManager(StagingManagerMixin, CommonKVManager):
|
|||||||
executor.submit(process_sliced_draft, [params])
|
executor.submit(process_sliced_draft, [params])
|
||||||
for params in sliced_draft_params
|
for params in sliced_draft_params
|
||||||
)
|
)
|
||||||
|
try:
|
||||||
return self._await_transfer_futures(futures)
|
return self._await_transfer_futures(futures)
|
||||||
|
finally:
|
||||||
|
if pack_buffer is not None:
|
||||||
|
concurrent.futures.wait(futures)
|
||||||
|
|
||||||
transfer_blocks = []
|
transfer_blocks = []
|
||||||
for layer_params in layers_params:
|
for layer_params in layers_params:
|
||||||
@@ -2492,7 +2528,9 @@ class MooncakeKVManager(StagingManagerMixin, CommonKVManager):
|
|||||||
decode_kv_args.dst_dcp_size,
|
decode_kv_args.dst_dcp_size,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
self._init_dcp_pack_buffers_once(decode_kv_args.dst_dcp_size)
|
self._init_dcp_pack_buffers_once(
|
||||||
|
decode_kv_args.dst_dcp_size, include_draft=True
|
||||||
|
)
|
||||||
self.decode_kv_args_table[mooncake_session_id] = decode_kv_args
|
self.decode_kv_args_table[mooncake_session_id] = decode_kv_args
|
||||||
with self.session_lock:
|
with self.session_lock:
|
||||||
if mooncake_session_id in self.failed_sessions:
|
if mooncake_session_id in self.failed_sessions:
|
||||||
|
|||||||
@@ -1,8 +1,13 @@
|
|||||||
|
import concurrent.futures
|
||||||
import unittest
|
import unittest
|
||||||
|
from types import SimpleNamespace
|
||||||
|
|
||||||
|
import numpy as np
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
from sglang.kernels.ops.kvcache.pd_dcp_gather import copy_mla_rows_into_pack
|
from sglang.kernels.ops.kvcache.pd_dcp_gather import copy_mla_rows_into_pack
|
||||||
|
from sglang.srt.disaggregation.common.staging_buffer import StagingBuffer
|
||||||
|
from sglang.srt.disaggregation.mooncake.conn import MooncakeKVManager
|
||||||
from sglang.test.ci.ci_register import register_cuda_ci
|
from sglang.test.ci.ci_register import register_cuda_ci
|
||||||
from sglang.test.test_utils import CustomTestCase
|
from sglang.test.test_utils import CustomTestCase
|
||||||
|
|
||||||
@@ -36,6 +41,135 @@ class TestPdDcpGather(CustomTestCase):
|
|||||||
torch.testing.assert_close(packed0, kv0[row_indices], rtol=0, atol=0)
|
torch.testing.assert_close(packed0, kv0[row_indices], rtol=0, atol=0)
|
||||||
torch.testing.assert_close(packed1, kv1[row_indices], rtol=0, atol=0)
|
torch.testing.assert_close(packed1, kv1[row_indices], rtol=0, atol=0)
|
||||||
|
|
||||||
|
def test_packed_tp2_pp2_to_dcp4_preserves_kv(self):
|
||||||
|
"""Packing must preserve both target rows and draft head shards across PP stages."""
|
||||||
|
for custom_pool in (False, True):
|
||||||
|
for capacity in (256 * (2 * 64 + 2 * 512), 256 * 2 * 64 // 4):
|
||||||
|
for rank in range(4):
|
||||||
|
with self.subTest(
|
||||||
|
custom_pool=custom_pool, capacity=capacity, rank=rank
|
||||||
|
):
|
||||||
|
self._check_packed_transfer(rank, custom_pool, capacity)
|
||||||
|
|
||||||
|
def _check_packed_transfer(self, rank, custom_pool, capacity):
|
||||||
|
page, tokens, chunk = 64, 521, 256
|
||||||
|
src_pages = np.array([7, 1, 9, 3, 4, 11, 2, 5, 8], dtype=np.int32)
|
||||||
|
dst_pages = np.array([4, 1, 6], dtype=np.int32)
|
||||||
|
layers, widths = [3, 11, 19, 27, 28, 28], [64] * 4 + [256] * 2
|
||||||
|
logical = torch.arange(tokens, device="cuda")
|
||||||
|
src_rows = (
|
||||||
|
torch.as_tensor(src_pages, device="cuda")[logical // page] * page
|
||||||
|
+ logical % page
|
||||||
|
)
|
||||||
|
values = [
|
||||||
|
(
|
||||||
|
(logical[:, None] + 256) * 13
|
||||||
|
+ torch.arange(width, device="cuda") * 7
|
||||||
|
+ entry * 31
|
||||||
|
)
|
||||||
|
.remainder(251)
|
||||||
|
.to(torch.uint8)
|
||||||
|
for entry, width in enumerate([64] * 4 + [1024] * 2)
|
||||||
|
]
|
||||||
|
destinations = [
|
||||||
|
torch.full((2048, w), 165, dtype=torch.uint8, device="cuda") for w in widths
|
||||||
|
]
|
||||||
|
expected = [x.clone() for x in destinations]
|
||||||
|
owned = logical[rank::4]
|
||||||
|
target_rows = (
|
||||||
|
torch.as_tensor(dst_pages, device="cuda")[owned // 256] * page
|
||||||
|
+ owned % 256 // 4
|
||||||
|
)
|
||||||
|
draft_rows = (
|
||||||
|
torch.as_tensor(dst_pages, device="cuda")[logical // 256] * 256
|
||||||
|
+ logical % 256
|
||||||
|
)
|
||||||
|
for entry in range(4):
|
||||||
|
expected[entry][target_rows] = values[entry][owned]
|
||||||
|
for entry in (4, 5):
|
||||||
|
expected[entry][draft_rows] = values[entry][
|
||||||
|
:, rank * 256 : (rank + 1) * 256
|
||||||
|
]
|
||||||
|
|
||||||
|
pack = StagingBuffer(capacity, "cuda:0", 0)
|
||||||
|
with concurrent.futures.ThreadPoolExecutor(max_workers=3) as executor:
|
||||||
|
for stage, entries in enumerate(([0, 1], [2, 3, 4, 5])):
|
||||||
|
sources = []
|
||||||
|
for entry in entries:
|
||||||
|
data = values[entry]
|
||||||
|
if entry >= 4:
|
||||||
|
start = (rank // 2) * 512
|
||||||
|
data = data[:, start : start + 512]
|
||||||
|
source = torch.full(
|
||||||
|
(1024, data.shape[1]), 165, dtype=torch.uint8, device="cuda"
|
||||||
|
)
|
||||||
|
source[src_rows] = data
|
||||||
|
sources.append(source)
|
||||||
|
buffers = sources + destinations + [pack.buffer]
|
||||||
|
|
||||||
|
def transfer(session, blocks):
|
||||||
|
def view(ptr, size):
|
||||||
|
for tensor in buffers:
|
||||||
|
offset = ptr - tensor.data_ptr()
|
||||||
|
if 0 <= offset and offset + size <= tensor.numel():
|
||||||
|
return tensor.flatten()[offset : offset + size]
|
||||||
|
raise AssertionError(
|
||||||
|
f"Transfer outside registered buffers: {ptr}, {size}"
|
||||||
|
)
|
||||||
|
|
||||||
|
for src, dst, size in blocks:
|
||||||
|
view(dst, size).copy_(view(src, size))
|
||||||
|
torch.cuda.synchronize()
|
||||||
|
return 0
|
||||||
|
|
||||||
|
manager = SimpleNamespace(
|
||||||
|
is_mla_backend=False,
|
||||||
|
kv_args=SimpleNamespace(
|
||||||
|
page_size=page,
|
||||||
|
kv_layer_ids=[layers[e] for e in entries],
|
||||||
|
kv_data_ptrs=[x.data_ptr() for x in sources],
|
||||||
|
num_draft_entries=2 if stage else 0,
|
||||||
|
engine_rank=stage * 2 + rank // 2,
|
||||||
|
),
|
||||||
|
attn_tp_size=2,
|
||||||
|
max_transfer_batch_indices=37,
|
||||||
|
enable_custom_mem_pool=custom_pool,
|
||||||
|
enable_deferred_decode_kv_release=False,
|
||||||
|
_transfer_data=transfer,
|
||||||
|
)
|
||||||
|
manager._await_transfer_futures = lambda futures: (
|
||||||
|
MooncakeKVManager._await_transfer_futures(manager, futures)
|
||||||
|
)
|
||||||
|
for start in range(0, tokens, chunk):
|
||||||
|
count = min(chunk, tokens - start)
|
||||||
|
result = MooncakeKVManager.send_kvcache_dcp(
|
||||||
|
manager,
|
||||||
|
"session",
|
||||||
|
src_pages[start // page : (start + count + page - 1) // page],
|
||||||
|
[x.data_ptr() for x in destinations],
|
||||||
|
dst_pages,
|
||||||
|
dcp_token_item_lens=[x.shape[1] for x in sources],
|
||||||
|
dst_dcp_size=4,
|
||||||
|
dst_dcp_rank=rank,
|
||||||
|
src_page_offset=start // page,
|
||||||
|
decode_prefix_len=256,
|
||||||
|
num_kv_tokens=count,
|
||||||
|
executor=executor,
|
||||||
|
dst_layer_ids=layers,
|
||||||
|
pack_buffer=pack,
|
||||||
|
dst_kv_item_lens=[
|
||||||
|
page * w * (4 if e >= 4 else 1)
|
||||||
|
for e, w in enumerate(widths)
|
||||||
|
],
|
||||||
|
dst_tp_rank=rank,
|
||||||
|
dst_attn_tp_size=4,
|
||||||
|
)
|
||||||
|
self.assertEqual(result, 0)
|
||||||
|
for entry in entries:
|
||||||
|
torch.testing.assert_close(
|
||||||
|
destinations[entry], expected[entry], rtol=0, atol=0
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
@@ -1,12 +1,14 @@
|
|||||||
import concurrent.futures
|
import concurrent.futures
|
||||||
import unittest
|
import unittest
|
||||||
|
from threading import Event
|
||||||
from types import SimpleNamespace
|
from types import SimpleNamespace
|
||||||
from unittest.mock import MagicMock, call
|
from unittest.mock import MagicMock, call, patch
|
||||||
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
|
|
||||||
from sglang.srt.disaggregation.mooncake.conn import MooncakeKVManager
|
from sglang.srt.disaggregation.mooncake.conn import MooncakeKVManager
|
||||||
from sglang.test.ci.ci_register import register_cpu_ci
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
|
from sglang.test.test_utils import CustomTestCase
|
||||||
|
|
||||||
register_cpu_ci(est_time=1, suite="base-a-test-cpu")
|
register_cpu_ci(est_time=1, suite="base-a-test-cpu")
|
||||||
|
|
||||||
@@ -286,5 +288,67 @@ class TestDcpDraftHeadTransfer(unittest.TestCase):
|
|||||||
self.assertFalse(dst_buffers[1000000].any())
|
self.assertFalse(dst_buffers[1000000].any())
|
||||||
|
|
||||||
|
|
||||||
|
class TestDcpPackLifetime(CustomTestCase):
|
||||||
|
def test_failed_transfer_drains_before_pack_buffer_reuse(self):
|
||||||
|
"""A failed layer must not release the pack buffer while another transfer reads it."""
|
||||||
|
manager = TestMooncakeTransferBatching._make_manager(
|
||||||
|
enable_custom_mem_pool=True
|
||||||
|
)
|
||||||
|
manager.kv_args = SimpleNamespace(
|
||||||
|
page_size=1, kv_layer_ids=[], kv_data_ptrs=[1000, 2000], num_draft_entries=0
|
||||||
|
)
|
||||||
|
source = np.array([11], dtype=np.uint8)
|
||||||
|
observed = []
|
||||||
|
running, release = Event(), Event()
|
||||||
|
|
||||||
|
def transfer(session, blocks):
|
||||||
|
if blocks[0][0] == 1000:
|
||||||
|
self.assertTrue(running.wait(10))
|
||||||
|
return 17
|
||||||
|
running.set()
|
||||||
|
self.assertTrue(release.wait(10))
|
||||||
|
observed.append(int(source[0]))
|
||||||
|
return 0
|
||||||
|
|
||||||
|
def send(executor):
|
||||||
|
result = MooncakeKVManager.send_kvcache_dcp(
|
||||||
|
manager,
|
||||||
|
"session",
|
||||||
|
np.array([0, 1], dtype=np.int32),
|
||||||
|
[5000, 6000],
|
||||||
|
np.array([0], dtype=np.int32),
|
||||||
|
dcp_token_item_lens=[1, 1],
|
||||||
|
dst_dcp_size=2,
|
||||||
|
dst_dcp_rank=0,
|
||||||
|
src_page_offset=0,
|
||||||
|
decode_prefix_len=0,
|
||||||
|
num_kv_tokens=2,
|
||||||
|
executor=executor,
|
||||||
|
dst_layer_ids=[],
|
||||||
|
pack_buffer=object(),
|
||||||
|
)
|
||||||
|
source[0] = 22
|
||||||
|
return result
|
||||||
|
|
||||||
|
manager._transfer_data = transfer
|
||||||
|
with (
|
||||||
|
patch(
|
||||||
|
"sglang.srt.disaggregation.common.dcp_pack.try_pack_dcp_src",
|
||||||
|
return_value=([1000, 2000], np.array([0], dtype=np.int64)),
|
||||||
|
),
|
||||||
|
concurrent.futures.ThreadPoolExecutor(max_workers=2) as transfers,
|
||||||
|
concurrent.futures.ThreadPoolExecutor(max_workers=1) as worker,
|
||||||
|
):
|
||||||
|
future = worker.submit(send, transfers)
|
||||||
|
try:
|
||||||
|
self.assertTrue(running.wait(10))
|
||||||
|
with self.assertRaises(concurrent.futures.TimeoutError):
|
||||||
|
future.result(timeout=1)
|
||||||
|
finally:
|
||||||
|
release.set()
|
||||||
|
self.assertEqual(future.result(timeout=10), 17)
|
||||||
|
self.assertEqual(observed, [11])
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
Reference in New Issue
Block a user