Files
enginex-ascend-910-vllm/vllm_ascend/patch/worker/patch_qwen3_dflash.py
Sun Ruoxi 7f8a1b1f7a init v0.23.0
Signed-off-by: Sun Ruoxi <sunruoxi@4paradigm.com>
2026-08-27 15:11:51 +08:00

63 lines
2.2 KiB
Python

import torch
import torch.nn.functional as F
from vllm.model_executor.models.qwen3_dflash import DFlashQwen3Model
def precompute_and_store_context_kv(
self,
context_states: torch.Tensor,
context_positions: torch.Tensor,
context_slot_mapping: torch.Tensor | None = None,
) -> None:
if not hasattr(self, "_num_attn_layers"):
self._build_fused_kv_buffers()
num_ctx = context_states.shape[0]
L = self._num_attn_layers
kv = self._kv_size
hd = self._head_dim
nkv = self._num_kv_heads
# --- Fused KV projection (one GEMM for all layers) ---
normed_context_states = self.hidden_norm(context_states)
all_kv_flat = F.linear(normed_context_states, self._fused_kv_weight, self._fused_kv_bias)
# Single contiguous copy that separates K/V and transposes to
# layer-major layout. Result: [2, L, num_ctx, nkv, hd] contiguous.
# Indexing dim-0 gives contiguous [L, num_ctx, nkv, hd] for K and V.
all_kv = all_kv_flat.view(num_ctx, L, 2, nkv, hd).permute(2, 1, 0, 3, 4).contiguous()
all_k = all_kv[0] # [L, num_ctx, nkv, hd], contiguous
all_v = all_kv[1] # [L, num_ctx, nkv, hd], contiguous
# --- Per-layer RMSNorm K (3D: [num_ctx, nkv, hd] per layer) ---
all_k_normed = torch.empty_like(all_k)
for i in range(L):
k_norm_layer = self.layers[i].self_attn.k_norm
all_k_normed[i] = k_norm_layer(all_k[i])
# --- Fused RoPE across all layers ---
# View as [L * num_ctx, kv] so RoPE sees one big batch (no copy).
# In-place RoPE: pass K as the "query" arg with key=None.
all_k_flat = all_k_normed.view(L * num_ctx, kv)
positions_repeated = context_positions.repeat(L)
tmpv = all_k_flat.clone()
self.layers[0].self_attn.rotary_emb(positions_repeated, all_k_flat, tmpv)
if context_slot_mapping is None:
return
# --- Per-layer cache insert ---
all_k_final = all_k_flat.view(L, num_ctx, nkv, hd)
for i in range(L):
attn = self._attn_layers[i]
kv_cache = attn.kv_cache
attn.impl.do_kv_cache_update(
attn,
all_k_final[i],
all_v[i],
kv_cache,
context_slot_mapping,
)
DFlashQwen3Model.precompute_and_store_context_kv = precompute_and_store_context_kv