Skip to content

Commit 686732c

Browse files
authored
[Feature][KVOffload][P/D] Recompute offload supports linear hybrid model (vllm-project#12881)
### What this PR does / why we need it? Recompute offload supports linear hybrid model, such as Qwen3.5. When the MTP accept length changes, the offload Mamba cache block corresponding to the accept length needs to be placed at the start position of mamba cache groups, and this logic is implemented in this PR. Specifically, when preemption occurs, the last `1+num_spec_tokens` mamba blocks are offloaded. When the blocks are reloaded, the index of the mamba block that needs to be reloaded is calculated based on the actual length of the last scheduling, that is, `accept_token_idx = self.num_spec_tokens - (state.num_computed_tokens - request.num_tokens + 1)`. In this case, we only need to reload the block to the 0th mamba block index of the resumed request. ### Does this PR introduce _any_ user-facing change? No. ### How was this patch tested? In local self-verification, I use Qwen3.6 35B-A3B for test. The GPU memory usage of D node is set to 0.75, and the DP1TP2 parallel configuration is used to trigger recompute offload at a high frequency. The verification result shows that the score of the GPQA dataset meets the expectation. ``` | dataset | version | metric | mode | vllm-api-general-chat | |----- | ----- | ----- | ----- | -----| | GPQA_diamond | b1ed2c | accuracy | gen | 85.86 | ``` - vLLM version: v0.25.1 - vLLM main: vllm-project/vllm@fe784ff --------- Signed-off-by: nwpu-zxr <zhouxuerong2@huawei.com>
1 parent 3f9cfbe commit 686732c

3 files changed

Lines changed: 99 additions & 79 deletions

File tree

tests/ut/kv_offload/test_recompute_cpu_offload.py

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -197,6 +197,7 @@ def test_recompute_cpu_offload_scheduler_aligns_sliding_window_blocks():
197197
def test_recompute_cpu_offload_scheduler_d2h_keeps_sliding_window_offsets():
198198
scheduler = RecomputeCPUOffloadScheduler.__new__(RecomputeCPUOffloadScheduler)
199199
scheduler._group_is_sliding_window = [True]
200+
scheduler._group_is_mamba = [False]
200201
scheduler.cpu_kv_cache_config = SimpleNamespace(
201202
kv_cache_groups=[SimpleNamespace(kv_cache_spec=SimpleNamespace(block_size=16))]
202203
)
@@ -231,6 +232,7 @@ def test_recompute_cpu_offload_scheduler_d2h_keeps_sliding_window_offsets():
231232
def test_recompute_cpu_offload_scheduler_h2d_skips_sliding_window_null_blocks():
232233
scheduler = RecomputeCPUOffloadScheduler.__new__(RecomputeCPUOffloadScheduler)
233234
scheduler._group_is_sliding_window = [True]
235+
scheduler._group_is_mamba = [False]
234236
scheduler.cpu_kv_cache_config = SimpleNamespace(
235237
kv_cache_groups=[SimpleNamespace(kv_cache_spec=SimpleNamespace(block_size=16))]
236238
)
@@ -262,6 +264,7 @@ def test_recompute_cpu_offload_scheduler_h2d_skips_sliding_window_null_blocks():
262264
def test_recompute_cpu_offload_scheduler_h2d_clips_mtp_tail_blocks():
263265
scheduler = RecomputeCPUOffloadScheduler.__new__(RecomputeCPUOffloadScheduler)
264266
scheduler._group_is_sliding_window = [False]
267+
scheduler._group_is_mamba = [False]
265268
scheduler.cpu_kv_cache_config = SimpleNamespace(
266269
kv_cache_groups=[SimpleNamespace(kv_cache_spec=SimpleNamespace(block_size=16))]
267270
)

vllm_ascend/core/recompute_scheduler.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -993,7 +993,7 @@ def update_from_output(
993993
for req_id, num_tokens_scheduled in num_scheduled_tokens.items():
994994
assert num_tokens_scheduled > 0
995995
request = self.requests.get(req_id)
996-
if request is not None:
996+
if not vllm_version_is("0.25.1") and request is not None:
997997
request.num_in_flight_tokens -= num_tokens_scheduled
998998
if failed_kv_load_req_ids and req_id in failed_kv_load_req_ids:
999999
# skip failed or rescheduled requests from KV load failure

vllm_ascend/distributed/kv_transfer/kv_pool/recompute_cpu_offload/manager.py

Lines changed: 95 additions & 78 deletions
Original file line numberDiff line numberDiff line change
@@ -18,7 +18,7 @@
1818
)
1919
from vllm.v1.core.kv_cache_utils import resolve_kv_cache_block_sizes
2020
from vllm.v1.core.sched.output import SchedulerOutput
21-
from vllm.v1.kv_cache_interface import SlidingWindowSpec, UniformTypeKVCacheSpecs
21+
from vllm.v1.kv_cache_interface import MambaSpec, SlidingWindowSpec, UniformTypeKVCacheSpecs
2222
from vllm.v1.outputs import KVConnectorOutput
2323

2424
from vllm_ascend.distributed.kv_transfer.kv_pool.recompute_cpu_offload.metadata import (
@@ -70,9 +70,13 @@ def __init__(
7070
assert kv_cache_config is not None
7171
self.vllm_config = vllm_config
7272
self.enable_offload_prefix_caching = enable_offload_prefix_caching
73+
self.num_spec_tokens = (
74+
vllm_config.speculative_config.num_speculative_tokens if vllm_config.speculative_config else 0
75+
)
7376
self.cpu_kv_cache_config = self._derive_cpu_config(kv_cache_config, cpu_capacity_bytes)
7477
self.num_cpu_blocks = self.cpu_kv_cache_config.num_blocks
7578
self._group_is_sliding_window = self._get_group_is_sliding_window(kv_cache_config)
79+
self._group_is_mamba = self._get_group_is_mamba(kv_cache_config)
7680
self.enable_kv_cache_events = (
7781
vllm_config.kv_events_config is not None and vllm_config.kv_events_config.enable_kv_cache_events
7882
)
@@ -129,6 +133,18 @@ def _get_group_is_sliding_window(kv_cache_config: "KVCacheConfig") -> list[bool]
129133
group_is_sliding_window.append(isinstance(group.kv_cache_spec, SlidingWindowSpec))
130134
return group_is_sliding_window
131135

136+
@staticmethod
137+
def _get_group_is_mamba(kv_cache_config: "KVCacheConfig") -> list[bool]:
138+
group_is_mamba: list[bool] = []
139+
for group in kv_cache_config.kv_cache_groups:
140+
if isinstance(group.kv_cache_spec, UniformTypeKVCacheSpecs):
141+
group_is_mamba.append(
142+
any(isinstance(spec, MambaSpec) for spec in group.kv_cache_spec.kv_cache_specs.values())
143+
)
144+
else:
145+
group_is_mamba.append(isinstance(group.kv_cache_spec, MambaSpec))
146+
return group_is_mamba
147+
132148
@staticmethod
133149
def _derive_cpu_config(gpu_config: "KVCacheConfig", cpu_capacity_bytes: int) -> "KVCacheConfig":
134150
from vllm.v1.kv_cache_interface import KVCacheConfig as KVCacheConfigCls
@@ -247,53 +263,49 @@ def _create_preempt_state(
247263
group_gpu_hashes: list[list[BlockHashWithGroupId | None]] = []
248264
missing_hashes: set[BlockHashWithGroupId] = set()
249265
num_unhashed = 0
266+
num_mamba_blocks = 0
250267

251268
for g, group_gpu_ids in enumerate(block_ids_by_group):
252-
group_block_size = kv_cache_groups[g].kv_cache_spec.block_size
253-
logical_num_blocks = cdiv(num_computed_tokens, group_block_size)
254-
aligned_group_gpu_ids = self._align_group_block_ids(g, group_gpu_ids, logical_num_blocks)
255-
eviction_group_gpu_ids = self._align_group_block_ids(
256-
g,
257-
group_gpu_ids,
258-
max(logical_num_blocks, len(group_gpu_ids)),
259-
)
260269
gpu_blocks: list[KVCacheBlock | None] = []
261270
effective_hashes: list[BlockHashWithGroupId | None] = []
262-
263-
for block_idx, block_id in enumerate(eviction_group_gpu_ids):
264-
if block_id <= 0:
265-
continue
266-
gpu_block = self._gpu_block_pool.blocks[block_id]
267-
block_is_computed = (block_idx + 1) * group_block_size <= num_computed_tokens
268-
if not block_is_computed and gpu_block.block_hash is not None:
269-
# allocate_slots() may assign a hash using tokens planned
270-
# for this scheduling step. If the request is then
271-
# preempted before forward, that block does not contain the
272-
# hashed KV and must not remain in the GPU prefix cache.
273-
self._gpu_block_pool._maybe_evict_cached_block(gpu_block)
274-
275-
for block_idx, block_id in enumerate(aligned_group_gpu_ids):
276-
if block_id <= 0:
277-
gpu_blocks.append(None)
278-
effective_hashes.append(None)
279-
continue
280-
281-
gpu_block = self._gpu_block_pool.blocks[block_id]
282-
block_is_computed = (block_idx + 1) * group_block_size <= num_computed_tokens
283-
block_hash = gpu_block.block_hash if block_is_computed and self.enable_offload_prefix_caching else None
284-
gpu_blocks.append(gpu_block)
285-
effective_hashes.append(block_hash)
286-
if block_hash is None:
287-
num_unhashed += 1
288-
elif (
289-
self.cpu_block_pool.cached_block_hash_to_block.get_one_block(block_hash) is None
290-
and block_hash not in self._pending_hash_blocks
291-
):
292-
missing_hashes.add(block_hash)
271+
if self._group_is_mamba[g]:
272+
# For Mamba cache, only the last `1 + num_spec_tokens` blocks need to offload
273+
offload_start_idx = len(group_gpu_ids) - self.num_spec_tokens - 1
274+
for block_idx, block_id in enumerate(group_gpu_ids):
275+
if block_idx >= offload_start_idx:
276+
num_mamba_blocks += 1
277+
gpu_block = self._gpu_block_pool.blocks[block_id]
278+
gpu_blocks.append(gpu_block)
279+
effective_hashes.append(None)
280+
else:
281+
group_block_size = kv_cache_groups[g].kv_cache_spec.block_size
282+
logical_num_blocks = cdiv(num_computed_tokens, group_block_size)
283+
aligned_group_gpu_ids = self._align_group_block_ids(g, group_gpu_ids, logical_num_blocks)
284+
285+
for block_idx, block_id in enumerate(aligned_group_gpu_ids):
286+
if block_id <= 0:
287+
gpu_blocks.append(None)
288+
effective_hashes.append(None)
289+
continue
290+
291+
gpu_block = self._gpu_block_pool.blocks[block_id]
292+
block_is_computed = (block_idx + 1) * group_block_size <= num_computed_tokens
293+
block_hash = (
294+
gpu_block.block_hash if block_is_computed and self.enable_offload_prefix_caching else None
295+
)
296+
gpu_blocks.append(gpu_block)
297+
effective_hashes.append(block_hash)
298+
if block_hash is None:
299+
num_unhashed += 1
300+
elif (
301+
self.cpu_block_pool.cached_block_hash_to_block.get_one_block(block_hash) is None
302+
and block_hash not in self._pending_hash_blocks
303+
):
304+
missing_hashes.add(block_hash)
293305
group_gpu_blocks.append(gpu_blocks)
294306
group_gpu_hashes.append(effective_hashes)
295307

296-
num_needed = num_unhashed + len(missing_hashes)
308+
num_needed = num_unhashed + len(missing_hashes) + num_mamba_blocks
297309
if not any(any(gpu_block is not None for gpu_block in group) for group in group_gpu_blocks):
298310
return False
299311
if num_needed > self.cpu_block_pool.get_num_free_blocks():
@@ -410,45 +422,50 @@ def _prepare_preempt_load_after_alloc(
410422
gpu_block_ids: list[int] = []
411423
cpu_block_ids: list[int] = []
412424
for g, group_cpu_ids in enumerate(state.cpu_block_ids):
413-
group_block_size = self.cpu_kv_cache_config.kv_cache_groups[g].kv_cache_spec.block_size
414-
start_block = load_start_tokens // group_block_size
415-
end_block = min(
416-
len(group_cpu_ids),
417-
len(
418-
self._align_group_block_ids(
419-
g,
420-
block_ids_by_group[g],
421-
max(
422-
cdiv(load_end_tokens, group_block_size),
423-
len(block_ids_by_group[g]),
424-
),
425-
)
426-
),
427-
cdiv(load_end_tokens, group_block_size),
428-
)
429-
if end_block == start_block:
430-
continue
431-
if end_block < start_block:
432-
raise RuntimeError(
433-
"Recompute H2D produced an empty block range: "
434-
f"req_id={request.request_id}, group={g}, "
435-
f"start_block={start_block}, end_block={end_block}, "
436-
f"gpu_blocks={len(block_ids_by_group[g])}, "
437-
f"cpu_blocks={len(group_cpu_ids)}"
425+
if self._group_is_mamba[g]:
426+
accept_token_idx = self.num_spec_tokens - (state.num_computed_tokens - request.num_tokens + 1)
427+
cpu_block_ids.append(group_cpu_ids[accept_token_idx])
428+
gpu_block_ids.append(block_ids_by_group[g][0])
429+
else:
430+
group_block_size = self.cpu_kv_cache_config.kv_cache_groups[g].kv_cache_spec.block_size
431+
start_block = load_start_tokens // group_block_size
432+
end_block = min(
433+
len(group_cpu_ids),
434+
len(
435+
self._align_group_block_ids(
436+
g,
437+
block_ids_by_group[g],
438+
max(
439+
cdiv(load_end_tokens, group_block_size),
440+
len(block_ids_by_group[g]),
441+
),
442+
)
443+
),
444+
cdiv(load_end_tokens, group_block_size),
438445
)
439-
440-
aligned_group_gpu_ids = self._align_group_block_ids(
441-
g,
442-
block_ids_by_group[g],
443-
end_block,
444-
)
445-
for block_idx in range(start_block, end_block):
446-
cpu_block_id = group_cpu_ids[block_idx]
447-
gpu_block_id = aligned_group_gpu_ids[block_idx]
448-
if cpu_block_id <= 0 or gpu_block_id <= 0:
446+
if end_block == start_block:
449447
continue
450-
cpu_block_ids.append(cpu_block_id)
451-
gpu_block_ids.append(gpu_block_id)
448+
if end_block < start_block:
449+
raise RuntimeError(
450+
"Recompute H2D produced an empty block range: "
451+
f"req_id={request.request_id}, group={g}, "
452+
f"start_block={start_block}, end_block={end_block}, "
453+
f"gpu_blocks={len(block_ids_by_group[g])}, "
454+
f"cpu_blocks={len(group_cpu_ids)}"
455+
)
456+
457+
aligned_group_gpu_ids = self._align_group_block_ids(
458+
g,
459+
block_ids_by_group[g],
460+
end_block,
461+
)
462+
for block_idx in range(start_block, end_block):
463+
cpu_block_id = group_cpu_ids[block_idx]
464+
gpu_block_id = aligned_group_gpu_ids[block_idx]
465+
if cpu_block_id <= 0 or gpu_block_id <= 0:
466+
continue
467+
cpu_block_ids.append(cpu_block_id)
468+
gpu_block_ids.append(gpu_block_id)
452469

453470
if not cpu_block_ids or len(cpu_block_ids) != len(gpu_block_ids):
454471
raise RuntimeError(

0 commit comments

Comments
 (0)