[spec decoding] support inkling dspark (#31847)

This commit is contained in:
Qiaolin Yu
2026-08-10 13:21:35 -07:00
committed by GitHub
parent e54c153ba6
commit 733c05c887
6 changed files with 687 additions and 22 deletions
@@ -469,6 +469,210 @@ def fused_conv_window_scatter_with_mask(
)
_CONV_MULTI_MAX_TYPES = 8
_CONV_MULTI_META_COLS = 12
_conv_multi_meta_cache: dict = {}
@triton.jit
def _fused_conv_window_scatter_multi_kernel(
meta_ptr, # int64 [num_types, 12]: src_ptr, dst_ptr, elem, s_l, s_r, s_s, s_d, s_w, d_l, d_r, block_start, last_axis
idx1_ptr,
step1_ptr,
idx2_ptr,
step2_ptr,
n1,
src_req_size,
src_step_size,
dst_req_size,
NUM_TYPES: tl.constexpr,
META_COLS: tl.constexpr,
BLOCK_SIZE: tl.constexpr,
):
pid_req = tl.program_id(0)
pid_layer = tl.program_id(1).to(tl.int64)
pid_block = tl.program_id(2).to(tl.int64)
is1 = pid_req < n1
is2 = pid_req >= n1
off1 = pid_req
off2 = pid_req - n1
s1 = tl.load(step1_ptr + off1, mask=is1, other=-1).to(tl.int64)
s2 = tl.load(step2_ptr + off2, mask=is2, other=-1).to(tl.int64)
step_idx = tl.where(is2, s2, s1)
if step_idx < 0:
return
d1 = tl.load(idx1_ptr + off1, mask=is1, other=-1).to(tl.int64)
d2 = tl.load(idx2_ptr + off2, mask=is2, other=-1).to(tl.int64)
dst_idx = tl.where(is2, d2, d1)
src_idx = tl.where(is2, off2, off1).to(tl.int64)
if not (
(dst_idx >= 0)
& (dst_idx < dst_req_size)
& (src_idx < src_req_size)
& (step_idx < src_step_size)
):
return
for t in tl.static_range(NUM_TYPES):
block_start = tl.load(meta_ptr + t * META_COLS + 10)
block_end = tl.load(
meta_ptr + (t + 1) * META_COLS + 10,
mask=t + 1 < NUM_TYPES,
other=2147483647,
)
if (pid_block >= block_start) & (pid_block < block_end):
src_ptr = tl.load(meta_ptr + t * META_COLS + 0).to(
tl.pointer_type(tl.bfloat16)
)
dst_ptr = tl.load(meta_ptr + t * META_COLS + 1).to(
tl.pointer_type(tl.bfloat16)
)
elem_per_entry = tl.load(meta_ptr + t * META_COLS + 2)
src_layer_stride = tl.load(meta_ptr + t * META_COLS + 3)
src_req_stride = tl.load(meta_ptr + t * META_COLS + 4)
src_step_stride = tl.load(meta_ptr + t * META_COLS + 5)
src_dim_stride = tl.load(meta_ptr + t * META_COLS + 6)
src_win_stride = tl.load(meta_ptr + t * META_COLS + 7)
dst_layer_stride = tl.load(meta_ptr + t * META_COLS + 8)
dst_req_stride = tl.load(meta_ptr + t * META_COLS + 9)
last_axis = tl.load(meta_ptr + t * META_COLS + 11)
start = (pid_block - block_start) * BLOCK_SIZE
e = start + tl.arange(0, BLOCK_SIZE)
mask = e < elem_per_entry
d = e // last_axis
w = e % last_axis
src_off = (
pid_layer * src_layer_stride
+ src_idx * src_req_stride
+ step_idx * src_step_stride
+ d * src_dim_stride
+ w * src_win_stride
)
dst_off = pid_layer * dst_layer_stride + dst_idx * dst_req_stride + e
data = tl.load(src_ptr + src_off, mask=mask, other=0.0)
tl.store(dst_ptr + dst_off, data, mask=mask)
def _conv_multi_build_meta(pairs, block_size: int):
rows = []
block_start = 0
for dst, src in pairs:
elem = dst.shape[2] * dst.shape[3]
rows.append(
[
src.data_ptr(),
dst.data_ptr(),
elem,
src.stride(0),
src.stride(1),
src.stride(2),
src.stride(3),
src.stride(4),
dst.stride(0),
dst.stride(1),
block_start,
dst.shape[3],
]
)
block_start += triton.cdiv(elem, block_size)
meta = torch.tensor(rows, dtype=torch.int64, device=pairs[0][0].device)
return meta, block_start
def _conv_multi_eligible(pairs) -> bool:
if not (0 < len(pairs) <= _CONV_MULTI_MAX_TYPES):
return False
layers = pairs[0][0].shape[0]
for dst, src in pairs:
if dst.dtype != torch.bfloat16 or src.dtype != torch.bfloat16:
return False
if dst.ndim != 4 or src.ndim != 5:
return False
if dst.shape[0] != layers:
return False
if src.shape[0] != layers or src.shape[3:] != dst.shape[2:]:
return False
if not dst.is_contiguous():
return False
if src.shape[1:3] != pairs[0][1].shape[1:3]:
return False
if dst.shape[1] != pairs[0][0].shape[1]:
return False
return True
def fused_conv_window_scatter_multi(
pairs,
dst_indices_raw: torch.Tensor,
step_indices_raw: torch.Tensor,
dst_indices2_raw: torch.Tensor | None = None,
step_indices2_raw: torch.Tensor | None = None,
) -> None:
"""Single-launch variant of ``fused_conv_window_scatter_with_mask`` over
multiple (dst, src) conv-type pairs and up to two request-index sets (the
accept commit plus the optional interval-crossing track set)."""
n1 = step_indices_raw.shape[0]
n2 = 0 if step_indices2_raw is None else step_indices2_raw.shape[0]
if n1 + n2 == 0:
return
BLOCK_SIZE = 1024
key = tuple(
(dst.data_ptr(), src.data_ptr()) + tuple(src.stride()) + tuple(dst.shape)
for dst, src in pairs
)
cached = _conv_multi_meta_cache.get(key)
if cached is None:
cached = _conv_multi_build_meta(pairs, BLOCK_SIZE)
_conv_multi_meta_cache.clear()
_conv_multi_meta_cache[key] = cached
meta, total_blocks = cached
idx1 = (
dst_indices_raw
if dst_indices_raw.is_contiguous()
else dst_indices_raw.contiguous()
)
st1 = (
step_indices_raw
if step_indices_raw.is_contiguous()
else step_indices_raw.contiguous()
)
if n2 > 0:
idx2 = (
dst_indices2_raw
if dst_indices2_raw.is_contiguous()
else dst_indices2_raw.contiguous()
)
st2 = (
step_indices2_raw
if step_indices2_raw.is_contiguous()
else step_indices2_raw.contiguous()
)
else:
idx2, st2 = idx1, st1
dst0, src0 = pairs[0]
grid = (n1 + n2, dst0.shape[0], total_blocks)
_fused_conv_window_scatter_multi_kernel[grid](
meta,
idx1,
st1,
idx2,
st2,
n1,
src0.shape[1],
src0.shape[2],
dst0.shape[1],
NUM_TYPES=len(pairs),
META_COLS=_CONV_MULTI_META_COLS,
BLOCK_SIZE=BLOCK_SIZE,
)
def scatter_mamba_states_after_mtp_verify(
mamba_caches,
state_indices_tensor: torch.Tensor,
@@ -488,28 +692,38 @@ def scatter_mamba_states_after_mtp_verify(
state_indices_tensor,
last_correct_step_indices,
)
for conv_states, intermediate_conv_window_cache in zip(
mamba_caches.conv, mamba_caches.intermediate_conv_window
):
fused_conv_window_scatter_with_mask(
conv_states,
intermediate_conv_window_cache,
state_indices_tensor,
last_correct_step_indices,
)
if mamba_track_indices is not None:
assert mamba_steps_to_track is not None
if ssm_states.numel() > 0:
if mamba_track_indices is not None:
assert mamba_steps_to_track is not None
fused_mamba_state_scatter_with_mask(
ssm_states,
intermediate_state_cache,
mamba_track_indices,
mamba_steps_to_track,
)
for conv_states, intermediate_conv_window_cache in zip(
mamba_caches.conv, mamba_caches.intermediate_conv_window
):
pairs = list(zip(mamba_caches.conv, mamba_caches.intermediate_conv_window))
if not pairs:
return
if mamba_track_indices is not None:
assert mamba_steps_to_track is not None
if _conv_multi_eligible(pairs):
fused_conv_window_scatter_multi(
pairs,
state_indices_tensor,
last_correct_step_indices,
mamba_track_indices,
mamba_steps_to_track,
)
return
for conv_states, intermediate_conv_window_cache in pairs:
fused_conv_window_scatter_with_mask(
conv_states,
intermediate_conv_window_cache,
state_indices_tensor,
last_correct_step_indices,
)
if mamba_track_indices is not None:
for conv_states, intermediate_conv_window_cache in pairs:
fused_conv_window_scatter_with_mask(
conv_states,
intermediate_conv_window_cache,
@@ -518,6 +732,65 @@ def scatter_mamba_states_after_mtp_verify(
)
@triton.jit
def _fused_commit_track_indices_kernel(
accept_index_ptr,
accept_lens_ptr,
seq_lens_ptr,
last_correct_out_ptr,
track_steps_out_ptr,
dtn,
interval,
HAS_TRACK: tl.constexpr,
):
b = tl.program_id(0).to(tl.int64)
al = tl.load(accept_lens_ptr + b).to(tl.int64)
base = b * dtn
last = tl.load(accept_index_ptr + base + al - 1).to(tl.int64) - base
tl.store(last_correct_out_ptr + b, last)
if HAS_TRACK:
pre = tl.load(seq_lens_ptr + b).to(tl.int64)
post = pre + al
cross = (pre // interval) != (post // interval)
tp = (post // interval) * interval
ti = tp - pre - 1
ti = tl.where(ti < 0, 0, ti)
cand = tl.load(accept_index_ptr + base + ti).to(tl.int64) - base
tl.store(track_steps_out_ptr + b, tl.where(cross, cand, -1))
def fused_commit_track_indices(
accept_index: torch.Tensor,
accept_lens: torch.Tensor,
seq_lens: torch.Tensor | None,
draft_token_num: int,
mamba_track_interval: int,
):
"""Single-launch replacement for the eager index math in
``commit_mamba_states_after_verify`` (index ranges, gathers, floordiv chain)."""
bs = accept_lens.shape[0]
last_correct_step_indices = torch.empty(
bs, dtype=torch.int64, device=accept_lens.device
)
has_track = seq_lens is not None
mamba_steps_to_track = (
torch.empty(bs, dtype=torch.int64, device=accept_lens.device)
if has_track
else last_correct_step_indices
)
_fused_commit_track_indices_kernel[(bs,)](
accept_index,
accept_lens,
seq_lens if has_track else accept_lens,
last_correct_step_indices,
mamba_steps_to_track,
draft_token_num,
mamba_track_interval,
HAS_TRACK=has_track,
)
return last_correct_step_indices, (mamba_steps_to_track if has_track else None)
@triton.jit
def track_mamba_states_all_layers_kernel(
conv_states_ptr, # [num_layers, pool_size, ...] full conv pool
@@ -0,0 +1,126 @@
from typing import Optional
import torch
import triton
import triton.language as tl
@triton.jit
def _fused_kv_norm_rope_write_kernel(
kv_ptr,
meta_ptr,
knw_ptr,
cos_sin_ptr,
pos_ptr,
loc_ptr,
commit_lens_ptr,
locs_row_width,
KV: tl.constexpr,
D: tl.constexpr,
NH: tl.constexpr,
L: tl.constexpr,
EPS: tl.constexpr,
HAS_COMMIT_LENS: tl.constexpr,
):
t = tl.program_id(0).to(tl.int64)
l = tl.program_id(1).to(tl.int64)
if HAS_COMMIT_LENS:
row_b = t // locs_row_width
col = t - row_b * locs_row_width
num_commit = tl.load(commit_lens_ptr + row_b).to(tl.int64)
if col >= num_commit:
return
loc = tl.load(loc_ptr + t).to(tl.int64)
if loc < 0:
return
pos = tl.load(pos_ptr + t).to(tl.int64)
HALF: tl.constexpr = D // 2
half_ar = tl.arange(0, HALF)
d_ar = tl.arange(0, D)
cos = tl.load(cos_sin_ptr + pos * D + half_ar).to(tl.float32)
sin = tl.load(cos_sin_ptr + pos * D + HALF + half_ar).to(tl.float32)
knw1 = tl.load(knw_ptr + l * D + half_ar).to(tl.float32)
knw2 = tl.load(knw_ptr + l * D + HALF + half_ar).to(tl.float32)
k_buf = tl.load(meta_ptr + l * 4 + 0).to(tl.pointer_type(tl.bfloat16))
v_buf = tl.load(meta_ptr + l * 4 + 1).to(tl.pointer_type(tl.bfloat16))
ks0 = tl.load(meta_ptr + l * 4 + 2)
vs0 = tl.load(meta_ptr + l * 4 + 3)
row = kv_ptr + t * (L * 2 * KV) + l * (2 * KV)
for h in tl.static_range(NH):
k = tl.load(row + h * D + d_ar).to(tl.float32)
ms = tl.sum(k * k, 0) / D
inv = 1.0 / tl.sqrt(ms + EPS)
k1 = tl.load(row + h * D + half_ar).to(tl.float32) * inv * knw1
k2 = tl.load(row + h * D + HALF + half_ar).to(tl.float32) * inv * knw2
k1 = k1.to(tl.bfloat16).to(tl.float32)
k2 = k2.to(tl.bfloat16).to(tl.float32)
o1 = k1 * cos - k2 * sin
o2 = k2 * cos + k1 * sin
tl.store(k_buf + loc * ks0 + h * D + half_ar, o1.to(tl.bfloat16))
tl.store(k_buf + loc * ks0 + h * D + HALF + half_ar, o2.to(tl.bfloat16))
v = tl.load(row + KV + h * D + d_ar)
tl.store(v_buf + loc * vs0 + h * D + d_ar, v)
def fused_kv_norm_rope_write(
kv: torch.Tensor,
meta: torch.Tensor,
k_norm_weights: torch.Tensor,
cos_sin_cache: torch.Tensor,
positions: torch.Tensor,
locs: torch.Tensor,
num_layers: int,
kv_size: int,
head_dim: int,
eps: float,
commit_lens: Optional[torch.Tensor] = None,
locs_row_width: Optional[int] = None,
) -> None:
"""Write per-layer normed+roped K and raw V rows into the KV pools.
Rows with loc < 0 are skipped. When commit_lens is given, locs is the
flattened [bs, locs_row_width] verify window and only the first
commit_lens[b] columns of each row are written — the in-kernel
replacement for masking the tail columns to -1 on the host.
"""
T = kv.shape[0]
if T == 0:
return
has_commit_lens = commit_lens is not None
if has_commit_lens != (locs_row_width is not None):
raise ValueError(
"commit_lens and locs_row_width must be passed together, got "
f"commit_lens={'set' if has_commit_lens else None}, "
f"locs_row_width={locs_row_width}."
)
if has_commit_lens:
if commit_lens.numel() * locs_row_width != locs.numel():
raise ValueError(
f"locs must be a flattened [{commit_lens.numel()}, "
f"{locs_row_width}] window, got numel={locs.numel()}."
)
commit_lens_arg = commit_lens.contiguous()
else:
locs_row_width = 1
commit_lens_arg = locs
grid = (T, num_layers)
_fused_kv_norm_rope_write_kernel[grid](
kv,
meta,
k_norm_weights,
cos_sin_cache,
positions.to(torch.int64).contiguous(),
locs.to(torch.int64).contiguous(),
commit_lens_arg,
locs_row_width,
KV=kv_size,
D=head_dim,
NH=kv_size // head_dim,
L=num_layers,
EPS=eps,
HAS_COMMIT_LENS=has_commit_lens,
)
+22
View File
@@ -149,6 +149,13 @@ class DFlashAttention(nn.Module):
)
self.scaling = head_dim**-0.5
rotary = self.rotary_emb
self.use_table_qk_norm_rope = (
not _is_npu
and hasattr(rotary, "cos_sin_cache")
and getattr(rotary, "rotary_dim", None) == head_dim
and getattr(rotary, "is_neox_style", False)
)
self.sliding_window_size, self.attn_type = _get_dflash_layer_attention_params(
config, layer_id
)
@@ -191,6 +198,21 @@ class DFlashAttention(nn.Module):
qkv, _ = self.qkv_proj(hidden_states)
if _is_npu:
q, k, v = self.forward_prepare_npu(positions, hidden_states)
elif self.use_table_qk_norm_rope and qkv.dtype == torch.bfloat16:
from sglang.srt.speculative.dflash_utils import table_qk_norm_rope_
table_qk_norm_rope_(
qkv,
positions,
self.q_norm.weight,
self.k_norm.weight,
self.rotary_emb.cos_sin_cache,
self.num_heads,
self.num_kv_heads,
self.head_dim,
self.q_norm.variance_epsilon,
)
q, k, v = qkv.split([self.q_size, self.kv_size, self.kv_size], dim=-1)
else:
q, k, v = qkv.split([self.q_size, self.kv_size, self.kv_size], dim=-1)
q, k = apply_qk_norm(q, k, self.q_norm, self.k_norm, self.head_dim)
+115
View File
@@ -364,6 +364,8 @@ class DSparkDraftMixin:
def __init__(self, config, quant_config=None, prefix: str = "") -> None:
super().__init__(config=config, quant_config=quant_config, prefix=prefix)
self._fused_kv_write_cache = None
self.logits_mup_width_multiplier = None
dspark_config = parse_dspark_draft_config(draft_hf_config=config)
if not dspark_config.require_markov():
raise ValueError(
@@ -390,11 +392,22 @@ class DSparkDraftMixin:
def compute_base_logits(
self, hidden: torch.Tensor
) -> tuple[torch.Tensor, Optional[torch.Tensor]]:
"""Project the draft's raw final hidden through the target lm_head.
muP targets (Inkling) train the draft against a FOLDED head (weights
pre-divided by logits_mup_width_multiplier) while serving attaches the
target's unfolded head, so the division happens here — exactly once,
keeping base logits in the scale the markov bias and confidence head
were trained against. DSparkWorkerV2 wires the multiplier from the
target config; it stays None for non-muP targets.
"""
if self.lm_head is None:
raise ValueError(
"DSpark dense draft requires the target lm_head "
"(call attach_shared_modules first)."
)
if self.logits_mup_width_multiplier:
hidden = hidden / self.logits_mup_width_multiplier
weight = self.lm_head.weight
if hidden.dtype != weight.dtype:
hidden = hidden.to(weight.dtype)
@@ -466,6 +479,75 @@ class DSparkDraftMixin:
f"or disable the confidence head (enable_confidence_head=False)."
)
def _fused_kv_write_bundle(self, pool):
cached = self._fused_kv_write_cache
if cached is not None and cached[0] == id(pool):
return cached[1]
bundle = self._build_fused_kv_write_bundle(pool)
self._fused_kv_write_cache = (id(pool), bundle)
return bundle
def _build_fused_kv_write_bundle(self, pool):
layers = list(self.layers)
if not layers:
return None
if not (hasattr(pool, "get_key_buffer") and hasattr(pool, "get_value_buffer")):
return None
attn0 = layers[0].self_attn
head_dim = attn0.head_dim
kv_size = attn0.kv_size
rotary = attn0.rotary_emb
if type(rotary).__name__ != "RotaryEmbedding":
return None
if not getattr(rotary, "is_neox_style", False):
return None
if getattr(rotary, "rotary_dim", None) != head_dim:
return None
eps = attn0.k_norm.variance_epsilon
weights, knws, meta_rows = [], [], []
for layer in layers:
attn = layer.self_attn
ok, _ = can_dflash_slice_qkv_weight(attn.qkv_proj)
if not ok:
return None
if attn.qkv_proj.bias is not None:
return None
if attn.attn.k_scale is not None or attn.attn.v_scale is not None:
return None
if attn.head_dim != head_dim or attn.kv_size != kv_size:
return None
if attn.rotary_emb is not rotary and not torch.equal(
attn.rotary_emb.cos_sin_cache, rotary.cos_sin_cache
):
return None
if attn.k_norm.variance_epsilon != eps:
return None
k_buf = pool.get_key_buffer(attn.attn.layer_id)
v_buf = pool.get_value_buffer(attn.attn.layer_id)
nh = kv_size // head_dim
for buf in (k_buf, v_buf):
if buf.dtype != torch.bfloat16:
return None
if buf.shape[1:] != (nh, head_dim):
return None
if buf.stride(1) != head_dim or buf.stride(2) != 1:
return None
kv_slice = slice(attn.q_size, attn.q_size + 2 * attn.kv_size)
w = attn.qkv_proj.weight[kv_slice]
if w.dtype != torch.bfloat16:
return None
weights.append(w)
knws.append(attn.k_norm.weight.data)
meta_rows.append(
[k_buf.data_ptr(), v_buf.data_ptr(), k_buf.stride(0), v_buf.stride(0)]
)
device = weights[0].device
w_all = torch.cat(weights, dim=0).contiguous()
knw = torch.stack(knws).to(device)
meta = torch.tensor(meta_rows, dtype=torch.int64, device=device)
cos_sin = rotary.cos_sin_cache.to(device)
return (w_all, meta, knw, cos_sin, eps, len(layers), kv_size, head_dim)
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).
@@ -513,6 +595,39 @@ class DSparkDraftMixin:
commit_lens: Optional[torch.Tensor] = None,
) -> None:
ctx_hidden = self.project_target_hidden(target_hidden)
bundle = self._fused_kv_write_bundle(pool)
if bundle is not None:
from sglang.kernels.ops.speculative.dspark.fused_kv_write import (
fused_kv_norm_rope_write,
)
w_all, meta, knw, cos_sin, eps, num_layers, kv_size, head_dim = bundle
kv_all = F.linear(ctx_hidden, w_all)
if cache_loc_2d is not None and commit_lens is not None:
locs = cache_loc_2d.reshape(-1)
write_commit_lens = commit_lens
locs_row_width = cache_loc_2d.shape[1]
else:
locs = cache_loc
write_commit_lens = None
locs_row_width = None
fused_kv_norm_rope_write(
kv_all,
meta,
knw,
cos_sin,
positions,
locs,
num_layers,
kv_size,
head_dim,
eps,
commit_lens=write_commit_lens,
locs_row_width=locs_row_width,
)
return
stacked = self._stacked_ctx_kv_params()
if stacked is not None:
k_all, v_all = self._project_ctx_kv_stacked(
+121 -1
View File
@@ -8,6 +8,8 @@ from typing import Any, List, Optional, Tuple
import torch
import torch.nn.functional as F
import triton
import triton.language as tl
from sglang.srt.layers.quantization.unquant import UnquantizedLinearMethod
from sglang.srt.layers.sampler import apply_custom_logit_processor
@@ -570,6 +572,29 @@ def can_dflash_use_fused_qkv_proj(qkv_proj: Any) -> Tuple[bool, str]:
return True, ""
@triton.jit
def _fused_correct_drafts_and_bonus_kernel(
candidates_ptr,
target_predict_ptr,
num_correct_drafts_ptr,
bonus_tokens_ptr,
block_size,
BLOCK: tl.constexpr,
):
b = tl.program_id(0).to(tl.int64)
offs = tl.arange(0, BLOCK)
in_row = offs < block_size - 1
drafts = tl.load(candidates_ptr + b * block_size + 1 + offs, mask=in_row, other=-1)
targets = tl.load(target_predict_ptr + b * block_size + offs, mask=in_row, other=-2)
eq = (drafts == targets) & in_row
# Leading-match count = index of the first mismatch lane; lanes past the
# row and all-match rows both resolve to block_size - 1 via the min.
num_correct = tl.min(tl.where(eq, BLOCK, offs), 0)
bonus_token = tl.load(target_predict_ptr + b * block_size + num_correct)
tl.store(num_correct_drafts_ptr + b, num_correct.to(tl.int32))
tl.store(bonus_tokens_ptr + b, bonus_token.to(tl.int64))
def compute_dflash_correct_drafts_and_bonus(
*,
candidates: torch.Tensor,
@@ -605,10 +630,25 @@ def compute_dflash_correct_drafts_and_bonus(
if block_size <= 0:
raise ValueError(f"block_size must be positive, got {block_size}.")
if candidates.is_cuda:
num_correct_drafts = torch.empty(
bs, dtype=torch.int32, device=candidates.device
)
bonus_tokens = torch.empty(bs, dtype=torch.int64, device=candidates.device)
_fused_correct_drafts_and_bonus_kernel[(bs,)](
candidates.contiguous(),
target_predict.contiguous(),
num_correct_drafts,
bonus_tokens,
block_size,
BLOCK=triton.next_power_of_2(max(block_size - 1, 1)),
)
return num_correct_drafts, bonus_tokens
matches = candidates[:, 1:] == target_predict[:, :-1]
correct_len = matches.to(torch.int32).cumprod(dim=1).sum(dim=1)
bonus = target_predict[torch.arange(bs, device=target_predict.device), correct_len]
return correct_len, bonus.to(torch.int64)
return correct_len.to(torch.int32), bonus.to(torch.int64)
def apply_dflash_simulated_acceptance(
@@ -861,3 +901,83 @@ def validate_dflash_request(req: Req, enable_overlap: bool) -> Optional[str]:
return "DFLASH speculative decoding does not support return_hidden_states yet."
return None
@triton.jit
def _table_qk_norm_rope_kernel(
qkv_ptr,
q_weight_ptr,
k_weight_ptr,
cos_sin_ptr,
pos_ptr,
row_stride,
q_size,
NHQ: tl.constexpr,
D: tl.constexpr,
EPS: tl.constexpr,
):
t = tl.program_id(0).to(tl.int64)
h = tl.program_id(1)
pos = tl.load(pos_ptr + t).to(tl.int64)
HALF: tl.constexpr = D // 2
half_ar = tl.arange(0, HALF)
d_ar = tl.arange(0, D)
cos = tl.load(cos_sin_ptr + pos * D + half_ar).to(tl.float32)
sin = tl.load(cos_sin_ptr + pos * D + HALF + half_ar).to(tl.float32)
is_q = h < NHQ
col0 = tl.where(is_q, h * D, q_size + (h - NHQ) * D).to(tl.int64)
w_ptr = tl.where(is_q, q_weight_ptr.to(tl.int64), k_weight_ptr.to(tl.int64)).to(
tl.pointer_type(tl.bfloat16)
)
row = qkv_ptr + t * row_stride + col0
x = tl.load(row + d_ar).to(tl.float32)
ms = tl.sum(x * x, 0) / D
inv = 1.0 / tl.sqrt(ms + EPS)
w1 = tl.load(w_ptr + half_ar).to(tl.float32)
w2 = tl.load(w_ptr + HALF + half_ar).to(tl.float32)
x1 = tl.load(row + half_ar).to(tl.float32) * inv * w1
x2 = tl.load(row + HALF + half_ar).to(tl.float32) * inv * w2
x1 = x1.to(tl.bfloat16).to(tl.float32)
x2 = x2.to(tl.bfloat16).to(tl.float32)
o1 = x1 * cos - x2 * sin
o2 = x2 * cos + x1 * sin
tl.store(row + half_ar, o1.to(tl.bfloat16))
tl.store(row + HALF + half_ar, o2.to(tl.bfloat16))
def table_qk_norm_rope_(
qkv: torch.Tensor,
positions: torch.Tensor,
q_weight: torch.Tensor,
k_weight: torch.Tensor,
cos_sin_cache: torch.Tensor,
num_q_heads: int,
num_k_heads: int,
head_dim: int,
eps: float,
) -> None:
"""In-place QK RMSNorm + table-lookup neox RoPE on the fused QKV tensor.
Reads cos/sin from the SAME rotary table as the unfused path, so there is
no large-position angle drift (unlike theta-recompute kernels). V columns
are untouched.
"""
T = qkv.shape[0]
if T == 0:
return
grid = (T, num_q_heads + num_k_heads)
_table_qk_norm_rope_kernel[grid](
qkv,
q_weight,
k_weight,
cos_sin_cache,
positions,
qkv.stride(0),
num_q_heads * head_dim,
NHQ=num_q_heads,
D=head_dim,
EPS=eps,
)
@@ -126,19 +126,27 @@ class DSparkWorkerV2(BaseSpecWorker):
self.draft_model_runner = bundle.draft_model_runner
self.draft_model = bundle.draft_model
self._draft_sampler = None
self._linear_accept_index_cache = None
# The mask token needs an embedding row, not a tokenizer entry, so bound it
# by the embedding width. A padded vocab reserves rows past the real tokens
# and drafts place the mask there (Inkling: 200058 real, 201024 padded).
# The mask token is input-only (it is embedded, never sampled), so its
# bound is the embedding-table row count: the PADDED vocab when the
# target pads its embedding (e.g. Inkling true vocab 200058, padded
# 201024, mask 200064), else the plain vocab size.
target_model_config = self.target_worker.model_runner.model_config
target_vocab_size = (
target_embed_rows = (
getattr(target_model_config.hf_text_config, "padded_vocab_size", None)
or target_model_config.vocab_size
)
# muP targets declare logits_mup_width_multiplier; the draft was
# trained against the folded head, so compute_base_logits divides.
self.draft_model.logits_mup_width_multiplier = getattr(
target_model_config.hf_text_config, "logits_mup_width_multiplier", None
)
self._target_is_mambaish = mambaish_config(target_model_config) is not None
runtime_config = resolve_runtime_config(
draft_hf_config=self.draft_model_runner.model_config.hf_config,
speculative_num_draft_tokens=server_args.speculative_num_draft_tokens,
target_vocab_size=int(target_vocab_size),
target_vocab_size=int(target_embed_rows),
)
self.gamma = runtime_config.gamma
self.verify_num_draft_tokens = runtime_config.verify_num_draft_tokens
@@ -806,6 +814,7 @@ class DSparkWorkerV2(BaseSpecWorker):
mamba_track_indices=batch.mamba_track_indices,
mamba_steps_to_track=mamba_steps_to_track,
model=self.target_worker.model_runner.model,
req_pool_indices=batch.req_pool_indices,
)
def get_confidence_budget_prepare(self):