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
Original file line number Diff line number Diff line change
Expand Up @@ -173,6 +173,10 @@ def match_outputs_with_goldens(outputs: list[tuple[list[int], str]], goldens: Se
"additional_config": {
"enable_dsa_cp": True,
},
"speculative_config": {
"num_speculative_tokens": 3,
"method": "mtp",
},
},
),
]
Expand Down
88 changes: 68 additions & 20 deletions vllm_ascend/attention/context_parallel/dsa_cp.py
Original file line number Diff line number Diff line change
Expand Up @@ -63,6 +63,9 @@
# Legacy DSA-CP implementation (TP/SP group)
# =============================================================================

# Fixed-size contract required by the underlying _C_ascend ops (must be exactly 1024).
SAS_METADATA_SIZE = 1024


def hadamard_transform_ref(
x: torch.Tensor,
Expand Down Expand Up @@ -255,8 +258,8 @@ def __init__(
),
)
self.start_pos_prefill = torch.zeros(scheduler_config.max_num_seqs, dtype=torch.int32, device=self.device)
self.req_sas_metadata = torch.zeros(1024, dtype=torch.int32, device=self.device)
self.req_qli_metadata = torch.zeros(1024, dtype=torch.int32, device=self.device)
self.req_sas_metadata = torch.zeros(SAS_METADATA_SIZE, dtype=torch.int32, device=self.device)
self.req_qli_metadata = torch.zeros(SAS_METADATA_SIZE, dtype=torch.int32, device=self.device)
# Full-decode graphs pad the request count beyond max_num_seqs
# (cudagraph capture sizes plus the FIA dummy request), so size the
# per-request QLI buffers for the graph-mode maximum.
Expand Down Expand Up @@ -301,6 +304,13 @@ def __init__(
torch.zeros(scheduler_config.max_num_seqs, dtype=torch.int32, device=self.device)
for _ in range(spec_token_num)
]
self.spec_sas_metadata = [
torch.zeros(SAS_METADATA_SIZE, dtype=torch.int32, device=self.device) for _ in range(spec_token_num)
]
self.spec_start_pos = [
torch.zeros(scheduler_config.max_num_seqs, dtype=torch.int32, device=self.device)
for _ in range(spec_token_num)
]
self.decode_threshold += spec_token_num
assert self.decode_threshold <= 16, (
f"decode_threshold exceeded \
Expand Down Expand Up @@ -431,6 +441,7 @@ def build_for_drafting(
**kwargs,
) -> AscendDSAMetadata:
assert self.compressor_ratio <= 1, "vLLM-Ascend only support SWA-layer for Deepseek-V4 now."
# NOTE(Csrayz): num_reqs here is the padded token count from graph dispatch
num_reqs = common_attn_metadata.num_reqs
num_input_tokens = common_attn_metadata.num_input_tokens
# Cross-kv-cache-group metadata cache. The spec-decode proposer passes
Expand Down Expand Up @@ -464,9 +475,10 @@ def build_for_drafting(
treat_short_extends_as_decodes=False,
)
input_positions = common_attn_metadata.positions[:num_input_tokens].long()
# Draft steps update positions independently. Reusing the global RoPE
# cache can let later draft steps overwrite step-0 metadata.
cos, sin = get_cos_and_sin_dsa(input_positions, use_cache=False)
# Use per-draft-index RoPE buffer so tensor addresses stay stable
# across graph capture/replay; the per-step cache below then lets
# sibling kv-cache groups reuse the same stable tensors.
cos, sin = get_cos_and_sin_dsa(input_positions, use_cache=True, draft_index=draft_index)
if metadata_cache is not None:
metadata_cache.update(
num_decodes=num_decodes,
Expand Down Expand Up @@ -494,6 +506,13 @@ def build_for_drafting(
self.seq_lens_cpu = self.seq_lens.cpu()
if metadata_cache is not None:
metadata_cache["seq_lens_cpu"] = self.seq_lens_cpu
# In aclgraph mode num_reqs is the padded request count; keep the
# tensor's batch dim stable at num_reqs across capture and replay.
num_real_reqs = self.seq_lens_cpu.shape[0]
if num_real_reqs < num_reqs:
self.seq_lens_cpu = F.pad(self.seq_lens_cpu, (0, num_reqs - num_real_reqs))
else:
self.seq_lens_cpu = self.seq_lens_cpu[:num_reqs]

slot_mapping = common_attn_metadata.slot_mapping[:num_input_tokens]

Expand Down Expand Up @@ -551,19 +570,25 @@ def build_req_metadata_for_drafting(

if metadata_cache is not None and "local_query_start_loc" in metadata_cache:
# Hit: local token metadata computed by a sibling kv-cache group
# of the same draft step. Values are identical across groups, so
# the cached (cloned) tensors are reused directly.
# of the same draft step. Values are identical across groups; copy
# them into this builder's stable per-draft buffers so the returned
# views keep fixed addresses across aclgraph capture/replay.
local_start = metadata_cache["local_start"]
local_end_with_pad = metadata_cache["local_end_with_pad"]
tokens_per_rank = metadata_cache["tokens_per_rank"]
num_tokens_pad = metadata_cache["num_tokens_pad"]
local_query_start_loc = metadata_cache["local_query_start_loc"]
local_seq_lens = metadata_cache["local_seq_lens"]
max_local_query_len = metadata_cache["max_local_query_len"]
max_local_seq_lens = metadata_cache["max_local_seq_lens"]
local_cos = metadata_cache["local_cos"]
local_sin = metadata_cache["local_sin"]
start_pos = metadata_cache["start_pos"]
self.spec_local_query_start_loc[draft_index - 1][: num_reqs + 1].copy_(
metadata_cache["local_query_start_loc"]
)
self.spec_local_seq_lens[draft_index - 1][:num_reqs].copy_(metadata_cache["local_seq_lens"])
self.spec_start_pos[draft_index - 1][:num_reqs].copy_(metadata_cache["start_pos"])
local_query_start_loc = self.spec_local_query_start_loc[draft_index - 1][: num_reqs + 1]
local_seq_lens = self.spec_local_seq_lens[draft_index - 1][:num_reqs]
start_pos = self.spec_start_pos[draft_index - 1][:num_reqs]
else:
(
local_start,
Expand All @@ -581,8 +606,11 @@ def build_req_metadata_for_drafting(
local_seq_lens=self.spec_local_seq_lens[draft_index - 1],
is_noncausal=is_noncausal,
)
local_query_start_loc = local_query_start_loc.clone()
local_seq_lens = local_seq_lens.clone()
# NOTE: no .clone() here. ACL graph capture bakes the addresses of
# these tensors into the draft graph, so every replay must read the
# freshly built values from the same (stable) buffer addresses.
# Returning a fresh clone would leave the graph reading stale
# capture-time metadata and cause illegal device memory accesses.
local_cos = cos.pad_to(num_tokens_pad)[local_start:local_end_with_pad]
local_sin = sin.pad_to(num_tokens_pad)[local_start:local_end_with_pad]

Expand All @@ -597,21 +625,24 @@ def build_req_metadata_for_drafting(
max_local_query_len = max(1, int(local_seq_lens_q_cpu.max().item()))
max_local_seq_lens = max(1, int(local_seq_lens_cpu.max().item()))

start_pos = self.seq_lens[:num_reqs] - seq_lens_q
start_pos = self.spec_start_pos[draft_index - 1][:num_reqs]
start_pos.copy_(self.seq_lens[:num_reqs] - seq_lens_q)

if metadata_cache is not None:
# Store value snapshots: the persistent buffers are refilled on
# every step, so sibling groups must copy the values now.
metadata_cache.update(
local_start=local_start,
local_end_with_pad=local_end_with_pad,
tokens_per_rank=tokens_per_rank,
num_tokens_pad=num_tokens_pad,
local_query_start_loc=local_query_start_loc,
local_seq_lens=local_seq_lens,
local_query_start_loc=local_query_start_loc.clone(),
local_seq_lens=local_seq_lens.clone(),
max_local_query_len=max_local_query_len,
max_local_seq_lens=max_local_seq_lens,
local_cos=local_cos,
local_sin=local_sin,
start_pos=start_pos,
start_pos=start_pos.clone(),
)

dspark_swa_indices = None
Expand Down Expand Up @@ -702,6 +733,11 @@ def build_req_metadata_for_drafting(
sas_head_dim=head_dim,
sas_metadata=sas_metadata,
)
# Cache sas_metadata in the per-draft-index buffer so the tensor
# address stays stable across aclgraph capture/replay, regardless of
# whether it came from the device metadata kernel or a sibling group.
self.spec_sas_metadata[draft_index - 1][:SAS_METADATA_SIZE].copy_(sas_metadata[:SAS_METADATA_SIZE])
Comment thread
Csrayz marked this conversation as resolved.
sas_metadata = self.spec_sas_metadata[draft_index - 1]

cp_metadata = DSACPMetadata(
local_query_start_loc=local_query_start_loc,
Expand Down Expand Up @@ -1260,8 +1296,8 @@ def _build_sas_metadata(

metadata = metadata_op(**kw)
self.common_ratio_to_sas_metadata[cache_key] = metadata
self.req_sas_metadata[:1024] = metadata
return self.req_sas_metadata[:1024]
self.req_sas_metadata[:SAS_METADATA_SIZE] = metadata
Comment thread
Csrayz marked this conversation as resolved.
return self.req_sas_metadata[:SAS_METADATA_SIZE]

def _build_qli_metadata(
self,
Expand Down Expand Up @@ -1310,8 +1346,8 @@ def _build_qli_metadata(
device=str(self.seqused_q.device),
)
self.common_ratio_to_sas_metadata[cache_key] = metadata
self.req_qli_metadata[:1024] = metadata
return self.req_qli_metadata[:1024]
self.req_qli_metadata[:SAS_METADATA_SIZE] = metadata
Comment thread
Csrayz marked this conversation as resolved.
return self.req_qli_metadata[:SAS_METADATA_SIZE]

def build_for_graph_capture(
self,
Expand Down Expand Up @@ -1495,6 +1531,18 @@ def process_weights_after_loading(self, act_dtype: torch.dtype):
if self.enable_dsa_cp_full_o_proj:
self._enable_o_proj_full_weight_switch()

@staticmethod
def update_graph_params(
update_stream,
forward_context,
num_tokens,
vllm_config=None,
speculative_config=None,
draft_attn_metadatas=None,
):
# DSA-CP does not need to update graph params.
pass

@staticmethod
def _get_weight_switch_method(layer: torch.nn.Module) -> WeightSwitchMixin:
quant_method = layer.quant_method
Expand Down
Loading