File size: 10,158 Bytes
a493cdb | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 | """DSpark-style low-rank Markov sequential head (vendored).
=============================================================================
ATTRIBUTION
=============================================================================
Adapted from DeepSpec (https://github.com/deepseek-ai/DeepSpec),
file ``deepspec/modeling/dspark/markov_head.py`` (DSpark / DeepSeek-V4 draft
head). DeepSpec is released under the MIT License, Copyright (c) 2026 The
DeepSpec Authors. Only the *VanillaMarkov* head is vendored here (the
+16-18% accepted-length default that ships in DeepSeek-V4); the gated / RNN
variants are intentionally omitted to keep the surface minimal.
=============================================================================
WHAT THIS IS
=============================================================================
A parallel block-drafter (our DFlashDraftModel) predicts every block position
in ONE forward from mask-token inputs, so position k cannot see what was
actually sampled at position k-1 -- this is the "suffix decay" we measured
([72,57,45,35,27,22,19,16] top-1 by position).
The Markov head fixes that *cheaply* by adding a per-position logit bias that
conditions on the previous token only:
B(x_{k-1}, :) = W2( W1[x_{k-1}] ) W1 in R^{V x r}, W2 in R^{r x V}
The corrected logit for position k is U_k + B(x_{k-1}, :) where U_k is the
backbone's base logit (lm_head(hidden_k)). At TRAIN time x_{k-1} is the
teacher-forced ground-truth predecessor (apply_block_logits); at INFERENCE
time x_{k-1} is the actually-sampled draft token, so the block is sampled
LEFT-TO-RIGHT (sample_block_tokens). This is CHAIN mode (single-block verify)
-> hybrid-safe (no per-branch SSM-state-fork tax).
The head is fully self-contained: it carries its OWN W1/W2 and never touches
the backbone's (borrowed) embed_tokens / lm_head.
=============================================================================
"""
from __future__ import annotations
from typing import Optional
import torch
from torch import nn
def _sample_tokens(logits: torch.Tensor, temperature: float = 0.0) -> torch.Tensor:
"""Greedy (temperature < 1e-5) or temperature multinomial sample.
logits: (..., vocab) -> returns (...) long token ids. Mirrors dflash.sample
semantics so the Markov resample matches the rest of the pipeline."""
if temperature is None or temperature < 1e-5:
return torch.argmax(logits, dim=-1)
*lead, vocab = logits.shape
flat = (logits / temperature).reshape(-1, vocab)
probs = torch.softmax(flat, dim=-1)
return torch.multinomial(probs, num_samples=1).reshape(*lead)
class VanillaMarkov(nn.Module):
"""Memoryless low-rank transition bias B(x_{k-1}) = W2(W1[x_{k-1}])."""
def __init__(self, *, vocab_size: int, markov_rank: int):
super().__init__()
self.vocab_size = int(vocab_size)
self.markov_rank = int(markov_rank)
self.markov_head_type = "vanilla"
assert self.markov_rank > 0, (
f"VanillaMarkov requires markov_rank > 0, got {self.markov_rank}."
)
self.markov_w1 = nn.Embedding(self.vocab_size, self.markov_rank)
self.markov_w2 = nn.Linear(self.markov_rank, self.vocab_size, bias=False)
def get_prev_embeddings(self, token_ids: torch.Tensor) -> torch.Tensor:
return self.markov_w1(token_ids.long())
def project_bias(self, latent_states: torch.Tensor) -> torch.Tensor:
return self.markov_w2(latent_states)
def compute_step_bias(
self,
token_ids: torch.Tensor,
hidden_states: Optional[torch.Tensor] = None,
) -> torch.Tensor:
# vanilla head ignores hidden_states (pure function of the prev token)
del hidden_states
return self.project_bias(self.get_prev_embeddings(token_ids))
def apply_step_logits(
self,
logits: torch.Tensor, # (B, V)
*,
token_ids: torch.Tensor, # (B,)
hidden_states: Optional[torch.Tensor] = None,
) -> torch.Tensor:
bias = self.compute_step_bias(token_ids, hidden_states)
return logits + bias.to(logits.dtype)
def apply_block_logits(
self,
base_logits: torch.Tensor, # (..., M, V)
*,
token_ids: torch.Tensor, # (..., M) teacher-forced prev tokens
hidden_states: Optional[torch.Tensor] = None,
) -> torch.Tensor:
"""Train-time teacher-forced bias. Shape-agnostic over leading dims:
works for our (R, M, V) layout AND the DSpark (B, num_blocks, bs, V)
layout, since W1/W2 act only on the last dim."""
if base_logits.numel() == 0 or base_logits.shape[-2] == 0:
return base_logits
bias = self.compute_step_bias(token_ids, hidden_states)
return base_logits + bias.to(base_logits.dtype)
@torch.no_grad()
def sample_block_tokens(
self,
base_logits: torch.Tensor, # (B, M, V) backbone base logits
*,
first_prev_token_ids: torch.Tensor, # (B,) verified token before pos 0
hidden_states: Optional[torch.Tensor] = None,
temperature: float = 0.0,
) -> tuple[torch.Tensor, torch.Tensor]:
"""Inference-time LEFT-TO-RIGHT block sampling. Each position's logit is
biased by the token actually sampled at the previous position.
Returns (sampled_tokens (B, M), corrected_logits (B, M, V))."""
batch_size, proposal_len = base_logits.shape[:2]
if proposal_len == 0:
empty = torch.empty(batch_size, 0, dtype=torch.long, device=base_logits.device)
return empty, base_logits
sampled_tokens = []
corrected_logits = []
prev_token_ids = first_prev_token_ids.long()
for step_idx in range(proposal_len):
step_logits = self.apply_step_logits(
base_logits[:, step_idx, :],
token_ids=prev_token_ids,
# Forward the per-position hidden so a GatedMarkovHead actually
# gates during LEFT-TO-RIGHT sampling. Was hard-coded None, which
# silently dropped the gate (vanilla fallback) even when the
# caller had the hidden -> offline accept sims measured gated as
# vanilla. None-passing callers keep the vanilla path unchanged.
hidden_states=(
hidden_states[:, step_idx, :]
if hidden_states is not None
else None
),
)
corrected_logits.append(step_logits.unsqueeze(1))
next_token_ids = _sample_tokens(step_logits, temperature=temperature)
sampled_tokens.append(next_token_ids)
prev_token_ids = next_token_ids
return torch.stack(sampled_tokens, dim=1), torch.cat(corrected_logits, dim=1)
class GatedMarkovHead(VanillaMarkov):
"""Gated DSpark Markov head (official DeepSpec GatedMarkovHead).
Uses a sigmoid gate conditioned on [hidden_state; prev_embedding] to
modulate the markov bias. Unlike VanillaMarkov which ignores hidden_states,
GatedMarkovHead uses the backbone hidden state to adaptively gate the
bigram bias -- stronger when the backbone is uncertain, weaker when it's
confident. This should help with the serve pos0 gap we observed.
"""
def __init__(self, *, vocab_size: int, markov_rank: int, hidden_size: int):
super().__init__(vocab_size=vocab_size, markov_rank=markov_rank)
self.markov_head_type = "gated"
self.gate_proj = nn.Linear(hidden_size + markov_rank, markov_rank)
def compute_gate(
self,
token_ids: torch.Tensor,
hidden_states: torch.Tensor,
) -> torch.Tensor:
prev_embeddings = self.get_prev_embeddings(token_ids)
# Defensive dtype align: nn.Linear (gate_proj) requires its input in the
# weight dtype, and torch.cat requires both operands to share a dtype.
# The caller may hand us hidden_states in a dtype that differs from the
# head params (e.g. a float32 --head-dtype head fed a bf16 forward hidden,
# or the serve-side bf16 gate fed an fp32 draft hidden). Cast BOTH cat
# operands to gate_proj.weight.dtype so neither the concat nor the matmul
# can raise a dtype mismatch. This is a no-op in the intended paths (whole
# head is a single dtype), so it changes no numerics.
w_dtype = self.gate_proj.weight.dtype
gate_inputs = torch.cat(
[hidden_states.to(w_dtype), prev_embeddings.to(w_dtype)], dim=-1
)
return torch.sigmoid(self.gate_proj(gate_inputs))
def compute_step_bias(
self,
token_ids: torch.Tensor,
hidden_states: Optional[torch.Tensor] = None,
) -> torch.Tensor:
if hidden_states is None:
# Fallback to vanilla when hidden_states not available (e.g., offline eval)
return self.project_bias(self.get_prev_embeddings(token_ids))
prev_embeddings = self.get_prev_embeddings(token_ids)
gate = self.compute_gate(token_ids, hidden_states).to(dtype=prev_embeddings.dtype)
return self.project_bias(gate * prev_embeddings)
def build_markov_head(
*,
markov_rank: int,
vocab_size: int,
hidden_size: Optional[int] = None,
head_type: str = "vanilla",
) -> Optional[nn.Module]:
"""Return a Markov head, or None when markov_rank == 0 (head disabled)."""
markov_rank = int(markov_rank)
assert markov_rank >= 0, f"markov_rank must be >= 0, got {markov_rank}"
if markov_rank == 0:
return None
head_type = str(head_type).lower()
if head_type == "vanilla":
return VanillaMarkov(vocab_size=vocab_size, markov_rank=markov_rank)
if head_type == "gated":
assert hidden_size is not None, "GatedMarkovHead requires hidden_size"
return GatedMarkovHead(
vocab_size=vocab_size, markov_rank=markov_rank, hidden_size=hidden_size
)
raise ValueError(
f"Unsupported markov_head_type={head_type!r}; only 'vanilla' and 'gated' are vendored."
)
__all__ = ["VanillaMarkov", "build_markov_head"]
|