Hikari07jp's picture
Upload folder using huggingface_hub
a493cdb verified
Raw
History Blame
10.2 kB
"""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"]