[Perf] Stack dspark dense draft per-layer ctx KV projection into one GEMM (#31986)
This commit is contained in:
@@ -4,11 +4,13 @@ import logging
|
||||
from typing import Callable, Iterable, Optional, Tuple
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from torch import nn
|
||||
|
||||
from sglang.srt.distributed.communication_op import tensor_model_parallel_all_gather
|
||||
from sglang.srt.model_loader.weight_utils import default_weight_loader
|
||||
from sglang.srt.models.dflash import DFlashDraftModel
|
||||
from sglang.srt.speculative.dflash_utils import can_dflash_slice_qkv_weight
|
||||
from sglang.srt.speculative.dspark_components.dspark_config import (
|
||||
parse_dspark_draft_config,
|
||||
)
|
||||
@@ -464,6 +466,42 @@ class DSparkDraftMixin:
|
||||
f"or disable the confidence head (enable_confidence_head=False)."
|
||||
)
|
||||
|
||||
def _stacked_ctx_kv_params(self) -> Optional[dict]:
|
||||
"""Stack every layer's KV projection into one weight (exact: the input
|
||||
hidden is shared, so concatenating output columns is equivalent).
|
||||
Cached; None (per-layer fallback) when a QKV weight cannot be sliced
|
||||
(quantized) or layers disagree on norm epsilon / bias presence.
|
||||
"""
|
||||
cached = getattr(self, "_stacked_ctx_kv_cache", False)
|
||||
if cached is not False:
|
||||
return cached
|
||||
weights, biases, k_norm_weights = [], [], []
|
||||
eps = None
|
||||
for layer in self.layers:
|
||||
attn = layer.self_attn
|
||||
can_slice, _ = can_dflash_slice_qkv_weight(attn.qkv_proj)
|
||||
if not can_slice or eps not in (None, attn.k_norm.variance_epsilon):
|
||||
self._stacked_ctx_kv_cache = None
|
||||
return None
|
||||
eps = attn.k_norm.variance_epsilon
|
||||
kv_slice = slice(attn.q_size, attn.q_size + 2 * attn.kv_size)
|
||||
weights.append(attn.qkv_proj.weight[kv_slice])
|
||||
biases.append(
|
||||
attn.qkv_proj.bias[kv_slice] if attn.qkv_proj.bias is not None else None
|
||||
)
|
||||
k_norm_weights.append(attn.k_norm.weight)
|
||||
has_bias = [b is not None for b in biases]
|
||||
if any(has_bias) and not all(has_bias):
|
||||
self._stacked_ctx_kv_cache = None
|
||||
return None
|
||||
self._stacked_ctx_kv_cache = {
|
||||
"weight": torch.cat(weights, dim=0),
|
||||
"bias": torch.cat(biases, dim=0) if all(has_bias) else None,
|
||||
"k_norm_weight": torch.stack(k_norm_weights, dim=0).float(),
|
||||
"eps": eps,
|
||||
}
|
||||
return self._stacked_ctx_kv_cache
|
||||
|
||||
def write_target_hidden_kv(
|
||||
self,
|
||||
*,
|
||||
@@ -475,13 +513,22 @@ class DSparkDraftMixin:
|
||||
commit_lens: Optional[torch.Tensor] = None,
|
||||
) -> None:
|
||||
ctx_hidden = self.project_target_hidden(target_hidden)
|
||||
for layer in self.layers:
|
||||
stacked = self._stacked_ctx_kv_params()
|
||||
if stacked is not None:
|
||||
k_all, v_all = self._project_ctx_kv_stacked(
|
||||
ctx_hidden=ctx_hidden, positions=positions, stacked=stacked
|
||||
)
|
||||
for i, layer in enumerate(self.layers):
|
||||
attn = layer.self_attn
|
||||
k, v = attn.kv_proj_only(ctx_hidden)
|
||||
k = attn.apply_k_norm(k)
|
||||
k = attn.apply_k_rope(positions, k)
|
||||
k = k.view(-1, attn.num_kv_heads, attn.head_dim)
|
||||
v = v.view(-1, attn.num_kv_heads, attn.head_dim)
|
||||
if stacked is not None:
|
||||
k = k_all[i]
|
||||
v = v_all[i]
|
||||
else:
|
||||
k, v = attn.kv_proj_only(ctx_hidden)
|
||||
k = attn.apply_k_norm(k)
|
||||
k = attn.apply_k_rope(positions, k)
|
||||
k = k.view(-1, attn.num_kv_heads, attn.head_dim)
|
||||
v = v.view(-1, attn.num_kv_heads, attn.head_dim)
|
||||
if cache_loc_2d is not None and commit_lens is not None:
|
||||
pool.set_kv_buffer_prefix_valid(
|
||||
attn.attn,
|
||||
@@ -502,6 +549,50 @@ class DSparkDraftMixin:
|
||||
attn.attn.v_scale,
|
||||
)
|
||||
|
||||
def _project_ctx_kv_stacked(
|
||||
self,
|
||||
*,
|
||||
ctx_hidden: torch.Tensor,
|
||||
positions: torch.Tensor,
|
||||
stacked: dict,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
attn0 = self.layers[0].self_attn
|
||||
num_layers = len(self.layers)
|
||||
kv_size = attn0.kv_size
|
||||
head_dim = attn0.head_dim
|
||||
num_kv_heads = attn0.num_kv_heads
|
||||
tokens = ctx_hidden.shape[0]
|
||||
|
||||
kv_all = F.linear(ctx_hidden, stacked["weight"], stacked["bias"])
|
||||
kv_all = kv_all.view(tokens, num_layers, 2, kv_size)
|
||||
# Batched per-head k-norm across layers (fp32 variance + weight, cast back).
|
||||
k32 = (
|
||||
kv_all[:, :, 0, :]
|
||||
.reshape(tokens, num_layers, num_kv_heads, head_dim)
|
||||
.to(torch.float32)
|
||||
)
|
||||
variance = k32.pow(2).mean(dim=-1, keepdim=True)
|
||||
k32 = k32 * torch.rsqrt(variance + stacked["eps"])
|
||||
k32 = k32 * stacked["k_norm_weight"].view(1, num_layers, 1, head_dim)
|
||||
k_all = k32.to(ctx_hidden.dtype)
|
||||
# One RoPE over all layers' heads (shared rotary params + positions).
|
||||
k_flat = k_all.reshape(tokens, num_layers * kv_size)
|
||||
dummy_q = k_flat.new_empty(k_flat.shape)
|
||||
_, k_flat = attn0.rotary_emb(positions, dummy_q, k_flat)
|
||||
# [layers, tokens, heads, dim]: per-layer slices are contiguous views.
|
||||
k_all = (
|
||||
k_flat.view(tokens, num_layers, num_kv_heads, head_dim)
|
||||
.permute(1, 0, 2, 3)
|
||||
.contiguous()
|
||||
)
|
||||
v_all = (
|
||||
kv_all[:, :, 1, :]
|
||||
.view(tokens, num_layers, num_kv_heads, head_dim)
|
||||
.permute(1, 0, 2, 3)
|
||||
.contiguous()
|
||||
)
|
||||
return k_all, v_all
|
||||
|
||||
|
||||
class DSparkDraftModel(DSparkDraftMixin, DFlashDraftModel):
|
||||
|
||||
|
||||
Reference in New Issue
Block a user