Skip to content

Commit 10b244a

Browse files
authored
[BugFix][P/D] Fix MooncakeLayerwiseConnector Mamba block_ids bug (vllm-project#11063)
### What this PR does / why we need it? This pull request updates the Mooncake layerwise connector to correctly calculate and use `remote_transfer_idx` instead of hardcoding `-1` when indexing `remote_block_ids` for Mamba speculative tokens. This ensures the correct remote blocks are targeted during KV transfer. ### Does this PR introduce _any_ user-facing change? No. ### How was this patch tested? CI tests. - vLLM version: v0.23.0 - vLLM main: vllm-project/vllm@967c5c3 --------- Signed-off-by: nwpu-zxr <zhouxuerong2@huawei.com>
1 parent ec1f3f7 commit 10b244a

1 file changed

Lines changed: 17 additions & 10 deletions

File tree

vllm_ascend/distributed/kv_transfer/kv_p2p/mooncake_layerwise_connector.py

Lines changed: 17 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -297,24 +297,25 @@ def get_transfer_meta(self, send_task: SendTask, req_id: str, req_meta: ReqMeta,
297297
if isinstance(layer_kv_cache_spec, MambaSpec):
298298
# only support one block transfer for mamba
299299
if self.mamba_cache_mode == "align":
300-
transfer_block_idx = len(local_block_ids) - self.num_speculative_tokens - 1
300+
local_transfer_idx = len(local_block_ids) - self.num_speculative_tokens - 1
301301
else:
302-
transfer_block_idx = 0
302+
local_transfer_idx = 0
303+
remote_transfer_idx = len(remote_block_ids) - self.num_speculative_tokens - 1
303304
local_conv_addr, local_ssm_addr = local_layer_metadata.kv_caches_base_addr
304305
remote_conv_addr, remote_ssm_addr = remote_layer_metadata.kv_caches_base_addr
305306
local_conv_len, local_ssm_len = local_layer_metadata.block_len
306307
tp_ratio = self.tp_size // req_meta.remote_tp_size
307308
if tp_ratio == 1:
308309
src_list.extend(
309310
[
310-
local_conv_addr + local_block_ids[transfer_block_idx] * local_conv_len,
311-
local_ssm_addr + local_block_ids[transfer_block_idx] * local_ssm_len,
311+
local_conv_addr + local_block_ids[local_transfer_idx] * local_conv_len,
312+
local_ssm_addr + local_block_ids[local_transfer_idx] * local_ssm_len,
312313
]
313314
)
314315
dst_list.extend(
315316
[
316-
remote_conv_addr + remote_block_ids[-1] * local_conv_len,
317-
remote_ssm_addr + remote_block_ids[-1] * local_ssm_len,
317+
remote_conv_addr + remote_block_ids[remote_transfer_idx] * local_conv_len,
318+
remote_ssm_addr + remote_block_ids[remote_transfer_idx] * local_ssm_len,
318319
]
319320
)
320321
length_list.extend([local_conv_len, local_ssm_len])
@@ -347,14 +348,20 @@ def get_transfer_meta(self, send_task: SendTask, req_id: str, req_meta: ReqMeta,
347348
+ (self.tp_rank % tp_ratio) * local_conv_size
348349
) * get_dtype_size(conv_dtype)
349350
src_list.append(
350-
local_conv_addr + local_block_ids[transfer_block_idx] * local_conv_len + local_addr_offset
351+
local_conv_addr + local_block_ids[local_transfer_idx] * local_conv_len + local_addr_offset
352+
)
353+
dst_list.append(
354+
remote_conv_addr
355+
+ remote_block_ids[remote_transfer_idx] * remote_conv_len
356+
+ remote_addr_offset
351357
)
352-
dst_list.append(remote_conv_addr + remote_block_ids[-1] * remote_conv_len + remote_addr_offset)
353358
length_list.append(local_conv_size * get_dtype_size(conv_dtype))
354359
# ssm
355360
remote_addr_offset = (self.tp_rank % tp_ratio) * math.prod(ssm_shape) * get_dtype_size(ssm_dtype)
356-
src_list.append(local_ssm_addr + local_block_ids[transfer_block_idx] * local_ssm_len)
357-
dst_list.append(remote_ssm_addr + remote_block_ids[-1] * remote_ssm_len + remote_addr_offset)
361+
src_list.append(local_ssm_addr + local_block_ids[local_transfer_idx] * local_ssm_len)
362+
dst_list.append(
363+
remote_ssm_addr + remote_block_ids[remote_transfer_idx] * remote_ssm_len + remote_addr_offset
364+
)
358365
length_list.append(local_ssm_len)
359366
else:
360367
if self.pd_head_ratio == 1:

0 commit comments

Comments
 (0)