Skip to content
Closed
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
4 changes: 2 additions & 2 deletions vllm_ascend/attention/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -496,7 +496,7 @@ def wait_for_kv_layer_from_connector(layer_name: str):

forward_context: ForwardContext = get_forward_context()
attn_metadata = forward_context.attn_metadata
if attn_metadata is None:
if attn_metadata is None or not connector.has_connector_metadata():
return
# TODO: assert ascendMetadata
connector.wait_for_layer_load(layer_name)
Expand All @@ -513,7 +513,7 @@ def maybe_save_kv_layer_to_connector(

forward_context: ForwardContext = get_forward_context()
attn_metadata = forward_context.attn_metadata
if attn_metadata is None:
if attn_metadata is None or not connector.has_connector_metadata():
return
# TODO: assert ascendMetadata
connector.save_kv_layer(layer_name, kv_cache_layer, attn_metadata)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -115,6 +115,15 @@ def _build_transfer_arrays(
(np.zeros(1, dtype=np.int64), np.cumsum(layer_block_len[:-1], dtype=np.int64))
)
rank_layer_offset = layer_id * self.page_size_bytes
if base_gvas_arr.size > 0 and np.any(base_gvas_arr <= 0):
zero_count = int(np.sum(base_gvas_arr <= 0))
logger.warning(
"[KVPOOL] build_transfer layer=%d detected %d zero/negative base_gvas "
"(base_gvas_sample=%s); these blocks will be skipped in batch_copy",
layer_id,
zero_count,
base_gvas_arr[:5].tolist(),
)
logger.debug(
"[KVPOOL] build_transfer layer=%d page_size=%d caches_per_layer=%d "
"rank_layer_offset=%d layer_block_len=%s layer_inner_offsets=%s "
Expand Down Expand Up @@ -229,7 +238,19 @@ def build_shared(self, task: LayerTransferTask, is_save: bool = True) -> SharedB
block_gvas_arr[offset] = request.last_block_gva
offset += 1

block_ids_arr, block_gvas_arr = self._dedupe_transfer_blocks(block_ids_arr[:offset], block_gvas_arr[:offset])
block_ids_slice = block_ids_arr[:offset]
block_gvas_slice = block_gvas_arr[:offset]
valid_mask = block_gvas_slice > 0
if not np.all(valid_mask):
skip_count = int(np.sum(~valid_mask))
logger.warning(
"[KVPOOL] build_shared skipping %d blocks with invalid gva (gva<=0)",
skip_count,
)
block_ids_slice = block_ids_slice[valid_mask]
block_gvas_slice = block_gvas_slice[valid_mask]

block_ids_arr, block_gvas_arr = self._dedupe_transfer_blocks(block_ids_slice, block_gvas_slice)

logger.debug(
"[KVPOOL] build_shared req_ids=%s block_gvas_arr=%s block_ids_arr=%s",
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -987,6 +987,8 @@ def touch_sending_mamba_blocks(self, req_meta: ReqMeta):
"""
if not self.use_hybrid or len(self.mamba_group_ids) == 0 or not req_meta.can_save:
return
if self.use_layerwise:
return
Comment on lines +990 to +991

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

critical

Returning early here when self.use_layerwise is True prevents the sending mamba blocks from being touched (i.e., pinned/referenced). Since request_finished and request_finished_all_groups also return False immediately when use_layerwise is True, these blocks will be freed immediately by the block manager when the request finishes. If the layerwise sending thread (KVCacheStoreLayerSendingThread) is still asynchronously reading from these blocks in the background, they can be reassigned to a new request and overwritten, leading to a critical race condition and data corruption. Consider implementing a proper synchronization or event-based completion mechanism for layerwise sending to safely defer freeing these blocks until the transfer is fully complete.

using_event_id = self.get_sending_event_id()
req_meta.event_id = using_event_id
current_step_sending: list[int] = []
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -800,8 +800,6 @@ def register_kv_caches(self, kv_caches: dict[str, torch.Tensor]):

# Initialize store, register buffers, and start transfer threads
# directly here (like main) — no separate init_backend handshake.
if hasattr(self.m_store, "init_store"):
self.m_store.init_store()
self.m_store.register_buffer(ptrs, lengths)
self._start_kv_transfer_threads()

Expand Down Expand Up @@ -1145,8 +1143,9 @@ def _alloc_gvas_for_save(self, requests: list[ReqMeta]) -> None:
sum(1 for g in new_gvas if g <= 0),
)
for pos, key, gva in zip(new_positions, new_keys, new_gvas):
block_gvas[pos] = gva
self._allocated_gvas[key] = gva
if gva > 0:
block_gvas[pos] = gva
self._allocated_gvas[key] = gva

logger.info(
"alloc_gvas: req=%s group=%d eff_bs=%d save_blocks=[%d,%d) "
Expand Down Expand Up @@ -1212,7 +1211,10 @@ def _prepare_load_gvas(self, requests: list[ReqMeta]) -> None:
if (request.block_ids_by_group_np is not None and group_id < len(request.block_ids_by_group_np))
else request.block_ids_np
)
full_len = len(block_ids_by_group) if block_ids_by_group is not None else 0
if block_ids_by_group is None:
all_group_load_gvas.append(np.zeros(0, dtype=np.int64))
continue
full_len = len(block_ids_by_group)

if load_start_block >= full_blocks:
all_group_load_gvas.append(np.zeros(full_len, dtype=np.int64))
Expand All @@ -1227,12 +1229,67 @@ def _prepare_load_gvas(self, requests: list[ReqMeta]) -> None:
continue

key_infos = self.m_store.batch_get_key_info(keys)
lease_results = self.m_store.batch_add_lease(keys, LAYERWISE_READ_LEASE_TTL_MS)
gvas = []
for ki in key_infos:
valid_keys_for_lease = []
valid_block_ids = []
invalid_block_ids: list[int] = []
for idx, (ki, key) in enumerate(zip(key_infos, keys)):
sizes = ki.size()
gvas.append(ki.gva_list()[0] if sizes and sizes > 0 else 0)
all_group_load_keys.extend(keys)
gva = ki.gva_list()[0] if sizes and sizes > 0 else 0
gvas.append(gva)
if gva > 0:
valid_keys_for_lease.append(key)
block_idx = load_start_block + idx
if block_idx < len(block_ids_by_group):
valid_block_ids.append(int(block_ids_by_group[block_idx]))
else:
block_idx = load_start_block + idx
if block_idx < len(block_ids_by_group):
invalid_block_ids.append(int(block_ids_by_group[block_idx]))
logger.warning(
"load_gvas: req=%s group=%d got invalid gva=%d (size=%d), block_id=%s load failed",
request.req_id,
group_id,
gva,
sizes if sizes else 0,
int(block_ids_by_group[block_idx]) if block_idx < len(block_ids_by_group) else "N/A",
)

# Only call batch_add_lease for keys with valid size
if valid_keys_for_lease:
lease_results = self.m_store.batch_add_lease(valid_keys_for_lease, LAYERWISE_READ_LEASE_TTL_MS)
# Report lease failures as invalid blocks
for i, lease_res in enumerate(lease_results):
if lease_res != 0 and i < len(valid_block_ids):
invalid_block_ids.append(valid_block_ids[i])
logger.warning(
"load_gvas: req=%s group=%d lease failed result=%d, block_id=%d load failed",
request.req_id,
group_id,
lease_res,
valid_block_ids[i],
)
else:
lease_results = []

# Report invalid blocks to scheduler for recompute.
# Single-group models can safely report individual block IDs.
# Multi-group (hybrid) models must not report partial group
# failures, as the scheduler cannot handle inconsistent KV
# cache state across groups (see PR #9701 for rationale).
if invalid_block_ids:
if len(request.block_ids_by_group) == 1:
with self._invalid_block_ids_lock:
self._invalid_block_ids.update(invalid_block_ids)
else:
logger.error(
"KV load failed for hybrid request %s. "
"Skip invalid-block fallback to avoid scheduler crash. "
"failed_blocks=%s",
request.req_id,
invalid_block_ids,
)
all_group_load_keys.extend(valid_keys_for_lease)

logger.info(
"load_gvas: req=%s group=%d eff_bs=%d load_blocks=[%d,%d) keys=%d valid_gvas=%d lease_fail=%d",
Expand Down
8 changes: 7 additions & 1 deletion vllm_ascend/ops/gdn.py
Original file line number Diff line number Diff line change
Expand Up @@ -27,8 +27,12 @@
from vllm.v1.attention.backends.gdn_attn import GDNAttentionMetadata
from vllm.v1.attention.backends.utils import PAD_SLOT_ID

from vllm_ascend.attention.utils import maybe_save_kv_layer_to_connector
from vllm_ascend.attention.utils import (
maybe_save_kv_layer_to_connector,
wait_for_kv_layer_from_connector,
)
from vllm_ascend.device.device_op import DeviceOperator
from vllm_ascend.memcache_comm_fence import record_attention_compute_start
from vllm_ascend.ops.gdn_attn_builder import AscendGDNAttentionBackend
from vllm_ascend.ops.triton.fla.chunk import chunk_gated_delta_rule
from vllm_ascend.ops.triton.fla.fused_qkvzba_split_reshape import fused_qkvzba_split_reshape_cat
Expand Down Expand Up @@ -75,6 +79,7 @@ def forward(
2. Core attention (custom op)
3. Output projection
"""
wait_for_kv_layer_from_connector(self.prefix)
num_tokens = hidden_states.size(0)
if hasattr(self, "in_proj_qkv"):
mixed_qkv, _ = self.in_proj_qkv(hidden_states)
Expand Down Expand Up @@ -122,6 +127,7 @@ def forward(
device=hidden_states.device,
)

record_attention_compute_start()
torch.ops.vllm.qwen_gdn_attention_core(
mixed_qkv,
b,
Expand Down
Loading