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"]