Skip to content
Open
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
78 changes: 69 additions & 9 deletions vllm_gaudi/attention/backends/hpu_attn.py
Original file line number Diff line number Diff line change
Expand Up @@ -260,6 +260,10 @@ def __init__(
f"heads in the layer. Sinks shape: {sinks.shape}, "
f"num_heads: {num_heads}.")


self.topk_indices_buffer = kwargs.get('topk_indices_buffer')
self.is_sparse = self.topk_indices_buffer is not None

def forward_mha( # type: ignore
self, q: torch.Tensor, latent_vec_k: torch.Tensor, k_cache: torch.Tensor,
attn_metadata: HPUAttentionMetadata) -> torch.Tensor:
Expand Down Expand Up @@ -353,19 +357,75 @@ def forward_mqa( # type: ignore
kv_lora_rank=self.kv_lora_rank)
return output

@torch.compiler.disable
def forward_mqa_sparse(self, q, k_cache, attn_metadata, topk_indices):
"""Sparse MLA decode: attend to only the top-K cache entries.

topk_indices uses -1 as the "no token" sentinel (matching upstream's
SparseAttnIndexer convention: see vllm.model_executor.layers.
sparse_attn_indexer). Rows/slots equal to -1 are padding: their cache
read is discarded and their attention score is masked out below,
instead of relying on seq_lens_tensor/context_lens_tensor, which are
always None on HPU decode for models without mamba-like layers.
"""
if isinstance(k_cache, tuple):
k_cache = k_cache[0]
batch_size = q.shape[0]
topk = topk_indices.shape[1]
topk_indices = topk_indices[:batch_size]
pad_mask = topk_indices == -1

# Gather top-K latent KV from cache using physical slot indices.
# Clamp the -1 sentinel to a safe index; the gathered content for
# padded slots is discarded by the mask below.
flat_idx = topk_indices.clamp(min=0).reshape(-1)
selected = k_cache[flat_idx].view(batch_size, topk, -1)

# Decompress KV
k_c, k_pe = selected.split([self.kv_lora_rank, self.qk_rope_head_dim], dim=-1)
kv_nope = self.kv_b_proj(k_c.reshape(-1, self.kv_lora_rank))[0]
kv_nope = kv_nope.view(batch_size, topk, self.num_heads, self.qk_nope_head_dim + self.v_head_dim)
k_nope, v = kv_nope.split([self.qk_nope_head_dim, self.v_head_dim], dim=-1)
k_pe_exp = k_pe.unsqueeze(2).expand(-1, -1, self.num_heads, -1)
key = torch.cat([k_nope, k_pe_exp], dim=-1)

# q: [B, H, D] -> [B, H, 1, D]
# key: [B, T, H, D] -> [B, H, T, D]
key = key.permute(0, 2, 1, 3)
v = v.permute(0, 2, 1, 3)
attn = torch.matmul(q.unsqueeze(2), key.transpose(-1, -2))

# Use finite minimum value to avoid NaN from softmax(-inf / -inf)
# when an entire row is masked (empty/padded request).
attn = attn.masked_fill(pad_mask.unsqueeze(1).unsqueeze(2), torch.finfo(attn.dtype).min)

attn = torch.softmax(attn * self.scale, dim=-1)
# Zero output for rows where every top-k slot is the -1 sentinel.
# `.to(attn.dtype)` (not `.float()`): attn is bf16 here, and
# torch.matmul below does not type-promote against a bf16 `v`.
# 4D view [B, 1, 1, 1] to broadcast correctly against attn's
# [B, H, 1, T]; a 3D [B, 1, 1] view would silently misalign against
# the heads dim instead of the batch dim.
empty_mask = pad_mask.all(dim=-1).view(batch_size, 1, 1, 1)
attn = attn * (~empty_mask).to(attn.dtype)
out = torch.matmul(attn, v).squeeze(2)
return out.reshape(-1, self.num_heads * self.v_head_dim)

# NOTE(Xinyu): Make the loaded weight contiguous to avoid the transpose
# during each graph execution
def process_weights_after_loading(self, act_dtype: torch.dtype):
super().process_weights_after_loading(act_dtype)
# Since vllm#48251 the base process_weights_after_loading registers
# W_UV and W_UK_T as nn.Parameter (via replace_parameter), so a plain
# tensor reassignment raises TypeError. Route the contiguous/.to("hpu")
# update back through replace_parameter to preserve the registration.
# When INC CPU-first loading is active the source weights live on CPU,
# making these derived tensors CPU-resident too — which then causes a
# device mismatch at the bmm calls in forward. Explicitly place on HPU.
replace_parameter(self, "W_UV", self.W_UV.contiguous().to("hpu"))
replace_parameter(self, "W_UK_T", self.W_UK_T.contiguous().to("hpu"))
# Upstream moved MLA weight packing off the attention impl onto the
# MLAAttention module, so the base process_weights_after_loading is now
# the AttentionImplBase no-op and W_UV/W_UK_T are no longer registered on
# the impl. Only massage them when this impl actually owns them: they are
# registered as nn.Parameter, so route through replace_parameter to keep
# the registration; .to("hpu") avoids a CPU/HPU bmm device mismatch under
# INC CPU-first loading.
for name in ("W_UV", "W_UK_T"):
weight = getattr(self, name, None)
if weight is not None:
replace_parameter(self, name, weight.contiguous().to("hpu"))

# NOTE(Chendi): PR25184 using output buffer as default, which can't be used in HPU Graph,
# so we override and always return a new tensor
Expand Down
29 changes: 21 additions & 8 deletions vllm_gaudi/attention/oot_mla.py
Original file line number Diff line number Diff line change
Expand Up @@ -55,6 +55,21 @@ def __init__(self, *args, **kwargs):
self.fused_scaled_dot_product_attention = None if HPUFusedSDPA is None \
else ModuleFusedSDPA(HPUFusedSDPA)

def bind_kv_cache(self, kv_cache) -> None:
"""Store the HPU per-layer KV-cache tuple unchanged.

vllm#51718 changed ``vllm.v1.worker.utils.bind_kv_cache`` to delegate
binding to each layer's own ``bind_kv_cache`` and added
``MLAAttention.bind_kv_cache`` whose body is
``self.kv_cache = kv_cache.squeeze(1)`` — it assumes a single packed
``[B, H=1, N, C]`` tensor. On HPU the model runner allocates a per-layer
``(key_cache, value_cache, key_scales, value_scales)`` tuple for MLA
layers and ``forward_impl`` consumes it directly (``kv_cache[0]``), so
the upstream ``.squeeze`` raises ``AttributeError: 'tuple' object has no
attribute 'squeeze'``. Keep the tuple as-is for the HPU path.
"""
self.kv_cache = kv_cache

def forward(
self,
q: torch.Tensor,
Expand Down Expand Up @@ -150,13 +165,14 @@ def forward_impl(
if is_prefill:
output = self.impl.forward_mha(q, latent_vec_k, kv_cache, attn_metadata)
return output
elif self.use_sparse and getattr(self.impl, 'topk_indices_buffer', None) is not None:
output = self.impl.forward_mqa_sparse(q, kv_cache, attn_metadata, self.impl.topk_indices_buffer)
return output
else:
output = self.impl.forward_mqa(decode_ql_nope, q_pe, kv_cache, attn_metadata)
output = self._v_up_proj(output)
return output
# NOTE(Xinyu): Make the loaded weight contiguous to avoid the transpose

# during each graph execution
def process_weights_after_loading(self, act_dtype: torch.dtype):
# HPU-specific: when VLLM_HPU_FORCE_CHANNEL_FP8=True (default), block-quantized
# FP8 weights (e.g. kv_b_proj in DeepSeek-R1) are converted to channel-wise FP8.
Expand Down Expand Up @@ -295,10 +311,7 @@ def __init__(
# None for DeepSeek-V2/R1 (no gate proj), leaving the HPU path unchanged.
self.g_proj = mla_modules.g_proj

# DSA sparse attention is not implemented on HPU: sparse layers run as
# dense MLA and the indexer must never be invoked (its kernels and the
# DeepseekV32IndexerBackend are CUDA-only).
self.skip_topk = skip_topk or self.is_sparse
self.skip_topk = skip_topk
# vllm#45964 (DCP query replication) added `self.dcp_q_replicate`, which
# the base MultiHeadLatentAttentionWrapper.forward (inherited here, since
# we do not override forward) reads and forwards to mla_attn. Because we
Expand Down Expand Up @@ -328,8 +341,8 @@ def __init__(
quant_config=quant_config,
prefix=layer_name,
kv_b_proj=self.kv_b_proj,
# Dense-MLA fallback: never request a sparse backend on HPU.
use_sparse=False,
use_sparse=self.is_sparse,
indexer=self.indexer,
topk_indices_buffer=mla_modules.topk_indices_buffer,
non_causal_multi_token_decode=non_causal_multi_token_decode,
)
4 changes: 3 additions & 1 deletion vllm_gaudi/extension/features.py
Original file line number Diff line number Diff line change
Expand Up @@ -97,7 +97,9 @@ def get_features():
ValueFromList('prompt_attn_impl', supported_attn_impls),
Value('skip_warmup', False),
Value('merged_prefill', False),
Value('use_contiguous_pa', Disabled('prefix_caching'), env_var='VLLM_CONTIGUOUS_PA'),
Value('use_contiguous_pa',
All(Disabled('prefix_caching'), Not(ModelType('glm_moe_dsa'))),
env_var='VLLM_CONTIGUOUS_PA'),
Value('use_bucketing', True, env_var='VLLM_ENABLE_BUCKETING'),
Value('bucketing_strategy',
'exp',
Expand Down
65 changes: 51 additions & 14 deletions vllm_gaudi/models/deepseek_v2.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,8 @@

from vllm.distributed import get_pp_group
from vllm.model_executor.models import deepseek_v2
from vllm.model_executor.models.deepseek_v2 import (DeepseekV32IndexerCache, Indexer)
from vllm.model_executor.layers.sparse_attn_indexer import SparseAttnIndexer
from vllm.sequence import IntermediateTensors


Expand Down Expand Up @@ -87,24 +89,59 @@ def _hpu_deepseek_v2_model_forward(
# Applies to DeepseekV2/V3/Deepseek/GlmMoe/DSA — all share model_cls = DeepseekV2Model.
deepseek_v2.DeepseekV2Model.forward = _hpu_deepseek_v2_model_forward

_orig_deepseek_v2_model_load_weights = deepseek_v2.DeepseekV2Model.load_weights
# ---------------------------------------------------------------------------
# DSA / Indexer enablement on HPU
# ---------------------------------------------------------------------------

# --- IndexerCache: BF16 storage instead of FP8 uint8 -----------------------
_orig_indexer_cache_init = DeepseekV32IndexerCache.__init__

def _hpu_deepseek_v2_model_load_weights(self, weights):
"""Drop GLM-5 DSA shared-indexer projection weights (`indexers_proj`).

vLLM's DeepseekV2Model has no module for them, and on HPU DSA layers run
as dense MLA with the indexer never executed, so the projection that
shares indexer K caches across layers is dead weight here.
"""
def _hpu_indexer_cache_init(self, head_dim, dtype, prefix, cache_config):
if dtype == torch.uint8:
head_dim = head_dim * 128 // (128 + 4)
dtype = torch.bfloat16
_orig_indexer_cache_init(self, head_dim, dtype, prefix, cache_config)


DeepseekV32IndexerCache.__init__ = _hpu_indexer_cache_init


def _hpu_indexer_cache_get_attn_backend(self):
from vllm_gaudi.attention.backends.hpu_attn import HPUMLAAttentionBackend
return HPUMLAAttentionBackend


DeepseekV32IndexerCache.get_attn_backend = _hpu_indexer_cache_get_attn_backend


# --- Indexer.forward: BF16 path, skip FP8 quantization ---------------------
def _hpu_indexer_forward(self, hidden_states, qr, positions, rotary_emb):
q, _ = self.wq_b(qr)
q = q.view(-1, self.n_head, self.head_dim)
q_pe, q_nope = torch.split(q, [self.rope_dim, self.head_dim - self.rope_dim], dim=-1)
kw, _ = self.wk_weights_proj(hidden_states)
kw = kw.reshape(-1, kw.shape[-1])
k, weights = torch.split(kw, [self.head_dim, self.n_head], dim=-1)
k = self.k_norm(k.contiguous())
k_pe, k_nope = torch.split(k, [self.rope_dim, self.head_dim - self.rope_dim], dim=-1)
q_pe, k_pe = rotary_emb(positions, q_pe, k_pe.unsqueeze(1))
q_pe = q_pe.reshape(-1, self.n_head, self.rope_dim)
k_pe = k_pe.reshape(-1, self.rope_dim)
k_nope = k_nope.reshape(-1, self.head_dim - self.rope_dim)
q = torch.cat([q_pe, q_nope], dim=-1)
k = torch.cat([k_pe, k_nope], dim=-1)
weights = weights.reshape(-1, self.n_head) * self.softmax_scale * self.n_head_scale
return self.indexer_op(hidden_states, q, k, weights)


Indexer.forward = _hpu_indexer_forward

def _filtered(ws):
for name, weight in ws:
if ".indexers_proj." in name:
continue
yield name, weight

return _orig_deepseek_v2_model_load_weights(self, _filtered(weights))
# --- SparseAttnIndexer: dispatch to HPU forward -----------------------------
def _hpu_sparse_indexer_forward_native(self, hidden_states, q_quant, k, weights):
from vllm_gaudi.ops.hpu_sparse_attn_indexer import forward_hpu
return forward_hpu(self, hidden_states, q_quant, k, weights)


deepseek_v2.DeepseekV2Model.load_weights = _hpu_deepseek_v2_model_load_weights
SparseAttnIndexer.forward_native = _hpu_sparse_indexer_forward_native
Comment thread
jkaniecki marked this conversation as resolved.
88 changes: 88 additions & 0 deletions vllm_gaudi/ops/hpu_sparse_attn_indexer.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,88 @@
# SPDX-License-Identifier: Apache-2.0

import torch
from vllm.forward_context import get_forward_context


def _fill_invalid(buf, n, device):
"""Fill topk_indices_buffer rows 0..n-1 with the upstream -1 sentinel.

-1 means "no token" (see vllm.model_executor.layers.sparse_attn_indexer),
consumed by forward_mqa_sparse to mask out padding instead of gathering it.
"""
buf[:n, :].fill_(-1)


@torch.compiler.disable
def forward_hpu(self, hidden_states, q, k, weights):
"""HPU SparseAttnIndexer: per-request QK BF16 scoring + torch.topk."""
forward_context = get_forward_context()
attn_metadata = forward_context.attn_metadata
kv_cache = self.k_cache.kv_cache
if isinstance(kv_cache, tuple):
kv_cache = kv_cache[0]
block_size = attn_metadata.block_size
slot_mapping = attn_metadata.slot_mapping.flatten()

if kv_cache is None or kv_cache.numel() == 0:

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

On this branch I believe this early return is the only path this function ever takes. get_kv_cache_spec in vllm_gaudi/v1/worker/hpu_model_runner.py:1717 handles only MambaBase, Attention and MLAAttention, and DeepseekV32IndexerCache is none of those - it is a plain AttentionLayerBase (upstream vllm/model_executor/models/deepseek_v2.py:604). So the indexer cache layer never gets a spec, never gets an allocation, and self.k_cache.kv_cache keeps the torch.tensor([]) it is initialised with at deepseek_v2.py:609. kv_cache.numel() == 0 is then true on every step: the index_copy_ below never inserts K, and the returned indices are always 0..topk_tokens-1.

The main-branch sibling (#1760) adds exactly the missing branch to get_kv_cache_spec, and this PR is otherwise byte-identical to it in all five shared files. The new check_runner_kv_caches_multi_layer no-op in platform.py only has an effect once the indexer cache is in kv_caches, which also points at that change being intended here.

Please cherry-pick the elif isinstance(attn_module, AttentionLayerBase) branch from #1760. One thing worth confirming when you do: it adds a second entry per DSA layer to runner_kv_caches, which is the case upstream's own TODO at vllm/v1/worker/utils.py:529-541 describes as unanalyzed - have you checked that nothing on the HPU side indexes runner_kv_caches positionally?

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

runner_kv_caches is not consumed by any function / read by any other variable.
kv_caches is also binded to model_runner.kv_Caches, which is consumed by defragmenter. I switched off the defrag for this model (and cont pa) in features.py - we don't want it to manipulate indexer cache. model_runner.kv_caches is also used in hpu_worker.py:L531-550, but there only first item from kv cache is taken and it's not indexer cache.

_fill_invalid(self.topk_indices_buffer, q.shape[0], q.device)
return self.topk_indices_buffer

if not self.skip_k_cache_insert:
kv_cache.index_copy_(0, slot_mapping[:k.shape[0]], k)

if attn_metadata.is_prompt:
_fill_invalid(self.topk_indices_buffer, q.shape[0], q.device)
return self.topk_indices_buffer

batch_size = q.shape[0]
block_list = attn_metadata.block_list
block_groups = attn_metadata.block_groups
block_usage = attn_metadata.block_usage

if block_list is None or block_groups is None or block_usage is None:
# No block-table metadata available for decode (unexpected); fall back
# to the sentinel fill so shapes stay valid instead of crashing.
_fill_invalid(self.topk_indices_buffer, batch_size, q.device)
return self.topk_indices_buffer

pos_range = torch.arange(block_size, device=block_list.device)
# block_usage is stored in model dtype (see hpu_model_runner.py); round to
# get an exact per-block valid-token count.
block_usage_long = block_usage.round().long()

for i in range(batch_size):
# Select this request's physical blocks directly via block_groups
# rather than assuming block_list is an unpadded per-request
# concatenation: with contiguous PA, blocks are scattered/reordered by
# physical block id, not laid out sequentially per request.
request_mask = block_groups == i
request_blocks = block_list[request_mask]
if request_blocks.numel() == 0:
self.topk_indices_buffer[i] = -1
continue
request_usage = block_usage_long[request_mask]

all_slots = (request_blocks.unsqueeze(1) * block_size + pos_range.unsqueeze(0)).reshape(-1)
valid_mask = (pos_range.unsqueeze(0) < request_usage.unsqueeze(1)).reshape(-1)
valid_slots = all_slots[valid_mask]
seq_len = valid_slots.shape[0]

if seq_len == 0:
self.topk_indices_buffer[i] = -1
continue

if seq_len <= self.topk_tokens:
self.topk_indices_buffer[i, :seq_len] = valid_slots
if seq_len < self.topk_tokens:
self.topk_indices_buffer[i, seq_len:] = -1
continue

k_all = kv_cache[valid_slots].to(torch.float32)
q_i = q[i].to(torch.float32)
logits = torch.mm(q_i.reshape(q_i.shape[0], -1), k_all.T)
scores = (torch.sigmoid(logits) * weights[i].to(torch.float32).unsqueeze(-1)).sum(0)
_, local_indices = torch.topk(scores, self.topk_tokens)
self.topk_indices_buffer[i] = valid_slots[local_indices]

return self.topk_indices_buffer
7 changes: 5 additions & 2 deletions vllm_gaudi/platform.py
Original file line number Diff line number Diff line change
Expand Up @@ -87,8 +87,7 @@ def get_attn_backend_cls(
if attn_selector_config.use_sparse:
if not attn_selector_config.use_mla:
raise NotImplementedError("Sparse Attention is not supported on HPU.")
logger.warning("Sparse attention (DSA) is not implemented on HPU; running DSA layers as dense MLA "
"(exact for sequences up to index_topk tokens, approximate beyond).")
logger.info("Using HPU DSA (Dynamic Sparse Attention) with BF16 indexer.")

if attn_selector_config.use_mla:
logger.info("Using HPUAttentionMLA backend.")
Expand All @@ -99,6 +98,10 @@ def get_attn_backend_cls(
return ("vllm_gaudi.v1.attention.backends."
"hpu_attn.HPUAttentionBackendV1")

@classmethod
def check_runner_kv_caches_multi_layer(cls) -> None:
pass # DSA indexer cache shares layer index with MLA attention

@classmethod
def is_async_output_supported(cls, enforce_eager: Optional[bool]) -> bool:
return True
Expand Down
4 changes: 4 additions & 0 deletions vllm_gaudi/v1/worker/hpu_model_runner.py
Original file line number Diff line number Diff line change
Expand Up @@ -1772,6 +1772,10 @@ def get_kv_cache_spec(self) -> dict[str, KVCacheSpec]:
dtype=self.kv_cache_dtype,
cache_dtype_str=cache_dtype_str,
)
elif isinstance(attn_module, AttentionLayerBase):
spec = attn_module.get_kv_cache_spec(self.vllm_config)
if spec is not None:
kv_cache_spec[layer_name] = spec

return kv_cache_spec

Expand Down